fix(alerting): make anomaly alerts opt-in, average sparse baselines over full window, validate numeric settings

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-08-27 01:25:18 +00:00
commit f545c04447
35 changed files with 799 additions and 550 deletions

View file

@ -1364,8 +1364,6 @@ X_LITELLM_DISABLE_CALLBACKS: Final = "x-litellm-disable-callbacks"
LITELLM_METADATA_FIELD: Final = "litellm_metadata"
OLD_LITELLM_METADATA_FIELD: Final = "metadata"
RETURN_RAW_MODEL_NAME_METADATA_KEY: Final = "_complexity_router_return_raw_model_name"
AUTO_ROUTED_REQUEST_METADATA_KEY: Final = "_auto_routed_request"
ROUTER_MODEL_NAME_RESPONSE_FIELD: Final = "router_model_name"
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY: Final = "_session_deployment_affinity_ttl"
CONSUMED_REQUEST_TAGS_METADATA_KEY: Final = "_consumed_request_tags"
INTERNAL_CALL_ORIGIN_METADATA_KEY: Final = "internal_call_origin"

View file

@ -7,6 +7,7 @@ import base64
import os
from collections.abc import Awaitable, Callable, Generator
from datetime import timedelta
from importlib import metadata
from typing import Any, Final, TypeVar
import httpx
@ -21,6 +22,18 @@ try:
streamable_http_client = getattr(streamable_http_module, "streamable_http_client", None)
except ImportError:
pass
MCP_STREAMABLE_HTTP_REQUIREMENT: Final = "mcp>=1.28.1"
def missing_streamable_http_client_error() -> ImportError:
return ImportError(
f"MCP streamable HTTP transport requires {MCP_STREAMABLE_HTTP_REQUIREMENT}, but the installed "
f"mcp {metadata.version('mcp')} does not provide streamable_http_client. "
"Fix with: pip install 'litellm[mcp]' (or upgrade mcp directly: pip install -U mcp)"
)
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
from mcp.types import CallToolResult as MCPCallToolResult
from mcp.types import (
@ -323,7 +336,7 @@ class MCPClient:
)
# HTTP transport (default)
if streamable_http_client is None:
raise ImportError("streamable_http_client is not available. Please install mcp with HTTP support.")
raise missing_streamable_http_client_error()
headers = self._get_auth_headers()
httpx_client_factory = self._create_httpx_client_factory()
verbose_logger.debug("litellm headers for streamable_http_client: %s", headers)

View file

@ -2,6 +2,7 @@ import asyncio
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy.pass_through_endpoints.success_handler import (
PassThroughEndpointLogging,
@ -65,6 +66,7 @@ class BaseGoogleGenAIGenerateContentStreamingIterator:
litellm_logging_obj: LiteLLMLoggingObj,
request_body: dict,
model: str,
custom_llm_provider: str,
hidden_params: dict[str, Any] | None = None,
):
self.litellm_logging_obj = litellm_logging_obj
@ -72,6 +74,10 @@ class BaseGoogleGenAIGenerateContentStreamingIterator:
self.start_time = datetime.now()
self.collected_chunks: list[bytes] = []
self.model = model
self.custom_llm_provider = custom_llm_provider
self.endpoint_type: Final = (
EndpointType.GEMINI if custom_llm_provider == litellm.LlmProviders.GEMINI.value else EndpointType.VERTEX_AI
)
self._hidden_params: dict[str, Any] = hidden_params or {}
async def _handle_async_streaming_logging(
@ -89,7 +95,7 @@ class BaseGoogleGenAIGenerateContentStreamingIterator:
passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ,
url_route="/v1/generateContent",
request_body=self.request_body or {},
endpoint_type=EndpointType.VERTEX_AI,
endpoint_type=self.endpoint_type,
start_time=self.start_time,
raw_bytes=self.collected_chunks,
end_time=end_time,
@ -118,13 +124,13 @@ class GoogleGenAIGenerateContentStreamingIterator(BaseGoogleGenAIGenerateContent
litellm_logging_obj=logging_obj,
request_body=request_body or {},
model=model,
custom_llm_provider=custom_llm_provider,
hidden_params=hidden_params,
)
self.response = response
self.model = model
self.generate_content_provider_config = generate_content_provider_config
self.litellm_metadata = litellm_metadata
self.custom_llm_provider = custom_llm_provider
# Gemini streamGenerateContent uses SSE line framing; iter_lines keeps
# large inlineData payloads (e.g. image/jpeg) intact within one event.
self.stream_iterator = response.iter_lines()
@ -169,13 +175,13 @@ class AsyncGoogleGenAIGenerateContentStreamingIterator(BaseGoogleGenAIGenerateCo
litellm_logging_obj=logging_obj,
request_body=request_body or {},
model=model,
custom_llm_provider=custom_llm_provider,
hidden_params=hidden_params,
)
self.response = response
self.model = model
self.generate_content_provider_config = generate_content_provider_config
self.litellm_metadata = litellm_metadata
self.custom_llm_provider = custom_llm_provider
# Gemini streamGenerateContent uses SSE line framing; aiter_lines keeps
# large inlineData payloads (e.g. image/jpeg) intact within one event.
self.stream_iterator = response.aiter_lines()

View file

@ -20,8 +20,7 @@ SELECT
user_id,
COALESCE(SUM(spend) FILTER (WHERE date = $1), 0)::float AS daily_spend,
COALESCE(SUM(spend) FILTER (WHERE date >= $2), 0)::float AS monthly_spend,
COALESCE(SUM(spend) FILTER (WHERE date >= $3 AND date < $1), 0)::float AS baseline_spend,
COUNT(DISTINCT date) FILTER (WHERE date >= $3 AND date < $1 AND spend > 0)::int AS baseline_days
COALESCE(SUM(spend) FILTER (WHERE date >= $3 AND date < $1), 0)::float AS baseline_spend
FROM "LiteLLM_DailyUserSpend"
WHERE date >= LEAST($2, $3) AND user_id IS NOT NULL
GROUP BY user_id
@ -35,7 +34,6 @@ class UserSpendRow:
daily_spend: float
monthly_spend: float
baseline_spend: float
baseline_days: int
@dataclass(frozen=True, slots=True)
@ -101,8 +99,8 @@ def _monthly_threshold_event(row: UserSpendRow, args: SlackAlertingArgs, month_s
def _anomaly_event(row: UserSpendRow, args: SlackAlertingArgs, today_str: str) -> UserSpendAlertEvent | None:
if row.daily_spend < args.spend_anomaly_min_spend:
return None
baseline_daily_avg: Final = row.baseline_spend / row.baseline_days if row.baseline_days > 0 else 0.0
if row.baseline_days > 0 and row.daily_spend <= args.spend_anomaly_multiplier * baseline_daily_avg:
baseline_daily_avg: Final = row.baseline_spend / args.spend_anomaly_baseline_days
if row.baseline_spend > 0 and row.daily_spend <= args.spend_anomaly_multiplier * baseline_daily_avg:
return None
return UserSpendAlertEvent(
kind="anomaly",

View file

@ -209,6 +209,8 @@ class DotpromptManager(CustomPromptManagement):
prompt_spec=prompt_spec,
prompt_label=prompt_label,
prompt_version=prompt_version,
ignore_prompt_manager_model=ignore_prompt_manager_model,
ignore_prompt_manager_optional_params=ignore_prompt_manager_optional_params,
)
async def async_get_chat_completion_prompt(

View file

@ -416,17 +416,8 @@ class GenericPromptManager(CustomPromptManagement):
tools=tools,
prompt_label=prompt_label,
prompt_version=prompt_version,
ignore_prompt_manager_model=(
ignore_prompt_manager_model or prompt_spec.litellm_params.ignore_prompt_manager_model
if prompt_spec
else False
),
ignore_prompt_manager_optional_params=(
ignore_prompt_manager_optional_params
or prompt_spec.litellm_params.ignore_prompt_manager_optional_params
if prompt_spec
else False
),
ignore_prompt_manager_model=ignore_prompt_manager_model,
ignore_prompt_manager_optional_params=ignore_prompt_manager_optional_params,
)
def get_chat_completion_prompt(
@ -457,17 +448,8 @@ class GenericPromptManager(CustomPromptManagement):
prompt_spec=prompt_spec,
prompt_label=prompt_label,
prompt_version=prompt_version,
ignore_prompt_manager_model=(
ignore_prompt_manager_model or prompt_spec.litellm_params.ignore_prompt_manager_model
if prompt_spec
else False
),
ignore_prompt_manager_optional_params=(
ignore_prompt_manager_optional_params
or prompt_spec.litellm_params.ignore_prompt_manager_optional_params
if prompt_spec
else False
),
ignore_prompt_manager_model=ignore_prompt_manager_model,
ignore_prompt_manager_optional_params=ignore_prompt_manager_optional_params,
)
def clear_cache(self) -> None:

View file

@ -19,6 +19,19 @@ class PromptManagementClient(TypedDict):
completed_messages: list[AllMessageValues] | None
def resolve_prompt_manager_ignore_flags(
prompt_spec: PromptSpec | None,
ignore_prompt_manager_model: bool | None,
ignore_prompt_manager_optional_params: bool | None,
) -> tuple[bool, bool]:
spec_params: Final = prompt_spec.litellm_params if prompt_spec is not None else None
return (
bool(ignore_prompt_manager_model) or bool(spec_params is not None and spec_params.ignore_prompt_manager_model),
bool(ignore_prompt_manager_optional_params)
or bool(spec_params is not None and spec_params.ignore_prompt_manager_optional_params),
)
class PromptManagementBase(ABC):
@property
@abstractmethod
@ -182,13 +195,18 @@ class PromptManagementBase(ABC):
prompt_version=prompt_version,
)
resolved_ignore_model, resolved_ignore_optional_params = resolve_prompt_manager_ignore_flags(
prompt_spec=prompt_spec,
ignore_prompt_manager_model=ignore_prompt_manager_model,
ignore_prompt_manager_optional_params=ignore_prompt_manager_optional_params,
)
return self.post_compile_prompt_processing(
prompt_template=prompt_template,
messages=messages,
non_default_params=non_default_params,
model=model,
ignore_prompt_manager_model=ignore_prompt_manager_model,
ignore_prompt_manager_optional_params=ignore_prompt_manager_optional_params,
ignore_prompt_manager_model=resolved_ignore_model,
ignore_prompt_manager_optional_params=resolved_ignore_optional_params,
)
async def async_get_chat_completion_prompt(
@ -224,11 +242,16 @@ class PromptManagementBase(ABC):
prompt_version=prompt_version,
)
resolved_ignore_model, resolved_ignore_optional_params = resolve_prompt_manager_ignore_flags(
prompt_spec=prompt_spec,
ignore_prompt_manager_model=ignore_prompt_manager_model,
ignore_prompt_manager_optional_params=ignore_prompt_manager_optional_params,
)
return self.post_compile_prompt_processing(
prompt_template=prompt_template,
messages=messages,
non_default_params=non_default_params,
model=model,
ignore_prompt_manager_model=ignore_prompt_manager_model,
ignore_prompt_manager_optional_params=ignore_prompt_manager_optional_params,
ignore_prompt_manager_model=resolved_ignore_model,
ignore_prompt_manager_optional_params=resolved_ignore_optional_params,
)

View file

@ -1063,15 +1063,17 @@ def get_token_type_cost_breakdown(
reasoning_tokens = _coerce_token_count(getattr(usage, "reasoning_tokens", 0))
# Reasoning is billed at the selected tier's reasoning rate for tiered models,
# else at the explicit per-reasoning-token rate when the model defines one,
# otherwise at the standard output-token rate - this mirrors how the total
# completion cost is computed, so the breakdown can never diverge from it.
# else at the service-tier-aware per-reasoning-token rate - this mirrors how the
# total completion cost is computed, so the breakdown can never diverge from it.
tiered_reasoning_rate: Final = _get_tiered_reasoning_rate(model_info=model_info, usage=usage)
flat_reasoning_rate: Final = _get_cost_per_unit(model_info, "output_cost_per_reasoning_token", None)
reasoning_rate: Final = (
tiered_reasoning_rate
if tiered_reasoning_rate is not None
else (flat_reasoning_rate if flat_reasoning_rate is not None else completion_base_cost)
else _resolve_reasoning_token_cost(
model_info=model_info,
service_tier=service_tier,
completion_base_cost=completion_base_cost,
)
)
reasoning_cost = float(reasoning_tokens) * reasoning_rate

View file

@ -178,7 +178,7 @@ def update_response_metadata(
- response._hidden_params["litellm_overhead_time_ms"]
- response.response_time_ms
"""
if result is None:
if result is None or not hasattr(result, "_hidden_params"):
return
metadata: Final = ResponseMetadata(result)

View file

@ -21,14 +21,12 @@ import litellm
from litellm._logging import _redact_string, verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import (
AUTO_ROUTED_REQUEST_METADATA_KEY,
DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE,
DEFAULT_MAX_RECURSE_DEPTH,
LITELLM_DETAILED_TIMING,
LITELLM_HTTP_STATUS_CLIENT_DISCONNECTED,
MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG,
RETURN_RAW_MODEL_NAME_METADATA_KEY,
ROUTER_MODEL_NAME_RESPONSE_FIELD,
STREAM_SSE_DATA_PREFIX,
STREAM_SSE_KEEPALIVE_PING_BYTES,
UNSAFE_PROXY_RESPONSE_HEADERS,
@ -2036,54 +2034,6 @@ class ProxyBaseLLMRequestProcessing:
return deployment
return None
@staticmethod
def get_router_selected_model_name(
litellm_logging_obj: LiteLLMLoggingObj | None,
) -> str | None:
"""Model group an auto-routing strategy selected, or None if none fired.
The marker and ``deployment_model_name`` are written by different bucket
resolvers (``get_or_create_metadata_bucket`` vs
``_get_router_metadata_variable_name``), so they can land in different
buckets on the same request. Resolve each across both.
"""
litellm_params: Final = getattr(litellm_logging_obj, "litellm_params", None)
if not isinstance(litellm_params, dict):
return None
buckets: Final = tuple(
bucket for key in ("litellm_metadata", "metadata") if isinstance(bucket := litellm_params.get(key), dict)
)
if not any(bucket.get(AUTO_ROUTED_REQUEST_METADATA_KEY) is True for bucket in buckets):
return None
return next(
(
model_group
for bucket in buckets
if isinstance(model_group := bucket.get("deployment_model_name"), str) and model_group
),
None,
)
@staticmethod
def set_router_selected_model_field(
*,
response_obj: object,
router_model_name: str | None,
) -> None:
if not router_model_name:
return
if isinstance(response_obj, dict):
response_obj[ROUTER_MODEL_NAME_RESPONSE_FIELD] = router_model_name
return
try:
setattr(response_obj, ROUTER_MODEL_NAME_RESPONSE_FIELD, router_model_name)
except (AttributeError, TypeError, ValueError):
verbose_proxy_logger.debug(
"Could not set %s on response object of type %s",
ROUTER_MODEL_NAME_RESPONSE_FIELD,
type(response_obj),
)
@staticmethod
def _response_cost_from_logging_obj(
*,
@ -2582,10 +2532,6 @@ class ProxyBaseLLMRequestProcessing:
log_context=f"litellm_call_id={logging_obj.litellm_call_id}",
return_raw_model_name=_should_return_raw_model_name(self.data),
)
self.set_router_selected_model_field(
response_obj=response,
router_model_name=self.get_router_selected_model_name(logging_obj),
)
hidden_params = get_hidden_params_dict(response) # get any updated response headers
additional_headers = hidden_params.get("additional_headers", {}) or {}

View file

@ -615,7 +615,7 @@ class VertexPassthroughLoggingHandler:
response_cost: Final = litellm.completion_cost(
completion_response=litellm_model_response,
model=model,
custom_llm_provider="vertex_ai",
custom_llm_provider=custom_llm_provider,
vertex_location=vertex_location,
)

View file

@ -17,6 +17,9 @@ from litellm.types.utils import StandardPassThroughResponseObject
from .llm_provider_handlers.anthropic_passthrough_logging_handler import (
AnthropicPassthroughLoggingHandler,
)
from .llm_provider_handlers.gemini_passthrough_logging_handler import (
GeminiPassthroughLoggingHandler,
)
from .llm_provider_handlers.openai_passthrough_logging_handler import (
OpenAIPassthroughLoggingHandler,
)
@ -243,6 +246,26 @@ class PassThroughStreamingHandler:
)
standard_logging_response_object = vertex_passthrough_logging_handler_result["result"]
kwargs = vertex_passthrough_logging_handler_result["kwargs"]
elif endpoint_type == EndpointType.GEMINI:
gemini_passthrough_logging_handler_result: Final = (
GeminiPassthroughLoggingHandler._handle_logging_gemini_collected_chunks( # pyright: ignore[reportPrivateUsage] # mirrors sibling handler dispatch
litellm_logging_obj=litellm_logging_obj,
passthrough_success_handler_obj=passthrough_success_handler_obj,
url_route=url_route,
request_body=request_body,
endpoint_type=endpoint_type,
start_time=start_time,
all_chunks=all_chunks,
end_time=end_time,
model=model,
)
)
standard_logging_response_object = ( # rebind-ok: branch bind in shared if/elif dispatch
gemini_passthrough_logging_handler_result["result"]
)
kwargs = ( # rebind-ok: branch bind in shared if/elif dispatch
gemini_passthrough_logging_handler_result["kwargs"]
)
elif endpoint_type == EndpointType.OPENAI:
openai_passthrough_logging_handler_result: Final = (
OpenAIPassthroughLoggingHandler._handle_logging_openai_collected_chunks(

View file

@ -248,7 +248,6 @@ from litellm.constants import (
PROXY_BUDGET_RESCHEDULER_MAX_TIME,
PROXY_BUDGET_RESCHEDULER_MIN_TIME,
PROXY_CONFIG_RELOAD_INTERVAL_SECONDS,
ROUTER_MODEL_NAME_RESPONSE_FIELD,
USER_SPEND_ALERTS_JOB_ID,
WEEKLY_SPEND_REPORT_JOB_ID,
)
@ -8045,10 +8044,6 @@ def _fast_serialize_simple_model_response_stream(
for top_level_key in ("id", "object", "created"):
if payload[top_level_key] is None:
payload.pop(top_level_key)
router_model_name: Final = getattr(chunk, ROUTER_MODEL_NAME_RESPONSE_FIELD, None)
if router_model_name is not None:
payload[ROUTER_MODEL_NAME_RESPONSE_FIELD] = router_model_name
return orjson.dumps(payload)
@ -8346,9 +8341,6 @@ async def async_data_generator(
model_mismatch_logged = False
fallback_metadata_event_sent = False
include_fallback_errors: Final = _should_include_fallback_errors(request_data)
# Fallbacks resolve on the first ``__anext__``, so the selected group is read
# per chunk off this object rather than snapshotted here.
router_logging_obj: Final = request_data.get("litellm_logging_obj")
# Use a running string instead of list + join to avoid O(n^2) overhead.
# Previously "".join(str_so_far_parts) was called every chunk, re-joining
# the entire accumulated response. String += is O(n) amortized total.
@ -8438,10 +8430,6 @@ async def async_data_generator(
fallback_was_attempted=fallback_was_attempted,
fallback_model_from_metadata=fallback_model_from_metadata,
)
ProxyBaseLLMRequestProcessing.set_router_selected_model_field(
response_obj=chunk,
router_model_name=ProxyBaseLLMRequestProcessing.get_router_selected_model_name(router_logging_obj),
)
if strip_stream_usage and _is_injected_stream_usage_artifact(chunk):
if pending_fallback_event:

View file

@ -1402,6 +1402,7 @@ class ProxyLogging:
get_latest_version_prompt_id,
)
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.utils import get_non_default_completion_params
if prompt_version is None:
@ -1420,13 +1421,20 @@ class ProxyLogging:
data.pop("prompt_id", None)
if custom_logger and prompt_spec is not None:
is_responses_call: Final = call_type == "aresponses"
original_responses_input: Final = data.get("input", "") if is_responses_call else ""
client_messages: Final = (
ResponsesAPIRequestUtils.responses_input_to_chat_messages(original_responses_input)
if is_responses_call
else data.get("messages", [])
)
(
model,
messages,
optional_params,
) = await litellm_logging_obj.async_get_chat_completion_prompt(
model=data.get("model", ""),
messages=data.get("messages", []),
messages=client_messages,
non_default_params=get_non_default_completion_params(kwargs=data) or {},
prompt_id=litellm_prompt_id,
prompt_spec=prompt_spec,
@ -1438,7 +1446,14 @@ class ProxyLogging:
data.update(optional_params)
data["model"] = model
data["messages"] = messages
if is_responses_call:
data["input"] = ResponsesAPIRequestUtils.merge_prompt_management_input(
original_input=original_responses_input,
client_input=client_messages,
merged_input=messages,
)
else:
data["messages"] = messages
# prevent re-processing the prompt template
data.pop("prompt_id", None)
data.pop("prompt_variables", None)
@ -1653,7 +1668,7 @@ class ProxyLogging:
not guardrails_only
and litellm_logging_obj is not None
and prompt_id is not None
and (call_type == "completion" or call_type == "acompletion")
and (call_type == "completion" or call_type == "acompletion" or call_type == "aresponses")
):
await self._process_prompt_template(
data=data,

View file

@ -28,7 +28,6 @@ from litellm.responses.litellm_completion_transformation.handler import (
)
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import (
AllMessageValues,
PromptObject,
Reasoning,
ResponseIncludable,
@ -519,10 +518,7 @@ async def aresponses(
if isinstance(
litellm_logging_obj, LiteLLMLoggingObj
) and litellm_logging_obj.should_run_prompt_management_hooks(prompt_id=prompt_id, non_default_params=kwargs):
if isinstance(input, str):
client_input: list[AllMessageValues] = [{"role": "user", "content": input}]
else:
client_input = [item for item in input if isinstance(item, dict) and "role" in item]
client_input: Final = ResponsesAPIRequestUtils.responses_input_to_chat_messages(input)
with _prompt_management_sees_a_provisional_message_list(
kwargs,
bridged=_will_bridge_to_chat_completions(
@ -551,7 +547,13 @@ async def aresponses(
),
)
if model != original_model:
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model)
custom_llm_provider = _resolve_prompt_swapped_provider(
original_model=original_model,
swapped_model=model,
custom_llm_provider=custom_llm_provider,
kwargs=kwargs,
prompt_id=prompt_id,
)
kwargs.pop("prompt_id", None)
kwargs["_async_prompt_merged_params"] = merged_optional_params
@ -621,6 +623,35 @@ async def aresponses(
)
def _resolve_prompt_swapped_provider(
original_model: str,
swapped_model: str,
custom_llm_provider: str | None,
kwargs: Mapping[str, object],
prompt_id: str | None,
) -> str:
swapped_provider: Final = litellm.get_llm_provider(model=swapped_model)[1]
if kwargs.get("api_key") is None and kwargs.get("api_base") is None:
return swapped_provider
try:
original_provider: Final = custom_llm_provider or litellm.get_llm_provider(model=original_model)[1]
except litellm.BadRequestError:
return swapped_provider
if swapped_provider == original_provider:
return swapped_provider
raise litellm.BadRequestError(
message=(
f"prompt_id '{prompt_id}' swaps model '{original_model}' -> '{swapped_model}', which changes the "
f"provider from '{original_provider}' to '{swapped_provider}' after credentials for "
f"'{original_provider}' were already resolved. Refusing to send them to '{swapped_provider}'. "
"Point the request at a model whose provider matches the prompt's metadata.model, or set "
"ignore_prompt_manager_model on the prompt to keep the requested model."
),
model=swapped_model,
llm_provider=swapped_provider,
)
def _apply_prompt_management_to_responses_call(
input: str | ResponseInputParam,
model: str,
@ -640,10 +671,7 @@ def _apply_prompt_management_to_responses_call(
prompt_variables: Final = cast(dict | None, kwargs.get("prompt_variables", None))
original_model: Final = model
if isinstance(input, str):
client_input: list[AllMessageValues] = [{"role": "user", "content": input}]
else:
client_input = [item for item in input if isinstance(item, dict) and "role" in item]
client_input: Final = ResponsesAPIRequestUtils.responses_input_to_chat_messages(input)
if isinstance(litellm_logging_obj, LiteLLMLoggingObj) and litellm_logging_obj.should_run_prompt_management_hooks(
prompt_id=prompt_id, non_default_params=kwargs
@ -676,7 +704,13 @@ def _apply_prompt_management_to_responses_call(
local_vars["input"] = input
local_vars["model"] = model
if model != original_model:
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model)
custom_llm_provider = _resolve_prompt_swapped_provider(
original_model=original_model,
swapped_model=model,
custom_llm_provider=custom_llm_provider,
kwargs=kwargs,
prompt_id=prompt_id,
)
local_vars["custom_llm_provider"] = custom_llm_provider
for key, value in merged_optional_params.items():
local_vars[key] = value
@ -994,6 +1028,33 @@ def responses(
# Update local_vars to include the converted text parameter
local_vars["text"] = text
#########################################################
# PROMPT MANAGEMENT
# If aresponses() already ran the async hook, it pops prompt_id and
# passes the result via _async_prompt_merged_params — apply those
# directly and skip the sync hook to avoid double-merging.
#########################################################
_stripped_model, _from_chat_completions_prefix = _normalize_openai_chat_completions_responses_model(model)
model = _stripped_model
local_vars["model"] = model
use_chat_completions_api = use_chat_completions_api or _from_chat_completions_prefix
if custom_llm_provider is None:
_, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=model, api_base=local_vars.get("base_url", None)
)
local_vars["custom_llm_provider"] = custom_llm_provider
input, model, custom_llm_provider = _apply_prompt_management_to_responses_call(
input=input,
model=model,
custom_llm_provider=custom_llm_provider,
litellm_logging_obj=litellm_logging_obj,
kwargs=kwargs,
local_vars=local_vars,
use_chat_completions_api=use_chat_completions_api,
)
# get llm provider logic
litellm_params: Final = GenericLiteLLMParams(**kwargs)
@ -1003,11 +1064,6 @@ def responses(
if litellm_params.mock_response and isinstance(litellm_params.mock_response, str):
return mock_responses_api_response(mock_response=litellm_params.mock_response)
_stripped_model, _from_chat_completions_prefix = _normalize_openai_chat_completions_responses_model(model)
model = _stripped_model
local_vars["model"] = model
use_chat_completions_api = use_chat_completions_api or _from_chat_completions_prefix
model, custom_llm_provider = _resolve_model_provider_for_responses(
model=model,
custom_llm_provider=custom_llm_provider,
@ -1015,22 +1071,6 @@ def responses(
local_vars=local_vars,
)
#########################################################
# PROMPT MANAGEMENT
# If aresponses() already ran the async hook, it pops prompt_id and
# passes the result via _async_prompt_merged_params — apply those
# directly and skip the sync hook to avoid double-merging.
#########################################################
input, model, custom_llm_provider = _apply_prompt_management_to_responses_call(
input=input,
model=model,
custom_llm_provider=custom_llm_provider,
litellm_logging_obj=litellm_logging_obj,
kwargs=kwargs,
local_vars=local_vars,
use_chat_completions_api=use_chat_completions_api,
)
#########################################################
# Update input and tools with provider-specific file IDs if managed files are used
#########################################################

View file

@ -72,6 +72,16 @@ class ResponsesAPIRequestUtils:
shaped_content: Final = [_as_input_text_part(part) for part in content] # mutable-ok: Responses-shaped copy
return {**message, "content": shaped_content} # mutable-ok: copy, the hook's message stays untouched
@staticmethod
def responses_input_to_chat_messages(
input: str | ResponseInputParam | None,
) -> list[AllMessageValues]:
if input is None:
return []
if isinstance(input, str):
return [{"role": "user", "content": input}]
return [item for item in input if isinstance(item, dict) and "role" in item]
@staticmethod
def merge_prompt_management_input(
original_input: str | ResponseInputParam,

View file

@ -45,7 +45,6 @@ from litellm.caching.caching import (
RedisClusterCache,
)
from litellm.constants import (
AUTO_ROUTED_REQUEST_METADATA_KEY,
CONSUMED_REQUEST_TAGS_METADATA_KEY,
DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS,
DEFAULT_HEALTH_CHECK_INTERVAL,
@ -12164,9 +12163,6 @@ class Router:
self._stamp_or_clear_metadata_key(
request_kwargs=request_kwargs, key=CONSUMED_REQUEST_TAGS_METADATA_KEY, value=None
)
self._stamp_or_clear_metadata_key(
request_kwargs=request_kwargs, key=AUTO_ROUTED_REQUEST_METADATA_KEY, value=None
)
return None
pre_routing_hook_response: Final = await selected_strategy.strategy.async_pre_routing_hook(
@ -12194,13 +12190,6 @@ class Router:
request_tags=_get_tags_from_request_kwargs(request_kwargs),
),
)
# Gates the proxy's `router_model_name` response field; the body `model` is
# always restamped back to the alias the client sent.
self._stamp_or_clear_metadata_key(
request_kwargs=request_kwargs,
key=AUTO_ROUTED_REQUEST_METADATA_KEY,
value=(True if pre_routing_hook_response is not None else None),
)
# `model` (the alias, e.g. "smart-router") is never the deployment actually
# called - apply the router marker's own litellm_params to the request,

View file

@ -93,26 +93,32 @@ class SlackAlertingArgs(LiteLLMPydanticObjectBase):
)
daily_spend_per_user_threshold: float | None = Field(
default=None,
gt=0,
description="Alert when a user's spend for the current day (UTC) crosses this USD amount. Off by default.",
)
monthly_spend_per_user_threshold: float | None = Field(
default=None,
gt=0,
description="Alert when a user's spend for the current calendar month (UTC) crosses this USD amount. Off by default.",
)
spend_anomaly_multiplier: float = Field(
default=3.0,
gt=0,
description="Flag a user's spend as anomalous when today's spend exceeds this multiple of their trailing daily average.",
)
spend_anomaly_baseline_days: int = Field(
default=7,
ge=1,
description="Number of trailing days used to compute a user's daily average spend for anomaly detection.",
)
spend_anomaly_min_spend: float = Field(
default=10.0,
gt=0,
description="Minimum spend (USD) a user must reach today before an anomaly alert can fire. Reduces false positives.",
)
user_spend_check_interval: int = Field(
default=3600,
ge=60,
description="How often (in seconds) to check per-user spend thresholds and anomalies. Default is hourly.",
)
@ -209,7 +215,6 @@ DEFAULT_ALERT_TYPES: Final[list[AlertType]] = [
AlertType.spend_reports,
AlertType.failed_tracking_spend,
AlertType.user_spend_thresholds,
AlertType.user_spend_anomalies,
# Database alerts
AlertType.db_exceptions,
# Report alerts

View file

@ -22,6 +22,7 @@ LITELLM_PASS_THROUGH_ENDPOINT_MARKER: Final = "__litellm_pass_through_endpoint__
class EndpointType(str, Enum):
VERTEX_AI = "vertex-ai"
GEMINI = "gemini"
ANTHROPIC = "anthropic"
OPENAI = "openai"
GENERIC = "generic"

View file

@ -106,6 +106,7 @@ utils = [
"numpydoc>=1.8.0,<2.0",
]
caching = ["diskcache>=5.6.3,<6.0"]
mcp = ["mcp>=1.28.1,<2.0"]
# SAML SSO for the admin UI. python3-saml pulls in xmlsec/lxml, whose wheels
# bundle the native libxmlsec1/libxml2 libraries, so no system packages are
# required. Kept out of the base `proxy` extra so it stays optional.

View file

@ -2,6 +2,8 @@ import asyncio
import base64
import os
import sys
from importlib import metadata
from pathlib import Path
from unittest.mock import AsyncMock, MagicMock, patch
import anyio
@ -24,9 +26,11 @@ from mcp.types import (
import litellm.experimental_mcp_client.client as mcp_client_module
from litellm.experimental_mcp_client.client import (
MCP_STREAMABLE_HTTP_REQUIREMENT,
MCPClient,
_as_read_timeout,
_first_non_cancelled_cause,
missing_streamable_http_client_error,
strip_auth_scheme,
)
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
@ -1047,3 +1051,47 @@ def test_openapi_byok_auth_header_emits_exactly_one_scheme(auth_type, auth_value
assert server.is_byok is False
assert _format_byok_openapi_auth_header(server, auth_value) == expected
def test_missing_streamable_http_client_error_names_requirement_and_remedy():
message = str(missing_streamable_http_client_error())
assert MCP_STREAMABLE_HTTP_REQUIREMENT in message
assert "pip install 'litellm[mcp]'" in message
assert metadata.version("mcp") in message
@pytest.mark.asyncio
async def test_http_transport_without_streamable_http_client_raises_actionable_import_error():
client = MCPClient(
server_url="https://mcp-server.example.com",
transport_type=MCPTransport.http,
)
with patch.object( # test-quality-ok: simulates mcp<1.24.0 whose module lacks this import-time symbol
mcp_client_module, "streamable_http_client", None
):
with pytest.raises(ImportError, match=r"pip install 'litellm\[mcp\]'"):
await client.list_tools(raise_on_error=True)
def test_mcp_extra_matches_proxy_extra_and_supports_streamable_http():
try:
import tomllib
except ImportError:
tomllib = pytest.importorskip("tomli")
from packaging.requirements import Requirement
pyproject_path = Path(__file__).parents[3] / "pyproject.toml"
with pyproject_path.open("rb") as f:
extras = tomllib.load(f)["project"]["optional-dependencies"]
mcp_extra = extras["mcp"]
assert len(mcp_extra) == 1
proxy_mcp_requirements = [req for req in extras["proxy"] if Requirement(req).name == "mcp"]
assert mcp_extra == proxy_mcp_requirements
specifier = Requirement(mcp_extra[0]).specifier
assert not specifier.contains("1.23.0")
assert specifier.contains("1.28.1")

View file

@ -8,6 +8,36 @@ from litellm.google_genai.streaming_iterator import (
GoogleGenAIGenerateContentStreamingIterator,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType
@pytest.mark.parametrize(
"custom_llm_provider, expected_endpoint_type",
[("gemini", EndpointType.GEMINI), ("vertex_ai", EndpointType.VERTEX_AI)],
)
@pytest.mark.parametrize(
"iterator_cls",
[
AsyncGoogleGenAIGenerateContentStreamingIterator,
GoogleGenAIGenerateContentStreamingIterator,
],
)
def test_streaming_logging_targets_the_provider_that_served_the_request(
iterator_cls: type,
custom_llm_provider: str,
expected_endpoint_type: EndpointType,
):
"""Routing every google stream through the vertex handler bills gemini/* at vertex_ai/ rates."""
iterator = iterator_cls(
response=MagicMock(),
model="gemini-3.1-flash-image",
logging_obj=MagicMock(spec=LiteLLMLoggingObj),
generate_content_provider_config=MagicMock(),
litellm_metadata={},
custom_llm_provider=custom_llm_provider,
)
assert iterator.endpoint_type is expected_endpoint_type
def _large_inline_data_event() -> str:
@ -53,9 +83,7 @@ async def test_async_streaming_iterator_yields_complete_sse_events():
assert chunk.startswith(b"data: ")
assert chunk.endswith(b"\n\n")
assert (
json.loads(chunk[len(b"data: ") : -2])["candidates"][0]["content"]["parts"][0][
"inlineData"
]["mimeType"]
json.loads(chunk[len(b"data: ") : -2])["candidates"][0]["content"]["parts"][0]["inlineData"]["mimeType"]
== "image/jpeg"
)
@ -76,9 +104,9 @@ def test_sync_streaming_iterator_yields_complete_sse_events():
chunk = next(iterator)
assert chunk.startswith(b"data: ")
assert chunk.endswith(b"\n\n")
assert json.loads(chunk[len(b"data: ") : -2])["candidates"][0]["content"]["parts"][
0
]["inlineData"]["data"].startswith("A")
assert json.loads(chunk[len(b"data: ") : -2])["candidates"][0]["content"]["parts"][0]["inlineData"][
"data"
].startswith("A")
@pytest.mark.asyncio

View file

@ -9,7 +9,11 @@ from litellm.integrations.SlackAlerting.user_spend_alerts import (
UserSpendRow,
evaluate_user_spend,
)
from litellm.types.integrations.slack_alerting import AlertType, SlackAlertingArgs
from litellm.types.integrations.slack_alerting import (
DEFAULT_ALERT_TYPES,
AlertType,
SlackAlertingArgs,
)
TODAY: Final = datetime.date(2026, 8, 15)
@ -18,14 +22,12 @@ def _row(
daily_spend: float = 0.0,
monthly_spend: float = 0.0,
baseline_spend: float = 0.0,
baseline_days: int = 0,
) -> UserSpendRow:
return UserSpendRow(
user_id="user-1",
daily_spend=daily_spend,
monthly_spend=monthly_spend,
baseline_spend=baseline_spend,
baseline_days=baseline_days,
)
@ -78,7 +80,7 @@ def test_thresholds_disabled_suppresses_threshold_events():
def test_anomaly_detected_above_multiple_of_baseline():
args: Final = SlackAlertingArgs(spend_anomaly_multiplier=3.0, spend_anomaly_min_spend=10.0)
events: Final = _evaluate(
_row(daily_spend=70.0, monthly_spend=100.0, baseline_spend=70.0, baseline_days=7), args
_row(daily_spend=70.0, monthly_spend=100.0, baseline_spend=70.0), args
)
assert [e.kind for e in events] == ["anomaly"]
assert events[0].alert_type == AlertType.user_spend_anomalies
@ -89,13 +91,13 @@ def test_anomaly_detected_above_multiple_of_baseline():
def test_no_anomaly_within_baseline_multiple():
args: Final = SlackAlertingArgs(spend_anomaly_multiplier=3.0, spend_anomaly_min_spend=10.0)
assert (
_evaluate(_row(daily_spend=25.0, monthly_spend=100.0, baseline_spend=70.0, baseline_days=7), args) == ()
_evaluate(_row(daily_spend=25.0, monthly_spend=100.0, baseline_spend=70.0), args) == ()
)
def test_no_anomaly_below_min_spend_floor():
args: Final = SlackAlertingArgs(spend_anomaly_multiplier=3.0, spend_anomaly_min_spend=10.0)
assert _evaluate(_row(daily_spend=9.0, monthly_spend=9.0, baseline_spend=0.1, baseline_days=1), args) == ()
assert _evaluate(_row(daily_spend=9.0, monthly_spend=9.0, baseline_spend=0.1), args) == ()
def test_anomaly_for_new_user_without_baseline():
@ -104,6 +106,28 @@ def test_anomaly_for_new_user_without_baseline():
assert [e.kind for e in events] == ["anomaly"]
def test_sparse_baseline_averages_over_full_window():
args: Final = SlackAlertingArgs(
spend_anomaly_multiplier=3.0, spend_anomaly_min_spend=10.0, spend_anomaly_baseline_days=7
)
events: Final = _evaluate(_row(daily_spend=13.0, monthly_spend=20.0, baseline_spend=7.0), args)
assert [e.kind for e in events] == ["anomaly"]
def test_anomalies_not_in_default_alert_types():
assert AlertType.user_spend_anomalies not in DEFAULT_ALERT_TYPES
assert AlertType.user_spend_thresholds in DEFAULT_ALERT_TYPES
def test_invalid_config_rejected():
with pytest.raises(ValueError):
SlackAlertingArgs(daily_spend_per_user_threshold=0)
with pytest.raises(ValueError):
SlackAlertingArgs(spend_anomaly_baseline_days=0)
with pytest.raises(ValueError):
SlackAlertingArgs(user_spend_check_interval=10)
def test_anomalies_disabled_suppresses_anomaly_events():
args: Final = SlackAlertingArgs(spend_anomaly_multiplier=3.0, spend_anomaly_min_spend=10.0)
assert _evaluate(_row(daily_spend=500.0, monthly_spend=500.0), args, anomalies=False) == ()
@ -123,14 +147,12 @@ async def test_send_user_spend_alerts_sends_and_dedupes():
"daily_spend": 75.0,
"monthly_spend": 75.0,
"baseline_spend": 0.0,
"baseline_days": 0,
},
{
"user_id": "user-2",
"daily_spend": 60.0,
"monthly_spend": 60.0,
"baseline_spend": 0.0,
"baseline_days": 0,
},
]
)

View file

@ -12,7 +12,9 @@ from unittest.mock import MagicMock, Mock, patch
import httpx
import litellm
from litellm.integrations.dotprompt.dotprompt_manager import DotpromptManager
from litellm.integrations.dotprompt.prompt_manager import PromptManager, PromptTemplate
from litellm.types.prompts.init_prompts import PromptLiteLLMParams, PromptSpec
def test_prompt_manager_initialization():
@ -657,3 +659,90 @@ def test_prompt_initializer_registers_flat_db_prompt_under_base_id():
template = dotprompt_manager.prompt_manager.get_prompt("agent-prompt")
assert template is not None
assert template.content == "AHOY {{name}}"
def _swap_prompt_manager_and_spec(ignore_prompt_manager_model: bool) -> tuple[DotpromptManager, PromptSpec]:
manager = DotpromptManager(
prompt_data={"content": "You are a pirate assistant.", "metadata": {"model": "gpt-4o-mini"}},
prompt_id="swap-prompt",
)
spec = PromptSpec(
prompt_id="swap-prompt",
litellm_params=PromptLiteLLMParams(
prompt_id="swap-prompt",
prompt_integration="dotprompt",
ignore_prompt_manager_model=ignore_prompt_manager_model,
),
)
return manager, spec
@pytest.mark.asyncio
async def test_async_prompt_spec_ignore_prompt_manager_model_keeps_requested_model():
from litellm.types.utils import StandardCallbackDynamicParams
manager, spec = _swap_prompt_manager_and_spec(ignore_prompt_manager_model=True)
model, messages, _ = await manager.async_get_chat_completion_prompt(
model="anthropic/claude-haiku-4-5",
messages=[{"role": "user", "content": "hi"}],
non_default_params={},
prompt_id="swap-prompt",
prompt_variables=None,
dynamic_callback_params=StandardCallbackDynamicParams(),
litellm_logging_obj=MagicMock(),
prompt_spec=spec,
)
assert model == "anthropic/claude-haiku-4-5"
assert len(messages) == 2
assert "pirate" in str(messages[0]["content"])
@pytest.mark.asyncio
async def test_async_prompt_spec_without_ignore_flag_swaps_model():
from litellm.types.utils import StandardCallbackDynamicParams
manager, spec = _swap_prompt_manager_and_spec(ignore_prompt_manager_model=False)
model, _, _ = await manager.async_get_chat_completion_prompt(
model="anthropic/claude-haiku-4-5",
messages=[{"role": "user", "content": "hi"}],
non_default_params={},
prompt_id="swap-prompt",
prompt_variables=None,
dynamic_callback_params=StandardCallbackDynamicParams(),
litellm_logging_obj=MagicMock(),
prompt_spec=spec,
)
assert model == "gpt-4o-mini"
def test_sync_prompt_spec_ignore_prompt_manager_model_keeps_requested_model():
from litellm.types.utils import StandardCallbackDynamicParams
manager, spec = _swap_prompt_manager_and_spec(ignore_prompt_manager_model=True)
model, _, _ = manager.get_chat_completion_prompt(
model="anthropic/claude-haiku-4-5",
messages=[{"role": "user", "content": "hi"}],
non_default_params={},
prompt_id="swap-prompt",
prompt_variables=None,
dynamic_callback_params=StandardCallbackDynamicParams(),
prompt_spec=spec,
)
assert model == "anthropic/claude-haiku-4-5"
def test_sync_caller_ignore_flag_survives_missing_prompt_spec():
from litellm.types.utils import StandardCallbackDynamicParams
manager, _ = _swap_prompt_manager_and_spec(ignore_prompt_manager_model=False)
model, _, _ = manager.get_chat_completion_prompt(
model="anthropic/claude-haiku-4-5",
messages=[{"role": "user", "content": "hi"}],
non_default_params={},
prompt_id="swap-prompt",
prompt_variables=None,
dynamic_callback_params=StandardCallbackDynamicParams(),
prompt_spec=None,
ignore_prompt_manager_model=True,
)
assert model == "anthropic/claude-haiku-4-5"

View file

@ -2764,6 +2764,46 @@ def test_token_type_cost_breakdown_matches_real_gemini_numbers(_local_model_cost
assert breakdown.cache_creation_cost == 0.0
def test_token_type_cost_breakdown_flex_tier_prices_reasoning_at_flex_rate(_local_model_cost_map):
"""Regression for the flex-tier breakdown drift: gemini-3.5-flash defines a flat
output_cost_per_reasoning_token (9e-06, the standard output rate) but no _flex
variant, so the breakdown priced reasoning at the standard rate on flex requests
while the total billed it at the flex output rate (4.5e-06). The reasoning
sub-cost then exceeded the entire flex completion cost."""
usage = Usage(
prompt_tokens=7,
completion_tokens=320,
total_tokens=327,
completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=315, text_tokens=5),
)
breakdown = get_token_type_cost_breakdown(
model="gemini-3.5-flash",
custom_llm_provider="vertex_ai",
usage=usage,
service_tier="flex",
)
assert breakdown.reasoning_cost == pytest.approx(315 * 4.5e-06)
_, flex_completion_cost = generic_cost_per_token(
model="gemini-3.5-flash",
usage=usage,
custom_llm_provider="vertex_ai",
service_tier="flex",
)
assert breakdown.reasoning_cost <= flex_completion_cost
standard_breakdown = get_token_type_cost_breakdown(
model="gemini-3.5-flash",
custom_llm_provider="vertex_ai",
usage=usage,
service_tier=None,
)
assert standard_breakdown.reasoning_cost == pytest.approx(315 * 9e-06)
def test_token_type_cost_breakdown_xai_at_exactly_200k_uses_higher_tier_rates(_local_model_cost_map):
usage = Usage(

View file

@ -92,6 +92,39 @@ class TestCallbackDurationMs:
assert hidden.get("litellm_overhead_time_ms") is not None
class TestDictResultsSkipMetadataUpdate:
"""Regression for /v1/messages cost-breakdown clobbering: AnthropicMessagesResponse
is a TypedDict, so apply() can never attach _hidden_params to it and the whole
metadata pass is discarded - except the cost recompute, whose only observable
effect was overwriting the logging object's already-correct cost breakdown with a
service-tier-less, reasoning-less recompute on the adapted response."""
def test_update_response_metadata_skips_cost_recompute_for_dict_results(self):
anthropic_response = {
"id": "msg_123",
"type": "message",
"role": "assistant",
"content": [{"type": "text", "text": "hi"}],
"usage": {"input_tokens": 7, "output_tokens": 320},
}
logging_obj = MagicMock()
logging_obj.model_call_details = {}
logging_obj.caching_details = None
logging_obj.litellm_call_id = "test-call-id"
update_response_metadata(
result=anthropic_response,
logging_obj=logging_obj,
model="vertex_ai/gemini-3.5-flash",
kwargs={},
start_time=datetime.datetime(2025, 1, 1, 0, 0, 0),
end_time=datetime.datetime(2025, 1, 1, 0, 0, 1),
)
logging_obj._response_cost_calculator.assert_not_called()
assert "_hidden_params" not in anthropic_response
class TestCallbackDurationInCustomHeaders:
"""Test that callback_duration_ms flows into get_custom_headers."""

View file

@ -0,0 +1,130 @@
import json
from collections.abc import Iterator
from datetime import datetime
from unittest.mock import MagicMock
import pytest
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler import (
VertexPassthroughLoggingHandler,
)
from litellm.proxy.pass_through_endpoints.streaming_handler import (
PassThroughStreamingHandler,
)
from litellm.proxy.pass_through_endpoints.success_handler import (
PassThroughEndpointLogging,
)
from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType
MODEL = "gemini-stream-pricing-probe"
PROMPT_TOKENS = 1000
COMPLETION_TOKENS = 1000
GEMINI_INPUT_RATE = 1e-07
GEMINI_OUTPUT_RATE = 4e-07
VERTEX_INPUT_RATE = 1.5e-07
VERTEX_OUTPUT_RATE = 6e-07
GEMINI_COST = PROMPT_TOKENS * GEMINI_INPUT_RATE + COMPLETION_TOKENS * GEMINI_OUTPUT_RATE
VERTEX_COST = PROMPT_TOKENS * VERTEX_INPUT_RATE + COMPLETION_TOKENS * VERTEX_OUTPUT_RATE
@pytest.fixture(autouse=True)
def divergent_rate_cards(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
monkeypatch.setitem(
litellm.model_cost,
f"gemini/{MODEL}",
{
"input_cost_per_token": GEMINI_INPUT_RATE,
"output_cost_per_token": GEMINI_OUTPUT_RATE,
"litellm_provider": "gemini",
"mode": "chat",
},
)
monkeypatch.setitem(
litellm.model_cost,
f"vertex_ai/{MODEL}",
{
"input_cost_per_token": VERTEX_INPUT_RATE,
"output_cost_per_token": VERTEX_OUTPUT_RATE,
"litellm_provider": "vertex_ai",
"mode": "chat",
},
)
litellm.get_model_info.cache_clear()
yield
litellm.get_model_info.cache_clear()
def _chunks() -> list[str]:
payload = {
"candidates": [
{
"content": {"parts": [{"text": "hi"}], "role": "model"},
"finishReason": "STOP",
"index": 0,
}
],
"usageMetadata": {
"promptTokenCount": PROMPT_TOKENS,
"candidatesTokenCount": COMPLETION_TOKENS,
"totalTokenCount": PROMPT_TOKENS + COMPLETION_TOKENS,
},
"modelVersion": MODEL,
}
return [f"data: {json.dumps(payload)}"]
def _logging_obj() -> LiteLLMLoggingObj:
logging_obj = MagicMock(spec=LiteLLMLoggingObj)
logging_obj.model_call_details = {}
logging_obj.optional_params = {}
logging_obj.litellm_call_id = "test-call-id"
return logging_obj
@pytest.mark.parametrize(
"endpoint_type, expected_provider, expected_cost",
[
(EndpointType.GEMINI, "gemini", GEMINI_COST),
(EndpointType.VERTEX_AI, "vertex_ai", VERTEX_COST),
],
)
def test_streaming_generate_content_bills_against_the_requested_provider(
endpoint_type, expected_provider, expected_cost
):
logging_obj = _logging_obj()
_, kwargs = PassThroughStreamingHandler._build_passthrough_logging_result(
litellm_logging_obj=logging_obj,
passthrough_success_handler_obj=PassThroughEndpointLogging(),
url_route="/v1/generateContent",
request_body={},
endpoint_type=endpoint_type,
start_time=datetime.now(),
raw_bytes=[chunk.encode("utf-8") for chunk in _chunks()],
end_time=datetime.now(),
model=MODEL,
)
assert kwargs["response_cost"] == pytest.approx(expected_cost)
assert logging_obj.model_call_details["custom_llm_provider"] == expected_provider
def test_vertex_generate_content_payload_prices_gemini_urls_at_gemini_rates():
logging_obj = _logging_obj()
result = VertexPassthroughLoggingHandler._handle_logging_vertex_collected_chunks(
litellm_logging_obj=logging_obj,
passthrough_success_handler_obj=PassThroughEndpointLogging(),
url_route=f"https://generativelanguage.googleapis.com/v1beta/models/{MODEL}:streamGenerateContent",
request_body={},
endpoint_type=EndpointType.VERTEX_AI,
start_time=datetime.now(),
all_chunks=_chunks(),
model=MODEL,
end_time=datetime.now(),
)
assert result["kwargs"]["response_cost"] == pytest.approx(GEMINI_COST)
assert logging_obj.model_call_details["custom_llm_provider"] == "gemini"

View file

@ -13,11 +13,7 @@ from fastapi.responses import JSONResponse, StreamingResponse
import litellm
from litellm._uuid import uuid
from litellm.constants import (
AUTO_ROUTED_REQUEST_METADATA_KEY,
RETURN_RAW_MODEL_NAME_METADATA_KEY,
ROUTER_MODEL_NAME_RESPONSE_FIELD,
)
from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.opentelemetry import UserAPIKeyAuth
from litellm.proxy.common_request_processing import (
@ -7223,128 +7219,6 @@ async def test_a_broken_hook_does_not_replace_the_real_error_with_its_own_bug():
assert "audit backend" not in collected[-2].decode()
class TestRouterModelNameOnNonStreamingResponse:
"""
The proxy restamps the response body `model` back to the client-requested
alias, so an auto-routed request (auto_router / complexity_router /
adaptive_router / quality_router) had no body-level surface naming the model
group that actually served it. `router_model_name` is now set on the response
whenever the router marked the request as auto-routed.
"""
@staticmethod
def _logging_obj(*, metadata_bucket, bucket_name="metadata"):
logging_obj = MagicMock()
logging_obj.litellm_call_id = "call-auto-routed"
logging_obj.cost_breakdown = None
logging_obj.model_call_details = {}
logging_obj.litellm_params = {bucket_name: metadata_bucket}
logging_obj._enqueue_deferred_logging = None
logging_obj._on_deferred_stream_complete = None
return logging_obj
async def _drive(self, *, monkeypatch, logging_obj):
import litellm.proxy.common_request_processing as crp
from litellm.proxy._types import UserAPIKeyAuth as RealUserAPIKeyAuth
from litellm.types.utils import ModelResponse
response = ModelResponse(
model="deep-model",
choices=[{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
)
async def fake_route_request(**kwargs):
async def _llm_call():
return response
return _llm_call()
monkeypatch.setattr(crp, "route_request", fake_route_request)
async def fake_post_call_success_hook(data, user_api_key_dict, response):
return response
proxy_logging_obj = MagicMock(spec=ProxyLogging)
proxy_logging_obj.during_call_hook = AsyncMock(return_value=None)
proxy_logging_obj.update_request_status = AsyncMock(return_value=None)
proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value={})
proxy_logging_obj.post_call_success_hook = fake_post_call_success_hook
processing_obj = ProxyBaseLLMRequestProcessing(
data={"model": "smart-route", "litellm_logging_obj": logging_obj}
)
with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails", return_value=False):
return await processing_obj.base_process_llm_request(
request=MagicMock(spec=Request, headers={}),
fastapi_response=Response(),
user_api_key_dict=RealUserAPIKeyAuth(api_key="sk-test"),
route_type="acompletion",
proxy_logging_obj=proxy_logging_obj,
general_settings={},
proxy_config=MagicMock(spec=ProxyConfig),
select_data_generator=None,
llm_router=None,
skip_pre_call_logic=True,
)
@pytest.mark.asyncio
async def test_auto_routed_request_carries_router_model_name(self, monkeypatch):
result = await self._drive(
monkeypatch=monkeypatch,
logging_obj=self._logging_obj(
metadata_bucket={
AUTO_ROUTED_REQUEST_METADATA_KEY: True,
"deployment_model_name": "deep-model",
}
),
)
assert result.model == "smart-route"
assert result.model_dump(exclude_none=True, exclude_unset=True)[ROUTER_MODEL_NAME_RESPONSE_FIELD] == (
"deep-model"
)
@pytest.mark.asyncio
async def test_marker_and_model_name_in_different_buckets(self, monkeypatch):
logging_obj = self._logging_obj(metadata_bucket={AUTO_ROUTED_REQUEST_METADATA_KEY: True})
logging_obj.litellm_params["litellm_metadata"] = {"deployment_model_name": "deep-model"}
result = await self._drive(monkeypatch=monkeypatch, logging_obj=logging_obj)
assert result.model_dump(exclude_none=True, exclude_unset=True)[ROUTER_MODEL_NAME_RESPONSE_FIELD] == (
"deep-model"
)
@pytest.mark.asyncio
async def test_plain_model_group_request_has_no_router_model_name(self, monkeypatch):
result = await self._drive(
monkeypatch=monkeypatch,
logging_obj=self._logging_obj(metadata_bucket={"deployment_model_name": "deep-model"}),
)
assert ROUTER_MODEL_NAME_RESPONSE_FIELD not in result.model_dump(exclude_none=True, exclude_unset=True)
@pytest.mark.asyncio
async def test_typeddict_response_gets_router_model_name(self):
from litellm.types.utils import AnthropicMessagesResponse
response: AnthropicMessagesResponse = {"id": "msg_1", "model": "smart-route", "type": "message"}
ProxyBaseLLMRequestProcessing.set_router_selected_model_field(
response_obj=response,
router_model_name=ProxyBaseLLMRequestProcessing.get_router_selected_model_name(
self._logging_obj(
metadata_bucket={
AUTO_ROUTED_REQUEST_METADATA_KEY: True,
"deployment_model_name": "deep-model",
}
)
),
)
assert response[ROUTER_MODEL_NAME_RESPONSE_FIELD] == "deep-model"
@pytest.mark.parametrize(
"exc,expect_traceback",
[

View file

@ -11535,157 +11535,6 @@ class TestEmbeddingsFailureHookRequestData:
assert hook_request_data["litellm_logging_obj"] is logging_obj_sentinel
class TestRouterModelNameOnStreamingChunks:
"""
Streaming chunks get the body `model` restamped to the client-requested alias
just like non-streaming responses, so an auto-routed request had no way to
name the model group that served it without reading response headers. Every
emitted chunk now carries `router_model_name`.
These assert on the serialized SSE bytes, not on the chunk objects. The fast
path (`_fast_serialize_simple_model_response_stream`) hand-builds a
closed-set dict, so a chunk object can carry the field while the wire drops
it, and an object-level assertion would pass against that bug.
"""
@staticmethod
def _chunk(*, with_usage=False):
from litellm.types.utils import ModelResponseStream
return ModelResponseStream(
model="smart-route",
choices=[{"index": 0, "delta": {"role": "assistant", "content": "hi"}}],
usage={"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2} if with_usage else None,
)
@staticmethod
def _request_data(*, auto_routed):
from litellm.constants import AUTO_ROUTED_REQUEST_METADATA_KEY
logging_obj = MagicMock()
logging_obj.litellm_params = {
"metadata": {
**({AUTO_ROUTED_REQUEST_METADATA_KEY: True} if auto_routed else {}),
"deployment_model_name": "deep-model",
}
}
return {"model": "smart-route", "litellm_logging_obj": logging_obj}
async def _drive(self, *, chunks, request_data, on_yield=None):
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import async_data_generator
from litellm.proxy.utils import ProxyLogging
class MockStream:
def __aiter__(self):
return self._stream()
async def _stream(self):
for index, chunk in enumerate(chunks):
if on_yield is not None:
on_yield(index)
yield chunk
mock_response = MockStream()
mock_response.aclose = AsyncMock()
proxy_logging_obj = MagicMock(spec=ProxyLogging)
proxy_logging_obj.has_streaming_callbacks.return_value = False
proxy_logging_obj.needs_iterator_wrap.return_value = False
proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock()
proxy_logging_obj.async_post_call_streaming_hook = AsyncMock()
proxy_logging_obj.post_call_failure_hook = AsyncMock()
with patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj):
with patch.object(ProxyLogging, "_fire_deferred_stream_logging"):
return [
data
async for data in async_data_generator(mock_response, MagicMock(spec=UserAPIKeyAuth), request_data)
]
@staticmethod
def _data_frames(emitted):
return [
frame.decode() if isinstance(frame, bytes) else frame
for frame in emitted
if b"[DONE]" not in (frame if isinstance(frame, bytes) else frame.encode())
]
@pytest.mark.asyncio
async def test_fast_path_chunk_carries_router_model_name_on_the_wire(self):
emitted = await self._drive(chunks=[self._chunk()], request_data=self._request_data(auto_routed=True))
frames = self._data_frames(emitted)
assert frames
assert all('"router_model_name":"deep-model"' in frame for frame in frames)
assert all('"model":"smart-route"' in frame for frame in frames)
@pytest.mark.asyncio
async def test_slow_path_chunk_carries_router_model_name_on_the_wire(self):
emitted = await self._drive(
chunks=[self._chunk(with_usage=True)], request_data=self._request_data(auto_routed=True)
)
frames = self._data_frames(emitted)
assert frames
assert all('"router_model_name":"deep-model"' in frame for frame in frames)
@pytest.mark.asyncio
async def test_plain_model_group_stream_has_no_router_model_name(self):
emitted = await self._drive(
chunks=[self._chunk(), self._chunk(with_usage=True)],
request_data=self._request_data(auto_routed=False),
)
frames = self._data_frames(emitted)
assert frames
assert all("router_model_name" not in frame for frame in frames)
@pytest.mark.asyncio
async def test_fallback_out_of_the_routed_group_drops_the_field(self):
from litellm.constants import AUTO_ROUTED_REQUEST_METADATA_KEY
request_data = self._request_data(auto_routed=True)
bucket = request_data["litellm_logging_obj"].litellm_params["metadata"]
def fall_back(index):
if index == 1:
bucket.pop(AUTO_ROUTED_REQUEST_METADATA_KEY)
bucket["deployment_model_name"] = "backup-model"
emitted = await self._drive(
chunks=[self._chunk(), self._chunk(), self._chunk()],
request_data=request_data,
on_yield=fall_back,
)
frames = self._data_frames(emitted)
assert len(frames) >= 3
assert '"router_model_name":"deep-model"' in frames[0]
assert all("router_model_name" not in frame for frame in frames[1:])
@pytest.mark.asyncio
async def test_fallback_to_another_auto_router_reports_the_new_tier(self):
request_data = self._request_data(auto_routed=True)
bucket = request_data["litellm_logging_obj"].litellm_params["metadata"]
def fall_back(index):
if index == 1:
bucket["deployment_model_name"] = "backup-tier"
emitted = await self._drive(
chunks=[self._chunk(), self._chunk(), self._chunk()],
request_data=request_data,
on_yield=fall_back,
)
frames = self._data_frames(emitted)
assert len(frames) >= 3
assert '"router_model_name":"deep-model"' in frames[0]
assert all('"router_model_name":"backup-tier"' in frame for frame in frames[1:])
@pytest.mark.asyncio
async def test_authoritative_floor_spend_keeps_a_reset_marker_written_during_the_db_read():
"""A team-member spend reset writes the post-reset floor to the spend_db_floor marker

View file

@ -818,3 +818,50 @@ async def test_process_prompt_template_async_get_prompt_error_raises(proxy_loggi
prompt_version=None,
call_type="completion",
)
@pytest.mark.asyncio
async def test_process_prompt_template_aresponses_swaps_model_and_merges_input(proxy_logging, monkeypatch):
from litellm.proxy.prompts import prompt_registry
custom_logger = MagicMock()
prompt_spec = MagicMock()
prompt_spec.litellm_params = MagicMock(prompt_id="resolved-id")
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY,
"get_prompt_callback_by_id",
lambda *a, **kw: custom_logger,
)
monkeypatch.setattr(
prompt_registry.IN_MEMORY_PROMPT_REGISTRY, "get_prompt_by_id", lambda *a, **kw: prompt_spec
)
logging_obj = MagicMock()
logging_obj.async_get_chat_completion_prompt = AsyncMock(
return_value=(
"gpt-4o-mini",
[
{"role": "user", "content": "You are a pirate."},
{"role": "user", "content": "Who are you?"},
],
{},
)
)
data: dict[str, object] = {"input": "Who are you?", "model": "anthropic-haiku-4-5", "prompt_id": "x"}
await proxy_logging._process_prompt_template(
data=data,
litellm_logging_obj=logging_obj,
prompt_id="x",
prompt_version=None,
call_type="aresponses",
)
assert data["model"] == "gpt-4o-mini"
assert data["input"] == [
{"role": "user", "content": "You are a pirate."},
{"role": "user", "content": "Who are you?"},
]
assert "messages" not in data
assert "prompt_id" not in data
hook_kwargs = logging_obj.async_get_chat_completion_prompt.await_args.kwargs
assert hook_kwargs["messages"] == [{"role": "user", "content": "Who are you?"}]
assert hook_kwargs["prompt_spec"] is prompt_spec

View file

@ -298,6 +298,22 @@ async def test_default_path_still_applies_prompt_templates(proxy_logging, make_u
process.assert_awaited_once()
@pytest.mark.asyncio
async def test_aresponses_call_type_applies_prompt_templates_before_routing(proxy_logging, make_user_api_key_auth, monkeypatch):
"""The responses surface must process registry prompts pre-routing so credentials follow the swapped model."""
monkeypatch.setattr(litellm, "callbacks", [])
proxy_logging.slack_alerting_instance = MagicMock(alerting=None)
process = AsyncMock()
monkeypatch.setattr(proxy_logging, "_process_prompt_template", process)
await proxy_logging.pre_call_hook(
user_api_key_dict=make_user_api_key_auth(),
data={"input": "hi", "model": "m", "prompt_id": "p1", "litellm_logging_obj": MagicMock()},
call_type="aresponses",
)
process.assert_awaited_once()
# ---------------------------------------------------------------------------
# enforces_request_content: which CustomLoggers a guardrails-only walk reaches
# ---------------------------------------------------------------------------

View file

@ -52,12 +52,19 @@ def _make_logging_obj(
return logging_obj
def _provider_by_model(model: str, **_: object) -> tuple[str, str, None, None]:
provider, _, bare_model = model.partition("/")
if not bare_model:
return (model, "anthropic" if "claude" in model else "openai", None, None)
return (bare_model, provider, None, None)
def _patch_responses_dispatch():
"""Patch everything after the prompt management block so tests stay unit-level."""
return [
patch(
"litellm.responses.main.litellm.get_llm_provider",
return_value=("gpt-4o", "openai", None, None),
side_effect=_provider_by_model,
),
patch(
"litellm.responses.mcp.litellm_proxy_mcp_handler."
@ -278,7 +285,7 @@ class TestResponsesAPIPromptManagement:
# The model passed to the downstream handler should be the overridden one
handler_call_kwargs = mock_handler.call_args.kwargs
assert handler_call_kwargs.get("model") == "openai/gpt-4o-mini"
assert handler_call_kwargs.get("model") == "gpt-4o-mini"
def test_non_message_input_items_filtered(self):
"""[F] Non-message items in ResponseInputParam (e.g. function_call_output) are
@ -388,10 +395,7 @@ class TestResponsesAPIPromptManagement:
with (
patch(
"litellm.responses.main.litellm.get_llm_provider",
side_effect=[
("gpt-4o", "openai", None, None),
("claude-3-5-sonnet", "anthropic", None, None),
],
side_effect=_provider_by_model,
),
patches[1],
patches[2],
@ -539,3 +543,102 @@ class TestAsyncResponsesAPIPromptManagement:
assert sent_input[0]["cache_control"] == {"type": "ephemeral"}
assert sent_input[1] == reasoning_item
assert sent_input[2]["id"] == "msg_1"
# ---------------------------------------------------------------------------
# Cross-provider model swap guard (prompt swaps model after credential resolution)
# ---------------------------------------------------------------------------
def test_resolve_prompt_swapped_provider_raises_cross_provider_with_credentials():
import litellm
from litellm.responses.main import _resolve_prompt_swapped_provider
with pytest.raises(litellm.BadRequestError, match="Refusing to send"):
_resolve_prompt_swapped_provider(
original_model="anthropic/claude-haiku-4-5",
swapped_model="gpt-4o-mini",
custom_llm_provider="anthropic",
kwargs={"api_key": "sk-ant-test"},
prompt_id="p1",
)
def test_resolve_prompt_swapped_provider_allows_swap_without_credentials():
from litellm.responses.main import _resolve_prompt_swapped_provider
assert (
_resolve_prompt_swapped_provider(
original_model="anthropic/claude-haiku-4-5",
swapped_model="gpt-4o-mini",
custom_llm_provider="anthropic",
kwargs={},
prompt_id="p1",
)
== "openai"
)
def test_resolve_prompt_swapped_provider_allows_same_provider_swap_with_credentials():
from litellm.responses.main import _resolve_prompt_swapped_provider
assert (
_resolve_prompt_swapped_provider(
original_model="openai/gpt-4o",
swapped_model="gpt-4o-mini",
custom_llm_provider="openai",
kwargs={"api_key": "sk-test", "api_base": "https://api.openai.com/v1"},
prompt_id="p1",
)
== "openai"
)
def test_sync_prompt_swap_resolves_credentials_for_swapped_provider(monkeypatch: pytest.MonkeyPatch):
import litellm
monkeypatch.setenv("XAI_API_KEY", "sk-xai-test")
logging_obj = _make_logging_obj("gpt-4o-mini", [{"role": "user", "content": "hi"}])
with patch( # test-quality-ok: handler boundary stub proves creds resolve for the swapped provider without network
"litellm.responses.main.base_llm_http_handler.response_api_handler", return_value=MagicMock()
) as mock_handler:
litellm.responses(input="hi", model="xai/grok-4", prompt_id="p1", litellm_logging_obj=logging_obj)
handler_kwargs = mock_handler.call_args.kwargs
assert handler_kwargs["model"] == "gpt-4o-mini"
assert handler_kwargs["custom_llm_provider"] == "openai"
assert handler_kwargs["litellm_params"].api_base is None
assert handler_kwargs["litellm_params"].api_key != "sk-xai-test"
def test_sync_prompt_swap_cross_provider_with_credentials_raises():
import litellm
from litellm.responses.main import _apply_prompt_management_to_responses_call
logging_obj = _make_logging_obj("gpt-4o-mini", [{"role": "user", "content": "hi"}])
with pytest.raises(litellm.BadRequestError, match="Refusing to send"):
_apply_prompt_management_to_responses_call(
input="hi",
model="anthropic/claude-haiku-4-5",
custom_llm_provider="anthropic",
litellm_logging_obj=logging_obj,
kwargs={"prompt_id": "p1", "api_key": "sk-ant-test"},
local_vars={},
use_chat_completions_api=False,
)
@pytest.mark.asyncio
async def test_aresponses_prompt_swap_cross_provider_with_credentials_raises():
import litellm
logging_obj = _make_logging_obj("gpt-4o-mini", [{"role": "user", "content": "hi"}])
logging_obj.async_failure_handler = AsyncMock()
with pytest.raises(litellm.BadRequestError, match="Refusing to send"):
await litellm.aresponses(
input="hi",
model="anthropic/claude-haiku-4-5",
litellm_logging_obj=logging_obj,
prompt_id="p1",
api_key="sk-ant-test",
)

View file

@ -724,3 +724,20 @@ class TestMergePromptManagementInputReshape:
)
assert result == merged
class TestResponsesInputToChatMessages:
def test_none_input_returns_empty_list(self):
assert ResponsesAPIRequestUtils.responses_input_to_chat_messages(None) == []
def test_str_input_becomes_user_message(self):
assert ResponsesAPIRequestUtils.responses_input_to_chat_messages("hi") == [
{"role": "user", "content": "hi"}
]
def test_list_input_keeps_only_role_items(self):
reasoning_item = {"type": "reasoning", "id": "rs_1", "summary": []}
user_message = {"role": "user", "content": "hi"}
assert ResponsesAPIRequestUtils.responses_input_to_chat_messages(
[reasoning_item, user_message, "stray"]
) == [user_message]

View file

@ -8765,96 +8765,6 @@ def test_get_router_model_info_keeps_explicit_pricing_overrides():
assert litellm.get_model_info(model="anthropic/claude-sonnet-4-5")["input_cost_per_token"] != 1e-08
class TestAutoRoutedRequestMarker:
"""The proxy exposes the routed model group in the response body only when an
auto-routing strategy actually picked it. The marker is what separates that from
ordinary model-group routing, so it must clear on any re-entry (fallbacks reuse the
same request_kwargs) that routes plainly."""
class _RewriteStrategy:
async def async_pre_routing_hook(
self, model, request_kwargs, messages=None, input=None, specific_deployment=False
):
from litellm.types.router import PreRoutingHookResponse
return PreRoutingHookResponse(model="gemini-flash", messages=messages)
class _AbstainStrategy:
async def async_pre_routing_hook(
self, model, request_kwargs, messages=None, input=None, specific_deployment=False
):
return None
@classmethod
def _router(cls, strategy) -> "litellm.Router":
from litellm.types.router import TaggedPreRoutingStrategy
router = litellm.Router(
model_list=[
{"model_name": "smart-route", "litellm_params": {"model": "openai/gpt-4o"}},
{"model_name": "gemini-flash", "litellm_params": {"model": "gemini/gemini-3.6-flash"}},
],
)
router.auto_routers = {"smart-route": [TaggedPreRoutingStrategy(tags=(), strategy=strategy)]}
return router
@pytest.mark.asyncio
async def test_marks_the_request_when_an_auto_routing_strategy_picked_the_group(self):
from litellm.constants import AUTO_ROUTED_REQUEST_METADATA_KEY
router = self._router(self._RewriteStrategy())
request_kwargs = {"metadata": {}}
await router.async_pre_routing_hook(model="smart-route", request_kwargs=request_kwargs)
assert request_kwargs["metadata"][AUTO_ROUTED_REQUEST_METADATA_KEY] is True
@pytest.mark.asyncio
async def test_marks_into_litellm_metadata_when_the_request_uses_that_bucket(self):
from litellm.constants import AUTO_ROUTED_REQUEST_METADATA_KEY
router = self._router(self._RewriteStrategy())
request_kwargs = {"litellm_metadata": {}}
await router.async_pre_routing_hook(model="smart-route", request_kwargs=request_kwargs)
assert request_kwargs["litellm_metadata"][AUTO_ROUTED_REQUEST_METADATA_KEY] is True
@pytest.mark.asyncio
async def test_no_marker_when_the_group_has_no_auto_routing_strategy(self):
from litellm.constants import AUTO_ROUTED_REQUEST_METADATA_KEY
router = self._router(self._RewriteStrategy())
request_kwargs = {"metadata": {}}
await router.async_pre_routing_hook(model="gemini-flash", request_kwargs=request_kwargs)
assert AUTO_ROUTED_REQUEST_METADATA_KEY not in request_kwargs["metadata"]
@pytest.mark.asyncio
async def test_no_marker_when_the_strategy_declined_to_route(self):
from litellm.constants import AUTO_ROUTED_REQUEST_METADATA_KEY
router = self._router(self._AbstainStrategy())
request_kwargs = {"metadata": {}}
await router.async_pre_routing_hook(model="smart-route", request_kwargs=request_kwargs)
assert AUTO_ROUTED_REQUEST_METADATA_KEY not in request_kwargs["metadata"]
@pytest.mark.asyncio
async def test_fallback_reentry_with_a_plain_group_clears_the_stale_marker(self):
from litellm.constants import AUTO_ROUTED_REQUEST_METADATA_KEY
router = self._router(self._RewriteStrategy())
request_kwargs = {"metadata": {}}
await router.async_pre_routing_hook(model="smart-route", request_kwargs=request_kwargs)
await router.async_pre_routing_hook(model="gemini-flash", request_kwargs=request_kwargs)
assert AUTO_ROUTED_REQUEST_METADATA_KEY not in request_kwargs["metadata"]
class TestModelGroupAliasReachesPreRoutingStrategies:
"""A `model_group_alias` whose target is a strategy router must dispatch exactly like the
router's own model_name. The four strategy registries are keyed by the marker deployment's
@ -8912,8 +8822,6 @@ class TestModelGroupAliasReachesPreRoutingStrategies:
@pytest.mark.parametrize("registry_name", REGISTRY_NAMES)
@pytest.mark.asyncio
async def test_alias_dispatches_to_the_strategy_registered_under_the_target(self, registry_name):
from litellm.constants import AUTO_ROUTED_REQUEST_METADATA_KEY
router = self._router(registry_name)
request_kwargs = {"metadata": {}}
@ -8923,7 +8831,6 @@ class TestModelGroupAliasReachesPreRoutingStrategies:
assert response is not None
assert response.model == "gemini-flash"
assert request_kwargs["metadata"][AUTO_ROUTED_REQUEST_METADATA_KEY] is True
@pytest.mark.asyncio
async def test_alias_call_still_forwards_the_marker_own_params_to_the_routed_tier(self):

8
uv.lock generated
View file

@ -10,7 +10,7 @@ resolution-markers = [
]
[options]
exclude-newer = "2026-08-23T02:27:57.028643Z"
exclude-newer = "2026-08-23T20:15:58.934396Z"
exclude-newer-span = "P3D"
[manifest]
@ -4315,6 +4315,9 @@ google = [
grpc = [
{ name = "grpcio" },
]
mcp = [
{ name = "mcp" },
]
mlflow = [
{ name = "mlflow" },
]
@ -4526,6 +4529,7 @@ requires-dist = [
{ name = "litellm-proxy-extras", marker = "extra == 'proxy'", editable = "litellm-proxy-extras" },
{ name = "llm-sandbox", marker = "extra == 'proxy-runtime'", specifier = ">=0.3.39,<1.0" },
{ name = "mangum", marker = "extra == 'proxy-runtime'", specifier = ">=0.17.0,<1.0" },
{ name = "mcp", marker = "extra == 'mcp'", specifier = ">=1.28.1,<2.0" },
{ name = "mcp", marker = "extra == 'proxy'", specifier = ">=1.28.1,<2.0" },
{ name = "mlflow", marker = "extra == 'mlflow'", specifier = ">=3.11.1,<4.0" },
{ name = "numpy", marker = "extra == 'stt-nvidia-riva'", specifier = ">=1.26.0" },
@ -4569,7 +4573,7 @@ requires-dist = [
{ name = "uvloop", marker = "sys_platform != 'win32' and extra == 'proxy'", specifier = ">=0.21.0,<1.0" },
{ name = "websockets", marker = "extra == 'proxy'", specifier = ">=15.0.1,<16.0" },
]
provides-extras = ["proxy", "cli", "extra-proxy", "utils", "caching", "saml", "semantic-router", "mlflow", "grpc", "stt-nvidia-riva", "google", "bedrock-realtime", "proxy-runtime"]
provides-extras = ["proxy", "cli", "extra-proxy", "utils", "caching", "mcp", "saml", "semantic-router", "mlflow", "grpc", "stt-nvidia-riva", "google", "bedrock-realtime", "proxy-runtime"]
[package.metadata.requires-dev]
ci = [