fix(fusion): retain latest budget reservation safeguards

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Moe Khalil 2026-09-22 03:01:36 +00:00
commit 7dbc21db87
35 changed files with 2602 additions and 325 deletions

View file

@ -6267,6 +6267,63 @@
],
"title": "Spend update queue sizes (litellm_<queue>_size)",
"type": "timeseries"
},
{
"datasource": {
"type": "prometheus",
"uid": "${DS_PROMETHEUS}"
},
"description": "Requests that carried usage but were logged at $0 on a model whose pricing entry has a non-zero rate, by requested model and reason",
"fieldConfig": {
"defaults": {
"color": {
"mode": "palette-classic"
},
"custom": {
"drawStyle": "line",
"fillOpacity": 10,
"lineWidth": 1,
"showPoints": "never",
"spanNulls": false
},
"unit": "short"
},
"overrides": []
},
"gridPos": {
"h": 8,
"w": 12,
"x": 0,
"y": 430
},
"id": 110,
"options": {
"legend": {
"calcs": [],
"displayMode": "list",
"placement": "bottom",
"showLegend": true
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"targets": [
{
"datasource": {
"type": "prometheus",
"uid": "${DS_PROMETHEUS}"
},
"editorMode": "code",
"expr": "sum(rate(litellm_zero_cost_requests_total[$__rate_interval])) by (requested_model, reason)",
"legendFormat": "{{requested_model}} / {{reason}}",
"range": true,
"refId": "A"
}
],
"title": "litellm_zero_cost_requests rate",
"type": "timeseries"
}
],
"preload": false,

View file

@ -948,7 +948,7 @@ def _extract_service_tier(source: object) -> str | None:
return None
def _get_usage_object(
def get_usage_object(
completion_response: object,
) -> Usage | None:
usage_obj: Final = cast(
@ -1336,7 +1336,7 @@ def completion_cost(
cache_creation_input_tokens: int | None = None
cache_read_input_tokens: int | None = None
audio_transcription_file_duration: float = 0.0
provider_usage_object: Final = _get_usage_object(completion_response=completion_response)
provider_usage_object: Final = get_usage_object(completion_response=completion_response)
cost_per_token_usage_object: Final[Usage | None] = (
_without_provider_stated_cost(provider_usage_object) if custom_pricing else provider_usage_object
)
@ -2033,6 +2033,45 @@ def _cost_map_model_info(model: str, custom_llm_provider: str | None) -> ModelIn
return None
def _raw_cost_map_entry(key: str) -> Mapping[str, object] | None:
raw_entry: Final = litellm.model_cost.get(key)
return raw_entry if isinstance(raw_entry, Mapping) else None
def pricing_entry_for_cost_calc(
model: str | None,
completion_response: object | None,
custom_llm_provider: str | None,
custom_pricing: bool | None,
base_model: str | None,
router_model_id: str | None,
region_name: str | None,
litellm_logging_obj: LitellmLoggingObject | None,
) -> tuple[str, Mapping[str, object]] | None:
deployment_entry: Final = _deployment_model_info(litellm_logging_obj, custom_pricing, router_model_id)
deployment_key: Final = router_model_id or model
if deployment_entry is not None and deployment_key is not None:
registered_entry: Final = _raw_cost_map_entry(router_model_id) if router_model_id is not None else None
return deployment_key, registered_entry or deployment_entry
selected_model: Final = _select_model_name_for_cost_calc(
model=model,
completion_response=completion_response,
base_model=base_model,
custom_pricing=custom_pricing,
custom_llm_provider=custom_llm_provider,
router_model_id=router_model_id,
region_name=region_name,
)
candidates: Final = (selected_model, _get_response_model(completion_response), model)
resolved: Final = next(
(info for info in (_cost_map_model_info(name, custom_llm_provider) for name in candidates if name) if info),
None,
)
if resolved is None:
return None
return resolved["key"], _raw_cost_map_entry(resolved["key"]) or resolved
def ocr_cost(
model: str,
custom_llm_provider: str | None,

View file

@ -10,6 +10,7 @@ import sys
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import replace
from datetime import datetime, timedelta
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeAlias, TypeVar, cast
from pydantic import BaseModel
@ -66,6 +67,7 @@ from litellm.types.proxy.carried_budget_state import (
from litellm.types.utils import (
StandardLoggingGuardrailInformation,
StandardLoggingPayload,
StandardLoggingZeroCostDiagnostic,
)
if TYPE_CHECKING:
@ -713,6 +715,15 @@ class PrometheusLogger(CustomLogger):
labelnames=self.get_labels_for_metric("litellm_requests_metric"),
)
self.litellm_zero_cost_requests_total = self._counter_factory(
name="litellm_zero_cost_requests_total",
documentation=(
"Requests that carried usage but were logged at $0 on a model whose pricing entry "
"has a non-zero rate, by reason (missing_pricing_key, pricing_not_applied, cost_calculation_error)"
),
labelnames=self.get_labels_for_metric("litellm_zero_cost_requests_total"),
)
# Cache metrics
self.litellm_cache_hits_metric = self._counter_factory(
name="litellm_cache_hits_metric",
@ -1410,6 +1421,11 @@ class PrometheusLogger(CustomLogger):
enum_values=enum_values,
label_context=label_context,
)
self._increment_zero_cost_requests_metric(
zero_cost_diagnostic=standard_logging_payload.get("zero_cost_diagnostic"),
enum_values=enum_values,
label_context=label_context,
)
# input, output, total token metrics
self._increment_token_metrics(
@ -1983,6 +1999,30 @@ class PrometheusLogger(CustomLogger):
amount=float(response_cost),
)
def _increment_zero_cost_requests_metric(
self,
zero_cost_diagnostic: StandardLoggingZeroCostDiagnostic | None,
enum_values: UserAPIKeyLabelValues,
label_context: PrometheusLabelFactoryContext,
) -> None:
if zero_cost_diagnostic is None:
return
supported_labels: Final = self.get_labels_for_metric("litellm_zero_cost_requests_total")
reason_label: Final = (
MappingProxyType({ZERO_COST_REASON_LABEL: zero_cost_diagnostic["reason"]})
if ZERO_COST_REASON_LABEL in supported_labels
else MappingProxyType({})
)
labels: Final = MappingProxyType(
{
**prometheus_label_factory(
supported_enum_labels=supported_labels, enum_values=enum_values, label_context=label_context
),
**reason_label,
}
)
self.litellm_zero_cost_requests_total.labels(**labels).inc()
@staticmethod
def _get_remaining_from_v3_rate_limit_headers(
standard_logging_payload: StandardLoggingPayload | None,
@ -2333,6 +2373,8 @@ class PrometheusLogger(CustomLogger):
team_alias=user_api_team_alias,
user=user_id,
model_id=standard_logging_payload.get("model_id", ""),
requested_model=standard_logging_payload.get("model_group"),
api_provider=standard_logging_payload.get("custom_llm_provider"),
custom_metadata_labels=get_custom_labels_from_metadata(
metadata=_get_combined_custom_metadata_from_standard_logging_payload(
standard_logging_payload=standard_logging_payload
@ -2345,6 +2387,11 @@ class PrometheusLogger(CustomLogger):
"litellm_llm_api_failed_requests_metric",
enum_values,
)
self._increment_zero_cost_requests_metric(
zero_cost_diagnostic=standard_logging_payload.get("zero_cost_diagnostic"),
enum_values=enum_values,
label_context=PrometheusLabelFactoryContext(enum_values),
)
self.set_llm_deployment_failure_metrics(kwargs)
await self._set_org_budget_metrics_after_api_request(
org_id=user_api_key_org_id,

View file

@ -359,6 +359,46 @@ def get_litellm_metadata_from_kwargs(kwargs: dict):
return {}
def _budget_reservation_on_auth_object(user_api_key_auth: object) -> object:
if isinstance(user_api_key_auth, Mapping):
return user_api_key_auth.get("budget_reservation")
return getattr(user_api_key_auth, "budget_reservation", None)
def budget_reservation_from_metadata(metadata: Mapping[str, object]) -> dict | None:
stamped: Final = metadata.get("user_api_key_budget_reservation")
if isinstance(stamped, dict):
return stamped
on_auth_object: Final = _budget_reservation_on_auth_object(metadata.get("user_api_key_auth"))
return on_auth_object if isinstance(on_auth_object, dict) else None
def _stamp_budget_reservation_callback_bound(litellm_params: Mapping[str, object], callback_bound: bool) -> None:
for metadata_variable_name in ("metadata", "litellm_metadata"):
metadata = litellm_params.get(metadata_variable_name)
if not isinstance(metadata, Mapping):
continue
budget_reservation = budget_reservation_from_metadata(metadata)
if budget_reservation is not None:
budget_reservation["callback_bound"] = callback_bound
def bind_budget_reservation_to_callbacks(litellm_params: Mapping[str, object]) -> None:
"""Mark the request's budget reservation as owned by the success callbacks of this call.
The proxy releases any reservation still unbound when the request ends; one bound here
is left for the cost callback, which may finish after the response has been sent. Bind
only where a success handler is guaranteed to run: a logging object merely existing is
not that, since the proxy builds one for every route before calling anything.
"""
_stamp_budget_reservation_callback_bound(litellm_params, True)
def unbind_budget_reservation_from_callbacks(litellm_params: Mapping[str, object]) -> None:
"""Hand a failed call's reservation back to the request-end release: failure handlers never settle it."""
_stamp_budget_reservation_callback_bound(litellm_params, False)
def reconstruct_model_name(
model_name: str,
custom_llm_provider: str | None,

View file

@ -50,6 +50,8 @@ from litellm.cost_calculator import (
RealtimeAPITokenUsageProcessor,
ResponsesWebSocketTokenUsageProcessor,
_select_model_name_for_cost_calc,
get_usage_object,
pricing_entry_for_cost_calc,
)
from litellm.exceptions import (
BudgetExceededError,
@ -89,6 +91,10 @@ from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
from litellm.litellm_core_utils.llm_cost_calc.usage_object_transformation import (
InteractionsUsageObjectTransformation,
)
from litellm.litellm_core_utils.llm_cost_calc.zero_cost_diagnostic import (
diagnose_zero_cost,
zero_cost_warning,
)
from litellm.litellm_core_utils.logging_utils import (
truncate_base64_in_messages,
truncate_base64_in_messages_async,
@ -157,6 +163,7 @@ from litellm.types.utils import (
StandardLoggingPayloadStatusFields,
StandardLoggingPromptManagementMetadata,
StandardLoggingVectorStoreRequest,
StandardLoggingZeroCostDiagnostic,
TextCompletionResponse,
TranscriptionResponse,
Usage,
@ -614,6 +621,7 @@ class Logging(LiteLLMLoggingBaseClass):
self.truncated_messages_for_logging: str | list | dict | None = None # mutable-ok: logged messages shape
## TIME TO FIRST TOKEN LOGGING ##
self.completion_start_time: datetime.datetime | None = None
self.zero_cost_warned: bool = False
self._llm_caching_handler: LLMCachingHandler | None = None
# INITIAL LITELLM_PARAMS
@ -1764,11 +1772,6 @@ class Logging(LiteLLMLoggingBaseClass):
)
result_hidden_params: Final = getattr(priced_result, "_hidden_params", None) or MappingProxyType({})
result_additional_headers: Final = (
result_hidden_params.get("additional_headers")
if isinstance(result_hidden_params, dict)
else getattr(result_hidden_params, "additional_headers", None)
)
if isinstance(priced_result, (BaseModel, HttpxBinaryResponseContent)) and hasattr(
priced_result, "_hidden_params"
):
@ -1776,6 +1779,12 @@ class Logging(LiteLLMLoggingBaseClass):
if (
"response_cost" in hidden_params and hidden_params["response_cost"] is not None
): # use cost if already calculated
self._record_zero_cost_diagnostic(
priced_result,
hidden_params["response_cost"],
litellm_model_name=litellm_model_name,
router_model_id=router_model_id or hidden_params.get("model_id"),
)
return hidden_params["response_cost"]
elif router_model_id is None and "model_id" in hidden_params: # use model_id if not already set
router_model_id = hidden_params["model_id"]
@ -1787,18 +1796,7 @@ class Logging(LiteLLMLoggingBaseClass):
router_model_id = self.get_router_model_id()
## RESPONSE COST ##
spilled_over: Final = is_spilled_over_ptu_request(
model_info=_deployment_model_info(self.litellm_params if hasattr(self, "litellm_params") else None),
response_headers=self.model_call_details.get("response_headers"),
additional_headers=result_additional_headers,
)
custom_pricing: Final = (
False
if spilled_over
else use_custom_pricing_for_model(
litellm_params=(self.litellm_params if hasattr(self, "litellm_params") else None)
)
)
custom_pricing: Final = self._custom_pricing_for(priced_result)
prompt = self._prompt_for_cost_calculation()
@ -1850,9 +1848,18 @@ class Logging(LiteLLMLoggingBaseClass):
verbose_logger.debug("response_cost: %s", response_cost)
additional_response_cost: Final[object] = self.model_call_details.get("additional_response_cost")
if isinstance(additional_response_cost, (int, float)) and additional_response_cost > 0:
return (response_cost or 0.0) + additional_response_cost
return response_cost
total_response_cost: Final = (
(response_cost or 0.0) + additional_response_cost
if isinstance(additional_response_cost, (int, float)) and additional_response_cost > 0
else response_cost
)
self._record_zero_cost_diagnostic(
priced_result,
total_response_cost,
litellm_model_name=litellm_model_name,
router_model_id=router_model_id,
)
return total_response_cost
except Exception as e: # error calculating cost
debug_info = StandardLoggingModelCostFailureDebugInformation(
error_str=str(e),
@ -1866,9 +1873,108 @@ class Logging(LiteLLMLoggingBaseClass):
)
verbose_logger.debug("response_cost_failure_debug_information: %s", debug_info)
self.model_call_details["response_cost_failure_debug_information"] = debug_info
self._record_zero_cost_diagnostic(
priced_result,
None,
calculation_failed=True,
litellm_model_name=litellm_model_name,
router_model_id=router_model_id,
)
return None
def _record_zero_cost_diagnostic(
self,
result: object,
response_cost: float | None,
*,
calculation_failed: bool = False,
litellm_model_name: str | None = None,
router_model_id: str | None = None,
) -> None:
if response_cost is None and not calculation_failed:
return
if self.model_call_details.get("cache_hit") is True:
self.model_call_details["zero_cost_diagnostic"] = None
return
try:
finding: Final = self._zero_cost_finding(
result,
response_cost,
calculation_failed=calculation_failed,
litellm_model_name=litellm_model_name,
router_model_id=router_model_id,
)
except Exception as e: # noqa: BLE001 # the pricing helpers raise plain Exception and a diagnostic must never break cost tracking
verbose_logger.debug("zero_cost_diagnostic skipped: %s", e)
return
self.model_call_details["zero_cost_diagnostic"] = finding[0] if finding is not None else None
if finding is None or self.zero_cost_warned:
return
self.zero_cost_warned = True
verbose_logger.warning(finding[1])
def _zero_cost_finding(
self,
result: object,
response_cost: float | None,
*,
calculation_failed: bool,
litellm_model_name: str | None,
router_model_id: str | None,
) -> tuple[StandardLoggingZeroCostDiagnostic, str] | None:
metadata: Final = StandardLoggingPayloadSetup.merge_litellm_metadata(self.litellm_params)
if response_cost or is_unbilled_non_inference_call(self.call_type, metadata, result):
return None
usage: Final = get_usage_object(completion_response=result)
if usage is None:
return None
model: Final = litellm_model_name or self.model
custom_llm_provider: Final = self.model_call_details.get("custom_llm_provider")
pricing: Final = pricing_entry_for_cost_calc(
model=model,
completion_response=result,
custom_llm_provider=custom_llm_provider,
custom_pricing=self._custom_pricing_for(result),
base_model=_get_base_model_from_metadata(model_call_details=self.model_call_details),
router_model_id=router_model_id or self.get_router_model_id(),
region_name=_resolve_mantle_region_for_cost(
custom_llm_provider=custom_llm_provider,
litellm_params=self.model_call_details.get("litellm_params"),
),
litellm_logging_obj=self,
)
if pricing is None:
return None
diagnostic: Final = diagnose_zero_cost(
usage=usage, pricing_model=pricing[0], pricing_entry=pricing[1], calculation_failed=calculation_failed
)
if diagnostic is None:
return None
model_group: Final = metadata.get("model_group")
return diagnostic, zero_cost_warning(
diagnostic,
model_group=model_group if isinstance(model_group, str) else None,
model=model,
custom_llm_provider=custom_llm_provider,
usage=usage,
)
def _custom_pricing_for(self, result: object) -> bool:
litellm_params: Final = getattr(self, "litellm_params", None)
result_hidden_params: Final = getattr(result, "_hidden_params", None) or MappingProxyType({})
additional_headers: Final = (
result_hidden_params.get("additional_headers")
if isinstance(result_hidden_params, dict)
else getattr(result_hidden_params, "additional_headers", None)
)
spilled_over: Final = is_spilled_over_ptu_request(
model_info=_deployment_model_info(litellm_params),
response_headers=self.model_call_details.get("response_headers"),
additional_headers=additional_headers,
)
return False if spilled_over else use_custom_pricing_for_model(litellm_params=litellm_params)
def _prompt_for_cost_calculation(self) -> str:
"""
The raw input string is only priced directly for text-to-speech, which bills per character.
@ -2213,6 +2319,7 @@ class Logging(LiteLLMLoggingBaseClass):
self.model_call_details["response_cost"] = 0.0
elif "response_cost" in hidden_params:
self.model_call_details["response_cost"] = hidden_params["response_cost"]
self._record_zero_cost_diagnostic(logging_result, hidden_params["response_cost"])
elif (existing_cost := self.model_call_details.get("response_cost")) is not None and existing_cost != 0:
# Preserve response_cost if already calculated (e.g., by pass-through
# handlers like Gemini/Vertex which call completion_cost directly).
@ -5386,7 +5493,7 @@ def request_model_access_groups_from_litellm_params(litellm_params: Mapping[str,
"""Access groups the auth layer stamped onto this request, from whichever metadata field carries them.
Detached internal sub-calls only inherit the identity keys, so the auth object is the
fallback there, exactly as _get_budget_reservation_from_metadata does for reservations.
fallback there, exactly as budget_reservation_from_metadata does for reservations.
"""
for metadata_variable_name in ("metadata", "litellm_metadata"):
metadata = litellm_params.get(metadata_variable_name)
@ -6507,6 +6614,7 @@ def get_standard_logging_object_payload(
error_str=error_str,
error_information=error_information,
response_cost_failure_debug_info=kwargs.get("response_cost_failure_debug_information"),
zero_cost_diagnostic=kwargs.get("zero_cost_diagnostic"),
guardrail_information=metadata.get("standard_logging_guardrail_information", None),
standard_built_in_tools_params=standard_built_in_tools_params,
)
@ -6685,6 +6793,7 @@ def create_dummy_standard_logging_payload() -> StandardLoggingPayload:
response_cost=response_cost,
autorouter_savings=None,
response_cost_failure_debug_info=None,
zero_cost_diagnostic=None,
status="success",
total_tokens=int(DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT + DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT),
prompt_tokens=int(DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT),

View file

@ -0,0 +1,146 @@
from collections.abc import Mapping
from functools import reduce
from typing import Final
from pydantic import TypeAdapter, ValidationError
from typing_extensions import assert_never
from litellm.types.utils import StandardLoggingZeroCostDiagnostic, Usage
ZERO_COST_COUNTER_NAME: Final = "litellm_zero_cost_requests_total"
_TEXT_INPUT_RATE: Final = "input_cost_per_token"
_AUDIO_INPUT_RATE: Final = "input_cost_per_audio_token"
_TEXT_OUTPUT_RATE: Final = "output_cost_per_token"
_AUDIO_OUTPUT_RATE: Final = "output_cost_per_audio_token"
_RATE_KEY_MARKERS: Final = ("cost", "pricing")
_NESTED_PRICING: Final = TypeAdapter(Mapping[str, object] | tuple[object, ...])
_MAX_PRICING_DEPTH: Final = 4
def _audio_tokens(details: object) -> int:
audio_tokens: Final = getattr(details, "audio_tokens", None)
return audio_tokens if isinstance(audio_tokens, int) and audio_tokens > 0 else 0
def _tokens(value: object) -> int:
return value if isinstance(value, int) and value > 0 else 0
def used_pricing_keys(usage: Usage) -> tuple[str, ...]:
prompt_audio: Final = _audio_tokens(usage.prompt_tokens_details)
completion_audio: Final = _audio_tokens(usage.completion_tokens_details)
prompt_text: Final = _tokens(usage.prompt_tokens) - prompt_audio
completion_text: Final = _tokens(usage.completion_tokens) - completion_audio
components: Final = (
(_TEXT_INPUT_RATE, prompt_text),
(_AUDIO_INPUT_RATE, prompt_audio),
(_TEXT_OUTPUT_RATE, completion_text),
(_AUDIO_OUTPUT_RATE, completion_audio),
)
return tuple(key for key, count in components if count > 0)
def _nested_pricing(value: object) -> Mapping[str, object] | tuple[object, ...] | None:
try:
return _NESTED_PRICING.validate_python(value)
except ValidationError:
return None
def _is_rate_key(key: str) -> bool:
return any(marker in key for marker in _RATE_KEY_MARKERS)
def _rate_values(value: object) -> tuple[object, ...]:
nested: Final = _nested_pricing(value)
if isinstance(nested, Mapping):
return tuple(child for key, child in nested.items() if _is_rate_key(key))
if nested is None:
return (value,)
return nested
def _expand_rate_values(values: tuple[object, ...], _depth: int) -> tuple[object, ...]:
return tuple(nested for value in values for nested in _rate_values(value))
def _is_positive_number(value: object) -> bool:
return not isinstance(value, bool) and isinstance(value, (int, float)) and value > 0
def _declares_a_rate(pricing_entry: Mapping[str, object]) -> bool:
leaves: Final = reduce(_expand_rate_values, range(_MAX_PRICING_DEPTH), (pricing_entry,))
return any(_is_positive_number(leaf) for leaf in leaves)
def _is_explicit_zero(value: object) -> bool:
return not isinstance(value, bool) and isinstance(value, (int, float)) and value == 0
def diagnose_zero_cost(
usage: Usage,
pricing_model: str,
pricing_entry: Mapping[str, object],
calculation_failed: bool,
) -> StandardLoggingZeroCostDiagnostic | None:
used_keys: Final = used_pricing_keys(usage)
if not used_keys:
return None
missing_keys: Final = tuple(key for key in used_keys if pricing_entry.get(key) is None)
if not missing_keys and all(_is_explicit_zero(pricing_entry[key]) for key in used_keys):
return None
if not _declares_a_rate(pricing_entry):
return None
if calculation_failed:
return StandardLoggingZeroCostDiagnostic(
reason="cost_calculation_error", pricing_model=pricing_model, missing_pricing_keys=()
)
if missing_keys:
return StandardLoggingZeroCostDiagnostic(
reason="missing_pricing_key", pricing_model=pricing_model, missing_pricing_keys=missing_keys
)
return StandardLoggingZeroCostDiagnostic(
reason="pricing_not_applied", pricing_model=pricing_model, missing_pricing_keys=()
)
def _cause(diagnostic: StandardLoggingZeroCostDiagnostic) -> str:
reason: Final = diagnostic["reason"]
match reason:
case "missing_pricing_key":
return (
f"pricing entry '{diagnostic['pricing_model']}' has no {', '.join(diagnostic['missing_pricing_keys'])}. "
"Set the missing rate in the deployment's model_info or in the model cost map, "
"or set every rate to 0 to mark the model free"
)
case "pricing_not_applied":
return (
f"pricing entry '{diagnostic['pricing_model']}' declares non-zero rates for this usage, "
"but the cost calculator returned $0"
)
case "cost_calculation_error":
return (
f"cost calculation raised for pricing entry '{diagnostic['pricing_model']}', "
"see response_cost_failure_debug_information"
)
case _:
return assert_never(reason)
def zero_cost_warning(
diagnostic: StandardLoggingZeroCostDiagnostic,
*,
model_group: str | None,
model: str,
custom_llm_provider: str | None,
usage: Usage,
) -> str:
request: Final = (
f"model_group={model_group or model} model={model} provider={custom_llm_provider or 'unknown'} "
f"prompt_tokens={_tokens(usage.prompt_tokens)} completion_tokens={_tokens(usage.completion_tokens)}"
)
return (
f"Billable request priced at $0 and logged as such ({request}): {_cause(diagnostic)}. "
f'Counted in {ZERO_COST_COUNTER_NAME}{{reason="{diagnostic["reason"]}"}}'
)

View file

@ -650,6 +650,7 @@ async def user_api_key_auth_websocket_for_model(websocket: WebSocket, model: str
"type": "http",
"headers": scope_headers,
"path": ws_scope.get("path", ""),
"state": ws_scope.setdefault("state", {}), # mutable-ok: Starlette's socket state, shared with the request
}
for key in ("root_path", "app_root_path"):
if key in ws_scope:
@ -3086,31 +3087,30 @@ async def _reserve_budget_after_common_checks(
request: Request | None = None,
) -> None:
user_api_key_auth_obj.budget_reservation = None
if skip_budget_checks:
return
if general_settings.get("disable_budget_reservation") is True:
return
if not skip_budget_checks and general_settings.get("disable_budget_reservation") is not True:
from litellm.proxy.spend_tracking.budget_reservation import (
reserve_budget_for_request,
)
from litellm.proxy.spend_tracking.budget_reservation import (
reserve_budget_for_request,
)
user_api_key_auth_obj.budget_reservation = await reserve_budget_for_request(
request_body=request_data,
route=route,
llm_router=llm_router,
valid_token=user_api_key_auth_obj,
team_object=team_object,
user_object=user_object,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
end_user_id=end_user_id,
end_user_object=end_user_object,
apply_user_budget_to_team_keys=general_settings.get("apply_user_budget_to_team_keys") is True,
fail_closed_budget_enforcement=general_settings.get("fail_closed_budget_enforcement") is True,
raw_body=await read_raw_json_body(request=request),
)
user_api_key_auth_obj.budget_reservation = await reserve_budget_for_request(
request_body=request_data,
route=route,
llm_router=llm_router,
valid_token=user_api_key_auth_obj,
team_object=team_object,
user_object=user_object,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
end_user_id=end_user_id,
end_user_object=end_user_object,
apply_user_budget_to_team_keys=general_settings.get("apply_user_budget_to_team_keys") is True,
fail_closed_budget_enforcement=general_settings.get("fail_closed_budget_enforcement") is True,
raw_body=await read_raw_json_body(request=request),
)
if request is not None:
reservation: Final = user_api_key_auth_obj.budget_reservation
request.state.budget_reservation = reservation # rebind-ok: read by the release middleware
def _should_skip_budget_checks(

View file

@ -17,6 +17,7 @@ from litellm.constants import (
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
budget_reservation_from_metadata,
get_litellm_metadata_from_kwargs,
)
from litellm.litellm_core_utils.fusion_budget import (
@ -781,17 +782,7 @@ def _metadata_keys(metadata: object) -> tuple[str, ...]:
def _get_budget_reservation_from_metadata(metadata: dict) -> dict | None:
metadata_budget_reservation: Final = metadata.get("user_api_key_budget_reservation")
if isinstance(metadata_budget_reservation, dict):
return metadata_budget_reservation
user_api_key_auth_obj: Final = metadata.get("user_api_key_auth")
if user_api_key_auth_obj is None:
return None
if isinstance(user_api_key_auth_obj, dict):
budget_reservation: Final = user_api_key_auth_obj.get("budget_reservation")
return budget_reservation if isinstance(budget_reservation, dict) else None
return getattr(user_api_key_auth_obj, "budget_reservation", None)
return budget_reservation_from_metadata(metadata)
def _get_request_tags_for_cost_tracking(

View file

@ -0,0 +1,33 @@
from collections.abc import Awaitable, Callable, Mapping
from typing import Final
from starlette.types import ASGIApp, Receive, Scope, Send
_SCOPES_AUTH_STAMPS: Final = frozenset({"http", "websocket"})
class BudgetReservationReleaseMiddleware:
"""Releases the budget reservation auth made for a request once no callback owns it.
Auth stamps the reservation on the request or socket state; a call that starts
claims it for the cost callbacks, which settle it on success or failure. When the
response has been sent or the socket has closed and the reservation is still
unclaimed, nothing else ever would, so it is released here instead of pinning the
spend counter until its TTL.
"""
def __init__(self, app: ASGIApp, release: Callable[[Mapping[str, object]], Awaitable[None]]) -> None:
self.app = app
self.release = release
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
if scope["type"] not in _SCOPES_AUTH_STAMPS:
await self.app(scope, receive, send)
return
try:
await self.app(scope, receive, send)
finally:
state: Final = scope.get("state")
budget_reservation: Final = state.get("budget_reservation") if isinstance(state, Mapping) else None
if isinstance(budget_reservation, Mapping):
await self.release(budget_reservation)

View file

@ -47,6 +47,7 @@ from litellm.constants import (
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import (
bind_budget_reservation_to_callbacks,
get_metadata_variable_name_from_kwargs,
get_or_create_metadata_bucket,
)
@ -1637,6 +1638,7 @@ async def pass_through_request(
**kwargs,
)
)
bind_budget_reservation_to_callbacks(logging_obj.litellm_params)
## CUSTOM HEADERS - `x-litellm-*`
custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
@ -2561,6 +2563,7 @@ async def websocket_passthrough_request(
**success_kwargs,
)
)
bind_budget_reservation_to_callbacks(logging_obj.litellm_params)
# Call the proxy logging success hook
if proxy_logging_obj:
@ -2732,6 +2735,7 @@ async def _relay_passthrough_response_bytes(
**success_handler_kwargs,
)
)
bind_budget_reservation_to_callbacks(logging_obj.litellm_params)
def _extract_model_from_vertex_ai_setup(setup_response: Mapping[str, object]) -> str | None:

View file

@ -9,6 +9,7 @@ import httpx
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.asyncify import asyncify
from litellm.litellm_core_utils.core_helpers import bind_budget_reservation_to_callbacks
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.proxy._types import PassThroughEndpointLoggingResultValues
@ -218,6 +219,7 @@ class PassThroughStreamingHandler:
and response.status_code < 400
):
logging_scheduled = True
bind_budget_reservation_to_callbacks(litellm_logging_obj.litellm_params)
litellm_logging_obj._deferred_stream_complete_args = (_build_logging_coroutine(),)
except Exception as e:
verbose_proxy_logger.error("Error in chunk_processor: %s", e)
@ -250,6 +252,8 @@ class PassThroughStreamingHandler:
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(async_coroutine=_build_logging_coroutine())
except Exception as e:
verbose_proxy_logger.error("Error scheduling chunk_processor logging: %s", e)
else:
bind_budget_reservation_to_callbacks(litellm_logging_obj.litellm_params)
@staticmethod
async def _route_streaming_logging_to_handler(

View file

@ -654,6 +654,9 @@ from litellm.proxy.middleware.billable_request_metrics_middleware import (
BillableRequestMetricsMiddleware,
BillingRecorder,
)
from litellm.proxy.middleware.budget_reservation_release_middleware import (
BudgetReservationReleaseMiddleware,
)
from litellm.proxy.plugin_routes import (
register_plugins_from_config,
)
@ -729,7 +732,11 @@ from litellm.proxy.shutdown.scheduled_jobs import (
pause_scheduled_jobs,
stop_in_flight_scheduler_jobs,
)
from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start
from litellm.proxy.spend_tracking.budget_reservation import (
get_budget_window_start,
get_reserved_counter_keys,
release_unbound_budget_reservation,
)
from litellm.proxy.spend_tracking.daily_global_spend_rollup import (
run_scheduled_daily_global_spend_reconcile,
)
@ -2358,6 +2365,7 @@ app.add_middleware(
# it sees prisma_client as of the first request rather than import time.
sink_factory=lambda: gateway_request_accumulator if prisma_client is not None else None,
)
app.add_middleware(BudgetReservationReleaseMiddleware, release=release_unbound_budget_reservation)
app.add_middleware(InFlightRequestsMiddleware)
app.add_middleware(SecurityHeadersMiddleware)
@ -3244,8 +3252,6 @@ async def increment_fusion_model_access_group_spend_counters(
provider call preserves attribution while skipping any group already reserved
for the virtual Fusion model itself.
"""
from litellm.proxy.spend_tracking.budget_reservation import get_reserved_counter_keys
results: Final = await _prepare_model_access_group_spend_increments(
model_access_groups=model_access_groups,
response_cost=response_cost,

View file

@ -371,6 +371,7 @@ async def reserve_budget_for_request(
"reserved_cost": reservation_cost,
"entries": applied_entries,
"finalized": False,
"callback_bound": False,
"input_cost": min(float(input_cost or 0.0), reservation_cost),
"input_tokens": max(input_token_counts.values(), default=None),
}
@ -500,6 +501,19 @@ async def release_or_invalidate_budget_reservation(
budget_reservation["finalized"] = True
async def release_unbound_budget_reservation(budget_reservation: Mapping[str, object]) -> None:
"""Release a reservation no logging callback took ownership of, once the request ended.
A handler whose litellm call never builds a logging object (batch cancel, file
content, anything without the client decorator) runs no cost callback, so nothing
else would ever reconcile its reservation. A bound reservation is left alone: its
success or failure handler settles it, possibly after the response has been sent.
"""
if not isinstance(budget_reservation, dict) or budget_reservation.get("callback_bound") is True:
return
await release_or_invalidate_budget_reservation(budget_reservation=budget_reservation)
async def _get_budget_counters(
request_body: dict,
valid_token: UserAPIKeyAuth,

View file

@ -59,9 +59,17 @@ def setup(
}
supplied: Final = arguments.get("litellm_logging_obj")
if isinstance(supplied, Logging):
return CallSetup(supplied, arguments)
return _claim_budget_reservation(CallSetup(supplied, arguments), asynchronous)
logger, prepared = function_setup(call_type, Rules(), start_time, *args, is_async_call=asynchronous, **arguments)
return CallSetup(logger, prepared)
return _claim_budget_reservation(CallSetup(logger, prepared), asynchronous)
def _claim_budget_reservation(call_setup: CallSetup, asynchronous: bool) -> CallSetup:
from litellm.litellm_core_utils.core_helpers import bind_budget_reservation_to_callbacks
if asynchronous and not is_internal_call():
bind_budget_reservation_to_callbacks(call_setup.logger.litellm_params)
return call_setup
def check_limits(kwargs: Mapping[str, object]) -> None:
@ -96,6 +104,9 @@ def finalize(
class LoggingSurface(Protocol):
@property
def litellm_params(self) -> Mapping[str, object]: ...
def update_from_kwargs(
self,
kwargs: dict[str, object],
@ -236,8 +247,12 @@ def sync_success_for_async_call(
def failure_handler(
logger: LoggingSurface, error: Exception, start: datetime.datetime, end: datetime.datetime, asynchronous: bool
) -> Coroutine[object, object, None] | None:
from litellm.litellm_core_utils.core_helpers import unbind_budget_reservation_from_callbacks
trace: Final = "".join(traceback.format_exception(error))
if asynchronous:
if not is_internal_call():
unbind_budget_reservation_from_callbacks(logger.litellm_params)
return logger.async_failure_handler(error, trace, start, end)
logger.failure_handler(error, trace, start, end)
return None

View file

@ -131,6 +131,7 @@ EXCEPTION_STATUS: Final = "exception_status"
EXCEPTION_CLASS: Final = "exception_class"
RATE_LIMIT_CATEGORY: Final = "rate_limit_category"
RATE_LIMIT_TYPE: Final = "rate_limit_type"
ZERO_COST_REASON_LABEL: Final = "reason"
STATUS_CODE: Final = "status_code"
EXCEPTION_LABELS: Final = [EXCEPTION_STATUS, EXCEPTION_CLASS]
LATENCY_BUCKETS: Final = (
@ -279,6 +280,7 @@ DEFINED_PROMETHEUS_METRICS = Literal[
"litellm_guardrail_latency_seconds",
"litellm_guardrail_errors_total",
"litellm_guardrail_requests_total",
"litellm_zero_cost_requests_total",
# Cache metrics
"litellm_cache_hits_metric",
"litellm_cache_misses_metric",
@ -590,6 +592,14 @@ class PrometheusMetricLabels:
UserAPIKeyLabelNames.SERVICE_TIER.value,
]
litellm_zero_cost_requests_total = (
UserAPIKeyLabelNames.REQUESTED_MODEL.value,
UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value,
UserAPIKeyLabelNames.MODEL_ID.value,
UserAPIKeyLabelNames.API_PROVIDER.value,
ZERO_COST_REASON_LABEL,
)
litellm_input_tokens_metric = [
UserAPIKeyLabelNames.END_USER.value,
UserAPIKeyLabelNames.API_KEY_HASH.value,

View file

@ -3216,6 +3216,15 @@ class StandardLoggingModelCostFailureDebugInformation(TypedDict, total=False):
custom_pricing: bool | None
ZeroCostReason = Literal["missing_pricing_key", "pricing_not_applied", "cost_calculation_error"]
class StandardLoggingZeroCostDiagnostic(TypedDict):
reason: ReadOnly[ZeroCostReason]
pricing_model: ReadOnly[str]
missing_pricing_keys: ReadOnly[tuple[str, ...]]
class StandardLoggingPayloadErrorInformation(TypedDict, total=False):
error_code: str | None
error_class: str | None
@ -3534,6 +3543,7 @@ class StandardLoggingPayload(ClassifierAudit):
autorouter_savings_estimate: ReadOnly[Mapping[str, JsonValue] | None]
autorouter_baseline_observation: ReadOnly[str | None]
response_cost_failure_debug_info: StandardLoggingModelCostFailureDebugInformation | None
zero_cost_diagnostic: NotRequired[ReadOnly[StandardLoggingZeroCostDiagnostic | None]]
status: StandardLoggingPayloadStatus
status_fields: StandardLoggingPayloadStatusFields
custom_llm_provider: str | None

View file

@ -81,7 +81,11 @@ from litellm.constants import (
PROVIDERS_THAT_AUTHENTICATE_ON_PROVIDER_INFO,
TOOL_CHOICE_OBJECT_TOKEN_COUNT,
)
from litellm.litellm_core_utils.core_helpers import normalize_drop_params
from litellm.litellm_core_utils.core_helpers import (
bind_budget_reservation_to_callbacks,
normalize_drop_params,
unbind_budget_reservation_from_callbacks,
)
from litellm.litellm_core_utils.fallback_generalizations import (
match_capability_generalizations,
match_fill_missing_generalizations,
@ -1880,6 +1884,8 @@ def client(original_function):
# Type assertion: logging_obj is guaranteed to be non-None after function_setup
assert logging_obj is not None, "logging_obj should not be None after function_setup"
if not _is_litellm_internal_call:
bind_budget_reservation_to_callbacks(logging_obj.litellm_params)
kwargs["litellm_logging_obj"] = logging_obj
modified_kwargs: Final = await async_pre_call_deployment_hook(kwargs, call_type)
@ -2081,6 +2087,7 @@ def client(original_function):
# the failure hook ran, so a slow callback doesn't inflate the reported duration.
end_time = _deployment_call_end_time if _deployment_call_end_time is not None else datetime.datetime.now() # noqa: DTZ005 # matches the naive datetimes this whole function already times start_time/end_time with
if logging_obj and not _is_litellm_internal_call:
unbind_budget_reservation_from_callbacks(logging_obj.litellm_params)
try:
logging_obj.failure_handler(
e, traceback_exception, start_time, end_time

View file

@ -89,9 +89,6 @@
- {id: llm.chat_completions.together_ai.multi_turn.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together tool result round trip"}
- {id: llm.chat_completions.together_ai.basic.nonstream.cost_logged, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: basic, streaming: nonstream, assertions: [cost_logged], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together cost header and spend row match the registry price"}
- {id: llm.chat_completions.together_ai.thinking.nonstream.effort_none_disables, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: nonstream, assertions: [effort_none_disables], source: "llm_translation/test_together_ai_e2e.py", rationale: "reasoning_effort=none maps to Together's reasoning disable toggle on hybrid models"}
- {id: llm.chat_completions.xiaomi_mimo.basic.nonstream.cost_logged, module: llm, tier: P1, subject_endpoint: chat_completions, route: xiaomi_mimo, capability: basic, streaming: nonstream, assertions: [cost_logged], source: "llm_translation/test_xiaomi_mimo_e2e.py", rationale: "Native MiMo v2.6 rows price the cost header and spend row from the cost map"}
- {id: llm.chat_completions.xiaomi_mimo.thinking.stream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: xiaomi_mimo, capability: thinking, streaming: stream, assertions: [works], source: "llm_translation/test_xiaomi_mimo_e2e.py", rationale: "MiMo reasoning deltas stream as reasoning_content"}
- {id: llm.chat_completions.xiaomi_mimo.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: xiaomi_mimo, capability: tool_use, streaming: nonstream, assertions: [works], source: "llm_translation/test_xiaomi_mimo_e2e.py", rationale: "MiMo tool calls are not dropped"}
- {id: llm.chat_completions.together_ai.structured_output.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: structured_output, streaming: nonstream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "response_format json_schema reaches Together and constrains the reply"}
- {id: llm.chat_completions.together_ai.prompt_cache_5m.nonstream.cost_logged, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: prompt_cache_5m, streaming: nonstream, assertions: [cache_hit, cost_logged], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together prefix-cache reads bill at cache_read_input_token_cost, not full input price"}
- {id: llm.messages.together_ai.basic.stream.works, module: llm, tier: P1, subject_endpoint: messages, route: together_ai, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together over /v1/messages streaming"}

View file

@ -1,214 +0,0 @@
"""Live e2e: Xiaomi MiMo v2.6 through the gateway on /chat/completions.
Both native ``xiaomi_mimo/`` v2.6 rows (pro and flash) are registered via
``/model/new`` and driven against Xiaomi's own endpoint. What the gateway owes
us is that the reasoning chain surfaces as ``reasoning_content``, tool calls
survive translation, and the cost header plus spend row follow the proxy's own
cost-map price for the row (read back from ``/model/info``, never pinned here).
Requires XIAOMI_MIMO_API_KEY on the proxy; no skip gate.
"""
from __future__ import annotations
from typing import Final
import pytest
from e2e_config import unique_marker
from e2e_http import StreamingResponse, require_successful_call, unwrap
from lifecycle import ResourceManager
from models import (
ChatBody,
ChatMessage,
ChatResponse,
ChatTool,
ChatToolFunction,
CostMapEntry,
LiteLLMParamsBody,
OutMessage,
SpendLogRow,
)
from passthrough_client import PassthroughClient
from pydantic import BaseModel
pytestmark = pytest.mark.e2e
BACKENDS: Final = ("xiaomi_mimo/mimo-v2.6-pro", "xiaomi_mimo/mimo-v2.6-flash")
ARITHMETIC_PROMPT = "What is 17 + 26? Answer with just the number."
WEATHER_PROMPT = "What is the weather in Paris? Use the tool."
COUNTING_PROMPT = "Count from 1 to 50, one number per line."
WEATHER_TOOL = ChatTool(
function=ChatToolFunction(
name="get_weather",
description="Get the current weather for a location.",
parameters={
"type": "object",
"properties": {"location": {"type": "string"}},
"required": ["location"],
},
)
)
class _WeatherArgs(BaseModel):
location: str
class _StreamDelta(BaseModel):
content: str | None = None
reasoning_content: str | None = None
class _StreamChoice(BaseModel):
delta: _StreamDelta | None = None
class _StreamChunk(BaseModel):
choices: list[_StreamChoice] = []
def _approx_equal(actual: float, expected: float) -> bool:
return abs(actual - expected) <= max(1e-9, abs(expected) * 1e-2)
@pytest.fixture(scope="module")
def registry(client: PassthroughClient) -> dict[str, CostMapEntry]:
return client.proxy.model_cost_map()
def _register(client: PassthroughClient, resources: ResourceManager, backend: str) -> tuple[str, str]:
model = f"e2e-xiaomi-{unique_marker()}"
model_id = client.proxy.create_model(
model, LiteLLMParamsBody(model=backend, api_key="os.environ/XIAOMI_MIMO_API_KEY")
)
resources.defer(lambda: client.proxy.delete_model(model_id))
return model, resources.key()
def _message(response: ChatResponse) -> OutMessage:
assert response.choices, f"Xiaomi returned no choices: {response}"
message = response.choices[0].message
assert message is not None, f"Xiaomi choice has no message: {response}"
return message
def _deltas(result: StreamingResponse) -> list[_StreamDelta]:
require_successful_call(result)
assert result.is_streaming, f"response was not streamed: {result.headers}"
assert not result.stream_error, f"stream errored: {result.stream_error}"
assert result.stream_done, f"stream never reached [DONE]: {result.stream_events[-3:]}"
return [
choice.delta
for event in result.stream_events
for choice in _StreamChunk.model_validate_json(event).choices
if choice.delta is not None
]
@pytest.mark.parametrize("backend", BACKENDS)
class TestXiaomiMimoChatCompletions:
@pytest.mark.covers("llm.chat_completions.xiaomi_mimo.basic.nonstream.cost_logged")
def test_cost_header_and_spend_row_match_the_registry_price(
self,
client: PassthroughClient,
resources: ResourceManager,
registry: dict[str, CostMapEntry],
backend: str,
) -> None:
price = registry.get(backend)
assert price is not None, f"{backend} has no row in the proxy's cost map, so native calls would bill $0"
assert price.litellm_provider == "xiaomi_mimo", f"{backend} is filed under the wrong provider: {price}"
assert price.input_cost_per_token and price.output_cost_per_token, f"{backend} carries no price: {price}"
model, key = _register(client, resources, backend)
result = client.proxy.transport.send(
"/chat/completions",
headers=client.proxy.transport.bearer(key),
json=ChatBody(
model=model,
messages=[ChatMessage(role="user", content=f"{ARITHMETIC_PROMPT} {unique_marker()}")],
max_tokens=1024,
),
)
require_successful_call(result)
response = ChatResponse.model_validate_json(result.body)
message = _message(response)
assert message.content and "43" in message.content, f"answer lost: {message}"
assert message.reasoning_content, f"{backend} reasons, but no reasoning_content came back: {message}"
usage = response.usage
assert usage is not None and usage.prompt_tokens and usage.completion_tokens, (
f"response carries no usage, so the cost cannot be real: {result.body[:300]}"
)
header_cost = result.response_cost
assert header_cost is not None and header_cost > 0, (
f"x-litellm-response-cost header missing or non-positive: {result.headers}"
)
cached = (usage.prompt_tokens_details.cached_tokens or 0) if usage.prompt_tokens_details else 0
expected = (
(usage.prompt_tokens - cached) * price.input_cost_per_token
+ cached * (price.cache_read_input_token_cost or 0.0)
+ usage.completion_tokens * price.output_cost_per_token
)
assert _approx_equal(header_cost, expected), (
f"header cost {header_cost} disagrees with the registry price for {backend} at {usage}: expected {expected}"
)
def _priced(rows: list[SpendLogRow]) -> bool:
return any(row.spend is not None and row.spend > 0 for row in rows)
rows = client.proxy.poll_logs_for_key(key, predicate=_priced)
priced = [row for row in rows if row.spend is not None and row.spend > 0]
assert priced, f"no priced spend row landed for key {key}; got {rows}"
row = priced[0]
assert row.custom_llm_provider == "xiaomi_mimo", f"spend row misattributed: {row}"
assert row.spend is not None and _approx_equal(row.spend, header_cost), (
f"logged spend {row.spend} disagrees with the x-litellm-response-cost header {header_cost}"
)
@pytest.mark.covers("llm.chat_completions.xiaomi_mimo.thinking.stream.works")
def test_reasoning_and_answer_stream_as_deltas(
self, client: PassthroughClient, resources: ResourceManager, backend: str
) -> None:
model, key = _register(client, resources, backend)
deltas = _deltas(
client.proxy.chat_stream(
key,
ChatBody(
model=model,
messages=[ChatMessage(role="user", content=COUNTING_PROMPT)],
max_tokens=2048,
stream=True,
),
)
)
reasoning = "".join(delta.reasoning_content or "" for delta in deltas)
content = "".join(delta.content or "" for delta in deltas)
assert reasoning, f"stream carried no reasoning_content deltas: {deltas[:5]}"
assert "50" in content, f"streamed answer lost: {content[:300]!r}"
@pytest.mark.covers("llm.chat_completions.xiaomi_mimo.tool_use.nonstream.works")
def test_tool_call_is_returned(self, client: PassthroughClient, resources: ResourceManager, backend: str) -> None:
model, key = _register(client, resources, backend)
message = _message(
unwrap(
client.proxy.chat(
key,
ChatBody(
model=model,
messages=[ChatMessage(role="user", content=WEATHER_PROMPT)],
tools=[WEATHER_TOOL],
max_tokens=1024,
),
)
)
)
assert message.tool_calls, f"{backend} dropped the tool call: {message}"
call = message.tool_calls[0]
assert call.id, f"tool call carries no id, so a tool result cannot answer it: {call}"
assert call.function.name == "get_weather", f"wrong tool called: {call}"
assert call.function.arguments, f"tool call carries no arguments: {call}"
args = _WeatherArgs.model_validate_json(call.function.arguments)
assert "paris" in args.location.lower(), f"tool arguments lost the location: {args}"

View file

@ -190,6 +190,18 @@
"tests/integration/providers/test_fal_ai_chat_wire.py::test_fal_moondream3_chat_sends_prompt_image_and_reasoning": [
"other.provider_wire.fal_ai.moondream3_chat_query_wire_and_token_pricing"
],
"tests/integration/providers/test_xiaomi_mimo_wire.py::test_xiaomi_mimo_nonstream_surfaces_reasoning_and_charges_registry_price[mimo-v2.6-pro]": [
"other.provider_wire.xiaomi_mimo.reasoning_content_and_registry_pricing"
],
"tests/integration/providers/test_xiaomi_mimo_wire.py::test_xiaomi_mimo_nonstream_surfaces_reasoning_and_charges_registry_price[mimo-v2.6-flash]": [
"other.provider_wire.xiaomi_mimo.reasoning_content_and_registry_pricing"
],
"tests/integration/providers/test_xiaomi_mimo_wire.py::test_xiaomi_mimo_stream_delivers_reasoning_then_answer_deltas": [
"other.provider_wire.xiaomi_mimo.reasoning_and_answer_stream_as_deltas"
],
"tests/integration/providers/test_xiaomi_mimo_wire.py::test_xiaomi_mimo_tool_call_is_forwarded_and_returned": [
"other.provider_wire.xiaomi_mimo.tool_call_survives_translation"
],
"tests/integration/providers/test_fal_ai_video_wire.py::test_fal_h3_video_create_uses_canonical_body_and_status_path": [
"other.provider_wire.fal_ai.video_queue_create_status_and_content_download"
],

View file

@ -0,0 +1,258 @@
import json
import uuid
from collections.abc import Mapping
from pathlib import Path
from typing import Final
import pytest
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.wire import Reply, Request, wire_server
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter
_BACKENDS: Final = ("mimo-v2.6-pro", "mimo-v2.6-flash")
_API_KEY: Final = "synthetic-xiaomi-key"
_ARITHMETIC_PROMPT: Final = "What is 17 + 26? Answer with just the number."
_WEATHER_PROMPT: Final = "What is the weather in Paris? Use the tool."
_COUNTING_PROMPT: Final = "Count from 1 to 5, one number per line."
_WEATHER_TOOL: Final[JsonValue] = {
"type": "function",
"function": {
"name": "get_weather",
"description": "Get the current weather for a city",
"parameters": {
"type": "object",
"properties": {"city": {"type": "string"}},
"required": ["city"],
},
},
}
_COST_MAP_PATH: Final = Path(__file__).resolve().parents[3] / "model_prices_and_context_window.json"
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_COST_MAP: Final = TypeAdapter(dict[str, dict[str, object]])
class _Delta(BaseModel):
model_config = ConfigDict(extra="ignore")
content: str | None = None
reasoning_content: str | None = None
class _Choice(BaseModel):
model_config = ConfigDict(extra="ignore")
delta: _Delta
finish_reason: str | None = None
class _Chunk(BaseModel):
model_config = ConfigDict(extra="ignore")
id: str
choices: tuple[_Choice, ...]
def _catalog_cost(backend: str, field: str) -> float:
cost_map: Final = _COST_MAP.validate_json(_COST_MAP_PATH.read_bytes())
cost_value: Final = cost_map[f"xiaomi_mimo/{backend}"][field]
assert isinstance(cost_value, (int, float))
return float(cost_value)
def _approx(value: float) -> object:
return pytest.approx(value, rel=1e-6) # pyright: ignore[reportUnknownMemberType] # pytest lacks typed approx stubs
def _completion(identity: str, backend: str, message: Mapping[str, object], finish: str) -> bytes:
return json.dumps(
{
"id": identity,
"object": "chat.completion",
"created": 1,
"model": backend,
"choices": [{"index": 0, "message": message, "finish_reason": finish}],
"usage": {"prompt_tokens": 23, "completion_tokens": 41, "total_tokens": 64},
}
).encode()
def _frame(identity: str, backend: str, delta: Mapping[str, object], finish: str | None = None) -> bytes:
value: Final = {
"id": identity,
"object": "chat.completion.chunk",
"created": 1,
"model": backend,
"choices": [{"index": 0, "delta": delta, "finish_reason": finish}],
}
return b"data: " + json.dumps(value).encode() + b"\n\n"
def _assert_provider_request(request: Request, backend: str, prompt: str) -> dict[str, JsonValue]:
assert request.method == "POST"
assert request.target == "/chat/completions"
assert request.headers["authorization"] == f"Bearer {_API_KEY}"
assert request.headers["content-type"] == "application/json"
body: Final = _JSON_OBJECT.validate_json(request.body)
assert body["model"] == backend
assert body["messages"] == [{"role": "user", "content": prompt}]
return body
@pytest.mark.covers("other.provider_wire.xiaomi_mimo.reasoning_content_and_registry_pricing")
@pytest.mark.parametrize("backend", _BACKENDS)
def test_xiaomi_mimo_nonstream_surfaces_reasoning_and_charges_registry_price(gateway: Gateway, backend: str) -> None:
identity: Final = f"xiaomi-cost-{uuid.uuid4().hex}"
def respond(request: Request) -> Reply:
body: Final = _assert_provider_request(request, backend, _ARITHMETIC_PROMPT)
assert body["max_tokens"] == 256
assert "max_completion_tokens" not in body
return Reply(
body=_completion(
identity,
backend,
{"role": "assistant", "content": "43", "reasoning_content": "17 plus 26 is 43."},
"stop",
)
)
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"xiaomi_mimo/{backend}", api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": _ARITHMETIC_PROMPT}],
"max_completion_tokens": 256,
},
)
assert response.status_code == 200, response.text
payload: Final = _JSON_OBJECT.validate_json(response.content)
assert payload["id"] == identity
assert payload["choices"] == [
{
"finish_reason": "stop",
"index": 0,
"message": {
"role": "assistant",
"content": "43",
"reasoning_content": "17 plus 26 is 43.",
"provider_specific_fields": {"refusal": None},
},
"provider_specific_fields": {},
}
]
assert payload["usage"] == {"prompt_tokens": 23, "completion_tokens": 41, "total_tokens": 64}
expected_cost: Final = 23 * _catalog_cost(backend, "input_cost_per_token") + 41 * _catalog_cost(
backend, "output_cost_per_token"
)
assert float(response.headers["x-litellm-response-cost"]) == _approx(expected_cost)
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")]
rows: Final = eventually(
lambda: read_rows(
'SELECT spend, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
(identity,),
),
lambda values: len(values) == 1,
seconds=70,
)
assert (rows[0]["prompt_tokens"], rows[0]["completion_tokens"]) == (23, 41)
spend: Final = rows[0]["spend"]
assert isinstance(spend, (int, float, str))
assert float(spend) == _approx(expected_cost)
@pytest.mark.covers("other.provider_wire.xiaomi_mimo.reasoning_and_answer_stream_as_deltas")
def test_xiaomi_mimo_stream_delivers_reasoning_then_answer_deltas(gateway: Gateway) -> None:
backend: Final = _BACKENDS[0]
identity: Final = f"xiaomi-stream-{uuid.uuid4().hex}"
frames: Final = (
_frame(identity, backend, {"role": "assistant", "reasoning_content": "Count "}),
_frame(identity, backend, {"reasoning_content": "up by one."}),
_frame(identity, backend, {"content": "1\n2\n"}),
_frame(identity, backend, {"content": "3\n4\n5"}),
_frame(identity, backend, {}, finish="stop"),
b"data: [DONE]\n\n",
)
def respond(request: Request) -> Reply:
body: Final = _assert_provider_request(request, backend, _COUNTING_PROMPT)
assert body["stream"] is True
return Reply(content_type="text/event-stream", chunks=frames)
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"xiaomi_mimo/{backend}", api_base=wire.url, api_key=_API_KEY)
with gateway.client.stream(
"POST",
"/v1/chat/completions",
json={"model": model, "messages": [{"role": "user", "content": _COUNTING_PROMPT}], "stream": True},
headers={"Authorization": f"Bearer {gateway.key}"},
) as response:
assert response.status_code == 200, response.read()
lines: Final = tuple(line for line in response.iter_lines() if line.startswith("data: "))
assert lines[-1] == "data: [DONE]"
chunks: Final = tuple(_Chunk.model_validate_json(line.removeprefix("data: ")) for line in lines[:-1])
assert {chunk.id for chunk in chunks} == {identity}
choices: Final = tuple(choice for chunk in chunks for choice in chunk.choices)
assert "".join(choice.delta.reasoning_content or "" for choice in choices) == "Count up by one."
assert "".join(choice.delta.content or "" for choice in choices) == "1\n2\n3\n4\n5"
assert tuple(choice.finish_reason for choice in choices if choice.finish_reason) == ("stop",)
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")]
@pytest.mark.covers("other.provider_wire.xiaomi_mimo.tool_call_survives_translation")
def test_xiaomi_mimo_tool_call_is_forwarded_and_returned(gateway: Gateway) -> None:
backend: Final = _BACKENDS[1]
identity: Final = f"xiaomi-tool-{uuid.uuid4().hex}"
tool_call: Final = {
"id": "call_paris",
"type": "function",
"function": {"name": "get_weather", "arguments": json.dumps({"city": "Paris"})},
}
def respond(request: Request) -> Reply:
body: Final = _assert_provider_request(request, backend, _WEATHER_PROMPT)
assert body["tools"] == [_WEATHER_TOOL]
assert body["tool_choice"] == "auto"
return Reply(
body=_completion(
identity,
backend,
{
"role": "assistant",
"content": None,
"reasoning_content": "Need the tool.",
"tool_calls": [tool_call],
},
"tool_calls",
)
)
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"xiaomi_mimo/{backend}", api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": _WEATHER_PROMPT}],
"tools": [_WEATHER_TOOL],
"tool_choice": "auto",
},
)
assert response.status_code == 200, response.text
payload: Final = _JSON_OBJECT.validate_json(response.content)
assert payload["choices"] == [
{
"finish_reason": "tool_calls",
"index": 0,
"message": {
"role": "assistant",
"content": None,
"reasoning_content": "Need the tool.",
"tool_calls": [tool_call],
"provider_specific_fields": {"refusal": None},
},
"provider_specific_fields": {},
}
]
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")]

View file

@ -53,6 +53,7 @@ async def test_vertex_ai_anthropic_streaming_cost_injection_enabled():
# Setup logging object with model info
litellm_logging_obj = MagicMock(spec=LiteLLMLoggingObj)
litellm_logging_obj.litellm_params = {}
litellm_logging_obj.model_call_details = {"model": "claude-sonnet-4@20250514"}
litellm_logging_obj.completion_start_time = None
litellm_logging_obj.async_success_handler = AsyncMock()
@ -132,6 +133,7 @@ async def test_vertex_ai_anthropic_streaming_cost_injection_disabled():
response.aiter_bytes = mock_aiter_bytes
litellm_logging_obj = MagicMock(spec=LiteLLMLoggingObj)
litellm_logging_obj.litellm_params = {}
litellm_logging_obj.model_call_details = {"model": "claude-sonnet-4@20250514"}
litellm_logging_obj.completion_start_time = None
litellm_logging_obj.async_success_handler = AsyncMock()
@ -194,6 +196,7 @@ async def test_vertex_ai_anthropic_streaming_cost_injection_no_usage_chunk():
response.aiter_bytes = mock_aiter_bytes
litellm_logging_obj = MagicMock(spec=LiteLLMLoggingObj)
litellm_logging_obj.litellm_params = {}
litellm_logging_obj.model_call_details = {"model": "claude-sonnet-4@20250514"}
litellm_logging_obj.completion_start_time = None
litellm_logging_obj.async_success_handler = AsyncMock()
@ -249,6 +252,7 @@ async def test_vertex_ai_anthropic_streaming_model_extraction():
response.aiter_bytes = mock_aiter_bytes
litellm_logging_obj = MagicMock(spec=LiteLLMLoggingObj)
litellm_logging_obj.litellm_params = {}
litellm_logging_obj.model_call_details = {}
litellm_logging_obj.completion_start_time = None
litellm_logging_obj.async_success_handler = AsyncMock()

View file

@ -0,0 +1,179 @@
import datetime
from typing import Final
import pytest
from prometheus_client import REGISTRY
from prometheus_client.samples import Sample
import litellm
from litellm.integrations.prometheus import PrometheusLogger
from litellm.types.utils import StandardLoggingZeroCostDiagnostic
METRIC: Final = "litellm_zero_cost_requests_total"
MISSING_KEY_DIAGNOSTIC: Final[StandardLoggingZeroCostDiagnostic] = {
"reason": "missing_pricing_key",
"pricing_model": "dep-1",
"missing_pricing_keys": ("input_cost_per_token", "output_cost_per_token"),
}
def _clear_prometheus_registry() -> None:
for collector in list(REGISTRY._collector_to_names.keys()):
try:
REGISTRY.unregister(collector)
except Exception:
pass
def _samples(metric_name: str) -> list[Sample]:
return [sample for metric in REGISTRY.collect() for sample in metric.samples if sample.name == metric_name]
def _payload(zero_cost_diagnostic: StandardLoggingZeroCostDiagnostic | None) -> dict[str, object]:
return {
"id": "t",
"call_type": "completion",
"response_cost": 0.0,
"status": "success",
"total_tokens": 30,
"prompt_tokens": 20,
"completion_tokens": 10,
"startTime": 1.0,
"endTime": 2.0,
"completionStartTime": 1.5,
"model": "openai/gpt-5.4-nano",
"model_id": "dep-1",
"model_group": "per-second-priced-chat",
"api_base": "https://api.openai.com",
"custom_llm_provider": "openai",
"request_tags": [],
"end_user": None,
"cache_hit": False,
"stream": False,
"response": {"id": "chatcmpl-1"},
"model_parameters": {},
"zero_cost_diagnostic": zero_cost_diagnostic,
"metadata": {
"user_api_key_hash": "h",
"user_api_key_alias": "a",
"user_api_key_team_id": "t",
"user_api_key_team_alias": "ta",
"user_api_key_user_id": "u",
"user_api_key_user_email": "e@x.com",
"user_api_key_org_id": None,
"user_api_key_org_alias": None,
"requester_metadata": None,
"user_api_key_end_user_id": None,
"usage_object": None,
},
"hidden_params": {"litellm_overhead_time_ms": None, "additional_headers": None},
}
async def _log_success(
logger: PrometheusLogger, zero_cost_diagnostic: StandardLoggingZeroCostDiagnostic | None
) -> None:
now: Final = datetime.datetime.now()
kwargs: Final = {
"model": "openai/gpt-5.4-nano",
"litellm_params": {"metadata": {}},
"standard_logging_object": _payload(zero_cost_diagnostic),
"stream": False,
"start_time": now - datetime.timedelta(seconds=3),
"api_call_start_time": now - datetime.timedelta(seconds=2),
"completion_start_time": now - datetime.timedelta(seconds=1),
"end_time": now,
}
await logger.async_log_success_event(kwargs, None, now, now)
async def _log_failure(
logger: PrometheusLogger, zero_cost_diagnostic: StandardLoggingZeroCostDiagnostic | None
) -> None:
now: Final = datetime.datetime.now()
kwargs: Final = {
"model": "openai/gpt-5.4-nano",
"litellm_params": {"metadata": {}},
"standard_logging_object": {**_payload(zero_cost_diagnostic), "status": "failure"},
"exception": Exception("stream cut off after the usage chunk"),
"stream": True,
"start_time": now - datetime.timedelta(seconds=3),
"end_time": now,
}
await logger.async_log_failure_event(kwargs, None, now, now)
@pytest.mark.asyncio
async def test_failure_event_counts_a_zero_cost_request_by_model_and_reason() -> None:
_clear_prometheus_registry()
try:
logger: Final = PrometheusLogger()
await _log_failure(logger, None)
assert _samples(METRIC) == []
await _log_failure(logger, MISSING_KEY_DIAGNOSTIC)
samples: Final = _samples(METRIC)
assert len(samples) == 1
assert samples[0].labels == {
"requested_model": "per-second-priced-chat",
"model": "openai/gpt-5.4-nano",
"model_id": "dep-1",
"api_provider": "openai",
"reason": "missing_pricing_key",
}
assert samples[0].value == 1.0
finally:
_clear_prometheus_registry()
@pytest.mark.asyncio
async def test_success_event_counts_a_zero_cost_request_by_model_and_reason() -> None:
_clear_prometheus_registry()
try:
logger: Final = PrometheusLogger()
await _log_success(logger, MISSING_KEY_DIAGNOSTIC)
await _log_success(logger, MISSING_KEY_DIAGNOSTIC)
samples: Final = _samples(METRIC)
assert len(samples) == 1
assert samples[0].labels == {
"requested_model": "per-second-priced-chat",
"model": "openai/gpt-5.4-nano",
"model_id": "dep-1",
"api_provider": "openai",
"reason": "missing_pricing_key",
}
assert samples[0].value == 2.0
finally:
_clear_prometheus_registry()
@pytest.mark.asyncio
async def test_request_without_a_diagnostic_leaves_the_counter_untouched() -> None:
_clear_prometheus_registry()
try:
await _log_success(PrometheusLogger(), None)
assert _samples(METRIC) == []
finally:
_clear_prometheus_registry()
@pytest.mark.asyncio
async def test_label_filter_that_drops_reason_still_counts_the_request() -> None:
_clear_prometheus_registry()
previous_config: Final = litellm.prometheus_metrics_config
litellm.prometheus_metrics_config = [
{"group": "zero_cost", "metrics": [METRIC], "include_labels": ["requested_model"]}
]
try:
await _log_success(PrometheusLogger(), MISSING_KEY_DIAGNOSTIC)
samples: Final = _samples(METRIC)
assert len(samples) == 1
assert samples[0].labels == {"requested_model": "per-second-priced-chat"}
assert samples[0].value == 1.0
finally:
litellm.prometheus_metrics_config = previous_config
_clear_prometheus_registry()

View file

@ -0,0 +1,157 @@
from collections.abc import Mapping
from typing import Final
import pytest
from litellm.litellm_core_utils.llm_cost_calc.zero_cost_diagnostic import (
ZERO_COST_COUNTER_NAME,
diagnose_zero_cost,
used_pricing_keys,
zero_cost_warning,
)
from litellm.types.utils import CompletionTokensDetailsWrapper, PromptTokensDetailsWrapper, Usage
PER_SECOND_ENTRY: Final = {"input_cost_per_second": 0.00042, "output_cost_per_second": 0.00042}
FREE_ENTRY: Final = {"input_cost_per_token": 0, "output_cost_per_token": 0, "cache_read_input_token_cost": 2e-08}
PRICED_ENTRY: Final = {"input_cost_per_token": 1e-06, "output_cost_per_token": 2e-06}
TEXT_USAGE: Final = Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30)
def test_missing_pricing_key_names_every_rate_the_usage_needs() -> None:
diagnostic = diagnose_zero_cost(
usage=TEXT_USAGE, pricing_model="dep-1", pricing_entry=PER_SECOND_ENTRY, calculation_failed=False
)
assert diagnostic == {
"reason": "missing_pricing_key",
"pricing_model": "dep-1",
"missing_pricing_keys": ("input_cost_per_token", "output_cost_per_token"),
}
def test_only_the_absent_rate_is_reported() -> None:
diagnostic = diagnose_zero_cost(
usage=TEXT_USAGE, pricing_model="dep-1", pricing_entry={"input_cost_per_token": 1e-06}, calculation_failed=False
)
assert diagnostic is not None
assert diagnostic["missing_pricing_keys"] == ("output_cost_per_token",)
def test_free_model_stays_silent() -> None:
assert (
diagnose_zero_cost(usage=TEXT_USAGE, pricing_model="dep-1", pricing_entry=FREE_ENTRY, calculation_failed=False)
is None
)
@pytest.mark.parametrize("calculation_failed", [False, True])
def test_request_without_usage_stays_silent(calculation_failed: bool) -> None:
usage = Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0)
assert (
diagnose_zero_cost(
usage=usage, pricing_model="dep-1", pricing_entry=PER_SECOND_ENTRY, calculation_failed=calculation_failed
)
is None
)
@pytest.mark.parametrize(
"entry",
[
{"litellm_provider": "openai", "mode": "chat", "supports_prompt_caching": True},
{"tiered_pricing": [{"range": [0, 128000], "input_cost_per_token": 0, "output_cost_per_token": 0}]},
{"tiered_pricing": "not a tier table", "litellm_provider": "openai"},
],
)
def test_entry_that_declares_no_rate_stays_silent(entry: Mapping[str, object]) -> None:
assert (
diagnose_zero_cost(usage=TEXT_USAGE, pricing_model="dep-1", pricing_entry=entry, calculation_failed=False)
is None
)
def test_tiered_rate_counts_as_a_declared_rate() -> None:
entry = {"tiered_pricing": [{"range": [0, 128000], "input_cost_per_token": 1e-06, "output_cost_per_token": 2e-06}]}
diagnostic = diagnose_zero_cost(
usage=TEXT_USAGE, pricing_model="dep-1", pricing_entry=entry, calculation_failed=False
)
assert diagnostic is not None
assert diagnostic["reason"] == "missing_pricing_key"
def test_priced_entry_that_still_prices_to_zero_is_pricing_not_applied() -> None:
diagnostic = diagnose_zero_cost(
usage=TEXT_USAGE, pricing_model="dep-1", pricing_entry=PRICED_ENTRY, calculation_failed=False
)
assert diagnostic == {"reason": "pricing_not_applied", "pricing_model": "dep-1", "missing_pricing_keys": ()}
def test_calculator_failure_on_a_priced_entry_is_cost_calculation_error() -> None:
diagnostic = diagnose_zero_cost(
usage=TEXT_USAGE, pricing_model="dep-1", pricing_entry=PRICED_ENTRY, calculation_failed=True
)
assert diagnostic == {"reason": "cost_calculation_error", "pricing_model": "dep-1", "missing_pricing_keys": ()}
def test_calculator_failure_on_a_free_entry_stays_silent() -> None:
assert (
diagnose_zero_cost(usage=TEXT_USAGE, pricing_model="dep-1", pricing_entry=FREE_ENTRY, calculation_failed=True)
is None
)
def test_calculator_failure_on_an_entry_that_declares_no_rate_stays_silent() -> None:
entry: Final = {"litellm_provider": "openai", "mode": "chat", "supports_prompt_caching": True}
assert (
diagnose_zero_cost(usage=TEXT_USAGE, pricing_model="dep-1", pricing_entry=entry, calculation_failed=True)
is None
)
def test_audio_tokens_need_the_audio_rates() -> None:
usage = Usage(
prompt_tokens=10,
completion_tokens=20,
total_tokens=30,
prompt_tokens_details=PromptTokensDetailsWrapper(audio_tokens=10, text_tokens=0),
completion_tokens_details=CompletionTokensDetailsWrapper(audio_tokens=5, text_tokens=15),
)
assert used_pricing_keys(usage) == (
"input_cost_per_audio_token",
"output_cost_per_token",
"output_cost_per_audio_token",
)
diagnostic = diagnose_zero_cost(
usage=usage, pricing_model="gemini-audio", pricing_entry=PRICED_ENTRY, calculation_failed=False
)
assert diagnostic is not None
assert diagnostic["missing_pricing_keys"] == ("input_cost_per_audio_token", "output_cost_per_audio_token")
def test_warning_names_the_request_the_entry_the_missing_keys_and_the_counter() -> None:
diagnostic = diagnose_zero_cost(
usage=TEXT_USAGE, pricing_model="dep-1", pricing_entry=PER_SECOND_ENTRY, calculation_failed=False
)
assert diagnostic is not None
message = zero_cost_warning(
diagnostic,
model_group="per-second-priced-chat",
model="openai/gpt-5.4-nano",
custom_llm_provider="openai",
usage=TEXT_USAGE,
)
assert "model_group=per-second-priced-chat" in message
assert "model=openai/gpt-5.4-nano" in message
assert "provider=openai" in message
assert "prompt_tokens=10 completion_tokens=20" in message
assert "pricing entry 'dep-1' has no input_cost_per_token, output_cost_per_token" in message
assert f'{ZERO_COST_COUNTER_NAME}{{reason="missing_pricing_key"}}' in message

View file

@ -6,6 +6,8 @@ import pytest
from litellm.litellm_core_utils.core_helpers import (
_FINISH_REASON_MAP,
bind_budget_reservation_to_callbacks,
budget_reservation_from_metadata,
drop_params_env_flag,
drop_params_flag,
get_or_create_metadata_bucket,
@ -13,7 +15,60 @@ from litellm.litellm_core_utils.core_helpers import (
normalize_drop_params,
reconstruct_model_name,
redact_nested_match_and_regex_keys,
unbind_budget_reservation_from_callbacks,
)
from litellm.proxy._types import UserAPIKeyAuth
class TestBudgetReservationBinding:
"""The request-end release skips a reservation a cost callback has claimed, so the claim
must land on the one dict auth stamped, through whichever metadata field or auth object
carries it, and a failed call must be able to hand it back."""
@staticmethod
def _reservation() -> dict:
return {"reserved_cost": 0.5, "entries": [], "finalized": False, "callback_bound": False}
@pytest.mark.parametrize("metadata_variable_name", ["metadata", "litellm_metadata"])
def test_reservation_stamped_on_the_metadata_is_bound(self, metadata_variable_name: str):
reservation = self._reservation()
bind_budget_reservation_to_callbacks({metadata_variable_name: {"user_api_key_budget_reservation": reservation}})
assert reservation["callback_bound"] is True
def test_reservation_reachable_only_through_the_auth_object_is_bound(self):
reservation = self._reservation()
user_api_key_auth = UserAPIKeyAuth(token="hashed")
user_api_key_auth.budget_reservation = reservation
bind_budget_reservation_to_callbacks({"metadata": {"user_api_key_auth": user_api_key_auth}})
assert reservation["callback_bound"] is True
def test_reservation_reachable_only_through_a_dumped_auth_object_is_bound(self):
reservation = self._reservation()
bind_budget_reservation_to_callbacks({"metadata": {"user_api_key_auth": {"budget_reservation": reservation}}})
assert reservation["callback_bound"] is True
def test_unbind_hands_a_claimed_reservation_back(self):
reservation = self._reservation()
litellm_params = {"litellm_metadata": {"user_api_key_budget_reservation": reservation}}
bind_budget_reservation_to_callbacks(litellm_params)
unbind_budget_reservation_from_callbacks(litellm_params)
assert reservation["callback_bound"] is False
def test_request_without_a_reservation_binds_nothing(self):
metadata = {"user_api_key_auth": UserAPIKeyAuth(token="hashed")}
bind_budget_reservation_to_callbacks({"metadata": metadata, "litellm_metadata": None})
assert budget_reservation_from_metadata(metadata) is None
assert "user_api_key_budget_reservation" not in metadata
class TestGetOrCreateMetadataBucket:

View file

@ -1,9 +1,10 @@
import asyncio
import contextlib
import datetime
import logging
import os
import sys
from collections.abc import Callable
from collections.abc import Callable, Iterator, Mapping
from typing import Final, Literal
from unittest.mock import AsyncMock, MagicMock, patch
@ -25,6 +26,7 @@ from litellm.litellm_core_utils.litellm_logging import (
set_callbacks,
)
from litellm.llms.base_llm.ocr.transformation import OCRUsageInfo
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.llms.openai import ResponseAPIUsage, ResponseCompletedEvent, ResponsesAPIResponse
from litellm.types.utils import (
CallTypes,
@ -302,6 +304,398 @@ def test_response_cost_calculator_uses_router_model_id_from_litellm_metadata():
litellm.model_cost.pop(custom_model_id, None)
class TestZeroCostDiagnostic:
DEPLOYMENT_ID: Final = "lit7898-per-second-priced-deployment"
MODEL_GROUP: Final = "per-second-priced-chat"
PER_SECOND_PRICING: Final = {"input_cost_per_second": 0.00042, "output_cost_per_second": 0.00042}
FREE_PRICING: Final = {"input_cost_per_token": 0, "output_cost_per_token": 0}
@pytest.fixture(params=["per_second", "free"])
def deployment_pricing(self, request: pytest.FixtureRequest) -> Iterator[Mapping[str, float]]:
pricing: Final = self.PER_SECOND_PRICING if request.param == "per_second" else self.FREE_PRICING
litellm.register_model(model_cost={self.DEPLOYMENT_ID: pricing}, persist_across_reloads=False)
try:
yield pricing
finally:
litellm.model_cost.pop(self.DEPLOYMENT_ID, None)
def _logging_obj(
self,
pricing: Mapping[str, object],
stream: bool = False,
model: str = "openai/gpt-5.4-nano",
call_type: str = "completion",
deployment_id: str | None = DEPLOYMENT_ID,
custom_llm_provider: str = "openai",
) -> LitellmLogging:
logging_obj: Final = LitellmLogging(
model=model,
messages=[{"role": "user", "content": "Hi"}],
stream=stream,
call_type=call_type,
start_time=time.time(),
litellm_call_id="lit7898",
function_id="fn",
)
self._route_to_deployment(
logging_obj, pricing, model=model, deployment_id=deployment_id, custom_llm_provider=custom_llm_provider
)
return logging_obj
def _route_to_deployment(
self,
logging_obj: LitellmLogging,
pricing: Mapping[str, object],
model: str = "openai/gpt-5.4-nano",
deployment_id: str | None = DEPLOYMENT_ID,
custom_llm_provider: str = "openai",
) -> None:
model_info: Final = pricing if deployment_id is None else {"id": deployment_id, **pricing}
logging_obj.update_environment_variables(
model=model,
user="",
optional_params={},
litellm_params={"metadata": {"model_group": self.MODEL_GROUP, "model_info": model_info}},
custom_llm_provider=custom_llm_provider,
)
@staticmethod
def _response(
usage: litellm.Usage | None = None, model: str = "gpt-5.4-nano", **hidden_params: object
) -> ModelResponse:
response: Final = ModelResponse(
model=model,
choices=[litellm.Choices(message=litellm.Message(role="assistant", content="hello"))],
usage=usage,
)
response._hidden_params = {"custom_llm_provider": "openai", **hidden_params}
return response
@staticmethod
def _zero_cost_warnings(caplog: pytest.LogCaptureFixture) -> list[str]:
return [
record.getMessage()
for record in caplog.records
if record.name == "LiteLLM" and record.levelno == logging.WARNING and "priced at $0" in record.getMessage()
]
def _assert_flagged(self, logging_obj: LitellmLogging, caplog: pytest.LogCaptureFixture) -> None:
assert logging_obj.model_call_details["zero_cost_diagnostic"] == {
"reason": "missing_pricing_key",
"pricing_model": self.DEPLOYMENT_ID,
"missing_pricing_keys": ("input_cost_per_token", "output_cost_per_token"),
}
warnings: Final = self._zero_cost_warnings(caplog)
assert len(warnings) == 1
assert f"model_group={self.MODEL_GROUP}" in warnings[0]
assert f"pricing entry '{self.DEPLOYMENT_ID}' has no input_cost_per_token, output_cost_per_token" in warnings[0]
def test_zero_cost_with_a_missing_rate_warns_once_and_is_recorded(
self, deployment_pricing: Mapping[str, float], caplog: pytest.LogCaptureFixture
) -> None:
usage: Final = litellm.Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30)
logging_obj: Final = self._logging_obj(deployment_pricing)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
first_cost: Final = logging_obj._response_cost_calculator(result=self._response(usage))
second_cost: Final = logging_obj._response_cost_calculator(result=self._response(usage))
assert first_cost == 0.0
assert second_cost == 0.0
if deployment_pricing is self.FREE_PRICING:
assert logging_obj.model_call_details["zero_cost_diagnostic"] is None
assert self._zero_cost_warnings(caplog) == []
return
self._assert_flagged(logging_obj, caplog)
def test_usage_less_stream_chunk_does_not_hide_the_final_response_diagnostic(
self, deployment_pricing: Mapping[str, float], caplog: pytest.LogCaptureFixture
) -> None:
usage: Final = litellm.Usage(prompt_tokens=8, completion_tokens=2, total_tokens=10)
logging_obj: Final = self._logging_obj(deployment_pricing, stream=True)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
logging_obj._response_cost_calculator(result=self._response(usage=None))
logging_obj._response_cost_calculator(result=self._response(usage))
if deployment_pricing is self.FREE_PRICING:
assert logging_obj.model_call_details["zero_cost_diagnostic"] is None
assert self._zero_cost_warnings(caplog) == []
return
self._assert_flagged(logging_obj, caplog)
def test_terminal_responses_stream_event_is_judged_by_its_inner_response(
self, deployment_pricing: Mapping[str, float], caplog: pytest.LogCaptureFixture
) -> None:
logging_obj: Final = self._logging_obj(deployment_pricing, stream=True, call_type="aresponses")
event: Final = ResponseCompletedEvent(
type="response.completed",
response=ResponsesAPIResponse(
id="resp-lit7898",
created_at=1,
object="response",
status="completed",
model="gpt-5.4-nano",
output=[],
usage=ResponseAPIUsage(input_tokens=10, output_tokens=20, total_tokens=30),
),
)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
cost: Final = logging_obj._response_cost_calculator(result=event)
assert cost == 0.0
if deployment_pricing is self.FREE_PRICING:
assert logging_obj.model_call_details["zero_cost_diagnostic"] is None
assert self._zero_cost_warnings(caplog) == []
return
self._assert_flagged(logging_obj, caplog)
def test_precomputed_zero_hidden_cost_is_flagged_and_lands_in_the_payload(
self, deployment_pricing: Mapping[str, float], caplog: pytest.LogCaptureFixture
) -> None:
usage: Final = litellm.Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30)
logging_obj: Final = self._logging_obj(deployment_pricing)
response: Final = self._response(usage, response_cost=0.0, model_id=self.DEPLOYMENT_ID)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
logging_obj._process_hidden_params_and_response_cost(
response, start_time=datetime.datetime.now(), end_time=datetime.datetime.now()
)
payload: Final = logging_obj.model_call_details["standard_logging_object"]
assert payload["response_cost"] == 0.0
if deployment_pricing is self.FREE_PRICING:
assert payload["zero_cost_diagnostic"] is None
assert self._zero_cost_warnings(caplog) == []
return
self._assert_flagged(logging_obj, caplog)
assert payload["zero_cost_diagnostic"] == logging_obj.model_call_details["zero_cost_diagnostic"]
def test_uncomputed_hidden_cost_is_not_a_zero_cost(
self, deployment_pricing: Mapping[str, float], caplog: pytest.LogCaptureFixture
) -> None:
usage: Final = litellm.Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30)
logging_obj: Final = self._logging_obj(deployment_pricing)
response: Final = self._response(usage, response_cost=None, model_id=self.DEPLOYMENT_ID)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
logging_obj._process_hidden_params_and_response_cost(
response, start_time=datetime.datetime.now(), end_time=datetime.datetime.now()
)
assert logging_obj.model_call_details["standard_logging_object"]["zero_cost_diagnostic"] is None
assert self._zero_cost_warnings(caplog) == []
def test_unbilled_read_route_with_usage_stays_silent(
self, deployment_pricing: Mapping[str, float], caplog: pytest.LogCaptureFixture
) -> None:
usage: Final = litellm.Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30)
logging_obj: Final = self._logging_obj(deployment_pricing, call_type="aget_responses")
response: Final = self._response(usage, response_cost=0.0, model_id=self.DEPLOYMENT_ID)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
logging_obj._process_hidden_params_and_response_cost(
response, start_time=datetime.datetime.now(), end_time=datetime.datetime.now()
)
assert logging_obj.model_call_details["standard_logging_object"]["zero_cost_diagnostic"] is None
assert self._zero_cost_warnings(caplog) == []
def test_unmapped_model_that_fails_cost_calculation_stays_silent(self, caplog: pytest.LogCaptureFixture) -> None:
usage: Final = litellm.Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30)
logging_obj: Final = self._logging_obj(
{}, model="openai/lit7898-unmapped-model", deployment_id="lit7898-unmapped-deployment"
)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
cost: Final = logging_obj._response_cost_calculator(
result=self._response(usage, model="lit7898-unmapped-model")
)
assert cost is None
assert logging_obj.model_call_details["response_cost_failure_debug_information"] is not None
assert logging_obj.model_call_details.get("zero_cost_diagnostic") is None
assert self._zero_cost_warnings(caplog) == []
def test_malformed_usage_never_raises_out_of_the_cost_calculator(
self, deployment_pricing: Mapping[str, float], caplog: pytest.LogCaptureFixture
) -> None:
logging_obj: Final = self._logging_obj(deployment_pricing)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
cost: Final = logging_obj._response_cost_calculator(
result={"model": "gpt-5.4-nano", "usage": {"prompt_tokens": "n/a", "completion_tokens": 3}}
)
assert cost is None
assert logging_obj.model_call_details.get("zero_cost_diagnostic") is None
assert self._zero_cost_warnings(caplog) == []
def test_usage_less_evaluation_between_two_zero_cost_findings_does_not_warn_twice(
self, deployment_pricing: Mapping[str, float], caplog: pytest.LogCaptureFixture
) -> None:
usage: Final = litellm.Usage(prompt_tokens=8, completion_tokens=2, total_tokens=10)
logging_obj: Final = self._logging_obj(deployment_pricing, stream=True, call_type="anthropic_messages")
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
logging_obj._response_cost_calculator(result=self._response(usage=None))
logging_obj._response_cost_calculator(result=self._response(usage))
logging_obj._response_cost_calculator(result=self._response(usage=None))
logging_obj._response_cost_calculator(result=self._response(usage))
if deployment_pricing is self.FREE_PRICING:
assert logging_obj.model_call_details["zero_cost_diagnostic"] is None
assert self._zero_cost_warnings(caplog) == []
return
self._assert_flagged(logging_obj, caplog)
def test_retry_that_prices_clears_the_diagnostic_and_a_later_zero_cost_is_recorded_silently(
self, caplog: pytest.LogCaptureFixture
) -> None:
priced_id: Final = "lit7898-priced-deployment"
priced_pricing: Final = {"input_cost_per_token": 1e-06, "output_cost_per_token": 2e-06}
usage: Final = litellm.Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30)
litellm.register_model(
model_cost={self.DEPLOYMENT_ID: self.PER_SECOND_PRICING, priced_id: priced_pricing},
persist_across_reloads=False,
)
try:
logging_obj: Final = self._logging_obj(self.PER_SECOND_PRICING)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
assert logging_obj._response_cost_calculator(result=self._response(usage)) == 0.0
self._assert_flagged(logging_obj, caplog)
self._route_to_deployment(logging_obj, priced_pricing, deployment_id=priced_id)
assert logging_obj._response_cost_calculator(result=self._response(usage)) == pytest.approx(5e-05)
assert logging_obj.model_call_details["zero_cost_diagnostic"] is None
self._route_to_deployment(logging_obj, self.PER_SECOND_PRICING)
assert logging_obj._response_cost_calculator(result=self._response(usage)) == 0.0
assert logging_obj.model_call_details["zero_cost_diagnostic"]["reason"] == "missing_pricing_key"
assert len(self._zero_cost_warnings(caplog)) == 1
finally:
litellm.model_cost.pop(self.DEPLOYMENT_ID, None)
litellm.model_cost.pop(priced_id, None)
def test_one_request_evaluated_against_two_cost_map_entries_warns_once(
self, caplog: pytest.LogCaptureFixture
) -> None:
dated_model: Final = "lit7898-nano-2026-03-17"
requested_model: Final = "lit7898-nano"
usage: Final = litellm.Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30)
cost_map_entry: Final = {"litellm_provider": "openai", "mode": "chat", **self.PER_SECOND_PRICING}
litellm.register_model(
model_cost={dated_model: cost_map_entry, requested_model: cost_map_entry}, persist_across_reloads=False
)
try:
logging_obj: Final = self._logging_obj(
{}, model=f"openai/{requested_model}", deployment_id="lit7898-cost-map-deployment"
)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
assert logging_obj._response_cost_calculator(result=self._response(usage, model=dated_model)) == 0.0
assert logging_obj._response_cost_calculator(result=self._response(usage, model=requested_model)) == 0.0
assert logging_obj.model_call_details["zero_cost_diagnostic"]["pricing_model"] == requested_model
warnings: Final = self._zero_cost_warnings(caplog)
assert len(warnings) == 1
assert f"pricing entry '{dated_model}' has no input_cost_per_token, output_cost_per_token" in warnings[0]
finally:
litellm.model_cost.pop(dated_model, None)
litellm.model_cost.pop(requested_model, None)
def test_free_deployment_without_a_router_id_is_judged_by_its_own_pricing(
self, caplog: pytest.LogCaptureFixture
) -> None:
global_model: Final = "lit7898-priced-global"
usage: Final = litellm.Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30)
litellm.register_model(
model_cost={
global_model: {
"litellm_provider": "openai",
"mode": "chat",
"input_cost_per_token": 1e-06,
"output_cost_per_token": 2e-06,
}
},
persist_across_reloads=False,
)
try:
logging_obj: Final = self._logging_obj(self.FREE_PRICING, model=global_model, deployment_id=None)
response: Final = self._response(usage, model=global_model, response_cost=0.0)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
logging_obj._process_hidden_params_and_response_cost(
response, start_time=datetime.datetime.now(), end_time=datetime.datetime.now()
)
assert logging_obj.model_call_details["zero_cost_diagnostic"] is None
assert self._zero_cost_warnings(caplog) == []
finally:
litellm.model_cost.pop(global_model, None)
def test_cache_hit_priced_for_saved_cost_stays_silent(
self, deployment_pricing: Mapping[str, float], caplog: pytest.LogCaptureFixture
) -> None:
usage: Final = litellm.Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30)
logging_obj: Final = self._logging_obj(deployment_pricing)
logging_obj.model_call_details["cache_hit"] = True
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
assert logging_obj._response_cost_calculator(result=self._response(usage), cache_hit=False) == 0.0
assert logging_obj.model_call_details["zero_cost_diagnostic"] is None
assert self._zero_cost_warnings(caplog) == []
@pytest.mark.parametrize("spilled_over", [True, False])
def test_ptu_deployment_is_judged_by_the_entry_the_calculator_priced_with(
self, spilled_over: bool, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
) -> None:
router_model_id: Final = "lit7898-ptu-router-model-id"
served_model: Final = "azure/lit7898-ptu-served-model"
ptu_model_info: Final = {
"team_id": "team-1",
"ptu_count": 100,
"cost_per_ptu_per_hour": 1.0,
"ptu_effective_from": "2026-01-01",
**self.FREE_PRICING,
}
usage: Final = litellm.Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30)
litellm.register_model(
model_cost={
router_model_id: {**self.FREE_PRICING, "litellm_provider": "azure", "mode": "chat"},
served_model: {**self.PER_SECOND_PRICING, "litellm_provider": "azure", "mode": "chat"},
},
persist_across_reloads=False,
)
monkeypatch.setenv("LITELLM_ENABLE_PTU_COST_ATTRIBUTION", "True")
try:
logging_obj: Final = self._logging_obj(
ptu_model_info, model=served_model, deployment_id=router_model_id, custom_llm_provider="azure"
)
spillover_headers: Final = {"llm_provider-x-ms-is-spilled-over": "true"} if spilled_over else {}
response: Final = self._response(
usage, model=served_model, custom_llm_provider="azure", additional_headers=spillover_headers
)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
assert logging_obj._response_cost_calculator(result=response) == 0.0
warnings: Final = self._zero_cost_warnings(caplog)
if not spilled_over:
assert logging_obj.model_call_details["zero_cost_diagnostic"] is None
assert warnings == []
return
assert logging_obj.model_call_details["zero_cost_diagnostic"] == {
"reason": "missing_pricing_key",
"pricing_model": served_model,
"missing_pricing_keys": ("input_cost_per_token", "output_cost_per_token"),
}
assert len(warnings) == 1
assert f"pricing entry '{served_model}' has no input_cost_per_token, output_cost_per_token" in warnings[0]
finally:
litellm.model_cost.pop(router_model_id, None)
litellm.model_cost.pop(served_model, None)
class TestGetRouterModelId:
"""Tests for the get_router_model_id helper method."""
@ -407,7 +801,6 @@ class TestGetRouterDeploymentModelInfo:
logging_obj.litellm_params = {"api_base": ""}
assert logging_obj.get_router_deployment_model_info() is None
def test_a_published_batch_rate_never_displaces_a_declared_standard_rate(self) -> None:
"""Ownership is per token direction, not per field.
@ -1111,7 +1504,9 @@ async def test_arealtime_marks_litellm_params_async(monkeypatch):
@pytest.mark.asyncio
async def test_aresponses_websocket_hands_back_the_provider_failure_without_a_success_log(monkeypatch: pytest.MonkeyPatch):
async def test_aresponses_websocket_hands_back_the_provider_failure_without_a_success_log(
monkeypatch: pytest.MonkeyPatch,
):
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.responses.main import base_llm_http_handler
@ -7066,22 +7461,41 @@ async def test_classifier_audit_matches_provider_transport(provider: str) -> Non
return httpx.Response(200, json=mock_responses_api_response(content).model_dump())
if provider == "anthropic":
return httpx.Response(200, json={
"id": "msg-audit", "type": "message", "role": "assistant", "model": "claude-haiku-4-5",
"content": [{"type": "text", "text": content}], "stop_reason": "end_turn",
"usage": {"input_tokens": 10, "output_tokens": 5},
})
return httpx.Response(
200,
json={
"id": "msg-audit",
"type": "message",
"role": "assistant",
"model": "claude-haiku-4-5",
"content": [{"type": "text", "text": content}],
"stop_reason": "end_turn",
"usage": {"input_tokens": 10, "output_tokens": 5},
},
)
if provider == "bedrock":
return httpx.Response(200, json={
"output": {"message": {"role": "assistant", "content": [{"text": content}]}},
"stopReason": "end_turn", "usage": {"inputTokens": 10, "outputTokens": 5, "totalTokens": 15},
"metrics": {"latencyMs": 1},
})
return httpx.Response(200, json={
"id": "chatcmpl-audit", "object": "chat.completion", "created": 0, "model": "gpt-5.6",
"choices": [{"index": 0, "message": {"role": "assistant", "content": content}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
})
return httpx.Response(
200,
json={
"output": {"message": {"role": "assistant", "content": [{"text": content}]}},
"stopReason": "end_turn",
"usage": {"inputTokens": 10, "outputTokens": 5, "totalTokens": 15},
"metrics": {"latencyMs": 1},
},
)
return httpx.Response(
200,
json={
"id": "chatcmpl-audit",
"object": "chat.completion",
"created": 0,
"model": "gpt-5.6",
"choices": [
{"index": 0, "message": {"role": "assistant", "content": content}, "finish_reason": "stop"}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
},
)
async def capture(kwargs, response_obj, start_time, end_time):
logs.put_nowait(kwargs["standard_logging_object"])
@ -7092,11 +7506,15 @@ async def test_classifier_audit_matches_provider_transport(provider: str) -> Non
handler.client = http_client
client: Final = (
AsyncAzureOpenAI(
api_key="transport-only", azure_endpoint="https://azure.invalid",
api_version="2025-04-01-preview", http_client=http_client,
api_key="transport-only",
azure_endpoint="https://azure.invalid",
api_version="2025-04-01-preview",
http_client=http_client,
)
if provider == "azure" else AsyncOpenAI(api_key="transport-only", http_client=http_client)
if provider == "openai" else handler
if provider == "azure"
else AsyncOpenAI(api_key="transport-only", http_client=http_client)
if provider == "openai"
else handler
)
model: Final = {
"openai": "openai/gpt-5.6",
@ -7109,23 +7527,44 @@ async def test_classifier_audit_matches_provider_transport(provider: str) -> Non
async def run(marker: str) -> None:
if provider == "responses":
await litellm.aresponses(
model=model, api_key="transport-only", client=client, max_output_tokens=128,
instructions="classifier-rubric", input=marker,
model=model,
api_key="transport-only",
client=client,
max_output_tokens=128,
instructions="classifier-rubric",
input=marker,
metadata={"internal_call_origin": "autorouter_classifier"},
proxy_server_request={"body": {}, "originating_request_masked": {"input": f"source-only-{marker}"}},
success_callback=[capture], num_retries=0,
success_callback=[capture],
num_retries=0,
)
return
await litellm.acompletion(
model=model, api_key="transport-only", client=client, max_tokens=128,
aws_access_key_id="transport-only", aws_secret_access_key="transport-only", aws_region_name="us-east-1",
model=model,
api_key="transport-only",
client=client,
max_tokens=128,
aws_access_key_id="transport-only",
aws_secret_access_key="transport-only",
aws_region_name="us-east-1",
messages=[{"role": "system", "content": "classifier-rubric"}, {"role": "user", "content": marker}],
metadata={"internal_call_origin": "autorouter_classifier"},
proxy_server_request={"body": {}, "originating_request_masked": {"input": f"source-only-{marker}"}},
success_callback=[capture], num_retries=0,
**({"api_base": "https://azure.invalid", "api_version": "2025-04-01-preview"} if provider == "azure" else {}),
**({"extra_body": {"audit_context": "provider-extra"}, "extra_headers": {"X-Audit": "header-only-secret"}}
if provider in ("openai", "azure") else {}),
success_callback=[capture],
num_retries=0,
**(
{"api_base": "https://azure.invalid", "api_version": "2025-04-01-preview"}
if provider == "azure"
else {}
),
**(
{
"extra_body": {"audit_context": "provider-extra"},
"extra_headers": {"X-Audit": "header-only-secret"},
}
if provider in ("openai", "azure")
else {}
),
)
await asyncio.gather(run("request-one"), run("request-two"))
@ -7149,14 +7588,17 @@ async def test_classifier_audit_matches_provider_transport(provider: str) -> Non
@pytest.mark.parametrize("redaction", ["none", "global", "request", "header"])
@pytest.mark.parametrize("status", ["success", "failure"])
@pytest.mark.parametrize("call_type", ["completion", "acompletion", "responses", "aresponses"])
def test_classifier_audit_obeys_message_logging_before_payload_emission(logging_obj, monkeypatch, redaction, status, call_type):
def test_classifier_audit_obeys_message_logging_before_payload_emission(
logging_obj, monkeypatch, redaction, status, call_type
):
from litellm.litellm_core_utils.litellm_logging import get_standard_logging_object_payload
monkeypatch.setattr(litellm, "turn_off_message_logging", redaction == "global")
params: Final = {
"metadata": {"internal_call_origin": "autorouter_classifier", **(
{"headers": {"x-litellm-enable-message-redaction": "true"}} if redaction == "header" else {}
)},
"metadata": {
"internal_call_origin": "autorouter_classifier",
**({"headers": {"x-litellm-enable-message-redaction": "true"}} if redaction == "header" else {}),
},
"proxy_server_request": {"body": {}, "originating_request_masked": {"input": "source-only"}},
}
logging_obj.call_type = call_type
@ -7169,8 +7611,12 @@ def test_classifier_audit_obeys_message_logging_before_payload_emission(logging_
)
now: Final = datetime.datetime.now()
payload: Final = get_standard_logging_object_payload(
kwargs={**logging_obj.model_call_details, "call_type": call_type}, init_response_obj={},
start_time=now, end_time=now, logging_obj=logging_obj, status=status,
kwargs={**logging_obj.model_call_details, "call_type": call_type},
init_response_obj={},
start_time=now,
end_time=now,
logging_obj=logging_obj,
status=status,
)
assert payload is not None
if redaction == "none":
@ -7421,7 +7867,13 @@ def _completed_responses_event(usage: ResponseAPIUsage) -> ResponseCompletedEven
return ResponseCompletedEvent(
type="response.completed",
response=ResponsesAPIResponse(
id="resp-1", created_at=1, object="response", status="completed", model="codex-mini-latest", output=[], usage=usage
id="resp-1",
created_at=1,
object="response",
status="completed",
model="codex-mini-latest",
output=[],
usage=usage,
),
)
@ -7441,7 +7893,9 @@ def test_get_assembled_streaming_response_bills_a_provider_reported_usage_cost()
now = datetime.datetime.now()
assembled = logging_obj._get_assembled_streaming_response(
result=_completed_responses_event(ResponseAPIUsage(input_tokens=12, output_tokens=2, total_tokens=14, cost=0.0042)),
result=_completed_responses_event(
ResponseAPIUsage(input_tokens=12, output_tokens=2, total_tokens=14, cost=0.0042)
),
start_time=now,
end_time=now,
is_async=True,
@ -7488,3 +7942,19 @@ def test_response_cost_calculator_prices_terminal_responses_event_from_its_respo
assert event_cost is not None and event_cost > 0
assert event_cost == inner_cost
assert logging_obj.cost_breakdown["input_cost"] is not None and logging_obj.cost_breakdown["input_cost"] > 0
class TestBudgetReservationBinding:
"""The proxy builds a logging object for every route before calling anything, so a
logging object seeing the reservation is no promise that a cost callback will settle
it: the claim belongs to the call wrapper, and this object must leave it unbound."""
def test_update_environment_variables_leaves_the_reservation_unbound(self, logging_obj):
reservation: Final = {"reserved_cost": 0.5, "entries": [], "finalized": False, "callback_bound": False}
logging_obj.update_environment_variables(
litellm_params={"metadata": {"user_api_key_budget_reservation": reservation}}, optional_params={}
)
assert logging_obj.litellm_params["metadata"]["user_api_key_budget_reservation"] is reservation
assert reservation["callback_bound"] is False

View file

@ -12,6 +12,9 @@ from litellm.types.router import GenericLiteLLMParams
class FakeLogging:
def __init__(self) -> None:
self.litellm_params: dict = {}
def update_from_kwargs(self, **kwargs):
pass

View file

@ -59,6 +59,7 @@ from litellm.proxy.auth.user_api_key_auth import (
_user_api_key_auth_builder,
get_api_key,
user_api_key_auth,
user_api_key_auth_websocket_for_model,
)
from litellm.proxy.spend_tracking.carried_budget_state import carried_budget_metadata
@ -9043,3 +9044,90 @@ async def test_router_settings_model_group_alias_authorizes_target_for_team(monk
await authorize()
assert (await request.json())["model"] == target
assert get_client_requested_model(request) == "AgentX-LLM"
@pytest.mark.asyncio
async def test_reserve_budget_after_common_checks_hands_the_reservation_to_the_request_state():
from fastapi import Request
request = Request(scope={"type": "http"})
user_api_key_auth_obj = UserAPIKeyAuth(token="test_token")
reservation = {"reserved_cost": 0.5, "entries": [], "finalized": False, "callback_bound": False}
with patch(
"litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request",
new=AsyncMock(return_value=reservation),
):
await _reserve_budget_after_common_checks(
user_api_key_auth_obj=user_api_key_auth_obj,
request_data={"model": "gpt-4o"},
route="/v1/batches/batch_123/cancel",
llm_router=None,
team_object=None,
user_object=None,
prisma_client=None,
user_api_key_cache=MagicMock(),
proxy_logging_obj=MagicMock(),
skip_budget_checks=False,
general_settings={},
request=request,
)
assert user_api_key_auth_obj.budget_reservation is reservation
assert request.state.budget_reservation is reservation
assert request.scope["state"]["budget_reservation"] is reservation
@pytest.mark.asyncio
async def test_reserve_budget_after_common_checks_clears_the_request_state_when_budget_checks_skip():
from fastapi import Request
request = Request(scope={"type": "http", "state": {"budget_reservation": {"reserved_cost": 0.5}}})
await _reserve_budget_after_common_checks(
user_api_key_auth_obj=UserAPIKeyAuth(token="test_token"),
request_data={"model": "free-model"},
route="/v1/chat/completions",
llm_router=None,
team_object=None,
user_object=None,
prisma_client=None,
user_api_key_cache=MagicMock(),
proxy_logging_obj=MagicMock(),
skip_budget_checks=True,
general_settings={},
request=request,
)
assert request.state.budget_reservation is None
@pytest.mark.asyncio
async def test_websocket_auth_hands_the_reservation_to_the_socket_state():
from fastapi import WebSocket
reservation = {"reserved_cost": 0.5, "entries": [], "finalized": False, "callback_bound": False}
websocket = WebSocket(
scope={
"type": "websocket",
"path": "/v1/realtime",
"headers": [(b"authorization", b"Bearer sk-1234")],
"query_string": b"model=gpt-realtime",
},
receive=AsyncMock(),
send=AsyncMock(),
)
async def auth_that_reserves(request, api_key):
request.state.budget_reservation = reservation
return UserAPIKeyAuth(token="hashed", budget_reservation=reservation)
with patch(
"litellm.proxy.auth.user_api_key_auth.user_api_key_auth",
new=AsyncMock(side_effect=auth_that_reserves),
):
result = await user_api_key_auth_websocket_for_model(websocket, model="gpt-realtime")
assert result.budget_reservation == reservation
assert websocket.state.budget_reservation is reservation
assert websocket.scope["state"]["budget_reservation"] is reservation

View file

@ -882,7 +882,8 @@ def test_mixed_fusion_and_client_tool_calls_reconcile_on_the_initial_response():
@pytest.mark.asyncio
async def test_cached_fusion_hidden_call_accumulates_zero_cost():
@pytest.mark.parametrize("guardrail_cost", [0.0, 0.03])
async def test_cached_fusion_hidden_call_accumulates_only_guardrail_cost(guardrail_cost: float):
logger = _ProxyDBLogger()
reservation = {
"reserved_cost": 1.0,
@ -894,6 +895,7 @@ async def test_cached_fusion_hidden_call_accumulates_zero_cost():
"call_type": "acompletion",
"model": "panel",
"cache_hit": True,
"response_cost": 0.2,
"litellm_call_id": "cached-panel-call",
"litellm_params": {
"metadata": {
@ -904,7 +906,7 @@ async def test_cached_fusion_hidden_call_accumulates_zero_cost():
}
},
"standard_logging_object": {
"response_cost": 0.2,
"response_cost": guardrail_cost,
"request_tags": [],
"metadata": {},
},
@ -931,8 +933,8 @@ async def test_cached_fusion_hidden_call_accumulates_zero_cost():
end_time=datetime.now(),
)
assert reservation[FUSION_BUDGET_ACCUMULATED_COST_KEY] == 0.0
assert proxy_logging.db_spend_update_writer.update_database.await_args.kwargs["response_cost"] == 0.0
assert reservation[FUSION_BUDGET_ACCUMULATED_COST_KEY] == guardrail_cost
assert proxy_logging.db_spend_update_writer.update_database.await_args.kwargs["response_cost"] == guardrail_cost
increment.assert_not_awaited()

View file

@ -0,0 +1,349 @@
"""
Tests for BudgetReservationReleaseMiddleware.
Auth reserves budget before the handler runs and hands the reservation to the
request or socket state. A litellm call made through the async client wrapper
claims it for the cost callback that runs after the call; anything still unclaimed
when the response is done or the socket has closed would keep the spend counter
pinned until its TTL, so the middleware releases it.
"""
import asyncio
from collections.abc import AsyncIterator, Awaitable, Callable, Mapping
from datetime import datetime
from typing import Final
import pytest
from starlette.applications import Starlette
from starlette.requests import Request
from starlette.responses import JSONResponse, Response, StreamingResponse
from starlette.routing import Route
from starlette.types import ASGIApp, Message, Receive, Scope, Send
from starlette.websockets import WebSocket
import litellm
from litellm.caching import DualCache
from litellm.proxy import proxy_server
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.middleware.budget_reservation_release_middleware import (
BudgetReservationReleaseMiddleware,
)
from litellm.proxy.spend_tracking.budget_reservation import (
reconcile_budget_reservation,
release_unbound_budget_reservation,
reserve_budget_for_request,
)
from litellm.proxy.utils import ProxyLogging
from litellm.utils import Rules, function_setup
KEY_TOKEN: Final = "hashed-release-middleware-key"
COUNTER_KEY: Final = f"spend:key:{KEY_TOKEN}"
CHAT_BODY: Final = {"model": "gpt-4o", "messages": [{"role": "user", "content": "hello"}]}
@pytest.fixture
def spend_counter_cache(monkeypatch: pytest.MonkeyPatch) -> DualCache:
cache: Final = DualCache()
monkeypatch.setattr(proxy_server, "spend_counter_cache", cache)
monkeypatch.setattr(proxy_server, "prisma_client", None)
return cache
@pytest.fixture
def no_callbacks(monkeypatch: pytest.MonkeyPatch) -> None:
for callback_list_name in (
"callbacks",
"success_callback",
"failure_callback",
"_async_success_callback",
"_async_failure_callback",
):
monkeypatch.setattr(litellm, callback_list_name, [])
async def _reserve() -> dict:
reservation: Final = await reserve_budget_for_request(
request_body=CHAT_BODY,
route="/v1/chat/completions",
llm_router=None,
valid_token=UserAPIKeyAuth(token=KEY_TOKEN, max_budget=1.0, spend=0.0),
team_object=None,
user_object=None,
prisma_client=None,
user_api_key_cache=UserApiKeyCache(),
proxy_logging_obj=ProxyLogging(user_api_key_cache=UserApiKeyCache()),
)
assert reservation is not None
assert reservation["reserved_cost"] > 0
return reservation
async def _chat(reservation: dict, **kwargs: object) -> object:
return await litellm.acompletion(
**CHAT_BODY,
metadata={"user_api_key_budget_reservation": reservation},
**kwargs,
)
def _proxy_pre_call_setup(route_type: str, reservation: dict) -> None:
function_setup(
original_function=route_type,
rules_obj=Rules(),
start_time=datetime.now(),
**CHAT_BODY,
litellm_call_id="proxy-pre-call-setup",
metadata={"user_api_key_budget_reservation": reservation},
)
def _app(
handler: Callable[[Request], Awaitable[Response]],
release: Callable[[Mapping[str, object]], Awaitable[None]] = release_unbound_budget_reservation,
) -> Starlette:
app: Final = Starlette(routes=[Route("/", handler, methods=["POST"])])
app.add_middleware(BudgetReservationReleaseMiddleware, release=release)
return app
async def _post(app: ASGIApp) -> None:
scope: Final = {
"type": "http",
"method": "POST",
"path": "/",
"raw_path": b"/",
"headers": [],
"query_string": b"",
"scheme": "http",
"server": ("testserver", 80),
"client": ("testclient", 1),
}
body_delivered: Final = asyncio.Event()
client_never_disconnects: Final = asyncio.Event()
async def receive() -> Message:
if body_delivered.is_set():
await client_never_disconnects.wait()
body_delivered.set()
return {"type": "http.request", "body": b"", "more_body": False}
async def send(message: Message) -> None:
return None
await app(scope, receive, send)
def _counter(spend_counter_cache: DualCache) -> float | None:
return spend_counter_cache.in_memory_cache.get_cache(key=COUNTER_KEY)
@pytest.mark.asyncio
async def test_unbound_reservation_is_released_after_the_response(spend_counter_cache: DualCache):
reservation: Final = await _reserve()
assert _counter(spend_counter_cache) == pytest.approx(reservation["reserved_cost"])
async def handler(request: Request) -> Response:
request.state.budget_reservation = reservation
return JSONResponse({"id": "batch_123", "status": "cancelling"})
await _post(_app(handler))
assert _counter(spend_counter_cache) == pytest.approx(0.0)
assert reservation["finalized"] is True
@pytest.mark.asyncio
async def test_unbound_reservation_is_released_when_the_handler_raises(spend_counter_cache: DualCache):
reservation: Final = await _reserve()
async def handler(request: Request) -> Response:
request.state.budget_reservation = reservation
raise RuntimeError("upstream refused the cancel")
with pytest.raises(RuntimeError):
await _post(_app(handler))
assert _counter(spend_counter_cache) == pytest.approx(0.0)
assert reservation["finalized"] is True
@pytest.mark.asyncio
async def test_reservation_seen_only_by_the_proxy_pre_call_logging_object_is_released(
spend_counter_cache: DualCache, no_callbacks: None
):
reservation: Final = await _reserve()
async def cancel_batch_without_a_client_wrapper() -> dict:
return {"id": "batch_123", "status": "cancelling"}
async def handler(request: Request) -> Response:
request.state.budget_reservation = reservation
_proxy_pre_call_setup("acancel_batch", reservation)
return JSONResponse(await cancel_batch_without_a_client_wrapper())
await _post(_app(handler))
assert _counter(spend_counter_cache) == pytest.approx(0.0)
assert reservation["finalized"] is True
@pytest.mark.asyncio
async def test_reservation_of_a_failed_call_is_released_after_the_error_response(
spend_counter_cache: DualCache, no_callbacks: None
):
reservation: Final = await _reserve()
refused: Final = litellm.AuthenticationError(message="bad key", llm_provider="openai", model="gpt-4o")
async def handler(request: Request) -> Response:
request.state.budget_reservation = reservation
_proxy_pre_call_setup("acompletion", reservation)
try:
await _chat(reservation, mock_response=refused)
except litellm.AuthenticationError:
return JSONResponse({"error": {"message": "bad key"}}, status_code=401)
raise AssertionError("the mocked call must fail")
await _post(_app(handler))
assert _counter(spend_counter_cache) == pytest.approx(0.0)
assert reservation["finalized"] is True
@pytest.mark.asyncio
async def test_reservation_claimed_by_a_completed_call_is_left_for_the_callback(
spend_counter_cache: DualCache, no_callbacks: None
):
reservation: Final = await _reserve()
reserved_cost: Final = reservation["reserved_cost"]
async def handler(request: Request) -> Response:
request.state.budget_reservation = reservation
_proxy_pre_call_setup("acompletion", reservation)
response: Final = await _chat(reservation, mock_response="ok")
return JSONResponse(response.model_dump())
await _post(_app(handler))
assert _counter(spend_counter_cache) == pytest.approx(reserved_cost)
assert reservation["finalized"] is False
@pytest.mark.asyncio
async def test_reservation_claimed_by_a_streaming_call_is_left_for_the_callback_that_finishes_after_the_response(
spend_counter_cache: DualCache, no_callbacks: None
):
reservation: Final = await _reserve()
reserved_cost: Final = reservation["reserved_cost"]
async def handler(request: Request) -> Response:
request.state.budget_reservation = reservation
_proxy_pre_call_setup("acompletion", reservation)
stream: Final = await _chat(reservation, mock_response="ok", stream=True)
async def sse() -> AsyncIterator[bytes]:
async for chunk in stream:
yield f"data: {chunk.model_dump_json()}\n\n".encode()
yield b"data: [DONE]\n\n"
return StreamingResponse(sse(), media_type="text/event-stream")
await _post(_app(handler))
assert _counter(spend_counter_cache) == pytest.approx(reserved_cost)
assert reservation["finalized"] is False
actual_cost: Final = reserved_cost / 4
await reconcile_budget_reservation(budget_reservation=reservation, actual_cost=actual_cost)
assert _counter(spend_counter_cache) == pytest.approx(actual_cost)
assert reservation["finalized"] is True
@pytest.mark.asyncio
async def test_unbound_reservation_of_a_websocket_session_is_released_when_the_socket_closes(
spend_counter_cache: DualCache,
):
reservation: Final = await _reserve()
async def listen_without_a_provider_key(scope: Scope, receive: Receive, send: Send) -> None:
websocket: Final = WebSocket(scope, receive, send)
websocket.state.budget_reservation = reservation
await websocket.close(code=1011, reason="Required 'DEEPGRAM_API_KEY' in environment")
async def receive() -> Message:
return {"type": "websocket.connect"}
async def send(message: Message) -> None:
return None
middleware: Final = BudgetReservationReleaseMiddleware(
listen_without_a_provider_key, release=release_unbound_budget_reservation
)
await middleware({"type": "websocket", "path": "/deepgram/v1/listen", "headers": []}, receive, send)
assert _counter(spend_counter_cache) == pytest.approx(0.0)
assert reservation["finalized"] is True
@pytest.mark.asyncio
async def test_release_runs_once_per_request_with_the_stamped_reservation():
released: Final = []
reservation: Final = {"reserved_cost": 0.5, "entries": [], "finalized": False, "callback_bound": False}
async def release(budget_reservation: Mapping[str, object]) -> None:
released.append(budget_reservation)
async def handler(request: Request) -> Response:
request.state.budget_reservation = reservation
return JSONResponse({})
await _post(_app(handler, release=release))
assert released == [reservation]
assert released[0] is reservation
@pytest.mark.asyncio
async def test_request_without_a_reservation_releases_nothing():
released: Final = []
async def release(budget_reservation: Mapping[str, object]) -> None:
released.append(budget_reservation)
async def unauthenticated(request: Request) -> Response:
return JSONResponse({})
async def budget_checks_skipped(request: Request) -> Response:
request.state.budget_reservation = None
return JSONResponse({})
await _post(_app(unauthenticated, release=release))
await _post(_app(budget_checks_skipped, release=release))
assert released == []
@pytest.mark.asyncio
async def test_lifespan_scopes_pass_through():
released: Final = []
seen: Final = []
async def release(budget_reservation: Mapping[str, object]) -> None:
released.append(budget_reservation)
async def inner(scope: Scope, receive: Receive, send: Send) -> None:
seen.append(scope["type"])
async def receive() -> Message:
return {"type": "lifespan.startup"}
async def send(message: Message) -> None:
return None
middleware: Final = BudgetReservationReleaseMiddleware(inner, release=release)
await middleware({"type": "lifespan", "state": {"budget_reservation": {"reserved_cost": 1.0}}}, receive, send)
assert seen == ["lifespan"]
assert released == []

View file

@ -4256,6 +4256,112 @@ async def test_pass_through_request_non_streaming_success_unchanged():
mock_success_handler.assert_called_once()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"upstream_status_code, claimed_by_the_success_handler",
[(200, True), (500, False)],
ids=["success-claims-the-reservation", "upstream-error-leaves-it-for-the-request-end-release"],
)
async def test_pass_through_request_claims_the_budget_reservation_only_when_its_success_handler_runs(
upstream_status_code: int, claimed_by_the_success_handler: bool
):
reservation: Final = {"reserved_cost": 0.5, "entries": [], "finalized": False, "callback_bound": False}
user_api_key_dict: Final = UserAPIKeyAuth(api_key="hashed")
user_api_key_dict.budget_reservation = reservation
upstream_response: Final = httpx.Response(
status_code=upstream_status_code,
headers={"content-type": "application/json"},
content=b'{"status": "upstream"}',
request=httpx.Request("POST", "http://target-api.com/api/generate"),
)
with (
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging,
patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client") as mock_get_client,
patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.ProxyBaseLLMRequestProcessing"
) as mock_processing,
patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.GLOBAL_LOGGING_WORKER") as mock_worker,
):
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_processing.get_custom_headers.return_value = {}
mock_worker.ensure_initialized_and_enqueue = MagicMock(side_effect=lambda async_coroutine: async_coroutine.close())
async_client = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/mock-upstream/api/generate"
mock_request.body = AsyncMock(return_value=b'{"prompt": "hi"}')
mock_request.headers = Headers({"content-type": "application/json"})
mock_request.query_params = QueryParams({})
response = await pass_through_request(
request=mock_request,
target="http://target-api.com/api/generate",
custom_headers={},
user_api_key_dict=user_api_key_dict,
)
assert response.status_code == upstream_status_code
assert reservation["callback_bound"] is claimed_by_the_success_handler
assert mock_worker.ensure_initialized_and_enqueue.call_count == int(claimed_by_the_success_handler)
@pytest.mark.asyncio
async def test_pass_through_request_leaves_the_budget_reservation_for_the_request_end_release_when_its_success_handler_cannot_be_enqueued():
reservation: Final = {"reserved_cost": 0.5, "entries": [], "finalized": False, "callback_bound": False}
user_api_key_dict: Final = UserAPIKeyAuth(api_key="hashed")
user_api_key_dict.budget_reservation = reservation
upstream_response: Final = httpx.Response(
status_code=200,
headers={"content-type": "application/json"},
content=b'{"status": "upstream"}',
request=httpx.Request("POST", "http://target-api.com/api/generate"),
)
def refuse_to_enqueue(async_coroutine):
async_coroutine.close()
raise RuntimeError("logging worker is shutting down")
with (
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging,
patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client") as mock_get_client,
patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.ProxyBaseLLMRequestProcessing"
) as mock_processing,
patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.GLOBAL_LOGGING_WORKER") as mock_worker,
):
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_processing.get_custom_headers.return_value = {}
mock_worker.ensure_initialized_and_enqueue = MagicMock(side_effect=refuse_to_enqueue)
async_client = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/mock-upstream/api/generate"
mock_request.body = AsyncMock(return_value=b'{"prompt": "hi"}')
mock_request.headers = Headers({"content-type": "application/json"})
mock_request.query_params = QueryParams({})
with pytest.raises(ProxyException):
await pass_through_request(
request=mock_request,
target="http://target-api.com/api/generate",
custom_headers={},
user_api_key_dict=user_api_key_dict,
)
assert reservation["callback_bound"] is False
@pytest.mark.asyncio
async def test_pass_through_request_internal_failure_still_raises_proxy_exception():
"""

View file

@ -838,3 +838,72 @@ async def test_chunk_processor_bills_partial_google_usage_on_mid_stream_exceptio
assert failure_payload["completion_tokens"] == 12
assert failure_payload["response_cost"] > 12 * 3.75e-06
assert isinstance(recorder.failure_kwargs[0]["exception"], httpx.ReadTimeout)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"deferred_dispatch_armed",
[False, True],
ids=["enqueued-at-end-of-stream", "parked-for-deferred-dispatch"],
)
async def test_chunk_processor_claims_the_budget_reservation_before_handing_it_to_the_cost_callback(
deferred_dispatch_armed: bool,
):
reservation = {"reserved_cost": 0.5, "entries": [], "finalized": False, "callback_bound": False}
response = _make_streaming_response([b"event-1", b"event-2"])
logging_obj = _unarmed_logging_obj()
logging_obj.litellm_params = {"metadata": {"user_api_key_budget_reservation": reservation}}
if deferred_dispatch_armed:
logging_obj._on_deferred_stream_complete = AsyncMock()
claimed_when_the_callback_ran = []
async def cost_callback(**kwargs):
claimed_when_the_callback_ran.append(reservation["callback_bound"])
async for _ in PassThroughStreamingHandler.chunk_processor(
response=response,
request_body={"model": "claude-3-haiku"},
litellm_logging_obj=logging_obj,
endpoint_type=EndpointType.GENERIC,
start_time=datetime.now(),
passthrough_success_handler_obj=MagicMock(),
url_route="/bedrock/model/claude/invoke-with-response-stream",
route_streaming_logging=cost_callback,
):
pass
if deferred_dispatch_armed:
(parked_cost_callback,) = logging_obj._deferred_stream_complete_args
await parked_cost_callback
else:
await GLOBAL_LOGGING_WORKER.flush()
assert reservation["callback_bound"] is True
assert claimed_when_the_callback_ran == [True]
@pytest.mark.asyncio
async def test_chunk_processor_leaves_the_budget_reservation_for_the_request_end_release_when_the_cost_callback_cannot_be_enqueued():
reservation = {"reserved_cost": 0.5, "entries": [], "finalized": False, "callback_bound": False}
response = _make_streaming_response([b"event-1", b"event-2"])
logging_obj = _unarmed_logging_obj()
logging_obj.litellm_params = {"metadata": {"user_api_key_budget_reservation": reservation}}
def refuse_to_enqueue(async_coroutine):
async_coroutine.close()
raise RuntimeError("logging worker is shutting down")
with patch.object(GLOBAL_LOGGING_WORKER, "ensure_initialized_and_enqueue", side_effect=refuse_to_enqueue):
async for _ in PassThroughStreamingHandler.chunk_processor(
response=response,
request_body={"model": "claude-3-haiku"},
litellm_logging_obj=logging_obj,
endpoint_type=EndpointType.GENERIC,
start_time=datetime.now(),
passthrough_success_handler_obj=MagicMock(),
url_route="/bedrock/model/claude/invoke-with-response-stream",
route_streaming_logging=AsyncMock(),
):
pass
assert reservation["callback_bound"] is False

View file

@ -26,6 +26,7 @@ from litellm.proxy.spend_tracking.budget_reservation import (
_get_team_member_budget_counter,
count_request_input_tokens,
estimate_request_max_cost,
release_unbound_budget_reservation,
reserve_budget_for_request,
)
from litellm.proxy.utils import ProxyLogging
@ -546,3 +547,41 @@ async def test_team_member_reservation_counter_adds_temp_increase_to_live_team_d
assert counter is not None
assert counter.max_budget == expected_max_budget
assert counter.fallback_spend == 0.5
@pytest.mark.asyncio
async def test_reservation_starts_unbound_to_any_callback():
reservation: Final = await _reserve("/v1/responses")
assert reservation is not None
assert reservation["callback_bound"] is False
@pytest.mark.asyncio
async def test_release_unbound_budget_reservation_frees_the_counter(spend_counter_cache: DualCache):
counter_key: Final = f"spend:key:{TINY_BUDGET_KEY_TOKEN}"
reservation: Final = await _reserve_for_tiny_budget_key(
"/v1/chat/completions", {"model": "gpt-4o", "messages": [{"role": "user", "content": "hello"}]}
)
assert reservation is not None
assert spend_counter_cache.in_memory_cache.get_cache(key=counter_key) == pytest.approx(reservation["reserved_cost"])
await release_unbound_budget_reservation(reservation)
assert spend_counter_cache.in_memory_cache.get_cache(key=counter_key) == pytest.approx(0.0)
assert reservation["finalized"] is True
@pytest.mark.asyncio
async def test_release_unbound_budget_reservation_leaves_a_bound_one_to_its_callback(spend_counter_cache: DualCache):
counter_key: Final = f"spend:key:{TINY_BUDGET_KEY_TOKEN}"
reservation: Final = await _reserve_for_tiny_budget_key(
"/v1/chat/completions", {"model": "gpt-4o", "messages": [{"role": "user", "content": "hello"}]}
)
assert reservation is not None
reservation["callback_bound"] = True
await release_unbound_budget_reservation(reservation)
assert spend_counter_cache.in_memory_cache.get_cache(key=counter_key) == pytest.approx(reservation["reserved_cost"])
assert reservation["finalized"] is False

View file

@ -9,9 +9,10 @@ import pytest
from pydantic import TypeAdapter
import litellm
from litellm._internal_context import is_internal_call
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.rust_bridge import callbacks_legacy_python as legacy
from litellm.rust_bridge.callbacks_legacy_python import check_limits, setup
from litellm.rust_bridge.callbacks_legacy_python import check_limits, failure_handler, setup
_OCR_KWARGS: Final = MappingProxyType(
{
@ -81,6 +82,82 @@ def test_setup_builds_a_logger_when_none_is_supplied(call_type: str, kwargs: Map
assert result.logger.litellm_call_id == result.kwargs["litellm_call_id"]
def _budget_reservation() -> dict:
return {"reserved_cost": 0.5, "entries": [], "finalized": False, "callback_bound": False}
def _kwargs_with_a_budget_reservation(reservation: dict) -> dict[str, object]:
return {**_OCR_KWARGS, "metadata": {"user_api_key_budget_reservation": reservation}}
def test_setup_claims_the_budget_reservation_for_an_async_call() -> None:
reservation: Final = _budget_reservation()
setup("aocr", (), _kwargs_with_a_budget_reservation(reservation), datetime.datetime.now(), asynchronous=True)
assert reservation["callback_bound"] is True
def test_setup_claims_the_budget_reservation_a_supplied_logger_already_saw() -> None:
reservation: Final = _budget_reservation()
supplied: Final = _supplied_logger()
supplied.update_environment_variables(
litellm_params={"metadata": {"user_api_key_budget_reservation": reservation}}, optional_params={}
)
assert reservation["callback_bound"] is False
setup("aocr", (), {**_OCR_KWARGS, "litellm_logging_obj": supplied}, datetime.datetime.now(), asynchronous=True)
assert reservation["callback_bound"] is True
def test_setup_leaves_the_budget_reservation_alone_for_a_sync_call() -> None:
reservation: Final = _budget_reservation()
setup("aocr", (), _kwargs_with_a_budget_reservation(reservation), datetime.datetime.now(), asynchronous=False)
assert reservation["callback_bound"] is False
def test_setup_leaves_the_budget_reservation_alone_for_an_internal_call() -> None:
reservation: Final = _budget_reservation()
token: Final = is_internal_call.set(True)
try:
setup("aocr", (), _kwargs_with_a_budget_reservation(reservation), datetime.datetime.now(), asynchronous=True)
finally:
is_internal_call.reset(token)
assert reservation["callback_bound"] is False
def test_failure_handler_hands_the_budget_reservation_back_for_an_async_call() -> None:
reservation: Final = _budget_reservation()
now: Final = datetime.datetime.now()
result: Final = setup("aocr", (), _kwargs_with_a_budget_reservation(reservation), now, asynchronous=True)
assert reservation["callback_bound"] is True
pending: Final = failure_handler(result.logger, RuntimeError("upstream refused"), now, now, asynchronous=True)
assert reservation["callback_bound"] is False
assert pending is not None
pending.close()
def test_failure_handler_of_an_internal_call_leaves_the_outer_budget_reservation_claim_in_place() -> None:
reservation: Final = _budget_reservation()
now: Final = datetime.datetime.now()
result: Final = setup("aocr", (), _kwargs_with_a_budget_reservation(reservation), now, asynchronous=True)
token: Final = is_internal_call.set(True)
try:
pending: Final = failure_handler(result.logger, RuntimeError("inner step failed"), now, now, asynchronous=True)
finally:
is_internal_call.reset(token)
assert reservation["callback_bound"] is True
assert pending is not None
pending.close()
CONTRACT_PATH: Final = (
Path(__file__).parents[3] / "litellm-rust/crates/callbacks-legacy-python/python_contract.json"
)

View file

@ -4984,6 +4984,100 @@ async def test_wrapper_async_fires_post_call_failure_deployment_hook_on_internal
assert isinstance(recorder.calls[0][1], litellm.AuthenticationError)
def _budget_reservation(callback_bound: bool = False) -> dict:
return {"reserved_cost": 0.5, "entries": [], "finalized": False, "callback_bound": callback_bound}
_BUDGET_RESERVATION_CALL_KWARGS: Final = {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]}
_BUDGET_RESERVATION_REFUSAL: Final = litellm.AuthenticationError(message="bad key", llm_provider="openai", model="gpt-4o")
@pytest.mark.asyncio
async def test_wrapper_async_claims_the_budget_reservation_for_the_cost_callback() -> None:
reservation = _budget_reservation()
await litellm.acompletion(
**_BUDGET_RESERVATION_CALL_KWARGS,
mock_response="ok",
metadata={"user_api_key_budget_reservation": reservation},
)
assert reservation["callback_bound"] is True
@pytest.mark.asyncio
async def test_wrapper_async_claims_the_budget_reservation_before_the_stream_is_consumed() -> None:
reservation = _budget_reservation()
stream = await litellm.acompletion(
**_BUDGET_RESERVATION_CALL_KWARGS,
mock_response="ok",
stream=True,
metadata={"user_api_key_budget_reservation": reservation},
)
assert reservation["callback_bound"] is True
async for _ in stream:
pass
@pytest.mark.asyncio
async def test_wrapper_async_claims_the_budget_reservation_a_supplied_logging_object_already_saw() -> None:
reservation = _budget_reservation()
logging_obj, kwargs = litellm.utils.function_setup(
original_function="acompletion",
rules_obj=litellm.utils.Rules(),
start_time=datetime.now(),
**_BUDGET_RESERVATION_CALL_KWARGS,
litellm_call_id="proxy-pre-call-setup",
metadata={"user_api_key_budget_reservation": reservation},
)
assert reservation["callback_bound"] is False
await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="ok")
assert reservation["callback_bound"] is True
@pytest.mark.asyncio
async def test_wrapper_async_hands_the_budget_reservation_back_when_the_call_fails() -> None:
reservation = _budget_reservation()
with pytest.raises(litellm.AuthenticationError):
await litellm.acompletion(
**_BUDGET_RESERVATION_CALL_KWARGS,
mock_response=_BUDGET_RESERVATION_REFUSAL,
metadata={"user_api_key_budget_reservation": reservation},
)
assert reservation["callback_bound"] is False
@pytest.mark.asyncio
async def test_wrapper_async_leaves_the_budget_reservation_alone_on_internal_calls() -> None:
claimed_by_the_outer_call = _budget_reservation(callback_bound=True)
never_claimed = _budget_reservation()
token = is_internal_call.set(True)
try:
await litellm.acompletion(
**_BUDGET_RESERVATION_CALL_KWARGS,
mock_response="ok",
metadata={"user_api_key_budget_reservation": never_claimed},
)
with pytest.raises(litellm.AuthenticationError):
await litellm.acompletion(
**_BUDGET_RESERVATION_CALL_KWARGS,
mock_response=_BUDGET_RESERVATION_REFUSAL,
metadata={"user_api_key_budget_reservation": claimed_by_the_outer_call},
)
finally:
is_internal_call.reset(token)
assert never_claimed["callback_bound"] is False
assert claimed_by_the_outer_call["callback_bound"] is True
@pytest.mark.asyncio
async def test_wrapper_async_does_not_fire_failure_hook_for_pre_call_budget_error(
monkeypatch: pytest.MonkeyPatch,