mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
f3ba65d04e
commit
f65b5b1281
6 changed files with 136 additions and 3 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue