This commit is contained in:
devin-ai-integration[bot] 2026-10-05 19:49:35 +00:00 • committed by GitHub
commit 980eab30c7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
54 changed files with 3757 additions and 120 deletions

View file

@ -703,9 +703,10 @@ class CheckBatchCost:
later poll.
"""
from litellm.batches.batch_utils import (
count_error_file_failed_requests,
_get_file_content_as_dictionary,
calculate_batch_cost_and_usage,
count_error_file_failed_requests_from_content,
fetch_batch_error_file_content,
)
from litellm.files.main import afile_content
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
@ -848,7 +849,7 @@ class CheckBatchCost:
model_name=model_name,
model_info=deployment_model_info,
)
error_file_failed_requests: Final = await count_error_file_failed_requests(
error_file_content = await fetch_batch_error_file_content(
response,
custom_llm_provider=batch_file_provider,
litellm_params={
@ -856,6 +857,7 @@ class CheckBatchCost:
"_litellm_internal_model_credentials": MappingProxyType(dict(credentials)),
},
)
error_file_failed_requests: Final = count_error_file_failed_requests_from_content(error_file_content)
batch_result: Final = (
output_file_result
if not error_file_failed_requests
@ -895,6 +897,7 @@ class CheckBatchCost:
optional_params={},
custom_llm_provider=str(llm_provider) if llm_provider else None,
)
logging_obj._litellm_internal_model_credentials = MappingProxyType(dict(credentials)) # pyright: ignore[reportPrivateUsage] # trusted credentials transport, consumed by batch_line_item_logging
if not await self._claim_job_for_costing(job):
verbose_proxy_logger.info(
@ -913,6 +916,8 @@ class CheckBatchCost:
batch_failed_requests=batch_result.failed_requests,
batch_prompt_cost=batch_result.prompt_cost,
batch_completion_cost=batch_result.completion_cost,
batch_output_file_content=content_bytes if isinstance(content_bytes, bytes) else None,
batch_error_file_content=error_file_content,
)
except Exception:
await self._release_job_claim(job)

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,518 @@
import json
import uuid
from collections.abc import Iterator, Mapping
from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Literal, TypeAlias, cast, get_args
from typing_extensions import assert_never
from litellm._internal_context import with_service_target
from litellm._logging import verbose_logger
from litellm.batches.batch_utils import (
BatchResultFiles,
_batch_response_was_successful, # pyright: ignore[reportPrivateUsage] # batch-internal helper shared with the aggregate cost path by design
_fetch_batch_managed_file_content, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # same reuse; helper is untyped upstream
_get_response_from_batch_job_output_file, # pyright: ignore[reportPrivateUsage] # same reuse
_iter_batch_output_entries, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # same reuse; helper is untyped upstream
_safe_output_line_stats, # pyright: ignore[reportPrivateUsage] # same reuse
_uses_native_vertex_output, # pyright: ignore[reportPrivateUsage] # same reuse
)
from litellm.caching.caching import DualCache
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.utils import (
EmbeddingResponse,
LiteLLMBatch,
ModelInfo,
ModelResponse,
Usage,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging
_BatchLineProvider: TypeAlias = Literal[
"openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"
]
_SUPPORTED_LINE_PROVIDERS: Final = frozenset(get_args(_BatchLineProvider))
_SECRET_PARAM_KEYS: Final = frozenset(
{
"api_key",
"_litellm_internal_model_credentials",
"azure_ad_token",
"azure_ad_token_provider",
"vertex_credentials",
"aws_access_key_id",
"aws_secret_access_key",
"aws_session_token",
"aws_web_identity_token",
}
)
def _supported_line_provider(value: str) -> _BatchLineProvider | None:
if value in _SUPPORTED_LINE_PROVIDERS:
return cast("_BatchLineProvider", value) # cast-ok: membership in the literal's args was just checked
return None
_CALL_TYPE_BY_BATCH_URL: Final = MappingProxyType(
{
"/v1/chat/completions": "acompletion",
"/v1/embeddings": "aembedding",
"/v1/responses": "aresponses",
}
)
_EMPTY_BODY: Final[Mapping[str, object]] = MappingProxyType({})
_LINE_ITEM_CLAIM_TTL_SECONDS: Final = 30 * 24 * 60 * 60
batch_line_item_claim_cache: Final = DualCache()
_ClaimResult: TypeAlias = Literal["claimed", "already_claimed", "unavailable"]
@with_service_target("batch_line_items")
async def _claim_line_items(claim_cache: DualCache, claim_key: str, token: str) -> _ClaimResult:
redis_cache: Final = claim_cache.redis_cache
if redis_cache is None:
count: Final = await claim_cache.async_increment_cache(claim_key, 1, ttl=_LINE_ITEM_CLAIM_TTL_SECONDS) # pyright: ignore[reportUnknownMemberType] # DualCache.increment is untyped upstream
if count is None or count == 1:
return "claimed"
return "already_claimed"
try:
ok: Final = await redis_cache.async_set_cache(claim_key, token, ttl=_LINE_ITEM_CLAIM_TTL_SECONDS, nx=True) # pyright: ignore[reportUnknownMemberType] # RedisCache.set is untyped upstream
except Exception: # noqa: BLE001 # a redis outage must not emit duplicates; the next retrieve retries
return "unavailable"
if ok:
return "claimed"
return "already_claimed"
@with_service_target("batch_line_items")
async def _release_line_item_claim(claim_cache: DualCache, claim_key: str, token: str) -> None:
try:
redis_cache: Final = claim_cache.redis_cache
if redis_cache is None:
await claim_cache.async_delete_cache(claim_key)
return
owner: Final = await redis_cache.async_get_cache(claim_key) # pyright: ignore[reportUnknownMemberType] # RedisCache.get is untyped upstream
if owner == token:
await redis_cache.async_delete_cache(claim_key)
except Exception: # noqa: BLE001 # the claim release must never raise; worst case the batch stays claimed
verbose_logger.debug(
"batch line item claim release failed for %s, claim persists until ttl",
claim_key,
)
def _json_fallback(value: object) -> Mapping[str, object] | str:
mapping: Final = _as_object_mapping(value)
if mapping is None:
return str(value)
return dict(mapping)
class _BatchLineFailure(Exception):
"""A provider-reported per-line batch failure; carries the batch's hidden
params so the failure logging payload can attribute the line."""
def __init__(self, error_payload: object) -> None:
super().__init__(json.dumps(error_payload, default=_json_fallback))
self._hidden_params: dict[str, object] = {} # mutable-ok: plain-dict contract like response _hidden_params
def _as_object_mapping(value: object) -> Mapping[str, object] | None:
if isinstance(value, Mapping) and all(isinstance(key, str) for key in value): # pyright: ignore[reportUnknownVariableType] # keys of an unparameterized Mapping are unknown until checked here
return value # pyright: ignore[reportUnknownVariableType, reportReturnType] # every key was verified str above
return None
def _output_entries(file_content: bytes) -> Iterator[Mapping[str, object]]:
for entry in _iter_batch_output_entries(file_content): # pyright: ignore[reportUnknownVariableType] # entries are validated into typed mappings below
mapping = _as_object_mapping(entry) # pyright: ignore[reportUnknownArgumentType] # raw entry is unknown until validated here
if mapping is not None:
yield mapping
def _line_id(entry: Mapping[str, object]) -> str | None:
line_id: Final = entry.get("custom_id") or entry.get("recordId")
return line_id if isinstance(line_id, str) and line_id else None
def _requests_by_custom_id(input_file_content: bytes) -> Mapping[str, Mapping[str, object]]:
"""Parse the batch input JSONL into {line id: request line}, keyed by
custom_id or recordId, skipping malformed lines and lines without either."""
return MappingProxyType(
{line_id: entry for entry in _output_entries(input_file_content) if (line_id := _line_id(entry)) is not None}
)
def _request_body_for_entry(
entry: Mapping[str, object], request_line: Mapping[str, object] | None
) -> Mapping[str, object]:
if request_line is not None:
body: Final = _as_object_mapping(request_line.get("body"))
if body:
return body
params: Final = _as_object_mapping(request_line.get("params"))
if params:
return params
request_model_input: Final = _as_object_mapping(request_line.get("modelInput"))
if request_model_input:
return request_model_input
model_input: Final = _as_object_mapping(entry.get("modelInput"))
return model_input if model_input else _EMPTY_BODY
def _line_status_code(entry: Mapping[str, object], custom_llm_provider: str) -> int | None:
response: Final = _as_object_mapping(entry.get("response"))
status: Final = response.get("status_code") if response is not None else None
if isinstance(status, int):
return status
if custom_llm_provider == "anthropic":
result: Final = _as_object_mapping(entry.get("result"))
if result is not None and result.get("type") == "succeeded":
return 200
return None
def _line_error_payload(entry: Mapping[str, object], custom_llm_provider: _BatchLineProvider) -> object:
if custom_llm_provider == "anthropic":
return (
(_as_object_mapping(entry.get("result")) or _EMPTY_BODY).get("error") or entry.get("result") or _EMPTY_BODY
)
return entry.get("error") or entry.get("response") or _EMPTY_BODY
def _call_type_for_request(request_line: Mapping[str, object] | None) -> str:
url: Final = request_line.get("url") if request_line is not None else None
return _CALL_TYPE_BY_BATCH_URL.get(url if isinstance(url, str) else "", "acompletion")
def _call_type_for_line(request_line: Mapping[str, object] | None, result: "_BatchLineResult | None") -> str:
if isinstance(result, EmbeddingResponse):
return "aembedding"
if isinstance(result, ResponsesAPIResponse):
return "aresponses"
if isinstance(result, ModelResponse):
return "acompletion"
return _call_type_for_request(request_line)
def _line_messages(request_body: Mapping[str, object]) -> object:
return request_body.get("messages") or request_body.get("input") or ()
def _line_model(
response_body: Mapping[str, object],
request_body: Mapping[str, object],
parent: "Logging",
) -> str:
for candidate in (response_body.get("model"), request_body.get("model"), parent.model):
if isinstance(candidate, str) and candidate:
return candidate
return ""
_BatchLineResult: TypeAlias = "ModelResponse | EmbeddingResponse | ResponsesAPIResponse"
def _line_result(
call_type: str,
custom_llm_provider: _BatchLineProvider,
model: str,
response_body: Mapping[str, object],
) -> _BatchLineResult:
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
provider_config: Final = ProviderConfigManager.get_provider_batches_config(model, LlmProviders(custom_llm_provider))
if provider_config is not None:
transformed: Final = provider_config.transform_batch_output_line(response_body, model)
if transformed is not None:
return transformed
if custom_llm_provider == "bedrock":
raise ValueError(f"unrecognized bedrock batch output line shape. keys={sorted(response_body)}")
if call_type == "aembedding":
return EmbeddingResponse(**response_body) # pyright: ignore[reportArgumentType] # provider output bodies are dicts expanded as response ctor kwargs
if call_type == "aresponses":
return ResponsesAPIResponse(**response_body) # pyright: ignore[reportArgumentType] # same as above
if custom_llm_provider == "anthropic":
from litellm.llms.anthropic.chat.transformation import anthropic_message_to_model_response
return anthropic_message_to_model_response(response_body, None)
return ModelResponse(**response_body) # pyright: ignore[reportArgumentType] # same as above
def _line_result_or_none(
request_call_type: str,
custom_llm_provider: _BatchLineProvider,
model: str,
response_body: Mapping[str, object],
custom_id: object,
) -> "_BatchLineResult | None":
try:
return _line_result(request_call_type, custom_llm_provider, model, response_body)
except Exception: # noqa: BLE001 # one unparseable line must not drop the rest of the batch's line events
verbose_logger.warning(
"batch output line could not be reconstructed as a %s response, skipping it. custom_id=%s",
request_call_type,
custom_id,
)
return None
def _new_child_logging(
parent: "Logging",
model: str,
messages: object,
call_type: str,
start_time: datetime,
) -> "Logging":
from litellm.litellm_core_utils.litellm_logging import Logging
return Logging(
model=model,
messages=messages,
stream=False,
call_type=call_type,
start_time=start_time,
litellm_call_id=str(uuid.uuid4()),
function_id=str(uuid.uuid4()),
litellm_trace_id=parent.litellm_trace_id,
dynamic_success_callbacks=parent.dynamic_success_callbacks, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # Logging ctor takes these dynamic callback lists untyped
dynamic_async_success_callbacks=parent.dynamic_async_success_callbacks, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # same as above
dynamic_failure_callbacks=parent.dynamic_failure_callbacks, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # same as above
dynamic_async_failure_callbacks=parent.dynamic_async_failure_callbacks, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # same as above
kwargs={"litellm_session_id": parent.litellm_session_id},
)
def _line_hidden_params(
batch: LiteLLMBatch,
custom_id: object,
status_code: int | None,
response_cost: float | None = None,
) -> dict[str, object]: # mutable-ok: response objects declare _hidden_params as a plain dict
return {
"batch_id": batch.id,
"batch_custom_id": custom_id,
"batch_line_status_code": status_code,
"response_cost": response_cost,
}
def _optional_params_for_body(
request_body: Mapping[str, object],
) -> dict[str, object]: # mutable-ok: update_environment_variables takes a plain dict
return {key: value for key, value in request_body.items() if key not in ("model", "messages", "input")}
def _metadata_copy(params: Mapping[str, object]) -> dict[str, object]: # mutable-ok: dict out for litellm_params
metadata: Final = _as_object_mapping(params.get("metadata")) or _EMPTY_BODY
return {**metadata}
async def _emit_line_event(
entry: Mapping[str, object],
requests_by_id: Mapping[str, Mapping[str, object]],
batch: LiteLLMBatch,
custom_llm_provider: _BatchLineProvider,
parent: "Logging",
model_name: str | None,
model_info: ModelInfo | None,
) -> bool:
custom_id: Final = _line_id(entry)
request_line: Final = requests_by_id.get(custom_id or "")
request_body: Final = _request_body_for_entry(entry, request_line)
status_code: Final = _line_status_code(entry, custom_llm_provider)
request_call_type: Final = _call_type_for_request(request_line)
response_body: Final = _get_response_from_batch_job_output_file(entry, custom_llm_provider)
parent_start_time: Final = parent.start_time # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # Logging.start_time is untyped upstream
start_time: Final = parent_start_time if isinstance(parent_start_time, datetime) else datetime.now() # noqa: DTZ005 # naive to match the logging pipeline start_time
parent_params: Final = _as_object_mapping(parent.litellm_params) or _EMPTY_BODY # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # Logging.litellm_params is untyped upstream
model: Final = _line_model(response_body, request_body, parent)
successful: Final = _batch_response_was_successful(entry, custom_llm_provider)
stats: Final = _safe_output_line_stats(entry, custom_llm_provider, model_name, model_info) if successful else None
result: Final[_BatchLineResult | None] = (
_line_result_or_none(request_call_type, custom_llm_provider, model, response_body, custom_id)
if successful
else None
)
if successful and result is None:
return False
call_type: Final = _call_type_for_line(request_line, result)
child: Final = _new_child_logging(
parent=parent,
model=model,
messages=_line_messages(request_body),
call_type=call_type,
start_time=start_time,
)
child.update_environment_variables( # pyright: ignore[reportUnknownMemberType] # Logging.update_environment_variables is untyped upstream
litellm_params={
**parent_params,
"batch_parent_id": batch.id,
"metadata": _metadata_copy(parent_params),
},
optional_params=_optional_params_for_body(request_body),
model=child.model,
custom_llm_provider=custom_llm_provider,
)
for secret_key in _SECRET_PARAM_KEYS:
child.litellm_params.pop(secret_key, None) # pyright: ignore[reportUnknownMemberType] # Logging.litellm_params is untyped upstream
now: Final = datetime.now() # noqa: DTZ005 # naive to match the logging pipeline start_time
if result is None:
exception: Final = _BatchLineFailure(_line_error_payload(entry, custom_llm_provider))
exception._hidden_params = _line_hidden_params(batch, custom_id, status_code) # pyright: ignore[reportPrivateUsage] # _hidden_params is set on the exception instance itself
await child.async_failure_handler(
exception=exception,
traceback_exception="",
start_time=start_time,
end_time=now,
)
return True
result._hidden_params = _line_hidden_params( # pyright: ignore[reportPrivateUsage] # same hidden_params channel the aggregate batch event uses
batch,
custom_id,
status_code,
response_cost=stats.prompt_cost + stats.completion_cost if stats is not None else None,
)
if stats is not None and not response_body.get("usage") and isinstance(result, (ModelResponse, EmbeddingResponse)):
setattr( # noqa: B010 # ModelResponse.usage is set dynamically by its ctor, so setattr keeps parity for both response types
result,
"usage",
Usage(
prompt_tokens=stats.prompt_tokens,
completion_tokens=stats.completion_tokens,
total_tokens=stats.total_tokens,
),
)
await child.async_success_handler(
result=result,
start_time=start_time,
end_time=now,
cache_hit=False,
)
return True
async def _fetch_managed_file_or_empty(
file_id: str | None,
custom_llm_provider: _BatchLineProvider,
fetch_params: dict[str, object] | None, # mutable-ok: batch_utils file fetch takes the shared litellm_params dict
) -> bytes:
if file_id is None:
return b""
return await _fetch_batch_managed_file_content(
file_id,
custom_llm_provider=custom_llm_provider,
litellm_params=fetch_params, # pyright: ignore[reportArgumentType] # batch_utils types this param as an unparameterized dict
)
async def log_batch_line_items(
batch: LiteLLMBatch,
custom_llm_provider: str,
parent: "Logging",
model_name: str | None,
litellm_params: dict[str, object] | None, # mutable-ok: the logging object's shared litellm_params dict
model_info: ModelInfo | None,
result_files: BatchResultFiles | None = None,
claim_cache: DualCache = batch_line_item_claim_cache,
) -> int:
"""Emit one callback event per JSONL line of a completed batch (request
paired with its response/error), behind the opt-in
``litellm.store_batch_line_items_in_callbacks`` flag. The aggregate
aretrieve_batch event still bills the batch, so per-line events carry
``batch_parent_id`` and never update spend themselves. Any failure here
is logged and swallowed: aggregate accounting must be unaffected.
``result_files`` carries the output/error bytes the aggregate path already
fetched, so they are reused instead of refetched."""
line_provider: Final = _supported_line_provider(custom_llm_provider)
if line_provider is None:
verbose_logger.warning(
"batch line-item callbacks are not supported for provider %s, skipping. batch_id=%s",
custom_llm_provider,
batch.id,
)
return 0
claim_key: Final = f"batch_line_items_emitted:{batch.id}"
token: Final = uuid.uuid4().hex
claim: Final = await _claim_line_items(claim_cache, claim_key, token)
match claim:
case "already_claimed":
verbose_logger.debug("batch line items already emitted for batch_id=%s, skipping", batch.id)
return 0
case "unavailable":
verbose_logger.warning(
"batch line item claim backend unavailable for batch_id=%s, line items will be retried on the next retrieve",
batch.id,
)
return 0
case "claimed":
pass
case _:
assert_never(claim)
emitted = 0 # rebind-ok: loop accumulator for emitted line count
try:
internal_credentials: Final = parent._litellm_internal_model_credentials # pyright: ignore[reportPrivateUsage] # declared transport attribute on Logging
internal_mapping: Final = _as_object_mapping(internal_credentials)
fetch_params: Final[dict[str, object] | None] = ( # mutable-ok: file fetcher requires a plain dict
dict(internal_mapping) if internal_mapping is not None else litellm_params
)
input_file_content: Final = await _fetch_managed_file_or_empty(batch.input_file_id, line_provider, fetch_params)
requests_by_id: Final = _requests_by_custom_id(input_file_content)
output_content: Final = (
result_files.output
if result_files is not None and result_files.output is not None
else await _fetch_managed_file_or_empty(batch.output_file_id, line_provider, fetch_params)
)
first_row: Final = next(_output_entries(output_content), None)
if _uses_native_vertex_output(line_provider, model_name, first_row):
verbose_logger.warning(
"batch line-item callbacks do not support native vertex_ai batch output rows yet, skipping. batch_id=%s",
batch.id,
)
return 0
error_content: Final = (
result_files.error
if result_files is not None and result_files.error is not None
else await _fetch_managed_file_or_empty(batch.error_file_id, line_provider, fetch_params)
)
for content in (output_content, error_content):
for entry in _output_entries(content):
try:
emitted += await _emit_line_event(
entry=entry,
requests_by_id=requests_by_id,
batch=batch,
custom_llm_provider=line_provider,
parent=parent,
model_name=model_name,
model_info=model_info,
)
except Exception: # noqa: BLE001 # one bad line must not drop the rest of the batch's line events
verbose_logger.exception(
"batch line item logging failed for entry, continuing with remaining lines. batch_id=%s",
batch.id,
)
except Exception: # noqa: BLE001 # line-item logging must never break the aggregate aretrieve_batch accounting
verbose_logger.exception(
"batch line item logging failed for batch_id=%s; aggregate logging unaffected",
batch.id,
)
finally:
if emitted == 0:
await _release_line_item_claim(claim_cache, claim_key, token)
return emitted

View file

@ -34,6 +34,14 @@ class BatchCostUsageResult:
completion_cost: float = 0.0
@dataclass(frozen=True, slots=True)
class BatchResultFiles:
"""Raw JSONL bytes of a completed batch's result files; None means not fetched."""
output: bytes | None
error: bytes | None
_COMPLETED_BATCH_STATUSES: Final = frozenset({"completed", "complete"})
@ -71,7 +79,7 @@ def batch_cost_is_final(batch: Batch) -> bool:
async def calculate_batch_cost_and_usage(
file_content_dictionary: list[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
model_name: str | None = None,
model_info: ModelInfo | None = None,
) -> BatchCostUsageResult:
@ -98,7 +106,7 @@ async def calculate_batch_cost_and_usage(
async def _handle_completed_batch(
batch: Batch,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
model_name: str | None = None,
litellm_params: dict | None = None,
model_info: ModelInfo | None = None,
@ -116,29 +124,44 @@ async def _handle_completed_batch(
threaded through so a deployment's configured rates win over the
global cost map.
"""
# A completed batch whose request lines all failed has no output file - the
# results are written to a separate error_file_id and output_file_id is None.
# There is nothing to price or measure, so report an empty result set instead
# of calling _fetch_batch_output_file_content, which raises on a missing
# output file. Without this guard the logging worker crashes on every
# aretrieve_batch poll and the completed batch's zero-cost accounting is lost.
# The generic retrieval helper keeps raising for callers that explicitly ask
# for a missing output file.
return (
await _handle_completed_batch_with_files(
batch,
custom_llm_provider,
model_name=model_name,
litellm_params=litellm_params,
model_info=model_info,
)
)[0]
async def _handle_completed_batch_with_files(
batch: Batch,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
model_name: str | None = None,
litellm_params: dict[str, object] | None = None, # mutable-ok: file fetchers require a plain dict
model_info: ModelInfo | None = None,
) -> "tuple[BatchCostUsageResult, BatchResultFiles]":
"""_handle_completed_batch plus the raw output/error bytes it fetched, so
downstream line-item logging reuses them instead of refetching."""
error_file_content: Final = await fetch_batch_error_file_content(
batch, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params
)
error_file_failed_requests: Final = count_error_file_failed_requests_from_content(error_file_content)
if batch.output_file_id is None:
return BatchCostUsageResult(
cost=0.0,
usage=Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
models=[],
successful_requests=0,
failed_requests=await count_error_file_failed_requests(
batch, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params
return (
BatchCostUsageResult(
cost=0.0,
usage=Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0),
models=[],
successful_requests=0,
failed_requests=error_file_failed_requests,
),
BatchResultFiles(output=None, error=error_file_content),
)
file_content = await _fetch_batch_output_file_content(batch, custom_llm_provider, litellm_params=litellm_params)
error_file_failed_requests: Final = await count_error_file_failed_requests(
batch, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params
)
output_file_result: Final = (
calculate_vertex_ai_batch_cost_and_usage(
@ -155,11 +178,14 @@ async def _handle_completed_batch(
)
)
if not error_file_failed_requests:
return output_file_result
return dataclasses_replace(
output_file_result, failed_requests=output_file_result.failed_requests + error_file_failed_requests
result: Final = (
output_file_result
if not error_file_failed_requests
else dataclasses_replace(
output_file_result, failed_requests=output_file_result.failed_requests + error_file_failed_requests
)
)
return result, BatchResultFiles(output=file_content, error=error_file_content)
class _LineOutcome(Enum):
@ -435,7 +461,9 @@ def _provider_output_file_id(output_file_id: str) -> str:
async def _fetch_batch_managed_file_content(
file_id: str,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"] = "openai",
custom_llm_provider: Literal[
"openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"
] = "openai",
litellm_params: dict | None = None,
) -> bytes:
"""
@ -465,7 +493,9 @@ async def _fetch_batch_managed_file_content(
async def _fetch_batch_output_file_content(
batch: Batch,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"] = "openai",
custom_llm_provider: Literal[
"openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"
] = "openai",
litellm_params: dict | None = None,
) -> bytes:
"""
@ -485,9 +515,30 @@ async def _fetch_batch_output_file_content(
)
async def fetch_batch_error_file_content(
batch: Batch,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
litellm_params: dict[str, object] | None, # mutable-ok: file fetchers require a plain dict
) -> bytes | None:
"""Fetch the batch's separate error file bytes; None when it has none or the fetch fails."""
if batch.error_file_id is None:
return None
try:
return await _fetch_batch_managed_file_content(
batch.error_file_id, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params
)
except Exception as e: # noqa: BLE001 # a failed/missing error file must not abort cost tracking for the batch
verbose_logger.debug("Failed to fetch batch error file %s: %s", batch.error_file_id, e)
return None
def count_error_file_failed_requests_from_content(content: bytes | None) -> int:
return 0 if content is None else sum(1 for _ in _iter_batch_input_lines(content))
async def count_error_file_failed_requests(
batch: Batch,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "mistral"],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock", "mistral"],
litellm_params: dict | None,
) -> int:
"""Count failed requests reported only in the batch's separate error file.
@ -497,16 +548,11 @@ async def count_error_file_failed_requests(
``error_file_id`` - they never appear in the output file at all, so
counting failures from the output file alone silently undercounts them.
"""
if batch.error_file_id is None:
return 0
try:
error_file_content = await _fetch_batch_managed_file_content(
batch.error_file_id, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params
return count_error_file_failed_requests_from_content(
await fetch_batch_error_file_content(
batch, custom_llm_provider=custom_llm_provider, litellm_params=litellm_params
)
except Exception as e: # noqa: BLE001 # a failed/missing error file must not abort cost tracking for the batch
verbose_logger.debug("Failed to fetch batch error file %s: %s", batch.error_file_id, e)
return 0
return sum(1 for _ in _iter_batch_input_lines(error_file_content))
)
def _extract_file_access_credentials(litellm_params: dict | None) -> dict:

View file

@ -19,6 +19,7 @@ from litellm._logging import verbose_logger
from litellm.integrations.clickhouse.clickhouse_batch_logger import ClickHouseBatchLogger
from litellm.integrations.clickhouse.context import is_lens_analysis
from litellm.integrations.clickhouse.schema import SPEND_LOGS_TABLE
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.litellm_core_utils.sensitive_data_masker import redact_credentials_in_payload
from litellm.tracing.types import SpendLogRecord
from litellm.types.utils import StandardLoggingPayload
@ -197,9 +198,17 @@ class ClickHouseSpendLogger(ClickHouseBatchLogger):
table = SPEND_LOGS_TABLE
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None:
# Batch line items are billed by the aggregate aretrieve_batch row; a per-line
# spend row here would bill the batch twice.
if is_batch_line_item_event(kwargs):
return
self._log(kwargs)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None:
# Failed batch lines are part of the same aggregate row's request counts;
# per-line rows here would inflate spend-log request counts.
if is_batch_line_item_event(kwargs):
return
self._log(kwargs)
def _log(self, kwargs: Mapping[str, Any]) -> None:

View file

@ -14,6 +14,7 @@ from litellm.integrations.datadog.datadog_handler import (
get_datadog_service,
normalize_datadog_tag_value,
)
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
@ -72,6 +73,10 @@ class DatadogCostManagementLogger(CustomBatchLogger):
super().__init__(**kwargs)
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
# A batch line item is billed by the aggregate aretrieve_batch event; a
# per-line FOCUS BilledCost row would double-count cloud spend.
if is_batch_line_item_event(kwargs):
return
try:
standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None)

View file

@ -11,6 +11,7 @@ import litellm
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.llms.custom_httpx.http_handler import (
HTTPHandler,
get_async_httpx_client,
@ -121,6 +122,9 @@ class LagoLogger(CustomLogger):
return returned_val
def log_success_event(self, kwargs, response_obj, start_time, end_time):
# A batch line item is billed by the aggregate aretrieve_batch event.
if is_batch_line_item_event(kwargs):
return
_url = os.getenv("LAGO_API_BASE")
assert _url is not None and isinstance(_url, str), (
f"LAGO_API_BASE missing or not set correctly. LAGO_API_BASE={_url}"
@ -153,6 +157,8 @@ class LagoLogger(CustomLogger):
raise e
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
if is_batch_line_item_event(kwargs):
return
try:
verbose_logger.debug("ENTERS LAGO CALLBACK")
_url = os.getenv("LAGO_API_BASE")

View file

@ -34,6 +34,7 @@ from httpx import HTTPStatusError, Response
from litellm._logging import verbose_logger
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
@ -301,12 +302,17 @@ class NewRelicMetricsLogger(CustomBatchLogger):
await self._final_drain()
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None:
# A batch line item is metered by the aggregate aretrieve_batch event.
if is_batch_line_item_event(kwargs):
return
try:
await self._log_async_event(standard_logging_object=kwargs.get("standard_logging_object", None))
except Exception as e: # noqa: BLE001 # logging must never break the request path
verbose_logger.exception("New Relic Metrics Layer Error - %s\n%s", e, traceback.format_exc())
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None:
if is_batch_line_item_event(kwargs):
return
try:
await self._log_async_event(standard_logging_object=kwargs.get("standard_logging_object", None))
except Exception as e: # noqa: BLE001 # logging must never break the request path

View file

@ -9,6 +9,7 @@ import httpx
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.llms.custom_httpx.http_handler import (
HTTPHandler,
get_async_httpx_client,
@ -97,6 +98,9 @@ class OpenMeterLogger(CustomLogger):
}
def log_success_event(self, kwargs, response_obj, start_time, end_time):
# A batch line item is billed by the aggregate aretrieve_batch event.
if is_batch_line_item_event(kwargs):
return
_url = os.getenv("OPENMETER_API_ENDPOINT", "https://openmeter.cloud")
if _url.endswith("/"):
_url += "api/v1/events"
@ -123,6 +127,8 @@ class OpenMeterLogger(CustomLogger):
raise e
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
if is_batch_line_item_event(kwargs):
return
_url = os.getenv("OPENMETER_API_ENDPOINT", "https://openmeter.cloud")
if _url.endswith("/"):
_url += "api/v1/events"

View file

@ -30,6 +30,7 @@ from litellm.integrations.otel.model.db_endpoint import db_span_attributes
from litellm.integrations.otel.model.metadata import flatten_metadata
from litellm.integrations.otel.model.semconv import LiteLLM, Metric
from litellm.integrations.otel.plumbing.otlp_tls import resolve_otlp_http_tls
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.litellm_core_utils.internal_call_metadata import is_unbilled_non_inference_call_from_params
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.secret_redaction import redact_string
@ -1659,6 +1660,11 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
return True
def _record_metrics(self, kwargs, response_obj, start_time, end_time):
# A batch line item's tokens and cost are metered by the aggregate
# aretrieve_batch event; per-line samples here would double-count them.
# Spans for line items are still emitted by _handle_success.
if is_batch_line_item_event(kwargs):
return
duration_s: Final = (end_time - start_time).total_seconds()
params: Final = kwargs.get("litellm_params") or {}
provider: Final = _provider_label(params.get("custom_llm_provider"))

View file

@ -33,6 +33,7 @@ from litellm.integrations.otel.model.semconv import (
resolve_provider,
)
from litellm.integrations.otel.model.utils import to_seconds
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.litellm_core_utils.internal_call_metadata import is_unbilled_non_inference_call_from_params
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
@ -221,6 +222,12 @@ class GenAIMetricRecorder:
start_time: datetime,
end_time: datetime,
) -> None:
# A batch line item is metered by the aggregate aretrieve_batch event, and
# its child interval (parent retrieve start -> line emission) measures
# retrieval and callback processing rather than the line's model call, so
# skip every metric for it; spans for line items are still emitted.
if is_batch_line_item_event(kwargs):
return
common_attrs: Final = self._filter_attributes(self._bounded_attributes(kwargs))
duration_s: Final = (end_time - start_time).total_seconds()
usage_is_replayed: Final = is_unbilled_non_inference_call_from_params(
@ -246,6 +253,10 @@ class GenAIMetricRecorder:
start_time: datetime,
end_time: datetime,
) -> None:
# A batch line item's interval measures retrieval and emission, not the
# line's model call, so its duration would be synthetic here too.
if is_batch_line_item_event(kwargs):
return
"""Record the one metric a failed request can honestly report: the
operation's duration, tagged with ``error.type``.

View file

@ -35,6 +35,7 @@ from litellm.integrations.prometheus_helpers.bounded_prometheus_series_tracker i
from litellm.litellm_core_utils.core_helpers import (
get_litellm_metadata_from_kwargs,
get_metadata_variable_name_from_kwargs,
is_batch_line_item_event,
)
from litellm.litellm_core_utils.service_tier_utils import (
get_service_tier_from_standard_logging_payload,
@ -1345,6 +1346,10 @@ class PrometheusLogger(CustomLogger):
self._track_end_user_metric_series(counter, metric_name, _labels)
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
# A batch line item is metered by the aggregate aretrieve_batch event; a
# per-line sample here would double-count spend and requests.
if is_batch_line_item_event(kwargs):
return
# Define prometheus client
verbose_logger.debug(
"prometheus Logging - Enters success logging function (kwargs keys: %s)",
@ -2365,6 +2370,8 @@ class PrometheusLogger(CustomLogger):
)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
if is_batch_line_item_event(kwargs):
return
verbose_logger.debug(
"prometheus Logging - Enters failure logging function (kwargs keys: %s)",
list(kwargs.keys()) if isinstance(kwargs, dict) else type(kwargs).__name__,

View file

@ -18,9 +18,9 @@ if TYPE_CHECKING:
from litellm.types.utils import ModelResponseStream
Span = _Span | Any
Span = _Span | object
else:
Span = Any
Span = object
_CODEX_CLIENT_PREFIX_RE: Final = re.compile(r"^codex[-_ /]", re.IGNORECASE)
@ -270,7 +270,7 @@ def remove_index_from_tool_calls(
tool_call.pop("index", None)
def remove_items_at_indices(items: list[Any] | None, indices: Iterable[int]) -> None:
def remove_items_at_indices(items: list[object] | None, indices: Iterable[int]) -> None:
"""Remove items from a list in-place by index"""
if items is None:
return
@ -713,9 +713,9 @@ def filter_internal_params(data: dict, additional_internal_params: set | None =
def redact_nested_match_and_regex_keys(
payload: dict | list[Any] | str | None,
payload: dict | list[object] | str | None,
keys: Collection[str] = ("match", "regex"),
) -> dict | list[Any] | str | None:
) -> dict | list[object] | str | None:
"""
Deep-copy `payload` and replace every configured string field with "[REDACTED]"
anywhere in nested dict/list structures.
@ -725,7 +725,7 @@ def redact_nested_match_and_regex_keys(
if payload is None or isinstance(payload, str):
return payload
try:
redacted: Final[dict | list[Any] | str | None] = copy.deepcopy(payload)
redacted: Final[dict | list[object] | str | None] = copy.deepcopy(payload)
except Exception:
return payload
@ -772,6 +772,15 @@ def set_response_cost_in_hidden_params(response: _CarriesHiddenParams, cost: flo
hidden_params["additional_headers"] = merged
def is_batch_line_item_event(kwargs: object) -> bool:
if not isinstance(kwargs, Mapping):
return False
litellm_params: Final[object] = kwargs.get("litellm_params")
if not isinstance(litellm_params, Mapping):
return False
return bool(litellm_params.get("batch_parent_id"))
_HIDDEN_PARAMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
_PROVIDER_HEADERS_ADAPTER: Final = TypeAdapter(Mapping[str, str])

View file

@ -32,7 +32,11 @@ from litellm._logging import (
verbose_logger,
)
from litellm._uuid import uuid
from litellm.batches.batch_utils import _handle_completed_batch, batch_cost_is_final
from litellm.batches.batch_utils import (
BatchResultFiles,
_handle_completed_batch_with_files,
batch_cost_is_final,
)
from litellm.caching.caching import DualCache
from litellm.caching.caching_handler import LLMCachingHandler
from litellm.caching.redis_batch import flush_post_call_redis_batches
@ -652,6 +656,7 @@ class Logging(LiteLLMLoggingBaseClass):
self.streaming_chunks: list[object] = []
self.sync_streaming_chunks: list[object] = []
self.log_raw_request_response = log_raw_request_response
self._litellm_internal_model_credentials: Mapping[str, object] | None = None
self.raw_request_only = raw_request_only
# Initialize dynamic callbacks
@ -947,6 +952,10 @@ class Logging(LiteLLMLoggingBaseClass):
**additional_params,
}
)
# Provider params / additional kwargs are caller-influenced; they must not
# overwrite the internal litellm_params, or a request body carrying
# `litellm_params` could forge logging markers (e.g. batch_parent_id).
self.model_call_details["litellm_params"] = self.litellm_params
## check if stream options is set ## - used by CustomStreamWrapper for easy instrumentation
if "stream_options" in additional_params:
@ -3249,9 +3258,12 @@ class Logging(LiteLLMLoggingBaseClass):
batch_models = kwargs.get("batch_models", None)
batch_successful_requests: Final = kwargs.get("batch_successful_requests", None)
batch_failed_requests: Final = kwargs.get("batch_failed_requests", None)
batch_output_file_content: Final = kwargs.get("batch_output_file_content", None)
batch_error_file_content: Final = kwargs.get("batch_error_file_content", None)
has_explicit_batch_data: Final = all(x is not None for x in (batch_cost, batch_usage, batch_models))
should_compute_batch_data: Final = not has_explicit_batch_data and batch_cost_is_final(result)
result_files: BatchResultFiles | None = None # rebind-ok: one batch-data branch below supplies the bytes
if has_explicit_batch_data:
result._hidden_params["response_cost"] = batch_cost
result._hidden_params["batch_models"] = batch_models
@ -3271,15 +3283,21 @@ class Logging(LiteLLMLoggingBaseClass):
total_cost=batch_cost,
cost_for_built_in_tools_cost_usd_dollar=0.0,
)
if batch_output_file_content is not None or batch_error_file_content is not None:
result_files = BatchResultFiles(
output=batch_output_file_content if isinstance(batch_output_file_content, bytes) else None,
error=batch_error_file_content if isinstance(batch_error_file_content, bytes) else None,
)
elif should_compute_batch_data:
batch_result: Final = await _handle_completed_batch(
batch_result, fetched_result_files = await _handle_completed_batch_with_files(
batch=result,
custom_llm_provider=self.custom_llm_provider,
model_name=self.get_deployment_model_for_cost(),
litellm_params=self.litellm_params,
model_info=self.get_router_deployment_model_info(),
)
result_files = fetched_result_files
result._hidden_params["response_cost"] = batch_result.cost
result._hidden_params["batch_models"] = batch_result.models
@ -3293,6 +3311,28 @@ class Logging(LiteLLMLoggingBaseClass):
cost_for_built_in_tools_cost_usd_dollar=0.0,
)
if litellm.store_batch_line_items_in_callbacks and (has_explicit_batch_data or should_compute_batch_data):
from litellm.batches.batch_line_item_logging import log_batch_line_items
_line_item_litellm_params: Final = cast( # cast-ok: shared attribute is an untyped dict
"dict[str, object] | None", self.litellm_params
)
try:
await log_batch_line_items(
batch=result,
custom_llm_provider=self.custom_llm_provider,
parent=self,
model_name=self.get_deployment_model_for_cost(),
litellm_params=_line_item_litellm_params,
model_info=self.get_router_deployment_model_info(),
result_files=result_files,
)
except Exception: # noqa: BLE001 # line-item logging (claim step included) must never reach the aggregate aretrieve_batch path
verbose_logger.exception(
"batch line item logging failed for batch_id=%s; aggregate logging unaffected",
result.id,
)
self.truncated_messages_for_logging = await truncate_base64_in_messages_async(
StandardLoggingPayloadSetup.append_system_prompt_messages(
kwargs=self.model_call_details, messages=self.model_call_details.get("messages")
@ -4247,19 +4287,10 @@ class Logging(LiteLLMLoggingBaseClass):
verbose_logger.warning("LiteLLM: the anthropic_messages stream assembled no response, logging an empty one")
return litellm.ModelResponse(model=self.model)
else:
from litellm.types.llms.anthropic import AnthropicResponse
from litellm.llms.anthropic.chat.transformation import anthropic_message_to_model_response
pydantic_result: Final = AnthropicResponse.model_validate(result)
import httpx
result = litellm.AnthropicConfig().transform_parsed_response(
completion_response=pydantic_result.model_dump(),
raw_response=httpx.Response(
status_code=200,
headers={},
),
model_response=litellm.ModelResponse(id=provider_response_id),
json_mode=None,
result = anthropic_message_to_model_response(
cast(Mapping[str, object], result), # cast-ok: handler result is typed Any upstream
speed=self.optional_params.get("speed") if self.optional_params else None,
)
return result
@ -6041,6 +6072,9 @@ class StandardLoggingPayloadSetup:
batch_failed_requests=None,
litellm_model_name=None,
usage_object=None,
batch_id=None,
batch_custom_id=None,
batch_line_status_code=None,
)
if hidden_params is not None:
for key in StandardLoggingHiddenParams.__annotations__:
@ -6446,8 +6480,12 @@ def _extract_response_obj_and_hidden_params(
response_obj = {}
if original_exception is not None and hidden_params is None:
response_headers: Final = _get_response_headers(original_exception)
if response_headers is not None:
exception_hidden_params: Final[Mapping[str, object] | None] = getattr(
original_exception, "_hidden_params", None
)
if isinstance(exception_hidden_params, dict) and exception_hidden_params:
hidden_params = dict(exception_hidden_params)
elif (response_headers := _get_response_headers(original_exception)) is not None:
hidden_params = dict(
StandardLoggingHiddenParams(
additional_headers=StandardLoggingPayloadSetup.get_additional_headers(dict(response_headers)),

View file

@ -52,6 +52,7 @@ from litellm.types.llms.anthropic import (
AnthropicMessagesToolChoice,
AnthropicOutputSchema,
AnthropicOutputTokensDetails,
AnthropicResponse,
AnthropicSystemMessageContent,
AnthropicThinkingParam,
AnthropicWebSearchTool,
@ -2808,3 +2809,18 @@ def _valid_user_id(user_id: str) -> bool:
return False
return True
def anthropic_message_to_model_response(result: Mapping[str, object], speed: str | None) -> ModelResponse:
pydantic_result: Final = AnthropicResponse.model_validate(result)
result_id: Final = result.get("id")
return AnthropicConfig().transform_parsed_response(
completion_response=pydantic_result.model_dump(),
raw_response=httpx.Response(
status_code=200,
headers={},
),
model_response=ModelResponse(id=result_id if isinstance(result_id, str) and result_id else None),
json_mode=None,
speed=speed,
)

View file

@ -1,5 +1,6 @@
import types
from abc import ABC, abstractmethod
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any
import httpx
@ -9,7 +10,7 @@ from litellm.types.llms.openai import (
AllMessageValues,
CreateBatchRequest,
)
from litellm.types.utils import LiteLLMBatch, LlmProviders
from litellm.types.utils import EmbeddingResponse, LiteLLMBatch, LlmProviders, ModelResponse
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
@ -39,6 +40,14 @@ class BaseBatchesConfig(ABC):
def custom_llm_provider(self) -> LlmProviders:
"""Return the LLM provider type for this configuration."""
def transform_batch_output_line(
self, model_output: Mapping[str, object], model: str
) -> ModelResponse | EmbeddingResponse | None:
"""Reconstruct one provider batch output line into a litellm response, or
None when the line is OpenAI-shaped and the caller can use the generic
reconstruction."""
return None
@classmethod
def get_config(cls):
"""Get configuration dictionary for this class."""

View file

@ -27,7 +27,7 @@ from litellm.types.llms.openai import (
AllMessageValues,
CreateBatchRequest,
)
from litellm.types.utils import LiteLLMBatch, LlmProviders, Usage
from litellm.types.utils import EmbeddingResponse, LiteLLMBatch, LlmProviders, ModelResponse, Usage
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import (
@ -96,6 +96,47 @@ def titan_embedding_usage_from_batch_output(model_output: Mapping[str, object])
)
def bedrock_batch_line_to_response(
model_output: Mapping[str, object], model: str
) -> ModelResponse | EmbeddingResponse | None:
"""Reconstruct a Bedrock batch output line (the ``modelOutput`` object) into
the litellm response type its shape implies, or None when the shape is
unrecognized."""
if "embedding" in model_output:
embedding: Final = model_output.get("embedding")
return EmbeddingResponse(
model=model,
data=[
{
"object": "embedding",
"index": 0,
"embedding": embedding if isinstance(embedding, list) else [],
}
],
usage=titan_embedding_usage_from_batch_output(model_output),
)
if "output" in model_output:
from ..chat.converse_transformation import AmazonConverseConfig
return AmazonConverseConfig()._transform_response( # pyright: ignore[reportPrivateUsage] # same reconstruction the converse chat path performs on the live response
model=model,
response=Response(200, json=dict(model_output)),
model_response=ModelResponse(),
stream=False,
logging_obj=None,
optional_params={},
api_key=None,
data="",
messages=[],
encoding=None,
)
if "content" in model_output:
from litellm.llms.anthropic.chat.transformation import anthropic_message_to_model_response
return anthropic_message_to_model_response(model_output, None)
return None
class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
"""
Config for Bedrock Batches - handles batch job creation and management for Bedrock
@ -109,6 +150,11 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
def custom_llm_provider(self) -> LlmProviders:
return LlmProviders.BEDROCK
def transform_batch_output_line(
self, model_output: Mapping[str, object], model: str
) -> ModelResponse | EmbeddingResponse | None:
return bedrock_batch_line_to_response(model_output, model)
@classmethod
def _get_bare_model_name_from_s3_key(cls, object_key: str) -> str | None:
if not object_key.startswith(BEDROCK_MANAGED_S3_BATCH_PREFIX):

View file

@ -3037,6 +3037,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
None,
description="If True, stores request messages and responses in spend logs. Default is False.",
)
store_batch_line_items_in_callbacks: bool | None = Field(
None,
description="If True, a completed batch logged via aretrieve_batch also emits one callback event per JSONL line item (request paired with its response or error). The aggregate batch callback is unchanged. Default is False.",
)
disable_auto_add_proxy_admin_to_teams: bool | None = Field(
None,
description="By default, the user calling /team/new is automatically added to the new team as a team admin. If True, proxy admins are no longer auto-added; members explicitly listed in members_with_roles are unaffected. Default is False.",

View file

@ -15,6 +15,7 @@ from litellm._internal_context import with_service_target
from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.proxy_rate_limit_error import (
ProxyRateLimitError,
@ -697,6 +698,8 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
- model_saturation_check: Model-wide token tracking
- priority_model: Priority-specific token tracking
"""
if is_batch_line_item_event(kwargs):
return
from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
)

View file

@ -24,6 +24,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.caching.redis_cache import log_redis_failure
from litellm.exceptions import RateLimitType
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit
@ -134,6 +135,8 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
"""
After a successful LLM call, increment the session spend by the response cost.
"""
if is_batch_line_item_event(kwargs): # pyright: ignore[reportUnknownArgumentType] # hook kwargs arrive untyped from the logging dispatcher
return
try:
litellm_params: Final = kwargs.get("litellm_params") or {}
metadata: Final = litellm_params.get("metadata") or {}

View file

@ -13,6 +13,7 @@ from litellm._internal_context import with_service_target
from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import Span
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.llms.bedrock.common_utils import get_bedrock_base_model
from litellm.proxy._types import Litellm_EntityType, UserAPIKeyAuth
@ -499,6 +500,8 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
Example: key=sk-1234567890, model=gpt-4o, max_budget=100, time_period=1d
"""
if is_batch_line_item_event(kwargs):
return
verbose_proxy_logger.debug("in RouterBudgetLimiting.async_log_success_event")
standard_logging_payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None)
if standard_logging_payload is None:

View file

@ -12,7 +12,10 @@ from litellm._internal_context import with_service_target
from litellm._logging import verbose_proxy_logger
from litellm.exceptions import RateLimitType
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
is_batch_line_item_event,
)
from litellm.proxy._types import CommonProxyErrors, CurrentItemRateLimit, UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import (
get_key_model_rpm_limit,
@ -495,6 +498,8 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
@with_service_target("rate_limits")
async def async_log_success_event(self, kwargs, response_obj: object, start_time, end_time):
if is_batch_line_item_event(kwargs):
return
from litellm.proxy.common_utils.callback_utils import (
get_model_group_from_litellm_kwargs,
)
@ -701,6 +706,8 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
@with_service_target("rate_limits")
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
if is_batch_line_item_event(kwargs):
return
try:
self.print_verbose("Inside Max Parallel Request Failure Hook")
litellm_parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(kwargs=kwargs)

View file

@ -43,6 +43,7 @@ from litellm.caching.redis_batch import (
from litellm.caching.redis_cache import log_redis_failure
from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE, INTERNAL_CALL_ORIGIN_METADATA_KEY
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_str_from_messages,
)
@ -5001,6 +5002,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
"""
Update TPM usage on successful API calls by incrementing counters using pipeline
"""
if is_batch_line_item_event(kwargs):
return
from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
)
@ -5124,6 +5127,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
whose partial usage was recovered settles the reservation at that
usage instead of refunding it.
"""
if is_batch_line_item_event(kwargs):
return
from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
)

View file

@ -14,6 +14,7 @@ from litellm.litellm_core_utils.core_helpers import (
budget_reservation_from_metadata,
get_litellm_metadata_from_kwargs,
get_metadata_variable_name_from_kwargs,
is_batch_line_item_event,
)
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import guardrail_information_cost
@ -121,6 +122,8 @@ class _ProxyDBLogger(CustomLogger):
async def async_log_success_event(
self, kwargs: ObjectMapping, response_obj: object, start_time: datetime, end_time: datetime
) -> None:
if is_batch_line_item_event(kwargs):
return
if self.spend_event_producer is None or not is_offloadable_success(response_obj):
await self._PROXY_track_cost_callback(kwargs, response_obj, start_time, end_time)
return

View file

@ -270,6 +270,10 @@ import litellm._redis
from litellm import Router
from litellm._internal_context import service_target, with_service_target
from litellm._logging import _redact_string, verbose_proxy_logger, verbose_router_logger
from litellm.batches.batch_line_item_logging import (
_LINE_ITEM_CLAIM_TTL_SECONDS, # pyright: ignore[reportPrivateUsage] # the claim window is defined next to the cache it expires
batch_line_item_claim_cache,
)
from litellm.caching.caching import DualCache, RedisCache
from litellm.caching.dual_cache import DeclaredBatchRead
from litellm.caching.redis_batch import (
@ -4982,8 +4986,8 @@ def _attach_redis_usage_cache(redis_cache: RedisCache, enable_redis_auth_cache:
"""
Wires an established coordination Redis into the proxy-level caches that
consume it directly: the spend counter cache, the CLI SSO login-session
cache, the cluster-wide config cache, and (only when opted in) the
virtual-key auth cache.
cache, the batch line-item claim cache, the cluster-wide config cache,
and (only when opted in) the virtual-key auth cache.
The CLI SSO login-session cache is always backed by Redis when available so
that the browser SSO flow behind `lite login` survives landing on different
@ -4997,6 +5001,10 @@ def _attach_redis_usage_cache(redis_cache: RedisCache, enable_redis_auth_cache:
redis_cache,
default_redis_ttl=CLI_SSO_SESSION_TTL_SECONDS,
)
batch_line_item_claim_cache.attach_redis_cache(
redis_cache,
default_redis_ttl=_LINE_ITEM_CLAIM_TTL_SECONDS,
)
if enable_redis_auth_cache is True:
user_api_key_cache.attach_redis_cache(
redis_cache,
@ -6910,6 +6918,13 @@ class ProxyConfig:
health_check_interval = general_settings.get("health_check_interval", DEFAULT_HEALTH_CHECK_INTERVAL)
health_check_concurrency = general_settings.get("health_check_concurrency", None)
health_check_details = general_settings.get("health_check_details", True)
### BATCH LINE ITEM CALLBACKS ###
_store_batch_line_items: Final[object] = general_settings.get("store_batch_line_items_in_callbacks")
if _store_batch_line_items is not None:
if isinstance(_store_batch_line_items, str):
litellm.store_batch_line_items_in_callbacks = _store_batch_line_items.lower() == "true"
else:
litellm.store_batch_line_items_in_callbacks = bool(_store_batch_line_items)
# Health-check-driven routing (opt-in, passes through to Router later)
_enable_hc_routing = general_settings.get("enable_health_check_routing", False)
_hc_staleness = general_settings.get("health_check_staleness_threshold", None)
@ -7921,6 +7936,7 @@ class ProxyConfig:
self._apply_alerting_settings,
self._apply_pass_through_settings,
self._apply_boolean_settings,
self._apply_batch_line_items_setting,
partial(self._apply_cache_size_setting, cache_size_was_db=cache_size_was_db),
self._apply_store_model_in_db_setting,
partial(self._apply_retention_settings, previous_cleanup_schedule=previous_cleanup_schedule),
@ -7980,6 +7996,11 @@ class ProxyConfig:
if (value := self.settings.get(key)) is not None:
self.settings[key] = coerce_bool(value)
async def _apply_batch_line_items_setting(self, db_values: Mapping[str, SettingsJsonValue]) -> None:
key: Final = "store_batch_line_items_in_callbacks"
value: Final = coerce_bool(self.settings.get(key))
litellm.store_batch_line_items_in_callbacks = bool(value) if value is not None else False
async def _apply_cache_size_setting(
self,
db_values: Mapping[str, SettingsJsonValue],

View file

@ -80,6 +80,7 @@ from litellm.litellm_core_utils.core_helpers import (
get_litellm_metadata_from_kwargs,
get_metadata_variable_name_from_kwargs,
get_or_create_metadata_bucket,
is_batch_line_item_event,
)
from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
@ -8089,7 +8090,7 @@ class Router:
# WS session wrappers fire with result=None; per-turn costs tracked by inner calls.
if kwargs.get("call_type") in ("_aresponses_websocket", "_arealtime"):
return
if is_batch_retrieve_call_type(kwargs.get("call_type")):
if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs):
return
standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None)
if standard_logging_object is None:
@ -8237,7 +8238,7 @@ class Router:
- key: str - The key used to increment the cache
- None: if no key is found
"""
if is_batch_retrieve_call_type(kwargs.get("call_type")):
if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs):
return None
id = None
if kwargs["litellm_params"].get("metadata") is None:
@ -8367,7 +8368,7 @@ class Router:
"""
Update RPM usage for a deployment
"""
if is_batch_retrieve_call_type(kwargs.get("call_type")):
if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs):
return
deployment_name: Final = kwargs["litellm_params"]["metadata"].get(
"deployment", None

View file

@ -35,6 +35,7 @@ from litellm.caching.redis_cache import RedisCache, RedisPipelineIncrementOperat
from litellm.integrations.custom_logger import CustomLogger, Span
from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
is_batch_line_item_event,
)
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs
@ -495,6 +496,8 @@ class RouterBudgetLimiting(CustomLogger):
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
"""Original method now uses helper functions"""
verbose_router_logger.debug("in RouterBudgetLimiting.async_log_success_event")
if is_batch_line_item_event(kwargs):
return
# WS session wrappers fire with result=None; per-turn costs tracked by inner calls.
if kwargs.get("call_type") in ("_aresponses_websocket", "_arealtime"):
return

View file

@ -10,6 +10,7 @@ from litellm._logging import verbose_router_logger
from litellm.caching.caching import DualCache
from litellm.caching.redis_cache import log_redis_failure
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.router_utils.batch_utils import is_batch_retrieve_call_type
IN_FLIGHT_COUNT_TTL_SECONDS: Final = 60 * 60
@ -50,7 +51,7 @@ def _request_count_key(model_group: str, deployment_id: str) -> str:
def _deployment_ref(kwargs: Mapping[str, object]) -> tuple[str, str] | None:
if is_batch_retrieve_call_type(kwargs.get("call_type")):
if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs):
return None
try:
call: Final = _CALL_KWARGS.validate_python(kwargs)

View file

@ -9,6 +9,7 @@ from litellm._internal_context import with_service_target
from litellm._logging import verbose_router_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.router_utils.batch_utils import is_batch_retrieve_call_type
@ -22,7 +23,7 @@ class LowestCostLoggingHandler(CustomLogger):
@with_service_target("router_usage")
def log_success_event(self, kwargs, response_obj, start_time, end_time):
if is_batch_retrieve_call_type(kwargs.get("call_type")):
if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs):
return
try:
"""
@ -98,7 +99,7 @@ class LowestCostLoggingHandler(CustomLogger):
@with_service_target("router_usage")
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
if is_batch_retrieve_call_type(kwargs.get("call_type")):
if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs):
return
try:
"""

View file

@ -13,7 +13,11 @@ from litellm import ModelResponse, token_counter, verbose_logger
from litellm._internal_context import with_service_target
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs, safe_divide_seconds
from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
is_batch_line_item_event,
safe_divide_seconds,
)
from litellm.router_utils.batch_utils import is_batch_retrieve_call_type
from litellm.types.utils import LiteLLMPydanticObjectBase
@ -61,7 +65,7 @@ class LowestLatencyLoggingHandler(CustomLogger):
@with_service_target("router_usage")
def log_success_event(self, kwargs, response_obj, start_time, end_time):
if is_batch_retrieve_call_type(kwargs.get("call_type")):
if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs):
return
try:
"""
@ -188,7 +192,7 @@ class LowestLatencyLoggingHandler(CustomLogger):
"""
Check if Timeout Error, if timeout set deployment latency -> 100
"""
if is_batch_retrieve_call_type(kwargs.get("call_type")):
if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs):
return
try:
metadata_field: Final = self._select_metadata_field(kwargs)
@ -245,7 +249,7 @@ class LowestLatencyLoggingHandler(CustomLogger):
@with_service_target("router_usage")
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
if is_batch_retrieve_call_type(kwargs.get("call_type")):
if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs):
return
try:
"""

View file

@ -9,6 +9,7 @@ from litellm._internal_context import with_service_target
from litellm._logging import verbose_router_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
from litellm.router_utils.batch_utils import is_batch_retrieve_call_type
from litellm.types.utils import LiteLLMPydanticObjectBase
from litellm.utils import print_verbose
@ -30,7 +31,7 @@ class LowestTPMLoggingHandler(CustomLogger):
@with_service_target("router_usage")
def log_success_event(self, kwargs, response_obj, start_time, end_time):
if is_batch_retrieve_call_type(kwargs.get("call_type")):
if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs):
return
try:
"""
@ -85,7 +86,7 @@ class LowestTPMLoggingHandler(CustomLogger):
@with_service_target("router_usage")
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
if is_batch_retrieve_call_type(kwargs.get("call_type")):
if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs):
return
try:
"""

View file

@ -15,7 +15,7 @@ from litellm._internal_context import with_service_target
from litellm._logging import verbose_logger, verbose_router_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs, is_batch_line_item_event
from litellm.router_utils.batch_utils import is_batch_retrieve_call_type
from litellm.types.router import RouterErrors
from litellm.types.utils import LiteLLMPydanticObjectBase, StandardLoggingPayload
@ -254,7 +254,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
@with_service_target("router_usage")
def log_success_event(self, kwargs, response_obj, start_time, end_time):
if is_batch_retrieve_call_type(kwargs.get("call_type")):
if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs):
return
try:
"""
@ -297,7 +297,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
@with_service_target("router_usage")
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
if is_batch_retrieve_call_type(kwargs.get("call_type")):
if is_batch_retrieve_call_type(kwargs.get("call_type")) or is_batch_line_item_event(kwargs):
return
try:
"""

View file

@ -3228,6 +3228,9 @@ class StandardLoggingHiddenParams(TypedDict):
batch_failed_requests: ReadOnly[int | None]
litellm_model_name: str | None # the model name sent to the provider by litellm
usage_object: dict | None
batch_id: NotRequired[ReadOnly[str | None]]
batch_custom_id: NotRequired[ReadOnly[str | None]]
batch_line_status_code: NotRequired[ReadOnly[int | None]]
class StandardLoggingModelInformation(TypedDict):

View file

@ -222,8 +222,8 @@ async def test_batch_retrieve_cost_tracking_with_completed_batch_no_explicit_cos
)
logging_obj.custom_llm_provider = "openai"
# Mock _handle_completed_batch to return cost data
from litellm.batches.batch_utils import BatchCostUsageResult
# Mock _handle_completed_batch_with_files to return cost data
from litellm.batches.batch_utils import BatchCostUsageResult, BatchResultFiles
expected_cost = 0.05
expected_usage = litellm.Usage(
@ -234,14 +234,17 @@ async def test_batch_retrieve_cost_tracking_with_completed_batch_no_explicit_cos
expected_models = ["gpt-5-mini"]
with patch(
"litellm.litellm_core_utils.litellm_logging._handle_completed_batch",
"litellm.litellm_core_utils.litellm_logging._handle_completed_batch_with_files",
new=AsyncMock(
return_value=BatchCostUsageResult(
cost=expected_cost,
usage=expected_usage,
models=expected_models,
successful_requests=10,
failed_requests=0,
return_value=(
BatchCostUsageResult(
cost=expected_cost,
usage=expected_usage,
models=expected_models,
successful_requests=10,
failed_requests=0,
),
BatchResultFiles(output=None, error=None),
)
),
) as mock_handle_batch:
@ -382,7 +385,7 @@ async def test_batch_retrieve_cost_tracking_with_explicit_cost_data():
explicit_models = ["gpt-5-mini", "gpt-5.5"]
with patch(
"litellm.litellm_core_utils.litellm_logging._handle_completed_batch",
"litellm.litellm_core_utils.litellm_logging._handle_completed_batch_with_files",
new=AsyncMock(),
) as mock_handle_batch:
# Call async_success_handler with explicit cost data
@ -517,7 +520,7 @@ async def test_batch_retrieve_cost_tracking_with_unified_file_id_incomplete_batc
logging_obj.custom_llm_provider = "openai"
with patch(
"litellm.litellm_core_utils.litellm_logging._handle_completed_batch",
"litellm.litellm_core_utils.litellm_logging._handle_completed_batch_with_files",
new=AsyncMock(),
) as mock_handle_batch:
# Call async_success_handler with in_progress batch (unified file ID)
@ -603,17 +606,20 @@ async def test_batch_retrieve_cost_tracking_with_partial_explicit_data():
)
expected_models = ["gpt-5-mini"]
from litellm.batches.batch_utils import BatchCostUsageResult
from litellm.batches.batch_utils import BatchCostUsageResult, BatchResultFiles
with patch(
"litellm.litellm_core_utils.litellm_logging._handle_completed_batch",
"litellm.litellm_core_utils.litellm_logging._handle_completed_batch_with_files",
new=AsyncMock(
return_value=BatchCostUsageResult(
cost=expected_cost,
usage=expected_usage,
models=expected_models,
successful_requests=8,
failed_requests=0,
return_value=(
BatchCostUsageResult(
cost=expected_cost,
usage=expected_usage,
models=expected_models,
successful_requests=8,
failed_requests=0,
),
BatchResultFiles(output=None, error=None),
)
),
) as mock_handle_batch:

View file

@ -0,0 +1,303 @@
from __future__ import annotations
import json
import uuid
from hashlib import sha256
from pathlib import Path
from typing import Final
import pytest
import yaml
from integration._support.client import JSON_OBJECT, Gateway, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.process import owned_proxy
from integration._support.upstream import delete_scenario, register_scenario
from integration._support.wire import Reply, Request, wire_server
from integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse, TextResponse
from pydantic import JsonValue
PROMPT_TOKENS: Final = 10
COMPLETION_TOKENS: Final = 100
OUTPUT_SUCCESS_IDS: Final = ("r1", "r3")
OUTPUT_FAILED_IDS: Final = ("r2",)
ERROR_FILE_IDS: Final = ("r4", "r5")
ALL_CUSTOM_IDS: Final = ("r1", "r2", "r3", "r4", "r5")
def _successful_line(custom_id: str) -> str:
return json.dumps(
{
"id": f"batch_req_{custom_id}",
"custom_id": custom_id,
"response": {
"status_code": 200,
"request_id": f"$REQUEST_ID-{custom_id}",
"body": {
"id": f"chatcmpl-$REQUEST_ID-{custom_id}",
"object": "chat.completion",
"model": "gpt-4o-mini",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": f"answer {custom_id}"},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": PROMPT_TOKENS,
"completion_tokens": COMPLETION_TOKENS,
"total_tokens": PROMPT_TOKENS + COMPLETION_TOKENS,
},
},
},
"error": None,
},
separators=(",", ":"),
)
def _failed_line(custom_id: str) -> str:
return json.dumps(
{
"id": f"batch_req_{custom_id}",
"custom_id": custom_id,
"response": {
"status_code": 400,
"request_id": f"$REQUEST_ID-{custom_id}",
"body": {"error": {"message": f"synthetic line failure {custom_id}", "code": "bad_request"}},
},
"error": {"code": "bad_request", "message": f"synthetic line failure {custom_id}"},
},
separators=(",", ":"),
)
def _batch(status: str, *, files_ready: bool) -> dict[str, JsonValue]:
return {
"id": "batch-$REQUEST_ID",
"object": "batch",
"endpoint": "/v1/chat/completions",
"errors": None,
"input_file_id": "file-in-$REQUEST_ID",
"completion_window": "24h",
"status": status,
"output_file_id": "file-out-$REQUEST_ID" if files_ready else None,
"error_file_id": "file-err-$REQUEST_ID" if files_ready else None,
"created_at": 1,
"in_progress_at": 1,
"completed_at": 1 if files_ready else None,
"expires_at": 1,
"request_counts": {"total": 5, "completed": 2, "failed": 3},
"metadata": None,
}
def _input_lines(model_name: str, marker: str) -> tuple[str, ...]:
return tuple(
json.dumps(
{
"custom_id": custom_id,
"method": "POST",
"url": "/v1/chat/completions",
"body": {"model": model_name, "messages": [{"role": "user", "content": f"{marker} {custom_id}"}]},
},
separators=(",", ":"),
)
for custom_id in ALL_CUSTOM_IDS
)
def _provider_routes(input_lines: tuple[str, ...]) -> RoutedResponse:
output_lines: Final = (_successful_line("r1"), _failed_line("r2"), _successful_line("r3"))
error_lines: Final = tuple(_failed_line(custom_id) for custom_id in ERROR_FILE_IDS)
return RoutedResponse(
content_type="application/x-routed",
routes={
"POST /files": JsonResponse(
content_type="application/json",
body={
"id": "file-in-$REQUEST_ID",
"object": "file",
"purpose": "batch",
"bytes": 100,
"created_at": 1,
"filename": "in.jsonl",
"status": "processed",
},
),
"POST /batches": JsonResponse(
content_type="application/json", body=_batch("validating", files_ready=False)
),
"GET /batches/batch-$REQUEST_ID": JsonResponse(
content_type="application/json", body=_batch("completed", files_ready=True)
),
"GET /files/file-in-$REQUEST_ID/content": TextResponse(
content_type="application/jsonl", body="\n".join(input_lines) + "\n"
),
"GET /files/file-out-$REQUEST_ID/content": TextResponse(
content_type="application/jsonl", body="\n".join(output_lines) + "\n"
),
"GET /files/file-err-$REQUEST_ID/content": TextResponse(
content_type="application/jsonl", body="\n".join(error_lines) + "\n"
),
},
)
def _line_item_config(base: Path, destination: Path) -> Path:
config: Final = yaml.safe_load(base.read_text())
config["litellm_settings"].update({"callbacks": ["generic_api"], "DEFAULT_FLUSH_INTERVAL_SECONDS": 1})
config["general_settings"].update({"store_batch_line_items_in_callbacks": True})
destination.write_text(yaml.safe_dump(config))
return destination
def _spend_rows(key: str) -> tuple[dict[str, JsonValue], ...]:
return tuple(
read_rows(
'SELECT call_type, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE api_key=%s',
(sha256(key.encode()).hexdigest(),),
)
)
def _hidden(event: dict[str, JsonValue]) -> dict[str, JsonValue]:
return object_value(event["hidden_params"])
@pytest.mark.covers(
"spend.batches.completed_batch_emits_one_callback_event_per_jsonl_line_beside_the_aggregate",
"spend.batches.line_item_callback_events_do_not_bill_spend_twice",
)
def test_completed_batch_emits_paired_request_response_callback_events_per_jsonl_line(
gateway: Gateway, tmp_path: Path
) -> None:
marker: Final = "batch-line-items-" + uuid.uuid4().hex[:12]
sink_secret: Final = "synthetic-sink-secret-" + marker
provider_secret: Final = "synthetic-provider-secret-" + marker
def sink(request: Request) -> Reply:
assert request.headers["authorization"] == f"Bearer {sink_secret}"
return Reply()
with (
wire_server(sink) as endpoint,
owned_proxy(
gateway,
tmp_path,
{"GENERIC_LOGGER_ENDPOINT": endpoint.url, "GENERIC_LOGGER_HEADERS": f"Authorization=Bearer {sink_secret}"},
config=_line_item_config(Path("tests/integration/proxy_config.yaml"), tmp_path / "line_items.yaml"),
) as candidate,
candidate.scenario() as scenario,
):
routed_model: Final = scenario.model(api_base=f"{gateway.upstream_url}/{marker}", api_key=provider_secret)
input_lines: Final = _input_lines(routed_model, marker)
handle: Final = register_scenario(marker, _provider_routes(input_lines))
scenario.cleanups.callback(delete_scenario, handle)
key: Final = scenario.key(models=[routed_model])
file_response: Final = candidate.request_multipart(
"/v1/files",
{"purpose": "batch", "model": routed_model},
{"file": ("in.jsonl", ("\n".join(input_lines) + "\n").encode(), "application/jsonl")},
key=key,
)
assert file_response.status_code == 200, file_response.text
input_file_id: Final = string_value(JSON_OBJECT.validate_json(file_response.content)["id"])
batch_response: Final = candidate.request(
"POST",
"/v1/batches",
{
"input_file_id": input_file_id,
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"model": routed_model,
},
key=key,
)
assert batch_response.status_code == 200, batch_response.text
batch_id: Final = string_value(JSON_OBJECT.validate_json(batch_response.content)["id"])
retrieval: Final = eventually(
lambda: candidate.request("GET", f"/v1/batches/{batch_id}", key=key),
lambda response: response.status_code == 200 and response.json()["status"] == "completed",
seconds=30,
)
assert retrieval.status_code == 200, retrieval.text
key_hash: Final = sha256(key.encode()).hexdigest()
batches: Final[
list[Request]
] = [] # mutable-ok: drain() consumes the queue, later polls must keep earlier batches
def delivered() -> tuple[dict[str, JsonValue], ...]:
batches.extend(endpoint.drain())
return tuple(
event
for batch in batches
for event in json.loads(batch.body)
if object_value(event["metadata"]).get("user_api_key_hash") == key_hash
)
events: Final = eventually(
delivered,
lambda values: (
len([e for e in values if e["call_type"] == "acreate_batch"]) >= 1
and len([e for e in values if e["call_type"] == "aretrieve_batch"]) >= 1
and len([e for e in values if _hidden(e).get("batch_custom_id") is not None]) >= len(ALL_CUSTOM_IDS)
),
seconds=40,
)
for batch in batches:
for credential in (provider_secret, sink_secret, key):
assert credential.encode() not in batch.body
aggregate_events: Final = tuple(event for event in events if event["call_type"] == "aretrieve_batch")
line_events: Final = tuple(event for event in events if _hidden(event).get("batch_custom_id") is not None)
assert len(aggregate_events) == 1, [event["call_type"] for event in events]
other_events: Final = tuple(event for event in events if event not in aggregate_events + line_events)
assert sorted(event["call_type"] for event in other_events) == ["acreate_batch", "acreate_file"], other_events
assert sorted(string_value(_hidden(event)["batch_custom_id"]) for event in line_events) == list(
ALL_CUSTOM_IDS
), line_events
by_custom_id: Final = {string_value(_hidden(event)["batch_custom_id"]): event for event in line_events}
for custom_id, line in zip(ALL_CUSTOM_IDS, input_lines, strict=True):
event: Final = by_custom_id[custom_id]
hidden: Final = _hidden(event)
assert hidden["batch_id"] == batch_id, hidden
assert event["call_type"] == "acompletion", event["call_type"]
assert event["messages"] == json.loads(line)["body"]["messages"], (custom_id, event["messages"])
if custom_id in OUTPUT_SUCCESS_IDS:
assert event["status"] == "success", event
assert hidden["batch_line_status_code"] == 200, hidden
assert event["prompt_tokens"] == PROMPT_TOKENS and event["completion_tokens"] == COMPLETION_TOKENS
assert object_value(event["response"])["id"] == f"chatcmpl-{marker}-{custom_id}", event["response"]
assert object_value(object_value(event["response"])["choices"][0]["message"])["content"] == (
f"answer {custom_id}"
)
else:
assert event["status"] == "failure", event
assert hidden["batch_line_status_code"] == 400, hidden
assert f"synthetic line failure {custom_id}" in json.dumps(event["error_information"]), event
aggregate: Final = aggregate_events[0]
assert aggregate["prompt_tokens"] == len(OUTPUT_SUCCESS_IDS) * PROMPT_TOKENS, aggregate
assert aggregate["completion_tokens"] == len(OUTPUT_SUCCESS_IDS) * COMPLETION_TOKENS, aggregate
rows: Final = eventually(
lambda: _spend_rows(key),
lambda values: any(r["call_type"] == "aretrieve_batch" for r in values),
seconds=70,
)
batch_rows: Final = tuple(row for row in rows if row["call_type"] == "aretrieve_batch")
assert len(batch_rows) == 1, rows
assert not any(row["call_type"] == "acompletion" for row in rows), rows
assert batch_rows[0]["prompt_tokens"] == len(OUTPUT_SUCCESS_IDS) * PROMPT_TOKENS, rows
repeated_gets: Final = [candidate.request("GET", f"/v1/batches/{batch_id}", key=key) for _ in range(2)]
assert all(response.status_code == 200 for response in repeated_gets)
events_after_repeats: Final = eventually(
delivered,
lambda values: len([e for e in values if e["call_type"] == "aretrieve_batch"]) >= 1 + len(repeated_gets),
seconds=40,
)
repeat_line_events: Final = tuple(
event for event in events_after_repeats if _hidden(event).get("batch_custom_id") is not None
)
assert len(repeat_line_events) == len(ALL_CUSTOM_IDS), [
(event["call_type"], _hidden(event).get("batch_custom_id")) for event in events_after_repeats
]

View file

@ -523,6 +523,92 @@ async def test_trace_ingest_and_invalid_payload_do_not_write_spend():
storage.ensure_schema.assert_not_awaited()
@pytest.mark.asyncio
async def test_batch_line_item_success_event_does_not_write_spend(monkeypatch: pytest.MonkeyPatch):
# A batch line item carries call_type=acompletion + litellm_params.batch_parent_id.
# The aggregate aretrieve_batch row already bills the batch, so a per-line spend row
# would bill it twice.
monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", True, raising=False)
storage = MagicMock()
logger = ClickHouseSpendLogger(storage=storage)
now = datetime.now(timezone.utc)
await logger.async_log_success_event(
{
"standard_logging_object": _minimal_payload("chatcmpl-line-1", status="success", cost=0.25),
"call_type": "acompletion",
"litellm_params": {"batch_parent_id": "batch-1"},
},
None,
now,
now,
)
assert logger.log_queue == []
await logger.async_log_success_event(
{
"standard_logging_object": _minimal_payload("chatcmpl-live-1", status="success", cost=0.25),
"call_type": "acompletion",
"litellm_params": {},
},
None,
now,
now,
)
# The aggregate aretrieve_batch event bills the batch; it must still write one
# row with the batch's full cost even while line items are skipped.
await logger.async_log_success_event(
{
"standard_logging_object": {
**_minimal_payload("batch-1", status="success", cost=0.0001032),
"call_type": "aretrieve_batch",
},
"call_type": "aretrieve_batch",
"response_cost": 0.0001032,
"litellm_params": {},
},
None,
now,
now,
)
assert len(logger.log_queue) == 2
assert logger.log_queue[0]["request_id"] == "chatcmpl-live-1"
assert logger.log_queue[1]["call_type"] == "aretrieve_batch"
assert logger.log_queue[1]["request_id"] == "batch-1"
assert logger.log_queue[1]["spend"] == 0.0001032
# A FAILED batch line is still part of the aggregate row's request counts and
# must not add its own spend row; a genuine live failure still writes one.
await logger.async_log_failure_event(
{
"standard_logging_object": _minimal_payload("chatcmpl-line-2", status="failure", cost=0.0),
"call_type": "acompletion",
"litellm_params": {"batch_parent_id": "batch-1"},
},
None,
now,
now,
)
await logger.async_log_failure_event(
{
"standard_logging_object": _minimal_payload("chatcmpl-live-2", status="failure", cost=0.0),
"call_type": "acompletion",
"litellm_params": {},
},
None,
now,
now,
)
assert len(logger.log_queue) == 3
assert logger.log_queue[-1]["request_id"] == "chatcmpl-live-2"
assert logger.log_queue[-1]["status"] == "failure"
if logger._flush_task is not None:
logger._flush_task.cancel()
@pytest.mark.parametrize(
"status,llm_cost,guardrail_cost,expected",
[

File diff suppressed because it is too large Load diff

View file

@ -2485,3 +2485,57 @@ async def test_flag_sends_every_vertex_row_down_the_native_path_when_a_model_is_
assert calls == []
assert (result.successful_requests, result.failed_requests) == (0, 1)
@pytest.mark.asyncio
async def test_handle_completed_batch_with_files_returns_bytes_and_fetches_error_once(monkeypatch):
rows = [_success_row(model="gpt-4o", usage=_usage(10, 5))]
output_bytes = _vertex_jsonl(rows)
error_bytes = _vertex_jsonl([{"custom_id": "req-bad", "error": {"message": "rejected"}}])
fetched = []
async def fake_fetch(batch, custom_llm_provider, litellm_params=None):
return output_bytes
async def fake_afile_content(**kw):
fetched.append(kw["file_id"])
return type("R", (), {"content": error_bytes})()
import litellm.cost_calculator as cc
import litellm.files.main as files_main
monkeypatch.setattr(bu, "_fetch_batch_output_file_content", fake_fetch)
monkeypatch.setattr(files_main, "afile_content", fake_afile_content)
monkeypatch.setattr(cc, "batch_cost_calculator", lambda **kw: (2.0, 1.3))
batch = _batch("of").model_copy(update={"error_file_id": "ef"})
result, files = await bu._handle_completed_batch_with_files(batch, custom_llm_provider="openai")
assert result.cost == 3.3
assert result.failed_requests == 1
assert files.output == output_bytes
assert files.error == error_bytes
assert fetched == ["ef"]
@pytest.mark.asyncio
async def test_handle_completed_batch_with_files_no_output_returns_error_bytes(monkeypatch):
error_bytes = _vertex_jsonl([{"custom_id": "req-bad", "error": {"message": "rejected"}}])
async def fake_afile_content(**kw):
return type("R", (), {"content": error_bytes})()
import litellm.files.main as files_main
monkeypatch.setattr(files_main, "afile_content", fake_afile_content)
batch = _batch(None).model_copy(update={"error_file_id": "ef"})
result, files = await bu._handle_completed_batch_with_files(batch, custom_llm_provider="openai")
assert result.failed_requests == 1
assert result.usage.total_tokens == 0
assert files.output is None
assert files.error == error_bytes

View file

@ -0,0 +1,419 @@
"""
Guards that keep batch line-item callback events out of the built-in metering sinks.
With ``litellm.store_batch_line_items_in_callbacks`` on, a completed batch emits
one child callback per JSONL line on top of the aggregate ``aretrieve_batch``
event. Every line carries a real ``response_cost``, so any billing/metering sink
without the ``is_batch_line_item_event`` guard meters aggregate + per-line and
reports roughly 2x the true spend.
These tests put REAL sink instances on the REAL dispatch lists, drive the REAL
``Logging.async_success_handler`` for a completed 2-line batch (1 success line +
1 error line, aggregate cost $1.50, per-line cost $0.03), and assert each sink
meters exactly the aggregate. OTel spans for line items must still be emitted:
the guard belongs on cost/token metrics, not on tracing.
"""
import io
import json
import os
import time
import uuid
from contextlib import redirect_stdout
from datetime import datetime
from types import SimpleNamespace
from typing import Any, Final
from unittest.mock import AsyncMock, patch
import pytest
import litellm
from litellm.batches.batch_line_item_logging import batch_line_item_claim_cache
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.types.utils import LiteLLMBatch, Usage
AGGREGATE_COST: Final[float] = 1.5
INPUT_JSONL: Final[bytes] = b"\n".join(
[
json.dumps(
{
"custom_id": "a",
"method": "POST",
"url": "/v1/chat/completions",
"body": {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi a"}]},
}
).encode(),
json.dumps(
{
"custom_id": "b",
"method": "POST",
"url": "/v1/chat/completions",
"body": {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi b"}]},
}
).encode(),
]
)
OUTPUT_JSONL: Final[bytes] = json.dumps(
{
"custom_id": "a",
"response": {
"status_code": 200,
"body": {
"id": "chatcmpl-line-1",
"model": "gpt-4o",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
},
},
}
).encode()
ERROR_JSONL: Final[bytes] = json.dumps(
{
"custom_id": "b",
"response": {"status_code": 400, "body": {"error": {"message": "boom"}}},
"error": {"message": "boom"},
}
).encode()
_FILE_BYTES: Final[dict[str, bytes]] = {
"input-file-1": INPUT_JSONL,
"output-file-1": OUTPUT_JSONL,
"error-file-1": ERROR_JSONL,
}
def _file_content(file_id: str, **_kwargs: Any) -> SimpleNamespace:
return SimpleNamespace(content=_FILE_BYTES[file_id])
def _batch() -> LiteLLMBatch:
return LiteLLMBatch(
id=f"batch_{uuid.uuid4().hex[:8]}",
object="batch",
endpoint="/v1/chat/completions",
input_file_id="input-file-1",
output_file_id="output-file-1",
error_file_id="error-file-1",
status="completed",
completion_window="24h",
created_at=1,
)
def _parent_logging() -> Logging:
logging_obj = Logging(
model="gpt-4o",
messages=[{"role": "user", "content": "<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",
"user_api_key_user_id": "user-77",
"user_api_key_team_id": "team-7",
"user_api_key_team_alias": "team-seven",
"user_api_key_alias": "key-alias-1",
}
},
optional_params={},
custom_llm_provider="openai",
)
return logging_obj
class RecordingHTTP:
"""Stands in for a sink's HTTP egress object only; the sink logic is real."""
def __init__(self) -> None:
self.posts: list[dict[str, Any]] = []
async def post(self, url: str, data: Any = None, content: Any = None, **_kw: Any) -> SimpleNamespace:
self.posts.append({"url": url, "body": data if data is not None else content})
return SimpleNamespace(status_code=200, text="ok", raise_for_status=lambda: None)
async def put(self, url: str, content: Any = None, **_kw: Any) -> SimpleNamespace:
return SimpleNamespace(status_code=202, text="ok", raise_for_status=lambda: None)
async def _log_completed_batch(monkeypatch: pytest.MonkeyPatch, loggers: list, flag: bool = True) -> None:
batch_line_item_claim_cache.in_memory_cache.flush_cache()
monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", flag, raising=False)
saved_success = list(litellm._async_success_callback)
saved_failure = list(litellm._async_failure_callback)
litellm._async_success_callback = list(loggers)
litellm._async_failure_callback = list(loggers)
buf = io.StringIO()
try:
with (
patch("litellm.files.main.afile_content", new_callable=AsyncMock, side_effect=_file_content),
patch("litellm.cost_calculator.batch_cost_calculator", return_value=(0.01, 0.02)),
redirect_stdout(buf),
):
await _parent_logging().async_success_handler(
result=_batch(),
batch_cost=1.5,
batch_usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15),
batch_models=["gpt-4o"],
batch_successful_requests=1,
batch_failed_requests=1,
batch_prompt_cost=1.0,
batch_completion_cost=0.5,
)
finally:
litellm._async_success_callback = saved_success
litellm._async_failure_callback = saved_failure
def _openmeter_sinks() -> tuple[Any, RecordingHTTP]:
from litellm.integrations.openmeter import OpenMeterLogger
recorder = RecordingHTTP()
logger = OpenMeterLogger()
logger.async_http_handler = recorder
return logger, recorder
def _lago_sinks() -> tuple[Any, RecordingHTTP]:
from litellm.integrations.lago import LagoLogger
recorder = RecordingHTTP()
logger = LagoLogger()
logger.async_http_handler = recorder
return logger, recorder
@pytest.mark.asyncio
async def test_openmeter_meters_only_the_aggregate_batch_event(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("OPENMETER_API_KEY", "test-openmeter-key")
logger, recorder = _openmeter_sinks()
await _log_completed_batch(monkeypatch, [logger])
costs = [json.loads(p["body"])["data"]["cost"] for p in recorder.posts]
assert costs == [AGGREGATE_COST], (
f"OpenMeter metered per-line costs on top of the aggregate: {costs}; "
"line items must be billed only by the aggregate aretrieve_batch event"
)
@pytest.mark.asyncio
async def test_lago_bills_only_the_aggregate_batch_event(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LAGO_API_KEY", "test-lago-key")
monkeypatch.setenv("LAGO_API_BASE", "http://lago.invalid")
monkeypatch.setenv("LAGO_API_EVENT_CODE", "litellm-usage")
monkeypatch.setenv("LAGO_API_CHARGE_BY", "user_id")
logger, recorder = _lago_sinks()
await _log_completed_batch(monkeypatch, [logger])
costs = [json.loads(p["body"])["event"]["properties"]["response_cost"] for p in recorder.posts]
assert costs == [AGGREGATE_COST], (
f"Lago billed per-line costs on top of the aggregate: {costs}; "
"line items must be billed only by the aggregate aretrieve_batch event"
)
@pytest.mark.asyncio
async def test_datadog_cost_management_bills_only_the_aggregate_batch_event(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from litellm.integrations.datadog.datadog_cost_management import DatadogCostManagementLogger
monkeypatch.setenv("DD_API_KEY", "test-dd-key")
monkeypatch.setenv("DD_APP_KEY", "test-dd-app-key")
logger = DatadogCostManagementLogger(cost_tag_keys=[])
await _log_completed_batch(monkeypatch, [logger])
entries = list(logger.log_queue)
costs = [e.get("response_cost", 0) for e in entries]
assert costs == [AGGREGATE_COST], (
f"Datadog FOCUS BilledCost queued per-line entries on top of the aggregate: {costs}; "
"line items must be billed only by the aggregate aretrieve_batch event"
)
@pytest.mark.asyncio
async def test_newrelic_metrics_meters_only_the_aggregate_batch_event(monkeypatch: pytest.MonkeyPatch) -> None:
from litellm.integrations.newrelic.newrelic_metrics import NewRelicMetricsLogger, build_metric_payload
logger = NewRelicMetricsLogger(newrelic_api_key="test-nr-key")
await _log_completed_batch(monkeypatch, [logger])
records = tuple(logger.log_queue)
costs = [r.response_cost for r in records]
assert costs == [AGGREGATE_COST], (
f"New Relic metered per-line costs on top of the aggregate: {costs}; "
"line items must be metered only by the aggregate aretrieve_batch event"
)
now = time.time()
envelopes = build_metric_payload(records=records, window_start=now - 1, now=now)
cost_sum = 0.0
for envelope in envelopes:
for metric in envelope["metrics"]:
if "cost" in metric["name"]:
value = metric["value"]["sum"] if isinstance(metric["value"], dict) else metric["value"]
cost_sum += value
assert cost_sum == pytest.approx(AGGREGATE_COST)
@pytest.mark.asyncio
async def test_prometheus_meters_only_the_aggregate_batch_event(monkeypatch: pytest.MonkeyPatch) -> None:
prometheus_client = pytest.importorskip("prometheus_client")
from litellm.integrations.prometheus import PrometheusLogger
from prometheus_client import REGISTRY
saved_collectors = list(REGISTRY._collector_to_names.keys())
for collector in saved_collectors:
REGISTRY.unregister(collector)
logger = PrometheusLogger()
try:
await _log_completed_batch(monkeypatch, [logger])
spend_samples = []
for metric in REGISTRY.collect():
if metric.name == "litellm_spend_metric":
spend_samples = [sample.value for sample in metric.samples if sample.name.endswith("_total")]
finally:
for collector in list(REGISTRY._collector_to_names.keys()):
REGISTRY.unregister(collector)
for collector in saved_collectors:
try:
REGISTRY.register(collector)
except Exception: # noqa: BLE001 # already re-registered by another holder
pass
assert spend_samples, "litellm_spend_metric saw no samples at all"
assert sum(spend_samples) == pytest.approx(AGGREGATE_COST), (
f"Prometheus spend metric double-counted batch line items: {spend_samples}; "
"line items must be metered only by the aggregate aretrieve_batch event"
)
def _otel_v1(monkeypatch: pytest.MonkeyPatch):
otel_sdk = pytest.importorskip("opentelemetry.sdk")
from opentelemetry.sdk.metrics import MeterProvider
from opentelemetry.sdk.metrics.export import InMemoryMetricReader
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.sdk.trace.export import SimpleSpanProcessor
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
from litellm.integrations.opentelemetry import OpenTelemetry as OTelV1, OpenTelemetryConfig
reader = InMemoryMetricReader()
span_exporter = InMemorySpanExporter()
tracer_provider = TracerProvider()
tracer_provider.add_span_processor(SimpleSpanProcessor(span_exporter))
logger = OTelV1(
config=OpenTelemetryConfig(exporter="console", enable_metrics=True),
callback_name="batch_sink_guard_v1",
tracer_provider=tracer_provider,
meter_provider=MeterProvider(metric_readers=[reader]),
)
return logger, reader, span_exporter
def _cost_points(reader: Any) -> list[float]:
data = reader.get_metrics_data()
points: list[float] = []
if data is None:
return points
for resource_metrics in data.resource_metrics:
for scope_metrics in resource_metrics.scope_metrics:
for metric in scope_metrics.metrics:
if not metric.name.endswith("cost"):
continue
for data_point in metric.data.data_points:
points.append(getattr(data_point, "sum", getattr(data_point, "value", None)))
return points
@pytest.mark.asyncio
async def test_otel_v1_meter_cost_skips_line_items_but_spans_stay(
monkeypatch: pytest.MonkeyPatch,
) -> None:
logger, reader, span_exporter = _otel_v1(monkeypatch)
await _log_completed_batch(monkeypatch, [logger])
costs = _cost_points(reader)
assert sum(costs) == pytest.approx(AGGREGATE_COST), (
f"OTel v1 gen_ai.usage.cost double-counted batch line items: {costs}; "
"line items must be metered only by the aggregate aretrieve_batch event"
)
finished = span_exporter.get_finished_spans()
assert finished, "OTel v1 emitted no spans at all"
assert any("chatcmpl-line-1" in str(span.attributes) for span in finished), (
f"OTel v1 dropped the per-line span the feature exists to deliver: {[span.name for span in finished]}"
)
def _reader_snapshot(reader: Any) -> dict[str, list[float]]:
"""Every metric's datapoint values, keyed by metric name (sums for histograms)."""
data = reader.get_metrics_data()
out: dict[str, list[float]] = {}
if data is None:
return out
for resource_metrics in data.resource_metrics:
for scope_metrics in resource_metrics.scope_metrics:
for metric in scope_metrics.metrics:
for data_point in metric.data.data_points:
value = getattr(data_point, "sum", None)
if value is None:
value = getattr(data_point, "value", None)
if value is not None:
out.setdefault(metric.name, []).append(value)
return out
async def _run_batch_with_otel_v2(monkeypatch: pytest.MonkeyPatch, flag: bool) -> dict[str, list[float]]:
from opentelemetry.sdk.metrics import MeterProvider
from opentelemetry.sdk.metrics.export import InMemoryMetricReader
from opentelemetry.sdk.trace import TracerProvider
from litellm.integrations.otel.logger import OpenTelemetryV2
from litellm.integrations.otel.model.config import OpenTelemetryV2Config
reader = InMemoryMetricReader()
logger = OpenTelemetryV2(
config=OpenTelemetryV2Config(exporter="console", enable_metrics=True),
callback_name="batch_sink_guard_v2",
tracer_provider=TracerProvider(),
meter_provider=MeterProvider(metric_readers=[reader]),
)
await _log_completed_batch(monkeypatch, [logger], flag=flag)
return _reader_snapshot(reader)
@pytest.mark.asyncio
async def test_otel_v2_metrics_are_identical_with_the_flag_on_and_off(
monkeypatch: pytest.MonkeyPatch,
) -> None:
pytest.importorskip("opentelemetry.sdk")
off = await _run_batch_with_otel_v2(monkeypatch, flag=False)
on = await _run_batch_with_otel_v2(monkeypatch, flag=True)
for name in ("gen_ai.usage.cost", "gen_ai.client.token.usage"):
assert on.get(name) == off.get(name), (
f"OTel v2 {name} changed when the line-item flag was turned on; line items "
"must be metered only by the aggregate aretrieve_batch event. "
f"off={off.get(name)} on={on.get(name)}"
)
for name, values in off.items():
on_counts = [len(on.get(name, [])), len(values)]
assert on_counts[0] == on_counts[1], (
f"OTel v2 {name} gained synthetic per-line samples with the flag on: "
f"{on_counts[0]} datapoints on vs {on_counts[1]} off"
)
assert sum(off.get("gen_ai.usage.cost", [])) == pytest.approx(AGGREGATE_COST)

View file

@ -14,6 +14,7 @@ from litellm.litellm_core_utils.core_helpers import (
drop_params_flag,
get_or_create_metadata_bucket,
get_provider_response_headers_from_hidden_params,
is_batch_line_item_event,
map_finish_reason,
normalize_drop_params,
reconstruct_model_name,
@ -496,6 +497,16 @@ class TestIsExpectedClientError:
assert is_expected_client_error(vendor_limit) is False
def test_is_batch_line_item_event():
assert is_batch_line_item_event({"litellm_params": {"batch_parent_id": "batch_1"}}) is True
assert is_batch_line_item_event({"litellm_params": {"batch_parent_id": "batch_1", "metadata": {}}}) is True
assert is_batch_line_item_event({"litellm_params": {"metadata": {}}}) is False
assert is_batch_line_item_event({"litellm_params": {"batch_parent_id": None}}) is False
assert is_batch_line_item_event({}) is False
assert is_batch_line_item_event({"litellm_params": "not-a-mapping"}) is False
assert is_batch_line_item_event({"litellm_params": None}) is False
class TestProviderResponseHeadersInHiddenParams:
def test_records_raw_headers_and_the_processed_additional_headers(self):
response = ImageResponse()

View file

@ -1197,19 +1197,24 @@ class TestRetrieveBatchCostPassesModelIdentity:
captured: dict[str, object] = {}
from litellm.batches.batch_utils import BatchCostUsageResult
from litellm.batches.batch_utils import BatchCostUsageResult, BatchResultFiles
async def fake_handle_completed_batch(**kwargs: object) -> BatchCostUsageResult:
async def fake_handle_completed_batch(
**kwargs: object,
) -> tuple[BatchCostUsageResult, BatchResultFiles]:
captured.update(kwargs)
return BatchCostUsageResult(
cost=1.25,
usage=Usage(prompt_tokens=1800, completion_tokens=1000, total_tokens=2800),
models=["m"],
successful_requests=1,
failed_requests=0,
return (
BatchCostUsageResult(
cost=1.25,
usage=Usage(prompt_tokens=1800, completion_tokens=1000, total_tokens=2800),
models=["m"],
successful_requests=1,
failed_requests=0,
),
BatchResultFiles(output=None, error=None),
)
monkeypatch.setattr(logging_module, "_handle_completed_batch", fake_handle_completed_batch)
monkeypatch.setattr(logging_module, "_handle_completed_batch_with_files", fake_handle_completed_batch)
obj = LitellmLogging(
model="bedrock/global.anthropic.claude-sonnet-4-6",
@ -1292,7 +1297,7 @@ class TestRetrieveBatchPricesOnlyFinalBatches:
from litellm.litellm_core_utils import litellm_logging as logging_module
handle_completed_batch = AsyncMock()
monkeypatch.setattr(logging_module, "_handle_completed_batch", handle_completed_batch)
monkeypatch.setattr(logging_module, "_handle_completed_batch_with_files", handle_completed_batch)
batch = self._batch(status, output_file_id)
await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None)
@ -1302,20 +1307,23 @@ class TestRetrieveBatchPricesOnlyFinalBatches:
@pytest.mark.asyncio
async def test_completed_batch_with_output_is_priced(self, monkeypatch) -> None:
from litellm.batches.batch_utils import BatchCostUsageResult
from litellm.batches.batch_utils import BatchCostUsageResult, BatchResultFiles
from litellm.litellm_core_utils import litellm_logging as logging_module
from litellm.types.utils import Usage
handle_completed_batch = AsyncMock(
return_value=BatchCostUsageResult(
cost=8e-06,
usage=Usage(prompt_tokens=26, completion_tokens=9, total_tokens=35),
models=["gpt-5.6-luna"],
successful_requests=2,
failed_requests=0,
return_value=(
BatchCostUsageResult(
cost=8e-06,
usage=Usage(prompt_tokens=26, completion_tokens=9, total_tokens=35),
models=["gpt-5.6-luna"],
successful_requests=2,
failed_requests=0,
),
BatchResultFiles(output=None, error=None),
)
)
monkeypatch.setattr(logging_module, "_handle_completed_batch", handle_completed_batch)
monkeypatch.setattr(logging_module, "_handle_completed_batch_with_files", handle_completed_batch)
batch = self._batch("completed", "file-out")
await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None)
@ -9315,3 +9323,179 @@ def test_signoz_dispatch_requires_an_endpoint(monkeypatch):
logging_module._in_memory_loggers.clear()
monkeypatch.delenv("LITELLM_OTEL_V2", raising=False)
is_otel_v2_enabled.cache_clear()
class TestRetrieveBatchReusesFetchedResultFiles:
"""The aggregate path and log_batch_line_items must share one fetch of each
result file: output and error are fetched once total (not once each per
consumer), and forwarded bytes keep line-item logging off the wire."""
@staticmethod
def _batch_file_bytes() -> dict[str, bytes]:
input_jsonl = json.dumps(
{
"custom_id": "a",
"method": "POST",
"url": "/v1/chat/completions",
"body": {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]},
}
).encode()
output_jsonl = json.dumps(
{
"custom_id": "a",
"response": {
"status_code": 200,
"body": {
"id": "chatcmpl-1",
"model": "gpt-4o",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "hi back"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
},
},
}
).encode()
error_jsonl = json.dumps({"custom_id": "b", "error": {"message": "rejected"}}).encode()
return {"file-in": input_jsonl, "file-out": output_jsonl, "file-err": error_jsonl}
@staticmethod
def _logging_obj() -> LitellmLogging:
obj = LitellmLogging(
model="gpt-4o",
messages=[{"role": "user", "content": "Hey"}],
stream=False,
call_type="aretrieve_batch",
start_time=time.time(),
litellm_call_id="batch-call-reuse",
function_id="f",
)
obj.custom_llm_provider = "openai"
obj.update_environment_variables(
litellm_params={"metadata": {}},
optional_params={},
custom_llm_provider="openai",
)
return obj
@staticmethod
def _batch():
from litellm.types.utils import LiteLLMBatch
return LiteLLMBatch(
id="batch_reuse_fetched_files",
completion_window="24h",
created_at=1,
endpoint="/v1/chat/completions",
input_file_id="file-in",
object="batch",
status="completed",
output_file_id="file-out",
error_file_id="file-err",
)
@pytest.mark.asyncio
async def test_compute_path_fetches_each_result_file_once(self, monkeypatch) -> None:
from litellm.batches.batch_line_item_logging import batch_line_item_claim_cache
batch_line_item_claim_cache.in_memory_cache.flush_cache()
file_bytes = self._batch_file_bytes()
async def fake_afile_content(**kwargs):
from types import SimpleNamespace
return SimpleNamespace(content=file_bytes[kwargs["file_id"]])
file_mock = AsyncMock(side_effect=fake_afile_content)
monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", True, raising=False)
monkeypatch.setattr("litellm.files.main.afile_content", file_mock)
monkeypatch.setattr("litellm.cost_calculator.batch_cost_calculator", lambda **kw: (0.01, 0.02))
await self._logging_obj()._async_success_handler_body(
result=self._batch(), start_time=None, end_time=None
)
fetched = sorted(call.kwargs["file_id"] for call in file_mock.await_args_list)
assert fetched == ["file-err", "file-in", "file-out"], (
"each result file must be fetched exactly once across aggregate costing and line-item logging"
)
@pytest.mark.asyncio
async def test_explicit_kwargs_path_forwards_result_files(self, monkeypatch) -> None:
from litellm.batches.batch_line_item_logging import batch_line_item_claim_cache
batch_line_item_claim_cache.in_memory_cache.flush_cache()
from litellm.types.utils import Usage
file_bytes = self._batch_file_bytes()
async def fake_afile_content(**kwargs):
from types import SimpleNamespace
return SimpleNamespace(content=file_bytes[kwargs["file_id"]])
file_mock = AsyncMock(side_effect=fake_afile_content)
monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", True, raising=False)
monkeypatch.setattr("litellm.files.main.afile_content", file_mock)
await self._logging_obj().async_success_handler(
result=self._batch(),
batch_cost=1.5,
batch_usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15),
batch_models=["gpt-4o"],
batch_successful_requests=1,
batch_failed_requests=1,
batch_output_file_content=file_bytes["file-out"],
batch_error_file_content=file_bytes["file-err"],
)
assert file_mock.await_count == 1
assert file_mock.await_args.kwargs["file_id"] == "file-in", (
"forwarded output/error bytes must leave only the input file to fetch"
)
@pytest.mark.asyncio
async def test_line_item_logging_failure_leaves_aggregate_result_intact(self, monkeypatch) -> None:
from litellm.batches.batch_line_item_logging import batch_line_item_claim_cache
batch_line_item_claim_cache.in_memory_cache.flush_cache()
from litellm.types.utils import Usage
monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", True, raising=False)
monkeypatch.setattr(
"litellm.batches.batch_line_item_logging.log_batch_line_items",
AsyncMock(side_effect=RuntimeError("claim backend exploded")),
)
batch = self._batch()
await self._logging_obj().async_success_handler(
result=batch,
batch_cost=1.5,
batch_usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15),
batch_models=["gpt-4o"],
batch_successful_requests=1,
batch_failed_requests=1,
)
assert batch._hidden_params["response_cost"] == 1.5
assert batch.usage.total_tokens == 15
def test_caller_supplied_litellm_params_cannot_forge_logging_markers(logging_obj):
"""A request body key `litellm_params` flows into optional_params; the logging
kwargs must keep the internal litellm_params or a caller could forge markers
(batch_parent_id) that spend sinks and router counters use to skip events."""
from litellm.litellm_core_utils.core_helpers import is_batch_line_item_event
logging_obj.update_environment_variables(
litellm_params={"metadata": {"model_group": "gpt-4o"}},
optional_params={"litellm_params": {"batch_parent_id": "fake-batch"}},
)
assert logging_obj.model_call_details["litellm_params"] is logging_obj.litellm_params
assert "batch_parent_id" not in logging_obj.model_call_details["litellm_params"]
assert not is_batch_line_item_event(logging_obj.model_call_details)

View file

@ -20,7 +20,7 @@ import httpx
import pytest
from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig
from litellm.types.utils import LlmProviders
from litellm.types.utils import EmbeddingResponse, LlmProviders, ModelResponse
# AWS JobStatus -> OpenAI BatchJobStatus, exactly as encoded in transformation.py
# (both transform_create_batch_response and transform_retrieve_batch_response).
@ -941,3 +941,25 @@ def test_retrieve_request_accepts_partition_arns(config: BedrockBatchesConfig, a
batch_id=arn, optional_params={}, litellm_params={}
)
assert result["url"].startswith(expected_prefix)
def test_transform_batch_output_line_dispatches_on_shape(config: BedrockBatchesConfig) -> None:
titan = config.transform_batch_output_line(
{"embedding": [0.1, 0.2], "inputTextTokenCount": 4}, model="amazon.titan-embed"
)
assert isinstance(titan, EmbeddingResponse)
assert titan.data[0]["embedding"] == [0.1, 0.2]
assert titan.usage.prompt_tokens == 4
converse = config.transform_batch_output_line(
{
"output": {"message": {"role": "assistant", "content": [{"text": "hi"}]}},
"stopReason": "end_turn",
"usage": {"inputTokens": 5, "outputTokens": 6, "totalTokens": 11},
},
model="amazon.nova-lite",
)
assert isinstance(converse, ModelResponse)
assert converse.choices[0].message.content == "hi"
assert config.transform_batch_output_line({"foo": 1}, model="x") is None

View file

@ -1546,10 +1546,10 @@ class TestCheckBatchCost:
Logging, "async_success_handler", new_callable=AsyncMock
) as success_handler,
):
provider.get("https://api.openai.com/v1/files/file-output-123/content").mock(
output_route = provider.get("https://api.openai.com/v1/files/file-output-123/content").mock(
return_value=httpx.Response(200, content=f"{succeeded_line}\n{rejected_line}\n".encode())
)
provider.get("https://api.openai.com/v1/files/file-error-456/content").mock(
error_route = provider.get("https://api.openai.com/v1/files/file-error-456/content").mock(
return_value=httpx.Response(200, content=f"{error_file_lines}\n\n".encode())
)
await check_batch_cost_instance.check_batch_cost()
@ -1558,6 +1558,10 @@ class TestCheckBatchCost:
assert len(spend_log_calls) == 1
handler_kwargs = spend_log_calls[0]
assert handler_kwargs["batch_successful_requests"] == 1
assert output_route.call_count == 1
assert error_route.call_count == 1
assert handler_kwargs["batch_output_file_content"] == f"{succeeded_line}\n{rejected_line}\n".encode()
assert handler_kwargs["batch_error_file_content"] == f"{error_file_lines}\n\n".encode()
assert handler_kwargs["batch_failed_requests"] == 3, (
"2 error-file lines must add to the output file's 1 rejected request"
)

View file

@ -196,3 +196,17 @@ async def test_a_batch_polled_within_every_budget_window_is_never_charged_again(
await _poll(limiter, finished, BATCH_COST)
assert _local_spend(limiter, KEY_SPEND_KEY) == pytest.approx(BATCH_COST)
@pytest.mark.asyncio
async def test_batch_line_item_events_do_not_charge_the_model_budget():
"""Line events carry batch_parent_id; the aggregate aretrieve_batch event is
the one that already bills the batch, so children must not double-charge."""
limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache())
line_event = _event("acompletion", CHAT_COST)
line_event["litellm_params"]["batch_parent_id"] = "batch_first"
await limiter.async_log_success_event(line_event, response_obj=None, start_time=None, end_time=None)
assert await _spend(limiter, KEY_SPEND_KEY) == 0.0
assert await _spend(limiter, USER_SPEND_KEY) == 0.0

View file

@ -107,3 +107,29 @@ async def test_async_log_success_event_counts_non_chat_response_tokens(response_
f"expected 50 tokens counted for {scope_id}, "
f"got {current['current_tpm']}"
)
@pytest.mark.asyncio
async def test_async_log_failure_event_skips_batch_line_item_events():
"""Failed line children were never admitted by the limiter, so the failure
hook must not decrement request counters they never incremented."""
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
local_cache = parallel_request_handler.internal_usage_cache.dual_cache.in_memory_cache
await parallel_request_handler.async_log_failure_event(
kwargs={
"exception": "litellm.APIError: upstream 500",
"litellm_params": {
"batch_parent_id": "batch_1",
"metadata": {"user_api_key": hash_token("sk-line-item")},
},
"model": "gpt-3.5-turbo",
},
response_obj=None,
start_time=datetime.now(),
end_time=datetime.now(),
)
assert local_cache.cache_dict == {}

View file

@ -7667,3 +7667,59 @@ async def test_a2a_url_target_owns_invocation_fee_and_request_limit(
await _rpm_request(limiter, cache, auth, "a2a/cheap")
assert denied.value.status_code == 429
assert "expensive" in str(denied.value.detail)
@pytest.mark.asyncio
async def test_async_log_success_event_skips_batch_line_item_events():
"""Per-line batch events already ran through the limiter as the aggregate
aretrieve_batch; children must not increment TPM or request counters."""
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache))
await handler.async_log_success_event(
kwargs={
"standard_logging_object": {"metadata": {"user_api_key_hash": hash_token("sk-line-item")}},
"litellm_params": {
"batch_parent_id": "batch_1",
"metadata": {"user_api_key_hash": hash_token("sk-line-item"), "model_group": "gpt-3.5-turbo"},
},
"model": "gpt-3.5-turbo",
},
response_obj=ModelResponse(
id="x", object="chat.completion", created=1, model="gpt-3.5-turbo",
usage=Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), choices=[],
),
start_time=datetime.now(),
end_time=datetime.now(),
)
assert local_cache.in_memory_cache.cache_dict == {}
@pytest.mark.asyncio
async def test_async_log_failure_event_skips_batch_line_item_events():
"""Failed line children were never admitted by the limiter, so the failure
hook must not release slots or refund TPM they never reserved."""
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache))
stash = get_or_create_request_stash()
stash.owner_litellm_call_id = "call-line-item"
stash.reserved_tokens = 50
stash.reserved_scopes = frozenset({("api_key", hash_token("sk-line-item"))})
await handler.async_log_failure_event(
kwargs={
"litellm_call_id": "call-line-item",
"standard_logging_object": {"metadata": {"user_api_key_hash": hash_token("sk-line-item")}},
"litellm_params": {
"batch_parent_id": "batch_1",
"metadata": {"user_api_key_hash": hash_token("sk-line-item"), "model_group": "gpt-3.5-turbo"},
},
"model": "gpt-3.5-turbo",
},
response_obj=None,
start_time=datetime.now(),
end_time=datetime.now(),
)
assert local_cache.in_memory_cache.cache_dict == {}

View file

@ -2772,6 +2772,24 @@ async def test_async_post_call_failure_hook_persists_no_raw_model_on_an_unknown_
assert error_information["error_class"] == "ProxyModelNotFoundError"
@pytest.mark.asyncio
async def test_batch_line_item_event_never_updates_spend(): # test-quality-ok: the observable contract is exactly that the spend path is never invoked
logger: Final = _ProxyDBLogger()
kwargs: Final = {
"litellm_params": {"batch_parent_id": "batch_1", "metadata": {}},
"model": "gpt-4o",
"call_type": "acompletion",
}
with patch.object(logger, "_PROXY_track_cost_callback", new_callable=AsyncMock) as mock_track: # test-quality-ok: asserts the callback's own method is skipped; the DB writer is never reached
await logger.async_log_success_event(
kwargs=kwargs,
response_obj=ModelResponse(),
start_time=datetime.now(),
end_time=datetime.now(),
)
mock_track.assert_not_awaited()
class _NeverStringifiedMetadataValue:
def __repr__(self) -> str:
raise AssertionError("a request metadata value was stringified by the cost tracking failure path")

View file

@ -8112,6 +8112,43 @@ async def test_update_general_settings_store_model_in_db_false():
assert ps.general_settings["store_model_in_db"] is False
@pytest.mark.asyncio
async def test_update_general_settings_store_batch_line_items_in_callbacks():
"""
Verify _update_general_settings sets the litellm module flag when the DB
general_settings carries store_batch_line_items_in_callbacks, and that a
YAML-explicit value wins over the DB value.
"""
from litellm.proxy.proxy_server import ProxyConfig
proxy_config = ProxyConfig()
saved_flag = litellm.store_batch_line_items_in_callbacks
try:
with patch("litellm.proxy.proxy_server.general_settings", {}): # test-quality-ok: module-global seam
await proxy_config._update_general_settings(
db_general_settings={"store_batch_line_items_in_callbacks": True}
)
assert litellm.store_batch_line_items_in_callbacks is True
proxy_config._yaml_general_settings_keys = {"store_batch_line_items_in_callbacks"}
with patch( # test-quality-ok: module-global seam
"litellm.proxy.proxy_server.general_settings",
{"store_batch_line_items_in_callbacks": "false"},
):
await proxy_config._update_general_settings(
db_general_settings={"store_batch_line_items_in_callbacks": True}
)
assert litellm.store_batch_line_items_in_callbacks is False
proxy_config._yaml_general_settings_keys = set()
litellm.store_batch_line_items_in_callbacks = True # test-quality-ok: seed the prior opt-in so its removal is observable
with patch("litellm.proxy.proxy_server.general_settings", {}): # test-quality-ok: module-global seam
await proxy_config._update_general_settings(db_general_settings={})
assert litellm.store_batch_line_items_in_callbacks is False
finally:
litellm.store_batch_line_items_in_callbacks = saved_flag
@pytest.mark.asyncio
async def test_update_general_settings_propagates_apply_user_budget_to_team_keys():
"""The Admin UI toggle writes to the DB config, so the flag has to be in the
@ -13288,6 +13325,7 @@ def _patched_coordination_redis_module_state(
patch.object(proxy_server_module, "spend_counter_cache", spend_cache),
patch.object(proxy_server_module, "user_api_key_cache", DualCache()),
patch.object(proxy_server_module, "cli_sso_session_cache", DualCache()),
patch.object(proxy_server_module, "batch_line_item_claim_cache", DualCache()),
patch.object(proxy_server_module, "llm_router", None),
patch.object(proxy_server_module, "litellm_config_cache", config_cache),
patch.object(proxy_server_module, "RedisCache", redis_cache_class),

View file

@ -184,10 +184,12 @@ class TestRedisAuthCacheFlag:
ps.cli_sso_session_cache,
ps.user_api_key_cache,
ps.litellm_config_cache,
ps.batch_line_item_claim_cache,
)
with ExitStack() as detached:
for cache in touched_caches:
detached.enter_context(patch.object(cache, "redis_cache", None))
ps._attach_redis_usage_cache(fake_redis, enable_redis_auth_cache=False)
assert limiter_cache.redis_cache is fake_redis
assert ps.batch_line_item_claim_cache.redis_cache is fake_redis
assert ps.user_api_key_cache.redis_cache is None

View file

@ -391,6 +391,29 @@ async def test_push_task_failure_is_logged_once_and_not_leaked(disable_budget_sy
unretrieved.assert_not_called()
@pytest.mark.asyncio
async def test_batch_line_item_events_do_not_charge_the_provider_budget(disable_budget_sync):
limiter = RouterBudgetLimiting(
dual_cache=DualCache(),
provider_budget_config={"openai": BudgetConfig(max_budget=100.0, budget_duration="1d")},
)
await asyncio.gather(*(task for task in asyncio.all_tasks() if task is not asyncio.current_task()))
await limiter.async_log_success_event(
kwargs={
"call_type": "acompletion",
"litellm_params": {"custom_llm_provider": "openai", "batch_parent_id": "batch_x"},
"standard_logging_object": {"response_cost": 1.0, "model_id": "dep-1"},
},
response_obj=None,
start_time=None,
end_time=None,
)
assert limiter.dual_cache.in_memory_cache.get_cache("provider_spend:openai:1d") in (None, 0)
assert limiter.redis_increment_operation_queue == []
_SPEND_KEY = "provider_spend:openai:1d"

View file

@ -208,3 +208,28 @@ async def test_an_open_circuit_breaker_falls_back_without_a_warning_per_request(
assert picked is DEPLOYMENT_B
assert [record.getMessage() for record in caplog.records if record.levelno >= logging.WARNING] == []
assert sum("circuit breaker is open" in record.getMessage() for record in caplog.records) == 2
@pytest.mark.asyncio
async def test_batch_line_item_success_does_not_release_in_flight_slots(monkeypatch: pytest.MonkeyPatch) -> None:
"""A batch line item is historical batch traffic (call_type=acompletion +
litellm_params.batch_parent_id): it never started a live request, so its success
callback must not decrement the in-flight count a live request still holds."""
import litellm
monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", True, raising=False)
shared: Final = SharedRedisCounters()
worker: Final = _worker(shared)
worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
line_item_kwargs: Final = {
"call_type": "acompletion",
"litellm_params": {
"batch_parent_id": "batch-1",
"metadata": {"model_group": GROUP},
"model_info": {"id": "dep-a"},
},
}
await worker.async_log_success_event(line_item_kwargs, None, None, None)
assert shared.count(f"{GROUP}_request_count:dep-a") == 1

View file

@ -117,3 +117,40 @@ async def test_v2_subclass_overriding_async_get_available_deployments_with_the_o
f"from {HIGH_USAGE_DEPLOYMENT_ID}",
f"from {LOW_USAGE_DEPLOYMENT_ID}",
}
@pytest.mark.asyncio
async def test_usage_based_routing_handlers_skip_batch_line_items(monkeypatch: pytest.MonkeyPatch) -> None:
"""
Batch line-item callbacks carry call_type=acompletion plus
litellm_params.batch_parent_id, so the call-type-only batch guard does not fire.
They are historical batch traffic and must not update the TPM/RPM usage counters.
"""
import litellm
from litellm.router_strategy.lowest_tpm_rpm import LowestTPMLoggingHandler
monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", True, raising=False)
router: Final = Router(
model_list=[_deployment(HIGH_USAGE_DEPLOYMENT_ID)],
routing_strategy="usage-based-routing",
)
handler: Final = LowestTPMLoggingHandler(router_cache=router.cache, routing_args={})
line_item_kwargs: Final = {
"call_type": "acompletion",
"litellm_params": {
"batch_parent_id": "batch-1",
"metadata": {"model_group": MODEL_GROUP},
"model_info": {"id": HIGH_USAGE_DEPLOYMENT_ID},
},
}
response_obj: Final = {"usage": {"total_tokens": 600}}
handler.log_success_event(line_item_kwargs, response_obj, None, None)
await handler.async_log_success_event(line_item_kwargs, response_obj, None, None)
moved: Final = sorted(
f"{key}={router.cache.in_memory_cache.cache_dict[key]}"
for key in router.cache.in_memory_cache.cache_dict
if ":tpm:" in key or ":rpm:" in key
)
assert moved == []

View file

@ -1265,6 +1265,53 @@ def test_sync_deployment_callback_on_success_skips_batch_retrieves(
== expected_successes
)
@pytest.mark.asyncio
async def test_deployment_callbacks_skip_batch_line_items(monkeypatch: pytest.MonkeyPatch):
"""
Batch line-item callbacks carry call_type=acompletion plus
litellm_params.batch_parent_id, so the call-type-only batch guards do not fire.
They describe historical batch traffic already reported by the aggregate
aretrieve_batch event and must not consume live TPM/RPM quota.
"""
monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", True, raising=False)
router = litellm.Router(
model_list=[
{
"model_name": _BATCH_GROUP,
"litellm_params": {"model": _BATCH_DEPLOYMENT_MODEL, "api_base": _BATCH_API_BASE, "api_key": "sk-fake"},
"model_info": {"id": "batch-dep"},
}
]
)
now = datetime.now()
line_item_kwargs = {
"call_type": "acompletion",
"standard_logging_object": {"total_tokens": _BATCH_TOKENS_PER_ROW},
"litellm_params": {
"batch_parent_id": _BATCH_ID,
"metadata": {"model_group": _BATCH_GROUP, "deployment": _BATCH_DEPLOYMENT_MODEL},
"model_info": {"id": "batch-dep"},
},
}
await router.deployment_callback_on_success(
kwargs=line_item_kwargs, completion_response=None, start_time=now, end_time=now
)
sync_key = router.sync_deployment_callback_on_success(
kwargs=line_item_kwargs, completion_response=None, start_time=now, end_time=now
)
await router.async_deployment_callback_on_failure(
kwargs=line_item_kwargs, completion_response=None, start_time=now, end_time=now
)
assert sync_key is None
assert (
get_deployment_successes_for_current_minute(litellm_router_instance=router, deployment_id="batch-dep") == 0
)
assert await _moved_routing_counters(router) == []
assert await _router_usage_keys(router) == []
_ROUTING_STRATEGY_CACHE_MARKERS = ("_map", "_request_count", ":tpm:", ":rpm:")
@ -1352,6 +1399,81 @@ async def test_arouter_aretrieve_batch_does_not_feed_routing_strategies(
assert moved_counters == []
_BATCH_INPUT_JSONL = "\n".join(
json.dumps(
{
"custom_id": f"row-{row}",
"method": "POST",
"url": "/v1/chat/completions",
"body": {"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]},
}
)
for row in range(_BATCH_ROWS)
)
def _mock_batch_input_file(respx_mock):
respx_mock.get(f"{_BATCH_API_BASE}/files/file-in-1/content").mock(
return_value=httpx.Response(200, text=_BATCH_INPUT_JSONL)
)
async def _line_item_payloads(collector: "_BatchPayloadCollector", minimum: int, timeout: float = 5.0) -> list:
loop = asyncio.get_event_loop()
deadline = loop.time() + timeout
while loop.time() < deadline:
line_items = [p for p in collector.payloads if p and p.get("call_type") == "acompletion"]
if len(line_items) >= minimum:
return line_items
await asyncio.sleep(0.05)
raise AssertionError(f"expected at least {minimum} batch line-item payloads, saw {len(collector.payloads)} total")
@pytest.mark.parametrize(
"routing_strategy",
[
"usage-based-routing",
"usage-based-routing-v2",
"latency-based-routing",
"cost-based-routing",
"least-busy",
],
)
@pytest.mark.asyncio
async def test_arouter_aretrieve_batch_line_items_do_not_feed_routing_strategies(
monkeypatch: pytest.MonkeyPatch, routing_strategy: str
):
"""
With the line-item flag on, retrieving a completed batch also emits one child
callback per JSONL line (call_type=acompletion + litellm_params.batch_parent_id,
inheriting the parent deployment's model_info). Those children are historical
batch traffic: like the aggregate poll, they must not move the counters that
decide where the next live chat request goes.
"""
from litellm.batches.batch_line_item_logging import batch_line_item_claim_cache
batch_line_item_claim_cache.in_memory_cache.flush_cache()
monkeypatch.setattr(litellm, "store_batch_line_items_in_callbacks", True, raising=False)
collector = _BatchPayloadCollector()
monkeypatch.setattr(litellm, "callbacks", [collector])
monkeypatch.setattr(litellm, "input_callback", [])
router = _batch_fan_out_router(routing_strategy)
with respx.mock(assert_all_called=True) as respx_mock:
_mock_batch_provider(respx_mock)
_mock_batch_input_file(respx_mock)
respx_mock.get(f"{_UNRELATED_BATCH_API_BASE}/batches/{_BATCH_ID}").mock(
return_value=httpx.Response(404, json=_BATCH_NOT_FOUND)
)
response = await router.aretrieve_batch(batch_id=_BATCH_ID)
line_items = await _line_item_payloads(collector, minimum=_BATCH_ROWS)
moved_counters = await _moved_routing_counters(router)
assert response.id == _BATCH_ID
assert len(line_items) == _BATCH_ROWS
assert moved_counters == []
@pytest.mark.asyncio
async def test_arouter_aretrieve_file_content():
"""

View file

@ -29867,6 +29867,11 @@ export interface components {
scheduled_job_stagger?: components["schemas"]["ScheduledJobStaggerSettings"] | null;
/** @description Daily check of the spend LiteLLM captured against the provider's own bill (OpenAI via OPENAI_ADMIN_KEY). Publishes litellm_spend_capture_rate per provider and alerts when the ratio over the lookback window falls under the threshold (default 0.9). Off unless set. */
spend_capture_rate_check?: components["schemas"]["SpendCaptureRateCheckSettings"] | null;
/**
* Store Batch Line Items In Callbacks
* @description If True, a completed batch logged via aretrieve_batch also emits one callback event per JSONL line item (request paired with its response or error). The aggregate batch callback is unchanged. Default is False.
*/
store_batch_line_items_in_callbacks?: boolean | null;
/**
* Store Model In Db
* @description If True, models and config are stored in and loaded from the database. Default is False.