fix(batches): emit batch line-item callbacks once per batch

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-27 00:18:14 +00:00
parent f3ba65d04e
commit f65b5b1281
6 changed files with 136 additions and 3 deletions

View file

@ -14,6 +14,7 @@ from litellm.batches.batch_utils import (
_safe_output_line_stats, # pyright: ignore[reportPrivateUsage] # same reuse
_uses_native_vertex_output, # pyright: ignore[reportPrivateUsage] # same reuse
)
from litellm.caching.caching import DualCache
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.utils import (
EmbeddingResponse,
@ -63,6 +64,10 @@ _CALL_TYPE_BY_BATCH_URL: Final = MappingProxyType(
_EMPTY_BODY: Final[Mapping[str, object]] = MappingProxyType({})
_LINE_ITEM_CLAIM_TTL_SECONDS: Final = 30 * 24 * 60 * 60
batch_line_item_claim_cache: Final = DualCache()
class _BatchLineFailure(Exception):
"""A provider-reported per-line batch failure; carries the batch's hidden
@ -376,6 +381,7 @@ async def log_batch_line_items(
model_name: str | None,
litellm_params: dict[str, object] | None, # mutable-ok: the logging object's shared litellm_params dict
model_info: ModelInfo | None,
claim_cache: DualCache = batch_line_item_claim_cache,
) -> int:
"""Emit one callback event per JSONL line of a completed batch (request
paired with its response/error), behind the opt-in
@ -391,6 +397,12 @@ async def log_batch_line_items(
batch.id,
)
return 0
claim: Final = await claim_cache.async_increment_cache(
f"batch_line_items_emitted:{batch.id}", 1, ttl=_LINE_ITEM_CLAIM_TTL_SECONDS
)
if claim is not None and claim > 1:
verbose_logger.debug("batch line items already emitted for batch_id=%s, skipping", batch.id)
return 0
emitted = 0 # rebind-ok: loop accumulator for emitted line count
try:
internal_credentials: Final = parent._litellm_internal_model_credentials # pyright: ignore[reportPrivateUsage] # declared transport attribute on Logging

View file

@ -267,6 +267,10 @@ import litellm
import litellm._redis
from litellm import Router
from litellm._logging import _redact_string, verbose_proxy_logger, verbose_router_logger
from litellm.batches.batch_line_item_logging import (
_LINE_ITEM_CLAIM_TTL_SECONDS, # pyright: ignore[reportPrivateUsage] # the claim window is defined next to the cache it expires
batch_line_item_claim_cache,
)
from litellm.caching.caching import DualCache, RedisCache
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError, is_redis_timeout_failure
from litellm.caching.redis_cluster_cache import RedisClusterCache
@ -4733,8 +4737,8 @@ def _attach_redis_usage_cache(redis_cache: RedisCache, enable_redis_auth_cache:
"""
Wires an established coordination Redis into the proxy-level caches that
consume it directly: the spend counter cache, the CLI SSO login-session
cache, the cluster-wide config cache, and (only when opted in) the
virtual-key auth cache.
cache, the batch line-item claim cache, the cluster-wide config cache,
and (only when opted in) the virtual-key auth cache.
The CLI SSO login-session cache is always backed by Redis when available so
that the browser SSO flow behind `lite login` survives landing on different
@ -4748,6 +4752,10 @@ def _attach_redis_usage_cache(redis_cache: RedisCache, enable_redis_auth_cache:
redis_cache,
default_redis_ttl=CLI_SSO_SESSION_TTL_SECONDS,
)
batch_line_item_claim_cache.attach_redis_cache(
redis_cache,
default_redis_ttl=_LINE_ITEM_CLAIM_TTL_SECONDS,
)
if enable_redis_auth_cache is True:
user_api_key_cache.attach_redis_cache(
redis_cache,

View file

@ -239,7 +239,8 @@ def test_completed_batch_emits_paired_request_response_callback_events_per_jsonl
events: Final = eventually(
delivered,
lambda values: (
len([e for e in values if e["call_type"] == "aretrieve_batch"]) >= 1
len([e for e in values if e["call_type"] == "acreate_batch"]) >= 1
and len([e for e in values if e["call_type"] == "aretrieve_batch"]) >= 1
and len([e for e in values if _hidden(e).get("batch_custom_id") is not None]) >= len(ALL_CUSTOM_IDS)
),
seconds=40,
@ -286,3 +287,17 @@ def test_completed_batch_emits_paired_request_response_callback_events_per_jsonl
assert len(batch_rows) == 1, rows
assert not any(row["call_type"] == "acompletion" for row in rows), rows
assert batch_rows[0]["prompt_tokens"] == len(OUTPUT_SUCCESS_IDS) * PROMPT_TOKENS, rows
repeated_gets: Final = [candidate.request("GET", f"/v1/batches/{batch_id}", key=key) for _ in range(2)]
assert all(response.status_code == 200 for response in repeated_gets)
events_after_repeats: Final = eventually(
delivered,
lambda values: len([e for e in values if e["call_type"] == "aretrieve_batch"]) >= 1 + len(repeated_gets),
seconds=40,
)
repeat_line_events: Final = tuple(
event for event in events_after_repeats if _hidden(event).get("batch_custom_id") is not None
)
assert len(repeat_line_events) == len(ALL_CUSTOM_IDS), [
(event["call_type"], _hidden(event).get("batch_custom_id")) for event in events_after_repeats
]

View file

@ -13121,6 +13121,7 @@ def _patched_coordination_redis_module_state(
patch.object(proxy_server_module, "spend_counter_cache", spend_cache),
patch.object(proxy_server_module, "user_api_key_cache", DualCache()),
patch.object(proxy_server_module, "cli_sso_session_cache", DualCache()),
patch.object(proxy_server_module, "batch_line_item_claim_cache", DualCache()),
patch.object(proxy_server_module, "llm_router", None),
patch.object(proxy_server_module, "litellm_config_cache", config_cache),
patch.object(proxy_server_module, "RedisCache", redis_cache_class),

View file

@ -182,10 +182,12 @@ class TestRedisAuthCacheFlag:
ps.cli_sso_session_cache,
ps.user_api_key_cache,
ps.litellm_config_cache,
ps.batch_line_item_claim_cache,
)
with ExitStack() as detached:
for cache in touched_caches:
detached.enter_context(patch.object(cache, "redis_cache", None))
ps._attach_redis_usage_cache(fake_redis, enable_redis_auth_cache=False)
assert limiter_cache.redis_cache is fake_redis
assert ps.batch_line_item_claim_cache.redis_cache is fake_redis
assert ps.user_api_key_cache.redis_cache is None

View file

@ -21,6 +21,8 @@ from unittest.mock import AsyncMock, patch
import pytest
import litellm
from litellm.batches.batch_line_item_logging import batch_line_item_claim_cache, log_batch_line_items
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.types.utils import LiteLLMBatch, Usage
@ -120,6 +122,12 @@ class _RecordingLogger(CustomLogger):
self.failure_events.append(kwargs)
@pytest.fixture(autouse=True)
def _fresh_line_item_claim_cache():
batch_line_item_claim_cache.in_memory_cache.flush_cache()
yield
@pytest.fixture
def recorder():
logger = _RecordingLogger()
@ -797,3 +805,90 @@ async def test_line_items_native_vertex_rows_are_skipped(recorder):
assert len(recorder.success_events) == 1
assert len(recorder.failure_events) == 0
assert _hidden(recorder.success_events[0]).get("batch_custom_id") is None
class _NeverClaimsCache(DualCache):
async def async_increment_cache(self, *_args: object, **_kwargs: object) -> None:
return None
@pytest.mark.asyncio
async def test_line_items_emit_once_per_batch_id(recorder):
file_mock: Final = AsyncMock(side_effect=_file_content)
claim_cache: Final = DualCache()
parent: Final = _parent_logging()
batch: Final = _batch()
with patch("litellm.files.main.afile_content", file_mock): # test-quality-ok: afile_content is the provider boundary; no injection seam for managed file fetch
first: Final = await log_batch_line_items(
batch=batch,
custom_llm_provider="openai",
parent=parent,
model_name="gpt-4o",
litellm_params=None,
model_info=None,
claim_cache=claim_cache,
)
second: Final = await log_batch_line_items(
batch=batch,
custom_llm_provider="openai",
parent=parent,
model_name="gpt-4o",
litellm_params=None,
model_info=None,
claim_cache=claim_cache,
)
assert first == 2
assert second == 0
assert len(recorder.success_events) == 1
assert len(recorder.failure_events) == 1
@pytest.mark.asyncio
async def test_line_items_claim_is_scoped_to_the_batch_id(recorder):
file_mock: Final = AsyncMock(side_effect=_file_content)
claim_cache: Final = DualCache()
parent: Final = _parent_logging()
with patch("litellm.files.main.afile_content", file_mock): # test-quality-ok: afile_content is the provider boundary; no injection seam for managed file fetch
first: Final = await log_batch_line_items(
batch=_batch(),
custom_llm_provider="openai",
parent=parent,
model_name="gpt-4o",
litellm_params=None,
model_info=None,
claim_cache=claim_cache,
)
second: Final = await log_batch_line_items(
batch=_provider_batch("batch_other", "input-file-1", "output-file-1"),
custom_llm_provider="openai",
parent=parent,
model_name="gpt-4o",
litellm_params=None,
model_info=None,
claim_cache=claim_cache,
)
assert first == 2
assert second == 1
assert len(recorder.success_events) == 2
assert len(recorder.failure_events) == 1
@pytest.mark.asyncio
async def test_line_items_emit_when_the_claim_backend_returns_nothing(recorder):
file_mock: Final = AsyncMock(side_effect=_file_content)
with patch("litellm.files.main.afile_content", file_mock): # test-quality-ok: afile_content is the provider boundary; no injection seam for managed file fetch
emitted: Final = await log_batch_line_items(
batch=_batch(),
custom_llm_provider="openai",
parent=_parent_logging(),
model_name="gpt-4o",
litellm_params=None,
model_info=None,
claim_cache=_NeverClaimsCache(),
)
assert emitted == 2
assert len(recorder.success_events) == 1
assert len(recorder.failure_events) == 1