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 13e9e5093a8..457bcf8c147 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -856,6 +856,7 @@ class CheckBatchCost: } }, **({"api_base": mask_api_base_credentials(deployment_api_base)} if deployment_api_base else {}), + "_litellm_internal_model_credentials": MappingProxyType({**credentials}), "metadata": { **(await self._build_creator_attribution_metadata(job, batch_id)), # spend logs read the deployment identity off these metadata keys, so diff --git a/litellm/__init__.py b/litellm/__init__.py index 71857877e53..7aaef9b91d5 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..077457dd299 --- /dev/null +++ b/litellm/batches/batch_line_item_logging.py @@ -0,0 +1,341 @@ +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 + +from litellm._logging import verbose_logger +from litellm.batches.batch_utils import ( + _batch_response_was_successful, # pyright: ignore[reportPrivateUsage] # batch-internal helper shared with the aggregate cost path by design + _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 +) +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"] + +_CALL_TYPE_BY_BATCH_URL: Final = MappingProxyType( + { + "/v1/chat/completions": "acompletion", + "/v1/embeddings": "aembedding", + "/v1/responses": "aresponses", + } +) + +_EMPTY_BODY: Final[Mapping[str, object]] = MappingProxyType({}) + + +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)) + self._hidden_params: dict[str, object] = {} # mutable-ok: mirrors the plain-dict _hidden_params contract on litellm response objects + + +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 _requests_by_custom_id(input_file_content: bytes) -> Mapping[str, Mapping[str, object]]: + """Parse the batch input JSONL into {custom_id: request line}, skipping + malformed lines and lines without a custom_id.""" + return MappingProxyType( + { + custom_id: entry + for entry in _output_entries(input_file_content) + if isinstance((custom_id := entry.get("custom_id")), str) and custom_id + } + ) + + +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 + 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 _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 _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, response_body: Mapping[str, object]) -> _BatchLineResult: + 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 + return ModelResponse(**response_body) # pyright: ignore[reportArgumentType] # same as above + + +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}, # mutable-ok: Logging's kwargs param takes a plain dict + ) + + +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 { # mutable-ok: same contract + "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 { # mutable-ok: same contract + key: value for key, value in request_body.items() if key not in ("model", "messages", "input") + } + + +async def _emit_line_event( + entry: Mapping[str, object], + request_line: Mapping[str, object] | None, + batch: LiteLLMBatch, + custom_llm_provider: _BatchLineProvider, + parent: "Logging", + model_name: str | None, + model_info: ModelInfo | None, +) -> bool: + custom_id: Final = entry.get("custom_id") or entry.get("recordId") + request_body: Final = _request_body_for_entry(entry, request_line) + status_code: Final = _line_status_code(entry, custom_llm_provider) + 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 + + child: Final = _new_child_logging( + parent=parent, + model=_line_model(response_body, request_body, parent), + 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={ # mutable-ok: update_environment_variables takes a plain dict + **parent_params, + "batch_parent_id": batch.id, + "metadata": dict(_as_object_mapping(parent_params.get("metadata")) or {}), # mutable-ok: copy of the parent's metadata dict + }, + optional_params=_optional_params_for_body(request_body), + model=child.model, + custom_llm_provider=custom_llm_provider, + ) + + now: Final = datetime.now() # noqa: DTZ005 # naive to match the logging pipeline start_time + if not _batch_response_was_successful(entry, custom_llm_provider): + exception: Final = _BatchLineFailure( + entry.get("error") or entry.get("response") or {} # mutable-ok: fallback payload dict passed to Exception + ) + 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 + + stats: Final = _safe_output_line_stats(entry, custom_llm_provider, model_name, model_info) + try: + result: Final = _line_result(call_type, 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", + call_type, + custom_id, + ) + return False + + 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: _BatchLineProvider, + 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, +) -> 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.""" + emitted = 0 # rebind-ok: loop accumulator for emitted line count + try: + internal_credentials: Final = ( + litellm_params.get("_litellm_internal_model_credentials") if litellm_params else None + ) + internal_mapping: Final = _as_object_mapping(internal_credentials) + fetch_params: Final[dict[str, object] | None] = ( # mutable-ok: _fetch_batch_managed_file_content requires a plain dict + dict(internal_mapping) # mutable-ok: the file fetcher reads credential kwargs off a plain dict + if internal_mapping is not None + else litellm_params + ) + + input_file_content: Final = await _fetch_managed_file_or_empty( + batch.input_file_id, custom_llm_provider, fetch_params + ) + requests_by_id: Final = _requests_by_custom_id(input_file_content) + + output_content: Final = await _fetch_managed_file_or_empty( + batch.output_file_id, custom_llm_provider, fetch_params + ) + error_content: Final = await _fetch_managed_file_or_empty( + batch.error_file_id, custom_llm_provider, fetch_params + ) + for content in (output_content, error_content): + for entry in _output_entries(content): + try: + entry_key = entry.get("custom_id") or entry.get("recordId") # rebind-ok: per-iteration binding inside a loop cannot carry Final + request_line = requests_by_id.get(entry_key if isinstance(entry_key, str) else "") # rebind-ok: per-iteration binding inside a loop cannot carry Final + line_emitted = await _emit_line_event( # rebind-ok: same per-iteration binding + entry=entry, + request_line=request_line, + batch=batch, + custom_llm_provider=custom_llm_provider, + parent=parent, + model_name=model_name, + model_info=model_info, + ) + if line_emitted: + emitted += 1 + 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, + ) + return emitted diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 40621a2f68d..5e8139c1e8d 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -3027,6 +3027,18 @@ class Logging(LiteLLMLoggingBaseClass): cost_for_built_in_tools_cost_usd_dollar=0.0, ) + if litellm.store_batch_line_items_in_callbacks: + from litellm.batches.batch_line_item_logging import log_batch_line_items + + 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=self.litellm_params, + model_info=self.get_router_deployment_model_info(), + ) + 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") @@ -6102,8 +6114,10 @@ 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 = getattr(original_exception, "_hidden_params", None) + if isinstance(exception_hidden_params, dict) and exception_hidden_params: + hidden_params = dict(exception_hidden_params) # mutable-ok: hidden_params downstream expects a plain dict + 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/proxy/_types.py b/litellm/proxy/_types.py index 680c63393e8..dacbfb8ec81 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2797,6 +2797,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/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 1ae106be390..8ad4de38095 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -1,6 +1,6 @@ import asyncio import traceback -from collections.abc import Callable, Sequence +from collections.abc import Callable, Mapping, Sequence from datetime import datetime from typing import TYPE_CHECKING, Any, Final, cast @@ -105,6 +105,11 @@ class _ProxyDBLogger(CustomLogger): async def async_log_success_event( self, kwargs: ObjectMapping, response_obj: object, start_time: datetime, end_time: datetime ) -> None: + # Per-line batch events emitted under store_batch_line_items_in_callbacks + # never touch spend: the aggregate aretrieve_batch event already bills the batch. + litellm_params: Final = kwargs.get("litellm_params") + if isinstance(litellm_params, Mapping) and litellm_params.get("batch_parent_id"): + 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 d7d8413d2ce..e47361d1de5 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6044,6 +6044,13 @@ class ProxyConfig: litellm.use_legacy_interactions_schema = _use_legacy_interactions_schema.lower() == "true" else: litellm.use_legacy_interactions_schema = bool(_use_legacy_interactions_schema) + ### BATCH LINE ITEM CALLBACKS ### + _store_batch_line_items: Final = 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) @@ -7158,6 +7165,19 @@ class ProxyConfig: # For other types, convert to bool general_settings["store_prompts_in_spend_logs"] = bool(value) + if "store_batch_line_items_in_callbacks" in _general_settings: + store_line_items_value: Final = ( + general_settings.get("store_batch_line_items_in_callbacks") + if "store_batch_line_items_in_callbacks" in self._yaml_general_settings_keys + else _general_settings["store_batch_line_items_in_callbacks"] + ) + if store_line_items_value is not None: + litellm.store_batch_line_items_in_callbacks = ( + store_line_items_value.lower() == "true" + if isinstance(store_line_items_value, str) + else bool(store_line_items_value) + ) + if "disable_auto_add_proxy_admin_to_teams" in _general_settings: value = _general_settings["disable_auto_add_proxy_admin_to_teams"] if isinstance(value, str): diff --git a/litellm/types/utils.py b/litellm/types/utils.py index aaa16fd2d44..8db91ffcb52 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3098,6 +3098,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/test_litellm/batches/test_batch_line_item_logging.py b/tests/test_litellm/batches/test_batch_line_item_logging.py new file mode 100644 index 00000000000..7fc60ddba77 --- /dev/null +++ b/tests/test_litellm/batches/test_batch_line_item_logging.py @@ -0,0 +1,233 @@ +""" +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 SimpleNamespace +from unittest.mock import AsyncMock, patch + +import pytest + +import litellm +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 +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() -> 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="openai", + ) + 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 "batch_custom_id" not in _hidden(aggregate) + + 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_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 diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index dfc95db3e14..b42dc41f097 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -2536,3 +2536,21 @@ async def test_async_post_call_failure_hook_persists_no_raw_model_on_an_unknown_ == "/chat/completions: Invalid model name passed in. Call `/v1/models` to view available models for your key." ) 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()