mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
356b8d4074
commit
b38b867918
10 changed files with 643 additions and 3 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
341
litellm/batches/batch_line_item_logging.py
Normal file
341
litellm/batches/batch_line_item_logging.py
Normal 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
|
||||
|
|
@ -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)),
|
||||
|
|
|
|||
|
|
@ -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.",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
233
tests/test_litellm/batches/test_batch_line_item_logging.py
Normal file
233
tests/test_litellm/batches/test_batch_line_item_logging.py
Normal 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
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue