diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 985a9b6de38..191c983f721 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -703,9 +703,10 @@ class CheckBatchCost: later poll. """ from litellm.batches.batch_utils import ( - count_error_file_failed_requests, _get_file_content_as_dictionary, calculate_batch_cost_and_usage, + count_error_file_failed_requests_from_content, + fetch_batch_error_file_content, ) from litellm.files.main import afile_content from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider @@ -848,7 +849,7 @@ class CheckBatchCost: model_name=model_name, model_info=deployment_model_info, ) - error_file_failed_requests: Final = await count_error_file_failed_requests( + error_file_content = await fetch_batch_error_file_content( response, custom_llm_provider=batch_file_provider, litellm_params={ @@ -856,6 +857,7 @@ class CheckBatchCost: "_litellm_internal_model_credentials": MappingProxyType(dict(credentials)), }, ) + error_file_failed_requests: Final = count_error_file_failed_requests_from_content(error_file_content) batch_result: Final = ( output_file_result if not error_file_failed_requests @@ -895,6 +897,7 @@ class CheckBatchCost: optional_params={}, custom_llm_provider=str(llm_provider) if llm_provider else None, ) + logging_obj._litellm_internal_model_credentials = MappingProxyType(dict(credentials)) # pyright: ignore[reportPrivateUsage] # trusted credentials transport, consumed by batch_line_item_logging if not await self._claim_job_for_costing(job): verbose_proxy_logger.info( @@ -913,6 +916,8 @@ class CheckBatchCost: batch_failed_requests=batch_result.failed_requests, batch_prompt_cost=batch_result.prompt_cost, batch_completion_cost=batch_result.completion_cost, + batch_output_file_content=content_bytes if isinstance(content_bytes, bytes) else None, + batch_error_file_content=error_file_content, ) except Exception: await self._release_job_claim(job) diff --git a/litellm/__init__.py b/litellm/__init__.py index fea7a27a5fb..dfe5de00f41 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -346,6 +346,7 @@ anthropic_prompt_caching_ttl: Optional[Literal["5m", "1h"]] = ( ) openai_system_messages_first: bool = False disable_vertex_batch_output_transformation: bool = False +store_batch_line_items_in_callbacks: bool = False extra_spend_tag_headers: Optional[List[str]] = None in_memory_llm_clients_cache: "LLMClientCache" safe_memory_mode: bool = False diff --git a/litellm/batches/batch_line_item_logging.py b/litellm/batches/batch_line_item_logging.py new file mode 100644 index 00000000000..a53159eb6d5 --- /dev/null +++ b/litellm/batches/batch_line_item_logging.py @@ -0,0 +1,518 @@ +import json +import uuid +from collections.abc import Iterator, Mapping +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._internal_context import with_service_target +from litellm._logging import verbose_logger +from litellm.batches.batch_utils import ( + BatchResultFiles, + _batch_response_was_successful, # pyright: ignore[reportPrivateUsage] # batch-internal helper shared with the aggregate cost path by design + _fetch_batch_managed_file_content, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # same reuse; helper is untyped upstream + _get_response_from_batch_job_output_file, # pyright: ignore[reportPrivateUsage] # same reuse + _iter_batch_output_entries, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # same reuse; helper is untyped upstream + _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, + LiteLLMBatch, + ModelInfo, + ModelResponse, + Usage, +) + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging + +_BatchLineProvider: TypeAlias = Literal[ + "openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral" +] + +_SUPPORTED_LINE_PROVIDERS: Final = frozenset(get_args(_BatchLineProvider)) + +_SECRET_PARAM_KEYS: Final = frozenset( + { + "api_key", + "_litellm_internal_model_credentials", + "azure_ad_token", + "azure_ad_token_provider", + "vertex_credentials", + "aws_access_key_id", + "aws_secret_access_key", + "aws_session_token", + "aws_web_identity_token", + } +) + + +def _supported_line_provider(value: str) -> _BatchLineProvider | None: + if value in _SUPPORTED_LINE_PROVIDERS: + return cast("_BatchLineProvider", value) # cast-ok: membership in the literal's args was just checked + return None + + +_CALL_TYPE_BY_BATCH_URL: Final = MappingProxyType( + { + "/v1/chat/completions": "acompletion", + "/v1/embeddings": "aembedding", + "/v1/responses": "aresponses", + } +) + +_EMPTY_BODY: Final[Mapping[str, object]] = MappingProxyType({}) + +_LINE_ITEM_CLAIM_TTL_SECONDS: Final = 30 * 24 * 60 * 60 + +batch_line_item_claim_cache: Final = DualCache() + +_ClaimResult: TypeAlias = Literal["claimed", "already_claimed", "unavailable"] + + +@with_service_target("batch_line_items") +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" + + +@with_service_target("batch_line_items") +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, + ) + + +def _json_fallback(value: object) -> Mapping[str, object] | str: + mapping: Final = _as_object_mapping(value) + if mapping is None: + return str(value) + return dict(mapping) + + +class _BatchLineFailure(Exception): + """A provider-reported per-line batch failure; carries the batch's hidden + params so the failure logging payload can attribute the line.""" + + def __init__(self, error_payload: object) -> None: + super().__init__(json.dumps(error_payload, default=_json_fallback)) + self._hidden_params: dict[str, object] = {} # mutable-ok: plain-dict contract like response _hidden_params + + +def _as_object_mapping(value: object) -> Mapping[str, object] | None: + if isinstance(value, Mapping) and all(isinstance(key, str) for key in value): # pyright: ignore[reportUnknownVariableType] # keys of an unparameterized Mapping are unknown until checked here + return value # pyright: ignore[reportUnknownVariableType, reportReturnType] # every key was verified str above + return None + + +def _output_entries(file_content: bytes) -> Iterator[Mapping[str, object]]: + for entry in _iter_batch_output_entries(file_content): # pyright: ignore[reportUnknownVariableType] # entries are validated into typed mappings below + mapping = _as_object_mapping(entry) # pyright: ignore[reportUnknownArgumentType] # raw entry is unknown until validated here + if mapping is not None: + yield mapping + + +def _line_id(entry: Mapping[str, object]) -> str | None: + line_id: Final = entry.get("custom_id") or entry.get("recordId") + return line_id if isinstance(line_id, str) and line_id else None + + +def _requests_by_custom_id(input_file_content: bytes) -> Mapping[str, Mapping[str, object]]: + """Parse the batch input JSONL into {line id: request line}, keyed by + custom_id or recordId, skipping malformed lines and lines without either.""" + return MappingProxyType( + {line_id: entry for entry in _output_entries(input_file_content) if (line_id := _line_id(entry)) is not None} + ) + + +def _request_body_for_entry( + entry: Mapping[str, object], request_line: Mapping[str, object] | None +) -> Mapping[str, object]: + if request_line is not None: + body: Final = _as_object_mapping(request_line.get("body")) + if body: + return body + params: Final = _as_object_mapping(request_line.get("params")) + if params: + return params + request_model_input: Final = _as_object_mapping(request_line.get("modelInput")) + if request_model_input: + return request_model_input + model_input: Final = _as_object_mapping(entry.get("modelInput")) + return model_input if model_input else _EMPTY_BODY + + +def _line_status_code(entry: Mapping[str, object], custom_llm_provider: str) -> int | None: + response: Final = _as_object_mapping(entry.get("response")) + status: Final = response.get("status_code") if response is not None else None + if isinstance(status, int): + return status + if custom_llm_provider == "anthropic": + result: Final = _as_object_mapping(entry.get("result")) + if result is not None and result.get("type") == "succeeded": + return 200 + return None + + +def _line_error_payload(entry: Mapping[str, object], custom_llm_provider: _BatchLineProvider) -> object: + if custom_llm_provider == "anthropic": + return ( + (_as_object_mapping(entry.get("result")) or _EMPTY_BODY).get("error") or entry.get("result") or _EMPTY_BODY + ) + return entry.get("error") or entry.get("response") or _EMPTY_BODY + + +def _call_type_for_request(request_line: Mapping[str, object] | None) -> str: + url: Final = request_line.get("url") if request_line is not None else None + return _CALL_TYPE_BY_BATCH_URL.get(url if isinstance(url, str) else "", "acompletion") + + +def _call_type_for_line(request_line: Mapping[str, object] | None, result: "_BatchLineResult | None") -> str: + if isinstance(result, EmbeddingResponse): + return "aembedding" + if isinstance(result, ResponsesAPIResponse): + return "aresponses" + if isinstance(result, ModelResponse): + return "acompletion" + return _call_type_for_request(request_line) + + +def _line_messages(request_body: Mapping[str, object]) -> object: + return request_body.get("messages") or request_body.get("input") or () + + +def _line_model( + response_body: Mapping[str, object], + request_body: Mapping[str, object], + parent: "Logging", +) -> str: + for candidate in (response_body.get("model"), request_body.get("model"), parent.model): + if isinstance(candidate, str) and candidate: + return candidate + return "" + + +_BatchLineResult: TypeAlias = "ModelResponse | EmbeddingResponse | ResponsesAPIResponse" + + +def _line_result( + call_type: str, + custom_llm_provider: _BatchLineProvider, + model: str, + response_body: Mapping[str, object], +) -> _BatchLineResult: + from litellm.types.utils import LlmProviders + from litellm.utils import ProviderConfigManager + + provider_config: Final = ProviderConfigManager.get_provider_batches_config(model, LlmProviders(custom_llm_provider)) + if provider_config is not None: + transformed: Final = provider_config.transform_batch_output_line(response_body, model) + if transformed is not None: + return transformed + if custom_llm_provider == "bedrock": + raise ValueError(f"unrecognized bedrock batch output line shape. keys={sorted(response_body)}") + if call_type == "aembedding": + return EmbeddingResponse(**response_body) # pyright: ignore[reportArgumentType] # provider output bodies are dicts expanded as response ctor kwargs + if call_type == "aresponses": + return ResponsesAPIResponse(**response_body) # pyright: ignore[reportArgumentType] # same as above + if custom_llm_provider == "anthropic": + from litellm.llms.anthropic.chat.transformation import anthropic_message_to_model_response + + return anthropic_message_to_model_response(response_body, None) + return ModelResponse(**response_body) # pyright: ignore[reportArgumentType] # same as above + + +def _line_result_or_none( + request_call_type: str, + custom_llm_provider: _BatchLineProvider, + model: str, + response_body: Mapping[str, object], + custom_id: object, +) -> "_BatchLineResult | None": + try: + return _line_result(request_call_type, custom_llm_provider, model, response_body) + except Exception: # noqa: BLE001 # one unparseable line must not drop the rest of the batch's line events + verbose_logger.warning( + "batch output line could not be reconstructed as a %s response, skipping it. custom_id=%s", + request_call_type, + custom_id, + ) + return None + + +def _new_child_logging( + parent: "Logging", + model: str, + messages: object, + call_type: str, + start_time: datetime, +) -> "Logging": + from litellm.litellm_core_utils.litellm_logging import Logging + + return Logging( + model=model, + messages=messages, + stream=False, + call_type=call_type, + start_time=start_time, + litellm_call_id=str(uuid.uuid4()), + function_id=str(uuid.uuid4()), + litellm_trace_id=parent.litellm_trace_id, + dynamic_success_callbacks=parent.dynamic_success_callbacks, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # Logging ctor takes these dynamic callback lists untyped + dynamic_async_success_callbacks=parent.dynamic_async_success_callbacks, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # same as above + dynamic_failure_callbacks=parent.dynamic_failure_callbacks, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # same as above + dynamic_async_failure_callbacks=parent.dynamic_async_failure_callbacks, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # same as above + kwargs={"litellm_session_id": parent.litellm_session_id}, + ) + + +def _line_hidden_params( + batch: LiteLLMBatch, + custom_id: object, + status_code: int | None, + response_cost: float | None = None, +) -> dict[str, object]: # mutable-ok: response objects declare _hidden_params as a plain dict + return { + "batch_id": batch.id, + "batch_custom_id": custom_id, + "batch_line_status_code": status_code, + "response_cost": response_cost, + } + + +def _optional_params_for_body( + request_body: Mapping[str, object], +) -> dict[str, object]: # mutable-ok: update_environment_variables takes a plain dict + return {key: value for key, value in request_body.items() if key not in ("model", "messages", "input")} + + +def _metadata_copy(params: Mapping[str, object]) -> dict[str, object]: # mutable-ok: dict out for litellm_params + metadata: Final = _as_object_mapping(params.get("metadata")) or _EMPTY_BODY + return {**metadata} + + +async def _emit_line_event( + entry: Mapping[str, object], + requests_by_id: Mapping[str, Mapping[str, object]], + batch: LiteLLMBatch, + custom_llm_provider: _BatchLineProvider, + parent: "Logging", + model_name: str | None, + model_info: ModelInfo | None, +) -> bool: + custom_id: Final = _line_id(entry) + request_line: Final = requests_by_id.get(custom_id or "") + request_body: Final = _request_body_for_entry(entry, request_line) + status_code: Final = _line_status_code(entry, custom_llm_provider) + request_call_type: Final = _call_type_for_request(request_line) + response_body: Final = _get_response_from_batch_job_output_file(entry, custom_llm_provider) + parent_start_time: Final = parent.start_time # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # Logging.start_time is untyped upstream + start_time: Final = parent_start_time if isinstance(parent_start_time, datetime) else datetime.now() # noqa: DTZ005 # naive to match the logging pipeline start_time + parent_params: Final = _as_object_mapping(parent.litellm_params) or _EMPTY_BODY # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # Logging.litellm_params is untyped upstream + model: Final = _line_model(response_body, request_body, parent) + + successful: Final = _batch_response_was_successful(entry, custom_llm_provider) + stats: Final = _safe_output_line_stats(entry, custom_llm_provider, model_name, model_info) if successful else None + result: Final[_BatchLineResult | None] = ( + _line_result_or_none(request_call_type, custom_llm_provider, model, response_body, custom_id) + if successful + else None + ) + if successful and result is None: + return False + + call_type: Final = _call_type_for_line(request_line, result) + child: Final = _new_child_logging( + parent=parent, + model=model, + messages=_line_messages(request_body), + call_type=call_type, + start_time=start_time, + ) + child.update_environment_variables( # pyright: ignore[reportUnknownMemberType] # Logging.update_environment_variables is untyped upstream + litellm_params={ + **parent_params, + "batch_parent_id": batch.id, + "metadata": _metadata_copy(parent_params), + }, + optional_params=_optional_params_for_body(request_body), + model=child.model, + custom_llm_provider=custom_llm_provider, + ) + for secret_key in _SECRET_PARAM_KEYS: + child.litellm_params.pop(secret_key, None) # pyright: ignore[reportUnknownMemberType] # Logging.litellm_params is untyped upstream + + now: Final = datetime.now() # noqa: DTZ005 # naive to match the logging pipeline start_time + if result is None: + exception: Final = _BatchLineFailure(_line_error_payload(entry, custom_llm_provider)) + exception._hidden_params = _line_hidden_params(batch, custom_id, status_code) # pyright: ignore[reportPrivateUsage] # _hidden_params is set on the exception instance itself + await child.async_failure_handler( + exception=exception, + traceback_exception="", + start_time=start_time, + end_time=now, + ) + return True + + result._hidden_params = _line_hidden_params( # pyright: ignore[reportPrivateUsage] # same hidden_params channel the aggregate batch event uses + batch, + custom_id, + status_code, + response_cost=stats.prompt_cost + stats.completion_cost if stats is not None else None, + ) + if stats is not None and not response_body.get("usage") and isinstance(result, (ModelResponse, EmbeddingResponse)): + setattr( # noqa: B010 # ModelResponse.usage is set dynamically by its ctor, so setattr keeps parity for both response types + result, + "usage", + Usage( + prompt_tokens=stats.prompt_tokens, + completion_tokens=stats.completion_tokens, + total_tokens=stats.total_tokens, + ), + ) + await child.async_success_handler( + result=result, + start_time=start_time, + end_time=now, + cache_hit=False, + ) + return True + + +async def _fetch_managed_file_or_empty( + file_id: str | None, + custom_llm_provider: _BatchLineProvider, + fetch_params: dict[str, object] | None, # mutable-ok: batch_utils file fetch takes the shared litellm_params dict +) -> bytes: + if file_id is None: + return b"" + return await _fetch_batch_managed_file_content( + file_id, + custom_llm_provider=custom_llm_provider, + litellm_params=fetch_params, # pyright: ignore[reportArgumentType] # batch_utils types this param as an unparameterized dict + ) + + +async def log_batch_line_items( + batch: LiteLLMBatch, + custom_llm_provider: str, + parent: "Logging", + model_name: str | None, + litellm_params: dict[str, object] | None, # mutable-ok: the logging object's shared litellm_params dict + model_info: ModelInfo | None, + result_files: BatchResultFiles | None = 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 + ``litellm.store_batch_line_items_in_callbacks`` flag. The aggregate + aretrieve_batch event still bills the batch, so per-line events carry + ``batch_parent_id`` and never update spend themselves. Any failure here + is logged and swallowed: aggregate accounting must be unaffected. + ``result_files`` carries the output/error bytes the aggregate path already + fetched, so they are reused instead of refetched.""" + line_provider: Final = _supported_line_provider(custom_llm_provider) + if line_provider is None: + verbose_logger.warning( + "batch line-item callbacks are not supported for provider %s, skipping. batch_id=%s", + custom_llm_provider, + batch.id, + ) + return 0 + claim_key: Final = f"batch_line_items_emitted:{batch.id}" + 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 + internal_mapping: Final = _as_object_mapping(internal_credentials) + fetch_params: Final[dict[str, object] | None] = ( # mutable-ok: file fetcher requires a plain dict + dict(internal_mapping) if internal_mapping is not None else litellm_params + ) + + input_file_content: Final = await _fetch_managed_file_or_empty(batch.input_file_id, line_provider, fetch_params) + requests_by_id: Final = _requests_by_custom_id(input_file_content) + + output_content: Final = ( + result_files.output + if result_files is not None and result_files.output is not None + else await _fetch_managed_file_or_empty(batch.output_file_id, line_provider, fetch_params) + ) + first_row: Final = next(_output_entries(output_content), None) + if _uses_native_vertex_output(line_provider, model_name, first_row): + verbose_logger.warning( + "batch line-item callbacks do not support native vertex_ai batch output rows yet, skipping. batch_id=%s", + batch.id, + ) + return 0 + error_content: Final = ( + result_files.error + if result_files is not None and result_files.error is not None + else await _fetch_managed_file_or_empty(batch.error_file_id, line_provider, fetch_params) + ) + for content in (output_content, error_content): + for entry in _output_entries(content): + try: + emitted += await _emit_line_event( + entry=entry, + requests_by_id=requests_by_id, + batch=batch, + custom_llm_provider=line_provider, + parent=parent, + model_name=model_name, + model_info=model_info, + ) + except Exception: # noqa: BLE001 # one bad line must not drop the rest of the batch's line events + verbose_logger.exception( + "batch line item logging failed for entry, continuing with remaining lines. batch_id=%s", + batch.id, + ) + except Exception: # noqa: BLE001 # line-item logging must never break the aggregate aretrieve_batch accounting + verbose_logger.exception( + "batch line item logging failed for batch_id=%s; aggregate logging unaffected", + batch.id, + ) + finally: + if emitted == 0: + await _release_line_item_claim(claim_cache, claim_key, token) + return emitted diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 23ef2a99585..04cd0777cd3 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -34,6 +34,14 @@ class BatchCostUsageResult: completion_cost: float = 0.0 +@dataclass(frozen=True, slots=True) +class BatchResultFiles: + """Raw JSONL bytes of a completed batch's result files; None means not fetched.""" + + output: bytes | None + error: bytes | None + + _COMPLETED_BATCH_STATUSES: Final = frozenset({"completed", "complete"}) @@ -71,7 +79,7 @@ def batch_cost_is_final(batch: Batch) -> bool: async def calculate_batch_cost_and_usage( file_content_dictionary: list[dict], - custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"], + custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"], model_name: str | None = None, model_info: ModelInfo | None = None, ) -> BatchCostUsageResult: @@ -98,7 +106,7 @@ async def calculate_batch_cost_and_usage( async def _handle_completed_batch( batch: Batch, - custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"], + custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"], model_name: str | None = None, litellm_params: dict | None = None, model_info: ModelInfo | None = None, @@ -116,29 +124,44 @@ async def _handle_completed_batch( threaded through so a deployment's configured rates win over the global cost map. """ - # A completed batch whose request lines all failed has no output file - the - # results are written to a separate error_file_id and output_file_id is None. - # There is nothing to price or measure, so report an empty result set instead - # of calling _fetch_batch_output_file_content, which raises on a missing - # output file. Without this guard the logging worker crashes on every - # aretrieve_batch poll and the completed batch's zero-cost accounting is lost. - # The generic retrieval helper keeps raising for callers that explicitly ask - # for a missing output file. + return ( + await _handle_completed_batch_with_files( + batch, + custom_llm_provider, + model_name=model_name, + litellm_params=litellm_params, + model_info=model_info, + ) + )[0] + + +async def _handle_completed_batch_with_files( + batch: Batch, + custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"], + model_name: str | None = None, + litellm_params: dict[str, object] | None = None, # mutable-ok: file fetchers require a plain dict + model_info: ModelInfo | None = None, +) -> "tuple[BatchCostUsageResult, BatchResultFiles]": + """_handle_completed_batch plus the raw output/error bytes it fetched, so + downstream line-item logging reuses them instead of refetching.""" + error_file_content: Final = await fetch_batch_error_file_content( + batch, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params + ) + error_file_failed_requests: Final = count_error_file_failed_requests_from_content(error_file_content) + if batch.output_file_id is None: - return BatchCostUsageResult( - cost=0.0, - usage=Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0), - models=[], - successful_requests=0, - failed_requests=await count_error_file_failed_requests( - batch, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params + return ( + BatchCostUsageResult( + cost=0.0, + usage=Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0), + models=[], + successful_requests=0, + failed_requests=error_file_failed_requests, ), + BatchResultFiles(output=None, error=error_file_content), ) file_content = await _fetch_batch_output_file_content(batch, custom_llm_provider, litellm_params=litellm_params) - error_file_failed_requests: Final = await count_error_file_failed_requests( - batch, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params - ) output_file_result: Final = ( calculate_vertex_ai_batch_cost_and_usage( @@ -155,11 +178,14 @@ async def _handle_completed_batch( ) ) - if not error_file_failed_requests: - return output_file_result - return dataclasses_replace( - output_file_result, failed_requests=output_file_result.failed_requests + error_file_failed_requests + result: Final = ( + output_file_result + if not error_file_failed_requests + else dataclasses_replace( + output_file_result, failed_requests=output_file_result.failed_requests + error_file_failed_requests + ) ) + return result, BatchResultFiles(output=file_content, error=error_file_content) class _LineOutcome(Enum): @@ -435,7 +461,9 @@ def _provider_output_file_id(output_file_id: str) -> str: async def _fetch_batch_managed_file_content( file_id: str, - custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"] = "openai", + custom_llm_provider: Literal[ + "openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral" + ] = "openai", litellm_params: dict | None = None, ) -> bytes: """ @@ -465,7 +493,9 @@ async def _fetch_batch_managed_file_content( async def _fetch_batch_output_file_content( batch: Batch, - custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"] = "openai", + custom_llm_provider: Literal[ + "openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral" + ] = "openai", litellm_params: dict | None = None, ) -> bytes: """ @@ -485,9 +515,30 @@ async def _fetch_batch_output_file_content( ) +async def fetch_batch_error_file_content( + batch: Batch, + custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"], + litellm_params: dict[str, object] | None, # mutable-ok: file fetchers require a plain dict +) -> bytes | None: + """Fetch the batch's separate error file bytes; None when it has none or the fetch fails.""" + if batch.error_file_id is None: + return None + try: + return await _fetch_batch_managed_file_content( + batch.error_file_id, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params + ) + except Exception as e: # noqa: BLE001 # a failed/missing error file must not abort cost tracking for the batch + verbose_logger.debug("Failed to fetch batch error file %s: %s", batch.error_file_id, e) + return None + + +def count_error_file_failed_requests_from_content(content: bytes | None) -> int: + return 0 if content is None else sum(1 for _ in _iter_batch_input_lines(content)) + + async def count_error_file_failed_requests( batch: Batch, - custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"], + custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"], litellm_params: dict | None, ) -> int: """Count failed requests reported only in the batch's separate error file. @@ -497,16 +548,11 @@ async def count_error_file_failed_requests( ``error_file_id`` - they never appear in the output file at all, so counting failures from the output file alone silently undercounts them. """ - if batch.error_file_id is None: - return 0 - try: - error_file_content = await _fetch_batch_managed_file_content( - batch.error_file_id, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params + return count_error_file_failed_requests_from_content( + await fetch_batch_error_file_content( + batch, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params ) - except Exception as e: # noqa: BLE001 # a failed/missing error file must not abort cost tracking for the batch - verbose_logger.debug("Failed to fetch batch error file %s: %s", batch.error_file_id, e) - return 0 - return sum(1 for _ in _iter_batch_input_lines(error_file_content)) + ) def _extract_file_access_credentials(litellm_params: dict | None) -> dict: diff --git a/litellm/integrations/clickhouse/clickhouse_spend_logger.py b/litellm/integrations/clickhouse/clickhouse_spend_logger.py index 588f81ec50a..4a10e29c702 100644 --- a/litellm/integrations/clickhouse/clickhouse_spend_logger.py +++ b/litellm/integrations/clickhouse/clickhouse_spend_logger.py @@ -19,6 +19,7 @@ from litellm._logging import verbose_logger from litellm.integrations.clickhouse.clickhouse_batch_logger import ClickHouseBatchLogger from litellm.integrations.clickhouse.context import is_lens_analysis from litellm.integrations.clickhouse.schema import SPEND_LOGS_TABLE +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.litellm_core_utils.sensitive_data_masker import redact_credentials_in_payload from litellm.tracing.types import SpendLogRecord from litellm.types.utils import StandardLoggingPayload @@ -197,9 +198,17 @@ class ClickHouseSpendLogger(ClickHouseBatchLogger): table = SPEND_LOGS_TABLE async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: + # Batch line items are billed by the aggregate aretrieve_batch row; a per-line + # spend row here would bill the batch twice. + if is_batch_line_item_event(kwargs): + return self._log(kwargs) async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None: + # Failed batch lines are part of the same aggregate row's request counts; + # per-line rows here would inflate spend-log request counts. + if is_batch_line_item_event(kwargs): + return self._log(kwargs) def _log(self, kwargs: Mapping[str, Any]) -> None: diff --git a/litellm/integrations/datadog/datadog_cost_management.py b/litellm/integrations/datadog/datadog_cost_management.py index 538dd95abdd..73017b79f7b 100644 --- a/litellm/integrations/datadog/datadog_cost_management.py +++ b/litellm/integrations/datadog/datadog_cost_management.py @@ -14,6 +14,7 @@ from litellm.integrations.datadog.datadog_handler import ( get_datadog_service, normalize_datadog_tag_value, ) +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, @@ -72,6 +73,10 @@ class DatadogCostManagementLogger(CustomBatchLogger): super().__init__(**kwargs) async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + # A batch line item is billed by the aggregate aretrieve_batch event; a + # per-line FOCUS BilledCost row would double-count cloud spend. + if is_batch_line_item_event(kwargs): + return try: standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None) diff --git a/litellm/integrations/lago.py b/litellm/integrations/lago.py index 594427b1e0a..193f809a5fd 100644 --- a/litellm/integrations/lago.py +++ b/litellm/integrations/lago.py @@ -11,6 +11,7 @@ import litellm from litellm._logging import verbose_logger from litellm._uuid import uuid from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.llms.custom_httpx.http_handler import ( HTTPHandler, get_async_httpx_client, @@ -121,6 +122,9 @@ class LagoLogger(CustomLogger): return returned_val def log_success_event(self, kwargs, response_obj, start_time, end_time): + # A batch line item is billed by the aggregate aretrieve_batch event. + if is_batch_line_item_event(kwargs): + return _url = os.getenv("LAGO_API_BASE") assert _url is not None and isinstance(_url, str), ( f"LAGO_API_BASE missing or not set correctly. LAGO_API_BASE={_url}" @@ -153,6 +157,8 @@ class LagoLogger(CustomLogger): raise e async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + if is_batch_line_item_event(kwargs): + return try: verbose_logger.debug("ENTERS LAGO CALLBACK") _url = os.getenv("LAGO_API_BASE") diff --git a/litellm/integrations/newrelic/newrelic_metrics.py b/litellm/integrations/newrelic/newrelic_metrics.py index a2cedeba0f5..8bd77c981f6 100644 --- a/litellm/integrations/newrelic/newrelic_metrics.py +++ b/litellm/integrations/newrelic/newrelic_metrics.py @@ -34,6 +34,7 @@ from httpx import HTTPStatusError, Response from litellm._logging import verbose_logger from litellm.integrations.custom_batch_logger import CustomBatchLogger +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, @@ -301,12 +302,17 @@ class NewRelicMetricsLogger(CustomBatchLogger): await self._final_drain() async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: + # A batch line item is metered by the aggregate aretrieve_batch event. + if is_batch_line_item_event(kwargs): + return try: await self._log_async_event(standard_logging_object=kwargs.get("standard_logging_object", None)) except Exception as e: # noqa: BLE001 # logging must never break the request path verbose_logger.exception("New Relic Metrics Layer Error - %s\n%s", e, traceback.format_exc()) async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None: + if is_batch_line_item_event(kwargs): + return try: await self._log_async_event(standard_logging_object=kwargs.get("standard_logging_object", None)) except Exception as e: # noqa: BLE001 # logging must never break the request path diff --git a/litellm/integrations/openmeter.py b/litellm/integrations/openmeter.py index db2fe386dec..029c684f969 100644 --- a/litellm/integrations/openmeter.py +++ b/litellm/integrations/openmeter.py @@ -9,6 +9,7 @@ import httpx import litellm from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.llms.custom_httpx.http_handler import ( HTTPHandler, get_async_httpx_client, @@ -97,6 +98,9 @@ class OpenMeterLogger(CustomLogger): } def log_success_event(self, kwargs, response_obj, start_time, end_time): + # A batch line item is billed by the aggregate aretrieve_batch event. + if is_batch_line_item_event(kwargs): + return _url = os.getenv("OPENMETER_API_ENDPOINT", "https://openmeter.cloud") if _url.endswith("/"): _url += "api/v1/events" @@ -123,6 +127,8 @@ class OpenMeterLogger(CustomLogger): raise e async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + if is_batch_line_item_event(kwargs): + return _url = os.getenv("OPENMETER_API_ENDPOINT", "https://openmeter.cloud") if _url.endswith("/"): _url += "api/v1/events" diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 956e754944a..203e9a7b528 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -30,6 +30,7 @@ from litellm.integrations.otel.model.db_endpoint import db_span_attributes from litellm.integrations.otel.model.metadata import flatten_metadata from litellm.integrations.otel.model.semconv import LiteLLM, Metric from litellm.integrations.otel.plumbing.otlp_tls import resolve_otlp_http_tls +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.litellm_core_utils.internal_call_metadata import is_unbilled_non_inference_call_from_params from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.litellm_core_utils.secret_redaction import redact_string @@ -1659,6 +1660,11 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): return True def _record_metrics(self, kwargs, response_obj, start_time, end_time): + # A batch line item's tokens and cost are metered by the aggregate + # aretrieve_batch event; per-line samples here would double-count them. + # Spans for line items are still emitted by _handle_success. + if is_batch_line_item_event(kwargs): + return duration_s: Final = (end_time - start_time).total_seconds() params: Final = kwargs.get("litellm_params") or {} provider: Final = _provider_label(params.get("custom_llm_provider")) diff --git a/litellm/integrations/otel/plumbing/metrics.py b/litellm/integrations/otel/plumbing/metrics.py index 76b23467679..b6f96f9ff66 100644 --- a/litellm/integrations/otel/plumbing/metrics.py +++ b/litellm/integrations/otel/plumbing/metrics.py @@ -33,6 +33,7 @@ from litellm.integrations.otel.model.semconv import ( resolve_provider, ) from litellm.integrations.otel.model.utils import to_seconds +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.litellm_core_utils.internal_call_metadata import is_unbilled_non_inference_call_from_params from litellm.litellm_core_utils.safe_json_dumps import safe_dumps @@ -221,6 +222,12 @@ class GenAIMetricRecorder: start_time: datetime, end_time: datetime, ) -> None: + # A batch line item is metered by the aggregate aretrieve_batch event, and + # its child interval (parent retrieve start -> line emission) measures + # retrieval and callback processing rather than the line's model call, so + # skip every metric for it; spans for line items are still emitted. + if is_batch_line_item_event(kwargs): + return common_attrs: Final = self._filter_attributes(self._bounded_attributes(kwargs)) duration_s: Final = (end_time - start_time).total_seconds() usage_is_replayed: Final = is_unbilled_non_inference_call_from_params( @@ -246,6 +253,10 @@ class GenAIMetricRecorder: start_time: datetime, end_time: datetime, ) -> None: + # A batch line item's interval measures retrieval and emission, not the + # line's model call, so its duration would be synthetic here too. + if is_batch_line_item_event(kwargs): + return """Record the one metric a failed request can honestly report: the operation's duration, tagged with ``error.type``. diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 6468dc41ea5..47b5cae63fe 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -35,6 +35,7 @@ from litellm.integrations.prometheus_helpers.bounded_prometheus_series_tracker i from litellm.litellm_core_utils.core_helpers import ( get_litellm_metadata_from_kwargs, get_metadata_variable_name_from_kwargs, + is_batch_line_item_event, ) from litellm.litellm_core_utils.service_tier_utils import ( get_service_tier_from_standard_logging_payload, @@ -1345,6 +1346,10 @@ class PrometheusLogger(CustomLogger): self._track_end_user_metric_series(counter, metric_name, _labels) async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + # A batch line item is metered by the aggregate aretrieve_batch event; a + # per-line sample here would double-count spend and requests. + if is_batch_line_item_event(kwargs): + return # Define prometheus client verbose_logger.debug( "prometheus Logging - Enters success logging function (kwargs keys: %s)", @@ -2365,6 +2370,8 @@ class PrometheusLogger(CustomLogger): ) async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): + if is_batch_line_item_event(kwargs): + return verbose_logger.debug( "prometheus Logging - Enters failure logging function (kwargs keys: %s)", list(kwargs.keys()) if isinstance(kwargs, dict) else type(kwargs).__name__, diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index 39e95fbf687..3333c1106fe 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -18,9 +18,9 @@ if TYPE_CHECKING: from litellm.types.utils import ModelResponseStream - Span = _Span | Any + Span = _Span | object else: - Span = Any + Span = object _CODEX_CLIENT_PREFIX_RE: Final = re.compile(r"^codex[-_ /]", re.IGNORECASE) @@ -270,7 +270,7 @@ def remove_index_from_tool_calls( tool_call.pop("index", None) -def remove_items_at_indices(items: list[Any] | None, indices: Iterable[int]) -> None: +def remove_items_at_indices(items: list[object] | None, indices: Iterable[int]) -> None: """Remove items from a list in-place by index""" if items is None: return @@ -713,9 +713,9 @@ def filter_internal_params(data: dict, additional_internal_params: set | None = def redact_nested_match_and_regex_keys( - payload: dict | list[Any] | str | None, + payload: dict | list[object] | str | None, keys: Collection[str] = ("match", "regex"), -) -> dict | list[Any] | str | None: +) -> dict | list[object] | str | None: """ Deep-copy `payload` and replace every configured string field with "[REDACTED]" anywhere in nested dict/list structures. @@ -725,7 +725,7 @@ def redact_nested_match_and_regex_keys( if payload is None or isinstance(payload, str): return payload try: - redacted: Final[dict | list[Any] | str | None] = copy.deepcopy(payload) + redacted: Final[dict | list[object] | str | None] = copy.deepcopy(payload) except Exception: return payload @@ -772,6 +772,15 @@ def set_response_cost_in_hidden_params(response: _CarriesHiddenParams, cost: flo hidden_params["additional_headers"] = merged +def is_batch_line_item_event(kwargs: object) -> bool: + if not isinstance(kwargs, Mapping): + return False + litellm_params: Final[object] = kwargs.get("litellm_params") + if not isinstance(litellm_params, Mapping): + return False + return bool(litellm_params.get("batch_parent_id")) + + _HIDDEN_PARAMS_ADAPTER: Final = TypeAdapter(Mapping[str, object]) _PROVIDER_HEADERS_ADAPTER: Final = TypeAdapter(Mapping[str, str]) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index b2c22880f5a..d92f0be21a9 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -32,7 +32,11 @@ from litellm._logging import ( verbose_logger, ) from litellm._uuid import uuid -from litellm.batches.batch_utils import _handle_completed_batch, batch_cost_is_final +from litellm.batches.batch_utils import ( + BatchResultFiles, + _handle_completed_batch_with_files, + batch_cost_is_final, +) from litellm.caching.caching import DualCache from litellm.caching.caching_handler import LLMCachingHandler from litellm.caching.redis_batch import flush_post_call_redis_batches @@ -652,6 +656,7 @@ class Logging(LiteLLMLoggingBaseClass): self.streaming_chunks: list[object] = [] self.sync_streaming_chunks: list[object] = [] self.log_raw_request_response = log_raw_request_response + self._litellm_internal_model_credentials: Mapping[str, object] | None = None self.raw_request_only = raw_request_only # Initialize dynamic callbacks @@ -947,6 +952,10 @@ class Logging(LiteLLMLoggingBaseClass): **additional_params, } ) + # Provider params / additional kwargs are caller-influenced; they must not + # overwrite the internal litellm_params, or a request body carrying + # `litellm_params` could forge logging markers (e.g. batch_parent_id). + self.model_call_details["litellm_params"] = self.litellm_params ## check if stream options is set ## - used by CustomStreamWrapper for easy instrumentation if "stream_options" in additional_params: @@ -3249,9 +3258,12 @@ class Logging(LiteLLMLoggingBaseClass): batch_models = kwargs.get("batch_models", None) batch_successful_requests: Final = kwargs.get("batch_successful_requests", None) batch_failed_requests: Final = kwargs.get("batch_failed_requests", None) + batch_output_file_content: Final = kwargs.get("batch_output_file_content", None) + batch_error_file_content: Final = kwargs.get("batch_error_file_content", None) has_explicit_batch_data: Final = all(x is not None for x in (batch_cost, batch_usage, batch_models)) should_compute_batch_data: Final = not has_explicit_batch_data and batch_cost_is_final(result) + result_files: BatchResultFiles | None = None # rebind-ok: one batch-data branch below supplies the bytes if has_explicit_batch_data: result._hidden_params["response_cost"] = batch_cost result._hidden_params["batch_models"] = batch_models @@ -3271,15 +3283,21 @@ class Logging(LiteLLMLoggingBaseClass): total_cost=batch_cost, cost_for_built_in_tools_cost_usd_dollar=0.0, ) + if batch_output_file_content is not None or batch_error_file_content is not None: + result_files = BatchResultFiles( + output=batch_output_file_content if isinstance(batch_output_file_content, bytes) else None, + error=batch_error_file_content if isinstance(batch_error_file_content, bytes) else None, + ) elif should_compute_batch_data: - batch_result: Final = await _handle_completed_batch( + batch_result, fetched_result_files = await _handle_completed_batch_with_files( batch=result, custom_llm_provider=self.custom_llm_provider, model_name=self.get_deployment_model_for_cost(), litellm_params=self.litellm_params, model_info=self.get_router_deployment_model_info(), ) + result_files = fetched_result_files result._hidden_params["response_cost"] = batch_result.cost result._hidden_params["batch_models"] = batch_result.models @@ -3293,6 +3311,28 @@ class Logging(LiteLLMLoggingBaseClass): cost_for_built_in_tools_cost_usd_dollar=0.0, ) + if litellm.store_batch_line_items_in_callbacks and (has_explicit_batch_data or should_compute_batch_data): + from litellm.batches.batch_line_item_logging import log_batch_line_items + + _line_item_litellm_params: Final = cast( # cast-ok: shared attribute is an untyped dict + "dict[str, object] | None", self.litellm_params + ) + try: + await log_batch_line_items( + batch=result, + custom_llm_provider=self.custom_llm_provider, + parent=self, + model_name=self.get_deployment_model_for_cost(), + litellm_params=_line_item_litellm_params, + model_info=self.get_router_deployment_model_info(), + result_files=result_files, + ) + except Exception: # noqa: BLE001 # line-item logging (claim step included) must never reach the aggregate aretrieve_batch path + verbose_logger.exception( + "batch line item logging failed for batch_id=%s; aggregate logging unaffected", + result.id, + ) + self.truncated_messages_for_logging = await truncate_base64_in_messages_async( StandardLoggingPayloadSetup.append_system_prompt_messages( kwargs=self.model_call_details, messages=self.model_call_details.get("messages") @@ -4247,19 +4287,10 @@ class Logging(LiteLLMLoggingBaseClass): verbose_logger.warning("LiteLLM: the anthropic_messages stream assembled no response, logging an empty one") return litellm.ModelResponse(model=self.model) else: - from litellm.types.llms.anthropic import AnthropicResponse + from litellm.llms.anthropic.chat.transformation import anthropic_message_to_model_response - pydantic_result: Final = AnthropicResponse.model_validate(result) - import httpx - - result = litellm.AnthropicConfig().transform_parsed_response( - completion_response=pydantic_result.model_dump(), - raw_response=httpx.Response( - status_code=200, - headers={}, - ), - model_response=litellm.ModelResponse(id=provider_response_id), - json_mode=None, + result = anthropic_message_to_model_response( + cast(Mapping[str, object], result), # cast-ok: handler result is typed Any upstream speed=self.optional_params.get("speed") if self.optional_params else None, ) return result @@ -6041,6 +6072,9 @@ class StandardLoggingPayloadSetup: batch_failed_requests=None, litellm_model_name=None, usage_object=None, + batch_id=None, + batch_custom_id=None, + batch_line_status_code=None, ) if hidden_params is not None: for key in StandardLoggingHiddenParams.__annotations__: @@ -6446,8 +6480,12 @@ def _extract_response_obj_and_hidden_params( response_obj = {} if original_exception is not None and hidden_params is None: - response_headers: Final = _get_response_headers(original_exception) - if response_headers is not None: + exception_hidden_params: Final[Mapping[str, object] | None] = getattr( + original_exception, "_hidden_params", None + ) + if isinstance(exception_hidden_params, dict) and exception_hidden_params: + hidden_params = dict(exception_hidden_params) + elif (response_headers := _get_response_headers(original_exception)) is not None: hidden_params = dict( StandardLoggingHiddenParams( additional_headers=StandardLoggingPayloadSetup.get_additional_headers(dict(response_headers)), diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index fd67ebd5293..fb70e39c5f1 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -52,6 +52,7 @@ from litellm.types.llms.anthropic import ( AnthropicMessagesToolChoice, AnthropicOutputSchema, AnthropicOutputTokensDetails, + AnthropicResponse, AnthropicSystemMessageContent, AnthropicThinkingParam, AnthropicWebSearchTool, @@ -2808,3 +2809,18 @@ def _valid_user_id(user_id: str) -> bool: return False return True + + +def anthropic_message_to_model_response(result: Mapping[str, object], speed: str | None) -> ModelResponse: + pydantic_result: Final = AnthropicResponse.model_validate(result) + result_id: Final = result.get("id") + return AnthropicConfig().transform_parsed_response( + completion_response=pydantic_result.model_dump(), + raw_response=httpx.Response( + status_code=200, + headers={}, + ), + model_response=ModelResponse(id=result_id if isinstance(result_id, str) and result_id else None), + json_mode=None, + speed=speed, + ) diff --git a/litellm/llms/base_llm/batches/transformation.py b/litellm/llms/base_llm/batches/transformation.py index 34c622d4cf6..904eae2f4f2 100644 --- a/litellm/llms/base_llm/batches/transformation.py +++ b/litellm/llms/base_llm/batches/transformation.py @@ -1,5 +1,6 @@ import types from abc import ABC, abstractmethod +from collections.abc import Mapping from typing import TYPE_CHECKING, Any import httpx @@ -9,7 +10,7 @@ from litellm.types.llms.openai import ( AllMessageValues, CreateBatchRequest, ) -from litellm.types.utils import LiteLLMBatch, LlmProviders +from litellm.types.utils import EmbeddingResponse, LiteLLMBatch, LlmProviders, ModelResponse if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj @@ -39,6 +40,14 @@ class BaseBatchesConfig(ABC): def custom_llm_provider(self) -> LlmProviders: """Return the LLM provider type for this configuration.""" + def transform_batch_output_line( + self, model_output: Mapping[str, object], model: str + ) -> ModelResponse | EmbeddingResponse | None: + """Reconstruct one provider batch output line into a litellm response, or + None when the line is OpenAI-shaped and the caller can use the generic + reconstruction.""" + return None + @classmethod def get_config(cls): """Get configuration dictionary for this class.""" diff --git a/litellm/llms/bedrock/batches/transformation.py b/litellm/llms/bedrock/batches/transformation.py index e4001566b8c..32fdf7e8997 100644 --- a/litellm/llms/bedrock/batches/transformation.py +++ b/litellm/llms/bedrock/batches/transformation.py @@ -27,7 +27,7 @@ from litellm.types.llms.openai import ( AllMessageValues, CreateBatchRequest, ) -from litellm.types.utils import LiteLLMBatch, LlmProviders, Usage +from litellm.types.utils import EmbeddingResponse, LiteLLMBatch, LlmProviders, ModelResponse, Usage from ..base_aws_llm import BaseAWSLLM from ..common_utils import ( @@ -96,6 +96,47 @@ def titan_embedding_usage_from_batch_output(model_output: Mapping[str, object]) ) +def bedrock_batch_line_to_response( + model_output: Mapping[str, object], model: str +) -> ModelResponse | EmbeddingResponse | None: + """Reconstruct a Bedrock batch output line (the ``modelOutput`` object) into + the litellm response type its shape implies, or None when the shape is + unrecognized.""" + if "embedding" in model_output: + embedding: Final = model_output.get("embedding") + return EmbeddingResponse( + model=model, + data=[ + { + "object": "embedding", + "index": 0, + "embedding": embedding if isinstance(embedding, list) else [], + } + ], + usage=titan_embedding_usage_from_batch_output(model_output), + ) + if "output" in model_output: + from ..chat.converse_transformation import AmazonConverseConfig + + return AmazonConverseConfig()._transform_response( # pyright: ignore[reportPrivateUsage] # same reconstruction the converse chat path performs on the live response + model=model, + response=Response(200, json=dict(model_output)), + model_response=ModelResponse(), + stream=False, + logging_obj=None, + optional_params={}, + api_key=None, + data="", + messages=[], + encoding=None, + ) + if "content" in model_output: + from litellm.llms.anthropic.chat.transformation import anthropic_message_to_model_response + + return anthropic_message_to_model_response(model_output, None) + return None + + class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): """ Config for Bedrock Batches - handles batch job creation and management for Bedrock @@ -109,6 +150,11 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): def custom_llm_provider(self) -> LlmProviders: return LlmProviders.BEDROCK + def transform_batch_output_line( + self, model_output: Mapping[str, object], model: str + ) -> ModelResponse | EmbeddingResponse | None: + return bedrock_batch_line_to_response(model_output, model) + @classmethod def _get_bare_model_name_from_s3_key(cls, object_key: str) -> str | None: if not object_key.startswith(BEDROCK_MANAGED_S3_BATCH_PREFIX): diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index b993a5d8d6d..e35e224545e 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -3037,6 +3037,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): None, description="If True, stores request messages and responses in spend logs. Default is False.", ) + store_batch_line_items_in_callbacks: bool | None = Field( + None, + description="If True, a completed batch logged via aretrieve_batch also emits one callback event per JSONL line item (request paired with its response or error). The aggregate batch callback is unchanged. Default is False.", + ) disable_auto_add_proxy_admin_to_teams: bool | None = Field( None, description="By default, the user calling /team/new is automatically added to the new team as a team admin. If True, proxy admins are no longer auto-added; members explicitly listed in members_with_roles are unaffected. Default is False.", diff --git a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py index d600e249754..2af19d1afb8 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py @@ -15,6 +15,7 @@ from litellm._internal_context import with_service_target from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.proxy_rate_limit_error import ( ProxyRateLimitError, @@ -697,6 +698,8 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): - model_saturation_check: Model-wide token tracking - priority_model: Priority-specific token tracking """ + if is_batch_line_item_event(kwargs): + return from litellm.litellm_core_utils.core_helpers import ( _get_parent_otel_span_from_kwargs, ) diff --git a/litellm/proxy/hooks/max_budget_per_session_limiter.py b/litellm/proxy/hooks/max_budget_per_session_limiter.py index 2ba42c43dcc..45c5b1c26f7 100644 --- a/litellm/proxy/hooks/max_budget_per_session_limiter.py +++ b/litellm/proxy/hooks/max_budget_per_session_limiter.py @@ -24,6 +24,7 @@ from litellm._logging import verbose_proxy_logger from litellm.caching.redis_cache import log_redis_failure from litellm.exceptions import RateLimitType from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit @@ -134,6 +135,8 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): """ After a successful LLM call, increment the session spend by the response cost. """ + if is_batch_line_item_event(kwargs): # pyright: ignore[reportUnknownArgumentType] # hook kwargs arrive untyped from the logging dispatcher + return try: litellm_params: Final = kwargs.get("litellm_params") or {} metadata: Final = litellm_params.get("metadata") or {} diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index d019271d404..9623e87b3f9 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -13,6 +13,7 @@ from litellm._internal_context import with_service_target from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import Span +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.llms.bedrock.common_utils import get_bedrock_base_model from litellm.proxy._types import Litellm_EntityType, UserAPIKeyAuth @@ -499,6 +500,8 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): Example: key=sk-1234567890, model=gpt-4o, max_budget=100, time_period=1d """ + if is_batch_line_item_event(kwargs): + return verbose_proxy_logger.debug("in RouterBudgetLimiting.async_log_success_event") standard_logging_payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None) if standard_logging_payload is None: diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index b4ce010dd27..21ebeb81d80 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -12,7 +12,10 @@ from litellm._internal_context import with_service_target from litellm._logging import verbose_proxy_logger from litellm.exceptions import RateLimitType from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs +from litellm.litellm_core_utils.core_helpers import ( + _get_parent_otel_span_from_kwargs, + is_batch_line_item_event, +) from litellm.proxy._types import CommonProxyErrors, CurrentItemRateLimit, UserAPIKeyAuth from litellm.proxy.auth.auth_utils import ( get_key_model_rpm_limit, @@ -495,6 +498,8 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): @with_service_target("rate_limits") async def async_log_success_event(self, kwargs, response_obj: object, start_time, end_time): + if is_batch_line_item_event(kwargs): + return from litellm.proxy.common_utils.callback_utils import ( get_model_group_from_litellm_kwargs, ) @@ -701,6 +706,8 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): @with_service_target("rate_limits") async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): + if is_batch_line_item_event(kwargs): + return try: self.print_verbose("Inside Max Parallel Request Failure Hook") litellm_parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(kwargs=kwargs) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 2bfaf57f0fc..8abc36ee89c 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -43,6 +43,7 @@ from litellm.caching.redis_batch import ( from litellm.caching.redis_cache import log_redis_failure from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE, INTERNAL_CALL_ORIGIN_METADATA_KEY from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.litellm_core_utils.prompt_templates.common_utils import ( get_str_from_messages, ) @@ -5001,6 +5002,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): """ Update TPM usage on successful API calls by incrementing counters using pipeline """ + if is_batch_line_item_event(kwargs): + return from litellm.litellm_core_utils.core_helpers import ( _get_parent_otel_span_from_kwargs, ) @@ -5124,6 +5127,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): whose partial usage was recovered settles the reservation at that usage instead of refunding it. """ + if is_batch_line_item_event(kwargs): + return from litellm.litellm_core_utils.core_helpers import ( _get_parent_otel_span_from_kwargs, ) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index f937b439042..73415825b19 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -14,6 +14,7 @@ from litellm.litellm_core_utils.core_helpers import ( budget_reservation_from_metadata, get_litellm_metadata_from_kwargs, get_metadata_variable_name_from_kwargs, + is_batch_line_item_event, ) from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import guardrail_information_cost @@ -121,6 +122,8 @@ class _ProxyDBLogger(CustomLogger): async def async_log_success_event( self, kwargs: ObjectMapping, response_obj: object, start_time: datetime, end_time: datetime ) -> None: + if is_batch_line_item_event(kwargs): + return if self.spend_event_producer is None or not is_offloadable_success(response_obj): await self._PROXY_track_cost_callback(kwargs, response_obj, start_time, end_time) return diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 6d1a04c2b5e..6be1475074a 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -270,6 +270,10 @@ import litellm._redis from litellm import Router from litellm._internal_context import service_target, with_service_target 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.dual_cache import DeclaredBatchRead from litellm.caching.redis_batch import ( @@ -4982,8 +4986,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 @@ -4997,6 +5001,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, @@ -6910,6 +6918,13 @@ class ProxyConfig: health_check_interval = general_settings.get("health_check_interval", DEFAULT_HEALTH_CHECK_INTERVAL) health_check_concurrency = general_settings.get("health_check_concurrency", None) health_check_details = general_settings.get("health_check_details", True) + ### BATCH LINE ITEM CALLBACKS ### + _store_batch_line_items: Final[object] = general_settings.get("store_batch_line_items_in_callbacks") + if _store_batch_line_items is not None: + if isinstance(_store_batch_line_items, str): + litellm.store_batch_line_items_in_callbacks = _store_batch_line_items.lower() == "true" + else: + litellm.store_batch_line_items_in_callbacks = bool(_store_batch_line_items) # Health-check-driven routing (opt-in, passes through to Router later) _enable_hc_routing = general_settings.get("enable_health_check_routing", False) _hc_staleness = general_settings.get("health_check_staleness_threshold", None) @@ -7921,6 +7936,7 @@ class ProxyConfig: self._apply_alerting_settings, self._apply_pass_through_settings, self._apply_boolean_settings, + self._apply_batch_line_items_setting, partial(self._apply_cache_size_setting, cache_size_was_db=cache_size_was_db), self._apply_store_model_in_db_setting, partial(self._apply_retention_settings, previous_cleanup_schedule=previous_cleanup_schedule), @@ -7980,6 +7996,11 @@ class ProxyConfig: if (value := self.settings.get(key)) is not None: self.settings[key] = coerce_bool(value) + async def _apply_batch_line_items_setting(self, db_values: Mapping[str, SettingsJsonValue]) -> None: + key: Final = "store_batch_line_items_in_callbacks" + value: Final = coerce_bool(self.settings.get(key)) + litellm.store_batch_line_items_in_callbacks = bool(value) if value is not None else False + async def _apply_cache_size_setting( self, db_values: Mapping[str, SettingsJsonValue], diff --git a/litellm/router.py b/litellm/router.py index afc842a05a0..ea5301e7c21 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -80,6 +80,7 @@ from litellm.litellm_core_utils.core_helpers import ( get_litellm_metadata_from_kwargs, get_metadata_variable_name_from_kwargs, get_or_create_metadata_bucket, + is_batch_line_item_event, ) from litellm.litellm_core_utils.coroutine_checker import coroutine_checker from litellm.litellm_core_utils.credential_accessor import CredentialAccessor @@ -8089,7 +8090,7 @@ class Router: # WS session wrappers fire with result=None; per-turn costs tracked by inner calls. if kwargs.get("call_type") in ("_aresponses_websocket", "_arealtime"): return - if is_batch_retrieve_call_type(kwargs.get("call_type")): + if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs): return standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None) if standard_logging_object is None: @@ -8237,7 +8238,7 @@ class Router: - key: str - The key used to increment the cache - None: if no key is found """ - if is_batch_retrieve_call_type(kwargs.get("call_type")): + if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs): return None id = None if kwargs["litellm_params"].get("metadata") is None: @@ -8367,7 +8368,7 @@ class Router: """ Update RPM usage for a deployment """ - if is_batch_retrieve_call_type(kwargs.get("call_type")): + if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs): return deployment_name: Final = kwargs["litellm_params"]["metadata"].get( "deployment", None diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index 631b0c3df3d..2ef783bc1cb 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -35,6 +35,7 @@ from litellm.caching.redis_cache import RedisCache, RedisPipelineIncrementOperat from litellm.integrations.custom_logger import CustomLogger, Span from litellm.litellm_core_utils.core_helpers import ( get_metadata_variable_name_from_kwargs, + is_batch_line_item_event, ) from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs @@ -495,6 +496,8 @@ class RouterBudgetLimiting(CustomLogger): async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): """Original method now uses helper functions""" verbose_router_logger.debug("in RouterBudgetLimiting.async_log_success_event") + if is_batch_line_item_event(kwargs): + return # WS session wrappers fire with result=None; per-turn costs tracked by inner calls. if kwargs.get("call_type") in ("_aresponses_websocket", "_arealtime"): return diff --git a/litellm/router_strategy/least_busy.py b/litellm/router_strategy/least_busy.py index 6f3e0936641..7cfada62ecc 100644 --- a/litellm/router_strategy/least_busy.py +++ b/litellm/router_strategy/least_busy.py @@ -10,6 +10,7 @@ from litellm._logging import verbose_router_logger from litellm.caching.caching import DualCache from litellm.caching.redis_cache import log_redis_failure from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.router_utils.batch_utils import is_batch_retrieve_call_type IN_FLIGHT_COUNT_TTL_SECONDS: Final = 60 * 60 @@ -50,7 +51,7 @@ def _request_count_key(model_group: str, deployment_id: str) -> str: def _deployment_ref(kwargs: Mapping[str, object]) -> tuple[str, str] | None: - if is_batch_retrieve_call_type(kwargs.get("call_type")): + if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs): return None try: call: Final = _CALL_KWARGS.validate_python(kwargs) diff --git a/litellm/router_strategy/lowest_cost.py b/litellm/router_strategy/lowest_cost.py index d567b6acccc..1bcf248331d 100644 --- a/litellm/router_strategy/lowest_cost.py +++ b/litellm/router_strategy/lowest_cost.py @@ -9,6 +9,7 @@ from litellm._internal_context import with_service_target from litellm._logging import verbose_router_logger from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.router_utils.batch_utils import is_batch_retrieve_call_type @@ -22,7 +23,7 @@ class LowestCostLoggingHandler(CustomLogger): @with_service_target("router_usage") def log_success_event(self, kwargs, response_obj, start_time, end_time): - if is_batch_retrieve_call_type(kwargs.get("call_type")): + if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs): return try: """ @@ -98,7 +99,7 @@ class LowestCostLoggingHandler(CustomLogger): @with_service_target("router_usage") async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - if is_batch_retrieve_call_type(kwargs.get("call_type")): + if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs): return try: """ diff --git a/litellm/router_strategy/lowest_latency.py b/litellm/router_strategy/lowest_latency.py index 622919e3443..d7980c09dff 100644 --- a/litellm/router_strategy/lowest_latency.py +++ b/litellm/router_strategy/lowest_latency.py @@ -13,7 +13,11 @@ from litellm import ModelResponse, token_counter, verbose_logger from litellm._internal_context import with_service_target from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs, safe_divide_seconds +from litellm.litellm_core_utils.core_helpers import ( + _get_parent_otel_span_from_kwargs, + is_batch_line_item_event, + safe_divide_seconds, +) from litellm.router_utils.batch_utils import is_batch_retrieve_call_type from litellm.types.utils import LiteLLMPydanticObjectBase @@ -61,7 +65,7 @@ class LowestLatencyLoggingHandler(CustomLogger): @with_service_target("router_usage") def log_success_event(self, kwargs, response_obj, start_time, end_time): - if is_batch_retrieve_call_type(kwargs.get("call_type")): + if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs): return try: """ @@ -188,7 +192,7 @@ class LowestLatencyLoggingHandler(CustomLogger): """ Check if Timeout Error, if timeout set deployment latency -> 100 """ - if is_batch_retrieve_call_type(kwargs.get("call_type")): + if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs): return try: metadata_field: Final = self._select_metadata_field(kwargs) @@ -245,7 +249,7 @@ class LowestLatencyLoggingHandler(CustomLogger): @with_service_target("router_usage") async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - if is_batch_retrieve_call_type(kwargs.get("call_type")): + if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs): return try: """ diff --git a/litellm/router_strategy/lowest_tpm_rpm.py b/litellm/router_strategy/lowest_tpm_rpm.py index 2d373e0c266..11f28e6c7a5 100644 --- a/litellm/router_strategy/lowest_tpm_rpm.py +++ b/litellm/router_strategy/lowest_tpm_rpm.py @@ -9,6 +9,7 @@ from litellm._internal_context import with_service_target from litellm._logging import verbose_router_logger from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event from litellm.router_utils.batch_utils import is_batch_retrieve_call_type from litellm.types.utils import LiteLLMPydanticObjectBase from litellm.utils import print_verbose @@ -30,7 +31,7 @@ class LowestTPMLoggingHandler(CustomLogger): @with_service_target("router_usage") def log_success_event(self, kwargs, response_obj, start_time, end_time): - if is_batch_retrieve_call_type(kwargs.get("call_type")): + if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs): return try: """ @@ -85,7 +86,7 @@ class LowestTPMLoggingHandler(CustomLogger): @with_service_target("router_usage") async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - if is_batch_retrieve_call_type(kwargs.get("call_type")): + if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs): return try: """ diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py index 9839c9be469..d3c5b897ea0 100644 --- a/litellm/router_strategy/lowest_tpm_rpm_v2.py +++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py @@ -15,7 +15,7 @@ from litellm._internal_context import with_service_target from litellm._logging import verbose_logger, verbose_router_logger from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs +from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs, is_batch_line_item_event from litellm.router_utils.batch_utils import is_batch_retrieve_call_type from litellm.types.router import RouterErrors from litellm.types.utils import LiteLLMPydanticObjectBase, StandardLoggingPayload @@ -254,7 +254,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): @with_service_target("router_usage") def log_success_event(self, kwargs, response_obj, start_time, end_time): - if is_batch_retrieve_call_type(kwargs.get("call_type")): + if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs): return try: """ @@ -297,7 +297,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): @with_service_target("router_usage") async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - if is_batch_retrieve_call_type(kwargs.get("call_type")): + if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs): return try: """ diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 142f9b14a72..332b606731a 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3228,6 +3228,9 @@ class StandardLoggingHiddenParams(TypedDict): batch_failed_requests: ReadOnly[int | None] litellm_model_name: str | None # the model name sent to the provider by litellm usage_object: dict | None + batch_id: NotRequired[ReadOnly[str | None]] + batch_custom_id: NotRequired[ReadOnly[str | None]] + batch_line_status_code: NotRequired[ReadOnly[int | None]] class StandardLoggingModelInformation(TypedDict): diff --git a/tests/batches_tests/test_batches_logging_unit_tests.py b/tests/batches_tests/test_batches_logging_unit_tests.py index 73adb391481..e8a90b385f2 100644 --- a/tests/batches_tests/test_batches_logging_unit_tests.py +++ b/tests/batches_tests/test_batches_logging_unit_tests.py @@ -222,8 +222,8 @@ async def test_batch_retrieve_cost_tracking_with_completed_batch_no_explicit_cos ) logging_obj.custom_llm_provider = "openai" - # Mock _handle_completed_batch to return cost data - from litellm.batches.batch_utils import BatchCostUsageResult + # Mock _handle_completed_batch_with_files to return cost data + from litellm.batches.batch_utils import BatchCostUsageResult, BatchResultFiles expected_cost = 0.05 expected_usage = litellm.Usage( @@ -234,14 +234,17 @@ async def test_batch_retrieve_cost_tracking_with_completed_batch_no_explicit_cos expected_models = ["gpt-5-mini"] with patch( - "litellm.litellm_core_utils.litellm_logging._handle_completed_batch", + "litellm.litellm_core_utils.litellm_logging._handle_completed_batch_with_files", new=AsyncMock( - return_value=BatchCostUsageResult( - cost=expected_cost, - usage=expected_usage, - models=expected_models, - successful_requests=10, - failed_requests=0, + return_value=( + BatchCostUsageResult( + cost=expected_cost, + usage=expected_usage, + models=expected_models, + successful_requests=10, + failed_requests=0, + ), + BatchResultFiles(output=None, error=None), ) ), ) as mock_handle_batch: @@ -382,7 +385,7 @@ async def test_batch_retrieve_cost_tracking_with_explicit_cost_data(): explicit_models = ["gpt-5-mini", "gpt-5.5"] with patch( - "litellm.litellm_core_utils.litellm_logging._handle_completed_batch", + "litellm.litellm_core_utils.litellm_logging._handle_completed_batch_with_files", new=AsyncMock(), ) as mock_handle_batch: # Call async_success_handler with explicit cost data @@ -517,7 +520,7 @@ async def test_batch_retrieve_cost_tracking_with_unified_file_id_incomplete_batc logging_obj.custom_llm_provider = "openai" with patch( - "litellm.litellm_core_utils.litellm_logging._handle_completed_batch", + "litellm.litellm_core_utils.litellm_logging._handle_completed_batch_with_files", new=AsyncMock(), ) as mock_handle_batch: # Call async_success_handler with in_progress batch (unified file ID) @@ -603,17 +606,20 @@ async def test_batch_retrieve_cost_tracking_with_partial_explicit_data(): ) expected_models = ["gpt-5-mini"] - from litellm.batches.batch_utils import BatchCostUsageResult + from litellm.batches.batch_utils import BatchCostUsageResult, BatchResultFiles with patch( - "litellm.litellm_core_utils.litellm_logging._handle_completed_batch", + "litellm.litellm_core_utils.litellm_logging._handle_completed_batch_with_files", new=AsyncMock( - return_value=BatchCostUsageResult( - cost=expected_cost, - usage=expected_usage, - models=expected_models, - successful_requests=8, - failed_requests=0, + return_value=( + BatchCostUsageResult( + cost=expected_cost, + usage=expected_usage, + models=expected_models, + successful_requests=8, + failed_requests=0, + ), + BatchResultFiles(output=None, error=None), ) ), ) as mock_handle_batch: diff --git a/tests/integration/spend/test_batch_line_item_callbacks.py b/tests/integration/spend/test_batch_line_item_callbacks.py new file mode 100644 index 00000000000..a707a7a4f7a --- /dev/null +++ b/tests/integration/spend/test_batch_line_item_callbacks.py @@ -0,0 +1,303 @@ +from __future__ import annotations + +import json +import uuid +from hashlib import sha256 +from pathlib import Path +from typing import Final + +import pytest +import yaml +from integration._support.client import JSON_OBJECT, Gateway, eventually, object_value, string_value +from integration._support.database import read_rows +from integration._support.process import owned_proxy +from integration._support.upstream import delete_scenario, register_scenario +from integration._support.wire import Reply, Request, wire_server +from integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse, TextResponse +from pydantic import JsonValue + +PROMPT_TOKENS: Final = 10 +COMPLETION_TOKENS: Final = 100 +OUTPUT_SUCCESS_IDS: Final = ("r1", "r3") +OUTPUT_FAILED_IDS: Final = ("r2",) +ERROR_FILE_IDS: Final = ("r4", "r5") +ALL_CUSTOM_IDS: Final = ("r1", "r2", "r3", "r4", "r5") + + +def _successful_line(custom_id: str) -> str: + return json.dumps( + { + "id": f"batch_req_{custom_id}", + "custom_id": custom_id, + "response": { + "status_code": 200, + "request_id": f"$REQUEST_ID-{custom_id}", + "body": { + "id": f"chatcmpl-$REQUEST_ID-{custom_id}", + "object": "chat.completion", + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": f"answer {custom_id}"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": PROMPT_TOKENS, + "completion_tokens": COMPLETION_TOKENS, + "total_tokens": PROMPT_TOKENS + COMPLETION_TOKENS, + }, + }, + }, + "error": None, + }, + separators=(",", ":"), + ) + + +def _failed_line(custom_id: str) -> str: + return json.dumps( + { + "id": f"batch_req_{custom_id}", + "custom_id": custom_id, + "response": { + "status_code": 400, + "request_id": f"$REQUEST_ID-{custom_id}", + "body": {"error": {"message": f"synthetic line failure {custom_id}", "code": "bad_request"}}, + }, + "error": {"code": "bad_request", "message": f"synthetic line failure {custom_id}"}, + }, + separators=(",", ":"), + ) + + +def _batch(status: str, *, files_ready: bool) -> dict[str, JsonValue]: + return { + "id": "batch-$REQUEST_ID", + "object": "batch", + "endpoint": "/v1/chat/completions", + "errors": None, + "input_file_id": "file-in-$REQUEST_ID", + "completion_window": "24h", + "status": status, + "output_file_id": "file-out-$REQUEST_ID" if files_ready else None, + "error_file_id": "file-err-$REQUEST_ID" if files_ready else None, + "created_at": 1, + "in_progress_at": 1, + "completed_at": 1 if files_ready else None, + "expires_at": 1, + "request_counts": {"total": 5, "completed": 2, "failed": 3}, + "metadata": None, + } + + +def _input_lines(model_name: str, marker: str) -> tuple[str, ...]: + return tuple( + json.dumps( + { + "custom_id": custom_id, + "method": "POST", + "url": "/v1/chat/completions", + "body": {"model": model_name, "messages": [{"role": "user", "content": f"{marker} {custom_id}"}]}, + }, + separators=(",", ":"), + ) + for custom_id in ALL_CUSTOM_IDS + ) + + +def _provider_routes(input_lines: tuple[str, ...]) -> RoutedResponse: + output_lines: Final = (_successful_line("r1"), _failed_line("r2"), _successful_line("r3")) + error_lines: Final = tuple(_failed_line(custom_id) for custom_id in ERROR_FILE_IDS) + return RoutedResponse( + content_type="application/x-routed", + routes={ + "POST /files": JsonResponse( + content_type="application/json", + body={ + "id": "file-in-$REQUEST_ID", + "object": "file", + "purpose": "batch", + "bytes": 100, + "created_at": 1, + "filename": "in.jsonl", + "status": "processed", + }, + ), + "POST /batches": JsonResponse( + content_type="application/json", body=_batch("validating", files_ready=False) + ), + "GET /batches/batch-$REQUEST_ID": JsonResponse( + content_type="application/json", body=_batch("completed", files_ready=True) + ), + "GET /files/file-in-$REQUEST_ID/content": TextResponse( + content_type="application/jsonl", body="\n".join(input_lines) + "\n" + ), + "GET /files/file-out-$REQUEST_ID/content": TextResponse( + content_type="application/jsonl", body="\n".join(output_lines) + "\n" + ), + "GET /files/file-err-$REQUEST_ID/content": TextResponse( + content_type="application/jsonl", body="\n".join(error_lines) + "\n" + ), + }, + ) + + +def _line_item_config(base: Path, destination: Path) -> Path: + config: Final = yaml.safe_load(base.read_text()) + config["litellm_settings"].update({"callbacks": ["generic_api"], "DEFAULT_FLUSH_INTERVAL_SECONDS": 1}) + config["general_settings"].update({"store_batch_line_items_in_callbacks": True}) + destination.write_text(yaml.safe_dump(config)) + return destination + + +def _spend_rows(key: str) -> tuple[dict[str, JsonValue], ...]: + return tuple( + read_rows( + 'SELECT call_type, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE api_key=%s', + (sha256(key.encode()).hexdigest(),), + ) + ) + + +def _hidden(event: dict[str, JsonValue]) -> dict[str, JsonValue]: + return object_value(event["hidden_params"]) + + +@pytest.mark.covers( + "spend.batches.completed_batch_emits_one_callback_event_per_jsonl_line_beside_the_aggregate", + "spend.batches.line_item_callback_events_do_not_bill_spend_twice", +) +def test_completed_batch_emits_paired_request_response_callback_events_per_jsonl_line( + gateway: Gateway, tmp_path: Path +) -> None: + marker: Final = "batch-line-items-" + uuid.uuid4().hex[:12] + sink_secret: Final = "synthetic-sink-secret-" + marker + provider_secret: Final = "synthetic-provider-secret-" + marker + + def sink(request: Request) -> Reply: + assert request.headers["authorization"] == f"Bearer {sink_secret}" + return Reply() + + with ( + wire_server(sink) as endpoint, + owned_proxy( + gateway, + tmp_path, + {"GENERIC_LOGGER_ENDPOINT": endpoint.url, "GENERIC_LOGGER_HEADERS": f"Authorization=Bearer {sink_secret}"}, + config=_line_item_config(Path("tests/integration/proxy_config.yaml"), tmp_path / "line_items.yaml"), + ) as candidate, + candidate.scenario() as scenario, + ): + routed_model: Final = scenario.model(api_base=f"{gateway.upstream_url}/{marker}", api_key=provider_secret) + input_lines: Final = _input_lines(routed_model, marker) + handle: Final = register_scenario(marker, _provider_routes(input_lines)) + scenario.cleanups.callback(delete_scenario, handle) + key: Final = scenario.key(models=[routed_model]) + file_response: Final = candidate.request_multipart( + "/v1/files", + {"purpose": "batch", "model": routed_model}, + {"file": ("in.jsonl", ("\n".join(input_lines) + "\n").encode(), "application/jsonl")}, + key=key, + ) + assert file_response.status_code == 200, file_response.text + input_file_id: Final = string_value(JSON_OBJECT.validate_json(file_response.content)["id"]) + batch_response: Final = candidate.request( + "POST", + "/v1/batches", + { + "input_file_id": input_file_id, + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "model": routed_model, + }, + key=key, + ) + assert batch_response.status_code == 200, batch_response.text + batch_id: Final = string_value(JSON_OBJECT.validate_json(batch_response.content)["id"]) + retrieval: Final = eventually( + lambda: candidate.request("GET", f"/v1/batches/{batch_id}", key=key), + lambda response: response.status_code == 200 and response.json()["status"] == "completed", + seconds=30, + ) + assert retrieval.status_code == 200, retrieval.text + key_hash: Final = sha256(key.encode()).hexdigest() + batches: Final[ + list[Request] + ] = [] # mutable-ok: drain() consumes the queue, later polls must keep earlier batches + + def delivered() -> tuple[dict[str, JsonValue], ...]: + batches.extend(endpoint.drain()) + return tuple( + event + for batch in batches + for event in json.loads(batch.body) + if object_value(event["metadata"]).get("user_api_key_hash") == key_hash + ) + + events: Final = eventually( + delivered, + lambda values: ( + 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, + ) + for batch in batches: + for credential in (provider_secret, sink_secret, key): + assert credential.encode() not in batch.body + aggregate_events: Final = tuple(event for event in events if event["call_type"] == "aretrieve_batch") + line_events: Final = tuple(event for event in events if _hidden(event).get("batch_custom_id") is not None) + assert len(aggregate_events) == 1, [event["call_type"] for event in events] + other_events: Final = tuple(event for event in events if event not in aggregate_events + line_events) + assert sorted(event["call_type"] for event in other_events) == ["acreate_batch", "acreate_file"], other_events + assert sorted(string_value(_hidden(event)["batch_custom_id"]) for event in line_events) == list( + ALL_CUSTOM_IDS + ), line_events + by_custom_id: Final = {string_value(_hidden(event)["batch_custom_id"]): event for event in line_events} + for custom_id, line in zip(ALL_CUSTOM_IDS, input_lines, strict=True): + event: Final = by_custom_id[custom_id] + hidden: Final = _hidden(event) + assert hidden["batch_id"] == batch_id, hidden + assert event["call_type"] == "acompletion", event["call_type"] + assert event["messages"] == json.loads(line)["body"]["messages"], (custom_id, event["messages"]) + if custom_id in OUTPUT_SUCCESS_IDS: + assert event["status"] == "success", event + assert hidden["batch_line_status_code"] == 200, hidden + assert event["prompt_tokens"] == PROMPT_TOKENS and event["completion_tokens"] == COMPLETION_TOKENS + assert object_value(event["response"])["id"] == f"chatcmpl-{marker}-{custom_id}", event["response"] + assert object_value(object_value(event["response"])["choices"][0]["message"])["content"] == ( + f"answer {custom_id}" + ) + else: + assert event["status"] == "failure", event + assert hidden["batch_line_status_code"] == 400, hidden + assert f"synthetic line failure {custom_id}" in json.dumps(event["error_information"]), event + aggregate: Final = aggregate_events[0] + assert aggregate["prompt_tokens"] == len(OUTPUT_SUCCESS_IDS) * PROMPT_TOKENS, aggregate + assert aggregate["completion_tokens"] == len(OUTPUT_SUCCESS_IDS) * COMPLETION_TOKENS, aggregate + rows: Final = eventually( + lambda: _spend_rows(key), + lambda values: any(r["call_type"] == "aretrieve_batch" for r in values), + seconds=70, + ) + batch_rows: Final = tuple(row for row in rows if row["call_type"] == "aretrieve_batch") + 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/integrations/clickhouse/test_clickhouse_spend_logger.py b/tests/test_litellm/integrations/clickhouse/test_clickhouse_spend_logger.py index 491f8651b52..37716f067e8 100644 --- a/tests/test_litellm/integrations/clickhouse/test_clickhouse_spend_logger.py +++ b/tests/test_litellm/integrations/clickhouse/test_clickhouse_spend_logger.py @@ -523,6 +523,92 @@ async def test_trace_ingest_and_invalid_payload_do_not_write_spend(): storage.ensure_schema.assert_not_awaited() +@pytest.mark.asyncio +async def test_batch_line_item_success_event_does_not_write_spend(monkeypatch: pytest.MonkeyPatch): + # A batch line item carries call_type=acompletion + litellm_params.batch_parent_id. + # The aggregate aretrieve_batch row already bills the batch, so a per-line spend row + # would bill it twice. + monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", True, raising=False) + storage = MagicMock() + logger = ClickHouseSpendLogger(storage=storage) + now = datetime.now(timezone.utc) + + await logger.async_log_success_event( + { + "standard_logging_object": _minimal_payload("chatcmpl-line-1", status="success", cost=0.25), + "call_type": "acompletion", + "litellm_params": {"batch_parent_id": "batch-1"}, + }, + None, + now, + now, + ) + + assert logger.log_queue == [] + + await logger.async_log_success_event( + { + "standard_logging_object": _minimal_payload("chatcmpl-live-1", status="success", cost=0.25), + "call_type": "acompletion", + "litellm_params": {}, + }, + None, + now, + now, + ) + + # The aggregate aretrieve_batch event bills the batch; it must still write one + # row with the batch's full cost even while line items are skipped. + await logger.async_log_success_event( + { + "standard_logging_object": { + **_minimal_payload("batch-1", status="success", cost=0.0001032), + "call_type": "aretrieve_batch", + }, + "call_type": "aretrieve_batch", + "response_cost": 0.0001032, + "litellm_params": {}, + }, + None, + now, + now, + ) + + assert len(logger.log_queue) == 2 + assert logger.log_queue[0]["request_id"] == "chatcmpl-live-1" + assert logger.log_queue[1]["call_type"] == "aretrieve_batch" + assert logger.log_queue[1]["request_id"] == "batch-1" + assert logger.log_queue[1]["spend"] == 0.0001032 + + # A FAILED batch line is still part of the aggregate row's request counts and + # must not add its own spend row; a genuine live failure still writes one. + await logger.async_log_failure_event( + { + "standard_logging_object": _minimal_payload("chatcmpl-line-2", status="failure", cost=0.0), + "call_type": "acompletion", + "litellm_params": {"batch_parent_id": "batch-1"}, + }, + None, + now, + now, + ) + await logger.async_log_failure_event( + { + "standard_logging_object": _minimal_payload("chatcmpl-live-2", status="failure", cost=0.0), + "call_type": "acompletion", + "litellm_params": {}, + }, + None, + now, + now, + ) + + assert len(logger.log_queue) == 3 + assert logger.log_queue[-1]["request_id"] == "chatcmpl-live-2" + assert logger.log_queue[-1]["status"] == "failure" + if logger._flush_task is not None: + logger._flush_task.cancel() + @pytest.mark.parametrize( "status,llm_cost,guardrail_cost,expected", [ diff --git a/tests/unit/batches/test_batch_line_item_logging.py b/tests/unit/batches/test_batch_line_item_logging.py new file mode 100644 index 00000000000..e338cad5d21 --- /dev/null +++ b/tests/unit/batches/test_batch_line_item_logging.py @@ -0,0 +1,1375 @@ +""" +Tests for litellm/batches/batch_line_item_logging.py and its hook in +Logging._async_success_handler_body. + +When ``litellm.store_batch_line_items_in_callbacks`` is on and a completed +batch is logged (call_type aretrieve_batch), litellm emits one callback event +per JSONL line (request paired with response/error) in addition to the +aggregate batch event. These tests run the real Logging.async_success_handler +the way the CheckBatchCost poller invokes it, with a recording CustomLogger on +the async success/failure lists, so a regression in pairing, hidden params, +cost, or error propagation fails here. +""" + +import json +import uuid +from datetime import datetime +from types import MappingProxyType, SimpleNamespace +from typing import Final +from unittest.mock import AsyncMock, patch + +import pytest + +import litellm +from litellm.batches.batch_line_item_logging import ( + _release_line_item_claim, + batch_line_item_claim_cache, + log_batch_line_items, +) +from litellm.batches.batch_utils import BatchResultFiles +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 + +INPUT_JSONL = b"\n".join( + [ + json.dumps( + { + "custom_id": "a", + "method": "POST", + "url": "/v1/chat/completions", + "body": { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi a"}], + "temperature": 0.2, + }, + } + ).encode(), + json.dumps( + { + "custom_id": "b", + "method": "POST", + "url": "/v1/chat/completions", + "body": { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi b"}], + }, + } + ).encode(), + ] +) + +OUTPUT_JSONL = json.dumps( + { + "custom_id": "a", + "response": { + "status_code": 200, + "body": { + "id": "chatcmpl-1", + "model": "gpt-4o", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "hello back"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + }, + }, + } +).encode() + +ERROR_JSONL = json.dumps( + { + "custom_id": "b", + "response": {"status_code": 400, "body": {"error": {"message": "bad request boom"}}}, + "error": {"message": "bad request boom"}, + } +).encode() + +_FILE_BYTES = { + "input-file-1": INPUT_JSONL, + "output-file-1": OUTPUT_JSONL, + "error-file-1": ERROR_JSONL, +} + + +def _batch() -> LiteLLMBatch: + return LiteLLMBatch( + id="batch_1", + object="batch", + endpoint="/v1/chat/completions", + input_file_id="input-file-1", + output_file_id="output-file-1", + error_file_id="error-file-1", + status="completed", + completion_window="24h", + created_at=1, + ) + + +def _file_content(file_id: str, **_kwargs): + return SimpleNamespace(content=_FILE_BYTES[file_id]) + + +class _RecordingLogger(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.success_events = [] + self.failure_events = [] + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + self.success_events.append(kwargs) + + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): + 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() + saved_flag = ( + litellm.store_batch_line_items_in_callbacks + ) # test-quality-ok: process-wide opt-in flag; restored in teardown + saved_success = list( + litellm._async_success_callback + ) # test-quality-ok: the feature dispatches through this global list; restored in teardown + saved_failure = list(litellm._async_failure_callback) # test-quality-ok: same dispatch seam, restored in teardown + litellm._async_success_callback = [ + logger + ] # test-quality-ok: there is no injection seam for callback lists; teardown restores + litellm._async_failure_callback = [logger] # test-quality-ok: same dispatch seam, restored in teardown + yield logger + litellm.store_batch_line_items_in_callbacks = saved_flag # test-quality-ok: teardown restoring the value set above + litellm._async_success_callback = saved_success # test-quality-ok: teardown restoring the value set above + litellm._async_failure_callback = saved_failure # test-quality-ok: teardown restoring the value set above + + +def _parent_logging(custom_llm_provider: str = "openai") -> Logging: + logging_obj = Logging( + model="gpt-4o", + messages=[{"role": "user", "content": ""}], + stream=False, + call_type="aretrieve_batch", + start_time=datetime.now(), + litellm_call_id=str(uuid.uuid4()), + function_id=str(uuid.uuid4()), + ) + logging_obj.update_environment_variables( + litellm_params={"metadata": {"model_info": {"id": "dep-1"}, "model_group": "gpt-4o"}}, + optional_params={}, + custom_llm_provider=custom_llm_provider, + ) + return logging_obj + + +async def _log_completed_batch(logging_obj: Logging, batch: LiteLLMBatch) -> None: + await logging_obj.async_success_handler( + result=batch, + batch_cost=1.5, + batch_usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), + batch_models=["gpt-4o"], + batch_successful_requests=1, + batch_failed_requests=1, + batch_prompt_cost=1.0, + batch_completion_cost=0.5, + ) + + +def _payload(event: dict) -> dict: + return event["standard_logging_object"] + + +def _hidden(event: dict) -> dict: + return _payload(event)["hidden_params"] + + +@pytest.mark.asyncio +async def test_line_items_emitted_alongside_aggregate(recorder): + litellm.store_batch_line_items_in_callbacks = ( + True # test-quality-ok: the flag under test is a module global; fixture restores it + ) + batch = _batch() + with ( + patch( + "litellm.files.main.afile_content", new_callable=AsyncMock, side_effect=_file_content + ), # test-quality-ok: afile_content is the provider boundary; there is no HTTP transport or injection seam for managed file fetch + patch( + "litellm.cost_calculator.batch_cost_calculator", return_value=(0.01, 0.02) + ), # test-quality-ok: the pricing table boundary, same seam existing batch_utils tests patch + ): + await _log_completed_batch(_parent_logging(), batch) + + assert len(recorder.success_events) == 2 + assert len(recorder.failure_events) == 1 + + aggregate = next(e for e in recorder.success_events if _hidden(e).get("batch_custom_id") is None) + assert _payload(aggregate)["response_cost"] == 1.5 + assert _hidden(aggregate)["batch_custom_id"] is None + + line = next(e for e in recorder.success_events if _hidden(e).get("batch_custom_id") == "a") + hidden = _hidden(line) + assert hidden["batch_id"] == batch.id + assert hidden["batch_line_status_code"] == 200 + assert line["litellm_params"]["batch_parent_id"] == batch.id + + payload = _payload(line) + assert payload["response_cost"] == pytest.approx(0.03) + assert payload["prompt_tokens"] == 10 + assert payload["completion_tokens"] == 5 + assert payload["model_parameters"]["temperature"] == 0.2 + assert any(m.get("content") == "hi a" for m in payload["messages"]) + assert payload["response"]["choices"][0]["message"]["content"] == "hello back" + + failure = recorder.failure_events[0] + assert _hidden(failure)["batch_custom_id"] == "b" + assert _hidden(failure)["batch_id"] == batch.id + assert _hidden(failure)["batch_line_status_code"] == 400 + assert "bad request boom" in _payload(failure)["error_str"] + + +@pytest.mark.asyncio +async def test_flag_off_emits_only_aggregate(recorder): + assert litellm.store_batch_line_items_in_callbacks is False + file_mock = 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 + await _log_completed_batch(_parent_logging(), _batch()) + + assert len(recorder.success_events) == 1 + assert len(recorder.failure_events) == 0 + file_mock.assert_not_called() + + +@pytest.mark.asyncio +async def test_line_items_skipped_for_unsupported_provider(recorder): + litellm.store_batch_line_items_in_callbacks = ( + True # test-quality-ok: the flag under test is a module global; fixture restores it + ) + 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 + await _log_completed_batch(_parent_logging(custom_llm_provider="cohere"), _batch()) + + file_mock.assert_not_awaited() + assert len(recorder.success_events) == 1 + assert _payload(recorder.success_events[0])["response_cost"] == 1.5 + assert _hidden(recorder.success_events[0])["batch_custom_id"] is None + assert len(recorder.failure_events) == 0 + + +@pytest.mark.asyncio +async def test_in_progress_batch_poll_emits_no_line_events(recorder): + litellm.store_batch_line_items_in_callbacks = ( + True # test-quality-ok: the flag under test is a module global; fixture restores it + ) + in_progress: Final = LiteLLMBatch( + id="batch_wip", + object="batch", + endpoint="/v1/chat/completions", + input_file_id="input-file-1", + output_file_id=None, + error_file_id=None, + status="in_progress", + completion_window="24h", + created_at=1, + ) + 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 + await _parent_logging().async_success_handler(result=in_progress) + + file_mock.assert_not_called() + assert all(_hidden(e).get("batch_custom_id") is None for e in recorder.success_events) + assert len(recorder.failure_events) == 0 + + +@pytest.mark.asyncio +async def test_input_fetch_failure_still_emits_aggregate(recorder): + litellm.store_batch_line_items_in_callbacks = ( + True # test-quality-ok: the flag under test is a module global; fixture restores it + ) + 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 + await _log_completed_batch(_parent_logging(), _batch()) + + assert len(recorder.success_events) == 1 + assert len(recorder.failure_events) == 0 + assert _payload(recorder.success_events[0])["response_cost"] == 1.5 + + +EDGE_INPUT_JSONL = b"\n".join( + [ + json.dumps( + { + "custom_id": "e", + "method": "POST", + "url": "/v1/embeddings", + "body": {"model": "text-embedding-3-small", "input": "embed me", "encoding_format": "float"}, + } + ).encode(), + json.dumps( + { + "custom_id": "r", + "method": "POST", + "url": "/v1/responses", + "body": {"model": "gpt-4o", "input": "respond to me"}, + } + ).encode(), + json.dumps( + { + "custom_id": "badresp", + "method": "POST", + "url": "/v1/responses", + "body": {"model": "gpt-4o", "input": "unreconstructable"}, + } + ).encode(), + json.dumps( + { + "custom_id": "n", + "method": "POST", + "url": "/v1/chat/completions", + "body": {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi n"}]}, + } + ).encode(), + ] +) + +EDGE_OUTPUT_JSONL = b"\n".join( + [ + json.dumps( + { + "custom_id": "e", + "response": { + "status_code": 200, + "body": { + "object": "list", + "data": [{"object": "embedding", "embedding": [0.1], "index": 0}], + "model": "text-embedding-3-small", + "usage": {"prompt_tokens": 3, "total_tokens": 3}, + }, + }, + } + ).encode(), + json.dumps( + { + "custom_id": "r", + "response": { + "status_code": 200, + "body": { + "id": "resp_1", + "object": "response", + "created_at": 1, + "status": "completed", + "output": [], + "model": "gpt-4o", + }, + }, + } + ).encode(), + json.dumps({"custom_id": "badresp", "response": {"status_code": 200, "body": {}}}).encode(), + json.dumps({"custom_id": "n", "response": {"body": {"error": {"message": "no status here"}}}}).encode(), + json.dumps({"custom_id": "boom", "response": "not-a-dict"}).encode(), + json.dumps( + { + "custom_id": "mi", + "modelInput": {"messages": [{"role": "user", "content": "hi mi"}]}, + "response": { + "status_code": 200, + "body": { + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "mi back"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 2, "completion_tokens": 1, "total_tokens": 3}, + }, + }, + } + ).encode(), + ] +) + +ANTHROPIC_INPUT_JSONL = json.dumps( + { + "custom_id": "b2", + "params": {"model": "claude-3", "max_tokens": 5, "messages": [{"role": "user", "content": "hi b2"}]}, + } +).encode() + +ANTHROPIC_OUTPUT_JSONL = json.dumps( + { + "custom_id": "b2", + "result": { + "type": "succeeded", + "message": { + "id": "msg_1", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "hello b2"}], + "model": "claude-3", + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 1, "output_tokens": 2}, + }, + }, + } +).encode() + + +BEDROCK_INPUT_JSONL = b"\n".join( + [ + json.dumps( + { + "recordId": "br-anth", + "modelInput": {"messages": [{"role": "user", "content": "hi br-anth"}]}, + } + ).encode(), + json.dumps( + { + "recordId": "br-titan", + "modelInput": {"inputText": "embed me"}, + } + ).encode(), + json.dumps( + { + "recordId": "br-conv", + "modelInput": {"messages": [{"role": "user", "content": "hi br-conv"}]}, + } + ).encode(), + ] +) + +BEDROCK_OUTPUT_JSONL = b"\n".join( + [ + json.dumps( + { + "recordId": "br-anth", + "modelOutput": { + "id": "msg_br", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "bedrock claude hi"}], + "model": "anthropic.claude-3", + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 2, "output_tokens": 3}, + }, + } + ).encode(), + json.dumps( + { + "recordId": "br-titan", + "modelOutput": {"embedding": [0.1, 0.2], "inputTextTokenCount": 4}, + } + ).encode(), + json.dumps( + { + "recordId": "br-conv", + "modelOutput": { + "output": {"message": {"role": "assistant", "content": [{"text": "converse hi"}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 5, "outputTokens": 6, "totalTokens": 11}, + "metrics": {"latencyMs": 12}, + }, + } + ).encode(), + ] +) + +MISTRAL_INPUT_JSONL = json.dumps( + { + "custom_id": "mis", + "body": {"model": "mistral-small", "messages": [{"role": "user", "content": "hi mis"}]}, + } +).encode() + +MISTRAL_OUTPUT_JSONL = json.dumps( + { + "custom_id": "mis", + "response": { + "status_code": 200, + "body": { + "id": "mis-1", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "mistral hi"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5}, + }, + }, + } +).encode() + + +def _edge_file_content(file_id: str, **_kwargs): + return SimpleNamespace( + content={ + "input-2": EDGE_INPUT_JSONL, + "output-2": EDGE_OUTPUT_JSONL, + "input-anth": ANTHROPIC_INPUT_JSONL, + "output-anth": ANTHROPIC_OUTPUT_JSONL, + "input-bed": BEDROCK_INPUT_JSONL, + "output-bed": BEDROCK_OUTPUT_JSONL, + "input-mist": MISTRAL_INPUT_JSONL, + "output-mist": MISTRAL_OUTPUT_JSONL, + }[file_id] + ) + + +def _parent_logging_with_params(litellm_params: dict) -> Logging: + logging_obj = _parent_logging() + logging_obj.update_environment_variables( + litellm_params=litellm_params, + optional_params={}, + custom_llm_provider="openai", + ) + return logging_obj + + +@pytest.mark.asyncio +async def test_line_items_edge_shapes_and_edge_cases(recorder): + litellm.store_batch_line_items_in_callbacks = ( + True # test-quality-ok: the flag under test is a module global; fixture restores it + ) + batch = LiteLLMBatch( + id="batch_edge", + object="batch", + endpoint="/v1/chat/completions", + input_file_id="input-2", + output_file_id="output-2", + error_file_id=None, + status="completed", + completion_window="24h", + created_at=1, + ) + file_mock: Final = AsyncMock(side_effect=_edge_file_content) + parent: Final = _parent_logging_with_params( + { + "metadata": {"model_info": {"id": "dep-1"}, "model_group": "gpt-4o"}, + } + ) + parent._litellm_internal_model_credentials = { + "api_key": "sk-line-items-marker" + } # test-quality-ok: private transport attribute, same channel the batch cost tracker uses + 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 + patch( + "litellm.cost_calculator.batch_cost_calculator", return_value=(0.01, 0.02) + ), # test-quality-ok: the pricing table boundary, same seam existing batch_utils tests patch + ): + await _log_completed_batch(parent, batch) + + by_custom_id = {_hidden(e).get("batch_custom_id"): e for e in recorder.success_events} + aggregate = next(e for e in recorder.success_events if _hidden(e).get("batch_custom_id") is None) + assert _payload(aggregate)["response_cost"] == 1.5 + + assert by_custom_id["e"]["litellm_params"]["batch_parent_id"] == batch.id + assert by_custom_id["e"]["call_type"] == "aembedding" + assert _payload(by_custom_id["e"])["response"]["data"][0]["embedding"] == [0.1] + + assert by_custom_id["r"]["call_type"] == "aresponses" + assert _payload(by_custom_id["r"])["response"]["id"] == "resp_1" + + assert by_custom_id["mi"]["call_type"] == "acompletion" + assert _hidden(by_custom_id["mi"])["batch_line_status_code"] == 200 + + assert "badresp" not in by_custom_id + assert "boom" not in by_custom_id + + failure = recorder.failure_events[0] + assert _hidden(failure)["batch_custom_id"] == "n" + assert _hidden(failure)["batch_line_status_code"] is None + assert "no status here" in _payload(failure)["error_str"] + + assert "sk-line-items-marker" in str(file_mock.call_args_list) + assert "sk-line-items-marker" not in str(by_custom_id["e"]["litellm_params"]) + file_ids_fetched = [call.kwargs.get("file_id") or call.args[0] for call in file_mock.call_args_list] + assert "input-2" in file_ids_fetched and "output-2" in file_ids_fetched + assert not any(file_id is None for file_id in file_ids_fetched) + + +@pytest.mark.asyncio +async def test_line_items_child_params_drop_parent_credentials(recorder): + litellm.store_batch_line_items_in_callbacks = ( + True # test-quality-ok: the flag under test is a module global; fixture restores it + ) + file_mock: Final = AsyncMock(side_effect=_file_content) + parent: Final = _parent_logging_with_params( + { + "api_key": "sk-parent-secret", + "_litellm_internal_model_credentials": MappingProxyType({"api_key": "sk-parent-secret"}), + "api_base": "https://api.openai.com/v1", + "metadata": {"model_info": {"id": "dep-1"}, "model_group": "gpt-4o"}, + } + ) + 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 + patch( + "litellm.cost_calculator.batch_cost_calculator", return_value=(0.01, 0.02) + ), # test-quality-ok: the pricing table boundary, same seam existing batch_utils tests patch + ): + await _log_completed_batch(parent, batch) + + line_events = [ + e for e in [*recorder.success_events, *recorder.failure_events] if _hidden(e).get("batch_custom_id") is not None + ] + assert len(line_events) == 2 + for event in line_events: + params = event["litellm_params"] + assert "api_key" not in params + assert "_litellm_internal_model_credentials" not in params + assert params["batch_parent_id"] == batch.id + assert params["metadata"]["model_group"] == "gpt-4o" + + +@pytest.mark.asyncio +async def test_line_items_anthropic_shapes(recorder): + litellm.store_batch_line_items_in_callbacks = ( + True # test-quality-ok: the flag under test is a module global; fixture restores it + ) + batch = LiteLLMBatch( + id="batch_anth", + object="batch", + endpoint="/v1/messages", + input_file_id="input-anth", + output_file_id="output-anth", + error_file_id=None, + status="completed", + completion_window="24h", + created_at=1, + ) + file_mock: Final = AsyncMock(side_effect=_edge_file_content) + logging_obj = _parent_logging() + logging_obj.update_environment_variables( + litellm_params={"metadata": {"model_info": {"id": "dep-1"}, "model_group": "claude-3"}}, + optional_params={}, + custom_llm_provider="anthropic", + ) + 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 + patch( + "litellm.cost_calculator.batch_cost_calculator", return_value=(0.01, 0.02) + ), # test-quality-ok: the pricing table boundary, same seam existing batch_utils tests patch + ): + await _log_completed_batch(logging_obj, batch) + + line = next(e for e in recorder.success_events if _hidden(e).get("batch_custom_id") == "b2") + assert _hidden(line)["batch_line_status_code"] == 200 + assert line["litellm_params"]["batch_parent_id"] == batch.id + payload = _payload(line) + assert payload["response"]["choices"][0]["message"]["content"] == "hello b2" + assert payload["prompt_tokens"] == 1 + assert payload["completion_tokens"] == 2 + + +def _provider_batch(batch_id: str, input_file_id: str, output_file_id: str) -> LiteLLMBatch: + return LiteLLMBatch( + id=batch_id, + object="batch", + endpoint="/v1/chat/completions", + input_file_id=input_file_id, + output_file_id=output_file_id, + error_file_id=None, + status="completed", + completion_window="24h", + created_at=1, + ) + + +@pytest.mark.asyncio +async def test_line_items_bedrock_shapes(recorder): + litellm.store_batch_line_items_in_callbacks = ( + True # test-quality-ok: the flag under test is a module global; fixture restores it + ) + file_mock: Final = AsyncMock(side_effect=_edge_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 + patch( + "litellm.cost_calculator.batch_cost_calculator", return_value=(0.01, 0.02) + ), # test-quality-ok: the pricing table boundary, same seam existing batch_utils tests patch + ): + await _log_completed_batch( + _parent_logging(custom_llm_provider="bedrock"), + _provider_batch("batch_bed", "input-bed", "output-bed"), + ) + + by_id = {_hidden(e).get("batch_custom_id"): e for e in recorder.success_events} + + anthropic_event = by_id["br-anth"] + assert anthropic_event["call_type"] == "acompletion" + assert _hidden(anthropic_event)["batch_custom_id"] == "br-anth" + anthropic_line = _payload(anthropic_event) + assert any(m.get("content") == "hi br-anth" for m in anthropic_line["messages"]) + assert anthropic_line["response"]["choices"][0]["message"]["content"] == "bedrock claude hi" + assert anthropic_line["prompt_tokens"] == 2 + assert anthropic_line["completion_tokens"] == 3 + + titan_event = by_id["br-titan"] + assert titan_event["call_type"] == "aembedding" + assert _hidden(titan_event)["batch_custom_id"] == "br-titan" + titan_line = _payload(titan_event) + assert titan_event["optional_params"]["inputText"] == "embed me" + assert titan_line["response"]["data"][0]["embedding"] == [0.1, 0.2] + assert titan_line["prompt_tokens"] == 4 + + converse_event = by_id["br-conv"] + assert converse_event["call_type"] == "acompletion" + converse_line = _payload(converse_event) + assert any(m.get("content") == "hi br-conv" for m in converse_line["messages"]) + assert converse_line["response"]["choices"][0]["message"]["content"] == "converse hi" + assert converse_line["prompt_tokens"] == 5 + assert converse_line["completion_tokens"] == 6 + + +@pytest.mark.asyncio +async def test_line_items_mistral_shape(recorder): + litellm.store_batch_line_items_in_callbacks = ( + True # test-quality-ok: the flag under test is a module global; fixture restores it + ) + file_mock: Final = AsyncMock(side_effect=_edge_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 + patch( + "litellm.cost_calculator.batch_cost_calculator", return_value=(0.01, 0.02) + ), # test-quality-ok: the pricing table boundary, same seam existing batch_utils tests patch + ): + await _log_completed_batch( + _parent_logging(custom_llm_provider="mistral"), + _provider_batch("batch_mist", "input-mist", "output-mist"), + ) + + line = next(e for e in recorder.success_events if _hidden(e).get("batch_custom_id") == "mis") + assert _hidden(line)["batch_line_status_code"] == 200 + assert _payload(line)["response"]["choices"][0]["message"]["content"] == "mistral hi" + + +ANTHROPIC_ERR_INPUT_JSONL = json.dumps( + { + "custom_id": "aerr", + "params": {"model": "claude-3", "max_tokens": 5, "messages": [{"role": "user", "content": "hi aerr"}]}, + } +).encode() + +ANTHROPIC_ERR_OUTPUT_JSONL = json.dumps( + { + "custom_id": "aerr", + "result": { + "type": "errored", + "error": { + "type": "error", + "error": {"type": "invalid_request_error", "message": "max_tokens must be positive"}, + }, + }, + } +).encode() + +VERTEX_NATIVE_OUTPUT_JSONL = b"\n".join( + [ + json.dumps( + { + "request": {"contents": [{"role": "user", "parts": [{"text": "hi v1"}]}]}, + "response": { + "candidates": [{"content": {"role": "model", "parts": [{"text": "hi back"}]}}], + "usageMetadata": {"promptTokenCount": 3, "candidatesTokenCount": 2, "totalTokenCount": 5}, + }, + } + ).encode(), + json.dumps( + { + "request": {"contents": [{"role": "user", "parts": [{"text": "hi v2"}]}]}, + "status": "Internal error", + } + ).encode(), + ] +) + + +def _scoped_file_content(file_map): + return lambda file_id, **kwargs: SimpleNamespace(content=file_map[file_id]) + + +@pytest.mark.asyncio +async def test_line_items_anthropic_failure_keeps_provider_error(recorder): + litellm.store_batch_line_items_in_callbacks = ( + True # test-quality-ok: the flag under test is a module global; fixture restores it + ) + file_mock: Final = AsyncMock( + side_effect=_scoped_file_content( + { + "input-anth-err": ANTHROPIC_ERR_INPUT_JSONL, + "output-anth-err": ANTHROPIC_ERR_OUTPUT_JSONL, + } + ) + ) + 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 + patch( + "litellm.cost_calculator.batch_cost_calculator", return_value=(0.01, 0.02) + ), # test-quality-ok: the pricing table boundary, same seam existing batch_utils tests patch + ): + await _log_completed_batch( + _parent_logging(custom_llm_provider="anthropic"), + _provider_batch("batch_anth_err", "input-anth-err", "output-anth-err"), + ) + + assert len(recorder.failure_events) == 1 + failure = recorder.failure_events[0] + assert _hidden(failure)["batch_custom_id"] == "aerr" + assert _hidden(failure)["batch_line_status_code"] is None + assert "max_tokens must be positive" in _payload(failure)["error_str"] + + +BODYLESS_ERR_INPUT_JSONL = json.dumps( + {"custom_id": "bare", "url": "/v1/chat/completions", "body": {"model": "gpt-4.1-nano", "messages": []}} +).encode() + +BODYLESS_ERR_OUTPUT_JSONL = json.dumps({"custom_id": "bare", "response": None, "error": None}).encode() + + +@pytest.mark.asyncio +async def test_line_items_failure_without_error_or_response_body_still_reaches_callbacks(recorder): + litellm.store_batch_line_items_in_callbacks = ( + True # test-quality-ok: the flag under test is a module global; fixture restores it + ) + file_mock: Final = AsyncMock( + side_effect=_scoped_file_content( + { + "input-bare": BODYLESS_ERR_INPUT_JSONL, + "output-bare": BODYLESS_ERR_OUTPUT_JSONL, + } + ) + ) + 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 + patch( + "litellm.cost_calculator.batch_cost_calculator", return_value=(0.01, 0.02) + ), # test-quality-ok: the pricing table boundary, same seam existing batch_utils tests patch + ): + await _log_completed_batch(_parent_logging(), _provider_batch("batch_bare", "input-bare", "output-bare")) + + assert len(recorder.failure_events) == 1 + failure = recorder.failure_events[0] + assert _hidden(failure)["batch_custom_id"] == "bare" + assert _payload(failure)["error_str"] == "{}" + + +@pytest.mark.asyncio +async def test_line_items_native_vertex_rows_are_skipped(recorder): + litellm.store_batch_line_items_in_callbacks = ( + True # test-quality-ok: the flag under test is a module global; fixture restores it + ) + file_mock: Final = AsyncMock( + side_effect=_scoped_file_content( + { + "input-vtx": b"", + "output-vtx": VERTEX_NATIVE_OUTPUT_JSONL, + } + ) + ) + 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 + patch( + "litellm.cost_calculator.batch_cost_calculator", return_value=(0.01, 0.02) + ), # test-quality-ok: the pricing table boundary, same seam existing batch_utils tests patch + ): + await _log_completed_batch( + _parent_logging(custom_llm_provider="vertex_ai"), + _provider_batch("batch_vtx", "input-vtx", "output-vtx"), + ) + + 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 + + +_BROKEN_ERROR_FILE: Final = {**_FILE_BYTES, "error-file-1": "not bytes"} + + +def _broken_error_file_content(file_id: str, **_kwargs): + return SimpleNamespace(content=_BROKEN_ERROR_FILE[file_id]) + + +@pytest.mark.asyncio +async def test_line_items_retry_after_a_failed_fanout(recorder): + claim_cache: Final = DualCache() + 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 == 2 + assert len(recorder.success_events) == 1 + assert len(recorder.failure_events) == 1 + + +@pytest.mark.asyncio +async def test_line_items_claim_kept_after_partial_emission(recorder): + claim_cache: Final = DualCache() + batch: Final = _batch() + parent: Final = _parent_logging() + file_mock: Final = AsyncMock(side_effect=_broken_error_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 + 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 == 1 + 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 + + +@pytest.mark.asyncio +async def test_supplied_result_files_skip_the_output_and_error_fetch(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, + result_files=BatchResultFiles(output=OUTPUT_JSONL, error=ERROR_JSONL), + claim_cache=DualCache(), + ) + + assert emitted == 2 + assert len(recorder.success_events) == 1 + assert len(recorder.failure_events) == 1 + assert file_mock.await_count == 1 + assert file_mock.await_args.kwargs["file_id"] == "input-file-1" + + +@pytest.mark.asyncio +async def test_empty_result_files_fall_back_to_fetching_everything(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, + result_files=BatchResultFiles(output=None, error=None), + claim_cache=DualCache(), + ) + + assert emitted == 2 + fetched: Final = sorted(call.kwargs["file_id"] for call in file_mock.await_args_list) + assert fetched == ["error-file-1", "input-file-1", "output-file-1"] + + +@pytest.mark.asyncio +async def test_native_vertex_skip_releases_the_claim(recorder): + file_mock: Final = AsyncMock( + side_effect=_scoped_file_content( + { + "input-vtx-rel": b"", + "output-vtx-rel": VERTEX_NATIVE_OUTPUT_JSONL, + } + ) + ) + claim_cache: Final = DualCache() + batch: Final = _provider_batch("batch_vtx_rel", "input-vtx-rel", "output-vtx-rel") + parent: Final = _parent_logging(custom_llm_provider="vertex_ai") + 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="vertex_ai", + parent=parent, + model_name="vertex_ai/gemini-2.5-flash", + litellm_params=None, + model_info=None, + claim_cache=claim_cache, + ) + second: Final = await log_batch_line_items( + batch=batch, + custom_llm_provider="vertex_ai", + parent=parent, + model_name="vertex_ai/gemini-2.5-flash", + litellm_params=None, + model_info=None, + claim_cache=claim_cache, + ) + + assert first == 0 + assert second == 0 + assert file_mock.await_count == 4, ( + "a released claim must let the next retrieve reach the fetch step again; already_claimed would return early at 2 fetches" + ) + + +@pytest.mark.asyncio +async def test_fanout_emitting_nothing_releases_the_claim(recorder): + file_mock: Final = AsyncMock( + side_effect=_scoped_file_content( + { + "input-empty": b"", + "output-empty": b"", + "error-empty": b"", + } + ) + ) + claim_cache: Final = DualCache() + batch: Final = _provider_batch("batch_empty", "input-empty", "output-empty") + batch = batch.model_copy(update={"error_file_id": "error-empty"}) + 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=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 file_mock.await_count == 6, ( + "a completed fan-out that emits nothing must release the claim so the next retrieve retries" + ) diff --git a/tests/unit/batches/test_batch_utils.py b/tests/unit/batches/test_batch_utils.py index 8147626fc5e..2d77c405f9c 100644 --- a/tests/unit/batches/test_batch_utils.py +++ b/tests/unit/batches/test_batch_utils.py @@ -2485,3 +2485,57 @@ async def test_flag_sends_every_vertex_row_down_the_native_path_when_a_model_is_ assert calls == [] assert (result.successful_requests, result.failed_requests) == (0, 1) + + +@pytest.mark.asyncio +async def test_handle_completed_batch_with_files_returns_bytes_and_fetches_error_once(monkeypatch): + rows = [_success_row(model="gpt-4o", usage=_usage(10, 5))] + output_bytes = _vertex_jsonl(rows) + error_bytes = _vertex_jsonl([{"custom_id": "req-bad", "error": {"message": "rejected"}}]) + + fetched = [] + + async def fake_fetch(batch, custom_llm_provider, litellm_params=None): + return output_bytes + + async def fake_afile_content(**kw): + fetched.append(kw["file_id"]) + return type("R", (), {"content": error_bytes})() + + import litellm.cost_calculator as cc + import litellm.files.main as files_main + + monkeypatch.setattr(bu, "_fetch_batch_output_file_content", fake_fetch) + monkeypatch.setattr(files_main, "afile_content", fake_afile_content) + monkeypatch.setattr(cc, "batch_cost_calculator", lambda **kw: (2.0, 1.3)) + + batch = _batch("of").model_copy(update={"error_file_id": "ef"}) + + result, files = await bu._handle_completed_batch_with_files(batch, custom_llm_provider="openai") + + assert result.cost == 3.3 + assert result.failed_requests == 1 + assert files.output == output_bytes + assert files.error == error_bytes + assert fetched == ["ef"] + + +@pytest.mark.asyncio +async def test_handle_completed_batch_with_files_no_output_returns_error_bytes(monkeypatch): + error_bytes = _vertex_jsonl([{"custom_id": "req-bad", "error": {"message": "rejected"}}]) + + async def fake_afile_content(**kw): + return type("R", (), {"content": error_bytes})() + + import litellm.files.main as files_main + + monkeypatch.setattr(files_main, "afile_content", fake_afile_content) + + batch = _batch(None).model_copy(update={"error_file_id": "ef"}) + + result, files = await bu._handle_completed_batch_with_files(batch, custom_llm_provider="openai") + + assert result.failed_requests == 1 + assert result.usage.total_tokens == 0 + assert files.output is None + assert files.error == error_bytes diff --git a/tests/unit/integrations/test_batch_line_item_metering_sinks.py b/tests/unit/integrations/test_batch_line_item_metering_sinks.py new file mode 100644 index 00000000000..525f6ae15d0 --- /dev/null +++ b/tests/unit/integrations/test_batch_line_item_metering_sinks.py @@ -0,0 +1,419 @@ +""" +Guards that keep batch line-item callback events out of the built-in metering sinks. + +With ``litellm.store_batch_line_items_in_callbacks`` on, a completed batch emits +one child callback per JSONL line on top of the aggregate ``aretrieve_batch`` +event. Every line carries a real ``response_cost``, so any billing/metering sink +without the ``is_batch_line_item_event`` guard meters aggregate + per-line and +reports roughly 2x the true spend. + +These tests put REAL sink instances on the REAL dispatch lists, drive the REAL +``Logging.async_success_handler`` for a completed 2-line batch (1 success line + +1 error line, aggregate cost $1.50, per-line cost $0.03), and assert each sink +meters exactly the aggregate. OTel spans for line items must still be emitted: +the guard belongs on cost/token metrics, not on tracing. +""" + +import io +import json +import os +import time +import uuid +from contextlib import redirect_stdout +from datetime import datetime +from types import SimpleNamespace +from typing import Any, Final +from unittest.mock import AsyncMock, patch + +import pytest + +import litellm +from litellm.batches.batch_line_item_logging import batch_line_item_claim_cache +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.types.utils import LiteLLMBatch, Usage + +AGGREGATE_COST: Final[float] = 1.5 + +INPUT_JSONL: Final[bytes] = b"\n".join( + [ + json.dumps( + { + "custom_id": "a", + "method": "POST", + "url": "/v1/chat/completions", + "body": {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi a"}]}, + } + ).encode(), + json.dumps( + { + "custom_id": "b", + "method": "POST", + "url": "/v1/chat/completions", + "body": {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi b"}]}, + } + ).encode(), + ] +) + +OUTPUT_JSONL: Final[bytes] = json.dumps( + { + "custom_id": "a", + "response": { + "status_code": 200, + "body": { + "id": "chatcmpl-line-1", + "model": "gpt-4o", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + }, + }, + } +).encode() + +ERROR_JSONL: Final[bytes] = json.dumps( + { + "custom_id": "b", + "response": {"status_code": 400, "body": {"error": {"message": "boom"}}}, + "error": {"message": "boom"}, + } +).encode() + +_FILE_BYTES: Final[dict[str, bytes]] = { + "input-file-1": INPUT_JSONL, + "output-file-1": OUTPUT_JSONL, + "error-file-1": ERROR_JSONL, +} + + +def _file_content(file_id: str, **_kwargs: Any) -> SimpleNamespace: + return SimpleNamespace(content=_FILE_BYTES[file_id]) + + +def _batch() -> LiteLLMBatch: + return LiteLLMBatch( + id=f"batch_{uuid.uuid4().hex[:8]}", + object="batch", + endpoint="/v1/chat/completions", + input_file_id="input-file-1", + output_file_id="output-file-1", + error_file_id="error-file-1", + status="completed", + completion_window="24h", + created_at=1, + ) + + +def _parent_logging() -> Logging: + logging_obj = Logging( + model="gpt-4o", + messages=[{"role": "user", "content": ""}], + stream=False, + call_type="aretrieve_batch", + start_time=datetime.now(), + litellm_call_id=str(uuid.uuid4()), + function_id=str(uuid.uuid4()), + ) + logging_obj.update_environment_variables( + litellm_params={ + "metadata": { + "model_info": {"id": "dep-1"}, + "model_group": "gpt-4o", + "user_api_key_user_id": "user-77", + "user_api_key_team_id": "team-7", + "user_api_key_team_alias": "team-seven", + "user_api_key_alias": "key-alias-1", + } + }, + optional_params={}, + custom_llm_provider="openai", + ) + return logging_obj + + +class RecordingHTTP: + """Stands in for a sink's HTTP egress object only; the sink logic is real.""" + + def __init__(self) -> None: + self.posts: list[dict[str, Any]] = [] + + async def post(self, url: str, data: Any = None, content: Any = None, **_kw: Any) -> SimpleNamespace: + self.posts.append({"url": url, "body": data if data is not None else content}) + return SimpleNamespace(status_code=200, text="ok", raise_for_status=lambda: None) + + async def put(self, url: str, content: Any = None, **_kw: Any) -> SimpleNamespace: + return SimpleNamespace(status_code=202, text="ok", raise_for_status=lambda: None) + + +async def _log_completed_batch(monkeypatch: pytest.MonkeyPatch, loggers: list, flag: bool = True) -> None: + batch_line_item_claim_cache.in_memory_cache.flush_cache() + monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", flag, raising=False) + saved_success = list(litellm._async_success_callback) + saved_failure = list(litellm._async_failure_callback) + litellm._async_success_callback = list(loggers) + litellm._async_failure_callback = list(loggers) + buf = io.StringIO() + try: + with ( + patch("litellm.files.main.afile_content", new_callable=AsyncMock, side_effect=_file_content), + patch("litellm.cost_calculator.batch_cost_calculator", return_value=(0.01, 0.02)), + redirect_stdout(buf), + ): + await _parent_logging().async_success_handler( + result=_batch(), + batch_cost=1.5, + batch_usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), + batch_models=["gpt-4o"], + batch_successful_requests=1, + batch_failed_requests=1, + batch_prompt_cost=1.0, + batch_completion_cost=0.5, + ) + finally: + litellm._async_success_callback = saved_success + litellm._async_failure_callback = saved_failure + + +def _openmeter_sinks() -> tuple[Any, RecordingHTTP]: + from litellm.integrations.openmeter import OpenMeterLogger + + recorder = RecordingHTTP() + logger = OpenMeterLogger() + logger.async_http_handler = recorder + return logger, recorder + + +def _lago_sinks() -> tuple[Any, RecordingHTTP]: + from litellm.integrations.lago import LagoLogger + + recorder = RecordingHTTP() + logger = LagoLogger() + logger.async_http_handler = recorder + return logger, recorder + + +@pytest.mark.asyncio +async def test_openmeter_meters_only_the_aggregate_batch_event(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("OPENMETER_API_KEY", "test-openmeter-key") + logger, recorder = _openmeter_sinks() + + await _log_completed_batch(monkeypatch, [logger]) + + costs = [json.loads(p["body"])["data"]["cost"] for p in recorder.posts] + assert costs == [AGGREGATE_COST], ( + f"OpenMeter metered per-line costs on top of the aggregate: {costs}; " + "line items must be billed only by the aggregate aretrieve_batch event" + ) + + +@pytest.mark.asyncio +async def test_lago_bills_only_the_aggregate_batch_event(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LAGO_API_KEY", "test-lago-key") + monkeypatch.setenv("LAGO_API_BASE", "http://lago.invalid") + monkeypatch.setenv("LAGO_API_EVENT_CODE", "litellm-usage") + monkeypatch.setenv("LAGO_API_CHARGE_BY", "user_id") + logger, recorder = _lago_sinks() + + await _log_completed_batch(monkeypatch, [logger]) + + costs = [json.loads(p["body"])["event"]["properties"]["response_cost"] for p in recorder.posts] + assert costs == [AGGREGATE_COST], ( + f"Lago billed per-line costs on top of the aggregate: {costs}; " + "line items must be billed only by the aggregate aretrieve_batch event" + ) + + +@pytest.mark.asyncio +async def test_datadog_cost_management_bills_only_the_aggregate_batch_event( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.integrations.datadog.datadog_cost_management import DatadogCostManagementLogger + + monkeypatch.setenv("DD_API_KEY", "test-dd-key") + monkeypatch.setenv("DD_APP_KEY", "test-dd-app-key") + logger = DatadogCostManagementLogger(cost_tag_keys=[]) + + await _log_completed_batch(monkeypatch, [logger]) + + entries = list(logger.log_queue) + costs = [e.get("response_cost", 0) for e in entries] + assert costs == [AGGREGATE_COST], ( + f"Datadog FOCUS BilledCost queued per-line entries on top of the aggregate: {costs}; " + "line items must be billed only by the aggregate aretrieve_batch event" + ) + + +@pytest.mark.asyncio +async def test_newrelic_metrics_meters_only_the_aggregate_batch_event(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.integrations.newrelic.newrelic_metrics import NewRelicMetricsLogger, build_metric_payload + + logger = NewRelicMetricsLogger(newrelic_api_key="test-nr-key") + + await _log_completed_batch(monkeypatch, [logger]) + + records = tuple(logger.log_queue) + costs = [r.response_cost for r in records] + assert costs == [AGGREGATE_COST], ( + f"New Relic metered per-line costs on top of the aggregate: {costs}; " + "line items must be metered only by the aggregate aretrieve_batch event" + ) + now = time.time() + envelopes = build_metric_payload(records=records, window_start=now - 1, now=now) + cost_sum = 0.0 + for envelope in envelopes: + for metric in envelope["metrics"]: + if "cost" in metric["name"]: + value = metric["value"]["sum"] if isinstance(metric["value"], dict) else metric["value"] + cost_sum += value + assert cost_sum == pytest.approx(AGGREGATE_COST) + + +@pytest.mark.asyncio +async def test_prometheus_meters_only_the_aggregate_batch_event(monkeypatch: pytest.MonkeyPatch) -> None: + prometheus_client = pytest.importorskip("prometheus_client") + from litellm.integrations.prometheus import PrometheusLogger + from prometheus_client import REGISTRY + + saved_collectors = list(REGISTRY._collector_to_names.keys()) + for collector in saved_collectors: + REGISTRY.unregister(collector) + logger = PrometheusLogger() + + try: + await _log_completed_batch(monkeypatch, [logger]) + spend_samples = [] + for metric in REGISTRY.collect(): + if metric.name == "litellm_spend_metric": + spend_samples = [sample.value for sample in metric.samples if sample.name.endswith("_total")] + finally: + for collector in list(REGISTRY._collector_to_names.keys()): + REGISTRY.unregister(collector) + for collector in saved_collectors: + try: + REGISTRY.register(collector) + except Exception: # noqa: BLE001 # already re-registered by another holder + pass + + assert spend_samples, "litellm_spend_metric saw no samples at all" + assert sum(spend_samples) == pytest.approx(AGGREGATE_COST), ( + f"Prometheus spend metric double-counted batch line items: {spend_samples}; " + "line items must be metered only by the aggregate aretrieve_batch event" + ) + + +def _otel_v1(monkeypatch: pytest.MonkeyPatch): + otel_sdk = pytest.importorskip("opentelemetry.sdk") + from opentelemetry.sdk.metrics import MeterProvider + from opentelemetry.sdk.metrics.export import InMemoryMetricReader + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor + from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + + from litellm.integrations.opentelemetry import OpenTelemetry as OTelV1, OpenTelemetryConfig + + reader = InMemoryMetricReader() + span_exporter = InMemorySpanExporter() + tracer_provider = TracerProvider() + tracer_provider.add_span_processor(SimpleSpanProcessor(span_exporter)) + logger = OTelV1( + config=OpenTelemetryConfig(exporter="console", enable_metrics=True), + callback_name="batch_sink_guard_v1", + tracer_provider=tracer_provider, + meter_provider=MeterProvider(metric_readers=[reader]), + ) + return logger, reader, span_exporter + + +def _cost_points(reader: Any) -> list[float]: + data = reader.get_metrics_data() + points: list[float] = [] + if data is None: + return points + for resource_metrics in data.resource_metrics: + for scope_metrics in resource_metrics.scope_metrics: + for metric in scope_metrics.metrics: + if not metric.name.endswith("cost"): + continue + for data_point in metric.data.data_points: + points.append(getattr(data_point, "sum", getattr(data_point, "value", None))) + return points + + +@pytest.mark.asyncio +async def test_otel_v1_meter_cost_skips_line_items_but_spans_stay( + monkeypatch: pytest.MonkeyPatch, +) -> None: + logger, reader, span_exporter = _otel_v1(monkeypatch) + + await _log_completed_batch(monkeypatch, [logger]) + + costs = _cost_points(reader) + assert sum(costs) == pytest.approx(AGGREGATE_COST), ( + f"OTel v1 gen_ai.usage.cost double-counted batch line items: {costs}; " + "line items must be metered only by the aggregate aretrieve_batch event" + ) + finished = span_exporter.get_finished_spans() + assert finished, "OTel v1 emitted no spans at all" + assert any("chatcmpl-line-1" in str(span.attributes) for span in finished), ( + f"OTel v1 dropped the per-line span the feature exists to deliver: {[span.name for span in finished]}" + ) + + +def _reader_snapshot(reader: Any) -> dict[str, list[float]]: + """Every metric's datapoint values, keyed by metric name (sums for histograms).""" + data = reader.get_metrics_data() + out: dict[str, list[float]] = {} + if data is None: + return out + for resource_metrics in data.resource_metrics: + for scope_metrics in resource_metrics.scope_metrics: + for metric in scope_metrics.metrics: + for data_point in metric.data.data_points: + value = getattr(data_point, "sum", None) + if value is None: + value = getattr(data_point, "value", None) + if value is not None: + out.setdefault(metric.name, []).append(value) + return out + + +async def _run_batch_with_otel_v2(monkeypatch: pytest.MonkeyPatch, flag: bool) -> dict[str, list[float]]: + from opentelemetry.sdk.metrics import MeterProvider + from opentelemetry.sdk.metrics.export import InMemoryMetricReader + from opentelemetry.sdk.trace import TracerProvider + + from litellm.integrations.otel.logger import OpenTelemetryV2 + from litellm.integrations.otel.model.config import OpenTelemetryV2Config + + reader = InMemoryMetricReader() + logger = OpenTelemetryV2( + config=OpenTelemetryV2Config(exporter="console", enable_metrics=True), + callback_name="batch_sink_guard_v2", + tracer_provider=TracerProvider(), + meter_provider=MeterProvider(metric_readers=[reader]), + ) + await _log_completed_batch(monkeypatch, [logger], flag=flag) + return _reader_snapshot(reader) + + +@pytest.mark.asyncio +async def test_otel_v2_metrics_are_identical_with_the_flag_on_and_off( + monkeypatch: pytest.MonkeyPatch, +) -> None: + pytest.importorskip("opentelemetry.sdk") + + off = await _run_batch_with_otel_v2(monkeypatch, flag=False) + on = await _run_batch_with_otel_v2(monkeypatch, flag=True) + + for name in ("gen_ai.usage.cost", "gen_ai.client.token.usage"): + assert on.get(name) == off.get(name), ( + f"OTel v2 {name} changed when the line-item flag was turned on; line items " + "must be metered only by the aggregate aretrieve_batch event. " + f"off={off.get(name)} on={on.get(name)}" + ) + for name, values in off.items(): + on_counts = [len(on.get(name, [])), len(values)] + assert on_counts[0] == on_counts[1], ( + f"OTel v2 {name} gained synthetic per-line samples with the flag on: " + f"{on_counts[0]} datapoints on vs {on_counts[1]} off" + ) + assert sum(off.get("gen_ai.usage.cost", [])) == pytest.approx(AGGREGATE_COST) diff --git a/tests/unit/litellm_core_utils/test_core_helpers.py b/tests/unit/litellm_core_utils/test_core_helpers.py index 2a6dd347d5f..774045f9d13 100644 --- a/tests/unit/litellm_core_utils/test_core_helpers.py +++ b/tests/unit/litellm_core_utils/test_core_helpers.py @@ -14,6 +14,7 @@ from litellm.litellm_core_utils.core_helpers import ( drop_params_flag, get_or_create_metadata_bucket, get_provider_response_headers_from_hidden_params, + is_batch_line_item_event, map_finish_reason, normalize_drop_params, reconstruct_model_name, @@ -496,6 +497,16 @@ class TestIsExpectedClientError: assert is_expected_client_error(vendor_limit) is False +def test_is_batch_line_item_event(): + assert is_batch_line_item_event({"litellm_params": {"batch_parent_id": "batch_1"}}) is True + assert is_batch_line_item_event({"litellm_params": {"batch_parent_id": "batch_1", "metadata": {}}}) is True + assert is_batch_line_item_event({"litellm_params": {"metadata": {}}}) is False + assert is_batch_line_item_event({"litellm_params": {"batch_parent_id": None}}) is False + assert is_batch_line_item_event({}) is False + assert is_batch_line_item_event({"litellm_params": "not-a-mapping"}) is False + assert is_batch_line_item_event({"litellm_params": None}) is False + + class TestProviderResponseHeadersInHiddenParams: def test_records_raw_headers_and_the_processed_additional_headers(self): response = ImageResponse() diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index c726127f892..536078ad33e 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -1197,19 +1197,24 @@ class TestRetrieveBatchCostPassesModelIdentity: captured: dict[str, object] = {} - from litellm.batches.batch_utils import BatchCostUsageResult + from litellm.batches.batch_utils import BatchCostUsageResult, BatchResultFiles - async def fake_handle_completed_batch(**kwargs: object) -> BatchCostUsageResult: + async def fake_handle_completed_batch( + **kwargs: object, + ) -> tuple[BatchCostUsageResult, BatchResultFiles]: captured.update(kwargs) - return BatchCostUsageResult( - cost=1.25, - usage=Usage(prompt_tokens=1800, completion_tokens=1000, total_tokens=2800), - models=["m"], - successful_requests=1, - failed_requests=0, + return ( + BatchCostUsageResult( + cost=1.25, + usage=Usage(prompt_tokens=1800, completion_tokens=1000, total_tokens=2800), + models=["m"], + successful_requests=1, + failed_requests=0, + ), + BatchResultFiles(output=None, error=None), ) - monkeypatch.setattr(logging_module, "_handle_completed_batch", fake_handle_completed_batch) + monkeypatch.setattr(logging_module, "_handle_completed_batch_with_files", fake_handle_completed_batch) obj = LitellmLogging( model="bedrock/global.anthropic.claude-sonnet-4-6", @@ -1292,7 +1297,7 @@ class TestRetrieveBatchPricesOnlyFinalBatches: from litellm.litellm_core_utils import litellm_logging as logging_module handle_completed_batch = AsyncMock() - monkeypatch.setattr(logging_module, "_handle_completed_batch", handle_completed_batch) + monkeypatch.setattr(logging_module, "_handle_completed_batch_with_files", handle_completed_batch) batch = self._batch(status, output_file_id) await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None) @@ -1302,20 +1307,23 @@ class TestRetrieveBatchPricesOnlyFinalBatches: @pytest.mark.asyncio async def test_completed_batch_with_output_is_priced(self, monkeypatch) -> None: - from litellm.batches.batch_utils import BatchCostUsageResult + from litellm.batches.batch_utils import BatchCostUsageResult, BatchResultFiles from litellm.litellm_core_utils import litellm_logging as logging_module from litellm.types.utils import Usage handle_completed_batch = AsyncMock( - return_value=BatchCostUsageResult( - cost=8e-06, - usage=Usage(prompt_tokens=26, completion_tokens=9, total_tokens=35), - models=["gpt-5.6-luna"], - successful_requests=2, - failed_requests=0, + return_value=( + BatchCostUsageResult( + cost=8e-06, + usage=Usage(prompt_tokens=26, completion_tokens=9, total_tokens=35), + models=["gpt-5.6-luna"], + successful_requests=2, + failed_requests=0, + ), + BatchResultFiles(output=None, error=None), ) ) - monkeypatch.setattr(logging_module, "_handle_completed_batch", handle_completed_batch) + monkeypatch.setattr(logging_module, "_handle_completed_batch_with_files", handle_completed_batch) batch = self._batch("completed", "file-out") await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None) @@ -9315,3 +9323,179 @@ def test_signoz_dispatch_requires_an_endpoint(monkeypatch): logging_module._in_memory_loggers.clear() monkeypatch.delenv("LITELLM_OTEL_V2", raising=False) is_otel_v2_enabled.cache_clear() + + +class TestRetrieveBatchReusesFetchedResultFiles: + """The aggregate path and log_batch_line_items must share one fetch of each + result file: output and error are fetched once total (not once each per + consumer), and forwarded bytes keep line-item logging off the wire.""" + + @staticmethod + def _batch_file_bytes() -> dict[str, bytes]: + input_jsonl = json.dumps( + { + "custom_id": "a", + "method": "POST", + "url": "/v1/chat/completions", + "body": {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]}, + } + ).encode() + output_jsonl = json.dumps( + { + "custom_id": "a", + "response": { + "status_code": 200, + "body": { + "id": "chatcmpl-1", + "model": "gpt-4o", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "hi back"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + }, + }, + } + ).encode() + error_jsonl = json.dumps({"custom_id": "b", "error": {"message": "rejected"}}).encode() + return {"file-in": input_jsonl, "file-out": output_jsonl, "file-err": error_jsonl} + + @staticmethod + def _logging_obj() -> LitellmLogging: + obj = LitellmLogging( + model="gpt-4o", + messages=[{"role": "user", "content": "Hey"}], + stream=False, + call_type="aretrieve_batch", + start_time=time.time(), + litellm_call_id="batch-call-reuse", + function_id="f", + ) + obj.custom_llm_provider = "openai" + obj.update_environment_variables( + litellm_params={"metadata": {}}, + optional_params={}, + custom_llm_provider="openai", + ) + return obj + + @staticmethod + def _batch(): + from litellm.types.utils import LiteLLMBatch + + return LiteLLMBatch( + id="batch_reuse_fetched_files", + completion_window="24h", + created_at=1, + endpoint="/v1/chat/completions", + input_file_id="file-in", + object="batch", + status="completed", + output_file_id="file-out", + error_file_id="file-err", + ) + + @pytest.mark.asyncio + async def test_compute_path_fetches_each_result_file_once(self, monkeypatch) -> None: + from litellm.batches.batch_line_item_logging import batch_line_item_claim_cache + + batch_line_item_claim_cache.in_memory_cache.flush_cache() + file_bytes = self._batch_file_bytes() + + async def fake_afile_content(**kwargs): + from types import SimpleNamespace + + return SimpleNamespace(content=file_bytes[kwargs["file_id"]]) + + file_mock = AsyncMock(side_effect=fake_afile_content) + monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", True, raising=False) + monkeypatch.setattr("litellm.files.main.afile_content", file_mock) + monkeypatch.setattr("litellm.cost_calculator.batch_cost_calculator", lambda **kw: (0.01, 0.02)) + + await self._logging_obj()._async_success_handler_body( + result=self._batch(), start_time=None, end_time=None + ) + + fetched = sorted(call.kwargs["file_id"] for call in file_mock.await_args_list) + assert fetched == ["file-err", "file-in", "file-out"], ( + "each result file must be fetched exactly once across aggregate costing and line-item logging" + ) + + @pytest.mark.asyncio + async def test_explicit_kwargs_path_forwards_result_files(self, monkeypatch) -> None: + from litellm.batches.batch_line_item_logging import batch_line_item_claim_cache + + batch_line_item_claim_cache.in_memory_cache.flush_cache() + from litellm.types.utils import Usage + + file_bytes = self._batch_file_bytes() + + async def fake_afile_content(**kwargs): + from types import SimpleNamespace + + return SimpleNamespace(content=file_bytes[kwargs["file_id"]]) + + file_mock = AsyncMock(side_effect=fake_afile_content) + monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", True, raising=False) + monkeypatch.setattr("litellm.files.main.afile_content", file_mock) + + await self._logging_obj().async_success_handler( + result=self._batch(), + batch_cost=1.5, + batch_usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), + batch_models=["gpt-4o"], + batch_successful_requests=1, + batch_failed_requests=1, + batch_output_file_content=file_bytes["file-out"], + batch_error_file_content=file_bytes["file-err"], + ) + + assert file_mock.await_count == 1 + assert file_mock.await_args.kwargs["file_id"] == "file-in", ( + "forwarded output/error bytes must leave only the input file to fetch" + ) + + @pytest.mark.asyncio + async def test_line_item_logging_failure_leaves_aggregate_result_intact(self, monkeypatch) -> None: + from litellm.batches.batch_line_item_logging import batch_line_item_claim_cache + + batch_line_item_claim_cache.in_memory_cache.flush_cache() + from litellm.types.utils import Usage + + monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", True, raising=False) + monkeypatch.setattr( + "litellm.batches.batch_line_item_logging.log_batch_line_items", + AsyncMock(side_effect=RuntimeError("claim backend exploded")), + ) + + batch = self._batch() + await self._logging_obj().async_success_handler( + result=batch, + batch_cost=1.5, + batch_usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), + batch_models=["gpt-4o"], + batch_successful_requests=1, + batch_failed_requests=1, + ) + + assert batch._hidden_params["response_cost"] == 1.5 + assert batch.usage.total_tokens == 15 + + +def test_caller_supplied_litellm_params_cannot_forge_logging_markers(logging_obj): + """A request body key `litellm_params` flows into optional_params; the logging + kwargs must keep the internal litellm_params or a caller could forge markers + (batch_parent_id) that spend sinks and router counters use to skip events.""" + from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event + + logging_obj.update_environment_variables( + litellm_params={"metadata": {"model_group": "gpt-4o"}}, + optional_params={"litellm_params": {"batch_parent_id": "fake-batch"}}, + ) + + assert logging_obj.model_call_details["litellm_params"] is logging_obj.litellm_params + assert "batch_parent_id" not in logging_obj.model_call_details["litellm_params"] + assert not is_batch_line_item_event(logging_obj.model_call_details) diff --git a/tests/unit/llms/bedrock/batches/test_transformation.py b/tests/unit/llms/bedrock/batches/test_transformation.py index 5e987239c54..4a9b175714c 100644 --- a/tests/unit/llms/bedrock/batches/test_transformation.py +++ b/tests/unit/llms/bedrock/batches/test_transformation.py @@ -20,7 +20,7 @@ import httpx import pytest from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig -from litellm.types.utils import LlmProviders +from litellm.types.utils import EmbeddingResponse, LlmProviders, ModelResponse # AWS JobStatus -> OpenAI BatchJobStatus, exactly as encoded in transformation.py # (both transform_create_batch_response and transform_retrieve_batch_response). @@ -941,3 +941,25 @@ def test_retrieve_request_accepts_partition_arns(config: BedrockBatchesConfig, a batch_id=arn, optional_params={}, litellm_params={} ) assert result["url"].startswith(expected_prefix) + + +def test_transform_batch_output_line_dispatches_on_shape(config: BedrockBatchesConfig) -> None: + titan = config.transform_batch_output_line( + {"embedding": [0.1, 0.2], "inputTextTokenCount": 4}, model="amazon.titan-embed" + ) + assert isinstance(titan, EmbeddingResponse) + assert titan.data[0]["embedding"] == [0.1, 0.2] + assert titan.usage.prompt_tokens == 4 + + converse = config.transform_batch_output_line( + { + "output": {"message": {"role": "assistant", "content": [{"text": "hi"}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 5, "outputTokens": 6, "totalTokens": 11}, + }, + model="amazon.nova-lite", + ) + assert isinstance(converse, ModelResponse) + assert converse.choices[0].message.content == "hi" + + assert config.transform_batch_output_line({"foo": 1}, model="x") is None diff --git a/tests/unit/proxy/common_utils/test_check_batch_cost.py b/tests/unit/proxy/common_utils/test_check_batch_cost.py index 4e7effad3bf..41470cc0d2e 100644 --- a/tests/unit/proxy/common_utils/test_check_batch_cost.py +++ b/tests/unit/proxy/common_utils/test_check_batch_cost.py @@ -1546,10 +1546,10 @@ class TestCheckBatchCost: Logging, "async_success_handler", new_callable=AsyncMock ) as success_handler, ): - provider.get("https://api.openai.com/v1/files/file-output-123/content").mock( + output_route = provider.get("https://api.openai.com/v1/files/file-output-123/content").mock( return_value=httpx.Response(200, content=f"{succeeded_line}\n{rejected_line}\n".encode()) ) - provider.get("https://api.openai.com/v1/files/file-error-456/content").mock( + error_route = provider.get("https://api.openai.com/v1/files/file-error-456/content").mock( return_value=httpx.Response(200, content=f"{error_file_lines}\n\n".encode()) ) await check_batch_cost_instance.check_batch_cost() @@ -1558,6 +1558,10 @@ class TestCheckBatchCost: assert len(spend_log_calls) == 1 handler_kwargs = spend_log_calls[0] assert handler_kwargs["batch_successful_requests"] == 1 + assert output_route.call_count == 1 + assert error_route.call_count == 1 + assert handler_kwargs["batch_output_file_content"] == f"{succeeded_line}\n{rejected_line}\n".encode() + assert handler_kwargs["batch_error_file_content"] == f"{error_file_lines}\n\n".encode() assert handler_kwargs["batch_failed_requests"] == 3, ( "2 error-file lines must add to the output file's 1 rejected request" ) diff --git a/tests/unit/proxy/hooks/test_model_max_budget_limiter.py b/tests/unit/proxy/hooks/test_model_max_budget_limiter.py index ffb60fb4651..b062efb3e81 100644 --- a/tests/unit/proxy/hooks/test_model_max_budget_limiter.py +++ b/tests/unit/proxy/hooks/test_model_max_budget_limiter.py @@ -196,3 +196,17 @@ async def test_a_batch_polled_within_every_budget_window_is_never_charged_again( await _poll(limiter, finished, BATCH_COST) assert _local_spend(limiter, KEY_SPEND_KEY) == pytest.approx(BATCH_COST) + + +@pytest.mark.asyncio +async def test_batch_line_item_events_do_not_charge_the_model_budget(): + """Line events carry batch_parent_id; the aggregate aretrieve_batch event is + the one that already bills the batch, so children must not double-charge.""" + limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + + line_event = _event("acompletion", CHAT_COST) + line_event["litellm_params"]["batch_parent_id"] = "batch_first" + await limiter.async_log_success_event(line_event, response_obj=None, start_time=None, end_time=None) + + assert await _spend(limiter, KEY_SPEND_KEY) == 0.0 + assert await _spend(limiter, USER_SPEND_KEY) == 0.0 diff --git a/tests/unit/proxy/hooks/test_parallel_request_limiter.py b/tests/unit/proxy/hooks/test_parallel_request_limiter.py index a83ddc69863..c38968e47d3 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter.py @@ -107,3 +107,29 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_ f"expected 50 tokens counted for {scope_id}, " f"got {current['current_tpm']}" ) + + +@pytest.mark.asyncio +async def test_async_log_failure_event_skips_batch_line_item_events(): + """Failed line children were never admitted by the limiter, so the failure + hook must not decrement request counters they never incremented.""" + parallel_request_handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(DualCache()) + ) + local_cache = parallel_request_handler.internal_usage_cache.dual_cache.in_memory_cache + + await parallel_request_handler.async_log_failure_event( + kwargs={ + "exception": "litellm.APIError: upstream 500", + "litellm_params": { + "batch_parent_id": "batch_1", + "metadata": {"user_api_key": hash_token("sk-line-item")}, + }, + "model": "gpt-3.5-turbo", + }, + response_obj=None, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + assert local_cache.cache_dict == {} diff --git a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py index d3a723fe1d3..fdf17d65c5a 100644 --- a/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/unit/proxy/hooks/test_parallel_request_limiter_v3.py @@ -7667,3 +7667,59 @@ async def test_a2a_url_target_owns_invocation_fee_and_request_limit( await _rpm_request(limiter, cache, auth, "a2a/cheap") assert denied.value.status_code == 429 assert "expensive" in str(denied.value.detail) + + +@pytest.mark.asyncio +async def test_async_log_success_event_skips_batch_line_item_events(): + """Per-line batch events already ran through the limiter as the aggregate + aretrieve_batch; children must not increment TPM or request counters.""" + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) + + await handler.async_log_success_event( + kwargs={ + "standard_logging_object": {"metadata": {"user_api_key_hash": hash_token("sk-line-item")}}, + "litellm_params": { + "batch_parent_id": "batch_1", + "metadata": {"user_api_key_hash": hash_token("sk-line-item"), "model_group": "gpt-3.5-turbo"}, + }, + "model": "gpt-3.5-turbo", + }, + response_obj=ModelResponse( + id="x", object="chat.completion", created=1, model="gpt-3.5-turbo", + usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), choices=[], + ), + start_time=datetime.now(), + end_time=datetime.now(), + ) + + assert local_cache.in_memory_cache.cache_dict == {} + + +@pytest.mark.asyncio +async def test_async_log_failure_event_skips_batch_line_item_events(): + """Failed line children were never admitted by the limiter, so the failure + hook must not release slots or refund TPM they never reserved.""" + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) + stash = get_or_create_request_stash() + stash.owner_litellm_call_id = "call-line-item" + stash.reserved_tokens = 50 + stash.reserved_scopes = frozenset({("api_key", hash_token("sk-line-item"))}) + + await handler.async_log_failure_event( + kwargs={ + "litellm_call_id": "call-line-item", + "standard_logging_object": {"metadata": {"user_api_key_hash": hash_token("sk-line-item")}}, + "litellm_params": { + "batch_parent_id": "batch_1", + "metadata": {"user_api_key_hash": hash_token("sk-line-item"), "model_group": "gpt-3.5-turbo"}, + }, + "model": "gpt-3.5-turbo", + }, + response_obj=None, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + assert local_cache.in_memory_cache.cache_dict == {} diff --git a/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py b/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py index 7afa275c801..c3b099a4ff4 100644 --- a/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py @@ -2772,6 +2772,24 @@ async def test_async_post_call_failure_hook_persists_no_raw_model_on_an_unknown_ assert error_information["error_class"] == "ProxyModelNotFoundError" +@pytest.mark.asyncio +async def test_batch_line_item_event_never_updates_spend(): # test-quality-ok: the observable contract is exactly that the spend path is never invoked + logger: Final = _ProxyDBLogger() + kwargs: Final = { + "litellm_params": {"batch_parent_id": "batch_1", "metadata": {}}, + "model": "gpt-4o", + "call_type": "acompletion", + } + with patch.object(logger, "_PROXY_track_cost_callback", new_callable=AsyncMock) as mock_track: # test-quality-ok: asserts the callback's own method is skipped; the DB writer is never reached + await logger.async_log_success_event( + kwargs=kwargs, + response_obj=ModelResponse(), + start_time=datetime.now(), + end_time=datetime.now(), + ) + mock_track.assert_not_awaited() + + class _NeverStringifiedMetadataValue: def __repr__(self) -> str: raise AssertionError("a request metadata value was stringified by the cost tracking failure path") diff --git a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py index e2fbd5dcd98..84b3a4c421b 100644 --- a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py +++ b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py @@ -8112,6 +8112,43 @@ async def test_update_general_settings_store_model_in_db_false(): assert ps.general_settings["store_model_in_db"] is False +@pytest.mark.asyncio +async def test_update_general_settings_store_batch_line_items_in_callbacks(): + """ + Verify _update_general_settings sets the litellm module flag when the DB + general_settings carries store_batch_line_items_in_callbacks, and that a + YAML-explicit value wins over the DB value. + """ + from litellm.proxy.proxy_server import ProxyConfig + + proxy_config = ProxyConfig() + saved_flag = litellm.store_batch_line_items_in_callbacks + try: + with patch("litellm.proxy.proxy_server.general_settings", {}): # test-quality-ok: module-global seam + await proxy_config._update_general_settings( + db_general_settings={"store_batch_line_items_in_callbacks": True} + ) + assert litellm.store_batch_line_items_in_callbacks is True + + proxy_config._yaml_general_settings_keys = {"store_batch_line_items_in_callbacks"} + with patch( # test-quality-ok: module-global seam + "litellm.proxy.proxy_server.general_settings", + {"store_batch_line_items_in_callbacks": "false"}, + ): + await proxy_config._update_general_settings( + db_general_settings={"store_batch_line_items_in_callbacks": True} + ) + assert litellm.store_batch_line_items_in_callbacks is False + + proxy_config._yaml_general_settings_keys = set() + litellm.store_batch_line_items_in_callbacks = True # test-quality-ok: seed the prior opt-in so its removal is observable + with patch("litellm.proxy.proxy_server.general_settings", {}): # test-quality-ok: module-global seam + await proxy_config._update_general_settings(db_general_settings={}) + assert litellm.store_batch_line_items_in_callbacks is False + finally: + litellm.store_batch_line_items_in_callbacks = saved_flag + + @pytest.mark.asyncio async def test_update_general_settings_propagates_apply_user_budget_to_team_keys(): """The Admin UI toggle writes to the DB config, so the flag has to be in the @@ -13288,6 +13325,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/unit/proxy/test_redis_auth_cache_flag.py b/tests/unit/proxy/test_redis_auth_cache_flag.py index cb600bbb5fd..e1eb9d13bdc 100644 --- a/tests/unit/proxy/test_redis_auth_cache_flag.py +++ b/tests/unit/proxy/test_redis_auth_cache_flag.py @@ -184,10 +184,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/router_strategy/test_budget_limiter_hotpath.py b/tests/unit/router_strategy/test_budget_limiter_hotpath.py index a417b789397..2c3561ca67f 100644 --- a/tests/unit/router_strategy/test_budget_limiter_hotpath.py +++ b/tests/unit/router_strategy/test_budget_limiter_hotpath.py @@ -391,6 +391,29 @@ async def test_push_task_failure_is_logged_once_and_not_leaked(disable_budget_sy unretrieved.assert_not_called() +@pytest.mark.asyncio +async def test_batch_line_item_events_do_not_charge_the_provider_budget(disable_budget_sync): + limiter = RouterBudgetLimiting( + dual_cache=DualCache(), + provider_budget_config={"openai": BudgetConfig(max_budget=100.0, budget_duration="1d")}, + ) + await asyncio.gather(*(task for task in asyncio.all_tasks() if task is not asyncio.current_task())) + + await limiter.async_log_success_event( + kwargs={ + "call_type": "acompletion", + "litellm_params": {"custom_llm_provider": "openai", "batch_parent_id": "batch_x"}, + "standard_logging_object": {"response_cost": 1.0, "model_id": "dep-1"}, + }, + response_obj=None, + start_time=None, + end_time=None, + ) + + assert limiter.dual_cache.in_memory_cache.get_cache("provider_spend:openai:1d") in (None, 0) + assert limiter.redis_increment_operation_queue == [] + + _SPEND_KEY = "provider_spend:openai:1d" diff --git a/tests/unit/router_strategy/test_least_busy.py b/tests/unit/router_strategy/test_least_busy.py index c2fa41f4ca8..cd0537dd8f5 100644 --- a/tests/unit/router_strategy/test_least_busy.py +++ b/tests/unit/router_strategy/test_least_busy.py @@ -208,3 +208,28 @@ async def test_an_open_circuit_breaker_falls_back_without_a_warning_per_request( assert picked is DEPLOYMENT_B assert [record.getMessage() for record in caplog.records if record.levelno >= logging.WARNING] == [] assert sum("circuit breaker is open" in record.getMessage() for record in caplog.records) == 2 + + +@pytest.mark.asyncio +async def test_batch_line_item_success_does_not_release_in_flight_slots(monkeypatch: pytest.MonkeyPatch) -> None: + """A batch line item is historical batch traffic (call_type=acompletion + + litellm_params.batch_parent_id): it never started a live request, so its success + callback must not decrement the in-flight count a live request still holds.""" + import litellm + + monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", True, raising=False) + shared: Final = SharedRedisCounters() + worker: Final = _worker(shared) + worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a")) + + line_item_kwargs: Final = { + "call_type": "acompletion", + "litellm_params": { + "batch_parent_id": "batch-1", + "metadata": {"model_group": GROUP}, + "model_info": {"id": "dep-a"}, + }, + } + await worker.async_log_success_event(line_item_kwargs, None, None, None) + + assert shared.count(f"{GROUP}_request_count:dep-a") == 1 diff --git a/tests/unit/router_strategy/test_lowest_tpm_rpm.py b/tests/unit/router_strategy/test_lowest_tpm_rpm.py index 625f648bec4..55147b316db 100644 --- a/tests/unit/router_strategy/test_lowest_tpm_rpm.py +++ b/tests/unit/router_strategy/test_lowest_tpm_rpm.py @@ -117,3 +117,40 @@ async def test_v2_subclass_overriding_async_get_available_deployments_with_the_o f"from {HIGH_USAGE_DEPLOYMENT_ID}", f"from {LOW_USAGE_DEPLOYMENT_ID}", } + + +@pytest.mark.asyncio +async def test_usage_based_routing_handlers_skip_batch_line_items(monkeypatch: pytest.MonkeyPatch) -> None: + """ + Batch line-item callbacks carry call_type=acompletion plus + litellm_params.batch_parent_id, so the call-type-only batch guard does not fire. + They are historical batch traffic and must not update the TPM/RPM usage counters. + """ + import litellm + from litellm.router_strategy.lowest_tpm_rpm import LowestTPMLoggingHandler + + monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", True, raising=False) + router: Final = Router( + model_list=[_deployment(HIGH_USAGE_DEPLOYMENT_ID)], + routing_strategy="usage-based-routing", + ) + handler: Final = LowestTPMLoggingHandler(router_cache=router.cache, routing_args={}) + line_item_kwargs: Final = { + "call_type": "acompletion", + "litellm_params": { + "batch_parent_id": "batch-1", + "metadata": {"model_group": MODEL_GROUP}, + "model_info": {"id": HIGH_USAGE_DEPLOYMENT_ID}, + }, + } + response_obj: Final = {"usage": {"total_tokens": 600}} + + handler.log_success_event(line_item_kwargs, response_obj, None, None) + await handler.async_log_success_event(line_item_kwargs, response_obj, None, None) + + moved: Final = sorted( + f"{key}={router.cache.in_memory_cache.cache_dict[key]}" + for key in router.cache.in_memory_cache.cache_dict + if ":tpm:" in key or ":rpm:" in key + ) + assert moved == [] diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index d0115593e46..d8493785f5e 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -1265,6 +1265,53 @@ def test_sync_deployment_callback_on_success_skips_batch_retrieves( == expected_successes ) + +@pytest.mark.asyncio +async def test_deployment_callbacks_skip_batch_line_items(monkeypatch: pytest.MonkeyPatch): + """ + Batch line-item callbacks carry call_type=acompletion plus + litellm_params.batch_parent_id, so the call-type-only batch guards do not fire. + They describe historical batch traffic already reported by the aggregate + aretrieve_batch event and must not consume live TPM/RPM quota. + """ + monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", True, raising=False) + router = litellm.Router( + model_list=[ + { + "model_name": _BATCH_GROUP, + "litellm_params": {"model": _BATCH_DEPLOYMENT_MODEL, "api_base": _BATCH_API_BASE, "api_key": "sk-fake"}, + "model_info": {"id": "batch-dep"}, + } + ] + ) + now = datetime.now() + line_item_kwargs = { + "call_type": "acompletion", + "standard_logging_object": {"total_tokens": _BATCH_TOKENS_PER_ROW}, + "litellm_params": { + "batch_parent_id": _BATCH_ID, + "metadata": {"model_group": _BATCH_GROUP, "deployment": _BATCH_DEPLOYMENT_MODEL}, + "model_info": {"id": "batch-dep"}, + }, + } + + await router.deployment_callback_on_success( + kwargs=line_item_kwargs, completion_response=None, start_time=now, end_time=now + ) + sync_key = router.sync_deployment_callback_on_success( + kwargs=line_item_kwargs, completion_response=None, start_time=now, end_time=now + ) + await router.async_deployment_callback_on_failure( + kwargs=line_item_kwargs, completion_response=None, start_time=now, end_time=now + ) + + assert sync_key is None + assert ( + get_deployment_successes_for_current_minute(litellm_router_instance=router, deployment_id="batch-dep") == 0 + ) + assert await _moved_routing_counters(router) == [] + assert await _router_usage_keys(router) == [] + _ROUTING_STRATEGY_CACHE_MARKERS = ("_map", "_request_count", ":tpm:", ":rpm:") @@ -1352,6 +1399,81 @@ async def test_arouter_aretrieve_batch_does_not_feed_routing_strategies( assert moved_counters == [] +_BATCH_INPUT_JSONL = "\n".join( + json.dumps( + { + "custom_id": f"row-{row}", + "method": "POST", + "url": "/v1/chat/completions", + "body": {"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]}, + } + ) + for row in range(_BATCH_ROWS) +) + + +def _mock_batch_input_file(respx_mock): + respx_mock.get(f"{_BATCH_API_BASE}/files/file-in-1/content").mock( + return_value=httpx.Response(200, text=_BATCH_INPUT_JSONL) + ) + + +async def _line_item_payloads(collector: "_BatchPayloadCollector", minimum: int, timeout: float = 5.0) -> list: + loop = asyncio.get_event_loop() + deadline = loop.time() + timeout + while loop.time() < deadline: + line_items = [p for p in collector.payloads if p and p.get("call_type") == "acompletion"] + if len(line_items) >= minimum: + return line_items + await asyncio.sleep(0.05) + raise AssertionError(f"expected at least {minimum} batch line-item payloads, saw {len(collector.payloads)} total") + + +@pytest.mark.parametrize( + "routing_strategy", + [ + "usage-based-routing", + "usage-based-routing-v2", + "latency-based-routing", + "cost-based-routing", + "least-busy", + ], +) +@pytest.mark.asyncio +async def test_arouter_aretrieve_batch_line_items_do_not_feed_routing_strategies( + monkeypatch: pytest.MonkeyPatch, routing_strategy: str +): + """ + With the line-item flag on, retrieving a completed batch also emits one child + callback per JSONL line (call_type=acompletion + litellm_params.batch_parent_id, + inheriting the parent deployment's model_info). Those children are historical + batch traffic: like the aggregate poll, they must not move the counters that + decide where the next live chat request goes. + """ + from litellm.batches.batch_line_item_logging import batch_line_item_claim_cache + + batch_line_item_claim_cache.in_memory_cache.flush_cache() + monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", True, raising=False) + collector = _BatchPayloadCollector() + monkeypatch.setattr(litellm, "callbacks", [collector]) + monkeypatch.setattr(litellm, "input_callback", []) + router = _batch_fan_out_router(routing_strategy) + + with respx.mock(assert_all_called=True) as respx_mock: + _mock_batch_provider(respx_mock) + _mock_batch_input_file(respx_mock) + respx_mock.get(f"{_UNRELATED_BATCH_API_BASE}/batches/{_BATCH_ID}").mock( + return_value=httpx.Response(404, json=_BATCH_NOT_FOUND) + ) + response = await router.aretrieve_batch(batch_id=_BATCH_ID) + line_items = await _line_item_payloads(collector, minimum=_BATCH_ROWS) + moved_counters = await _moved_routing_counters(router) + + assert response.id == _BATCH_ID + assert len(line_items) == _BATCH_ROWS + assert moved_counters == [] + + @pytest.mark.asyncio async def test_arouter_aretrieve_file_content(): """ diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index bc407c30c70..5f92a47fa51 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -29867,6 +29867,11 @@ export interface components { scheduled_job_stagger?: components["schemas"]["ScheduledJobStaggerSettings"] | null; /** @description Daily check of the spend LiteLLM captured against the provider's own bill (OpenAI via OPENAI_ADMIN_KEY). Publishes litellm_spend_capture_rate per provider and alerts when the ratio over the lookback window falls under the threshold (default 0.9). Off unless set. */ spend_capture_rate_check?: components["schemas"]["SpendCaptureRateCheckSettings"] | null; + /** + * Store Batch Line Items In Callbacks + * @description If True, a completed batch logged via aretrieve_batch also emits one callback event per JSONL line item (request paired with its response or error). The aggregate batch callback is unchanged. Default is False. + */ + store_batch_line_items_in_callbacks?: boolean | null; /** * Store Model In Db * @description If True, models and config are stored in and loaded from the database. Default is False.