This commit is contained in:
tin-berri 2026-10-04 12:47:45 -07:00 • committed by GitHub
commit b72d644468
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
31 changed files with 2643 additions and 134 deletions

View file

@ -738,10 +738,13 @@ class LLMCachingHandler:
end_time (datetime): The end time of the operation.
cache_hit (bool): Whether it was a cache hit.
"""
from litellm.litellm_core_utils.litellm_logging import evaluation_logging_snapshot
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
receipt_logger: Final = evaluation_logging_snapshot(logging_obj)
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
async_coroutine=logging_obj.async_success_handler(
async_coroutine=receipt_logger.async_success_handler(
result=cached_result,
start_time=start_time,
end_time=end_time,

View file

@ -5,6 +5,8 @@ from collections.abc import Mapping
from datetime import datetime
from typing import Any, Final, cast
from pydantic import TypeAdapter
from litellm._logging import verbose_logger
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.integrations.datadog.datadog_handler import (
@ -14,6 +16,7 @@ from litellm.integrations.datadog.datadog_handler import (
get_datadog_service,
normalize_datadog_tag_value,
)
from litellm.litellm_core_utils.internal_call_metadata import project_evaluation_billing_kwargs
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
@ -43,6 +46,7 @@ _RESERVED_TAG_KEYS: Final[frozenset] = frozenset(
"model_group",
}
)
_COST_PAYLOAD_ADAPTER: Final = TypeAdapter(Mapping[str, object])
class DatadogCostManagementLogger(CustomBatchLogger):
@ -71,15 +75,23 @@ class DatadogCostManagementLogger(CustomBatchLogger):
super().__init__(**kwargs)
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
async def async_log_success_event(
self,
kwargs: Mapping[str, object],
response_obj: object,
start_time: datetime | float,
end_time: datetime | float,
) -> None:
try:
standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None)
if standard_logging_object is None:
receipt: Final = project_evaluation_billing_kwargs(kwargs)
payload: Final = receipt.get("standard_logging_object")
if not isinstance(payload, Mapping):
return
standard_logging_object: Final = _COST_PAYLOAD_ADAPTER.validate_python(payload)
response_cost: Final = standard_logging_object.get("response_cost")
# Only log if there is a cost associated
if standard_logging_object.get("response_cost", 0) > 0:
if isinstance(response_cost, (int, float)) and response_cost > 0:
self.log_queue.append(standard_logging_object)
if len(self.log_queue) >= self.batch_size:
@ -185,8 +197,9 @@ class DatadogCostManagementLogger(CustomBatchLogger):
metadata: Final[Mapping[str, object]] = cast(dict[str, Any], log.get("metadata") or {})
# Backwards-compat: team/user/model_group preserved regardless of allowlist.
if metadata.get("user_api_key_alias"):
tags["user"] = normalize_datadog_tag_value(metadata["user_api_key_alias"])
user_tag: Final = metadata.get("user_api_key_alias") or metadata.get("user_api_key_user_id")
if user_tag:
tags["user"] = normalize_datadog_tag_value(user_tag)
team_tag: Final = (
metadata.get("user_api_key_team_alias")
or metadata.get("team_alias")

View file

@ -11,6 +11,9 @@ import litellm
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.internal_call_metadata import (
get_evaluation_billing_owner_from_kwargs,
)
from litellm.llms.custom_httpx.http_handler import (
HTTPHandler,
get_async_httpx_client,
@ -61,6 +64,7 @@ class LagoLogger(CustomLogger):
raise Exception(f"Missing keys={missing_keys} in environment.")
def _common_logic(self, kwargs: dict, response_obj) -> dict:
billing_owner: Final = get_evaluation_billing_owner_from_kwargs(kwargs)
response_obj.get("id", kwargs.get("litellm_call_id"))
get_utc_datetime().isoformat()
cost: Final = kwargs.get("response_cost", None)
@ -96,7 +100,9 @@ class LagoLogger(CustomLogger):
else:
raise Exception("invalid LAGO_API_CHARGE_BY set")
if charge_by == "end_user_id":
if billing_owner is not None:
external_customer_id = billing_owner.user_id
elif charge_by == "end_user_id":
external_customer_id = end_user_id
elif charge_by == "team_id":
external_customer_id = team_id

View file

@ -9,6 +9,9 @@ import httpx
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.internal_call_metadata import (
get_evaluation_billing_owner_from_kwargs,
)
from litellm.llms.custom_httpx.http_handler import (
HTTPHandler,
get_async_httpx_client,
@ -49,6 +52,7 @@ class OpenMeterLogger(CustomLogger):
raise Exception(f"Missing keys={missing_keys} in environment.")
def _common_logic(self, kwargs: dict, response_obj):
billing_owner: Final = get_evaluation_billing_owner_from_kwargs(kwargs)
call_id: Final = response_obj.get("id", kwargs.get("litellm_call_id"))
dt: Final = get_utc_datetime().isoformat()
cost: Final = kwargs.get("response_cost", None)
@ -69,7 +73,8 @@ class OpenMeterLogger(CustomLogger):
# serving multi-tenant traffic enable this to prevent clients from
# forging attribution by setting `user` in the request body.
trust_request_user: Final = os.getenv("OPENMETER_TRUST_REQUEST_USER", "true").lower() != "false"
user_param = kwargs.get("user", None) if trust_request_user else None
request_user: Final = kwargs.get("user") if trust_request_user else None
user_param = billing_owner.user_id if billing_owner is not None else request_user
# If no user provided directly, try to get it from token user_id
if user_param is None:

View file

@ -30,7 +30,11 @@ from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.websearch_interception.tools import is_web_search_tool_responses
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs, independent_snapshot
from litellm.litellm_core_utils.internal_call_metadata import sanitized_forwardable_call_metadata
from litellm.litellm_core_utils.internal_call_metadata import (
EvaluationBillingOwner,
evaluation_billing_context,
sanitized_forwardable_call_metadata,
)
from litellm.litellm_core_utils.llm_judge import (
default_router_provider,
extract_text_from_content,
@ -76,6 +80,7 @@ _EMPTY_METADATA: Final[Mapping[str, object]] = MappingProxyType({})
_CHAT_REQUEST_ADAPTER: Final = TypeAdapter(Mapping[str, object])
_CHAT_MESSAGES_ADAPTER: Final = TypeAdapter(tuple[Mapping[str, object], ...])
_MESSAGE_ITEMS_ADAPTER: Final = TypeAdapter(tuple[object, ...])
_CREATOR_MODEL_BUDGET: Final = TypeAdapter(Mapping[str, object] | None)
def _chat_messages(kwargs: Mapping[str, object]) -> tuple[Mapping[str, object], ...]:
@ -747,6 +752,7 @@ class ActiveShadowEvalJob(BaseModel):
baseline_model: str | None = None
shadow_percentage: float
judge_model: str
created_by: str | None = None
max_turns: int
max_budget: float | None = None
ends_at: datetime
@ -1082,22 +1088,27 @@ class ShadowEvalLogger(CustomLogger):
if spend >= job.max_budget:
self._record_funnel(job.id, "withheld")
return
for arm_router in job.arm_router_names:
await self._run_shadow_arm(
prisma=prisma,
job=job,
arm_router=arm_router,
request_id=request_id,
messages=messages,
real_text=real_text,
real_model=real_model,
real_cost=real_cost,
real_classifier_cost=real_classifier_cost,
real_cache_hit=real_cache_hit,
control_tier=control_tier,
shadow_params=shadow_params,
parent_metadata=parent_metadata,
)
owner: Final = await _evaluation_billing_owner(prisma, job.created_by)
if owner is None:
self._record_funnel(job.id, "withheld")
return
with evaluation_billing_context(owner):
for arm_router in job.arm_router_names:
await self._run_shadow_arm(
prisma=prisma,
job=job,
arm_router=arm_router,
request_id=request_id,
messages=messages,
real_text=real_text,
real_model=real_model,
real_cost=real_cost,
real_classifier_cost=real_classifier_cost,
real_cache_hit=real_cache_hit,
control_tier=control_tier,
shadow_params=shadow_params,
parent_metadata=parent_metadata,
)
async def _run_shadow_arm(
self,
@ -1380,3 +1391,32 @@ def _default_prisma_provider() -> "PrismaClient | None":
except ImportError:
return None
return prisma_client
async def _evaluation_billing_owner(prisma: "PrismaClient", created_by: str | None) -> EvaluationBillingOwner | None:
from litellm.proxy.auth.auth_checks import get_user_object
from litellm.proxy.proxy_server import litellm_proxy_admin_name, proxy_logging_obj, user_api_key_cache
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
creator_id: Final = created_by or litellm_proxy_admin_name
try:
creator: Final = await get_user_object(
user_id=creator_id,
prisma_client=prisma,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
proxy_logging_obj=proxy_logging_obj,
)
except UserNotFoundError:
return EvaluationBillingOwner(creator_id) if creator_id == litellm_proxy_admin_name else None
except Exception as e: # noqa: BLE001 # optional evaluation work must not spend against unverifiable limits
verbose_logger.warning("shadow_eval: creator budget unavailable for %s: %s", creator_id, e)
return None
if creator is None:
return None
return EvaluationBillingOwner(
creator_id,
_CREATOR_MODEL_BUDGET.validate_python(creator.model_max_budget),
creator.max_budget,
creator.spend or 0.0,
)

View file

@ -1,9 +1,8 @@
"""Metadata a request forwards to the internal LLM sub-calls it triggers.
Internal features (the auto-router's classifier and embeddings, shadow eval's shadow and
judge calls) bill real provider spend that nobody typed a prompt for. That spend must land
on the same key/team/org/user as the request that caused it, so the sub-call carries the
caller's identity metadata, minus two things that must never be forwarded as-is:
Internal calls retain the caller's routing identity. Ordinary classifier and embedding
calls also retain its billing identity; shadow evaluations project billing receipts onto
the evaluation creator without changing routing metadata. Two fields need special handling:
* ``user_api_key_budget_reservation`` (and the reservation nested inside
``user_api_key_auth``) belongs to the parent completion. If a sub-call's cost callback
@ -17,16 +16,154 @@ caller's identity metadata, minus two things that must never be forwarded as-is:
from __future__ import annotations
from collections.abc import Mapping
from collections.abc import Generator, Mapping
from contextlib import contextmanager
from contextvars import ContextVar
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final
from pydantic import TypeAdapter
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY, NON_INFERENCE_CALL_TYPES
from litellm.litellm_core_utils.initialize_dynamic_callback_params import initialize_standard_callback_dynamic_params
from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN, InternalCallOrigin
BUDGET_RESERVATION_METADATA_KEYS: Final = frozenset({"user_api_key_budget_reservation"})
@dataclass(frozen=True, slots=True)
class EvaluationBillingOwner:
user_id: str
user_model_max_budget: Mapping[str, object] | None = None
max_budget: float | None = None
spend: float = 0.0
EVALUATION_BILLING_OWNER_KEY: Final = "_evaluation_billing_owner"
EVALUATION_BUDGET_RESERVATION_KEY: Final = "_evaluation_budget_reservation"
_EVALUATION_BILLING_OWNER: Final[ContextVar[EvaluationBillingOwner | None]] = ContextVar(
"evaluation_billing_owner", default=None
)
_BILLING_MAPPING: Final = TypeAdapter(Mapping[str, object])
_EMPTY_BILLING_FIELDS: Final[Mapping[str, object]] = MappingProxyType({})
_BILLING_IDENTITY_FIELDS: Final = frozenset(
{"user_api_key", "user_api_end_user_max_budget", "team_id", "team_alias", "agent_id", "billing_agent_id"}
)
def get_evaluation_billing_owner() -> EvaluationBillingOwner | None:
return _EVALUATION_BILLING_OWNER.get()
@contextmanager
def evaluation_billing_context(owner: EvaluationBillingOwner | None) -> Generator[None]:
token: Final = _EVALUATION_BILLING_OWNER.set(owner)
try:
yield
finally:
_EVALUATION_BILLING_OWNER.reset(token)
def get_evaluation_billing_owner_from_kwargs(kwargs: Mapping[str, object]) -> EvaluationBillingOwner | None:
owner: Final = kwargs.get(EVALUATION_BILLING_OWNER_KEY)
return owner if isinstance(owner, EvaluationBillingOwner) else None
def _billing_mapping(value: object) -> Mapping[str, object] | None:
return _BILLING_MAPPING.validate_python(value) if isinstance(value, Mapping) else None
def _evaluation_budget_reservation(kwargs: Mapping[str, object]) -> Mapping[str, object] | None:
handle: Final = kwargs.get(EVALUATION_BUDGET_RESERVATION_KEY)
if handle is None:
return None
from litellm.proxy.spend_tracking.evaluation_budget import EvaluationBudgetReservation
return handle.total if isinstance(handle, EvaluationBudgetReservation) else None
def _billing_metadata(
metadata: Mapping[str, object] | None, owner: EvaluationBillingOwner, reservation: Mapping[str, object] | None
) -> Mapping[str, object]:
return {
**{
key: None if key.startswith("user_api_key_") or key in _BILLING_IDENTITY_FIELDS else item
for key, item in (metadata or _EMPTY_BILLING_FIELDS).items()
},
"user_api_key_user_id": owner.user_id,
"user_api_key_user_model_max_budget": owner.user_model_max_budget,
"user_api_key_budget_reservation": reservation,
"tags": [],
}
def _billing_request(value: object) -> object:
request: Final = _billing_mapping(value)
if request is None:
return value
body: Final = _billing_mapping(request.get("body"))
if body is None:
return dict(request)
return {**request, "body": {**body, "user": None}}
def _billing_fields(
fields: Mapping[str, object], owner: EvaluationBillingOwner, reservation: Mapping[str, object] | None
) -> Mapping[str, object]:
alternate: Final = _billing_mapping(fields.get("litellm_metadata"))
return {
**fields,
"agent_id": None,
"billing_agent_id": None,
**({"user": owner.user_id} if "user" in fields else {}),
"metadata": _billing_metadata(_billing_mapping(fields.get("metadata")), owner, reservation),
**({"litellm_metadata": _billing_metadata(alternate, owner, reservation)} if alternate else {}),
**(
{"proxy_server_request": _billing_request(fields["proxy_server_request"])}
if "proxy_server_request" in fields
else {}
),
}
def project_evaluation_billing_kwargs(
kwargs: Mapping[str, object],
) -> dict[str, object]: # mutable-ok: existing SDK receipt consumers require dictionaries
"""Copy an evaluation's receipt onto its creator without changing request state."""
owner: Final = get_evaluation_billing_owner_from_kwargs(kwargs)
if owner is None:
return kwargs if isinstance(kwargs, dict) else dict(kwargs)
reservation: Final = _evaluation_budget_reservation(kwargs)
params: Final = _billing_mapping(kwargs.get("litellm_params")) or _EMPTY_BILLING_FIELDS
payload: Final = _billing_mapping(kwargs.get("standard_logging_object"))
receipt_fields: Final[Mapping[str, object]] = {
"user": owner.user_id,
"end_user": None,
"request_tags": [],
"request_model_access_groups": (),
}
return {
**_billing_fields(kwargs, owner, reservation),
**receipt_fields,
"user_api_key_end_user_id": None,
"litellm_params": {
**_billing_fields(params, owner, reservation),
"user_api_key_end_user_id": None,
},
**(
{
"standard_logging_object": {
**_billing_fields(payload, owner, reservation),
**receipt_fields,
}
}
if payload is not None
else {}
),
}
MODEL_ACCESS_GROUP_METADATA_KEY: Final = "user_api_key_matched_model_access_groups"
"""Where auth records the model access groups that authorized the request, for the spend writer.

View file

@ -11,13 +11,14 @@ import sys
import time
import traceback
from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime as dt_object
from functools import lru_cache
from types import MappingProxyType, TracebackType
from typing import TYPE_CHECKING, Any, Final, Literal, Union, cast
from httpx import Response
from pydantic import BaseModel, JsonValue
from pydantic import BaseModel, JsonValue, TypeAdapter
import litellm
from litellm import _custom_logger_compatible_callbacks_literal
@ -79,7 +80,11 @@ from litellm.litellm_core_utils.core_helpers import (
from litellm.litellm_core_utils.error_normalization import normalize_error
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
from litellm.litellm_core_utils.internal_call_metadata import (
EVALUATION_BILLING_OWNER_KEY,
EVALUATION_BUDGET_RESERVATION_KEY,
MODEL_ACCESS_GROUP_METADATA_KEY,
EvaluationBillingOwner,
get_evaluation_billing_owner,
is_unbilled_non_inference_call,
)
from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import (
@ -236,6 +241,7 @@ if TYPE_CHECKING:
from litellm.litellm_core_utils.llm_cost_calc.utils import BilledTokenRates
from litellm.llms.base_llm.passthrough.transformation import PassthroughStreamCollector
from litellm.proxy.hooks.autorouter_baseline_cache import BaselineCacheContext, CapturedBaselineObservation
from litellm.proxy.spend_tracking.evaluation_budget import EvaluationBudgetReservation
try:
from litellm_enterprise.enterprise_callbacks.callback_controls import (
EnterpriseCallbackControls,
@ -553,6 +559,19 @@ def _timestamp_seconds(moment: object) -> float | None:
return None
@dataclass(slots=True)
class EvaluationBudgetInvocation:
reservation: "EvaluationBudgetReservation | None" = None
def evaluation_logging_snapshot(logging_obj: "Logging") -> "Logging":
return (
copy.copy(logging_obj)
if isinstance(logging_obj.evaluation_billing_owner, EvaluationBillingOwner)
else logging_obj
)
class Logging(LiteLLMLoggingBaseClass):
global \
supabaseClient, \
@ -618,6 +637,9 @@ class Logging(LiteLLMLoggingBaseClass):
self.call_type = call_type
self.litellm_call_id = litellm_call_id
self.litellm_trace_id: str = litellm_trace_id if litellm_trace_id else str(uuid.uuid4())
self.evaluation_billing_owner: Final[EvaluationBillingOwner | None] = get_evaluation_billing_owner()
self.evaluation_budget_reservation: EvaluationBudgetReservation | None = None
self.evaluation_budget_invocation: EvaluationBudgetInvocation | None = None
# Capture the pre-call *value* (not a contextvars.Token) so restoration works
# even if this attempt's own logging ends up dispatched onto a different
@ -714,6 +736,8 @@ class Logging(LiteLLMLoggingBaseClass):
"litellm_params": litellm_params,
"applied_guardrails": applied_guardrails,
"model": model,
EVALUATION_BILLING_OWNER_KEY: self.evaluation_billing_owner,
EVALUATION_BUDGET_RESERVATION_KEY: self.evaluation_budget_reservation,
}
# Set by proxy request handlers to defer spend-log fire until after
@ -945,6 +969,8 @@ class Logging(LiteLLMLoggingBaseClass):
"standard_callback_dynamic_params": self.standard_callback_dynamic_params,
**self.optional_params,
**additional_params,
EVALUATION_BILLING_OWNER_KEY: self.evaluation_billing_owner,
EVALUATION_BUDGET_RESERVATION_KEY: self.evaluation_budget_reservation,
}
)
@ -2216,6 +2242,16 @@ class Logging(LiteLLMLoggingBaseClass):
if isinstance(usage, Usage):
self.record_partial_usage_for_failure(usage, self._response_cost_calculator(result=assembled) or 0.0)
def recover_failure_cost(self, result: object) -> float:
if isinstance(result, ModelResponse):
self.record_assembled_response_for_failure(result)
elif isinstance(result, ResponsesAPIResponse):
self.record_assembled_response_for_failure(self._translate_responses_api_response_to_model_response(result))
self.model_call_details["response_cost"] = self._response_cost_calculator(result=result)
elif result is not None and self.call_type == CallTypes.anthropic_messages.value:
self.record_assembled_response_for_failure(self._handle_anthropic_messages_response_logging(result))
return TypeAdapter(float).validate_python(self.model_call_details.get("response_cost") or 0.0)
async def dispatch_failure_handlers(
self,
exception: Exception,

View file

@ -39,6 +39,7 @@ from litellm.integrations.otel.model.config import is_otel_v2_enabled
from litellm.integrations.otel.runtime import phase_event, phase_span, seed_request_identity
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value
from litellm.litellm_core_utils.internal_call_metadata import get_evaluation_billing_owner
from litellm.proxy._types import *
from litellm.proxy.agent_endpoints.auth.agent_caller import agent_caller_from_headers
from litellm.proxy.auth.auth_checks import (
@ -3224,7 +3225,11 @@ async def _reserve_budget_after_common_checks(
request: Request | None = None,
) -> None:
user_api_key_auth_obj.budget_reservation = None
if not skip_budget_checks and general_settings.get("disable_budget_reservation") is not True:
if (
not skip_budget_checks
and general_settings.get("disable_budget_reservation") is not True
and get_evaluation_billing_owner() is None
):
from litellm.proxy.spend_tracking.budget_reservation import (
reserve_budget_for_request,
)

View file

@ -3,10 +3,13 @@ import json
import time
from collections.abc import Iterable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime
from types import MappingProxyType
from typing import Final
from typing import Annotated, Final
from openai.types import Batch
from pydantic import BeforeValidator, TypeAdapter
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._internal_context import with_service_target
@ -14,12 +17,17 @@ from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import Span
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.litellm_core_utils.internal_call_metadata import (
EVALUATION_BUDGET_RESERVATION_KEY,
get_evaluation_billing_owner_from_kwargs,
project_evaluation_billing_kwargs,
)
from litellm.llms.bedrock.common_utils import get_bedrock_base_model
from litellm.proxy._types import Litellm_EntityType, UserAPIKeyAuth
from litellm.router_strategy.budget_limiter import RouterBudgetLimiting
from litellm.router_utils.batch_utils import is_batch_retrieve_call_type
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import BudgetConfig, StandardLoggingPayload
from litellm.types.utils import BudgetConfig
VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX: Final = "virtual_key_spend"
END_USER_SPEND_CACHE_KEY_PREFIX: Final = "end_user_model_spend"
@ -49,6 +57,31 @@ _BUDGET_START_TIME_KEY_PREFIXES: Final = MappingProxyType(
)
class _ModelBudgetLogIdentity(TypedDict, total=False):
user_api_key_hash: ReadOnly[str | None]
user_api_key_user_id: ReadOnly[str | None]
user_api_key_team_id: ReadOnly[str | None]
user_api_key_end_user_id: ReadOnly[str | None]
def _optional_end_user(value: object) -> str | None:
return value if isinstance(value, str) else None
class _ModelBudgetLogPayload(TypedDict, total=False):
model_group: ReadOnly[str | None]
model: ReadOnly[str | None]
response_cost: ReadOnly[float | None]
metadata: ReadOnly[_ModelBudgetLogIdentity | None]
end_user: ReadOnly[Annotated[str | None, BeforeValidator(_optional_end_user)]]
_MODEL_BUDGET_LOG_PAYLOAD: Final = TypeAdapter(_ModelBudgetLogPayload)
_MODEL_BUDGET_MAPPING: Final = TypeAdapter(Mapping[str, object])
_EMPTY_MODEL_BUDGET_MAPPING: Final[Mapping[str, object]] = MappingProxyType({})
_EMPTY_MODEL_BUDGET_IDENTITY: Final[_ModelBudgetLogIdentity] = {}
@dataclass(frozen=True, slots=True)
class ResolvedModelBudget:
"""The `model_max_budget` entry a request resolved to.
@ -429,7 +462,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
if max_budget is None or max_budget < 0:
return True
current_spend: Final = await self._get_spend_for_model_budget(
current_spend: Final = await self.get_spend_for_model_budget(
entity_type=entity_type,
entity_id=entity_id,
model=model,
@ -445,7 +478,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
)
return True
async def _get_spend_for_model_budget(
async def get_spend_for_model_budget(
self,
entity_type: Litellm_EntityType,
entity_id: str | None,
@ -464,13 +497,19 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
budget_model=resolved.budget_model,
budget_duration=resolved.budget_config.budget_duration,
)
from litellm.proxy.spend_tracking.evaluation_budget import model_budget_spend
current_spend: Final = (
await model_budget_spend(self.dual_cache, spend_key)
if entity_type == Litellm_EntityType.USER
else _as_spend(await self._cached_spend(spend_key))
)
legacy_spend_key: Final = _legacy_request_model_spend_cache_key(
entity_type=entity_type,
entity_id=entity_id,
model=model,
resolved=resolved,
)
current_spend: Final = _as_spend(await self._cached_spend(spend_key))
if legacy_spend_key is None or legacy_spend_key == spend_key:
return current_spend
return current_spend + _as_spend(await self._cached_spend(legacy_spend_key))
@ -493,23 +532,35 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
return healthy_deployments
@with_service_target("model_budgets")
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
async def async_log_success_event(
self,
kwargs: Mapping[str, object],
response_obj: object,
start_time: datetime | None,
end_time: datetime | None,
) -> None:
"""
Track spend for virtual key + model in DualCache
Example: key=sk-1234567890, model=gpt-4o, max_budget=100, time_period=1d
"""
verbose_proxy_logger.debug("in RouterBudgetLimiting.async_log_success_event")
standard_logging_payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None)
if standard_logging_payload is None:
receipt: Final = project_evaluation_billing_kwargs(kwargs)
payload_value: Final = receipt.get("standard_logging_object")
if payload_value is None:
verbose_proxy_logger.debug(
"Skipping _PROXY_VirtualKeyModelMaxBudgetLimiter.async_log_success_event: standard_logging_payload is None"
)
return
_litellm_params: Final[dict] = kwargs.get("litellm_params", {}) or {}
_metadata: Final[dict] = _litellm_params.get("metadata", {}) or {}
payload_metadata: Final = standard_logging_payload.get("metadata") or {}
standard_logging_payload: Final = _MODEL_BUDGET_LOG_PAYLOAD.validate_python(payload_value)
_litellm_params: Final = _MODEL_BUDGET_MAPPING.validate_python(
receipt.get("litellm_params") or _EMPTY_MODEL_BUDGET_MAPPING
)
_metadata: Final = _MODEL_BUDGET_MAPPING.validate_python(
_litellm_params.get("metadata") or _EMPTY_MODEL_BUDGET_MAPPING
)
payload_metadata: Final = standard_logging_payload.get("metadata") or _EMPTY_MODEL_BUDGET_IDENTITY
# Use model_group (the user-facing model alias, e.g. "gpt-4o") when
# available. The enforcement path receives the model name from
@ -521,7 +572,16 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
if model is None:
return
response_cost: Final[float] = standard_logging_payload.get("response_cost", 0)
response_cost: Final = standard_logging_payload.get("response_cost", 0)
if response_cost is None:
return
if get_evaluation_billing_owner_from_kwargs(kwargs) is not None:
from litellm.proxy.spend_tracking.evaluation_budget import EvaluationBudgetReservation
reservation: Final = kwargs.get(EVALUATION_BUDGET_RESERVATION_KEY)
if isinstance(reservation, EvaluationBudgetReservation) and reservation.model is not None:
await reservation.model.settle(response_cost)
return
key_model_max_budget: Final = _metadata.get("user_api_key_model_max_budget")
entity_budgets: Final = (
(
@ -537,7 +597,9 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
if team_model_budget_applies(
model=model,
key_model_max_budget=(
key_model_max_budget if isinstance(key_model_max_budget, Mapping) else None
_MODEL_BUDGET_MAPPING.validate_python(key_model_max_budget)
if isinstance(key_model_max_budget, Mapping)
else None
),
)
else None
@ -565,7 +627,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
return
batch_id: Final = batch_id_to_charge_once(
call_type=kwargs.get("call_type"),
call_type=receipt.get("call_type"),
response_obj=response_obj,
response_cost=response_cost,
)
@ -586,6 +648,25 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
json.dumps(self.dual_cache.in_memory_cache.cache_dict, indent=4, default=str),
)
async def async_log_failure_event(
self,
kwargs: Mapping[str, object],
response_obj: object,
start_time: datetime | None,
end_time: datetime | None,
) -> None:
if get_evaluation_billing_owner_from_kwargs(kwargs) is None:
return
from litellm.proxy.spend_tracking.evaluation_budget import (
EvaluationBudgetReservation,
release_evaluation_budget,
)
reservation: Final = kwargs.get(EVALUATION_BUDGET_RESERVATION_KEY)
if isinstance(reservation, EvaluationBudgetReservation):
payload: Final = _MODEL_BUDGET_LOG_PAYLOAD.validate_python(kwargs.get("standard_logging_object") or kwargs)
await release_evaluation_budget(reservation, actual_cost=payload.get("response_cost") or 0.0)
async def _charge_entity(
self,
entity_type: Litellm_EntityType,

View file

@ -2,6 +2,7 @@ import asyncio
import traceback
from collections.abc import Callable, Mapping, Sequence
from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Protocol, cast
import litellm
@ -15,6 +16,12 @@ from litellm.litellm_core_utils.core_helpers import (
get_litellm_metadata_from_kwargs,
get_metadata_variable_name_from_kwargs,
)
from litellm.litellm_core_utils.internal_call_metadata import (
EVALUATION_BILLING_OWNER_KEY,
get_evaluation_billing_owner,
get_evaluation_billing_owner_from_kwargs,
project_evaluation_billing_kwargs,
)
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import guardrail_information_cost
from litellm.proxy._types import UserAPIKeyAuth
@ -121,11 +128,12 @@ class _ProxyDBLogger(CustomLogger):
async def async_log_success_event(
self, kwargs: ObjectMapping, response_obj: object, start_time: datetime, end_time: datetime
) -> None:
receipt: Final = project_evaluation_billing_kwargs(kwargs)
if self.spend_event_producer is None or not is_offloadable_success(response_obj):
await self._PROXY_track_cost_callback(kwargs, response_obj, start_time, end_time)
await self._PROXY_track_cost_callback(receipt, response_obj, start_time, end_time)
return
event: Final = build_spend_event(
kwargs,
receipt,
response_obj,
start_time,
end_time,
@ -133,7 +141,7 @@ class _ProxyDBLogger(CustomLogger):
)
if isinstance(event, SpendEventBuildError):
verbose_proxy_logger.warning("collector: tracking cost in-process, event not buildable: %s", event.reason)
await self._PROXY_track_cost_callback(kwargs, response_obj, start_time, end_time)
await self._PROXY_track_cost_callback(receipt, response_obj, start_time, end_time)
return
await self.spend_event_producer.publish(event)
@ -274,18 +282,23 @@ class _ProxyDBLogger(CustomLogger):
existing_metadata.get("standard_logging_guardrail_information")
)
billing_owner: Final = get_evaluation_billing_owner_from_kwargs(request_data) or get_evaluation_billing_owner()
receipt: Final = project_evaluation_billing_kwargs(
MappingProxyType({**request_data, EVALUATION_BILLING_OWNER_KEY: billing_owner})
)
await self._spend_writer().update_database(
token=LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict),
token=None if billing_owner is not None else LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict),
response_cost=recovered_response_cost,
user_id=user_api_key_dict.user_id,
end_user_id=user_api_key_dict.end_user_id,
team_id=user_api_key_dict.team_id,
kwargs=request_data,
user_id=billing_owner.user_id if billing_owner is not None else user_api_key_dict.user_id,
end_user_id=None if billing_owner is not None else user_api_key_dict.end_user_id,
team_id=None if billing_owner is not None else user_api_key_dict.team_id,
kwargs=receipt,
completion_response=original_exception,
start_time=actual_start_time,
end_time=datetime.now(),
org_id=user_api_key_dict.org_id,
project_id=user_api_key_dict.project_id,
org_id=None if billing_owner is not None else user_api_key_dict.org_id,
project_id=None if billing_owner is not None else user_api_key_dict.project_id,
)
async def _PROXY_track_cost_callback(

View file

@ -1684,10 +1684,10 @@ async def start_shadow_eval(
eval spend, the shadow and judge calls' own cost, reaches max_budget dollars, the
job's window ends, or the job is stopped, so one target running out of budget does
not end sampling for the others; sampling changes propagate to pods within about 10
seconds. Shadow and judge calls bill to the sampled request's own identity but are
seconds. Shadow and judge calls bill to the admin who started the job and are
excluded from request counts and auto-router adoption metrics.
"""
from litellm.proxy.proxy_server import llm_router, prisma_client
from litellm.proxy.proxy_server import litellm_proxy_admin_name, llm_router, prisma_client
_require_admin_writer(user_api_key_dict, "start a shadow eval")
if prisma_client is None:
@ -1805,7 +1805,7 @@ async def start_shadow_eval(
"shadow_percentage": data.shadow_percentage,
"max_turns": SHADOW_EVAL_TURN_VALVE,
"max_budget": data.max_budget,
"created_by": user_api_key_dict.user_id,
"created_by": user_api_key_dict.user_id or litellm_proxy_admin_name,
"created_at": now,
"ends_at": ends_at,
}

View file

@ -13,6 +13,10 @@ from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
inherit_message_logging_privacy,
initialize_standard_callback_dynamic_params,
)
from litellm.litellm_core_utils.internal_call_metadata import (
evaluation_billing_context,
get_evaluation_billing_owner,
)
from litellm.llms.custom_httpx.asgi_handler import get_async_asgi_client
from litellm.proxy.litellm_pre_call_utils import UNTRUSTED_REQUEST_HEADER_CONTROL_FIELDS
from litellm.router_strategy.complexity_router.context_compaction import (
@ -40,6 +44,7 @@ async def with_proxy_compaction_executor(call: Awaitable[_ResultT], request: Req
protocol: Literal["chat", "messages"], payload: Mapping[str, object], parent_model: str | None = None
) -> Mapping[str, object]:
logging_disabled: Final = initialize_standard_callback_dynamic_params().get("turn_off_message_logging") is True
billing_owner: Final = get_evaluation_billing_owner()
async def dispatch() -> Mapping[str, object]:
scope: Final = _JSON_OBJECT.validate_python(request.scope)
@ -55,6 +60,7 @@ async def with_proxy_compaction_executor(call: Awaitable[_ResultT], request: Req
with (
native_compaction_call(parent_model, str(payload["model"])),
inherit_message_logging_privacy(logging_disabled),
evaluation_billing_context(billing_owner),
):
with get_async_asgi_client(
app=_ASGI_APP.validate_python(scope["app"]),

View file

@ -111,7 +111,11 @@ def get_reserved_counter_keys(budget_reservation: dict | None) -> set:
_lease_renewals: Final[set[asyncio.Task[None]]] = set() # mutable-ok: asyncio only weak-refs pending tasks
def _start_reservation_lease_renewal(budget_reservation: Mapping[str, object], counter_keys: frozenset[str]) -> None:
def _start_reservation_lease_renewal(
budget_reservation: Mapping[str, object],
counter_keys: frozenset[str],
request_task: asyncio.Task[object] | None = None,
) -> None:
"""A reservation lives inside spend counter keys that expire on their Redis TTL. Renew the TTL
while the request is in flight so a request longer than the TTL does not drop its
reservation and admit concurrent requests against the DB floor on any worker."""
@ -124,7 +128,7 @@ def _start_reservation_lease_renewal(budget_reservation: Mapping[str, object], c
budget_reservation=budget_reservation,
counter_keys=counter_keys,
interval=spend_counter_cache.redis_cache.default_ttl / 2,
request_task=asyncio.current_task(),
request_task=request_task or asyncio.current_task(),
)
)
_lease_renewals.add(task)
@ -248,7 +252,7 @@ def _is_unbilled_route(route: str) -> bool:
async def reserve_budget_for_request(
request_body: dict,
request_body: dict[str, object], # mutable-ok: existing reservation helpers accept dictionaries
route: str,
llm_router: Router | None,
valid_token: UserAPIKeyAuth | None,
@ -262,7 +266,8 @@ async def reserve_budget_for_request(
apply_user_budget_to_team_keys: bool = False,
fail_closed_budget_enforcement: bool = False,
raw_body: bytes | None = None,
) -> dict | None:
request_task: asyncio.Task[object] | None = None,
) -> dict[str, object] | None: # mutable-ok: shared reservation is finalized by the spend writer
if valid_token is None or not RouteChecks.is_llm_api_route(route=route):
return None
if _is_unbilled_route(route):
@ -334,7 +339,7 @@ async def reserve_budget_for_request(
llm_router=llm_router,
input_token_counts=input_token_counts,
)
budget_reservation: Final = {
budget_reservation: Final[dict[str, object]] = { # mutable-ok: shared finalization state
"reserved_cost": reservation_cost,
"entries": applied_entries,
"finalized": False,
@ -345,6 +350,7 @@ async def reserve_budget_for_request(
_start_reservation_lease_renewal(
budget_reservation=budget_reservation,
counter_keys=frozenset(get_reserved_counter_keys(budget_reservation=budget_reservation)),
request_task=request_task,
)
return budget_reservation
@ -1302,7 +1308,7 @@ def _coerce_datetime(value: object) -> datetime | None:
def estimate_request_max_cost(
request_body: dict,
request_body: dict[str, object], # mutable-ok: existing cost helpers accept dictionaries
route: str,
llm_router: Router | None,
input_token_counts: Mapping[str, int] | None = None,
@ -1324,7 +1330,7 @@ def estimate_request_max_cost(
def estimate_request_input_cost(
request_body: dict,
request_body: dict[str, object], # mutable-ok: existing cost helpers accept dictionaries
route: str,
llm_router: Router | None,
input_token_counts: Mapping[str, int] | None = None,

View file

@ -0,0 +1,444 @@
from __future__ import annotations
import asyncio
import math
import time
import uuid
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass, field
from functools import lru_cache
from itertools import product
from string import ascii_letters, digits
from threading import Lock
from typing import Final, Literal, Protocol, TypeVar, runtime_checkable
from pydantic import ConfigDict, TypeAdapter
from redis.crc import key_slot
import litellm
from litellm._internal_context import with_service_target
from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.litellm_core_utils.internal_call_metadata import EvaluationBillingOwner
from litellm.proxy._types import Litellm_EntityType, LiteLLM_UserTable, UserAPIKeyAuth
from litellm.proxy.hooks.model_max_budget_limiter import (
ResolvedModelBudget,
model_budget_spend_cache_key,
model_budget_start_time_cache_key,
resolve_model_budget,
)
from litellm.proxy.spend_tracking import budget_reservation
from litellm.proxy.spend_tracking.budget_reservation import (
estimate_request_input_cost,
estimate_request_max_cost,
reserve_budget_for_request,
)
from litellm.router import Router
from litellm.router_utils.common_utils import resolve_model_group_alias
from litellm.types.utils import API_ROUTE_TO_CALL_TYPES, CallTypes
_NUMBER: Final = TypeAdapter(float)
_METADATA: Final = TypeAdapter(Mapping[str, object])
_REQUEST: Final = TypeAdapter(dict[str, object])
_HOLDS: Final = TypeAdapter(Mapping[str, float])
_RECONCILE: Final = TypeAdapter[Callable[[dict[str, object] | None, float | None], Awaitable[object]]](
Callable[[dict[str, object] | None, float | None], Awaitable[object]]
).validate_python(
budget_reservation.reconcile_budget_reservation # pyright: ignore[reportUnknownMemberType] # legacy reservation parameter is untyped
)
_LEASE_SECONDS: Final = 60
_LOCAL_HOLDS_LOCK: Final = Lock()
_LEASE_TASKS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: asyncio only weakly references pending tasks
_Result: Final = TypeVar("_Result")
_HOLD_SCRIPT: Final = """
local actual = tonumber(redis.call('GET', KEYS[2]) or '0')
if not actual then return redis.error_reply('Invalid model budget spend') end
local clock = redis.call('TIME')
local now = tonumber(clock[1]) + tonumber(clock[2]) / 1000000
local total = 0
for _, member in ipairs(redis.call('ZRANGEBYSCORE', KEYS[1], '(' .. now, '+inf')) do
if member ~= ARGV[2] then total = total + tonumber(string.match(member, ':([^:]+)$')) end
end
if ARGV[1] == 'renew' then
local expiry = redis.call('ZSCORE', KEYS[1], ARGV[2])
if not expiry or tonumber(expiry) <= now then
return redis.error_reply('Evaluation budget reservation expired')
end
end
if ARGV[1] == 'reserve' or ARGV[1] == 'renew' then
total = total + tonumber(string.match(ARGV[2], ':([^:]+)$'))
end
if ARGV[1] == 'reserve' and ARGV[6] ~= '' then
local proposed = actual + total
local estimate = tonumber(string.match(ARGV[2], ':([^:]+)$'))
if proposed > tonumber(ARGV[6]) or proposed - estimate >= tonumber(ARGV[6]) then
return tostring(proposed)
end
end
if ARGV[1] == 'settle' and tonumber(ARGV[4]) > 0 then
actual = tonumber(redis.call('INCRBYFLOAT', KEYS[2], ARGV[4]))
if redis.call('TTL', KEYS[2]) < 0 then redis.call('EXPIRE', KEYS[2], ARGV[5]) end
end
redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', now)
if ARGV[1] == 'reserve' or ARGV[1] == 'renew' then
redis.call('ZADD', KEYS[1], now + tonumber(ARGV[3]), ARGV[2])
redis.call('EXPIRE', KEYS[1], ARGV[3])
elseif ARGV[1] == 'settle' then
redis.call('ZREM', KEYS[1], ARGV[2])
end
return tostring(actual + total)
"""
@runtime_checkable
class _NumericCache(Protocol):
def get_cache(self, key: str) -> object: ...
def set_cache(self, key: str, value: object, *, ttl: int) -> object: ...
def delete_cache(self, key: str) -> object: ...
def increment_cache(self, key: str, value: float, *, ttl: int) -> float: ...
@runtime_checkable
class _Script(Protocol):
async def __call__(self, *, keys: Sequence[str], args: Sequence[str | int | float]) -> object: ...
_NUMERIC_CACHE: Final = TypeAdapter(_NumericCache, config=ConfigDict(arbitrary_types_allowed=True))
_SCRIPT: Final = TypeAdapter(_Script, config=ConfigDict(arbitrary_types_allowed=True))
@lru_cache(maxsize=4096)
def _model_hold_key(effective_spend_key: str) -> str:
prefix: Final = f"{effective_spend_key}:evaluation_holds:"
tagged: Final = f"{prefix}{{{effective_spend_key}}}"
slot: Final = key_slot(effective_spend_key.encode())
if key_slot(tagged.encode()) == slot:
return tagged
# Empty or unmatched braces in the effective namespace/key can disable Redis hash tags.
candidates: Final = (prefix + "".join(chars) for chars in product(ascii_letters + digits + "-_", repeat=3))
return next(candidate for candidate in candidates if key_slot(candidate.encode()) == slot)
@with_service_target("model_budgets")
async def model_budget_spend(
cache: DualCache,
spend_key: str,
*,
operation: Literal["read", "reserve", "renew", "settle"] = "read",
member: str = "",
adjustment: float = 0.0,
ttl: int = 1,
limit: float | None = None,
) -> float:
if cache.redis_cache is not None:
effective_spend_key: Final = cache.redis_cache.check_and_fix_namespace(spend_key)
script: Final = _SCRIPT.validate_python(cache.redis_cache.async_register_script(_HOLD_SCRIPT))
return _NUMBER.validate_python(
await script(
keys=(_model_hold_key(effective_spend_key), effective_spend_key),
args=(operation, member, _LEASE_SECONDS, adjustment, ttl, "" if limit is None else limit),
)
)
key: Final = _model_hold_key(spend_key)
local: Final = _NUMERIC_CACHE.validate_python(cache.in_memory_cache)
with _LOCAL_HOLDS_LOCK:
actual: Final = _NUMBER.validate_python(local.get_cache(spend_key) or 0.0)
now: Final = cache.in_memory_cache._clock() # pyright: ignore[reportPrivateUsage] # use the cache's injected clock for lease expiry
current: Final = _HOLDS.validate_python(local.get_cache(key) or {})
active: Final = {token: expiry for token, expiry in current.items() if expiry > now}
if operation == "renew" and member not in active:
raise RuntimeError("Evaluation budget reservation expired")
added: Final = (
{member: now + _LEASE_SECONDS}
if operation == "reserve" or (operation == "renew" and member in active)
else {}
)
kept: Final = {token: expiry for token, expiry in active.items() if token != member}
updated: Final = {**kept, **added}
held: Final = sum(float(token.rsplit(":", 1)[1]) for token in updated)
proposed: Final = actual + held
if (
operation == "reserve"
and limit is not None
and (proposed > limit or proposed - float(member.rsplit(":", 1)[1]) >= limit)
):
return proposed
settled: Final = (
local.increment_cache(spend_key, adjustment, ttl=ttl) if operation == "settle" and adjustment else actual
)
local.delete_cache(key)
if updated:
local.set_cache(key, updated, ttl=_LEASE_SECONDS)
return settled + held
async def _drain(task: asyncio.Future[_Result]) -> _Result:
while not task.done():
try:
await asyncio.shield(task)
except asyncio.CancelledError:
continue
return task.result()
async def _complete(operation: Awaitable[_Result]) -> _Result:
task: Final = asyncio.ensure_future(operation)
try:
return await asyncio.shield(task)
except asyncio.CancelledError:
await _drain(task)
raise
@with_service_target("model_budgets")
async def _model_window(cache: DualCache, start_key: str, duration: int) -> int:
if cache.redis_cache is not None:
script: Final = _SCRIPT.validate_python(
cache.redis_cache.async_register_script(
"local clock = redis.call('TIME'); "
"local now = tonumber(clock[1]) + tonumber(clock[2]) / 1000000; "
"local start = tonumber(redis.call('GET', KEYS[1]) or now); "
"if now - start >= tonumber(ARGV[1]) then start = now end; "
"local ttl = math.max(1, math.ceil(tonumber(ARGV[1]) - (now - start))); "
"redis.call('SET', KEYS[1], start, 'EX', ttl); "
"return ttl"
)
)
return math.ceil(_NUMBER.validate_python(await script(keys=(start_key,), args=(duration,))))
local: Final = _NUMERIC_CACHE.validate_python(cache.in_memory_cache)
now: Final = time.time()
cached: Final = local.get_cache(start_key)
previous: Final = _NUMBER.validate_python(now if cached is None else cached)
start: Final = now if now - previous >= duration else previous
ttl: Final = max(1, math.ceil(duration - (now - start)))
local.delete_cache(start_key)
local.set_cache(key=start_key, value=start, ttl=ttl)
return ttl
@dataclass(slots=True)
class EvaluationModelReservation:
cache: DualCache
spend_key: str
start_key: str
duration: int
reserved_cost: float
member: str = field(init=False)
settled_cost: float | None = None
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
def __post_init__(self) -> None:
self.member = f"{uuid.uuid4()}:{self.reserved_cost}"
async def settle(self, actual_cost: float) -> None:
await _complete(self._settle(actual_cost))
async def _settle(self, actual_cost: float) -> None:
async with self.lock:
adjustment: Final = max(actual_cost - (self.settled_cost or 0.0), 0.0)
ttl: Final = await _model_window(self.cache, self.start_key, self.duration) if adjustment else 1
await model_budget_spend(
self.cache, self.spend_key, operation="settle", member=self.member, adjustment=adjustment, ttl=ttl
)
self.settled_cost = max(actual_cost, self.settled_cost or 0.0)
async def renew(self, request_task: asyncio.Task[object] | None) -> None:
# Allow the request timeout again for queued success callbacks to settle the hold.
deadline: Final = time.monotonic() + 2 * litellm.request_timeout
while self.settled_cost is None:
await asyncio.sleep(_LEASE_SECONDS / 2)
async with self.lock:
if self.settled_cost is not None:
return
if time.monotonic() >= deadline:
if request_task is not None and not request_task.done():
request_task.cancel()
return
try:
await model_budget_spend(self.cache, self.spend_key, operation="renew", member=self.member)
except Exception: # noqa: BLE001 # any lease backend failure must stop an unreserved provider call
verbose_proxy_logger.exception("Unable to renew evaluation budget reservation")
if request_task is not None and not request_task.done():
request_task.cancel()
return
@dataclass(slots=True)
class EvaluationBudgetReservation:
total: dict[str, object] | None = None # mutable-ok: shared reservation is finalized by the spend writer
model: EvaluationModelReservation | None = None
input_cost: float = 0.0
settled_failure_cost: float = 0.0
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
async def settle_failure(self, actual_cost: float) -> None:
await _complete(self._settle_failure(actual_cost))
async def _settle_failure(self, actual_cost: float) -> None:
async with self.lock:
cost: Final = max(actual_cost, self.settled_failure_cost)
self.settled_failure_cost = cost
total: Final = {**self.total, "finalized": False} if self.total is not None else None
try:
await _RECONCILE(total, cost)
if self.total is not None:
self.total["finalized"] = True
finally:
if self.model is not None:
await self.model.settle(cost)
def _pricing_model(model: str, router: Router | None, model_info: Mapping[str, object]) -> str:
if router is None:
return model
deployment_id: Final = model_info.get("id")
deployment: Final = router.get_deployment(deployment_id) if isinstance(deployment_id, str) else None
if deployment is not None:
return deployment.model_name
return resolve_model_group_alias(router.model_group_alias, model) or model
async def reserve_evaluation_budget(
owner: EvaluationBillingOwner, request: Mapping[str, object], call_type: str
) -> EvaluationBudgetReservation | None:
task: Final = asyncio.create_task(_reserve_evaluation_budget(owner, request, call_type, asyncio.current_task()))
try:
return await asyncio.shield(task)
except asyncio.CancelledError:
try:
reservation: Final = await _drain(task)
except Exception: # noqa: BLE001 # acquisition already rolled back; preserve the caller's cancellation
raise asyncio.CancelledError from None
await _complete(release_evaluation_budget(reservation))
raise
async def _reserve_evaluation_budget(
owner: EvaluationBillingOwner,
request: Mapping[str, object],
call_type: str,
request_task: asyncio.Task[object] | None,
) -> EvaluationBudgetReservation | None:
from litellm.proxy.proxy_server import (
get_current_spend,
llm_router,
model_max_budget_limiter,
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
metadata: Final = _METADATA.validate_python(request.get("litellm_metadata") or request.get("metadata") or {})
model: Final = str(metadata.get("model_group") or request["model"])
resolved: Final = resolve_model_budget(model, owner.user_model_max_budget or {})
model_budget: Final = (
resolved
if resolved is not None
and resolved.budget_config.max_budget is not None
and math.isfinite(resolved.budget_config.max_budget)
and resolved.budget_config.max_budget >= 0
else None
)
total_budget: Final = owner.max_budget is not None and math.isfinite(owner.max_budget)
if not total_budget and model_budget is None:
return None
if total_budget:
current: Final = await get_current_spend(
counter_key=f"spend:user:{owner.user_id}", fallback_spend=owner.spend, max_budget=owner.max_budget
)
if current >= _NUMBER.validate_python(owner.max_budget):
raise litellm.BudgetExceededError(
current_cost=current,
max_budget=_NUMBER.validate_python(owner.max_budget),
entity_type=Litellm_EntityType.USER.value,
entity_id=owner.user_id,
)
route: Final = next(route for route, types in API_ROUTE_TO_CALL_TYPES.items() if CallTypes(call_type) in types)
model_info: Final = _METADATA.validate_python(request.get("model_info") or metadata.get("model_info") or {})
body: Final = _REQUEST.validate_python(
{
**request,
"model": _pricing_model(model, llm_router, model_info),
"metadata": {},
"litellm_metadata": {},
"tags": [],
}
)
reservation: Final = EvaluationBudgetReservation()
try:
reservation.total = await reserve_budget_for_request(
request_body=body,
route=route,
llm_router=llm_router,
valid_token=UserAPIKeyAuth(user_id=owner.user_id),
team_object=None,
user_object=LiteLLM_UserTable(user_id=owner.user_id, max_budget=owner.max_budget, spend=owner.spend),
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
fail_closed_budget_enforcement=True,
request_task=request_task,
)
estimate: Final = (
_NUMBER.validate_python(reservation.total["reserved_cost"])
if reservation.total is not None
else estimate_request_max_cost(body, route, llm_router)
)
if estimate is None or not math.isfinite(estimate) or estimate < 0:
raise ValueError("Evaluation budget cannot be checked for an unpriced model")
if model_budget is not None:
reservation.model = _model_reservation(owner, model_budget, estimate, model_max_budget_limiter.dual_cache)
limit: Final = _NUMBER.validate_python(model_budget.budget_config.max_budget)
spend: Final = await model_budget_spend(
reservation.model.cache,
reservation.model.spend_key,
operation="reserve",
member=reservation.model.member,
limit=limit,
)
if spend > limit or spend - estimate >= limit:
raise litellm.BudgetExceededError(
current_cost=spend - estimate,
max_budget=limit,
entity_type=Litellm_EntityType.USER.value,
entity_id=owner.user_id,
)
lease: Final = asyncio.create_task(reservation.model.renew(request_task))
_LEASE_TASKS.add(lease)
lease.add_done_callback(_LEASE_TASKS.discard)
reservation.input_cost = (
_NUMBER.validate_python(reservation.total["input_cost"])
if reservation.total is not None
else estimate_request_input_cost(body, route, llm_router) or 0.0
)
except Exception:
await reservation.settle_failure(0.0)
raise
return reservation
def _model_reservation(
owner: EvaluationBillingOwner, resolved: ResolvedModelBudget, estimate: float, cache: DualCache
) -> EvaluationModelReservation:
duration: Final = str(resolved.budget_config.budget_duration)
return EvaluationModelReservation(
cache=cache,
spend_key=model_budget_spend_cache_key(Litellm_EntityType.USER, owner.user_id, resolved.budget_model, duration),
start_key=model_budget_start_time_cache_key(
Litellm_EntityType.USER, owner.user_id, resolved.budget_model, duration
),
duration=duration_in_seconds(duration),
reserved_cost=estimate,
)
async def release_evaluation_budget(
reservation: EvaluationBudgetReservation | None, *, cancelled: bool = False, actual_cost: float = 0.0
) -> None:
if reservation is not None:
await reservation.settle_failure(max(reservation.input_cost if cancelled else 0.0, actual_cost))

View file

@ -45,7 +45,7 @@ from httpx import Proxy
from httpx._utils import get_environment_proxies
from openai.lib import _parsing, _pydantic
from openai.types.chat.completion_create_params import ResponseFormat
from pydantic import BaseModel
from pydantic import BaseModel, TypeAdapter
import litellm
import litellm.litellm_core_utils
@ -1379,11 +1379,14 @@ def _schedule_async_success_logging(
first-wins rule: the innermost wrapper's provider-shaped result is the one the spend log
reads usage from, and a later wrapper never swaps in its client-shaped translation.
"""
from litellm.litellm_core_utils.litellm_logging import evaluation_logging_snapshot
receipt_logger: Final = evaluation_logging_snapshot(logging_obj)
def _enqueue_async_logging() -> None:
asyncio.create_task(
_client_async_logging_helper(
logging_obj=logging_obj,
logging_obj=receipt_logger,
result=result,
start_time=start_time,
end_time=end_time,
@ -1965,11 +1968,25 @@ def client(original_function):
@wraps(original_function)
async def wrapper_async(*args, **kwargs):
from litellm.litellm_core_utils.internal_call_metadata import (
EVALUATION_BUDGET_RESERVATION_KEY,
EvaluationBillingOwner,
get_evaluation_billing_owner,
)
from litellm.litellm_core_utils.litellm_logging import EvaluationBudgetInvocation, Logging
print_args_passed_to_litellm(original_function, args, kwargs)
start_time: Final = datetime.datetime.now()
result = None
_update_response_metadata: Final[_ResponseMetadataUpdater] = litellm_utils.update_response_metadata
logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None)
evaluation_invocation: Final = (
EvaluationBudgetInvocation()
if get_evaluation_billing_owner() is not None
or isinstance(logging_obj, Logging)
and logging_obj.evaluation_billing_owner is not None
else None
)
LLMCachingHandler: Final = _get_cached_llm_caching_handler()
_llm_caching_handler: Final[LLMCachingHandler] = LLMCachingHandler(
original_function=original_function,
@ -1982,7 +1999,14 @@ def client(original_function):
kwargs["litellm_call_id"] = str(uuid.uuid4())
model: Final[str | None] = args[0] if len(args) > 0 else kwargs.get("model", None)
is_completion_with_fallbacks: Final = kwargs.get("fallbacks") is not None
is_completion_with_fallbacks: Final = (
kwargs.get("fallbacks")
or (
litellm.model_fallbacks
if call_type in (CallTypes.acompletion.value, CallTypes.atext_completion.value)
else None
)
) is not None
kwargs.pop("_is_litellm_internal_call", None) # discard if injected
_is_litellm_internal_call: Final = is_internal_call.get()
_deployment_call_end_time: datetime.datetime | None = None
@ -1993,6 +2017,24 @@ 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 (
evaluation_invocation is not None
and isinstance(logging_obj, Logging)
and isinstance(logging_obj.evaluation_billing_owner, EvaluationBillingOwner)
and not _is_litellm_internal_call
and not is_completion_with_fallbacks
and logging_obj.evaluation_budget_invocation is None
):
logging_obj.evaluation_budget_invocation = evaluation_invocation
logging_obj.evaluation_budget_reservation = None
prior_receipt: Final = TypeAdapter(Mapping[str, object]).validate_python(logging_obj.model_call_details)
logging_obj.model_call_details = {
key: value
for key, value in prior_receipt.items()
if not key.startswith("has_logged_")
and key not in ("response_cost", "standard_logging_object", "combined_usage_object")
}
logging_obj.model_call_details[EVALUATION_BUDGET_RESERVATION_KEY] = None
if not _is_litellm_internal_call:
bind_budget_reservation_to_callbacks(logging_obj.litellm_params)
@ -2089,6 +2131,23 @@ def client(original_function):
and _caching_handler_response.embedding_uncached_input is not None
else kwargs
)
if (
evaluation_invocation is not None
and isinstance(logging_obj, Logging)
and isinstance(logging_obj.evaluation_billing_owner, EvaluationBillingOwner)
and logging_obj.evaluation_budget_invocation is evaluation_invocation
):
from litellm.proxy.spend_tracking.evaluation_budget import reserve_evaluation_budget
evaluation_invocation.reservation = await reserve_evaluation_budget(
logging_obj.evaluation_billing_owner,
TypeAdapter(dict[str, object]).validate_python(
{**call_kwargs, "model": model, "messages": logging_obj.messages}
),
TypeAdapter(str).validate_python(call_type),
)
logging_obj.evaluation_budget_reservation = evaluation_invocation.reservation
logging_obj.model_call_details[EVALUATION_BUDGET_RESERVATION_KEY] = evaluation_invocation.reservation
try:
result = await original_function(*args, **call_kwargs)
except Exception as deployment_error:
@ -2196,7 +2255,23 @@ def client(original_function):
)
return result
except Exception as e:
except (Exception, asyncio.CancelledError) as e:
if (
evaluation_invocation is not None
and isinstance(logging_obj, Logging)
and logging_obj.evaluation_budget_invocation is evaluation_invocation
):
from litellm.proxy.spend_tracking.evaluation_budget import release_evaluation_budget
await asyncio.shield(
release_evaluation_budget(
evaluation_invocation.reservation,
cancelled=isinstance(e, asyncio.CancelledError),
actual_cost=logging_obj.recover_failure_cost(result),
)
)
if isinstance(e, asyncio.CancelledError):
raise
traceback_exception: Final = traceback.format_exc()
# Reuse the timestamp taken right when the deployment call itself failed, before
# the failure hook ran, so a slow callback doesn't inflate the reported duration.
@ -2216,7 +2291,10 @@ def client(original_function):
call_type = original_function.__name__
num_retries, kwargs = _get_wrapper_num_retries(kwargs=kwargs, exception=e)
if call_type == CallTypes.acompletion.value:
sdk_retries_enabled: Final = not isinstance(logging_obj, Logging) or not isinstance(
logging_obj.evaluation_billing_owner, EvaluationBillingOwner
)
if call_type == CallTypes.acompletion.value and sdk_retries_enabled:
context_window_fallback_dict: Final = kwargs.get("context_window_fallback_dict", {})
_is_litellm_router_call = "model_group" in (
@ -2251,7 +2329,7 @@ def client(original_function):
kwargs["model"] = context_window_fallback_dict[model]
result = await original_function(*args, **kwargs)
return result
elif call_type == CallTypes.aresponses.value:
elif call_type == CallTypes.aresponses.value and sdk_retries_enabled:
_is_litellm_router_call = "model_group" in (
kwargs.get("metadata") or {}
) # check if call from litellm.router/proxy
@ -2282,6 +2360,12 @@ def client(original_function):
raise e
finally:
if (
evaluation_invocation is not None
and isinstance(logging_obj, Logging)
and logging_obj.evaluation_budget_invocation is evaluation_invocation
):
logging_obj.evaluation_budget_invocation = None
# Restore trace_id/session_id contextvars to their pre-call value once
# this call (in this asyncio Task) is fully done - see
# request_correlation_in_logs. Unlike wrapper()'s sync path, it's safe to

View file

@ -0,0 +1,151 @@
from __future__ import annotations
import asyncio
import os
import uuid
from collections.abc import Iterator, Sequence
from types import SimpleNamespace
from typing import Final
import pytest
from redis import Redis
import litellm
from litellm.caching.caching import DualCache
from litellm.caching.redis_cache import RedisCache
from litellm.litellm_core_utils.internal_call_metadata import EvaluationBillingOwner
from litellm.proxy import proxy_server
from litellm.proxy.hooks.model_max_budget_limiter import _PROXY_VirtualKeyModelMaxBudgetLimiter
from litellm.proxy.spend_tracking.budget_reservation import estimate_request_max_cost
from litellm.proxy.spend_tracking.evaluation_budget import (
_SCRIPT,
_model_hold_key,
model_budget_spend,
release_evaluation_budget,
reserve_evaluation_budget,
)
MODEL: Final = "openai/evaluation-budget-integration"
REQUEST: Final = {"model": MODEL, "messages": [{"role": "user", "content": "hello"}], "max_tokens": 10}
@pytest.fixture
def budget(monkeypatch: pytest.MonkeyPatch) -> Iterator[tuple[DualCache, EvaluationBillingOwner, float]]:
namespace: Final = f"evaluation-budget-{uuid.uuid4().hex}"
cache: Final = DualCache(
redis_cache=RedisCache(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"]), namespace=namespace)
)
monkeypatch.setattr(proxy_server, "llm_router", None)
monkeypatch.setattr(proxy_server, "model_max_budget_limiter", _PROXY_VirtualKeyModelMaxBudgetLimiter(cache))
monkeypatch.setitem(
litellm.model_cost,
MODEL,
{"input_cost_per_token": 0.001, "output_cost_per_token": 0.002, "max_output_tokens": 1000},
)
estimate: Final = estimate_request_max_cost(REQUEST, "/chat/completions", None)
assert estimate is not None and estimate > 0
owner: Final = EvaluationBillingOwner("creator", {MODEL: {"max_budget": 2 * estimate, "budget_duration": "1d"}})
try:
yield cache, owner, estimate
finally:
with Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) as raw:
keys: Final = tuple(raw.scan_iter(match=f"{namespace}:*"))
if keys:
raw.delete(*keys)
@pytest.mark.asyncio
async def test_settlement_before_reserve_reply_does_not_count_a_request_twice(
budget: tuple[DualCache, EvaluationBillingOwner, float],
) -> None:
cache, owner, estimate = budget
earlier: Final = await reserve_evaluation_budget(owner, REQUEST, "acompletion")
assert earlier is not None and earlier.model is not None
backend: Final = cache.redis_cache
assert backend is not None
def delayed_register(source: str) -> object:
execute: Final = _SCRIPT.validate_python(backend.async_register_script(source))
async def execute_then_settle(*, keys: Sequence[str], args: Sequence[str | int | float]) -> object:
result: Final = await execute(keys=keys, args=args)
if args and args[0] == "reserve":
await release_evaluation_budget(earlier, actual_cost=estimate / 4)
return result
return execute_then_settle
cache.redis_cache = SimpleNamespace( # pyright: ignore[reportAttributeAccessIssue] # injected transport boundary forwards every script to real Redis
check_and_fix_namespace=backend.check_and_fix_namespace,
async_register_script=delayed_register,
)
later: Final = await reserve_evaluation_budget(owner, REQUEST, "acompletion")
assert later is not None
assert await model_budget_spend(cache, earlier.model.spend_key) == pytest.approx(estimate * 1.25)
await release_evaluation_budget(later)
assert await model_budget_spend(cache, earlier.model.spend_key) == pytest.approx(estimate / 4)
@pytest.mark.asyncio
@pytest.mark.parametrize("marker_age", (30, 86401))
async def test_persistent_budget_marker_preserves_the_remaining_spend_window(
budget: tuple[DualCache, EvaluationBillingOwner, float], marker_age: int
) -> None:
cache, owner, estimate = budget
reservation: Final = await reserve_evaluation_budget(owner, REQUEST, "acompletion")
assert reservation is not None and reservation.model is not None
backend: Final = cache.redis_cache
assert backend is not None
model: Final = reservation.model
start_key: Final = backend.check_and_fix_namespace(model.start_key)
spend_key: Final = backend.check_and_fix_namespace(model.spend_key)
with Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) as raw:
raw.set(start_key, raw.time()[0] - marker_age)
assert raw.ttl(start_key) == -1
await release_evaluation_budget(reservation, actual_cost=estimate / 4)
expected: Final = model.duration - marker_age if marker_age < model.duration else model.duration
assert expected - 2 <= raw.ttl(spend_key) <= expected
assert expected - 2 <= raw.ttl(start_key) <= expected
assert float(raw.get(spend_key)) == pytest.approx(estimate / 4)
assert raw.zcard(_model_hold_key(spend_key)) == 0
@pytest.mark.asyncio
async def test_cancelled_settlement_drains_the_atomic_redis_write(
budget: tuple[DualCache, EvaluationBillingOwner, float],
) -> None:
cache, owner, estimate = budget
reservation: Final = await reserve_evaluation_budget(owner, REQUEST, "acompletion")
assert reservation is not None and reservation.model is not None
backend: Final = cache.redis_cache
assert backend is not None
written: Final = asyncio.Event()
deliver: Final = asyncio.Event()
def delayed_register(source: str) -> object:
execute: Final = _SCRIPT.validate_python(backend.async_register_script(source))
async def execute_then_wait(*, keys: Sequence[str], args: Sequence[str | int | float]) -> object:
result: Final = await execute(keys=keys, args=args)
if args and args[0] == "settle":
written.set()
await deliver.wait()
return result
return execute_then_wait
cache.redis_cache = SimpleNamespace( # pyright: ignore[reportAttributeAccessIssue] # injected transport boundary delays only the real Redis reply
check_and_fix_namespace=backend.check_and_fix_namespace,
async_register_script=delayed_register,
)
pending: Final = asyncio.create_task(release_evaluation_budget(reservation, actual_cost=estimate / 4))
await asyncio.wait_for(written.wait(), 5)
pending.cancel()
await asyncio.sleep(0)
pending.cancel()
assert not pending.done()
deliver.set()
with pytest.raises(asyncio.CancelledError):
await pending
await release_evaluation_budget(reservation, actual_cost=estimate / 4)
assert await model_budget_spend(cache, reservation.model.spend_key) == pytest.approx(estimate / 4)

View file

@ -1,4 +1,5 @@
import time
from typing import Final
from unittest.mock import AsyncMock
import pytest
@ -7,6 +8,7 @@ from httpx import Request, Response
from litellm.integrations.datadog.datadog_cost_management import (
DatadogCostManagementLogger,
)
from litellm.litellm_core_utils import internal_call_metadata as billing
from litellm.types.utils import StandardLoggingPayload
@ -86,14 +88,26 @@ async def test_aggregate_costs(clean_env):
@pytest.mark.asyncio
async def test_async_log_success_event(clean_env):
@pytest.mark.parametrize("evaluation", (False, True))
async def test_async_log_success_event(clean_env: None, evaluation: bool) -> None:
"""
Test that logs are added to queue
"""
logger = DatadogCostManagementLogger(batch_size=10)
logger = DatadogCostManagementLogger(batch_size=10, cost_tag_keys=["billing_agent_id", "department"])
payload: Final = StandardLoggingPayload(
response_cost=0.01, total_tokens=5,
metadata={
"user_api_key_user_id": "sampled-user", "user_api_key_alias": "sampled-key",
"user_api_key_team_id": "sampled-team", "billing_agent_id": "sampled-agent",
},
request_tags=["department:sampled"],
)
await logger.async_log_success_event(
kwargs={"standard_logging_object": {"response_cost": 0.01}},
kwargs={
"standard_logging_object": payload,
billing.EVALUATION_BILLING_OWNER_KEY: billing.EvaluationBillingOwner("admin") if evaluation else None,
},
response_obj={},
start_time=time.time(),
end_time=time.time(),
@ -101,6 +115,15 @@ async def test_async_log_success_event(clean_env):
assert len(logger.log_queue) == 1
assert logger.log_queue[0]["response_cost"] == 0.01
assert logger.log_queue[0]["total_tokens"] == 5
entry: Final = logger._aggregate_costs(logger.log_queue)[0]
assert entry["BilledCost"] == 0.01
assert entry["Tags"] is not None
assert entry["Tags"]["user"] == ("admin" if evaluation else "sampled-key")
assert entry["Tags"].get("team") == (None if evaluation else "sampled-team")
assert entry["Tags"].get("billing_agent_id") == (None if evaluation else "sampled-agent")
assert entry["Tags"].get("department") == (None if evaluation else "sampled")
assert payload["metadata"]["user_api_key_user_id"] == "sampled-user"
# Test zero cost ignored
await logger.async_log_success_event(

View file

@ -1,11 +1,14 @@
import json
import os
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import litellm
from litellm.integrations.openmeter import OpenMeterLogger
from litellm.integrations.lago import LagoLogger
from litellm.litellm_core_utils import internal_call_metadata as billing
class TestOpenMeterIntegration:
@ -200,15 +203,31 @@ class TestOpenMeterIntegration:
assert isinstance(data["subject"], str)
assert data["data"]["model"] == "gpt-4"
def test_cloudevents_structure(self):
@pytest.mark.parametrize("evaluation", (False, True))
@pytest.mark.parametrize("trust_request_user", (False, True))
@pytest.mark.parametrize("charge_by", ("end_user_id", "user_id", "team_id"))
@patch.dict(
os.environ, LAGO_API_BASE="https://billing.test", LAGO_API_KEY="test",
LAGO_API_EVENT_CODE="eval",
)
def test_cloudevents_structure(
self, monkeypatch: pytest.MonkeyPatch, evaluation: bool, trust_request_user: bool, charge_by: str,
) -> None:
"""Test that the CloudEvents structure is correct"""
monkeypatch.setenv("OPENMETER_TRUST_REQUEST_USER", str(trust_request_user).lower())
monkeypatch.setenv("LAGO_API_CHARGE_BY", charge_by)
logger = OpenMeterLogger()
kwargs = {
billing.EVALUATION_BILLING_OWNER_KEY: billing.EvaluationBillingOwner("admin") if evaluation else None,
"user": "cloudevents-test-user",
"model": "gpt-3.5-turbo",
"response_cost": 0.001,
"litellm_call_id": "cloudevents-test-call-id",
"litellm_params": {
"metadata": {"user_api_key_user_id": "key-user", "user_api_key_team_id": "key-team"},
"proxy_server_request": {"body": {"user": "cloudevents-test-user"}},
},
}
response_data = {
@ -226,7 +245,9 @@ class TestOpenMeterIntegration:
assert result["source"] == "litellm-proxy"
assert "time" in result
assert isinstance(result["subject"], str)
assert result["subject"] == "cloudevents-test-user"
assert result["subject"] == (
"admin" if evaluation else "cloudevents-test-user" if trust_request_user else "key-user"
)
# Verify data structure
assert "data" in result
@ -234,6 +255,14 @@ class TestOpenMeterIntegration:
assert result["data"]["cost"] == 0.001
assert result["data"]["prompt_tokens"] == 15
assert result["data"]["completion_tokens"] == 8
event: Final = LagoLogger()._common_logic(kwargs, response_obj)["event"]
source_identity: Final = {
"end_user_id": "cloudevents-test-user", "user_id": "key-user", "team_id": "key-team",
}
assert event["external_subscription_id"] == ("admin" if evaluation else source_identity[charge_by])
assert event["properties"]["response_cost"] == result["data"]["cost"]
assert event["properties"]["total_tokens"] == response_obj.usage.total_tokens
assert kwargs["user"] == "cloudevents-test-user"
assert result["data"]["total_tokens"] == 23
def test_custom_event_type(self, monkeypatch):

View file

@ -2,7 +2,7 @@
the detached pipeline's single attempt-row write, and the cache-first job lookup."""
import asyncio
from collections.abc import Mapping
from collections.abc import Awaitable, Callable, Mapping
from datetime import datetime, timedelta, timezone
from typing import Final, Literal
from unittest.mock import AsyncMock, MagicMock
@ -12,6 +12,7 @@ from pydantic import ValidationError
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
from litellm.models.user import LiteLLM_UserTable
from litellm.integrations.shadow_eval_logger import (
_MAX_CONCURRENT_SHADOW_TASKS,
_MAX_ERROR_CHARS,
@ -52,6 +53,7 @@ def _job(**overrides) -> ActiveShadowEvalJob:
router_name="my-router",
shadow_percentage=100.0,
judge_model="judge-model",
created_by="evaluation-admin",
max_turns=200,
ends_at=datetime.now(timezone.utc) + timedelta(days=1),
attempts=0,
@ -74,6 +76,7 @@ def _prisma(jobs=(), attempt_counts=(), attempt_costs=()) -> MagicMock:
]
)
prisma.db.litellm_shadowevalattempt.create = AsyncMock()
prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=LiteLLM_UserTable(user_id="evaluation-admin"))
return prisma
@ -90,6 +93,7 @@ def _job_record(job: ActiveShadowEvalJob, target_type="key", target_id="key-hash
baseline_model=job.baseline_model,
shadow_percentage=job.shadow_percentage,
judge_model=job.judge_model,
created_by=job.created_by,
max_turns=job.max_turns,
max_budget=job.max_budget,
ends_at=job.ends_at,
@ -103,6 +107,8 @@ def _router(
judge_json='{"preference": "A", "confidence": 0.9, "reasoning": "x"}',
classifier_cost=None,
sibling_router_texts=None,
*,
on_call: Callable[[Mapping[str, object]], Awaitable[None]] | None = None,
):
"""One mock router serving the shadow call first, the judge call second, told apart by
the internal-origin stamp rather than the model, since a reverse job's shadow arm names
@ -114,6 +120,8 @@ def _router(
router.get_model_list = MagicMock(return_value=[{"litellm_params": {"model": "openai/gpt-4o-mini"}}])
async def acompletion(**kwargs):
if on_call is not None:
await on_call(kwargs)
if kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN:
return {"choices": [{"message": {"content": judge_json}}]}
if kwargs["model"] == "my-router":
@ -1440,6 +1448,7 @@ class TestActiveJobsCache:
assert [job.id for job in first[("key", "key-hash")]] == ["job-1"]
assert second[("key", "key-hash")][0].attempts == 7
assert second[("key", "key-hash")][0].created_by == job.created_by
assert prisma.db.litellm_shadowevaljob.find_many.await_count == 1
where = prisma.db.litellm_shadowevaljob.find_many.call_args.kwargs["where"]
assert where["stopped_at"] is None
@ -2521,3 +2530,91 @@ class TestSamplingFunnel:
assert logger._test_funnel == []
prisma.db.litellm_shadowevalattempt.create.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("call_type,creator,lookup", [
("acompletion", "initiating-admin", "found"),
("anthropic_messages", None, "found"),
("aresponses", "deleted-admin", "missing"),
("acompletion", None, "missing"),
("acompletion", "unreadable-admin", "error"),
])
async def test_evaluation_uses_current_creator_budgets_and_preserves_source(
monkeypatch: pytest.MonkeyPatch, call_type: str, creator: str | None, lookup: str,
) -> None:
from litellm.caching.caching import DualCache
from litellm.litellm_core_utils import internal_call_metadata as ownership
from litellm.models.user import LiteLLM_UserTable
from litellm.proxy import proxy_server
from litellm.proxy._types import Litellm_EntityType
from litellm.proxy.auth import auth_checks
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.hooks import model_max_budget_limiter as budgets
owner_id: Final = creator or proxy_server.litellm_proxy_admin_name
user_cache: Final = UserApiKeyCache()
monkeypatch.setattr(proxy_server, "user_api_key_cache", user_cache)
monkeypatch.setattr(auth_checks, "last_db_access_time", {})
limiter: Final = budgets._PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache())
prisma: Final = _prisma(jobs=[_job_record(_job(created_by=creator))])
source: Final = {"user_api_key_user_id": "sampled-user", "user_api_key_team_id": "sampled-team"}
if lookup == "missing":
prisma.db.litellm_usertable.find_unique.return_value = None
if lookup == "error":
prisma.db.litellm_usertable.find_unique.side_effect = RuntimeError("creator budget unavailable")
async def account_call(kwargs: Mapping[str, object]) -> None:
owner: Final = ownership.get_evaluation_billing_owner()
assert owner is not None and owner.user_id == owner_id
assert owner.user_model_max_budget == (creator_budget if lookup == "found" else None)
assert owner.max_budget == (1.0 if lookup == "found" else None)
assert owner.spend == (0.1 if lookup == "found" else 0.0)
metadata: Final = kwargs["metadata"]
assert isinstance(metadata, Mapping)
assert all(metadata[key] == value for key, value in source.items())
await limiter.async_log_success_event(
{
ownership.EVALUATION_BILLING_OWNER_KEY: owner,
"standard_logging_object": {
"model_group": kwargs["model"], "response_cost": 0.025, "metadata": metadata,
},
"litellm_params": {"metadata": metadata},
},
response_obj=None, start_time=None, end_time=None,
)
logger: Final = _logger(router=_router(on_call=account_call), prisma=prisma)
for sample, period in enumerate(("1h", "2h") if lookup == "found" else ("1h",), start=1):
creator_budget: Final = {
model: {"budget_limit": 0.02, "time_period": period} for model in ("my-router", "judge-model")
}
if lookup == "found":
row: Final = LiteLLM_UserTable(user_id=owner_id, model_max_budget=creator_budget, max_budget=1.0, spend=0.1)
if sample == 1:
prisma.db.litellm_usertable.find_unique.return_value = row
else:
await user_cache.async_set_cache(key=owner_id, value=row, model_type=LiteLLM_UserTable)
event: Final = _success_kwargs(request_id=f"sample-{sample}", call_type=call_type, request_metadata=source)
response: Final = RESPONSES_API_RESPONSE if call_type == "aresponses" else RESPONSE
await logger.async_log_success_event(event, response, None, None)
await _drain(logger)
usage: Final = await budgets.build_model_max_budget_usage(
Litellm_EntityType.USER, owner_id, creator_budget, limiter.dual_cache,
)
assert usage == {
model: {"current_spend": 0.025 if lookup == "found" else 0.0, "budget_limit": 0.02, "time_period": period}
for model in creator_budget
}
admitted: Final = lookup == "found" or (creator is None and lookup == "missing")
assert prisma.db.litellm_shadowevalattempt.create.await_count == (sample if admitted else 0)
assert logger._test_funnel == ([] if admitted else [("job-1", "withheld")])
if admitted:
assert prisma.db.litellm_shadowevalattempt.create.call_args.kwargs["data"]["outcome"] in ("real", "shadow", "tie")
assert event["litellm_params"]["metadata"] == source
assert ownership.get_evaluation_billing_owner() is None
prisma.db.litellm_usertable.find_unique.assert_awaited_once_with(
where={"user_id": owner_id}, include={"organization_memberships": True},
)
prisma.db.litellm_usertable.create.assert_not_called()
prisma.db.litellm_usertable.upsert.assert_not_called()

View file

@ -1,6 +1,14 @@
"""Unit tests for internal-call metadata forwarding: budget-reservation stripping and origin stamping."""
from collections.abc import Mapping
from copy import deepcopy
from typing import Final
import pytest
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
from litellm.litellm_core_utils import internal_call_metadata as billing
from litellm.litellm_core_utils.internal_call_metadata import (
forwarded_internal_call_metadata,
sanitized_forwardable_call_metadata,
@ -119,3 +127,83 @@ class TestSubCallMetadataSanitization:
assert sanitized_auth.team_id == "team-1"
assert sanitized_auth.api_key == auth.api_key
assert auth.budget_reservation == {"reserved_cost": 1.0}
@pytest.mark.parametrize(
"alternate",
[
{},
{"litellm_metadata": None},
{"litellm_metadata": {}},
{"litellm_metadata": {"model_group": "responses", "internal_call_origin": "shadow_eval_judge"}},
],
)
def test_evaluation_receipt_preserves_metadata_selection(alternate: Mapping[str, object]) -> None:
agent_identity: Final = {"agent_id": "sampled-agent", "billing_agent_id": "billed-agent"}
metadata: Final = {
**agent_identity,
**PARENT,
"user_api_key_user_id": "sampled",
"model_group": "chat",
"internal_call_origin": "shadow_eval_router",
}
kwargs: Final = {
billing.EVALUATION_BILLING_OWNER_KEY: billing.EvaluationBillingOwner("admin"),
**agent_identity,
"standard_logging_object": {**agent_identity, "metadata": metadata, "response_cost": 0.25},
"metadata": metadata,
**alternate,
"litellm_params": {
**agent_identity,
"metadata": metadata,
**alternate,
"proxy_server_request": {"body": {"user": "customer"}},
},
}
snapshot: Final = deepcopy(kwargs)
receipt: Final = billing.project_evaluation_billing_kwargs(kwargs)
resolved: Final = get_litellm_metadata_from_kwargs(receipt)
expected: Final = alternate.get("litellm_metadata") or metadata
assert isinstance(expected, dict)
assert resolved["user_api_key_user_id"] == receipt["user"] == "admin"
assert all(resolved[key] == expected[key] for key in ("model_group", "internal_call_origin"))
params: Final = receipt["litellm_params"]
assert isinstance(params, dict)
for bucket in (receipt, params):
if not alternate.get("litellm_metadata"):
assert bucket.get("litellm_metadata") == alternate.get("litellm_metadata")
assert ("litellm_metadata" in bucket) == ("litellm_metadata" in alternate)
payload: Final = receipt["standard_logging_object"]
assert isinstance(payload, dict)
for container in (receipt, params, payload, resolved, payload["metadata"]):
assert container.get("agent_id") is None
assert container.get("billing_agent_id") is None
assert payload["response_cost"] == 0.25
assert params["proxy_server_request"]["body"]["user"] is None
assert kwargs == snapshot
@pytest.mark.parametrize("marker", [None, "forged-admin", {"user_id": "forged-admin"}])
def test_evaluation_receipt_requires_a_captured_typed_owner(marker: object) -> None:
kwargs: Final = {billing.EVALUATION_BILLING_OWNER_KEY: marker, "user": "sampled-user"}
with billing.evaluation_billing_context(billing.EvaluationBillingOwner("ambient-admin")):
assert billing.project_evaluation_billing_kwargs(kwargs) is kwargs
@pytest.mark.parametrize("trusted", (False, True))
def test_evaluation_receipt_keeps_only_its_own_budget_reservation(trusted: bool) -> None:
from litellm.litellm_core_utils.core_helpers import budget_reservation_from_metadata
from litellm.proxy.spend_tracking.evaluation_budget import EvaluationBudgetReservation
reservation: Final = EvaluationBudgetReservation(total={"reserved_cost": 0.25}, model=None)
kwargs: Final = {
billing.EVALUATION_BILLING_OWNER_KEY: billing.EvaluationBillingOwner("admin"),
billing.EVALUATION_BUDGET_RESERVATION_KEY: reservation if trusted else {"total": reservation.total},
"litellm_params": {"metadata": PARENT, "litellm_metadata": {"internal_call_origin": "shadow_eval_judge"}},
}
receipt: Final = billing.project_evaluation_billing_kwargs(kwargs)
assert budget_reservation_from_metadata(get_litellm_metadata_from_kwargs(receipt)) is (
reservation.total if trusted else None
)
assert PARENT["user_api_key_budget_reservation"] == {"amount": 1.0}

View file

@ -4971,6 +4971,50 @@ def test_get_standard_logging_object_payload_carries_matched_access_groups(loggi
assert payload["request_model_access_groups"] == ("premium-pool", "shared-pool")
@pytest.mark.asyncio
@pytest.mark.parametrize("status", ["success", "failure"])
async def test_evaluation_logging_keeps_captured_owner_after_context_and_thread_handoff(
status: Literal["success", "failure"],
) -> None:
from litellm.litellm_core_utils import internal_call_metadata as billing
from litellm.litellm_core_utils.litellm_logging import get_standard_logging_object_payload
from litellm.types.utils import StandardLoggingPayload
now: Final = datetime.datetime.now()
owner: Final = billing.EvaluationBillingOwner("evaluation-admin")
with billing.evaluation_billing_context(owner):
logger: Final = LitellmLogging(
model="evaluation-model",
messages=[],
stream=False,
call_type="acompletion",
start_time=now,
litellm_call_id="evaluation-call",
function_id="evaluation-function",
)
def payload_on_logging_thread() -> StandardLoggingPayload | None:
with billing.evaluation_billing_context(billing.EvaluationBillingOwner("unrelated-admin")):
logger.update_environment_variables(
litellm_params={"metadata": {"user_api_key_user_id": "sampled-user"}},
optional_params={},
**{billing.EVALUATION_BILLING_OWNER_KEY: {"user_id": "forged-owner"}},
)
return get_standard_logging_object_payload(
kwargs=logger.model_call_details,
init_response_obj={},
start_time=now,
end_time=now,
logging_obj=logger,
status=status,
)
payload: Final = await asyncio.get_running_loop().run_in_executor(None, payload_on_logging_thread)
assert payload is not None
assert payload["metadata"]["user_api_key_user_id"] == "sampled-user"
assert logger.model_call_details[billing.EVALUATION_BILLING_OWNER_KEY] == owner
def test_get_standard_logging_object_payload_has_no_access_groups_when_unstamped(
logging_obj,
):

View file

@ -12,6 +12,7 @@ from functools import partial
from pathlib import Path
from textwrap import dedent
from types import SimpleNamespace
from typing import Final
from unittest.mock import ANY, AsyncMock, MagicMock, patch
@ -21,6 +22,7 @@ from fastapi import HTTPException, status
import litellm
import litellm.proxy.proxy_server
from litellm.caching.dual_cache import DualCache
from litellm.litellm_core_utils import internal_call_metadata as billing
from litellm.proxy._types import (
LiteLLMRoutes,
LiteLLM_JWTAuth,
@ -9651,20 +9653,24 @@ async def test_router_settings_model_group_alias_authorizes_target_for_team(monk
@pytest.mark.asyncio
async def test_reserve_budget_after_common_checks_hands_the_reservation_to_the_request_state():
@pytest.mark.parametrize("evaluation", (False, True))
async def test_reserve_budget_after_common_checks_hands_the_reservation_to_the_request_state(evaluation: bool) -> None:
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}
reservation: Final = None if evaluation else {
"reserved_cost": 0.5, "entries": [], "finalized": False, "callback_bound": False,
}
owner: Final = billing.EvaluationBillingOwner("evaluation-admin") if evaluation else None
with patch(
with billing.evaluation_billing_context(owner), patch(
"litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request",
new=AsyncMock(return_value=reservation),
):
) as reserve:
await _reserve_budget_after_common_checks(
user_api_key_auth_obj=user_api_key_auth_obj,
request_data={"model": "gpt-4o"},
request_data={"model": "gpt-4o", billing.EVALUATION_BILLING_OWNER_KEY: {"user_id": "forged-admin"}},
route="/v1/batches/batch_123/cancel",
llm_router=None,
team_object=None,
@ -9677,6 +9683,7 @@ async def test_reserve_budget_after_common_checks_hands_the_reservation_to_the_r
request=request,
)
assert reserve.await_count == (0 if evaluation else 1)
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

View file

@ -7,6 +7,7 @@ from typing import Final
import pytest
from litellm.caching.caching import DualCache
from litellm.litellm_core_utils import internal_call_metadata as billing
from litellm.proxy.hooks.model_max_budget_limiter import (
_PROXY_VirtualKeyModelMaxBudgetLimiter,
)
@ -20,6 +21,7 @@ BATCH_COST: Final = 2.925e-05
CHAT_COST: Final = 0.001
KEY_SPEND_KEY: Final = f"virtual_key_spend:{KEY_HASH}:{MODEL_GROUP}:1d"
USER_SPEND_KEY: Final = f"user_model_spend:{USER_ID}:{MODEL_GROUP}:1d"
TEAM_SPEND_KEY: Final = f"team_model_spend:sampled-team:{MODEL_GROUP}:1d"
def _batch(batch_id: str, status: str) -> LiteLLMBatch:
@ -35,20 +37,25 @@ def _batch(batch_id: str, status: str) -> LiteLLMBatch:
)
def _event(call_type: str, response_cost: float) -> dict[str, object]:
def _event(call_type: str, response_cost: float, team_budget: bool = False) -> dict[str, object]:
budget: Final = {MODEL_GROUP: {"budget_limit": 0.0001, "time_period": "1d"}}
return {
billing.EVALUATION_BILLING_OWNER_KEY: billing.get_evaluation_billing_owner(),
"call_type": call_type,
"standard_logging_object": {
"call_type": call_type,
"response_cost": response_cost,
"model": "openai/gpt-5.4-mini",
"model_group": MODEL_GROUP,
"metadata": {"user_api_key_hash": KEY_HASH, "user_api_key_user_id": USER_ID},
"metadata": {
"user_api_key_hash": KEY_HASH, "user_api_key_user_id": USER_ID, "user_api_key_team_id": "sampled-team",
},
},
"litellm_params": {
"metadata": {
"user_api_key_model_max_budget": {MODEL_GROUP: {"budget_limit": 0.0001, "time_period": "1d"}},
"user_api_key_user_model_max_budget": {MODEL_GROUP: {"budget_limit": 0.0001, "time_period": "1d"}},
"user_api_key_model_max_budget": None if team_budget else budget,
"user_api_key_user_model_max_budget": budget,
"user_api_key_team_model_max_budget": budget,
}
},
}
@ -60,9 +67,9 @@ async def _poll(limiter: _PROXY_VirtualKeyModelMaxBudgetLimiter, batch: LiteLLMB
)
async def _chat(limiter: _PROXY_VirtualKeyModelMaxBudgetLimiter) -> None:
async def _chat(limiter: _PROXY_VirtualKeyModelMaxBudgetLimiter, team_budget: bool = False) -> None:
await limiter.async_log_success_event(
_event("acompletion", CHAT_COST), response_obj=None, start_time=None, end_time=None
_event("acompletion", CHAT_COST, team_budget), response_obj=None, start_time=None, end_time=None
)
@ -157,16 +164,23 @@ async def test_polls_of_a_finished_batch_charge_each_per_model_budget_once():
@pytest.mark.asyncio
async def test_a_second_batch_and_chat_requests_still_charge_the_budget():
@pytest.mark.parametrize("evaluation", (False, True))
@pytest.mark.parametrize("team_budget", (False, True))
async def test_a_second_batch_and_chat_requests_still_charge_the_budget(evaluation: bool, team_budget: bool) -> None:
limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache())
await _poll(limiter, _batch("batch_first", "completed"), response_cost=BATCH_COST)
await _poll(limiter, _batch("batch_first", "completed"), response_cost=BATCH_COST)
await _poll(limiter, _batch("batch_second", "completed"), response_cost=BATCH_COST)
await _chat(limiter)
await _chat(limiter)
owner: Final = billing.EvaluationBillingOwner("evaluation-admin") if evaluation else None
with billing.evaluation_billing_context(owner):
await _chat(limiter, team_budget)
await _chat(limiter, team_budget)
assert await _spend(limiter, KEY_SPEND_KEY) == pytest.approx(2 * BATCH_COST + 2 * CHAT_COST)
chat_cost: Final = 0 if evaluation else 2 * CHAT_COST
assert await _spend(limiter, KEY_SPEND_KEY) == pytest.approx(2 * BATCH_COST + (0 if team_budget else chat_cost))
assert await _spend(limiter, USER_SPEND_KEY) == pytest.approx(2 * BATCH_COST + chat_cost)
assert await _spend(limiter, TEAM_SPEND_KEY) == pytest.approx(chat_cost if team_budget else 0)
@pytest.mark.asyncio
@ -196,3 +210,17 @@ async def test_a_batch_polled_within_every_budget_window_is_never_charged_again(
await _poll(limiter, finished, BATCH_COST)
assert _local_spend(limiter, KEY_SPEND_KEY) == pytest.approx(BATCH_COST)
@pytest.mark.asyncio
@pytest.mark.parametrize("end_user", (True, 7, ["forged"], {"user": "forged"}, None))
async def test_invalid_optional_end_user_does_not_drop_authenticated_model_spend(end_user: object) -> None:
limiter: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache())
event: Final = _event("acompletion", CHAT_COST, team_budget=True)
payload: Final = event["standard_logging_object"]
assert isinstance(payload, dict)
await limiter.async_log_success_event(
{**event, "standard_logging_object": {**payload, "end_user": end_user}}, None, None, None,
)
assert await _spend(limiter, USER_SPEND_KEY) == pytest.approx(CHAT_COST)
assert await _spend(limiter, TEAM_SPEND_KEY) == pytest.approx(CHAT_COST)

View file

@ -1,6 +1,7 @@
import asyncio
import json
import logging
from collections.abc import Mapping
from datetime import datetime, timedelta, timezone
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
@ -9,6 +10,8 @@ import httpx
import pytest
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils import internal_call_metadata as billing
from litellm.litellm_core_utils.core_helpers import budget_reservation_from_metadata, get_litellm_metadata_from_kwargs
from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP_METADATA_KEY
from litellm.proxy._types import SpendLogsPayload, UserAPIKeyAuth
from litellm.proxy.collector import SpendEventConsumer
@ -24,6 +27,7 @@ from litellm.proxy.hooks.proxy_track_cost_callback import (
)
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
from litellm.proxy.utils import ProxyUpdateSpend
from litellm.proxy.spend_tracking.evaluation_budget import EvaluationBudgetReservation
from litellm.proxy.spend_tracking.spend_event import SpendEventDecodeError, build_spend_event, decode_spend_event
from litellm.proxy.spend_tracking.spend_event_producer import SpendEventProducer, UnixAddress
from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload
@ -2127,7 +2131,8 @@ async def test_failure_hook_drops_error_information_traceback_when_env_set(
@pytest.mark.asyncio
async def test_async_post_call_failure_hook_records_recovered_partial_spend():
@pytest.mark.parametrize("evaluation,captured", ((False, False), (True, False), (True, True)))
async def test_async_post_call_failure_hook_records_recovered_partial_spend(evaluation: bool, captured: bool) -> None:
"""A stream that broke mid-flight still billed the provider. The failure
hook lifts the recovered cost onto request_data as ``response_cost``; this
hook must pass it through to update_database so the failure row records the
@ -2137,20 +2142,25 @@ async def test_async_post_call_failure_hook_records_recovered_partial_spend():
logger = _ProxyDBLogger()
user_api_key_dict = UserAPIKeyAuth(api_key="test_api_key", user_id="u", team_id="t")
owner: Final = billing.EvaluationBillingOwner("admin") if evaluation else None
request_data = {
billing.EVALUATION_BILLING_OWNER_KEY: owner if captured else None,
"model": "anthropic/claude-haiku-4-5",
"messages": [{"role": "user", "content": "Hello"}],
"metadata": {},
"metadata": {"agent_id": "sampled-agent", "billing_agent_id": "billed-agent"},
"proxy_server_request": {"request_id": "rid"},
"response_cost": 3.5e-05,
"combined_usage_object": Usage(prompt_tokens=30, completion_tokens=1, total_tokens=31),
}
with patch(
"litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database",
new_callable=AsyncMock,
) as mock_update_database:
with (
billing.evaluation_billing_context(owner),
patch(
"litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database",
new_callable=AsyncMock,
) as mock_update_database,
):
await logger.async_post_call_failure_hook(
request_data=request_data,
original_exception=Exception("MidStreamFallbackError: read timeout"),
@ -2159,6 +2169,17 @@ async def test_async_post_call_failure_hook_records_recovered_partial_spend():
mock_update_database.assert_called_once()
assert mock_update_database.call_args[1]["response_cost"] == 3.5e-05
written: Final = mock_update_database.call_args.kwargs
assert written["user_id"] == ("admin" if evaluation else "u")
assert written["token"] == (None if evaluation else "test_api_key")
assert written["team_id"] == (None if evaluation else "t")
row: Final = get_logging_payload(
written["kwargs"], written["completion_response"], written["start_time"], written["end_time"]
)
assert row["user"] == written["user_id"]
assert row["agent_id"] == (None if evaluation else "sampled-agent")
assert row["billing_agent_id"] == (None if evaluation else "billed-agent")
assert row["total_tokens"] == 31
@pytest.mark.asyncio
@ -2503,14 +2524,17 @@ async def test_spend_counters_keep_every_granted_group_when_the_deployment_is_un
assert charged == ("premium", "tier0")
def _offload_kwargs() -> dict:
big_prompt = "x" * 10_000
reservation = {"reserved_cost": 0.5, "entries": [{"counter_key": "key:hash-1", "reserved_cost": 0.5}]}
def _offload_kwargs(evaluation: bool = False, call_origin: str | None = None) -> Mapping[str, object]:
big_prompt: Final = "x" * 10_000
reservation: Final = {"reserved_cost": 0.5, "entries": [{"counter_key": "key:hash-1", "reserved_cost": 0.5}]}
return {
billing.EVALUATION_BILLING_OWNER_KEY: billing.EvaluationBillingOwner("admin") if evaluation else None,
"litellm_call_id": "call-1",
"call_type": "acompletion",
"model": "gpt-4o",
"custom_llm_provider": "openai",
"agent_id": "agent-1",
"billing_agent_id": "billed-agent-1",
"stream": False,
"cache_hit": None,
"response_cost": 0.0125,
@ -2523,11 +2547,15 @@ def _offload_kwargs() -> dict:
"proxy_server_request": {"body": {"messages": [{"role": "user", "content": big_prompt}]}},
"metadata": {
"user_api_key": "hash-1",
"internal_call_origin": call_origin or ("shadow_eval_judge" if evaluation else None),
"agent_id": "agent-1",
"billing_agent_id": "billed-agent-1",
"user_api_key_hash": "hash-1",
"user_api_key_alias": "alias-1",
"user_api_key_user_id": "user-1",
"user_api_key_team_id": "team-1",
"user_api_key_org_id": "org-1",
"user_api_key_project_id": "project-1",
"user_api_key_end_user_id": "end-user-1",
"user_api_key_auth": UserAPIKeyAuth(api_key="hash-1", budget_reservation=reservation),
"model_group": "gpt-4o",
@ -2536,6 +2564,8 @@ def _offload_kwargs() -> dict:
},
},
"standard_logging_object": {
"agent_id": "agent-1",
"billing_agent_id": "billed-agent-1",
"id": "chatcmpl-1",
"trace_id": "trace-1",
"response_cost": 0.0125,
@ -2553,6 +2583,8 @@ def _offload_kwargs() -> dict:
"response": {"choices": [{"message": {"content": "y" * 10_000}}]},
"model_parameters": {"temperature": 0.1},
"metadata": {
"agent_id": "agent-1",
"billing_agent_id": "billed-agent-1",
"user_api_key_hash": "hash-1",
"user_api_key_end_user_id": "end-user-1",
"usage_object": {"prompt_tokens": 5000, "completion_tokens": 4000, "total_tokens": 9000},
@ -2704,36 +2736,84 @@ async def _spend_row_written_by(run) -> tuple[SpendLogsPayload, dict, tuple[str,
@pytest.mark.asyncio
async def test_sidecar_writes_the_same_spend_row_and_counters_as_the_in_process_path():
@pytest.mark.parametrize(
"evaluation,call_origin,reserved",
((False, None, False), (True, "shadow_eval_judge", False), (True, "autorouter_compaction", True)),
)
async def test_sidecar_writes_the_same_spend_row_and_counters_as_the_in_process_path(
evaluation: bool, call_origin: str | None, reserved: bool
) -> None:
start_time = datetime(2026, 1, 1, 0, 0, 0)
end_time = datetime(2026, 1, 1, 0, 0, 2)
creator_reservation: Final = (
EvaluationBudgetReservation(
total={"reserved_cost": 0.5, "entries": [{"counter_key": "spend:user:admin", "reserved_cost": 0.5}]},
model=None,
)
if reserved
else None
)
kwargs: Final = {
**_offload_kwargs(evaluation, call_origin),
billing.EVALUATION_BUDGET_RESERVATION_KEY: creator_reservation,
}
expected_reservation: Final = (
creator_reservation.total
if creator_reservation is not None
else None
if evaluation
else budget_reservation_from_metadata(get_litellm_metadata_from_kwargs(kwargs))
)
async def in_process() -> None:
await _ProxyDBLogger().async_log_success_event(_offload_kwargs(), _offload_response(), start_time, end_time)
await _ProxyDBLogger().async_log_success_event(kwargs, _offload_response(), start_time, end_time)
async def via_sidecar() -> None:
line = build_spend_event(_offload_kwargs(), _offload_response(), start_time, end_time, store_bodies=False)
async def replay(line: bytes) -> None:
assert isinstance(line, bytes)
await run_spend_event(line)
async def via_sidecar() -> None:
producer: Final = SpendEventProducer(
address=UnixAddress(path="/nonexistent/spend.sock"),
on_unavailable="fallback",
buffer_size=10,
connect_timeout=0.1,
fallback=replay,
)
await _ProxyDBLogger(producer).async_log_success_event(kwargs, _offload_response(), start_time, end_time)
await producer.close(drain_timeout=5.0)
in_process_row, in_process_counters, in_process_tools = await _spend_row_written_by(in_process)
sidecar_row, sidecar_counters, sidecar_tools = await _spend_row_written_by(via_sidecar)
assert sidecar_row == in_process_row
assert in_process_row["spend"] == 0.0125
assert in_process_row["team_id"] == "team-1"
assert in_process_row["end_user"] == "end-user-1"
assert in_process_row["team_id"] == ("" if evaluation else "team-1")
assert in_process_row["end_user"] == ("" if evaluation else "end-user-1")
assert bool(in_process_row["api_key"]) == (not evaluation)
assert in_process_row["organization_id"] == ("" if evaluation else "org-1")
assert in_process_row["agent_id"] == (None if evaluation else "agent-1")
assert in_process_row["billing_agent_id"] == (None if evaluation else "billed-agent-1")
origin: Final = json.loads(in_process_row["metadata"])["internal_call_origin"]
assert origin == call_origin
assert in_process_row["total_tokens"] == 9000
assert (in_process_row["prompt_tokens"], in_process_row["completion_tokens"]) == (5000, 4000)
assert in_process_row["model_id"] == "deployment-1"
assert in_process_row["request_tags"] == '["tag-a"]'
assert in_process_row["request_tags"] == ("[]" if evaluation else '["tag-a"]')
assert in_process_row["messages"] == "{}"
assert in_process_row["response"] == "{}"
assert sidecar_counters == in_process_counters
assert in_process_counters["token"] == "hash-1"
assert in_process_row["user"] == in_process_counters["user_id"] == ("admin" if evaluation else "user-1")
assert in_process_counters["token"] == (None if evaluation else "hash-1")
assert in_process_counters["team_id"] == (None if evaluation else "team-1")
assert in_process_counters["org_id"] == (None if evaluation else "org-1")
assert in_process_counters["project_id"] == (None if evaluation else "project-1")
assert in_process_counters["end_user_id"] == (None if evaluation else "end-user-1")
assert in_process_counters["response_cost"] == 0.0125
assert in_process_counters["budget_reservation"]["reserved_cost"] == 0.5
assert in_process_counters["model_access_groups"] == ("premium",)
assert in_process_counters["budget_reservation"] is expected_reservation
assert in_process_counters["model_access_groups"] == (() if evaluation else ("premium",))
assert sidecar_tools == in_process_tools == ("get_weather",)
assert get_litellm_metadata_from_kwargs(dict(kwargs))["user_api_key_user_id"] == "user-1"
@pytest.mark.asyncio

View file

@ -132,11 +132,11 @@ async def test_is_key_within_model_budget(budget_limiter):
)
# Test when model is within budget
with patch.object(budget_limiter, "_get_spend_for_model_budget", return_value=50.0):
with patch.object(budget_limiter, "get_spend_for_model_budget", return_value=50.0):
assert await budget_limiter.is_key_within_model_budget(user_api_key, "gpt-4") is True
# Test when model exceeds budget
with patch.object(budget_limiter, "_get_spend_for_model_budget", return_value=150.0):
with patch.object(budget_limiter, "get_spend_for_model_budget", return_value=150.0):
with pytest.raises(litellm.BudgetExceededError):
await budget_limiter.is_key_within_model_budget(user_api_key, "gpt-4")
@ -144,9 +144,9 @@ async def test_is_key_within_model_budget(budget_limiter):
assert await budget_limiter.is_key_within_model_budget(user_api_key, "non-existent") is True
# Test _get_spend_for_model_budget
# Test get_spend_for_model_budget
@pytest.mark.asyncio
async def test_get_spend_for_model_budget_reads_the_configured_model_key(
async def testget_spend_for_model_budget_reads_the_configured_model_key(
budget_limiter,
):
from litellm.proxy.hooks.model_max_budget_limiter import (
@ -162,7 +162,7 @@ async def test_get_spend_for_model_budget_reads_the_configured_model_key(
return 50.0 if key == f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:test-key:gpt-4:1d" else None
with patch.object(budget_limiter.dual_cache, "async_get_cache", side_effect=_spend) as mock_get:
spend = await budget_limiter._get_spend_for_model_budget(
spend = await budget_limiter.get_spend_for_model_budget(
entity_type=Litellm_EntityType.KEY,
entity_id="test-key",
model="openai/gpt-4",
@ -218,7 +218,7 @@ async def test_async_log_success_event_uses_per_model_budget_duration(budget_lim
@pytest.mark.asyncio
async def test_is_end_user_within_model_budget(budget_limiter):
# Test when model is within budget
with patch.object(budget_limiter, "_get_spend_for_model_budget", return_value=50.0):
with patch.object(budget_limiter, "get_spend_for_model_budget", return_value=50.0):
assert (
await budget_limiter.is_end_user_within_model_budget(
"test-user",
@ -229,7 +229,7 @@ async def test_is_end_user_within_model_budget(budget_limiter):
)
# Test when model exceeds budget
with patch.object(budget_limiter, "_get_spend_for_model_budget", return_value=150.0):
with patch.object(budget_limiter, "get_spend_for_model_budget", return_value=150.0):
with pytest.raises(litellm.BudgetExceededError):
await budget_limiter.is_end_user_within_model_budget(
"test-user",
@ -248,7 +248,7 @@ async def test_is_end_user_within_model_budget(budget_limiter):
)
# Test _get_spend_for_model_budget for the end-user scope
# Test get_spend_for_model_budget for the end-user scope
@pytest.mark.asyncio
async def test_get_spend_for_end_user_model_budget(budget_limiter):
from litellm.proxy.hooks.model_max_budget_limiter import (
@ -262,7 +262,7 @@ async def test_get_spend_for_end_user_model_budget(budget_limiter):
return 50.0 if key == f"{END_USER_SPEND_CACHE_KEY_PREFIX}:test-user:gpt-4:1d" else None
with patch.object(budget_limiter.dual_cache, "async_get_cache", side_effect=_spend) as mock_get:
spend = await budget_limiter._get_spend_for_model_budget(
spend = await budget_limiter.get_spend_for_model_budget(
entity_type=Litellm_EntityType.END_USER,
entity_id="test-user",
model="openai/gpt-4",
@ -522,7 +522,7 @@ async def test_get_fallback_model_within_budget_returns_first_within_budget(
model_max_budget={"gpt-4o-mini": {"budget_limit": 100.0, "time_period": "1d"}},
budget_fallbacks={"gpt-4": ["gpt-4o-mini", "claude-haiku"]},
)
with patch.object(budget_limiter, "_get_spend_for_model_budget", return_value=1.0):
with patch.object(budget_limiter, "get_spend_for_model_budget", return_value=1.0):
result = await budget_limiter.get_fallback_model_within_budget(user_api_key, "gpt-4")
assert result == "gpt-4o-mini"
@ -545,7 +545,7 @@ async def test_get_fallback_model_within_budget_skips_exhausted_fallback(
with patch.object(
budget_limiter,
"_get_spend_for_model_budget",
"get_spend_for_model_budget",
side_effect=_spend_for_model,
):
result = await budget_limiter.get_fallback_model_within_budget(user_api_key, "gpt-4")
@ -564,7 +564,7 @@ async def test_get_fallback_model_within_budget_returns_none_when_chain_exhauste
},
budget_fallbacks={"gpt-4": ["gpt-4o-mini", "claude-haiku"]},
)
with patch.object(budget_limiter, "_get_spend_for_model_budget", return_value=150.0):
with patch.object(budget_limiter, "get_spend_for_model_budget", return_value=150.0):
result = await budget_limiter.get_fallback_model_within_budget(user_api_key, "gpt-4")
assert result is None

View file

@ -2040,7 +2040,10 @@ async def test_start_shadow_eval_reverse_records_its_arms_and_holds_its_own_slot
@pytest.mark.asyncio
async def test_start_shadow_eval_forward_leaves_the_baseline_column_empty(monkeypatch: pytest.MonkeyPatch):
@pytest.mark.parametrize("creator", ("admin", None))
async def test_start_shadow_eval_forward_leaves_the_baseline_column_empty(
monkeypatch: pytest.MonkeyPatch, creator: str | None,
) -> None:
import litellm.proxy.proxy_server as proxy_server
_configure_anthropic_sdk_judge(monkeypatch)
@ -2048,9 +2051,11 @@ async def test_start_shadow_eval_forward_leaves_the_baseline_column_empty(monkey
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
monkeypatch.setattr(proxy_server, "llm_router", _shadow_router())
await start_shadow_eval(_start_request(), ADMIN)
monkeypatch.setattr(proxy_server, "litellm_proxy_admin_name", "configured-admin")
await start_shadow_eval(_start_request(), ADMIN.model_copy(update={"user_id": creator}))
rows = prisma.db.litellm_shadowevaljob.create_many.call_args.kwargs["data"]
assert all(row["created_by"] == (creator or "configured-admin") for row in rows)
assert rows[0]["direction"] == "forward"
assert rows[0]["baseline_model"] is None

View file

@ -0,0 +1,675 @@
from __future__ import annotations
import asyncio
import time
from collections.abc import Mapping
from datetime import datetime
from types import SimpleNamespace
from typing import Final
import httpx
import pytest
from openai import AsyncOpenAI
from pydantic import TypeAdapter
from redis.crc import key_slot
import litellm
from litellm.caching.caching import DualCache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.litellm_core_utils.internal_call_metadata import (
EVALUATION_BILLING_OWNER_KEY,
EVALUATION_BUDGET_RESERVATION_KEY,
EvaluationBillingOwner,
evaluation_billing_context,
)
from litellm.proxy import proxy_server
from litellm.proxy._types import LiteLLM_BudgetTable, Litellm_EntityType, LiteLLM_TagTable, LiteLLM_UserTable
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, tag_cache_key
from litellm.proxy.hooks.model_max_budget_limiter import (
_PROXY_VirtualKeyModelMaxBudgetLimiter,
model_budget_spend_cache_key,
model_budget_start_time_cache_key,
resolve_model_budget,
)
from litellm.proxy.spend_tracking.budget_reservation import estimate_request_input_cost, estimate_request_max_cost
from litellm.proxy.spend_tracking.evaluation_budget import (
EvaluationBudgetReservation,
EvaluationModelReservation,
_model_hold_key,
model_budget_spend,
release_evaluation_budget,
reserve_evaluation_budget,
)
from litellm.types.router import RetryPolicy
MODEL: Final = "openai/evaluation-budget-fixture"
REQUEST: Final = {"model": MODEL, "messages": [{"role": "user", "content": "hello"}], "max_tokens": 10}
@pytest.mark.parametrize("namespace", ("", "gateway:", "{gateway}:", "gateway{}:", "gateway{:", "{}later{tag}:"))
@pytest.mark.parametrize("model", ("model", "model{}", "model{tag}", "model{", "模型"))
def test_model_holds_share_the_unchanged_numeric_counters_redis_slot(namespace: str, model: str) -> None:
effective: Final = f"{namespace}user_model_spend:creator:{model}:1d"
hold_key: Final = _model_hold_key(effective)
assert hold_key != effective
assert hold_key.startswith(namespace)
assert key_slot(hold_key.encode()) == key_slot(effective.encode())
@pytest.fixture
def cache(monkeypatch: pytest.MonkeyPatch) -> DualCache:
cache: Final = DualCache()
monkeypatch.setattr(proxy_server, "spend_counter_cache", cache)
monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache())
monkeypatch.setattr(proxy_server, "prisma_client", None)
monkeypatch.setattr(proxy_server, "general_settings", {})
monkeypatch.setattr(proxy_server, "llm_router", None)
monkeypatch.setattr(proxy_server, "model_max_budget_limiter", _PROXY_VirtualKeyModelMaxBudgetLimiter(cache))
monkeypatch.setitem(
litellm.model_cost,
MODEL,
{
"input_cost_per_token": 0.001,
"output_cost_per_token": 0.002,
"max_input_tokens": 1000,
"max_output_tokens": 1000,
"litellm_provider": "openai",
"mode": "chat",
},
)
return cache
async def _owner(scope: str, limit: float) -> EvaluationBillingOwner:
owner: Final = EvaluationBillingOwner(
"evaluation-admin",
{MODEL: {"max_budget": limit, "budget_duration": "1d"}} if scope in ("model", "both") else None,
max_budget=limit if scope in ("total", "both") else None,
)
await proxy_server.user_api_key_cache.async_set_cache(
key=owner.user_id, value=LiteLLM_UserTable(user_id=owner.user_id, max_budget=owner.max_budget, spend=0.0)
)
return owner
@pytest.mark.asyncio
@pytest.mark.parametrize("scope", ("total", "model", "both"))
async def test_concurrent_evaluations_reserve_the_creators_remaining_budget(cache: DualCache, scope: str) -> None:
estimate: Final = estimate_request_max_cost(REQUEST, "/chat/completions", None)
assert estimate is not None and estimate > 0
owner: Final = await _owner(scope, estimate * 1.5)
attempts: Final = await asyncio.gather(
reserve_evaluation_budget(owner, REQUEST, "acompletion"),
reserve_evaluation_budget(owner, REQUEST, "acompletion"),
return_exceptions=True,
)
admitted: Final = tuple(item for item in attempts if isinstance(item, EvaluationBudgetReservation))
assert len(admitted) == 1
assert sum(isinstance(item, litellm.BudgetExceededError) for item in attempts) == 1
reservation: Final = admitted[0]
assert (await cache.async_get_cache("spend:user:evaluation-admin") or 0.0) == pytest.approx(
estimate if scope in ("total", "both") else 0.0
)
await release_evaluation_budget(reservation)
retried: Final = await reserve_evaluation_budget(owner, REQUEST, "acompletion")
assert retried is not None
await release_evaluation_budget(retried)
if reservation.model is not None:
assert (await cache.async_get_cache(reservation.model.spend_key) or 0.0) == pytest.approx(0.0)
@pytest.mark.asyncio
@pytest.mark.parametrize("scope", ("total", "model"))
async def test_zero_creator_budget_blocks_evaluation(cache: DualCache, scope: str) -> None:
owner: Final = await _owner(scope, 0.0)
with pytest.raises(litellm.BudgetExceededError):
await reserve_evaluation_budget(owner, REQUEST, "acompletion")
@pytest.mark.asyncio
async def test_rejected_model_reservation_cannot_block_a_smaller_request(cache: DualCache) -> None:
assert await model_budget_spend(cache, "budget", operation="reserve", member="large:2", limit=1) == 2
assert await model_budget_spend(cache, "budget", operation="reserve", member="small:0.5", limit=1) == 0.5
assert await model_budget_spend(cache, "budget") == 0.5
@pytest.mark.asyncio
async def test_evaluation_reserves_and_charges_only_creator_without_sampled_tag_budget(
cache: DualCache, monkeypatch: pytest.MonkeyPatch
) -> None:
owner: Final = await _owner("total", 1.0)
tag: Final = LiteLLM_TagTable(
tag_name="sampled-tag", spend=0.3, litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0)
)
monkeypatch.setattr(proxy_server, "prisma_client", object())
await proxy_server.user_api_key_cache.async_set_cache(
key=tag_cache_key(tag.tag_name), value=tag, model_type=LiteLLM_TagTable
)
await cache.async_set_cache("spend:user:evaluation-admin", 0.0)
await cache.async_set_cache("spend:tag:sampled-tag", tag.spend)
reservation: Final = await reserve_evaluation_budget(owner, {**REQUEST, "tags": [tag.tag_name]}, "acompletion")
assert reservation is not None
estimate: Final = estimate_request_max_cost(REQUEST, "/chat/completions", None)
assert await cache.async_get_cache("spend:user:evaluation-admin") == pytest.approx(estimate)
assert await cache.async_get_cache("spend:tag:sampled-tag") == pytest.approx(tag.spend)
await release_evaluation_budget(reservation, actual_cost=0.01)
assert await cache.async_get_cache("spend:user:evaluation-admin") == pytest.approx(0.01)
assert await cache.async_get_cache("spend:tag:sampled-tag") == pytest.approx(tag.spend)
def _receipt(
owner: EvaluationBillingOwner, reservation: EvaluationBudgetReservation, cost: float
) -> Mapping[str, object]:
return {
EVALUATION_BILLING_OWNER_KEY: owner,
EVALUATION_BUDGET_RESERVATION_KEY: reservation,
"litellm_params": {"metadata": {"user_api_key_user_id": "sampled-user"}},
"standard_logging_object": {
"model": MODEL,
"response_cost": cost,
"metadata": {"user_api_key_user_id": "sampled-user"},
},
}
@pytest.mark.asyncio
@pytest.mark.parametrize("scope", ("model", "both"))
async def test_model_callback_settles_once_and_preserves_recovered_failure_spend(cache: DualCache, scope: str) -> None:
owner: Final = await _owner(scope, 1.0)
reservation: Final = await reserve_evaluation_budget(owner, REQUEST, "acompletion")
assert reservation is not None and reservation.model is not None
limiter: Final = proxy_server.model_max_budget_limiter
await release_evaluation_budget(reservation)
receipt: Final = _receipt(owner, reservation, 0.003)
await limiter.async_log_failure_event(receipt, None, None, None)
await limiter.async_log_success_event(receipt, None, None, None)
assert await cache.async_get_cache(reservation.model.spend_key) == pytest.approx(0.003)
assert (await cache.async_get_cache("spend:user:evaluation-admin") or 0.0) == pytest.approx(
0.003 if scope == "both" else 0.0
)
@pytest.mark.asyncio
async def test_failed_evaluation_keeps_incurred_cost_and_frees_unused_reservation(cache: DualCache) -> None:
owner: Final = await _owner("both", 1.0)
reservation: Final = await reserve_evaluation_budget(owner, REQUEST, "acompletion")
assert reservation is not None and reservation.model is not None
await release_evaluation_budget(reservation, actual_cost=0.005)
await release_evaluation_budget(reservation, actual_cost=0.005)
assert await cache.async_get_cache("spend:user:evaluation-admin") == pytest.approx(0.005)
assert await cache.async_get_cache(reservation.model.spend_key) == pytest.approx(0.005)
@pytest.mark.asyncio
async def test_sdk_blocks_concurrent_evaluation_before_dispatch_and_settles_actual_usage(cache: DualCache) -> None:
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
estimate: Final = estimate_request_max_cost(REQUEST, "/chat/completions", None)
assert estimate is not None
owner: Final = await _owner("both", estimate * 1.5)
litellm.logging_callback_manager.add_litellm_async_success_callback(proxy_server.model_max_budget_limiter)
entered: Final = asyncio.Event()
complete: Final = asyncio.Event()
requests: Final[asyncio.Queue[httpx.Request]] = asyncio.Queue()
receipts: Final[asyncio.Queue[Mapping[str, object]]] = asyncio.Queue()
async def upstream(request: httpx.Request) -> httpx.Response:
requests.put_nowait(request)
entered.set()
await complete.wait()
return httpx.Response(
200,
json={
"id": "evaluation-response",
"object": "chat.completion",
"created": 1,
"model": MODEL,
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12},
},
)
async def capture(kwargs: Mapping[str, object], response: object, start: datetime, end: datetime) -> None:
receipts.put_nowait(kwargs)
async with httpx.AsyncClient(transport=httpx.MockTransport(upstream)) as http_client:
client: Final = AsyncOpenAI(api_key="transport-only", http_client=http_client)
async def completion() -> object:
return await litellm.acompletion(
model=MODEL,
messages=[{"role": "user", "content": "hello"}],
max_tokens=10,
client=client,
api_key="transport-only",
num_retries=0,
success_callback=[capture],
)
with evaluation_billing_context(owner):
pending: Final = asyncio.create_task(completion())
try:
await asyncio.wait_for(entered.wait(), 30)
with pytest.raises(litellm.BudgetExceededError):
await asyncio.wait_for(completion(), 5)
finally:
complete.set()
await pending
receipt: Final = await asyncio.wait_for(receipts.get(), 30)
await GLOBAL_LOGGING_WORKER.flush()
reservation: Final = receipt[EVALUATION_BUDGET_RESERVATION_KEY]
assert isinstance(reservation, EvaluationBudgetReservation) and reservation.model is not None
actual: Final = 10 * 0.001 + 2 * 0.002
assert receipt.get("response_cost") == pytest.approx(actual)
assert reservation.model.settled_cost == pytest.approx(actual)
await proxy_server.increment_spend_counters(
token=None, team_id=None, user_id=owner.user_id, response_cost=actual, budget_reservation=reservation.total
)
assert requests.qsize() == 1
assert await cache.async_get_cache("spend:user:evaluation-admin") == pytest.approx(actual)
assert await cache.async_get_cache(reservation.model.spend_key) == pytest.approx(actual)
@pytest.mark.asyncio
@pytest.mark.parametrize("cancelled", (False, True))
async def test_sdk_evaluation_failure_releases_budget_without_unreserved_retries(
cache: DualCache, cancelled: bool
) -> None:
owner: Final = await _owner("both", 1.0)
entered: Final = asyncio.Event()
blocked: Final = asyncio.Event()
requests: Final[asyncio.Queue[httpx.Request]] = asyncio.Queue()
async def upstream(request: httpx.Request) -> httpx.Response:
requests.put_nowait(request)
entered.set()
if cancelled:
await blocked.wait()
return httpx.Response(500, json={"error": {"message": "upstream failed", "type": "server_error"}})
async with httpx.AsyncClient(transport=httpx.MockTransport(upstream)) as http_client:
client: Final = AsyncOpenAI(api_key="transport-only", http_client=http_client, max_retries=0)
with evaluation_billing_context(owner):
pending: Final = asyncio.create_task(
litellm.acompletion(
model=MODEL,
messages=[{"role": "user", "content": "hello"}],
max_tokens=10,
client=client,
api_key="transport-only",
num_retries=0,
retry_policy=RetryPolicy(InternalServerErrorRetries=2),
)
)
await asyncio.wait_for(entered.wait(), 30)
if cancelled:
pending.cancel()
with pytest.raises(asyncio.CancelledError if cancelled else litellm.InternalServerError):
await asyncio.wait_for(pending, 5)
incurred: Final = estimate_request_input_cost(REQUEST, "/chat/completions", None) if cancelled else 0.0
assert incurred is not None
model_key: Final = model_budget_spend_cache_key(Litellm_EntityType.USER, owner.user_id, MODEL, "1d")
assert requests.qsize() == 1
assert await cache.async_get_cache("spend:user:evaluation-admin") == pytest.approx(incurred)
assert (await cache.async_get_cache(model_key) or 0.0) == pytest.approx(incurred)
@pytest.mark.asyncio
@pytest.mark.parametrize("backend", (MODEL, "openai/private-evaluation-deployment"))
@pytest.mark.parametrize("routing", ("group", "hidden-alias", "deployment-id", "routing-group"))
async def test_evaluation_uses_configured_prices_for_the_selected_model_group(
cache: DualCache, monkeypatch: pytest.MonkeyPatch, backend: str, routing: str
) -> None:
router: Final = litellm.Router(
model_list=[
{
"model_name": "configured-judge",
"litellm_params": {"model": backend, "api_key": "transport-only"},
"model_info": {
"id": "judge-deployment",
"input_cost_per_token": 0.004,
"output_cost_per_token": 0.007,
"max_input_tokens": 1000,
"max_output_tokens": 1000,
},
}
],
model_group_alias={"hidden-judge": {"model": "configured-judge", "hidden": True}},
routing_groups=[
{"group_name": "evaluation-group", "models": ["configured-judge"], "routing_strategy": "simple-shuffle"}
],
)
monkeypatch.setattr(proxy_server, "llm_router", router)
group: Final = {
"group": "configured-judge",
"hidden-alias": "hidden-judge",
"deployment-id": "judge-deployment",
"routing-group": "evaluation-group",
}[routing]
request: Final = {
**REQUEST,
"model": backend,
"metadata": {"model_group": group},
"model_info": {"id": "judge-deployment"} if routing in ("deployment-id", "routing-group") else {},
}
owner: Final = EvaluationBillingOwner(
"configured-price-owner", {group: {"max_budget": 1.0, "budget_duration": "1d"}}, max_budget=1.0
)
estimate: Final = estimate_request_max_cost({**REQUEST, "model": "configured-judge"}, "/chat/completions", router)
assert estimate is not None
reservation: Final = await reserve_evaluation_budget(owner, request, "acompletion")
assert reservation is not None and reservation.model is not None and reservation.total is not None
assert reservation.total["reserved_cost"] == pytest.approx(estimate)
assert reservation.model.reserved_cost == pytest.approx(estimate)
assert reservation.model.spend_key == model_budget_spend_cache_key(
Litellm_EntityType.USER, owner.user_id, group, "1d"
)
await release_evaluation_budget(reservation)
@pytest.mark.asyncio
@pytest.mark.parametrize("configured_prices", (False, True))
async def test_auto_router_prices_the_selected_leaf_and_preserves_its_logical_budget(
cache: DualCache, monkeypatch: pytest.MonkeyPatch, configured_prices: bool
) -> None:
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
prices: Final = {"input_cost_per_token": 0.004, "output_cost_per_token": 0.007} if configured_prices else {}
backend: Final = "openai/private-auto-router-leaf" if configured_prices else MODEL
router: Final = litellm.Router(
model_list=[
{
"model_name": "priced-leaf",
"litellm_params": {"model": backend, "api_key": "transport-only", "max_tokens": 10},
"model_info": {"id": "selected-leaf", "max_input_tokens": 1000, "max_output_tokens": 10, **prices},
},
{
"model_name": "evaluation-router",
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": {
"classifier_type": "heuristic",
"tiers": dict.fromkeys(("SIMPLE", "MEDIUM", "COMPLEX", "REASONING"), "priced-leaf"),
},
},
},
],
num_retries=0,
)
monkeypatch.setattr(proxy_server, "llm_router", router)
owner: Final = EvaluationBillingOwner(
"auto-router-owner", {"evaluation-router": {"max_budget": 1.0, "budget_duration": "1d"}}, max_budget=1.0
)
estimate: Final = estimate_request_max_cost({**REQUEST, "model": "priced-leaf"}, "/chat/completions", router)
assert estimate is not None and estimate > 0
model_key: Final = model_budget_spend_cache_key(Litellm_EntityType.USER, owner.user_id, "evaluation-router", "1d")
receipts: Final[asyncio.Queue[Mapping[str, object]]] = asyncio.Queue()
litellm.logging_callback_manager.add_litellm_async_success_callback(proxy_server.model_max_budget_limiter)
async def upstream(request: httpx.Request) -> httpx.Response:
payload: Final = TypeAdapter(Mapping[str, object]).validate_json(request.content)
assert payload.get("model") == backend.removeprefix("openai/")
assert payload.get("max_tokens") == 10
assert await cache.async_get_cache(f"spend:user:{owner.user_id}") == pytest.approx(estimate)
assert await model_budget_spend(cache, model_key) == pytest.approx(estimate)
return httpx.Response(
200,
json={
"id": "auto-router-evaluation-response",
"object": "chat.completion",
"created": 1,
"model": backend,
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12},
},
)
async def capture(kwargs: Mapping[str, object], response: object, start: datetime, end: datetime) -> None:
receipts.put_nowait(kwargs)
async with httpx.AsyncClient(transport=httpx.MockTransport(upstream)) as http_client:
client: Final = AsyncOpenAI(api_key="transport-only", http_client=http_client)
with evaluation_billing_context(owner):
await router.acompletion(
model="evaluation-router",
messages=[{"role": "user", "content": "hello"}],
max_tokens=10,
client=client,
fallbacks=[],
success_callback=[capture],
)
receipt: Final = await asyncio.wait_for(receipts.get(), 30)
await GLOBAL_LOGGING_WORKER.flush()
reservation: Final = receipt[EVALUATION_BUDGET_RESERVATION_KEY]
assert isinstance(reservation, EvaluationBudgetReservation) and reservation.model is not None
assert reservation.model.spend_key == model_key
assert reservation.model.reserved_cost == pytest.approx(estimate)
actual: Final = 10 * prices.get("input_cost_per_token", 0.001) + 2 * prices.get("output_cost_per_token", 0.002)
assert await cache.async_get_cache(model_key) == pytest.approx(actual)
await proxy_server.increment_spend_counters(
token=None, team_id=None, user_id=owner.user_id, response_cost=actual, budget_reservation=reservation.total
)
assert await cache.async_get_cache(f"spend:user:{owner.user_id}") == pytest.approx(actual)
@pytest.mark.asyncio
@pytest.mark.parametrize("window_change", ("missing-start", "replaced-start", "rollover"))
async def test_model_settlement_never_refunds_another_window_or_request(cache: DualCache, window_change: str) -> None:
owner: Final = await _owner("model", 1.0)
spend_key: Final = model_budget_spend_cache_key(Litellm_EntityType.USER, owner.user_id, MODEL, "1d")
start_key: Final = model_budget_start_time_cache_key(Litellm_EntityType.USER, owner.user_id, MODEL, "1d")
await cache.async_set_cache(spend_key, 0.2)
await cache.async_set_cache(start_key, time.time())
earlier: Final = await reserve_evaluation_budget(owner, REQUEST, "acompletion")
assert earlier is not None and earlier.model is not None
cache.in_memory_cache.delete_cache(start_key)
if window_change == "replaced-start":
await cache.async_set_cache(start_key, time.time() + 1)
if window_change == "rollover":
cache.in_memory_cache.delete_cache(spend_key)
await cache.async_set_cache(spend_key, 0.3)
later: Final = await reserve_evaluation_budget(owner, REQUEST, "acompletion")
assert later is not None and later.model is not None
await release_evaluation_budget(earlier, actual_cost=0.01)
assert await model_budget_spend(cache, spend_key) == pytest.approx(
(0.3 if window_change == "rollover" else 0.2) + 0.01 + later.model.reserved_cost
)
assert await cache.async_get_cache(spend_key) == pytest.approx((0.3 if window_change == "rollover" else 0.2) + 0.01)
await release_evaluation_budget(later, actual_cost=0.02)
assert await model_budget_spend(cache, spend_key) == pytest.approx(
(0.3 if window_change == "rollover" else 0.2) + 0.03
)
assert await cache.async_get_cache(spend_key) == pytest.approx((0.3 if window_change == "rollover" else 0.2) + 0.03)
class _Clock:
seconds: float = 0.0
def now(self) -> float:
return self.seconds
def advance(self, seconds: float) -> None:
self.seconds += seconds
@pytest.mark.asyncio
async def test_expired_model_hold_cannot_release_or_renew_a_new_requests_budget() -> None:
clock: Final = _Clock()
cache: Final = DualCache(in_memory_cache=InMemoryCache(clock=clock.now))
await model_budget_spend(cache, "budget", operation="reserve", member="old:0.2")
clock.advance(30)
assert await model_budget_spend(cache, "budget", operation="reserve", member="new:0.3") == pytest.approx(0.5)
clock.advance(31)
assert await model_budget_spend(cache, "budget") == pytest.approx(0.3)
assert await model_budget_spend(cache, "budget", operation="settle", member="old:0.2") == pytest.approx(0.3)
with pytest.raises(RuntimeError, match="Evaluation budget reservation expired"):
await model_budget_spend(cache, "budget", operation="renew", member="old:0.2")
assert await model_budget_spend(cache, "budget") == pytest.approx(0.3)
await model_budget_spend(cache, "budget", operation="renew", member="new:0.3")
clock.advance(31)
assert await model_budget_spend(cache, "budget") == pytest.approx(0.3)
await model_budget_spend(cache, "budget", operation="settle", member="new:0.3")
assert await model_budget_spend(cache, "budget") == pytest.approx(0.0)
@pytest.mark.asyncio
async def test_ordinary_user_model_gate_includes_outstanding_evaluations(cache: DualCache) -> None:
estimate: Final = estimate_request_max_cost(REQUEST, "/chat/completions", None)
assert estimate is not None
owner: Final = await _owner("model", estimate)
reservation: Final = await reserve_evaluation_budget(owner, REQUEST, "acompletion")
assert reservation is not None
resolved: Final = resolve_model_budget(MODEL, owner.user_model_max_budget or {})
assert resolved is not None
limiter: Final = proxy_server.model_max_budget_limiter
with pytest.raises(litellm.BudgetExceededError):
await limiter.is_user_within_model_budget(owner.user_id, owner.user_model_max_budget or {}, MODEL)
await release_evaluation_budget(reservation)
assert await limiter.is_user_within_model_budget(owner.user_id, owner.user_model_max_budget or {}, MODEL)
@pytest.mark.asyncio
@pytest.mark.parametrize("phase", ("total-acquire", "total-settle"))
async def test_repeated_cancellation_drains_budget_writes_before_returning(cache: DualCache, phase: str) -> None:
owner: Final = await _owner("both", 1.0)
model_key: Final = model_budget_spend_cache_key(Litellm_EntityType.USER, owner.user_id, MODEL, "1d")
total_key: Final = f"spend:user:{owner.user_id}"
reservation: Final = (
await reserve_evaluation_budget(owner, REQUEST, "acompletion") if phase == "total-settle" else None
)
backend: Final = cache.in_memory_cache
entered: Final = asyncio.Event()
proceed: Final = asyncio.Event()
pause_key: Final = total_key
async def write_then_wait(key: str, value: float, **kwargs: object) -> float:
result: Final = await backend.async_increment(key, value, **kwargs)
if key == pause_key and value != 0 and not entered.is_set():
entered.set()
await proceed.wait()
return result
cache.in_memory_cache = SimpleNamespace( # pyright: ignore[reportAttributeAccessIssue] # injected storage boundary delegates every operation to the real cache
get_cache=backend.get_cache,
set_cache=backend.set_cache,
delete_cache=backend.delete_cache,
async_get_cache=backend.async_get_cache,
async_set_cache=backend.async_set_cache,
async_increment=write_then_wait,
increment_cache=backend.increment_cache,
_clock=backend._clock,
)
pending: Final = asyncio.create_task(
reserve_evaluation_budget(owner, REQUEST, "acompletion")
if reservation is None
else release_evaluation_budget(reservation, actual_cost=0.005)
)
await asyncio.wait_for(entered.wait(), 5)
pending.cancel()
await asyncio.sleep(0)
pending.cancel()
assert not pending.done()
proceed.set()
with pytest.raises(asyncio.CancelledError):
await pending
expected: Final = 0.0 if reservation is None else 0.005
assert (await cache.async_get_cache(total_key) or 0.0) == pytest.approx(expected)
assert (await cache.async_get_cache(model_key) or 0.0) == pytest.approx(expected)
assert await model_budget_spend(cache, model_key) == pytest.approx(expected)
await release_evaluation_budget(reservation, actual_cost=expected)
assert (await cache.async_get_cache(model_key) or 0.0) == pytest.approx(expected)
@pytest.mark.asyncio
async def test_evaluation_and_ordinary_requests_share_epoch_budget_window_markers(cache: DualCache) -> None:
clock: Final = _Clock()
cache.in_memory_cache = InMemoryCache(clock=clock.now)
owner: Final = await _owner("model", 1.0)
reservation: Final = await reserve_evaluation_budget(owner, REQUEST, "acompletion")
assert reservation is not None and reservation.model is not None
await release_evaluation_budget(reservation, actual_cost=0.02)
ordinary: Final = {
"litellm_params": {"metadata": {"user_api_key_user_model_max_budget": owner.user_model_max_budget}},
"standard_logging_object": {
"model": MODEL,
"response_cost": 0.01,
"metadata": {"user_api_key_user_id": owner.user_id},
},
}
await proxy_server.model_max_budget_limiter.async_log_success_event(ordinary, None, None, None)
await release_evaluation_budget(reservation, actual_cost=0.03)
await proxy_server.model_max_budget_limiter.async_log_success_event(ordinary, None, None, None)
assert await cache.async_get_cache(reservation.model.spend_key) == pytest.approx(0.05)
@pytest.mark.asyncio
async def test_total_storage_failure_does_not_strand_model_budget(
cache: DualCache, monkeypatch: pytest.MonkeyPatch
) -> None:
owner: Final = await _owner("both", 1.0)
reservation: Final = await reserve_evaluation_budget(owner, REQUEST, "acompletion")
assert reservation is not None and reservation.model is not None and reservation.total is not None
def unavailable(key: str) -> object:
raise OSError("total spend storage unavailable")
failed_cache: Final = DualCache(
in_memory_cache=SimpleNamespace(get_cache=unavailable) # pyright: ignore[reportArgumentType] # injected failing storage boundary
)
monkeypatch.setattr(proxy_server, "spend_counter_cache", failed_cache)
with pytest.raises(OSError, match="total spend storage unavailable"):
await release_evaluation_budget(reservation, actual_cost=0.005)
assert reservation.total["finalized"] is False
assert await model_budget_spend(cache, reservation.model.spend_key) == pytest.approx(0.005)
assert await cache.async_get_cache(reservation.model.spend_key) == pytest.approx(0.005)
monkeypatch.setattr(proxy_server, "spend_counter_cache", cache)
await release_evaluation_budget(reservation, actual_cost=0.0)
assert reservation.total["finalized"] is True
assert await cache.async_get_cache(f"spend:user:{owner.user_id}") == pytest.approx(0.005)
assert await cache.async_get_cache(reservation.model.spend_key) == pytest.approx(0.005)
@pytest.mark.asyncio
@pytest.mark.parametrize("case", ("queued-receipt", "lost-lease"))
async def test_model_lease_survives_queued_logging_and_cancels_an_unreserved_call(
cache: DualCache, monkeypatch: pytest.MonkeyPatch, case: str
) -> None:
clock: Final = _Clock()
cache.in_memory_cache = InMemoryCache(clock=clock.now)
model: Final = EvaluationModelReservation(cache, "model-spend", "model-start", 86400, 0.2)
await model_budget_spend(cache, model.spend_key, operation="reserve", member=model.member)
dispatched: Final = asyncio.Event()
request: Final = asyncio.create_task(dispatched.wait())
if case == "queued-receipt":
dispatched.set()
await request
intervals: Final[asyncio.Queue[asyncio.Event]] = asyncio.Queue()
sleep: Final = asyncio.sleep
async def tick(delay: float) -> None:
if asyncio.current_task() is not renewal:
await sleep(delay)
return
interval: Final = asyncio.Event()
intervals.put_nowait(interval)
await interval.wait()
with monkeypatch.context() as timers:
timers.setattr(asyncio, "sleep", tick)
renewal: Final = asyncio.create_task(model.renew(request))
first: Final = await asyncio.wait_for(intervals.get(), 5)
clock.advance(61 if case == "lost-lease" else 31)
first.set()
if case == "lost-lease":
await asyncio.wait_for(renewal, 5)
with pytest.raises(asyncio.CancelledError):
await request
assert await model_budget_spend(cache, model.spend_key) == pytest.approx(0.0)
else:
second: Final = await asyncio.wait_for(intervals.get(), 5)
clock.advance(31)
assert await model_budget_spend(cache, model.spend_key) == pytest.approx(0.2)
await model.settle(0.0)
second.set()
await asyncio.wait_for(renewal, 5)
assert await model_budget_spend(cache, model.spend_key) == pytest.approx(0.0)

View file

@ -10,6 +10,7 @@ from pydantic import TypeAdapter
import litellm
from litellm.caching.caching import DualCache
from litellm.exceptions import BadRequestError
from litellm.litellm_core_utils import internal_call_metadata as billing
from litellm.litellm_core_utils.initialize_dynamic_callback_params import inherit_message_logging_privacy
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.redact_messages import should_redact_message_logging
@ -83,8 +84,9 @@ async def test_child_preserves_credentials_and_isolates_context(protocol: Litera
@pytest.mark.asyncio
@pytest.mark.parametrize("protocol", ("chat", "messages"))
@pytest.mark.parametrize("policy", ("allowed", "denied", "forged", "router_alias", "unrelated_alias"))
@pytest.mark.parametrize("evaluation_owned", (False, True))
async def test_real_proxy_child_auth_privacy_and_body_policy(
monkeypatch: pytest.MonkeyPatch, protocol: Literal["chat", "messages"], policy: str,
monkeypatch: pytest.MonkeyPatch, protocol: Literal["chat", "messages"], policy: str, evaluation_owned: bool,
) -> None:
cache: Final = DualCache()
token: Final = proxy_server.hash_token("sk-compaction-fixture")
@ -93,6 +95,7 @@ async def test_real_proxy_child_auth_privacy_and_body_policy(
await cache.async_set_cache(key=token, value=auth)
dispatched: Final = asyncio.Event()
allowed: Final = policy in ("allowed", "router_alias")
owner: Final = billing.EvaluationBillingOwner("evaluation-admin") if evaluation_owned else None
async def route(
data: Mapping[str, object], llm_router: Router | None, user_model: str | None,
@ -100,6 +103,8 @@ async def test_real_proxy_child_auth_privacy_and_body_policy(
) -> Awaitable[ModelResponse]:
dispatched.set()
assert allowed
assert billing.get_evaluation_billing_owner() is owner
assert user_api_key_dict is not None and user_api_key_dict.api_key == token
if policy == "router_alias":
with pytest.raises(ProxyException):
await can_key_call_model("unrelated-compactor", None, auth, None)
@ -122,7 +127,7 @@ async def test_real_proxy_child_auth_privacy_and_body_policy(
monkeypatch.setattr(proxy_server, "llm_router", None)
monkeypatch.setattr(proxy_server, "general_settings", {})
monkeypatch.setattr(common_request_processing, "route_request", route)
with inherit_message_logging_privacy(True):
with inherit_message_logging_privacy(True), billing.evaluation_billing_context(owner):
call: Final = with_proxy_compaction_executor(
_child(protocol, policy == "forged", "auto" if policy.endswith("alias") else None), _request(proxy_server.app)
)
@ -133,6 +138,7 @@ async def test_real_proxy_child_auth_privacy_and_body_policy(
with pytest.raises(BadRequestError, match=rf"child request failed \(HTTP {status}\)"):
await call
assert dispatched.is_set() is allowed
assert billing.get_evaluation_billing_owner() is None
assert compaction_executor.get() is None
if policy.endswith("alias"):
with pytest.raises(ProxyException):

View file

@ -29,13 +29,19 @@ from litellm._logging import (
trace_id_var,
verbose_logger,
)
from litellm.caching.caching import Cache
from litellm.caching.caching import Cache, DualCache
from litellm.caching.caching_handler import _PENDING_CACHE_WRITES
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.internal_call_metadata import (
EVALUATION_BUDGET_RESERVATION_KEY,
EvaluationBillingOwner,
evaluation_billing_context,
)
from litellm.litellm_core_utils.thread_pool_executor import executor as logging_executor
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
from litellm.proxy.utils import is_valid_api_key
@ -5524,6 +5530,397 @@ async def test_wrapper_async_leaves_the_budget_reservation_alone_on_internal_cal
assert claimed_by_the_outer_call["callback_bound"] is True
_EVALUATION_BRIDGE_MODEL: Final = "hosted_vllm/evaluation-bridge-fixture"
_EVALUATION_FALLBACK_MODEL: Final = "hosted_vllm/evaluation-fallback-fixture"
_EVALUATION_BRIDGE_REQUEST: Final = {
"model": _EVALUATION_BRIDGE_MODEL,
"messages": [{"role": "user", "content": "hello"}],
"max_tokens": 10,
}
def _evaluation_provider_response(output_tokens: int = 2, model: str = "evaluation-bridge-fixture") -> httpx.Response:
return httpx.Response(
200,
json={
"id": f"evaluation-{output_tokens}",
"object": "chat.completion",
"created": 1,
"model": model,
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 10, "completion_tokens": output_tokens, "total_tokens": 10 + output_tokens},
},
)
def _deferred_evaluation_logger(receipts: asyncio.Queue[Mapping[str, object]]) -> Logging:
async def capture(kwargs: Mapping[str, object], response: object, start: datetime, end: datetime) -> None:
receipts.put_nowait(kwargs)
logger: Final = Logging(
model=_EVALUATION_BRIDGE_MODEL,
messages=_EVALUATION_BRIDGE_REQUEST["messages"],
stream=False,
call_type="acompletion",
start_time=datetime.now(),
litellm_call_id="evaluation",
function_id="test",
dynamic_async_success_callbacks=[capture],
)
logger._defer_async_logging = True
return logger
@pytest.fixture
def evaluation_bridge_budget(monkeypatch: pytest.MonkeyPatch) -> tuple[EvaluationBillingOwner, DualCache]:
from litellm.proxy import proxy_server
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.hooks.model_max_budget_limiter import _PROXY_VirtualKeyModelMaxBudgetLimiter
cache: Final = DualCache()
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setattr(proxy_server, "spend_counter_cache", cache)
monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache())
monkeypatch.setattr(proxy_server, "prisma_client", None)
monkeypatch.setattr(proxy_server, "general_settings", {})
monkeypatch.setattr(proxy_server, "model_max_budget_limiter", _PROXY_VirtualKeyModelMaxBudgetLimiter(cache))
for model in (_EVALUATION_BRIDGE_MODEL, _EVALUATION_FALLBACK_MODEL):
monkeypatch.setitem(
litellm.model_cost,
model,
{
"input_cost_per_token": 0.001 if model == _EVALUATION_BRIDGE_MODEL else 0.003,
"output_cost_per_token": 0.002 if model == _EVALUATION_BRIDGE_MODEL else 0.004,
"max_input_tokens": 1000,
"max_output_tokens": 1000,
"litellm_provider": "hosted_vllm",
"mode": "chat",
},
)
owner: Final = EvaluationBillingOwner(
"evaluation-admin",
{
model: {"max_budget": 1.0, "budget_duration": "1d"}
for model in (_EVALUATION_BRIDGE_MODEL, _EVALUATION_FALLBACK_MODEL)
},
max_budget=1.0,
)
return owner, cache
@pytest.mark.asyncio
@pytest.mark.parametrize("outcome", ("success", "error", "cancelled"))
async def test_nested_messages_evaluation_reserves_once_and_releases_its_own_budget(
evaluation_bridge_budget: tuple[EvaluationBillingOwner, DualCache],
outcome: str,
) -> None:
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.proxy import proxy_server
from litellm.proxy.spend_tracking.budget_reservation import estimate_request_input_cost, estimate_request_max_cost
from litellm.proxy.spend_tracking.evaluation_budget import EvaluationBudgetReservation
owner, cache = evaluation_bridge_budget
estimate: Final = estimate_request_max_cost(_EVALUATION_BRIDGE_REQUEST, "/v1/messages", None)
entered: Final = asyncio.Event()
complete: Final = asyncio.Event()
receipts: Final[asyncio.Queue[Mapping[str, object]]] = asyncio.Queue()
requests: Final[asyncio.Queue[httpx.Request]] = asyncio.Queue()
async def upstream(request: httpx.Request) -> httpx.Response:
requests.put_nowait(request)
entered.set()
await complete.wait()
if outcome == "error":
return httpx.Response(500, json={"error": {"message": "upstream failed", "type": "server_error"}})
return _evaluation_provider_response()
async def capture(kwargs: Mapping[str, object], response: object, start: datetime, end: datetime) -> None:
receipts.put_nowait(kwargs)
with respx.mock(assert_all_called=False) as transport, evaluation_billing_context(owner):
transport.post("https://evaluation.invalid/v1/chat/completions").mock(side_effect=upstream)
pending: Final = asyncio.create_task(
litellm.anthropic.messages.acreate(
**_EVALUATION_BRIDGE_REQUEST,
api_base="https://evaluation.invalid/v1",
api_key="transport-only",
num_retries=0,
max_retries=0,
fallbacks=[],
success_callback=[capture],
)
)
pending.add_done_callback(lambda task: entered.set())
try:
await asyncio.wait_for(entered.wait(), 10)
if pending.done():
await pending
assert await cache.async_get_cache("spend:user:evaluation-admin") == pytest.approx(estimate)
if outcome == "cancelled":
pending.cancel()
else:
complete.set()
if outcome == "success":
response: Final = await asyncio.wait_for(pending, 10)
assert response["content"][0]["text"] == "ok"
else:
with pytest.raises(asyncio.CancelledError if outcome == "cancelled" else litellm.InternalServerError):
await asyncio.wait_for(pending, 10)
finally:
complete.set()
if not pending.done():
pending.cancel()
with contextlib.suppress(asyncio.CancelledError):
await pending
assert requests.qsize() == 1
if outcome == "success":
receipt: Final = await asyncio.wait_for(receipts.get(), 10)
reservation: Final = receipt[EVALUATION_BUDGET_RESERVATION_KEY]
assert isinstance(reservation, EvaluationBudgetReservation)
assert reservation.model is not None
actual: Final = 10 * 0.001 + 2 * 0.002
await proxy_server.model_max_budget_limiter.async_log_success_event(receipt, None, None, None)
await proxy_server.increment_spend_counters(
token=None,
team_id=None,
user_id=owner.user_id,
response_cost=actual,
budget_reservation=reservation.total,
)
assert await cache.async_get_cache(reservation.model.spend_key) == pytest.approx(actual)
assert await cache.async_get_cache("spend:user:evaluation-admin") == pytest.approx(actual)
await GLOBAL_LOGGING_WORKER.flush()
else:
incurred: Final = (
estimate_request_input_cost(_EVALUATION_BRIDGE_REQUEST, "/v1/messages", None)
if outcome == "cancelled"
else 0.0
)
assert await cache.async_get_cache("spend:user:evaluation-admin") == pytest.approx(incurred)
assert (
await cache.async_get_cache(f"user_model_spend:{owner.user_id}:{_EVALUATION_BRIDGE_MODEL}:1d") or 0.0
) == pytest.approx(incurred)
@pytest.mark.asyncio
@pytest.mark.parametrize("api", ("chat", "messages"))
async def test_evaluation_fallback_rechecks_budget_before_the_next_provider_call(
evaluation_bridge_budget: tuple[EvaluationBillingOwner, DualCache],
api: str,
) -> None:
owner, cache = evaluation_bridge_budget
restricted: Final = EvaluationBillingOwner(
owner.user_id,
{
_EVALUATION_BRIDGE_MODEL: {"max_budget": 1.0, "budget_duration": "1d"},
_EVALUATION_FALLBACK_MODEL: {"max_budget": 0.0, "budget_duration": "1d"},
},
max_budget=1.0,
)
with respx.mock(assert_all_called=True) as transport, evaluation_billing_context(restricted):
route: Final = transport.post("https://evaluation.invalid/v1/chat/completions").respond(
500,
json={"error": {"message": "upstream failed", "type": "server_error"}},
)
create: Final = litellm.acompletion if api == "chat" else litellm.anthropic.messages.acreate
with pytest.raises(Exception, match="All fallback attempts failed"):
await create(
**_EVALUATION_BRIDGE_REQUEST,
fallbacks=[_EVALUATION_FALLBACK_MODEL],
api_base="https://evaluation.invalid/v1",
api_key="transport-only",
num_retries=0,
max_retries=0,
)
assert route.call_count == 1
assert await cache.async_get_cache("spend:user:evaluation-admin") == pytest.approx(0.0)
assert (
await cache.async_get_cache(f"user_model_spend:{owner.user_id}:{_EVALUATION_BRIDGE_MODEL}:1d") or 0.0
) == pytest.approx(0.0)
assert (
await cache.async_get_cache(f"user_model_spend:{owner.user_id}:{_EVALUATION_FALLBACK_MODEL}:1d") or 0.0
) == pytest.approx(0.0)
@pytest.mark.asyncio
@pytest.mark.parametrize("api", ("chat", "embedding"))
@pytest.mark.parametrize("scope", ("total", "model"))
async def test_paid_classifier_calls_enforce_creator_budget_before_dispatch(
evaluation_bridge_budget: tuple[EvaluationBillingOwner, DualCache],
api: str,
scope: str,
monkeypatch: pytest.MonkeyPatch,
) -> None:
from litellm.litellm_core_utils.internal_call_metadata import forwarded_internal_call_metadata
from litellm.types.utils import AUTOROUTER_CLASSIFIER_CALL_ORIGIN
owner, cache = evaluation_bridge_budget
restricted: Final = EvaluationBillingOwner(
owner.user_id,
{_EVALUATION_BRIDGE_MODEL: {"max_budget": 0.0, "budget_duration": "1d"}} if scope == "model" else None,
max_budget=0.0 if scope == "total" else None,
)
metadata: Final = forwarded_internal_call_metadata(
{"user_api_key_user_id": "sampled-user"}, AUTOROUTER_CLASSIFIER_CALL_ORIGIN
)
request: Final = (
_EVALUATION_BRIDGE_REQUEST if api == "chat" else {"model": _EVALUATION_BRIDGE_MODEL, "input": ["hello"]}
)
create: Final = litellm.acompletion if api == "chat" else litellm.aembedding
if api == "embedding":
monkeypatch.setattr(litellm, "model_fallbacks", [_EVALUATION_FALLBACK_MODEL])
with respx.mock(assert_all_called=False) as transport, evaluation_billing_context(restricted):
route: Final = transport.post(
"https://evaluation.invalid/v1/chat/completions" if api == "chat" else "https://evaluation.invalid/v1/embeddings"
).respond(500)
with pytest.raises(litellm.BudgetExceededError):
await create(
**request,
api_base="https://evaluation.invalid/v1",
api_key="transport-only",
metadata=metadata,
num_retries=0,
)
assert route.call_count == 0
assert (await cache.async_get_cache("spend:user:evaluation-admin") or 0.0) == pytest.approx(0.0)
@pytest.mark.asyncio
async def test_delayed_nested_evaluation_callbacks_keep_each_calls_reservation_and_cost(
evaluation_bridge_budget: tuple[EvaluationBillingOwner, DualCache],
) -> None:
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.proxy import proxy_server
from litellm.proxy.spend_tracking.evaluation_budget import EvaluationBudgetReservation
owner, cache = evaluation_bridge_budget
requests: Final[asyncio.Queue[httpx.Request]] = asyncio.Queue()
receipts: Final[asyncio.Queue[Mapping[str, object]]] = asyncio.Queue()
def upstream(request: httpx.Request) -> httpx.Response:
requests.put_nowait(request)
return _evaluation_provider_response(
requests.qsize(),
"evaluation-bridge-fixture" if requests.qsize() == 1 else "evaluation-fallback-fixture",
)
with respx.mock(assert_all_called=True) as transport, evaluation_billing_context(owner):
transport.post("https://evaluation.invalid/v1/chat/completions").mock(side_effect=upstream)
logging_obj: Final = _deferred_evaluation_logger(receipts)
callbacks: Final[asyncio.Queue[Callable[[], None]]] = asyncio.Queue()
reservations: Final[asyncio.Queue[EvaluationBudgetReservation]] = asyncio.Queue()
for model in (_EVALUATION_BRIDGE_MODEL, _EVALUATION_FALLBACK_MODEL):
await litellm.anthropic.messages.acreate(
**{**_EVALUATION_BRIDGE_REQUEST, "model": model},
litellm_logging_obj=logging_obj,
api_base="https://evaluation.invalid/v1",
api_key="transport-only",
num_retries=0,
)
enqueue: Final = logging_obj._enqueue_deferred_logging
assert enqueue is not None
reservation: Final = logging_obj.evaluation_budget_reservation
assert reservation is not None
callbacks.put_nowait(enqueue)
reservations.put_nowait(reservation)
logging_obj._enqueue_deferred_logging = None
callbacks.get_nowait()()
callbacks.get_nowait()()
received: Final = (
await asyncio.wait_for(receipts.get(), 10),
await asyncio.wait_for(receipts.get(), 10),
)
await GLOBAL_LOGGING_WORKER.flush()
first, second = reservations.get_nowait(), reservations.get_nowait()
assert first is not second
assert receipts.empty()
assert {id(receipt[EVALUATION_BUDGET_RESERVATION_KEY]): receipt["response_cost"] for receipt in received} == {
id(first): 10 * 0.001 + 0.002,
id(second): 10 * 0.003 + 2 * 0.004,
}
for receipt in received:
handle: Final = receipt[EVALUATION_BUDGET_RESERVATION_KEY]
assert isinstance(handle, EvaluationBudgetReservation)
cost: Final = receipt["response_cost"]
assert isinstance(cost, float)
await proxy_server.model_max_budget_limiter.async_log_success_event(receipt, None, None, None)
await proxy_server.increment_spend_counters(
token=None,
team_id=None,
user_id=owner.user_id,
response_cost=cost,
budget_reservation=handle.total,
)
assert await cache.async_get_cache("spend:user:evaluation-admin") == pytest.approx(
10 * 0.001 + 0.002 + 10 * 0.003 + 2 * 0.004
)
await cache.async_set_cache("spend:user:evaluation-admin", 1.0)
with respx.mock(assert_all_called=False), pytest.raises(litellm.BudgetExceededError):
await litellm.anthropic.messages.acreate(
**_EVALUATION_BRIDGE_REQUEST,
litellm_logging_obj=logging_obj,
api_base="https://evaluation.invalid/v1",
api_key="transport-only",
num_retries=0,
)
assert await cache.async_get_cache("spend:user:evaluation-admin") == pytest.approx(1.0)
@pytest.mark.asyncio
async def test_cached_evaluation_reusing_logger_does_not_release_prior_pending_spend(
evaluation_bridge_budget: tuple[EvaluationBillingOwner, DualCache],
monkeypatch: pytest.MonkeyPatch,
) -> None:
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.proxy import proxy_server
from litellm.proxy.spend_tracking.budget_reservation import estimate_request_max_cost
owner, cache = evaluation_bridge_budget
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
receipts: Final[asyncio.Queue[Mapping[str, object]]] = asyncio.Queue()
with respx.mock(assert_all_called=True) as transport, evaluation_billing_context(owner):
route: Final = transport.post("https://evaluation.invalid/v1/chat/completions").mock(
return_value=_evaluation_provider_response(),
)
logging_obj: Final = _deferred_evaluation_logger(receipts)
request: Final = {
**_EVALUATION_BRIDGE_REQUEST,
"api_base": "https://evaluation.invalid/v1",
"api_key": "transport-only",
"num_retries": 0,
"litellm_logging_obj": logging_obj,
}
await litellm.acompletion(**request)
original: Final = logging_obj.evaluation_budget_reservation
enqueue: Final = logging_obj._enqueue_deferred_logging
assert original is not None and enqueue is not None
logging_obj._enqueue_deferred_logging = None
await asyncio.gather(*tuple(_PENDING_CACHE_WRITES))
await litellm.acompletion(**request)
cached: Final = await asyncio.wait_for(receipts.get(), 10)
await GLOBAL_LOGGING_WORKER.flush()
assert route.call_count == 1
assert cached[EVALUATION_BUDGET_RESERVATION_KEY] is None
assert cached["response_cost"] == 0.0
assert await cache.async_get_cache("spend:user:evaluation-admin") == pytest.approx(
estimate_request_max_cost(_EVALUATION_BRIDGE_REQUEST, "/chat/completions", None)
)
enqueue()
receipt: Final = await asyncio.wait_for(receipts.get(), 10)
await GLOBAL_LOGGING_WORKER.flush()
assert receipt[EVALUATION_BUDGET_RESERVATION_KEY] is original
assert receipt["response_cost"] == pytest.approx(10 * 0.001 + 2 * 0.002)
await proxy_server.increment_spend_counters(
token=None,
team_id=None,
user_id=owner.user_id,
response_cost=0.014,
budget_reservation=original.total,
)
assert await cache.async_get_cache("spend:user:evaluation-admin") == pytest.approx(0.014)
@pytest.mark.asyncio
async def test_wrapper_async_does_not_fire_failure_hook_for_pre_call_budget_error(
monkeypatch: pytest.MonkeyPatch,

View file

@ -38,9 +38,9 @@ const DIRECTION_OPTIONS: readonly { value: ShadowEvalDirection; label: string }[
const START_FORM_DESCRIPTION: Record<ShadowEvalDirection, string> = {
forward:
"Duplicates a sampled slice of the selected targets' traffic (keys, teams, or users) through the auto-router and has an LLM judge compare both answers blind. Each target gets its own spend budget. The router's answers are never served to users; judge calls bill to the sampled traffic's own identity.",
"Duplicates a sampled slice of the selected targets' traffic (keys, teams, or users) through the auto-router and has an LLM judge compare both answers blind. Each target gets its own spend budget. The router's answers are never served to users; shadow and judge calls bill to the initiating admin.",
reverse:
"Duplicates a sampled slice of the traffic the auto-router already serves against a fixed baseline model and has an LLM judge compare both answers blind. Each target gets its own spend budget. The baseline's answers are never served to users; judge calls bill to the sampled traffic's own identity.",
"Duplicates a sampled slice of the traffic the auto-router already serves against a fixed baseline model and has an LLM judge compare both answers blind. Each target gets its own spend budget. The baseline's answers are never served to users; shadow and judge calls bill to the initiating admin.",
};
const DURATION_OPTIONS = [

View file

@ -1475,7 +1475,7 @@ export interface paths {
* eval spend, the shadow and judge calls' own cost, reaches max_budget dollars, the
* job's window ends, or the job is stopped, so one target running out of budget does
* not end sampling for the others; sampling changes propagate to pods within about 10
* seconds. Shadow and judge calls bill to the sampled request's own identity but are
* seconds. Shadow and judge calls bill to the admin who started the job and are
* excluded from request counts and auto-router adoption metrics.
*/
post: operations["start_shadow_eval_auto_router_shadow_eval_start_post"];