From a7c9dedbf5a3dc77c7b9ce1f373477e4b0b8e9a8 Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 29 Sep 2026 20:45:36 +0000 Subject: [PATCH] fix(batches): make the line-item claim release ownership-safe Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/batches/batch_line_item_logging.py | 64 ++++++-- .../batches/test_batch_line_item_logging.py | 151 +++++++++++++++++- 2 files changed, 203 insertions(+), 12 deletions(-) diff --git a/litellm/batches/batch_line_item_logging.py b/litellm/batches/batch_line_item_logging.py index 21dd23bdb59..94c2bb0bb95 100644 --- a/litellm/batches/batch_line_item_logging.py +++ b/litellm/batches/batch_line_item_logging.py @@ -5,6 +5,8 @@ from datetime import datetime from types import MappingProxyType from typing import TYPE_CHECKING, Final, Literal, TypeAlias, cast, get_args +from typing_extensions import assert_never + from litellm._logging import verbose_logger from litellm.batches.batch_utils import ( _batch_response_was_successful, # pyright: ignore[reportPrivateUsage] # batch-internal helper shared with the aggregate cost path by design @@ -68,6 +70,40 @@ _LINE_ITEM_CLAIM_TTL_SECONDS: Final = 30 * 24 * 60 * 60 batch_line_item_claim_cache: Final = DualCache() +_ClaimResult: TypeAlias = Literal["claimed", "already_claimed", "unavailable"] + + +async def _claim_line_items(claim_cache: DualCache, claim_key: str, token: str) -> _ClaimResult: + redis_cache: Final = claim_cache.redis_cache + if redis_cache is None: + count: Final = await claim_cache.async_increment_cache(claim_key, 1, ttl=_LINE_ITEM_CLAIM_TTL_SECONDS) # pyright: ignore[reportUnknownMemberType] # DualCache.increment is untyped upstream + if count is None or count == 1: + return "claimed" + return "already_claimed" + try: + ok: Final = await redis_cache.async_set_cache(claim_key, token, ttl=_LINE_ITEM_CLAIM_TTL_SECONDS, nx=True) # pyright: ignore[reportUnknownMemberType] # RedisCache.set is untyped upstream + except Exception: # noqa: BLE001 # a redis outage must not emit duplicates; the next retrieve retries + return "unavailable" + if ok: + return "claimed" + return "already_claimed" + + +async def _release_line_item_claim(claim_cache: DualCache, claim_key: str, token: str) -> None: + try: + redis_cache: Final = claim_cache.redis_cache + if redis_cache is None: + await claim_cache.async_delete_cache(claim_key) + return + owner: Final = await redis_cache.async_get_cache(claim_key) # pyright: ignore[reportUnknownMemberType] # RedisCache.get is untyped upstream + if owner == token: + await redis_cache.async_delete_cache(claim_key) + except Exception: # noqa: BLE001 # the claim release must never raise; worst case the batch stays claimed + verbose_logger.debug( + "batch line item claim release failed for %s, claim persists until ttl", + claim_key, + ) + class _BatchLineFailure(Exception): """A provider-reported per-line batch failure; carries the batch's hidden @@ -398,10 +434,22 @@ async def log_batch_line_items( ) return 0 claim_key: Final = f"batch_line_items_emitted:{batch.id}" - claim: Final = await claim_cache.async_increment_cache(claim_key, 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 + token: Final = uuid.uuid4().hex + claim: Final = await _claim_line_items(claim_cache, claim_key, token) + match claim: + case "already_claimed": + verbose_logger.debug("batch line items already emitted for batch_id=%s, skipping", batch.id) + return 0 + case "unavailable": + verbose_logger.warning( + "batch line item claim backend unavailable for batch_id=%s, line items will be retried on the next retrieve", + batch.id, + ) + return 0 + case "claimed": + pass + case _: + assert_never(claim) 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 @@ -447,11 +495,5 @@ async def log_batch_line_items( batch.id, ) if emitted == 0: - try: - await claim_cache.async_delete_cache(claim_key) - except Exception: # noqa: BLE001 # the claim release must never raise; worst case the batch stays claimed - verbose_logger.debug( - "batch line item claim release failed for batch_id=%s, claim persists until ttl", - batch.id, - ) + await _release_line_item_claim(claim_cache, claim_key, token) return emitted diff --git a/tests/unit/batches/test_batch_line_item_logging.py b/tests/unit/batches/test_batch_line_item_logging.py index f5514c5814c..e78c0a1c307 100644 --- a/tests/unit/batches/test_batch_line_item_logging.py +++ b/tests/unit/batches/test_batch_line_item_logging.py @@ -21,8 +21,13 @@ 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.batches.batch_line_item_logging import ( + _release_line_item_claim, + batch_line_item_claim_cache, + log_batch_line_items, +) from litellm.caching.caching import DualCache +from litellm.caching.redis_cache import RedisCache from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging from litellm.types.utils import LiteLLMBatch, Usage @@ -964,3 +969,147 @@ async def test_line_items_claim_kept_after_partial_emission(recorder): assert second == 0 assert len(recorder.success_events) == 1 assert len(recorder.failure_events) == 0 + + +class _FakeRedisCache(RedisCache): + """Dict-backed RedisCache stand-in honoring SET NX so the tokenized + line-item claim can be exercised without a Redis process.""" + + def __init__(self) -> None: + self._store: dict[str, object] = {} # mutable-ok: a dict-backed fake needs a mutable store + self.fail_next_set = False + self.swap_owner_on_get: str | None = None + + async def async_set_cache(self, key, value, nx=False, **_kwargs): + if self.fail_next_set: + self.fail_next_set = False + raise ConnectionError("redis down") + if nx and key in self._store: + return None + self._store[key] = value + return True + + async def async_get_cache(self, key, **_kwargs): + if self.swap_owner_on_get is not None and key in self._store: + self._store[key] = self.swap_owner_on_get + return self._store.get(key) + + async def async_delete_cache(self, key): + self._store.pop(key, None) + + +@pytest.mark.asyncio +async def test_line_items_emit_once_per_batch_id_over_redis(recorder): + claim_cache: Final = DualCache(redis_cache=_FakeRedisCache()) + file_mock: Final = AsyncMock(side_effect=_file_content) + 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_skip_when_redis_claim_backend_is_down(recorder): + fake: Final = _FakeRedisCache() + fake.fail_next_set = True + claim_cache: Final = DualCache(redis_cache=fake) + file_mock: Final = AsyncMock(side_effect=_file_content) + 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 == 0 + assert second == 2 + assert len(recorder.success_events) == 1 + assert len(recorder.failure_events) == 1 + + +@pytest.mark.asyncio +async def test_release_line_item_claim_only_deletes_owned_claims(): + fake: Final = _FakeRedisCache() + claim_cache: Final = DualCache(redis_cache=fake) + key: Final = "batch_line_items_emitted:batch_1" + fake._store[key] = "someone-else" # test-quality-ok: seeding the fake's store is the arrangement, like writing Redis directly + + await _release_line_item_claim(claim_cache, key, "my-token") + assert fake._store.get(key) == "someone-else" + + await _release_line_item_claim(claim_cache, key, "someone-else") + assert key not in fake._store + + +@pytest.mark.asyncio +async def test_line_items_failed_fanout_does_not_delete_another_workers_claim(recorder): + fake: Final = _FakeRedisCache() + fake.swap_owner_on_get = "other-worker" + claim_cache: Final = DualCache(redis_cache=fake) + key: Final = "batch_line_items_emitted:batch_1" + batch: Final = _batch() + parent: Final = _parent_logging() + with patch("litellm.files.main.afile_content", new_callable=AsyncMock, side_effect=ValueError("boom")): # 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, + ) + 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 + 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 == 0 + assert second == 0 + assert fake._store.get(key) == "other-worker" + assert len(recorder.success_events) == 0 + assert len(recorder.failure_events) == 0