mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
3538e87e45
commit
0338498067
22 changed files with 1608 additions and 277 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"] = ""
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
153
litellm/router_utils/baseline_request.py
Normal file
153
litellm/router_utils/baseline_request.py
Normal 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 {}
|
||||
),
|
||||
}
|
||||
)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
59
tests/unit/router_utils/test_baseline_request.py
Normal file
59
tests/unit/router_utils/test_baseline_request.py
Normal 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,
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue