diff --git a/litellm/batches/batch_line_item_logging.py b/litellm/batches/batch_line_item_logging.py index 3be093657ec..f24ae9fb432 100644 --- a/litellm/batches/batch_line_item_logging.py +++ b/litellm/batches/batch_line_item_logging.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 73c993f57da..bea6a9d2439 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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, diff --git a/tests/integration/spend/test_batch_line_item_callbacks.py b/tests/integration/spend/test_batch_line_item_callbacks.py index 97cad137b51..a707a7a4f7a 100644 --- a/tests/integration/spend/test_batch_line_item_callbacks.py +++ b/tests/integration/spend/test_batch_line_item_callbacks.py @@ -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 + ] diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index ffbcd1fba24..e3737b0b761 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -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), diff --git a/tests/test_litellm/proxy/test_redis_auth_cache_flag.py b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py index 573bfc40c96..c59c4f3422d 100644 --- a/tests/test_litellm/proxy/test_redis_auth_cache_flag.py +++ b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py @@ -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 diff --git a/tests/unit/batches/test_batch_line_item_logging.py b/tests/unit/batches/test_batch_line_item_logging.py index 97442f59a0d..9a855b602e5 100644 --- a/tests/unit/batches/test_batch_line_item_logging.py +++ b/tests/unit/batches/test_batch_line_item_logging.py @@ -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