diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index 6a4987298c4..45fce665bf0 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -909,7 +909,8 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac """ import litellm from litellm import Choices, Message, ModelResponse - from litellm.litellm_core_utils.classifier_logging import CLASSIFIER_AUDIT_FIELDS, without_classifier_audit + from litellm.litellm_core_utils.classifier_logging import CLASSIFIER_AUDIT_FIELDS + from litellm.litellm_core_utils.redact_messages import redacted_litellm_params turn_off_message_logging: Final[bool] = getattr(self, "turn_off_message_logging", False) excluded_fields: Final[list[str] | None] = getattr(litellm, "standard_logging_payload_excluded_fields", None) @@ -918,9 +919,15 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac if turn_off_message_logging is False and not excluded_fields: return model_call_details + params: Final = model_call_details.get("litellm_params") + redacted_params: Final = ( + MappingProxyType({"litellm_params": redacted_litellm_params(params)}) + if turn_off_message_logging and isinstance(params, Mapping) + else EMPTY_MAPPING + ) standard_logging_object: Final = model_call_details.get("standard_logging_object") if standard_logging_object is None: - return model_call_details.copy() + return {**model_call_details, **redacted_params} # Make a copy of just the standard_logging_object to avoid modifying the original standard_logging_object_copy: Final = { @@ -960,13 +967,6 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac model_response_dict: Final = model_response.model_dump() standard_logging_object_copy["response"] = model_response_dict - params: Final = model_call_details.get("litellm_params") - request: Final = params.get("proxy_server_request") if isinstance(params, dict) else None - redacted_params: Final = ( - MappingProxyType({"litellm_params": {**params, "proxy_server_request": without_classifier_audit(request)}}) - if turn_off_message_logging and isinstance(params, dict) and isinstance(request, dict) - else EMPTY_MAPPING - ) return { **model_call_details, **redacted_params, diff --git a/litellm/litellm_core_utils/logging_worker.py b/litellm/litellm_core_utils/logging_worker.py index 03420b84c22..57da3f8dabe 100644 --- a/litellm/litellm_core_utils/logging_worker.py +++ b/litellm/litellm_core_utils/logging_worker.py @@ -23,6 +23,19 @@ from litellm.constants import ( MAX_TIME_TO_CLEAR_QUEUE, ) +_CALLBACK_DEADLINE: Final[contextvars.ContextVar[float | None]] = contextvars.ContextVar( + "logging_callback_deadline", default=None +) + + +def optional_callback_budget(maximum: float, *, fraction: float = 0.25) -> float: + deadline: Final = _CALLBACK_DEADLINE.get() + return ( + maximum + if deadline is None + else max(0.0, min(maximum, (deadline - asyncio.get_running_loop().time()) * fraction)) + ) + def _coroutine_name(coroutine: Coroutine) -> str: return getattr(coroutine, "__qualname__", None) or getattr(coroutine, "__name__", None) or type(coroutine).__name__ @@ -100,12 +113,20 @@ class LoggingWorker: return len(revived) def _run_coroutine_silently(self, loop: asyncio.AbstractEventLoop, coroutine: Coroutine) -> bool: + token: Final = _CALLBACK_DEADLINE.set(loop.time() + self.timeout) try: loop.run_until_complete(asyncio.wait_for(coroutine, timeout=self.timeout)) except (Exception, asyncio.CancelledError): # noqa: BLE001 # atexit flush must never break the user's program return False + finally: + _CALLBACK_DEADLINE.reset(token) return True + def _create_callback_task(self, task: LoggingTask) -> asyncio.Task[object]: + context: Final = task["context"].copy() + context.run(_CALLBACK_DEADLINE.set, asyncio.get_running_loop().time() + self.timeout) + return context.run(asyncio.create_task, task["coroutine"]) + @staticmethod def _drain_pending(queue: "asyncio.Queue[LoggingTask]") -> tuple[LoggingTask, ...]: """Pop every task still queued, without awaiting them, so they can be moved to another queue.""" @@ -172,7 +193,7 @@ class LoggingWorker: try: if self._queue is not None: # Run the coroutine in its original context - callback_task: Final = task["context"].run(asyncio.create_task, task["coroutine"]) + callback_task: Final = self._create_callback_task(task) try: await asyncio.wait_for(callback_task, timeout=self.timeout) except asyncio.TimeoutError as e: @@ -424,7 +445,7 @@ class LoggingWorker: try: await asyncio.wait_for( - task["context"].run(asyncio.create_task, task["coroutine"]), + self._create_callback_task(task), timeout=self.timeout, ) except Exception: @@ -517,7 +538,7 @@ class LoggingWorker: # Await the coroutine to properly execute and avoid "never awaited" warnings try: await asyncio.wait_for( - task["context"].run(asyncio.create_task, task["coroutine"]), + self._create_callback_task(task), timeout=self.timeout, ) except Exception: diff --git a/litellm/litellm_core_utils/redact_messages.py b/litellm/litellm_core_utils/redact_messages.py index 85ed0a40687..7c6c39abb76 100644 --- a/litellm/litellm_core_utils/redact_messages.py +++ b/litellm/litellm_core_utils/redact_messages.py @@ -11,6 +11,7 @@ import asyncio import copy import inspect from collections.abc import Mapping +from dataclasses import replace from typing import TYPE_CHECKING, Any, Final import litellm @@ -26,6 +27,7 @@ from litellm.llms.vertex_ai.common_utils import ( redact_vertex_ai_metadata_from_logged_object, ) from litellm.secret_managers.main import str_to_bool +from litellm.types.router import BaselineRouteStamp from litellm.types.utils import StandardCallbackDynamicParams if TYPE_CHECKING: @@ -252,6 +254,26 @@ def _redact_model_response_dict_choices(choices, redacted_str: str): _redact_choice_content(choice) +def _redacted_baseline_metadata(metadata: Mapping[str, object]) -> Mapping[str, object]: + route: Final = metadata.get("_autorouter_baseline_route") + if not isinstance(route, BaselineRouteStamp): + return metadata + return {**metadata, "_autorouter_baseline_route": replace(route, request_parameters=None)} + + +def redacted_litellm_params(params: Mapping[str, object]) -> dict[str, object]: + request: Final = params.get("proxy_server_request") + return { + **params, + **{ + key: _redacted_baseline_metadata(value) + for key, value in params.items() + if key in ("metadata", "litellm_metadata") and isinstance(value, Mapping) + }, + **({"proxy_server_request": without_classifier_audit(request)} if isinstance(request, Mapping) else {}), + } + + def perform_redaction(model_call_details: dict, result, redact_streaming_responses: bool = True): """ Performs the actual redaction on the logging object and result. @@ -262,9 +284,8 @@ def perform_redaction(model_call_details: dict, result, redact_streaming_respons """ # Redact model_call_details params: Final = model_call_details.get("litellm_params") - request: Final = params.get("proxy_server_request") if isinstance(params, dict) else None - if isinstance(params, dict) and isinstance(request, Mapping): - model_call_details["litellm_params"] = {**params, "proxy_server_request": without_classifier_audit(request)} + if isinstance(params, Mapping): + model_call_details["litellm_params"] = redacted_litellm_params(params) model_call_details["messages"] = [{"role": "user", "content": REDACTED_BY_LITELLM}] model_call_details["prompt"] = "" model_call_details["input"] = "" diff --git a/litellm/llms/anthropic/pass_through/messages/handler.py b/litellm/llms/anthropic/pass_through/messages/handler.py index 6d13e38aa45..4e7a154be67 100644 --- a/litellm/llms/anthropic/pass_through/messages/handler.py +++ b/litellm/llms/anthropic/pass_through/messages/handler.py @@ -15,9 +15,6 @@ import litellm from litellm.litellm_core_utils.exception_mapping_utils import exception_type from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.anthropic.common_utils import ( - flatten_unencrypted_web_search_results_in_anthropic_messages, - sanitize_tool_use_ids_in_anthropic_messages, - strip_empty_content_blocks_from_anthropic_messages, strip_provider_specific_fields_from_anthropic_messages, ) from litellm.llms.base_llm.anthropic_messages.transformation import ( @@ -36,9 +33,8 @@ from litellm.utils import ProviderConfigManager, client from ..adapters.handler import LiteLLMMessagesToCompletionTransformationHandler from ..responses_adapters.handler import LiteLLMMessagesToResponsesAPIHandler -from ..utils import is_reasoning_auto_summary_enabled from .interceptors import get_messages_interceptors -from .utils import AnthropicMessagesRequestUtils, mock_response +from .utils import AnthropicMessagesRequestUtils, mock_response, prepare_native_messages __all__ = ("anthropic_messages", "anthropic_messages_handler") @@ -251,28 +247,7 @@ async def anthropic_messages( Runs the empty-content-block sanitizer before any backend dispatch. """ - # Anthropic's API rejects requests containing empty / whitespace-only - # text content blocks ("messages: text content blocks must be - # non-empty") and empty thinking blocks ("each thinking block must - # contain thinking"). Multi-turn tool-use clients (e.g. Claude Code) - # routinely loop assistant responses that contain such blocks — an empty - # text block alongside tool_use, or an empty thinking block from a turn - # a non-Anthropic reasoning model served through the bridge — back as - # conversation history, which then causes the next /v1/messages call to - # 400. /v1/chat/completions already handles this in - # anthropic_messages_pt; sanitize the native Anthropic Messages path - # here for the same guarantee. See #22930. - messages = strip_empty_content_blocks_from_anthropic_messages(messages) - # Replay of cross-provider tool history (e.g. kimi -> Anthropic) may carry - # ids like ``functions.Bash:0`` that violate Anthropic's id pattern. - messages = sanitize_tool_use_ids_in_anthropic_messages(messages) - messages = flatten_unencrypted_web_search_results_in_anthropic_messages(messages) - - from litellm.integrations.anthropic_cache_control_hook import ( - AnthropicCacheControlHook, - ) - - messages, system = AnthropicCacheControlHook.maybe_inject_cache_control( + messages, system = prepare_native_messages( messages, system, kwargs, model=model, custom_llm_provider=custom_llm_provider, tools=tools, api_base=api_base ) @@ -454,23 +429,15 @@ def anthropic_messages_handler( """ from litellm.types.utils import LlmProviders - # Sanitize empty text blocks so the sync entry point - # (litellm.messages.create -> anthropic_messages_handler) gets the same - # protection as the async wrapper. The async wrapper already sanitized and - # does not reassign messages before dispatch, so it sets - # ``_litellm_messages_presanitized`` to skip this redundant second - # full-messages scan. Pop it so it never leaks into provider params. - if not kwargs.pop("_litellm_messages_presanitized", False): - messages = strip_empty_content_blocks_from_anthropic_messages(messages) - messages = sanitize_tool_use_ids_in_anthropic_messages(messages) - messages = flatten_unencrypted_web_search_results_in_anthropic_messages(messages) - - from litellm.integrations.anthropic_cache_control_hook import ( - AnthropicCacheControlHook, - ) - - messages, system = AnthropicCacheControlHook.maybe_inject_cache_control( - messages, system, kwargs, model=model, custom_llm_provider=custom_llm_provider, tools=tools, api_base=api_base + messages, system = prepare_native_messages( + messages, + system, + kwargs, + model=model, + custom_llm_provider=custom_llm_provider, + tools=tools, + api_base=api_base, + presanitized=bool(kwargs.pop("_litellm_messages_presanitized", False)), ) metadata = validate_anthropic_api_metadata(metadata) @@ -645,14 +612,6 @@ def anthropic_messages_handler( custom_llm_provider=custom_llm_provider, ) ) - if is_reasoning_auto_summary_enabled(): - thinking_param: Final = anthropic_messages_optional_request_params.get("thinking") - if isinstance(thinking_param, dict) and thinking_param.get("type") != "disabled": - anthropic_messages_optional_request_params["thinking"] = { - **thinking_param, - "display": "summarized", - } - resolved_api_base: Final = ( dynamic_api_base if dynamic_api_base is not None and anthropic_messages_provider_config.uses_get_llm_provider_api_base() diff --git a/litellm/llms/anthropic/pass_through/messages/utils.py b/litellm/llms/anthropic/pass_through/messages/utils.py index dfab0af8eaa..8371615baea 100644 --- a/litellm/llms/anthropic/pass_through/messages/utils.py +++ b/litellm/llms/anthropic/pass_through/messages/utils.py @@ -2,6 +2,15 @@ from collections.abc import Iterable, Mapping, Sequence from functools import lru_cache from typing import TYPE_CHECKING, Any, Final, cast, get_type_hints +from pydantic import JsonValue + +from litellm.integrations.anthropic_cache_control_hook import AnthropicCacheControlHook +from litellm.llms.anthropic.common_utils import ( + flatten_unencrypted_web_search_results_in_anthropic_messages, + sanitize_tool_use_ids_in_anthropic_messages, + strip_empty_content_blocks_from_anthropic_messages, +) +from litellm.llms.anthropic.pass_through.utils import is_reasoning_auto_summary_enabled from litellm.types.llms.anthropic import ( AnthropicMessagesRequestOptionalParams, AnthropicStopDetails, @@ -119,8 +128,40 @@ def anthropic_system_to_openai_message(system: object) -> ChatCompletionSystemMe return ChatCompletionSystemMessage(role="system", content=system) +def prepare_native_messages( + messages: list[dict[str, JsonValue]], + system: str | list[dict[str, JsonValue]] | None, + kwargs: dict[str, object], + *, + model: str, + custom_llm_provider: str | None = None, + tools: list[dict[str, JsonValue]] | None = None, + api_base: str | None = None, + presanitized: bool = False, +) -> tuple[list[dict[str, JsonValue]], str | list[dict[str, JsonValue]] | None]: + normalized: Final = ( + messages + if presanitized + else flatten_unencrypted_web_search_results_in_anthropic_messages( + sanitize_tool_use_ids_in_anthropic_messages(strip_empty_content_blocks_from_anthropic_messages(messages)) + ) + ) + return cast( # cast-ok: legacy normalizers and injection preserve the JSON message and system shapes + tuple[list[dict[str, JsonValue]], str | list[dict[str, JsonValue]] | None], + AnthropicCacheControlHook.maybe_inject_cache_control( + normalized, + system, + kwargs, + model=model, + custom_llm_provider=custom_llm_provider, + tools=tools, + api_base=api_base, + ), + ) + + @lru_cache(maxsize=1) -def _anthropic_messages_optional_param_keys() -> frozenset[str]: +def anthropic_messages_optional_param_keys() -> frozenset[str]: """ Valid AnthropicMessagesRequestOptionalParams keys. @@ -152,7 +193,7 @@ class AnthropicMessagesRequestUtils: Returns: AnthropicMessagesRequestOptionalParams instance with only the valid parameters """ - valid_keys: Final = _anthropic_messages_optional_param_keys() + valid_keys: Final = anthropic_messages_optional_param_keys() filtered_params: Final = {k: v for k, v in params.items() if k in valid_keys and v is not None} if model is not None: from litellm.llms.anthropic.chat.transformation import AnthropicConfig @@ -174,6 +215,13 @@ class AnthropicMessagesRequestUtils: drop_params=drop_params, output_key=param, ) + if is_reasoning_auto_summary_enabled(): + thinking_param: Final = filtered_params.get("thinking") + if isinstance(thinking_param, dict) and thinking_param.get("type") != "disabled": + return cast( + AnthropicMessagesRequestOptionalParams, + {**filtered_params, "thinking": {**thinking_param, "display": "summarized"}}, + ) return cast(AnthropicMessagesRequestOptionalParams, filtered_params) diff --git a/litellm/llms/anthropic/prompt_cache_prediction.py b/litellm/llms/anthropic/prompt_cache_prediction.py index 00fe56a5e39..634eb264b38 100644 --- a/litellm/llms/anthropic/prompt_cache_prediction.py +++ b/litellm/llms/anthropic/prompt_cache_prediction.py @@ -5,6 +5,7 @@ import hashlib import json from collections.abc import Mapping, Sequence from dataclasses import dataclass, field +from functools import reduce from itertools import accumulate, groupby from types import MappingProxyType from typing import Annotated, Final, Literal, Protocol, TypeAlias @@ -13,20 +14,32 @@ import httpx from pydantic import ConfigDict, Field, JsonValue, StrictInt, TypeAdapter, ValidationError import litellm -from litellm.llms.anthropic.common_utils import AnthropicModelInfo, is_anthropic_oauth_key +from litellm.litellm_core_utils.dot_notation_indexing import delete_nested_value +from litellm.llms.anthropic.common_utils import ( + AnthropicModelInfo, + is_anthropic_oauth_key, + strip_provider_specific_fields_from_anthropic_messages, +) from litellm.llms.anthropic.count_tokens.handler import AnthropicCountTokensHandler from litellm.llms.anthropic.count_tokens.transformation import COUNT_TOKEN_OPTION_NAMES from litellm.llms.anthropic.pass_through.messages.transformation import ( DEFAULT_ANTHROPIC_API_VERSION, AnthropicMessagesConfig, ) +from litellm.llms.anthropic.pass_through.messages.utils import AnthropicMessagesRequestUtils, prepare_native_messages +from litellm.router_utils.baseline_request import ( + BASELINE_PARAMETERS, + capture_baseline_parameters, +) from litellm.types.llms.base import LiteLLMBaseModel -from litellm.types.router import LiteLLM_Params +from litellm.types.router import GenericLiteLLMParams, LiteLLM_Params from litellm.types.utils import ModelResponse from litellm.utils import supports_thinking_cache_preservation _JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) _HEADERS: Final = TypeAdapter(dict[str, str]) +_MESSAGES: Final = TypeAdapter(list[dict[str, JsonValue]]) +_SYSTEM: Final = TypeAdapter(str | list[dict[str, JsonValue]] | None) _counter: Final = AnthropicCountTokensHandler() @@ -325,7 +338,7 @@ def _entry_fingerprint(fingerprint: str, ttl_seconds: int) -> str: def parse_cache_plan(body: Mapping[str, JsonValue]) -> PromptCachePlan | UnsupportedCachePlan: try: - request: Final = _PlanRequest.model_validate(body) + request: Final = _PlanRequest.model_validate(dict(body)) positions: Final = _positions(body) except ValidationError: return UnsupportedCachePlan("unsupported_prompt_shape") @@ -618,13 +631,56 @@ def resolve_baseline_prediction_target(params: LiteLLM_Params) -> NativePredicti return _resolve_prediction_target(params, allow_configured_endpoint=True) +def prepare_native_baseline_body(request: Mapping[str, object], model: str) -> Mapping[str, JsonValue] | None: + parameters: Final = capture_baseline_parameters(request) + if parameters is None: + return None + source: Final = {**parameters, "messages": request.get("messages"), "stream": request.get("stream", False)} + try: + owned: Final = _JSON_OBJECT.validate_python(source) + context: Final = {**{k: v for k, v in request.items() if k not in ("metadata", "litellm_metadata")}, **owned} + resolved_model: Final = litellm.get_llm_provider(model=model, custom_llm_provider="anthropic")[0] + messages, system = prepare_native_messages( + _MESSAGES.validate_python(owned.get("messages")), + _SYSTEM.validate_python(owned.get("system")), + context, + model=resolved_model, + custom_llm_provider="anthropic", + tools=_MESSAGES.validate_python(owned.get("tools") or []), + ) + options: Final = AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param( + {**owned, "system": system}, + model=resolved_model, + custom_llm_provider="anthropic", + drop_params=owned.get("drop_params") is True, + ) + filtered: Final = reduce( + delete_nested_value, + TypeAdapter(tuple[str, ...]).validate_python(owned.get("additional_drop_params") or ()), + dict(options), + ) + body: Final = AnthropicMessagesConfig().transform_anthropic_messages_request( + model=resolved_model, + messages=strip_provider_specific_fields_from_anthropic_messages(messages), + anthropic_messages_optional_request_params=filtered, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + return MappingProxyType(_JSON_OBJECT.validate_python(body)) + except Exception: # noqa: BLE001 # an unsupported hypothetical request is unavailable, never an inference failure + return None + + def _resolve_prediction_target( params: LiteLLM_Params, *, allow_configured_endpoint: bool, ) -> NativePredictionTarget | UnsupportedPredictionTarget: configured_options: Final = frozenset(params.model_dump(exclude_defaults=True, exclude_none=True)) - if configured_options - _DEPLOYMENT_OPTIONS: + allowed: Final = ( + _DEPLOYMENT_OPTIONS | frozenset(BASELINE_PARAMETERS) if allow_configured_endpoint else _DEPLOYMENT_OPTIONS + ) + if configured_options - allowed: return UnsupportedPredictionTarget("unsupported_deployment_configuration") api_base: Final = AnthropicModelInfo.get_api_base(params.api_base) if not allow_configured_endpoint and api_base not in ( diff --git a/litellm/proxy/db/baseline_accounting.py b/litellm/proxy/db/baseline_accounting.py index b21b7c8a2b9..486717f9b93 100644 --- a/litellm/proxy/db/baseline_accounting.py +++ b/litellm/proxy/db/baseline_accounting.py @@ -223,7 +223,8 @@ ON CONFLICT (request_id) DO NOTHING _MARK_CONFLICT: Final = """ UPDATE "LiteLLM_AutoRouterBaselineObservation" SET conflicted = TRUE, revision = $4::bigint -WHERE request_id = $1 AND scope = $2 AND data <> $3 AND NOT conflicted +WHERE request_id = $1 AND scope = $2 AND NOT conflicted + AND (data::jsonb #- '{turn,turn_at}') <> ($3::jsonb #- '{turn,turn_at}') """ _READ_PAGE: Final = """ WITH times AS ( diff --git a/litellm/proxy/hooks/autorouter_baseline_cache.py b/litellm/proxy/hooks/autorouter_baseline_cache.py index e3cd6c67aa0..c4232609124 100644 --- a/litellm/proxy/hooks/autorouter_baseline_cache.py +++ b/litellm/proxy/hooks/autorouter_baseline_cache.py @@ -5,7 +5,7 @@ import hashlib import json import time from collections.abc import Callable, Mapping -from dataclasses import dataclass, replace +from dataclasses import dataclass, field, replace from datetime import datetime from types import MappingProxyType from typing import TYPE_CHECKING, Final @@ -16,9 +16,8 @@ from pydantic import ConfigDict, Field, JsonValue, TypeAdapter from litellm._logging import verbose_proxy_logger from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.core_helpers import ( - get_litellm_metadata_from_kwargs, # pyright: ignore[reportUnknownVariableType] # legacy metadata boundary validated below -) +from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs +from litellm.litellm_core_utils.logging_worker import optional_callback_budget from litellm.llms.anthropic.prompt_cache_prediction import ( CountedPromptCachePlan, NativePredictionTarget, @@ -28,6 +27,7 @@ from litellm.llms.anthropic.prompt_cache_prediction import ( count_cache_plan, count_prompt_tokens, parse_cache_plan, + prepare_native_baseline_body, resolve_baseline_prediction_target, supported_baseline_recipient, supported_prediction_headers, @@ -37,6 +37,8 @@ from litellm.proxy.spend_tracking.savings import ( _effective_model_info, # pyright: ignore[reportPrivateUsage] # existing deployment-price owner _proxy_llm_router, # pyright: ignore[reportPrivateUsage] # existing optional proxy-router owner ) +from litellm.router_strategy.complexity_router.context_compaction import compaction_applied +from litellm.router_utils.baseline_request import baseline_request from litellm.types.llms.base import LiteLLMBaseModel from litellm.types.router import BaselineRouteStamp from litellm.types.utils import CallTypes, ModelInfo, Usage @@ -66,6 +68,9 @@ class CapturedBaselineObservation(LiteLLMBaseModel): prices: ModelInfo | None observation: BaselineObservation + def with_observation(self, observation: BaselineObservation) -> CapturedBaselineObservation: + return self.model_copy(update={"observation": observation}) + @dataclass(frozen=True, slots=True) class BaselineCacheContext: @@ -73,7 +78,10 @@ class BaselineCacheContext: capture: CapturedBaselineObservation target: NativePredictionTarget | UnsupportedPredictionTarget baseline_deployment_id: str + baseline_body: Mapping[str, JsonValue] | None = field(default=None, repr=False) + selected_body_digest: str | None = field(default=None, repr=False) invalidated: str | None = None + finalization: asyncio.Task[CapturedBaselineObservation] | None = field(default=None, repr=False, compare=False) class _Metadata(LiteLLMBaseModel): @@ -102,6 +110,10 @@ def _digest(value: object) -> str: return hashlib.sha256(json.dumps(value, sort_keys=True, separators=(",", ":")).encode()).hexdigest() +def _native_body_digest(body: Mapping[str, JsonValue]) -> str: + return _digest({key: value for key, value in body.items() if key not in ("metadata", "stream")}) + + class AutoRouterBaselineCache(CustomLogger): def __init__( self, @@ -124,12 +136,15 @@ class AutoRouterBaselineCache(CustomLogger): if not isinstance(logging_obj, Logging) or call_type != CallTypes.anthropic_messages: return try: - metadata: Final = _METADATA.validate_python(get_litellm_metadata_from_kwargs({"litellm_params": kwargs})) + raw_metadata: Final = kwargs.get(get_metadata_variable_name_from_kwargs(kwargs)) + metadata: Final = _METADATA.validate_python(raw_metadata) if isinstance(raw_metadata, Mapping) else {} if metadata.get(INTERNAL_CALL_ORIGIN_METADATA_KEY): return if logging_obj.baseline_cache_context is not None: await invalidate_baseline_cache(logging_obj, "retried_request") return + if not isinstance(metadata.get("_autorouter_baseline_route"), BaselineRouteStamp): + return request: Final = _Metadata.model_validate(metadata) session: Final = kwargs.get("litellm_session_id") or request.session_id or logging_obj.litellm_session_id if not isinstance(session, str) or not session or len(session) > 256: @@ -142,13 +157,27 @@ class AutoRouterBaselineCache(CustomLogger): prices: Final = _PRICES.validate_python( _effective_model_info(router, request.route.baseline_deployment_id, request.route.baseline_model) ) + params: Final = ( + _METADATA.validate_python(deployment.litellm_params.model_dump(mode="json")) if deployment else {} + ) + projected: Final = ( + baseline_request( + kwargs, + request.route.request_parameters, + params, + include_extra_body=False, + ) + if request.route.request_parameters is not None + else None + ) scope: Final = "autorouter-baseline:v3:" + _digest( ( + "baseline_request_v4", request.user_api_key_hash, session, request.route.router_name, request.route.baseline_deployment_id, - deployment.litellm_params.model_dump(mode="json"), + params, prices, ) ) @@ -170,8 +199,21 @@ class AutoRouterBaselineCache(CustomLogger): reason="incomplete_response", ), ) + selected_model: Final = kwargs.get("model") + selected_body: Final = prepare_native_baseline_body( + kwargs, selected_model if isinstance(selected_model, str) else logging_obj.model + ) logging_obj.baseline_cache_context = BaselineCacheContext( - self, capture, target, request.route.baseline_deployment_id + self, + capture, + target, + request.route.baseline_deployment_id, + prepare_native_baseline_body(projected, target.model) + if projected is not None and isinstance(target, NativePredictionTarget) + else None, + _native_body_digest(selected_body) + if selected_body is not None and not compaction_applied(kwargs) + else None, ) except Exception: # noqa: BLE001 # optional observation cannot fail inference verbose_proxy_logger.warning("Auto-router baseline observation could not be initialized") @@ -197,14 +239,16 @@ class AutoRouterBaselineCache(CustomLogger): async def plan( self, target: NativePredictionTarget, wire: httpx.Request, body: Mapping[str, JsonValue], usage: Usage | None ) -> tuple[CountedPromptCachePlan | None, str | None]: + deadline: Final = asyncio.get_running_loop().time() + optional_callback_budget(_COUNT_TIMEOUT, fraction=0.75) if not supported_prediction_headers(wire.headers): return None, "unsupported_request_headers" plan: Final = parse_cache_plan(body) if isinstance(plan, UnsupportedCachePlan): return None, plan.reason details: Final = usage.prompt_tokens_details if usage is not None else None + selected: Final = parse_cache_plan(_JSON_BODY.validate_json(wire.content)) if ( - not plan.breakpoints + (isinstance(selected, UnsupportedCachePlan) or not selected.breakpoints) and details is not None and ((details.cached_tokens or 0) + (details.cache_creation_tokens or 0)) ): @@ -215,7 +259,8 @@ class AutoRouterBaselineCache(CustomLogger): try: counted: Final = await asyncio.wait_for( - count_cache_plan(target.model, target.api_key, plan, token_counter=count), timeout=_COUNT_TIMEOUT + count_cache_plan(target.model, target.api_key, plan, token_counter=count), + timeout=max(0.0, deadline - asyncio.get_running_loop().time()), ) return (None, counted.reason) if isinstance(counted, UnsupportedCachePlan) else (counted, None) except TimeoutError: @@ -229,111 +274,121 @@ async def invalidate_baseline_cache(logging_obj: Logging, reason: str, *, comple if context is not None: logging_obj.baseline_cache_context = replace(context, invalidated=reason) logging_obj.baseline_observation = context.capture.model_copy( - update=MappingProxyType( - { - "observation": context.capture.observation.model_copy( - update=MappingProxyType( - { - "available_at": max(context.capture.observation.started_at, context.collector.clock()), - "reason": reason, - } - ) - ), - } - ) + update={ + "observation": context.capture.observation.model_copy( + update={ + "available_at": max(context.capture.observation.started_at, context.collector.clock()), + "reason": reason, + } + ), + } ) async def finalize_baseline_cache(logging_obj: Logging, response_obj: object) -> None: context: Final = logging_obj.baseline_cache_context - if context is None: + if context is None or logging_obj.baseline_observation is not None: return + task: Final = context.finalization or asyncio.create_task(_capture(context, logging_obj, response_obj)) + active: Final = context if context.finalization is not None else replace(context, finalization=task) + if context.finalization is None: + task.add_done_callback(_consume_finalization) + logging_obj.baseline_cache_context = active try: - capture: Final = await _capture(context, logging_obj, response_obj) - if logging_obj.baseline_cache_context is context: - logging_obj.baseline_observation = capture # rebind-ok: attach only to the captured request owner - except Exception: # noqa: BLE001 # observation failures must preserve inference and billing + capture: Final = await asyncio.shield(task) + if logging_obj.baseline_cache_context is active: + logging_obj.baseline_observation = capture # rebind-ok: publish only for the current attempt + except Exception: # noqa: BLE001 # estimation must preserve inference and billing await invalidate_baseline_cache(logging_obj, "observation_unavailable") -async def _capture( +def _consume_finalization(task: asyncio.Task[CapturedBaselineObservation]) -> None: + if not task.cancelled(): + task.exception() + + +async def _capture_native( context: BaselineCacheContext, logging_obj: Logging, response_obj: object ) -> CapturedBaselineObservation: - original: Final = context.capture.observation - details: Final = _METADATA.validate_python(logging_obj.model_call_details) - if details.get("cache_hit") is True: - return context.capture.model_copy( - update=MappingProxyType( - { - "observation": original.model_copy( - update=MappingProxyType({"outcome": "response_cache", "reason": "response_cache_hit"}) - ) - } - ) - ) - event: Final = _WireEvent.model_validate(details) + capture: Final = context.capture + original: Final = capture.observation + event: Final = _WireEvent.model_validate(logging_obj.model_call_details) wire: Final = event.httpx_response.request usage: Final = _ResponseUsage.model_validate(response_obj).usage + available: Final = event.completion_start_time.timestamp() complete: Final = ( event.custom_llm_provider == "anthropic" and event.httpx_response.status_code == 200 and (not event.stream or event.prompt_cache_response_complete) ) - started: Final = original.started_at - available: Final = event.completion_start_time.timestamp() - if context.invalidated or not complete or not started <= available <= context.collector.clock(): - return context.capture.model_copy( - update=MappingProxyType( - { - "observation": original.model_copy( - update=MappingProxyType( - { - "available_at": max(started, context.collector.clock()), - "reason": context.invalidated or "incomplete_response", - } - ) - ) + if context.invalidated or not complete or not original.started_at <= available <= context.collector.clock(): + return capture.with_observation( + original.model_copy( + update={ + "available_at": max(original.started_at, context.collector.clock()), + "reason": context.invalidated or "incomplete_response", } ) ) target: Final = context.target if isinstance(target, UnsupportedPredictionTarget) or not supported_baseline_recipient(target, wire): - return context.capture.model_copy( - update=MappingProxyType( - { - "observation": original.model_copy( - update=MappingProxyType( - { - "available_at": available, - "reason": target.reason - if isinstance(target, UnsupportedPredictionTarget) - else "unsupported_baseline_recipient", - } - ) - ) + return capture.with_observation( + original.model_copy( + update={ + "available_at": available, + "reason": target.reason + if isinstance(target, UnsupportedPredictionTarget) + else "unsupported_baseline_recipient", } ) ) body: Final = _JSON_BODY.validate_json(wire.content) - same: Final = ( - logging_obj.get_router_model_id() == context.baseline_deployment_id and body.get("model") == target.model - ) - plan, reason = await context.collector.plan(target, wire, body, usage) - minimum: Final = get_prompt_cache_min_tokens(target.model) - return context.capture.model_copy( - update=MappingProxyType( - { - "observation": BaselineObservation( - request_id=original.request_id, - started_at=started, - available_at=available, - outcome="complete", - baseline_equivalent=same, - usage=usage, - plan=plan, - minimum_cache_tokens=minimum, - reason=reason, - ) - } + projected: Final = context.baseline_body + if projected is None or context.selected_body_digest != _native_body_digest(body): + return capture.with_observation( + original.model_copy( + update={ + "available_at": available, + "usage": usage, + "reason": "unsupported_baseline_settings" + if projected is None + else "unsupported_request_transformation", + } + ) + ) + same: Final = logging_obj.get_router_model_id() == context.baseline_deployment_id and _native_body_digest( + projected + ) == _native_body_digest(body) + plan, reason = await context.collector.plan(target, wire, projected, usage) + return capture.with_observation( + BaselineObservation( + request_id=original.request_id, + started_at=original.started_at, + available_at=available, + outcome="complete", + baseline_equivalent=same, + usage=usage.model_copy(update={key: projected.get(key) for key in ("speed", "inference_geo")}) + if usage is not None and not same + else usage, + plan=plan, + reason=reason, + minimum_cache_tokens=get_prompt_cache_min_tokens(target.model), ) ) + + +async def _capture( + context: BaselineCacheContext, logging_obj: Logging, response_obj: object +) -> CapturedBaselineObservation: + if _METADATA.validate_python(logging_obj.model_call_details).get("cache_hit") is True: + return context.capture.model_copy( + update={ + "observation": context.capture.observation.model_copy( + update={ + "outcome": "response_cache", + "reason": "response_cache_hit", + } + ), + } + ) + return await _capture_native(context, logging_obj, response_obj) diff --git a/litellm/proxy/spend_tracking/baseline_accounting.py b/litellm/proxy/spend_tracking/baseline_accounting.py index 46a38f71260..e344c0d7169 100644 --- a/litellm/proxy/spend_tracking/baseline_accounting.py +++ b/litellm/proxy/spend_tracking/baseline_accounting.py @@ -71,12 +71,12 @@ def _complete_usage(usage: Usage | None) -> bool: if usage is None or usage.prompt_tokens < 0 or usage.completion_tokens < 0: return False details: Final = usage.prompt_tokens_details - if details is None: + if details is None or not hasattr(details, "cache_creation_tokens"): return False values: Final = (details.text_tokens, details.cached_tokens, details.cache_creation_tokens) if any(value is None or value < 0 for value in values): return False - split: Final = details.cache_creation_token_details + split: Final = details.cache_creation_token_details if hasattr(details, "cache_creation_token_details") else None writes: Final = details.cache_creation_tokens or 0 return ( usage.total_tokens == usage.prompt_tokens + usage.completion_tokens @@ -133,10 +133,13 @@ def _matches(entry: CacheEntry, markers: tuple[CountedBreakpoint, ...], started: def _ambiguous(entry: CacheEntry, markers: tuple[CountedBreakpoint, ...], started: float) -> bool: - return entry.available_at <= started < entry.expires_at and any( - entry.content_fingerprint in marker.lookback_content_fingerprints - and (entry.uncertain or entry.ttl_seconds != marker.ttl_seconds) - for marker in markers + matching: Final = tuple( + marker for marker in markers if entry.content_fingerprint in marker.lookback_content_fingerprints + ) + return ( + entry.available_at <= started < entry.expires_at + and bool(matching) + and (entry.uncertain or all(entry.ttl_seconds != marker.ttl_seconds for marker in matching)) ) @@ -261,7 +264,7 @@ def _writes(history: BaselineHistory, observation: BaselineObservation) -> tuple observation.started_at + hit.ttl_seconds, ), ) - if hit is not None and all(marker.fingerprint != hit.fingerprint for marker in markers) + if hit is not None else () ) return ( @@ -277,6 +280,7 @@ def _writes(history: BaselineHistory, observation: BaselineObservation) -> tuple uncertain=bool(ambiguous), ) for marker in markers + if hit is None or marker.prefix_tokens > hit.tokens ), ) diff --git a/litellm/router.py b/litellm/router.py index 77cff6254ea..ab9d8884a1d 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -14582,17 +14582,21 @@ class Router: to the deployment that actually served the request. Every attempt therefore writes or clears, never just writes. """ + from litellm.router_utils.baseline_request import capture_baseline_parameters from litellm.types.router import BaselineRouteStamp phase_attributes(routing_decision_attributes(routing_decision)) baseline_model: Final = routing_decision.get("savings_baseline_model") if routing_decision else None baseline_id: Final = routing_decision.get("savings_baseline_deployment_id") if routing_decision else None router_name: Final = routing_decision.get("router_model_name") if routing_decision else None + caller_parameters: Final = ( + capture_baseline_parameters(request_kwargs) if router_name and baseline_model else None + ) Router._stamp_or_clear_metadata_key( request_kwargs=request_kwargs, key="_autorouter_baseline_route", value=( - BaselineRouteStamp(router_name, baseline_model, baseline_id) + BaselineRouteStamp(router_name, baseline_model, baseline_id, caller_parameters) if router_name and baseline_model and baseline_id else None ), diff --git a/litellm/router_strategy/complexity_router/context_compaction.py b/litellm/router_strategy/complexity_router/context_compaction.py index d82090f3b99..6a1a63ea4b6 100644 --- a/litellm/router_strategy/complexity_router/context_compaction.py +++ b/litellm/router_strategy/complexity_router/context_compaction.py @@ -164,6 +164,11 @@ def compaction_pending(kwargs: Mapping[str, object] | None) -> bool: return isinstance(state, CompactionState) and state.config is not None and not _client_managed(kwargs or _EMPTY) +def compaction_applied(kwargs: Mapping[str, object]) -> bool: + state: Final = kwargs.get(_STATE_KEY) + return isinstance(state, CompactionState) and state.summary is not None + + def _reject(model: str, reason: str) -> NoReturn: from litellm.exceptions import BadRequestError diff --git a/litellm/router_utils/baseline_request.py b/litellm/router_utils/baseline_request.py new file mode 100644 index 00000000000..2f6262197a1 --- /dev/null +++ b/litellm/router_utils/baseline_request.py @@ -0,0 +1,153 @@ +from __future__ import annotations + +from collections.abc import Iterator, Mapping +from itertools import accumulate +from types import MappingProxyType +from typing import Final, cast + +from pydantic import JsonValue, TypeAdapter, ValidationError + +from litellm.llms.anthropic.pass_through.messages.utils import anthropic_messages_optional_param_keys + +CACHE_SETTINGS: Final = ( + "system", + "instructions", + "tools", + "tool_choice", + "parallel_tool_calls", + "response_format", + "text", + "reasoning", + "reasoning_effort", + "thinking", + "verbosity", + "output_config", + "output_format", + "speed", + "prompt_cache_key", + "cache_key", + "cached_content", + "previous_response_id", + "conversation", + "context_management", + "compaction", +) +_GENERIC_PARAMETERS: Final = ( + *CACHE_SETTINGS, + "prompt_cache_options", + "prompt_cache_retention", + "cache_control", + "max_tokens", + "max_completion_tokens", + "max_output_tokens", + "temperature", + "top_p", + "top_k", + "stop_sequences", + "enable_prompt_caching", + "cache_control_injection_points", + "drop_params", + "additional_drop_params", +) +NATIVE_ONLY_PARAMETERS: Final = tuple( + key + for key in sorted(anthropic_messages_optional_param_keys()) + if key not in (*_GENERIC_PARAMETERS, "metadata", "stream") +) +BASELINE_PARAMETERS: Final = (*_GENERIC_PARAMETERS, *NATIVE_ONLY_PARAMETERS) +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_MAX_BYTES: Final = 4 * 1024 * 1024 +_MAX_NODES: Final = 32768 +_MAX_DEPTH: Final = 32 + + +def _json_cost(value: object, depth: int = 0) -> Iterator[int]: + if depth > _MAX_DEPTH: + yield _MAX_BYTES + 1 + elif isinstance(value, str): + yield (6 if value.isascii() else 12) * len(value) + 2 + elif isinstance(value, dict): + yield 2 + for key, item in cast(dict[object, object], value).items(): + yield from _json_cost(key, depth + 1) + yield from _json_cost(item, depth + 1) + yield 2 + elif isinstance(value, (list, tuple)): + yield 2 + for item in cast(list[object] | tuple[object, ...], value): + yield from _json_cost(item, depth + 1) + yield 1 + elif isinstance(value, int) and value.bit_length() > 64: + yield _MAX_BYTES + 1 + elif value is None or isinstance(value, (bool, int, float)): + yield 32 + else: + yield _MAX_BYTES + 1 + + +def within_baseline_budget(value: object) -> bool: + return all( + size <= _MAX_BYTES and nodes <= _MAX_NODES for nodes, size in enumerate(accumulate(_json_cost(value)), 1) + ) + + +def _parameters(value: object, *, envelope: bool = False) -> dict[str, object]: + if not isinstance(value, Mapping): + return {} + mapping: Final = cast(Mapping[str, object], value) + keys: Final = (*BASELINE_PARAMETERS, "messages") if envelope else BASELINE_PARAMETERS + return {key: mapping[key] for key in keys if key in mapping} + + +def capture_baseline_parameters( + kwargs: Mapping[str, object], *, include_extra_body: bool = True +) -> Mapping[str, JsonValue] | None: + extra: Final = ( + {"extra_body": _parameters(kwargs.get("extra_body"), envelope=True)} + if include_extra_body and "extra_body" in kwargs + else {} + ) + parameters: Final = {**_parameters(kwargs), **extra} + if not within_baseline_budget(parameters): + return None + try: + return MappingProxyType(_JSON_OBJECT.validate_python(parameters)) + except ValidationError: + return None + + +def baseline_request( + kwargs: Mapping[str, object], + caller: Mapping[str, JsonValue], + deployment: Mapping[str, object], + *, + include_extra_body: bool = True, +) -> Mapping[str, object] | None: + snapshot: Final = capture_baseline_parameters(deployment) + if snapshot is None: + return None + configured: Final = { + **_parameters(snapshot), + **(_parameters(snapshot.get("extra_body")) if include_extra_body else {}), + } + requested: Final = {**_parameters(caller), **(_parameters(caller.get("extra_body")) if include_extra_body else {})} + configured_tools: Final = configured.get("tools") or [] + caller_tools: Final = requested.get("tools") or [] + merged_tools: Final = ( + {"tools": [*configured_tools, *caller_tools]} + if (configured_tools or caller_tools) and isinstance(configured_tools, list) and isinstance(caller_tools, list) + else {} + ) + return MappingProxyType( + { + **{key: value for key, value in kwargs.items() if key not in (*BASELINE_PARAMETERS, "extra_body")}, + **configured, + **requested, + **merged_tools, + **( + {"extra_body": caller.get("extra_body", snapshot.get("extra_body"))} + if not include_extra_body and ("extra_body" in caller or "extra_body" in snapshot) + else {} + ), + } + ) diff --git a/litellm/types/router.py b/litellm/types/router.py index 66a5b3540f9..a66c4571b39 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -5,7 +5,7 @@ litellm.Router Types - includes RouterConfig, UpdateRouterConfig, ModelInfo etc import datetime import enum from collections.abc import Container, Mapping, Sequence -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import ( TYPE_CHECKING, Annotated, @@ -21,7 +21,7 @@ from typing import ( from zoneinfo import ZoneInfo, ZoneInfoNotFoundError import httpx -from pydantic import ConfigDict, Field, field_validator, model_validator +from pydantic import ConfigDict, Field, JsonValue, field_validator, model_validator from typing_extensions import Protocol, ReadOnly, Required, TypedDict, runtime_checkable from litellm._logging import verbose_logger @@ -1223,6 +1223,7 @@ class BaselineRouteStamp: router_name: str baseline_model: str baseline_deployment_id: str + request_parameters: Mapping[str, JsonValue] | None = field(default=None, repr=False) @dataclass(frozen=True, slots=True) diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index 0c53137cfab..7166891bb1a 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -2,6 +2,7 @@ import ast import os IGNORE_FUNCTIONS = [ + "_json_cost", # bounded at depth 32 and consumed under byte/node limits. "_format_type", "remove_additional_properties", "remove_strict_from_schema", diff --git a/tests/proxy_behavior/spend/test_baseline_accounting.py b/tests/proxy_behavior/spend/test_baseline_accounting.py index dbaf32d579f..4298b85cab4 100644 --- a/tests/proxy_behavior/spend/test_baseline_accounting.py +++ b/tests/proxy_behavior/spend/test_baseline_accounting.py @@ -3,7 +3,8 @@ import json import uuid from collections.abc import AsyncIterator, Callable from contextlib import asynccontextmanager -from datetime import datetime, timezone +from dataclasses import replace +from datetime import datetime, timedelta, timezone from typing import Final import pytest @@ -174,6 +175,10 @@ async def test_commit_ack_loss_and_concurrent_duplicate_delivery_are_idempotent( assert await _store(db, after_commit=True).append(event) == "unavailable" store: Final = _store(db) assert set(await asyncio.gather(*(store.append(event) for _ in range(4)))) == {"recorded"} + duplicate: Final = event.model_copy(update={ + "turn": replace(event.turn, turn_at=event.turn.turn_at + timedelta(seconds=1)) + }) + assert await store.append(duplicate) == "recorded" await _log(db, other) assert await store.append(other) == "recorded" if not attributed: diff --git a/tests/unit/litellm_core_utils/test_logging_worker.py b/tests/unit/litellm_core_utils/test_logging_worker.py index 5d4c9e65d9b..f48495d1062 100644 --- a/tests/unit/litellm_core_utils/test_logging_worker.py +++ b/tests/unit/litellm_core_utils/test_logging_worker.py @@ -6,6 +6,7 @@ import asyncio import contextvars import io import logging +from typing import Final from unittest.mock import AsyncMock, patch import pytest @@ -14,6 +15,61 @@ from litellm.constants import LOGGING_WORKER_AGGRESSIVE_CLEAR_COOLDOWN_SECONDS from litellm.litellm_core_utils.logging_worker import LoggingWorker +@pytest.mark.asyncio +@pytest.mark.parametrize("dispatch", ("worker", "flush", "extracted")) +async def test_optional_work_budget_preserves_callback_context_and_reserves_logging_time(dispatch: str) -> None: + from litellm.litellm_core_utils.logging_worker import optional_callback_budget + + worker: Final = LoggingWorker(timeout=1.0) + identity: Final = contextvars.ContextVar("test_callback_identity", default="outside") + results: Final[asyncio.Queue[tuple[str, float]]] = asyncio.Queue() + + async def callback() -> None: + results.put_nowait((identity.get(), optional_callback_budget(3.0))) + + token: Final = identity.set("request") + worker._ensure_queue() + worker.enqueue(callback()) + identity.reset(token) + try: + if dispatch == "worker": + worker.start() + elif dispatch == "flush": + await worker.flush() + else: + assert worker._queue is not None + await worker._process_single_task(worker._queue.get_nowait()) + restored_identity, budget = await asyncio.wait_for(results.get(), timeout=2) + assert restored_identity == "request" + assert 0 < budget <= worker.timeout / 4 + assert identity.get() == "outside" + assert optional_callback_budget(3.0) == 3.0 + finally: + await worker.stop() + + +def test_exit_flush_bounds_optional_work_and_restores_callers_budget() -> None: + from queue import SimpleQueue + + from litellm.litellm_core_utils.logging_worker import optional_callback_budget + + worker: Final = LoggingWorker(timeout=1.0) + observed: Final[SimpleQueue[float]] = SimpleQueue() + + async def callback() -> None: + observed.put(optional_callback_budget(3.0)) + + async def enqueue() -> None: + worker._ensure_queue() + worker.enqueue(callback()) + + asyncio.run(enqueue()) + worker._flush_on_exit() + assert observed.qsize() == 1 + assert 0 < observed.get_nowait() <= worker.timeout / 4 + assert optional_callback_budget(3.0) == 3.0 + + class _RecordCollector(logging.Handler): """Captures emitted log records so a test can assert on real logging output (level, message args, traceback) instead of patching the logger object.""" @@ -205,7 +261,9 @@ class TestLoggingWorker: asyncio.run(log_on_second_loop()) first_loop.run_until_complete(asyncio.sleep(0.1)) failures = [ - task.exception() for task in first_loop_tasks if task.done() and not task.cancelled() and task.exception() + task.exception() + for task in first_loop_tasks + if task.done() and not task.cancelled() and task.exception() ] finally: first_loop.close() diff --git a/tests/unit/litellm_core_utils/test_redact_messages.py b/tests/unit/litellm_core_utils/test_redact_messages.py index 3bb6b379873..6fbcba170f1 100644 --- a/tests/unit/litellm_core_utils/test_redact_messages.py +++ b/tests/unit/litellm_core_utils/test_redact_messages.py @@ -5,14 +5,26 @@ Covers the proxy flow where headers arrive in litellm_params["metadata"]["header but litellm_params["litellm_metadata"] is None. """ -import asyncio, httpx, importlib, json, os, pytest_asyncio, threading +import asyncio +import importlib +import json +import os +import threading +from collections.abc import AsyncIterator, Mapping +from datetime import datetime +from types import MappingProxyType, SimpleNamespace from typing import Final, Optional, Union -from types import SimpleNamespace +from unittest.mock import patch +import httpx import pytest +import pytest_asyncio +from pydantic import JsonValue import litellm +from litellm.constants import LOGGING_WORKER_MAX_TIME_PER_COROUTINE from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.litellm_core_utils.redact_messages import ( _redact_responses_api_output, perform_redaction, @@ -21,18 +33,14 @@ from litellm.litellm_core_utils.redact_messages import ( should_redact_message_logging, ) from litellm.responses.main import mock_responses_api_response -from collections.abc import AsyncIterator -from datetime import datetime -from litellm.constants import LOGGING_WORKER_MAX_TIME_PER_COROUTINE -from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER -from litellm.types.utils import( +from litellm.types.router import BaselineRouteStamp +from litellm.types.utils import ( ModelResponse, ResponsesAPIResponse, StandardLoggingPayload, TextCompletionResponse, ) from tests._vcr_conftest_common import install_live_call_probe, record_vcr_outcome -from unittest.mock import patch @pytest.fixture(autouse=True) @@ -113,9 +121,7 @@ class TestShouldRedactMessageLogging: def test_enable_redaction_via_header_in_litellm_metadata(self): """Headers inside litellm_metadata (SDK direct call) should work.""" details = _make_model_call_details( - litellm_metadata={ - "headers": {"x-litellm-enable-message-redaction": "true"} - }, + litellm_metadata={"headers": {"x-litellm-enable-message-redaction": "true"}}, ) assert should_redact_message_logging(details) is True @@ -217,21 +223,15 @@ class TestPerformRedaction: redacted = perform_redaction(details, result) - assert details["messages"] == [ - {"role": "user", "content": "redacted-by-litellm"} - ] + assert details["messages"] == [{"role": "user", "content": "redacted-by-litellm"}] assert details["prompt"] == "" assert details["input"] == "" logged_response = details["standard_logging_object"]["response"] assert logged_response["usage"] == {"total_tokens": 1} assert logged_response["output"][0]["text"] == "redacted-by-litellm" - assert logged_response["output"][1]["content"][0]["text"] == ( - "redacted-by-litellm" - ) - assert logged_response["output"][2]["summary"][0]["text"] == ( - "redacted-by-litellm" - ) + assert logged_response["output"][1]["content"][0]["text"] == ("redacted-by-litellm") + assert logged_response["output"][2]["summary"][0]["text"] == ("redacted-by-litellm") assert redacted["usage"] == {"total_tokens": 1} assert redacted["output"][0]["text"] == "redacted-by-litellm" @@ -444,9 +444,7 @@ class TestPerformRedaction: tool_call = redacted.choices[0].message.tool_calls[0] assert tool_call.function.arguments == "redacted-by-litellm" assert tool_call.function.name == "get_weather" - assert result.choices[0].message.tool_calls[0].function.arguments == ( - '{"city": "sensitive-city"}' - ) + assert result.choices[0].message.tool_calls[0].function.arguments == ('{"city": "sensitive-city"}') def test_redacts_tool_call_arguments_on_streaming_response_object(self): """Reproduces the Stream=True path where tool calls arrive as deltas.""" @@ -714,12 +712,8 @@ class TestPerformRedaction: } } ], - "vertex_ai_grounding_metadata": [ - {"webSearchQueries": ["sensitive search term"]} - ], - "vertex_ai_url_context_metadata": [ - {"urlMetadata": [{"retrievedUrl": "https://example.com"}]} - ], + "vertex_ai_grounding_metadata": [{"webSearchQueries": ["sensitive search term"]}], + "vertex_ai_url_context_metadata": [{"urlMetadata": [{"retrievedUrl": "https://example.com"}]}], }, } } @@ -749,9 +743,7 @@ class TestPerformRedaction: "vertex_ai_grounding_metadata", [{"webSearchQueries": ["sensitive search term"]}], ) - response._hidden_params["vertex_ai_grounding_metadata"] = [ - {"webSearchQueries": ["sensitive search term"]} - ] + response._hidden_params["vertex_ai_grounding_metadata"] = [{"webSearchQueries": ["sensitive search term"]}] details = { "stream": True, @@ -772,12 +764,8 @@ class TestPerformRedaction: "metadata": { "hidden_params": { "response_cost": 0.01, - "vertex_ai_grounding_metadata": [ - {"webSearchQueries": ["sensitive search term"]} - ], - "vertex_ai_url_context_metadata": [ - {"urlMetadata": [{"retrievedUrl": "https://example.com"}]} - ], + "vertex_ai_grounding_metadata": [{"webSearchQueries": ["sensitive search term"]}], + "vertex_ai_url_context_metadata": [{"urlMetadata": [{"retrievedUrl": "https://example.com"}]}], "vertex_ai_safety_ratings": [{"category": "HARM"}], "vertex_ai_citation_metadata": [{"citations": ["source"]}], } @@ -797,11 +785,7 @@ class TestPerformRedaction: def test_redact_async_complete_streaming_response(self): """Test that async_complete_streaming_response is properly redacted.""" response_obj = litellm.ModelResponse( - choices=[ - litellm.Choices( - message=litellm.Message(content="secret content", role="assistant") - ) - ] + choices=[litellm.Choices(message=litellm.Message(content="secret content", role="assistant"))] ) model_call_details = { @@ -820,11 +804,7 @@ class TestPerformRedaction: def test_redact_complete_streaming_response(self): """Test that complete_streaming_response is properly redacted.""" response_obj = litellm.ModelResponse( - choices=[ - litellm.Choices( - message=litellm.Message(content="secret content", role="assistant") - ) - ] + choices=[litellm.Choices(message=litellm.Message(content="secret content", role="assistant"))] ) model_call_details = { @@ -842,11 +822,7 @@ class TestPerformRedaction: def test_streaming_responses_untouched_when_disabled(self): response_obj = litellm.ModelResponse( - choices=[ - litellm.Choices( - message=litellm.Message(content="secret content", role="assistant") - ) - ] + choices=[litellm.Choices(message=litellm.Message(content="secret content", role="assistant"))] ) model_call_details = { @@ -909,11 +885,7 @@ class TestPerformRedaction: class TestRedactStreamingResponsesForCustomLogger: def _model_call_details(self): response_obj = litellm.ModelResponse( - choices=[ - litellm.Choices( - message=litellm.Message(content="secret content", role="assistant") - ) - ] + choices=[litellm.Choices(message=litellm.Message(content="secret content", role="assistant"))] ) return { "stream": True, @@ -947,7 +919,10 @@ class TestRedactStreamingResponsesForCustomLogger: @pytest.mark.parametrize("callback_only", [False, True]) def test_classifier_audit_redaction_removes_both_fields_and_source_carrier(callback_only: bool) -> None: - audit: Final = {"classifier_input": {"system": "private rubric"}, "originating_request_masked": {"input": "private source"}} + audit: Final = { + "classifier_input": {"system": "private rubric"}, + "originating_request_masked": {"input": "private source"}, + } standard_payload: Final = { **audit, "messages": [{"role": "user", "content": "private prompt"}], @@ -956,7 +931,9 @@ def test_classifier_audit_redaction_removes_both_fields_and_source_carrier(callb } details: Final = { "standard_logging_object": standard_payload, - "litellm_params": {"proxy_server_request": {"body": {}, "originating_request_masked": audit["originating_request_masked"]}}, + "litellm_params": { + "proxy_server_request": {"body": {}, "originating_request_masked": audit["originating_request_masked"]} + }, } logger: Final = CustomLogger() logger.turn_off_message_logging = True @@ -966,7 +943,10 @@ def test_classifier_audit_redaction_removes_both_fields_and_source_carrier(callb assert "originating_request_masked" not in redacted["standard_logging_object"] assert "originating_request_masked" not in redacted["litellm_params"]["proxy_server_request"] assert details["standard_logging_object"]["classifier_input"] == audit["classifier_input"] - assert details["litellm_params"]["proxy_server_request"]["originating_request_masked"] == audit["originating_request_masked"] + assert ( + details["litellm_params"]["proxy_server_request"]["originating_request_masked"] + == audit["originating_request_masked"] + ) else: perform_redaction(details, result=None) assert "classifier_input" not in details["standard_logging_object"] @@ -981,7 +961,9 @@ def test_classifier_audit_redaction_removes_both_fields_and_source_carrier(callb @pytest.mark.parametrize("excluded", [False, True]) def test_classifier_callback_redaction_preserves_exclusions(monkeypatch: pytest.MonkeyPatch, excluded: bool) -> None: - monkeypatch.setattr(litellm, "standard_logging_payload_excluded_fields", ["messages", "response"] if excluded else []) + monkeypatch.setattr( + litellm, "standard_logging_payload_excluded_fields", ["messages", "response"] if excluded else [] + ) payload: Final = { "classifier_input": {"system": "private rubric"}, "originating_request_masked": {"input": "private source"}, @@ -991,7 +973,9 @@ def test_classifier_callback_redaction_preserves_exclusions(monkeypatch: pytest. } logger: Final = CustomLogger() logger.turn_off_message_logging = True - redacted: Final = logger.redact_standard_logging_payload_from_model_call_details({"standard_logging_object": payload}) + redacted: Final = logger.redact_standard_logging_payload_from_model_call_details( + {"standard_logging_object": payload} + ) stored: Final = redacted["standard_logging_object"] assert "classifier_input" not in stored assert "originating_request_masked" not in stored @@ -1017,16 +1001,18 @@ class _SelfRedactingLogger(CustomLogger): @pytest.mark.parametrize("logger", [CustomLogger(), _SelfRedactingLogger()], ids=["default", "redacts_itself"]) -def test_field_exclusion_alone_leaves_messages_and_responses_intact(monkeypatch: pytest.MonkeyPatch, logger: CustomLogger) -> None: +def test_field_exclusion_alone_leaves_messages_and_responses_intact( + monkeypatch: pytest.MonkeyPatch, logger: CustomLogger +) -> None: monkeypatch.setattr(litellm, "standard_logging_payload_excluded_fields", ["model"]) payload: Final = { "messages": [{"role": "user", "content": "private prompt"}], "response": {"choices": [{"message": {"content": "private answer"}}]}, "model": "classifier", } - stored: Final = logger.redact_standard_logging_payload_from_model_call_details({"standard_logging_object": payload})[ - "standard_logging_object" - ] + stored: Final = logger.redact_standard_logging_payload_from_model_call_details( + {"standard_logging_object": payload} + )["standard_logging_object"] assert stored == {"messages": payload["messages"], "response": payload["response"]} @@ -1038,9 +1024,9 @@ def test_a_callback_that_redacts_itself_keeps_its_messages_but_not_the_classifie } logger: Final = _SelfRedactingLogger() logger.turn_off_message_logging = True - stored: Final = logger.redact_standard_logging_payload_from_model_call_details({"standard_logging_object": payload})[ - "standard_logging_object" - ] + stored: Final = logger.redact_standard_logging_payload_from_model_call_details( + {"standard_logging_object": payload} + )["standard_logging_object"] assert "classifier_input" not in stored assert stored["messages"] == payload["messages"] assert stored["response"] == payload["response"] @@ -1054,19 +1040,54 @@ def test_perform_redaction_drops_the_served_output_texts_from_the_callback_kwarg assert SERVED_OUTPUT_TEXTS_KEY not in details +@pytest.mark.parametrize("callback_only", (False, True)) +@pytest.mark.parametrize("with_standard_payload", (False, True)) +def test_baseline_snapshots_are_redacted_without_mutating_request_state( + callback_only: bool, with_standard_payload: bool +) -> None: + snapshot: Final[Mapping[str, JsonValue]] = MappingProxyType({"system": "private system"}) + route: Final = BaselineRouteStamp("router", "baseline", "deployment", snapshot) + metadata: Final = {"_autorouter_baseline_route": route, "session_id": "session"} + params: Final = {"metadata": metadata, "litellm_metadata": metadata} + details: Final = { + "litellm_params": params, + **({"standard_logging_object": {"model": "model"}} if with_standard_payload else {}), + } + logger: Final = CustomLogger() + logger.turn_off_message_logging = True + if not callback_only: + perform_redaction(details, None) + redacted: Final = ( + logger.redact_standard_logging_payload_from_model_call_details(details) if callback_only else details + ) + expected: Final = BaselineRouteStamp(route.router_name, route.baseline_model, route.baseline_deployment_id) + assert redacted["litellm_params"] == { + key: {"_autorouter_baseline_route": expected, "session_id": "session"} + for key in ("metadata", "litellm_metadata") + } + assert route.request_parameters is snapshot + assert params["metadata"]["_autorouter_baseline_route"] is route + assert params["litellm_metadata"]["_autorouter_baseline_route"] is route + if callback_only: + assert details["litellm_params"] is params + + @pytest.fixture() def _vcr_outcome_gate(request, vcr): install_live_call_probe(request, vcr) yield record_vcr_outcome(request, vcr) + @pytest_asyncio.fixture(loop_scope="function") async def drain_logging_worker(isolate_litellm_state: None) -> AsyncIterator[None]: yield await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=LOGGING_WORKER_DRAIN_TIMEOUT_SECONDS) + LOGGING_WORKER_DRAIN_TIMEOUT_SECONDS: Final = LOGGING_WORKER_MAX_TIME_PER_COROUTINE + 5.0 + @pytest.fixture(scope="function") def isolate_litellm_state(): """ @@ -1104,6 +1125,7 @@ def isolate_litellm_state(): if attr in _DEFAULTS: setattr(litellm, attr, _DEFAULTS[attr]) + _LIST_ATTRS = ( "callbacks", "success_callback", @@ -1131,6 +1153,7 @@ _SCALAR_ATTRS = ( _DEFAULTS: dict = {} + @pytest.fixture(scope="module") def setup_and_teardown(): """ @@ -1153,6 +1176,7 @@ def setup_and_teardown(): litellm.in_memory_llm_clients_cache.flush_cache() yield + class TestCustomLogger(CustomLogger): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) @@ -1164,6 +1188,7 @@ class TestCustomLogger(CustomLogger): self.logged_standard_logging_payload = standard_logging_payload self.response_obj = response_obj + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_global_redaction_on(): @@ -1187,6 +1212,7 @@ async def test_global_redaction_on(): json.dumps(standard_logging_payload, indent=2), ) + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.parametrize( "dynamic_turn_off, expect_redacted", @@ -1213,6 +1239,7 @@ async def test_dynamic_turn_off_message_logging_overrides_global_on(dynamic_turn assert standard_logging_payload["response"]["choices"][0]["message"]["content"] == expected_response_content assert standard_logging_payload["messages"][0]["content"] == expected_message_content + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.parametrize( "dynamic_turn_off, expect_redacted", @@ -1239,6 +1266,7 @@ async def test_dynamic_turn_off_message_logging_overrides_global_off(dynamic_tur assert standard_logging_payload["response"]["choices"][0]["message"]["content"] == expected_response_content assert standard_logging_payload["messages"][0]["content"] == expected_message_content + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_redaction_with_custom_logger_streaming(): @@ -1284,6 +1312,7 @@ async def test_redaction_with_custom_logger_streaming(): finally: litellm.turn_off_message_logging = False + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_streaming_redaction_scoped_to_opted_out_logger(): @@ -1311,6 +1340,7 @@ async def test_streaming_redaction_scoped_to_opted_out_logger(): finally: litellm.callbacks = [] + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_redaction_responses_api(): @@ -1355,6 +1385,7 @@ async def test_redaction_responses_api(): json.dumps(standard_logging_payload, indent=2), ) + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_redaction_responses_api_stream(): @@ -1430,6 +1461,7 @@ async def test_redaction_responses_api_stream(): json.dumps(standard_logging_payload, indent=2), ) + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_redaction_responses_api_with_reasoning_summary(): @@ -1490,6 +1522,7 @@ async def test_redaction_responses_api_with_reasoning_summary(): assert model_call_details["messages"][0]["content"] == "redacted-by-litellm", "Input messages should be redacted" + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_redaction_with_coroutine_objects(): @@ -1535,6 +1568,7 @@ async def test_redaction_with_coroutine_objects(): result = perform_redaction({}, mock_iter) assert result == {"text": "redacted-by-litellm"} + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_redaction_with_streaming_response(): @@ -1570,6 +1604,7 @@ async def test_redaction_with_streaming_response(): json.dumps(standard_logging_payload, indent=2), ) + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_disable_redaction_header_responses_api(): @@ -1605,6 +1640,7 @@ async def test_disable_redaction_header_responses_api(): assert response["output"][0]["content"][0]["text"] == "This is a test response" assert standard_logging_payload["messages"][0]["content"] == "hi" + @pytest.mark.usefixtures("_vcr_outcome_gate", "drain_logging_worker", "isolate_litellm_state", "setup_and_teardown") @pytest.mark.asyncio async def test_redaction_with_metadata_completion_api(): diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py index c3d4dba7376..ac55e8e9350 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py @@ -684,12 +684,12 @@ def _empty_block_msgs(): def test_handler_strips_when_no_presanitized_flag(): """Sync entry point (no async wrapper): handler must still sanitize.""" - from litellm.llms.anthropic.pass_through.messages import handler + from litellm.llms.anthropic.pass_through.messages import handler, utils with patch.object( - handler, + utils, "strip_empty_content_blocks_from_anthropic_messages", - wraps=handler.strip_empty_content_blocks_from_anthropic_messages, + wraps=utils.strip_empty_content_blocks_from_anthropic_messages, ) as spy: result = handler.anthropic_messages_handler( max_tokens=10, @@ -704,12 +704,12 @@ def test_handler_strips_when_no_presanitized_flag(): def test_handler_skips_strip_when_presanitized(): """Async wrapper already sanitized -> handler must NOT rescan.""" - from litellm.llms.anthropic.pass_through.messages import handler + from litellm.llms.anthropic.pass_through.messages import handler, utils with patch.object( - handler, + utils, "strip_empty_content_blocks_from_anthropic_messages", - wraps=handler.strip_empty_content_blocks_from_anthropic_messages, + wraps=utils.strip_empty_content_blocks_from_anthropic_messages, ) as spy: result = handler.anthropic_messages_handler( max_tokens=10, @@ -809,7 +809,7 @@ def test_presanitized_flag_not_leaked_to_provider_params(): @pytest.mark.asyncio async def test_async_wrapper_sets_presanitized_and_sanitizes_once(): """End-to-end: wrapper sanitizes (once) AND signals the handler to skip.""" - from litellm.llms.anthropic.pass_through.messages import handler + from litellm.llms.anthropic.pass_through.messages import handler, utils captured = {} @@ -825,9 +825,9 @@ async def test_async_wrapper_sets_presanitized_and_sanitizes_once(): patch.object(handler, "anthropic_messages_handler", side_effect=fake_handler), patch("asyncio.get_event_loop", return_value=fake_loop), patch.object( - handler, + utils, "strip_empty_content_blocks_from_anthropic_messages", - wraps=handler.strip_empty_content_blocks_from_anthropic_messages, + wraps=utils.strip_empty_content_blocks_from_anthropic_messages, ) as spy, ): await handler.anthropic_messages( diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_request_optional_param_utils.py b/tests/unit/llms/anthropic/pass_through/messages/test_request_optional_param_utils.py index dd744ca66a1..7a4f3aa5f90 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_request_optional_param_utils.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_request_optional_param_utils.py @@ -11,7 +11,7 @@ import pytest import litellm from litellm.llms.anthropic.pass_through.messages.utils import ( AnthropicMessagesRequestUtils, - _anthropic_messages_optional_param_keys, + anthropic_messages_optional_param_keys, ) @@ -30,16 +30,16 @@ def test_optional_param_filtering_unchanged(): def test_valid_keys_are_memoized(): - _anthropic_messages_optional_param_keys.cache_clear() - first = _anthropic_messages_optional_param_keys() + anthropic_messages_optional_param_keys.cache_clear() + first = anthropic_messages_optional_param_keys() for _ in range(50): AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param({"temperature": 0.1}) - info = _anthropic_messages_optional_param_keys.cache_info() + info = anthropic_messages_optional_param_keys.cache_info() # Resolved exactly once despite many calls. assert info.misses == 1 assert info.hits >= 50 # Stable identity (frozenset) returned each call. - assert _anthropic_messages_optional_param_keys() is first + assert anthropic_messages_optional_param_keys() is first assert isinstance(first, frozenset) assert "temperature" in first and "tools" in first diff --git a/tests/unit/proxy/hooks/test_autorouter_baseline_cache.py b/tests/unit/proxy/hooks/test_autorouter_baseline_cache.py index 0f4fd2ff5cb..1be02d2f269 100644 --- a/tests/unit/proxy/hooks/test_autorouter_baseline_cache.py +++ b/tests/unit/proxy/hooks/test_autorouter_baseline_cache.py @@ -3,6 +3,7 @@ import json from collections.abc import AsyncIterator, Callable, Generator, Mapping from contextlib import contextmanager from datetime import datetime +from itertools import product from types import MappingProxyType from typing import Final, cast from uuid import uuid4 @@ -42,7 +43,7 @@ _MESSAGES_JSON: Final = """[{"role":"user","content":[ _MODELS: Final = _MESSAGES.validate_json("""[ {"model_name":"test-router","litellm_params":{"model":"auto_router/complexity_router", "complexity_router_config":{"tiers":{"SIMPLE":"sonnet","MEDIUM":"sonnet","COMPLEX":"sonnet", - "REASONING":"opus"},"session_affinity":false, + "REASONING":{"model_name":"opus","litellm_params":{"max_tokens":16}}},"session_affinity":false, "keyword_tier_rules":[{"keywords":["USE_OPUS"],"tier":"REASONING"}]}}}, {"model_name":"sonnet","litellm_params":{"model":"anthropic/claude-sonnet-5","api_key":"test-selected"}, "model_info":{"id":"selected"}}, @@ -82,7 +83,9 @@ class _CallContext(TypedDict): def _kwargs(logging_obj: Logging, trusted: bool = True, *, explicit_logging: bool = True) -> _CallContext: - context: Final = _OBJECTS.validate_json('{"litellm_metadata":{"user_api_key_hash":"test-caller-hash"}}') + context: Final = _OBJECTS.validate_json( + '{"max_tokens":16,"litellm_metadata":{"user_api_key_hash":"test-caller-hash"}}' + ) Router._record_routing_decision( # pyright: ignore[reportUnknownMemberType, reportPrivateUsage] # production trusted stamp owner context, StandardLoggingRoutingDecision( @@ -133,7 +136,10 @@ def _upstream(request: httpx.Request) -> httpx.Response: assert isinstance(model, str) stream: Final = body.get("stream") is True content: Final = b"".join(_sse(model=model)) if stream else json.dumps(_message(True, model)).encode() - return httpx.Response(200, content=content, request=request, + return httpx.Response( + 200, + content=content, + request=request, headers=MappingProxyType({"content-type": "text/event-stream" if stream else "application/json"}), ) @@ -182,6 +188,7 @@ async def _call( stream: Final = cast(AsyncIterator[object], response) # cast-ok: iterator checked; all items satisfy object assert tuple([chunk async for chunk in stream]) + class _Capture(CustomLogger): def __init__(self, call_id: str) -> None: self.call_id: Final = call_id @@ -199,9 +206,20 @@ class _Capture(CustomLogger): class _Rig: - def __init__(self, monkeypatch: pytest.MonkeyPatch, *, retries: int = 0, count: TokenCounter = _count) -> None: - self.router: Final = Router(model_list=_MODELS, num_retries=retries, - retry_policy=RetryPolicy(RateLimitErrorRetries=retries), disable_cooldowns=True) + def __init__( + self, + monkeypatch: pytest.MonkeyPatch, + *, + retries: int = 0, + count: TokenCounter = _count, + models: list[dict[str, JsonValue]] = _MODELS, + ) -> None: + self.router: Final = Router( + model_list=models, + num_retries=retries, + retry_policy=RetryPolicy(RateLimitErrorRetries=retries), + disable_cooldowns=True, + ) def router() -> Router: return self.router @@ -218,9 +236,16 @@ class _Rig: monkeypatch.setattr(litellm, "_async_success_callback", [self.capture]) def logging(self, stream: bool = False) -> Logging: - return Logging(model="anthropic/claude-sonnet-5", messages=_MESSAGES.validate_json(_MESSAGES_JSON), - stream=stream, call_type=CallTypes.anthropic_messages.value, start_time=datetime.now(), - litellm_call_id=self.call_id, function_id=self.call_id, kwargs={"litellm_session_id":"baseline-session"}) + return Logging( + model="anthropic/claude-sonnet-5", + messages=_MESSAGES.validate_json(_MESSAGES_JSON), + stream=stream, + call_type=CallTypes.anthropic_messages.value, + start_time=datetime.now(), + litellm_call_id=self.call_id, + function_id=self.call_id, + kwargs={"litellm_session_id": "baseline-session"}, + ) def _observation(payload: Mapping[str, object]) -> CapturedBaselineObservation: @@ -232,7 +257,9 @@ def _observation(payload: Mapping[str, object]) -> CapturedBaselineObservation: @pytest.mark.parametrize("stream,baseline", ((False, False), (True, False), (False, True), (True, True))) async def test_native_logging_captures_usage_without_publishing_hypothetical_savings( - monkeypatch: pytest.MonkeyPatch, stream: bool, baseline: bool, + monkeypatch: pytest.MonkeyPatch, + stream: bool, + baseline: bool, ) -> None: rig: Final = _Rig(monkeypatch) messages: Final = _MESSAGES_JSON.replace("question", "question USE_OPUS") if baseline else _MESSAGES_JSON @@ -283,11 +310,14 @@ async def test_caller_cannot_forge_an_observation_scope(monkeypatch: pytest.Monk assert payload["autorouter_savings"] is None -@pytest.mark.parametrize("model,key,endpoint", ( - ("claude-sonnet-5", "test-first", None), - ("claude-opus-5", "test-second", None), - ("claude-opus-5", "test-first", "https://example.test"), -)) +@pytest.mark.parametrize( + "model,key,endpoint", + ( + ("claude-sonnet-5", "test-first", None), + ("claude-opus-5", "test-second", None), + ("claude-opus-5", "test-first", "https://example.test"), + ), +) async def test_count_memo_is_scoped_to_provider_recipient(model: str, key: str, endpoint: str | None) -> None: counts: Final = iter((5000, 6000)) @@ -304,7 +334,8 @@ async def test_count_memo_is_scoped_to_provider_recipient(model: str, key: str, @pytest.mark.parametrize("stream", (False, True)) async def test_provider_counting_does_not_hold_the_inference_response( - monkeypatch: pytest.MonkeyPatch, stream: bool, + monkeypatch: pytest.MonkeyPatch, + stream: bool, ) -> None: counting: Final = asyncio.Event() release: Final = asyncio.Event() @@ -324,3 +355,744 @@ async def test_provider_counting_does_not_hold_the_inference_response( assert _observation(await rig.capture.payload()).observation.plan is not None finally: release.set() + + +@pytest.mark.parametrize("baseline_effort", (None, "medium")) +@pytest.mark.parametrize( + "automatic_system, caching", + ( + (None, "explicit"), + ("stable system", "request"), + ([{"type": "text", "text": "stable system"}], "request"), + ("stable system", "global"), + ([{"type": "text", "text": "stable system"}], "configured"), + ), +) +async def test_native_tier_switch_uses_baseline_settings_and_preserves_history( + monkeypatch: pytest.MonkeyPatch, + baseline_effort: str | None, + automatic_system: str | list[dict[str, str]] | None, + caching: str, +) -> None: + from litellm.proxy.spend_tracking.baseline_accounting import BaselineHistory, advance_baseline_history + + models: Final = _MESSAGES.validate_python( + [ + { + "model_name": "test-router", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": { + "tiers": { + "SIMPLE": {"model_name": "sonnet", "litellm_params": {"reasoning_effort": "low"}}, + "MEDIUM": {"model_name": "sonnet", "litellm_params": {"reasoning_effort": "low"}}, + "COMPLEX": {"model_name": "sonnet", "litellm_params": {"reasoning_effort": "high"}}, + "REASONING": "opus", + }, + "session_affinity": False, + "keyword_tier_rules": [{"keywords": ["ESCALATE"], "tier": "COMPLEX"}], + }, + }, + }, + _MODELS[1], + { + "model_name": "opus", + "model_info": {"id": "baseline"}, + "litellm_params": { + "model": "anthropic/claude-opus-5", + "api_key": "test-selected", + **({"reasoning_effort": baseline_effort} if baseline_effort else {}), + }, + }, + ] + ) + rig: Final = _Rig(monkeypatch, models=models) + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", caching == "global") + controls: Final = ( + { + "cache_control_injection_points": [ + {"location": "message", "role": "system", "control": {"type": "ephemeral", "ttl": "1h"}}, + {"location": "message", "index": -1, "control": {"type": "ephemeral", "ttl": "1h"}}, + ] + } + if caching == "configured" + else {"enable_prompt_caching": caching == "request"} + ) + captures: Final[asyncio.Queue[CapturedBaselineObservation]] = asyncio.Queue() + with _transport(_upstream) as route: + for suffix in ("", " ESCALATE"): + log: Final = rig.logging() + await rig.router.anthropic_messages( + model="test-router", + max_tokens=4096, + messages=( + [{"role": "user", "content": "question" + suffix}] + if automatic_system is not None + else _MESSAGES.validate_json(_MESSAGES_JSON.replace("question", "question" + suffix)) + ), + system=automatic_system, + **controls, + litellm_logging_obj=log, + litellm_call_id=rig.call_id, + litellm_metadata={"user_api_key_hash": "test-caller-hash"}, + litellm_session_id="native-tiers", + ) + captures.put_nowait(_observation(await rig.capture.payload())) + first_wire, second_wire = (_JSON_OBJECT.validate_json(call.request.content) for call in route.calls) + assert (first_wire.get("thinking"), first_wire.get("output_config")) != ( + second_wire.get("thinking"), + second_wire.get("output_config"), + ) + first, second = (captures.get_nowait() for _ in range(2)) + assert first.scope == second.scope + assert first.observation.plan is not None and second.observation.plan is not None + assert first.observation.plan.breakpoints[0] == second.observation.plan.breakpoints[0] + assert len(first.observation.plan.breakpoints) == (2 if automatic_system is not None else 1) + history, _ = advance_baseline_history( + BaselineHistory(first_at=0.0), + (first.observation.model_copy(update={"request_id": "first", "started_at": 10000.0, "available_at": 10001.0}),), + ) + _, result = advance_baseline_history( + history, + ( + second.observation.model_copy( + update={"request_id": "second", "started_at": 10020.0, "available_at": 10021.0} + ), + ), + ) + assert result[0].usage is not None and result[0].usage.prompt_tokens_details.cached_tokens == 5000 + + +@pytest.mark.parametrize("call_type", (CallTypes.acompletion, CallTypes.aresponses, CallTypes.anthropic_messages)) +async def test_plain_requests_do_not_initialize_or_warn( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, + call_type: CallTypes, +) -> None: + rig: Final = _Rig(monkeypatch) + logging: Final = rig.logging() + await rig.hook.async_pre_call_deployment_hook( + { + "litellm_logging_obj": logging, + "litellm_metadata": {"session_id": "ordinary"}, + }, + call_type, + ) + assert logging.baseline_cache_context is None + assert "baseline observation could not be initialized" not in caplog.text + assert not rig.hook.counts + + +async def test_plain_fallback_invalidates_existing_autorouter_capture(monkeypatch: pytest.MonkeyPatch) -> None: + rig: Final = _Rig(monkeypatch) + logging: Final = rig.logging() + await rig.hook.async_pre_call_deployment_hook(_kwargs(logging), CallTypes.anthropic_messages) + assert logging.baseline_cache_context is not None + await rig.hook.async_pre_call_deployment_hook({"litellm_logging_obj": logging}, CallTypes.anthropic_messages) + assert logging.baseline_observation is not None + assert logging.baseline_observation.observation.reason == "retried_request" + + +async def test_native_count_finishing_after_quarter_worker_budget_keeps_plan_and_spend( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.litellm_core_utils import logging_worker + from litellm.litellm_core_utils.logging_worker import LoggingWorker + + release: Final = asyncio.Event() + + async def count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int: + if not release.is_set(): + asyncio.get_running_loop().call_later(2.3, release.set) + await release.wait() + return await _count(model, api_key, body) + + worker: Final = LoggingWorker(timeout=8.0) + monkeypatch.setattr(logging_worker, "GLOBAL_LOGGING_WORKER", worker) + rig: Final = _Rig(monkeypatch, count=count) + try: + with _transport(_upstream): + await _call(rig.router, rig.logging()) + payload: Final = await rig.capture.payload() + observed: Final = _observation(payload).observation + assert observed.outcome == "complete" and observed.reason is None + assert observed.plan is not None and observed.plan.breakpoints[0].prefix_tokens == 5000 + assert payload["response_cost"] is not None and worker._timeout_total == 0 + finally: + release.set() + await worker.stop() + + +@pytest.mark.parametrize( + "options,on_deployment", + ( + ({"thinking": {"type": "enabled", "budget_tokens": 2048}}, False), + ({"extra_body": {"speed": "fast", "output_config": {"effort": "high"}}}, False), + *product( + ( + {"container": {"id": "container_test"}}, + {"mcp_servers": [{"type": "url", "name": "test", "url": "https://example.com/mcp"}]}, + {"inference_geo": "us"}, + {"safeguards": [{"type": "default"}]}, + ), + (False, True), + ), + ), +) +async def test_native_baseline_identity_keeps_the_actual_transformed_body( + monkeypatch: pytest.MonkeyPatch, options: dict[str, JsonValue], on_deployment: bool +) -> None: + models: Final = _MESSAGES.validate_python( + [ + *_MODELS[:2], + { + **_MODELS[2], + "litellm_params": { + **_JSON_OBJECT.validate_python(_MODELS[2]["litellm_params"]), + **(options if on_deployment else {}), + }, + }, + ] + ) + rig: Final = _Rig(monkeypatch, models=models) + log: Final = rig.logging() + with _transport(_upstream): + await rig.router.anthropic_messages( + model="test-router", + max_tokens=16, + messages=_MESSAGES.validate_json(_MESSAGES_JSON.replace("question", "question USE_OPUS")), + litellm_logging_obj=log, + litellm_call_id=rig.call_id, + litellm_metadata={"user_api_key_hash": "test-caller-hash"}, + litellm_session_id="native-identical", + **({} if on_deployment else options), + ) + observed: Final = _observation(await rig.capture.payload()).observation + assert observed.outcome == "complete" + assert log.baseline_cache_context is not None + assert observed.baseline_equivalent and observed.usage is not None, ( + log.baseline_cache_context.baseline_body, + log.baseline_cache_context.selected_body_digest, + ) + + from litellm.proxy.spend_tracking.baseline_accounting import BaselineHistory, advance_baseline_history + + _, estimates = advance_baseline_history(BaselineHistory(), (observed,)) + assert estimates[0].provenance == "observed_identical" and estimates[0].usage == observed.usage + + +@pytest.mark.parametrize("tier_limit", (8, 16)) +@pytest.mark.parametrize("extra", ({}, {"max_tokens": 8})) +async def test_native_baseline_identity_respects_caller_limit_and_tier_override( + monkeypatch: pytest.MonkeyPatch, tier_limit: int, extra: dict[str, int] +) -> None: + models: Final = _MESSAGES.validate_python( + [ + { + "model_name": "test-router", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": { + "tiers": { + "SIMPLE": {"model_name": "opus", "litellm_params": {"max_tokens": tier_limit}}, + "MEDIUM": {"model_name": "opus", "litellm_params": {"max_tokens": tier_limit}}, + "COMPLEX": "opus", + "REASONING": "opus", + }, + "session_affinity": False, + }, + }, + }, + {**_MODELS[2], "litellm_params": {**_MODELS[2]["litellm_params"], "max_tokens": 64}}, + ] + ) + rig: Final = _Rig(monkeypatch, models=models) + with _transport(_upstream) as route: + await rig.router.anthropic_messages( + model="test-router", + max_tokens=8, + messages=_MESSAGES.validate_json(_MESSAGES_JSON), + litellm_logging_obj=rig.logging(), + litellm_call_id=rig.call_id, + litellm_metadata={"user_api_key_hash": "test-caller-hash"}, + litellm_session_id="native-limits", + extra_body=extra, + ) + observed: Final = _observation(await rig.capture.payload()).observation + wire: Final = _JSON_OBJECT.validate_json(route.calls.last.request.content) + assert wire["max_tokens"] == tier_limit + assert observed.baseline_equivalent == (tier_limit == 8) + + +@pytest.mark.parametrize("nested", (False, True)) +async def test_native_baseline_projection_matches_wire_parameter_placement( + monkeypatch: pytest.MonkeyPatch, + nested: bool, +) -> None: + counted: Final[asyncio.Queue[Mapping[str, JsonValue]]] = asyncio.Queue() + + async def count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int: + counted.put_nowait(body) + return await _count(model, api_key, body) + + models: Final = _MESSAGES.validate_python( + [ + { + **entry, + "litellm_params": { + **_JSON_OBJECT.validate_python(entry["litellm_params"]), + "model": "anthropic/claude-opus-5", + }, + } + if entry["model_name"] == "sonnet" + else entry + for entry in _MODELS + ] + ) + rig: Final = _Rig(monkeypatch, count=count, models=models) + settings: Final = {"speed": "standard", "thinking": {"type": "adaptive"}, "output_config": {"effort": "medium"}} + with _transport(_upstream) as route: + await rig.router.anthropic_messages( + model="test-router", + max_tokens=4096, + messages=_MESSAGES.validate_json(_MESSAGES_JSON), + litellm_logging_obj=rig.logging(), + litellm_call_id=rig.call_id, + litellm_metadata={"user_api_key_hash": "test-caller-hash"}, + litellm_session_id="native-placement", + **({"extra_body": settings} if nested else settings), + ) + captured: Final = _observation(await rig.capture.payload()) + wire: Final = _JSON_OBJECT.validate_json(route.calls.last.request.content) + assert captured.observation.plan is not None and not captured.observation.baseline_equivalent + projected: Final = counted.get_nowait() + assert {key: projected[key] for key in settings if key in projected} == { + key: wire[key] for key in settings if key in wire + } + assert {key: wire[key] for key in settings if key in wire} == ({} if nested else settings) + + +_PARITY_TOOL: Final = {"name": "custom", "input_schema": {"type": "object"}, "cache_control": {"type": "ephemeral"}} +_PARITY_SYSTEM: Final = [{"type": "text", "text": "stable system", "cache_control": {"type": "ephemeral"}}] +_PARITY_POINTS: Final = [{"location": "message", "role": "system"}, {"location": "message", "index": -1}] + + +@pytest.mark.parametrize( + "caller,selected,baseline,summary", + ( + pytest.param({"extra_body": {"cache_control": {"type": "ephemeral"}}}, {}, {}, False, id="envelope-control"), + pytest.param({"extra_body": {"system": _PARITY_SYSTEM}}, {}, {}, False, id="envelope-system"), + pytest.param( + {"extra_body": {"messages": _MESSAGES.validate_json(_MESSAGES_JSON)}}, {}, {}, False, id="envelope-messages" + ), + pytest.param({}, {"tools": [_PARITY_TOOL]}, {}, False, id="selected-tool-mark"), + pytest.param({}, {}, {"tools": [_PARITY_TOOL]}, False, id="baseline-tool-mark"), + pytest.param({}, {"system": _PARITY_SYSTEM}, {"system": "baseline system"}, False, id="selected-system-mark"), + pytest.param({}, {"system": "selected system"}, {"system": _PARITY_SYSTEM}, False, id="baseline-system-mark"), + pytest.param({"system": None}, {}, {"system": "configured system"}, False, id="null-system"), + pytest.param({"thinking": None}, {}, {"thinking": {"type": "adaptive"}}, False, id="null-thinking"), + pytest.param({"tools": None}, {}, {"tools": [_PARITY_TOOL]}, False, id="null-tools"), + pytest.param({"verbosity": "low", "instructions": "ignored"}, {}, {}, False, id="ignored-native-options"), + pytest.param( + {"messages": _MESSAGES.validate_json(_MESSAGES_JSON.replace("stable", " "))}, + {}, + {}, + False, + id="empty-marked-block", + ), + pytest.param({"thinking": {"type": "adaptive"}}, {}, {}, True, id="reasoning-summary"), + pytest.param( + {"thinking": {"type": "adaptive"}, "additional_drop_params": ["thinking.display"]}, + {}, + {}, + True, + id="drop-nested-option", + ), + pytest.param( + {"cache_control_injection_points": _PARITY_POINTS}, + {"tools": [{**_PARITY_TOOL, "name": f"custom_{index}"} for index in range(4)]}, + {}, + False, + id="configured-cap", + ), + ), +) +async def test_native_baseline_projection_matches_direct_baseline_request( + monkeypatch: pytest.MonkeyPatch, + caller: dict[str, JsonValue], + selected: dict[str, JsonValue], + baseline: dict[str, JsonValue], + summary: bool, +) -> None: + counted: Final[asyncio.Queue[Mapping[str, JsonValue]]] = asyncio.Queue() + + async def count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int: + counted.put_nowait(body) + return await _count(model, api_key, body) + + def upstream(request: httpx.Request) -> httpx.Response: + body: Final = _JSON_OBJECT.validate_json(request.content) + model: Final = body.get("model") + assert isinstance(model, str) + return httpx.Response( + 200, + request=request, + json={ + **_message(True, model), + "usage": {"input_tokens": 6000, "output_tokens": 10}, + }, + ) + + models: Final = _MESSAGES.validate_python( + [ + _MODELS[0], + { + **_MODELS[1], + "litellm_params": {**_JSON_OBJECT.validate_python(_MODELS[1]["litellm_params"]), **selected}, + }, + { + **_MODELS[2], + "litellm_params": {**_JSON_OBJECT.validate_python(_MODELS[2]["litellm_params"]), **baseline}, + }, + ] + ) + rig: Final = _Rig(monkeypatch, models=models, count=count) + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", False) + monkeypatch.setattr(litellm, "reasoning_auto_summary", summary) + monkeypatch.delenv("LITELLM_REASONING_AUTO_SUMMARY", raising=False) + request: Final = { + "messages": [{"role": "user", "content": "question"}], + **({"system": "stable system"} if "system" not in selected and "system" not in baseline else {}), + "max_tokens": 4096, + "enable_prompt_caching": True, + **caller, + } + with _transport(upstream) as route: + await rig.router.anthropic_messages(model="opus", **_JSON_OBJECT.validate_python(request)) + direct: Final = _JSON_OBJECT.validate_json(route.calls.last.request.content) + await rig.router.anthropic_messages( + model="test-router", + litellm_logging_obj=rig.logging(), + litellm_call_id=rig.call_id, + litellm_metadata={"user_api_key_hash": "test-caller-hash"}, + litellm_session_id="native-projection-parity", + **_JSON_OBJECT.validate_python(request), + ) + captured: Final = _observation(await rig.capture.payload()) + assert captured.observation.plan is not None, captured.observation.reason + projected: Final = counted.get_nowait() + assert {key: value for key, value in projected.items() if key not in ("metadata", "stream")} == { + key: value for key, value in direct.items() if key not in ("metadata", "stream") + } + + +@pytest.mark.parametrize( + "selected,baseline,usage_field,observed_value,multiplier", + ( + ({}, {"speed": "fast"}, "speed", "standard", 3.0), + ({"inference_geo": "us"}, {}, "inference_geo", "us", 1.0), + ), +) +async def test_native_baseline_prices_projected_settings_without_changing_actual_spend( + monkeypatch: pytest.MonkeyPatch, + selected: dict[str, JsonValue], + baseline: dict[str, JsonValue], + usage_field: str, + observed_value: str, + multiplier: float, +) -> None: + from litellm.proxy.spend_tracking.baseline_accounting import BaselineHistory, advance_baseline_history + from litellm.proxy.spend_tracking.savings import baseline_cost_snapshot, price_baseline_comparison + from litellm.types.utils import ModelInfo + + def upstream(request: httpx.Request) -> httpx.Response: + body: Final = _JSON_OBJECT.validate_json(request.content) + model: Final = body.get("model") + assert isinstance(model, str) + assert body.get(usage_field) == selected.get(usage_field) + return httpx.Response( + 200, + request=request, + json={ + **_message(True, model), + "usage": { + "input_tokens": 6000, + "output_tokens": 10, + usage_field: observed_value, + "cache_creation": {"ephemeral_5m_input_tokens": 0, "ephemeral_1h_input_tokens": 0}, + }, + }, + ) + + rig: Final = _Rig( + monkeypatch, + models=_MESSAGES.validate_python( + [ + _MODELS[0], + { + **_MODELS[1], + "litellm_params": { + **_JSON_OBJECT.validate_python(_MODELS[1]["litellm_params"]), + **selected, + }, + }, + { + **_MODELS[2], + "litellm_params": { + **_JSON_OBJECT.validate_python(_MODELS[2]["litellm_params"]), + **baseline, + }, + }, + ] + ), + ) + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", False) + with _transport(upstream): + response: Final = _JSON_OBJECT.validate_python( + await rig.router.anthropic_messages( + model="test-router", + max_tokens=16, + messages=[{"role": "user", "content": "question"}], + litellm_logging_obj=rig.logging(), + litellm_call_id=rig.call_id, + litellm_metadata={"user_api_key_hash": "test-caller-hash"}, + litellm_session_id="baseline-speed", + ) + ) + payload: Final = await rig.capture.payload() + captured: Final = _observation(payload) + _, estimates = advance_baseline_history(BaselineHistory(), (captured.observation,)) + estimate: Final = estimates[0] + assert captured.prices is not None + prices: Final[ModelInfo] = { + **captured.prices, + "input_cost_per_token": 1e-6, + "output_cost_per_token": 2e-6, + "provider_specific_entry": {"fast": 3.0, "us": 2.0}, + } + actual: Final = payload["response_cost"] + assert isinstance(actual, float) + snapshot: Final = baseline_cost_snapshot( + captured.model, + prices, + actual, + _OBJECTS.validate_python(payload["cost_breakdown"]), + None, + ) + comparison: Final = price_baseline_comparison(snapshot, estimate.usage, estimate.provenance) + assert comparison is not None and snapshot.actual_token_cost is not None, estimate.reason + assert comparison.baseline == pytest.approx( + actual + (6000 * 1e-6 + 10 * 2e-6) * multiplier - snapshot.actual_token_cost + ) + assert comparison.actual == actual + assert _JSON_OBJECT.validate_python(response["usage"])[usage_field] == observed_value + + +async def test_native_request_rewritten_after_capture_preserves_spend_without_guessing_baseline( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class RewriteSystem(CustomLogger): + async def async_pre_call_deployment_hook( + self, kwargs: Mapping[str, object], call_type: CallTypes | None + ) -> dict[str, object]: + return {**kwargs, "system": "hook system"} + + rig: Final = _Rig(monkeypatch) + monkeypatch.setattr(litellm, "callbacks", [rig.hook, RewriteSystem()]) + with _transport(_upstream) as route: + await _call(rig.router, rig.logging()) + payload: Final = await rig.capture.payload() + wire: Final = _JSON_OBJECT.validate_json(route.calls.last.request.content) + observed: Final = _observation(payload).observation + assert wire["system"] == "hook system" + assert observed.reason == "unsupported_request_transformation" and observed.plan is None + assert observed.usage is not None + actual: Final = payload["response_cost"] + assert isinstance(actual, float) and actual > 0 + + +@pytest.mark.parametrize("history", ("long_session", "non_ascii")) +async def test_native_baseline_models_long_and_non_ascii_history(monkeypatch: pytest.MonkeyPatch, history: str) -> None: + rounds: Final = tuple( + message + for index in range(1200) + for message in ( + {"role": "assistant", "content": [{"type": "tool_use", "id": f"t{index}", "name": "Read", "input": {}}]}, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": f"t{index}", "content": "ok"}]}, + ) + ) + prefix: Final = ( + [{"role": "user", "content": "start"}, *rounds] + if history == "long_session" + else [{"role": "user", "content": "a" * 400_000 + "é"}, {"role": "assistant", "content": "ok"}] + ) + messages: Final = json.dumps([*prefix, *_MESSAGES.validate_json(_MESSAGES_JSON)]) + rig: Final = _Rig(monkeypatch) + with _transport(_upstream): + await _call(rig.router, rig.logging(), messages=messages) + observed: Final = _observation(await rig.capture.payload()).observation + assert observed.outcome == "complete" and observed.plan is not None + + +async def test_native_baseline_abstains_after_selected_tier_compaction(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.router_strategy.complexity_router.context_compaction import compaction_executor + + monkeypatch.setitem( + litellm.model_cost, + "summary-fixture", + { + "litellm_provider": "anthropic", + "mode": "chat", + "max_input_tokens": 32000, + "max_output_tokens": 4096, + "supports_anthropic_compaction": True, + }, + ) + models: Final = _MESSAGES.validate_python( + [ + { + "model_name": "test-router", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": { + "tiers": {"SIMPLE": "sonnet", "MEDIUM": "opus", "COMPLEX": "opus", "REASONING": "opus"}, + "keyword_tier_rules": [{"keywords": ["answer"], "tier": "SIMPLE"}], + "session_affinity": False, + "enable_context_window_escalation": False, + "max_tokens_from_tier_model": False, + "context_compaction": {"model": "compactor", "max_tokens": 512}, + }, + }, + }, + { + "model_name": "sonnet", + "litellm_params": {"model": "anthropic/claude-sonnet-5", "api_key": "test-selected"}, + "model_info": {"id": "selected", "max_input_tokens": 512, "max_output_tokens": 64}, + }, + { + "model_name": "opus", + "litellm_params": {"model": "anthropic/claude-opus-5", "api_key": "test-selected"}, + "model_info": {"id": "baseline", "max_input_tokens": 200000, "max_output_tokens": 4096}, + }, + { + "model_name": "compactor", + "litellm_params": {"model": "anthropic/summary-fixture", "api_key": "test-compactor"}, + "model_info": {"id": "compactor"}, + }, + ] + ) + + async def summarize(protocol: object, request: object, parent_model: object = None) -> Mapping[str, object]: + return { + "stop_reason": "compaction", + "content": [{"type": "compaction", "content": "compacted", "signature": "s"}], + "usage": {"input_tokens": 0, "output_tokens": 0}, + } + + messages: Final = json.dumps( + [ + {"role": "user", "content": "Background detail. " * 300}, + {"role": "assistant", "content": "Recorded"}, + {"role": "user", "content": "Answer briefly"}, + ] + ) + rig: Final = _Rig(monkeypatch, models=models) + token: Final = compaction_executor.set(summarize) + try: + with _transport(_upstream) as route: + await _call(rig.router, rig.logging(), messages=messages) + observed: Final = _observation(await rig.capture.payload()).observation + wire: Final = route.calls.last.request.content.decode() + finally: + compaction_executor.reset(token) + assert "compacted" in wire and "Background detail" not in wire + assert observed.reason == "unsupported_request_transformation" and observed.plan is None + assert observed.usage is not None + + +async def test_selected_tier_cache_markers_do_not_hide_an_unmarked_baseline_plan( + monkeypatch: pytest.MonkeyPatch, +) -> None: + models: Final = _MESSAGES.validate_python( + [ + _MODELS[0], + { + "model_name": "sonnet", + "litellm_params": { + "model": "anthropic/claude-sonnet-5", + "api_key": "test-selected", + "cache_control_injection_points": [{"location": "message", "role": "user", "index": -1}], + }, + "model_info": {"id": "selected"}, + }, + _MODELS[2], + ] + ) + rig: Final = _Rig(monkeypatch, models=models) + with _transport(_upstream) as route: + await _call( + rig.router, + rig.logging(), + messages='[{"role":"user","content":[{"type":"text","text":"stable"},{"type":"text","text":"question"}]}]', + ) + observed: Final = _observation(await rig.capture.payload()).observation + wire: Final = route.calls.last.request.content.decode() + assert "cache_control" in wire + assert observed.outcome == "complete" and not observed.baseline_equivalent + assert observed.reason is None and observed.plan is not None and not observed.plan.breakpoints + + +@pytest.mark.parametrize("recovery", ("retry", "fallback")) +async def test_tier_pins_never_enter_the_caller_snapshot_on_later_routing_passes( + monkeypatch: pytest.MonkeyPatch, recovery: str +) -> None: + pinned: Final = {"model_name": "first", "litellm_params": {"reasoning_effort": "high", "max_tokens": 777}} + models: Final = _MESSAGES.validate_python( + [ + { + "model_name": "test-router", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": { + "tiers": {"SIMPLE": pinned, "MEDIUM": pinned, "COMPLEX": pinned, "REASONING": "opus"}, + "session_affinity": False, + }, + }, + }, + { + "model_name": "fallback-router", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": { + "tiers": {"SIMPLE": "sonnet", "MEDIUM": "sonnet", "COMPLEX": "sonnet", "REASONING": "opus"}, + "session_affinity": False, + }, + }, + }, + { + "model_name": "first", + "litellm_params": { + "model": "anthropic/claude-sonnet-5" if recovery == "retry" else "anthropic/claude-haiku-5", + "api_key": "test-selected", + }, + "model_info": {"id": "first"}, + }, + *_MODELS[1:], + ] + ) + rig: Final = _Rig(monkeypatch, models=models, retries=1 if recovery == "retry" else 0) + rig.router.fallbacks = [{"test-router": ["fallback-router"]}] + + def upstream(request: httpx.Request) -> httpx.Response: + return _upstream(request) if route.call_count else _error(request, 429, "first attempt") + + log: Final = rig.logging() + with _transport(upstream) as route: + await _call(rig.router, log) + await rig.capture.payload() + first_wire: Final = _JSON_OBJECT.validate_json(route.calls[0].request.content) + assert first_wire.get("output_config") == {"effort": "high"} and first_wire.get("max_tokens") == 777 + context: Final = log.baseline_cache_context + assert context is not None and context.baseline_body is not None + assert context.baseline_body.get("max_tokens") == 16 + assert "output_config" not in context.baseline_body and "thinking" not in context.baseline_body diff --git a/tests/unit/proxy/spend_tracking/test_baseline_accounting.py b/tests/unit/proxy/spend_tracking/test_baseline_accounting.py index a188d65502d..368704e7a75 100644 --- a/tests/unit/proxy/spend_tracking/test_baseline_accounting.py +++ b/tests/unit/proxy/spend_tracking/test_baseline_accounting.py @@ -111,7 +111,13 @@ def test_prefix_match_expiry_and_usage_pricing_fields(ttl: int) -> None: assert warm.usage.prompt_tokens_details.cached_tokens == 6000 assert cold.usage.prompt_tokens_details.cached_tokens == 0 assert cold.usage.prompt_tokens_details.cache_creation_tokens == 6000 - unaffected: Final = {"prompt_tokens", "total_tokens", "prompt_tokens_details", "cache_read_input_tokens", "cache_creation_input_tokens"} + unaffected: Final = { + "prompt_tokens", + "total_tokens", + "prompt_tokens_details", + "cache_read_input_tokens", + "cache_creation_input_tokens", + } assert warm.usage.model_dump(exclude=unaffected) == first.usage.model_dump(exclude=unaffected) assert cold.usage.model_dump(exclude=unaffected) == first.usage.model_dump(exclude=unaffected) @@ -124,7 +130,10 @@ def test_growth_lookback_and_mixed_ttl_keep_distinct_read_write_buckets(warm_tai second: Final = _replay(first, _observation("second", 10001.0, plan=grown))[-1] assert second.reason == "history_unavailable" history: Final = BaselineHistory( - first_at=1.0, last_at=10000.0, equivalent=False, uncertain_before=1.0, + first_at=1.0, + last_at=10000.0, + equivalent=False, + uncertain_before=1.0, entries=(CacheEntry("tail:300", "tail", 7000, 300, 10000.0, 10300.0),) if warm_tail else (), ) _, estimates = advance_baseline_history(history, (_observation("mixed", 10001.0, plan=grown),)) @@ -134,8 +143,12 @@ def test_growth_lookback_and_mixed_ttl_keep_distinct_read_write_buckets(warm_tai # Anthropic billing locations: B is the highest 1h breakpoint AFTER the highest hit A. # https://platform.claude.com/docs/en/build-with-claude/prompt-caching#mixing-different-ttls (2026-09-15) assert usage.prompt_tokens_details.cached_tokens == (7000 if warm_tail else 0) - assert usage.prompt_tokens_details.cache_creation_token_details.ephemeral_1h_input_tokens == (0 if warm_tail else 6500) - assert usage.prompt_tokens_details.cache_creation_token_details.ephemeral_5m_input_tokens == (0 if warm_tail else 500) + assert usage.prompt_tokens_details.cache_creation_token_details.ephemeral_1h_input_tokens == ( + 0 if warm_tail else 6500 + ) + assert usage.prompt_tokens_details.cache_creation_token_details.ephemeral_5m_input_tokens == ( + 0 if warm_tail else 500 + ) @pytest.mark.parametrize("change", ["prefix", "ttl", "unavailable", "failed", "response_cache"]) @@ -197,3 +210,62 @@ def test_modeled_read_cannot_recharge_the_original_private_write_count() -> None } input_cost, output_cost = cost_per_token("claude-opus-5", warm.usage, model_info=prices) assert input_cost + output_cost == pytest.approx((200 * 1e-6 + 6000 * 1e-7 + 30 * 2e-6) * 2.0 * 1.1) + + +def test_mixed_lifetime_lookback_preserves_a_compatible_native_hit() -> None: + marker: Final = _marker("prefix", 3600, 6000) + history: Final = BaselineHistory( + first_at=1.0, + last_at=10000.0, + equivalent=False, + uncertain_before=1.0, + entries=(CacheEntry(marker.fingerprint, marker.content_fingerprint, 6000, 3600, 10000.0, 13600.0),), + ) + plan: Final = CountedPromptCachePlan( + 7100, (_marker("grown", 3600, 6500, ("prefix",)), _marker("tail", 300, 7000, ("prefix",))) + ) + _, estimates = advance_baseline_history(history, (_observation("next", 10001.0, plan=plan),)) + usage: Final = estimates[0].usage + assert usage is not None, estimates[0].reason + assert usage.prompt_tokens_details.cached_tokens == 6000 + assert usage.prompt_tokens_details.text_tokens == 100 + assert usage.prompt_tokens_details.cache_creation_token_details == CacheCreationTokenDetails( + ephemeral_5m_input_tokens=500, ephemeral_1h_input_tokens=500 + ) + + +def test_short_lifetime_hit_cannot_seed_an_unpaid_long_lifetime_entry() -> None: + first: Final = _observation("initial", plan=CountedPromptCachePlan(6200, (_marker("5", 3600, 6000),))) + short: Final = _observation("short", 13700.0, plan=CountedPromptCachePlan(6200, (_marker("3", 300, 4600),))) + mixed: Final = _observation( + "mixed", + 13710.0, + plan=CountedPromptCachePlan(6200, (_marker("3", 3600, 4600), _marker("4", 300, 5500, ("3",)))), + ) + later: Final = _observation("later", 14710.0, plan=CountedPromptCachePlan(6200, (_marker("3", 3600, 4600),))) + _, _, upgrade, after_expiry = _replay(first, short, mixed, later) + assert upgrade.usage is not None and after_expiry.usage is not None + assert upgrade.usage.prompt_tokens_details.cached_tokens == 4600 + assert upgrade.usage.prompt_tokens_details.cache_creation_token_details.ephemeral_1h_input_tokens == 0 + assert after_expiry.usage.prompt_tokens_details.cached_tokens == 0 + assert after_expiry.usage.prompt_tokens_details.cache_creation_token_details.ephemeral_1h_input_tokens == 4600 + + +@pytest.mark.parametrize("writes", (0, 50, None)) +def test_cache_creation_split_is_optional_only_without_writes(writes: int | None) -> None: + usage: Final = Usage( + prompt_tokens=100, + completion_tokens=10, + total_tokens=110, + prompt_tokens_details=PromptTokensDetailsWrapper( + text_tokens=100 - (writes or 0), + cached_tokens=0, + cache_creation_tokens=writes, + ), + ) + observed: Final = _observation("no-split", usage=usage, plan=CountedPromptCachePlan(100, ())) + restored: Final = BaselineObservation.model_validate_json(observed.model_dump_json()) + estimate: Final = _replay(restored)[0] + assert (estimate.usage is not None) is (writes == 0) + if estimate.usage is not None: + assert estimate.usage.prompt_tokens == usage.prompt_tokens diff --git a/tests/unit/router_utils/test_baseline_request.py b/tests/unit/router_utils/test_baseline_request.py new file mode 100644 index 00000000000..37d21597752 --- /dev/null +++ b/tests/unit/router_utils/test_baseline_request.py @@ -0,0 +1,59 @@ +from typing import Final + +from litellm.router_utils.baseline_request import baseline_request, capture_baseline_parameters + + +def test_baseline_snapshot_owns_nested_caller_settings_and_overrides_routed_settings() -> None: + reasoning: Final = {"effort": "medium"} + snapshot: Final = capture_baseline_parameters({"reasoning": reasoning, "verbosity": "low"}) + assert snapshot is not None + reasoning["effort"] = "high" + projected: Final = baseline_request( + {"messages": [{"role": "user", "content": "hello"}], "reasoning": reasoning, "verbosity": "high"}, + snapshot, + {"verbosity": "medium"}, + ) + assert projected == { + "messages": [{"role": "user", "content": "hello"}], + "reasoning": {"effort": "medium"}, + "verbosity": "low", + } + + +def test_oversized_snapshot_fails_closed_before_json_validation() -> None: + assert capture_baseline_parameters({"output_config": {"format": "x" * 5_000_000}}) is None + + +def test_snapshot_retains_extra_body_settings_but_no_credentials() -> None: + assert capture_baseline_parameters({"api_key": "private", "extra_body": {"verbosity": "low"}}) == { + "extra_body": {"verbosity": "low"} + } + + +def test_chat_projection_applies_extra_body_after_top_level_parameters() -> None: + snapshot: Final = capture_baseline_parameters({"verbosity": "high", "extra_body": {"verbosity": "low"}}) + assert snapshot is not None + assert baseline_request({}, snapshot, {}) == {"verbosity": "low"} + + +def test_baseline_projection_keeps_caller_tools_and_request_parameter_precedence() -> None: + from litellm.router import Router + + deployment: Final = { + "tools": [{"type": "function", "function": {"name": "configured"}}], + "tool_choice": "required", + "max_tokens": 64, + } + caller: Final = { + "tools": [{"type": "function", "function": {"name": "caller"}}], + "tool_choice": "auto", + "max_tokens": 128, + } + actual_request: Final = dict(caller) + Router._merge_tools_from_deployment({"litellm_params": deployment}, actual_request) + snapshot: Final = capture_baseline_parameters(caller) + assert snapshot is not None + assert baseline_request({"tools": [{"name": "routed-only"}], "max_tokens": 4}, snapshot, deployment) == { + **deployment, + **actual_request, + }