mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
commit
7dbc21db87
35 changed files with 2602 additions and 325 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
146
litellm/litellm_core_utils/llm_cost_calc/zero_cost_diagnostic.py
Normal file
146
litellm/litellm_core_utils/llm_cost_calc/zero_cost_diagnostic.py
Normal 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"]}"}}'
|
||||
)
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
@ -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"
|
||||
],
|
||||
|
|
|
|||
258
tests/integration/providers/test_xiaomi_mimo_wire.py
Normal file
258
tests/integration/providers/test_xiaomi_mimo_wire.py
Normal 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")]
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue