fix(batches): make the line-item claim release ownership-safe

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-29 20:45:36 +00:00
parent cffa6405b4
commit a7c9dedbf5
2 changed files with 203 additions and 12 deletions

View file

@ -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

View file

@ -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