mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
cffa6405b4
commit
a7c9dedbf5
2 changed files with 203 additions and 12 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue