feat(batches): emit per-line JSONL batch records to callbacks

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
shivam 2026-09-17 22:48:28 +00:00
parent 356b8d4074
commit b38b867918
10 changed files with 643 additions and 3 deletions

View file

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

View file

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

View file

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

View file

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

View file

@ -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.",

View file

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

View file

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

View file

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

View file

@ -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": "<retrieve_batch>"}],
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

View file

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