fix(router): preserve native baseline identity and accounting (#44960)

* fix(router): preserve native baseline identity and accounting

* fix(router): preserve injected system caches in native baselines

* fix(router): prepare native baselines through shared request owners

* fix(router): capture native baseline fields from the provider schema

* refactor(router): reuse native provider parameter discovery

* fix(router): keep long native baselines and abstain after compaction

Message history no longer counts against the settings snapshot budget, so long
and non-ASCII native sessions keep modeled baselines. Selected-tier compaction
now abstains because the baseline would otherwise inherit the compacted history.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

* fix(router): judge implicit caching against the selected request

The implicit-cache guard compared the selected response's cache usage with the
projected baseline's breakpoints, so selected-tier cache markers made a usable
unmarked baseline plan look like unexplained caching. The guard now checks the
selected wire request. Also removes a stamp-reuse branch that could never run
because routing clears the stamp first; every pass already captures caller
settings from fresh kwargs.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

---------

Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
tin-berri 2026-10-07 11:00:31 -07:00 • committed by GitHub
parent 3538e87e45
commit 0338498067
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
22 changed files with 1608 additions and 277 deletions

View file

@ -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,

View file

@ -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:

View file

@ -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"] = ""

View file

@ -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()

View file

@ -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)

View file

@ -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 (

View file

@ -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 (

View file

@ -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)

View file

@ -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
),
)

View file

@ -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
),

View file

@ -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

View file

@ -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 {}
),
}
)

View file

@ -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)

View file

@ -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",

View file

@ -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:

View file

@ -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()

View file

@ -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():

View file

@ -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(

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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,
}