diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index f04ff8b6e78..f72adf8dbdc 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -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, diff --git a/litellm/integrations/datadog/datadog_cost_management.py b/litellm/integrations/datadog/datadog_cost_management.py index 538dd95abdd..7ad4867b568 100644 --- a/litellm/integrations/datadog/datadog_cost_management.py +++ b/litellm/integrations/datadog/datadog_cost_management.py @@ -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") diff --git a/litellm/integrations/lago.py b/litellm/integrations/lago.py index 594427b1e0a..370e66c7797 100644 --- a/litellm/integrations/lago.py +++ b/litellm/integrations/lago.py @@ -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 diff --git a/litellm/integrations/openmeter.py b/litellm/integrations/openmeter.py index db2fe386dec..a010c933ed4 100644 --- a/litellm/integrations/openmeter.py +++ b/litellm/integrations/openmeter.py @@ -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: diff --git a/litellm/integrations/shadow_eval_logger.py b/litellm/integrations/shadow_eval_logger.py index 7f5d9fedd3d..2cb9f0c9959 100644 --- a/litellm/integrations/shadow_eval_logger.py +++ b/litellm/integrations/shadow_eval_logger.py @@ -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, + ) diff --git a/litellm/litellm_core_utils/internal_call_metadata.py b/litellm/litellm_core_utils/internal_call_metadata.py index d844cbae367..84a540ab9be 100644 --- a/litellm/litellm_core_utils/internal_call_metadata.py +++ b/litellm/litellm_core_utils/internal_call_metadata.py @@ -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. diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 5a3a17f338c..acf5c2208d3 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 5395d817e33..d2fb5fe20bd 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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, ) diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index d019271d404..bcb1fa6830d 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -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, diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index f937b439042..130047d3283 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -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( diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index 33fb069afbd..3a8aca01dac 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -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, } diff --git a/litellm/proxy/native_compaction.py b/litellm/proxy/native_compaction.py index fd27e7fbbdc..f30c480c677 100644 --- a/litellm/proxy/native_compaction.py +++ b/litellm/proxy/native_compaction.py @@ -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"]), diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 2dd1c9419f1..73ab9ecc888 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -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, diff --git a/litellm/proxy/spend_tracking/evaluation_budget.py b/litellm/proxy/spend_tracking/evaluation_budget.py new file mode 100644 index 00000000000..09fa686edb2 --- /dev/null +++ b/litellm/proxy/spend_tracking/evaluation_budget.py @@ -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)) diff --git a/litellm/utils.py b/litellm/utils.py index d72588e2b00..95eb2f060de 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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 diff --git a/tests/integration/spend/test_evaluation_budget.py b/tests/integration/spend/test_evaluation_budget.py new file mode 100644 index 00000000000..9298ff753a5 --- /dev/null +++ b/tests/integration/spend/test_evaluation_budget.py @@ -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) diff --git a/tests/unit/integrations/datadog/test_datadog_cost_management.py b/tests/unit/integrations/datadog/test_datadog_cost_management.py index 1a50a6991da..226616815f4 100644 --- a/tests/unit/integrations/datadog/test_datadog_cost_management.py +++ b/tests/unit/integrations/datadog/test_datadog_cost_management.py @@ -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( diff --git a/tests/unit/integrations/test_openmeter.py b/tests/unit/integrations/test_openmeter.py index 2d09e1572db..d8b249ecca5 100644 --- a/tests/unit/integrations/test_openmeter.py +++ b/tests/unit/integrations/test_openmeter.py @@ -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): diff --git a/tests/unit/integrations/test_shadow_eval_logger.py b/tests/unit/integrations/test_shadow_eval_logger.py index e3f059a7941..c962c8cd5e7 100644 --- a/tests/unit/integrations/test_shadow_eval_logger.py +++ b/tests/unit/integrations/test_shadow_eval_logger.py @@ -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() diff --git a/tests/unit/litellm_core_utils/test_internal_call_metadata.py b/tests/unit/litellm_core_utils/test_internal_call_metadata.py index 73923dc75a5..ae897394e89 100644 --- a/tests/unit/litellm_core_utils/test_internal_call_metadata.py +++ b/tests/unit/litellm_core_utils/test_internal_call_metadata.py @@ -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} diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index 4b25da2ff79..e954f4feb18 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -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, ): diff --git a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py index fc8bc289735..1926d6db0c6 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py @@ -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, @@ -9661,20 +9663,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, @@ -9687,6 +9693,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 diff --git a/tests/unit/proxy/hooks/test_model_max_budget_limiter.py b/tests/unit/proxy/hooks/test_model_max_budget_limiter.py index ffb60fb4651..eebc1d47d46 100644 --- a/tests/unit/proxy/hooks/test_model_max_budget_limiter.py +++ b/tests/unit/proxy/hooks/test_model_max_budget_limiter.py @@ -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) diff --git a/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py b/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py index 7afa275c801..681db1a6d50 100644 --- a/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py @@ -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 diff --git a/tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py b/tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py index f194e43c74a..61e77908e4b 100644 --- a/tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py +++ b/tests/unit/proxy/hooks/test_unit_test_max_model_budget_limiter.py @@ -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 diff --git a/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py b/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py index 01e41e8b03f..1d516076cb9 100644 --- a/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_auto_router_endpoints.py @@ -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 diff --git a/tests/unit/proxy/spend_tracking/test_evaluation_budget.py b/tests/unit/proxy/spend_tracking/test_evaluation_budget.py new file mode 100644 index 00000000000..3f94b7fda13 --- /dev/null +++ b/tests/unit/proxy/spend_tracking/test_evaluation_budget.py @@ -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) diff --git a/tests/unit/proxy/test_native_compaction.py b/tests/unit/proxy/test_native_compaction.py index 778c1715530..d2409311639 100644 --- a/tests/unit/proxy/test_native_compaction.py +++ b/tests/unit/proxy/test_native_compaction.py @@ -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): diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index a72aec07766..e5069e4e4db 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -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, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalStartForm.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalStartForm.tsx index 06e332d2cfc..b332fa1ad91 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalStartForm.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalStartForm.tsx @@ -38,9 +38,9 @@ const DIRECTION_OPTIONS: readonly { value: ShadowEvalDirection; label: string }[ const START_FORM_DESCRIPTION: Record = { 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 = [ diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 3ed4103a227..27d061d78a3 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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"];