mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 1b02cb64fb into 9cdedf81cd
This commit is contained in:
commit
980eab30c7
54 changed files with 3757 additions and 120 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
518
litellm/batches/batch_line_item_logging.py
Normal file
518
litellm/batches/batch_line_item_logging.py
Normal 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
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
|
|
@ -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``.
|
||||
|
||||
|
|
|
|||
|
|
@ -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__,
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
||||
|
|
|
|||
|
|
@ -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)),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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.",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
303
tests/integration/spend/test_batch_line_item_callbacks.py
Normal file
303
tests/integration/spend/test_batch_line_item_callbacks.py
Normal 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
|
||||
]
|
||||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
1375
tests/unit/batches/test_batch_line_item_logging.py
Normal file
1375
tests/unit/batches/test_batch_line_item_logging.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -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
|
||||
|
|
|
|||
419
tests/unit/integrations/test_batch_line_item_metering_sinks.py
Normal file
419
tests/unit/integrations/test_batch_line_item_metering_sinks.py
Normal 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)
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 == {}
|
||||
|
|
|
|||
|
|
@ -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 == {}
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
5
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
5
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue