Merge branch 'litellm_internal_staging' into litellm_e2e_coverage_generated_denominator

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-07-17 18:32:53 +00:00
commit 2fe9d30f3b
38 changed files with 1900 additions and 152 deletions

View file

@ -315,6 +315,11 @@ disable_token_counter: bool = False
disable_add_transform_inline_image_block: bool = False
disable_add_user_agent_to_request_tags: bool = False
disable_anthropic_gemini_context_caching_transform: bool = False
enable_anthropic_prompt_caching: bool = os.getenv("LITELLM_ENABLE_ANTHROPIC_PROMPT_CACHING", "false").lower() == "true"
_anthropic_prompt_caching_ttl_env: Optional[str] = os.getenv("LITELLM_ANTHROPIC_PROMPT_CACHING_TTL")
anthropic_prompt_caching_ttl: Optional[Literal["5m", "1h"]] = (
"1h" if _anthropic_prompt_caching_ttl_env == "1h" else "5m" if _anthropic_prompt_caching_ttl_env == "5m" else None
)
disable_vertex_batch_output_transformation: bool = False
extra_spend_tag_headers: Optional[List[str]] = None
in_memory_llm_clients_cache: "LLMClientCache"

View file

@ -296,18 +296,148 @@ class AnthropicCacheControlHook(CustomPromptManagement):
return processed_messages, processed_system, remaining_points
@staticmethod
def _default_control() -> ChatCompletionCachedContent:
"""Build the cache_control block for auto-injected breakpoints.
Defaults to Anthropic's 5-minute ephemeral cache; honors the optional
``litellm.anthropic_prompt_caching_ttl`` override ("5m" or "1h").
"""
import litellm
ttl = litellm.anthropic_prompt_caching_ttl
if ttl == "5m" or ttl == "1h":
return ChatCompletionCachedContent(type="ephemeral", ttl=ttl)
return ChatCompletionCachedContent(type="ephemeral")
@staticmethod
def _request_has_cache_control(
messages: list[AllMessageValues],
system: str | list | None,
tools: list | None = None,
) -> bool:
"""Return True if the request already carries any client-supplied cache_control.
When the client (e.g. Claude Code) already marks its own breakpoints we
stand down entirely rather than add more, per the auto-caching contract.
Tools count: they are a breakpoint the client can mark, they count toward
the provider's four-block limit, and caching only the tool definitions is
a common pattern, so injecting alongside them can exceed the cap.
"""
if any(AnthropicCacheControlHook._count_cache_control_blocks(msg) for msg in messages):
return True
if isinstance(system, list):
if any(isinstance(block, dict) and block.get("cache_control") is not None for block in system):
return True
if tools is not None:
return any(isinstance(tool, dict) and tool.get("cache_control") is not None for tool in tools)
return False
@staticmethod
def get_default_injection_points(
messages: list[AllMessageValues],
system: str | list | None,
model: str,
custom_llm_provider: str | None,
tools: list | None = None,
) -> list[CacheControlInjectionPoint]:
"""Default breakpoints when ``litellm.enable_anthropic_prompt_caching`` is on.
Caches the system prompt and the trailing turn, so the stable prefix
(system + tools + history) is reused while the breakpoint advances with
the conversation. Returns [] (stand down) when the flag is off, the
provider does not consume cache_control breakpoints (only anthropic /
bedrock do), the model lacks prompt-caching support, or the request
already carries client-supplied cache_control.
"""
import litellm
if litellm.enable_anthropic_prompt_caching is not True:
return []
provider = custom_llm_provider
if provider is None:
from litellm.litellm_core_utils.get_llm_provider_logic import (
get_llm_provider,
)
try:
_, provider, _, _ = get_llm_provider(model=model)
except Exception: # noqa: BLE001 # unroutable model must never block the call, just skip auto-caching
return []
if provider not in ("anthropic", "bedrock"):
return []
from litellm.utils import supports_prompt_caching
if not supports_prompt_caching(model=model, custom_llm_provider=provider):
return []
if AnthropicCacheControlHook._request_has_cache_control(messages, system, tools):
return []
control = AnthropicCacheControlHook._default_control()
points: list[CacheControlInjectionPoint] = [
CacheControlMessageInjectionPoint(location="message", role="system", index=None, control=control),
CacheControlMessageInjectionPoint(location="message", role=None, index=-1, control=control),
]
return points
@staticmethod
def maybe_seed_default_injection_points(
non_default_params: dict[str, Any],
messages: list[AllMessageValues],
model: str,
custom_llm_provider: str | None,
tools: list | None = None,
) -> None:
"""For /chat/completions: add default injection points to the request params.
No-op when injection points are already configured (explicit config wins).
Seeding the param lets the existing prompt-management gate and the
AnthropicCacheControlHook run unchanged.
"""
if non_default_params.get("cache_control_injection_points"):
return
points = AnthropicCacheControlHook.get_default_injection_points(
messages=messages,
system=None,
model=model,
custom_llm_provider=custom_llm_provider,
tools=tools,
)
if points:
non_default_params["cache_control_injection_points"] = points
@staticmethod
def maybe_inject_cache_control(
messages: List[Dict],
system: str | list | None,
kwargs: Dict[str, Any],
model: str | None = None,
custom_llm_provider: str | None = None,
tools: list[dict] | None = None,
) -> Tuple[List[Dict], str | list | None]:
"""Extract cache_control_injection_points from kwargs and apply if present.
When none are configured but ``litellm.enable_anthropic_prompt_caching``
is on, synthesize default breakpoints for the native /v1/messages path.
Pops the key from kwargs; if remaining (non-message) points exist they
are written back so downstream transforms can handle them.
"""
injection_points = kwargs.pop("cache_control_injection_points", None)
configured = cast( # cast-ok: kwargs is untyped; this key only holds the documented injection-point list
list[CacheControlInjectionPoint] | None, kwargs.pop("cache_control_injection_points", None)
)
injection_points: list[CacheControlInjectionPoint] = configured or []
if not injection_points and model is not None:
injection_points = AnthropicCacheControlHook.get_default_injection_points(
messages=cast(list[AllMessageValues], messages), # cast-ok: Anthropic-shaped dicts from v1/messages
system=system,
tools=tools,
model=model,
custom_llm_provider=custom_llm_provider,
)
if not injection_points:
return messages, system

View file

@ -1453,6 +1453,9 @@ class Logging(LiteLLMLoggingBaseClass):
response_cost = litellm.response_cost_calculator(**response_cost_calculator_kwargs)
verbose_logger.debug(f"response_cost: {response_cost}")
additional_response_cost: object = self.model_call_details.get("additional_response_cost")
if isinstance(additional_response_cost, (int, float)) and additional_response_cost > 0:
return (response_cost or 0.0) + additional_response_cost
return response_cost
except Exception as e: # error calculating cost
debug_info = StandardLoggingModelCostFailureDebugInformation(

View file

@ -906,16 +906,17 @@ def strip_advisor_blocks_from_messages(messages: List[Any], replace_with_text: b
def is_anthropic_invalid_thinking_signature_error(error_text: str) -> bool:
"""
Detect Anthropic 400 when encrypted thinking signatures in history do not match
the current deployment (e.g. user rotated API key or switched model endpoint).
Detect Anthropic 400 errors caused by missing or invalid thinking signatures.
Example API message:
Known error formats:
{"message":"messages.2.content.0.thinking.signature.str: Input should be a valid string"}
messages.N.content.M.thinking.signature.str: Input should be a valid string
messages.N.content.M: Invalid `signature` in `thinking` block
"""
if not error_text:
return False
lower = error_text.lower()
return "invalid" in lower and "signature" in lower and "thinking" in lower and "block" in lower
return "thinking" in lower and "signature" in lower and ("invalid" in lower or "valid string" in lower)
def strip_thinking_blocks_from_anthropic_messages(messages: List[Any]) -> List[Any]:

View file

@ -237,7 +237,9 @@ async def anthropic_messages(
AnthropicCacheControlHook,
)
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(messages, system, kwargs)
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(
messages, system, kwargs, model=model, custom_llm_provider=custom_llm_provider, tools=tools
)
original_stream = stream or kwargs.get("_websearch_interception_converted_stream", False)
@ -426,7 +428,9 @@ def anthropic_messages_handler(
AnthropicCacheControlHook,
)
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(messages, system, kwargs)
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(
messages, system, kwargs, model=model, custom_llm_provider=custom_llm_provider, tools=tools
)
metadata = validate_anthropic_api_metadata(metadata)

View file

@ -75,10 +75,23 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]:
model_info = get_model_info(model=base_model, custom_llm_provider="fireworks_ai")
## CALCULATE INPUT COST
prompt_tokens_details = usage.prompt_tokens_details
cached_tokens: int = (
prompt_tokens_details.cached_tokens
if prompt_tokens_details is not None and prompt_tokens_details.cached_tokens is not None
else 0
)
input_cost_per_token: float = model_info["input_cost_per_token"] or 0.0
cache_read_input_token_cost = model_info.get("cache_read_input_token_cost")
cache_read_cost_per_token: float = (
cache_read_input_token_cost if cache_read_input_token_cost is not None else input_cost_per_token
)
non_cached_prompt_tokens: int = max(usage.prompt_tokens - cached_tokens, 0)
prompt_cost: float = usage["prompt_tokens"] * model_info["input_cost_per_token"]
prompt_cost: float = non_cached_prompt_tokens * input_cost_per_token + cached_tokens * cache_read_cost_per_token
## CALCULATE OUTPUT COST
completion_cost = usage["completion_tokens"] * model_info["output_cost_per_token"]
output_cost_per_token: float = model_info["output_cost_per_token"] or 0.0
completion_cost: float = usage.completion_tokens * output_cost_per_token
return prompt_cost, completion_cost

View file

@ -510,6 +510,20 @@ async def acompletion(
#########################################################
#########################################################
litellm_logging_obj = kwargs.get("litellm_logging_obj", None)
from litellm.integrations.anthropic_cache_control_hook import (
AnthropicCacheControlHook,
)
from litellm.types.llms.openai import AllMessageValues
AnthropicCacheControlHook.maybe_seed_default_injection_points(
non_default_params=kwargs,
messages=cast(list[AllMessageValues], messages), # cast-ok: acompletion types messages as a bare List
model=model,
custom_llm_provider=cast(Optional[str], custom_llm_provider), # cast-ok: read from untyped kwargs
tools=tools,
)
if isinstance(litellm_logging_obj, LiteLLMLoggingObj) and (
litellm_logging_obj.should_run_prompt_management_hooks(
prompt_id=kwargs.get("prompt_id", None),
@ -5055,6 +5069,19 @@ def completion( # type: ignore
litellm_params = {} # used to prevent unbound var errors
## PROMPT MANAGEMENT HOOKS ##
from litellm.integrations.anthropic_cache_control_hook import (
AnthropicCacheControlHook,
)
from litellm.types.llms.openai import AllMessageValues
AnthropicCacheControlHook.maybe_seed_default_injection_points(
non_default_params=non_default_params,
messages=cast(list[AllMessageValues], messages), # cast-ok: completion types messages as a bare List
model=model,
custom_llm_provider=cast(Optional[str], kwargs.get("custom_llm_provider")), # cast-ok: untyped kwargs
tools=tools,
)
if isinstance(litellm_logging_obj, LiteLLMLoggingObj) and (
litellm_logging_obj.should_run_prompt_management_hooks(
prompt_id=prompt_id, non_default_params=non_default_params

View file

@ -3451,7 +3451,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2.2e-05,
"output_cost_per_token": 2.64e-06,
"supports_audio_input": true,
@ -3470,7 +3470,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 0.00022,
"output_cost_per_token": 2.2e-05,
"supports_audio_input": true,
@ -3489,7 +3489,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2.2e-05,
"supported_modalities": [
@ -4687,7 +4687,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supports_audio_input": true,
@ -4707,7 +4707,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -4739,7 +4739,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -4771,7 +4771,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supported_endpoints": [
@ -4832,7 +4832,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 0.0002,
"output_cost_per_token": 2e-05,
"supports_audio_input": true,
@ -4850,7 +4850,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2e-05,
"supported_modalities": [
@ -7922,7 +7922,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2.2e-05,
"output_cost_per_token": 2.64e-06,
"supports_audio_input": true,
@ -7941,7 +7941,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 0.00022,
"output_cost_per_token": 2.2e-05,
"supports_audio_input": true,
@ -7960,7 +7960,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2.2e-05,
"supported_modalities": [
@ -22094,7 +22094,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supports_audio_input": true,
@ -22113,7 +22113,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supports_audio_input": true,
@ -22207,7 +22207,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2e-05,
"supports_audio_input": true,
@ -22225,7 +22225,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2e-05,
"supports_audio_input": true,
@ -22243,7 +22243,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2e-05,
"supports_audio_input": true,
@ -24438,7 +24438,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -24470,7 +24470,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -24502,7 +24502,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -24535,7 +24535,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 32000,
"max_tokens": 32000,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 2.4e-05,
"regional_processing_uplift_multiplier_eu": 1.1,
@ -24570,7 +24570,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"regional_processing_uplift_multiplier_eu": 1.1,
@ -24603,7 +24603,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supported_endpoints": [
@ -24635,7 +24635,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -43573,7 +43573,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supported_endpoints": [
@ -43606,7 +43606,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supported_endpoints": [

View file

@ -3708,22 +3708,22 @@ def _attach_redis_usage_cache(redis_cache: RedisCache, enable_redis_auth_cache:
litellm_config_cache.redis_cache = redis_cache
def resolve_complexity_router_plugins(
model_name: str,
complexity_router_config: dict,
def resolve_routing_plugins(
plugin_paths: list,
config_file_path: str | None,
) -> None:
source_label: str,
) -> list:
"""
Resolves `complexity_router_config["plugins"]` dotted-path strings to live
instances via `get_instance_fn` (the same convention `litellm_settings.callbacks`
uses), in place. Raises at config-load time if a path resolves to something that
doesn't implement `RoutingPlugin`, rather than deferring to a confusing
`AttributeError` on the first request that reaches the plugin pipeline.
Resolves a list of routing-plugin entries to live `RoutingPlugin` instances.
Each string entry is resolved through `get_instance_fn` (the same dotted-path
convention `litellm_settings.callbacks` uses, which resolves both local module
files next to the config and modules installed as Python packages); non-string
entries are assumed to already be instances and passed through. Raises at
config-load time if any entry resolves to something that doesn't implement
`RoutingPlugin`, rather than deferring to a confusing `AttributeError` on the
first request that reaches the plugin pipeline. `source_label` names the config
key being resolved so the error points the operator at the right place.
"""
plugin_paths = complexity_router_config.get("plugins")
if not isinstance(plugin_paths, list):
return
resolved_plugins = [
get_instance_fn(value=plugin_path, config_file_path=config_file_path)
if isinstance(plugin_path, str)
@ -3739,12 +3739,31 @@ def resolve_complexity_router_plugins(
getattr(resolved_plugin, "run", None)
):
raise ValueError(
f"complexity_router_config.plugins entry {plugin_path!r} on model {model_name!r} "
f"resolved to {resolved_plugin!r}, which does not implement the RoutingPlugin "
"interface (an async `run(context)` method). Fix the referenced module before "
"starting the proxy."
f"{source_label} entry {plugin_path!r} resolved to {resolved_plugin!r}, which does "
"not implement the RoutingPlugin interface (an async `run(context)` method). Fix the "
"referenced module before starting the proxy."
)
complexity_router_config["plugins"] = resolved_plugins
return resolved_plugins
def resolve_complexity_router_plugins(
model_name: str,
complexity_router_config: dict,
config_file_path: str | None,
) -> None:
"""
Resolves `complexity_router_config["plugins"]` dotted-path strings to live
instances in place, via `resolve_routing_plugins`.
"""
plugin_paths = complexity_router_config.get("plugins")
if not isinstance(plugin_paths, list):
return
complexity_router_config["plugins"] = resolve_routing_plugins(
plugin_paths=plugin_paths,
config_file_path=config_file_path,
source_label=f"complexity_router_config.plugins on model {model_name!r}",
)
class ProxyConfig:
@ -4874,6 +4893,12 @@ class ProxyConfig:
for k, v in router_settings.items():
if k in available_args:
if k == "plugins" and isinstance(v, list):
v = resolve_routing_plugins(
plugin_paths=v,
config_file_path=config_file_path,
source_label="router_settings.plugins",
)
router_params[k] = v
elif k in {"health_check_interval", "health_check_concurrency"}:
raise ValueError(

View file

@ -11,14 +11,16 @@ from typing import Any, Dict, Optional, Tuple
import orjson
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
from fastapi.responses import ORJSONResponse
from fastapi.responses import ORJSONResponse, StreamingResponse
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.proxy._types import *
from litellm.proxy.auth.auth_utils import is_request_body_safe
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
@ -604,6 +606,7 @@ async def rag_query(
general_settings,
llm_router,
proxy_config,
select_data_generator,
version,
)
@ -673,6 +676,31 @@ async def rag_query(
**request_data,
)
hidden_params = getattr(response, "_hidden_params", {}) or {}
custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=user_api_key_dict,
call_id=hidden_params.get("litellm_call_id", None) or "",
model_id=hidden_params.get("model_id", None) or "",
cache_key=hidden_params.get("cache_key", None) or "",
api_base=hidden_params.get("api_base", None) or "",
version=version,
response_cost=hidden_params.get("response_cost", None),
request_data=request_data,
)
if isinstance(response, CustomStreamWrapper):
return StreamingResponse(
select_data_generator(
response=response,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
request=request,
),
media_type="text/event-stream",
headers=custom_headers,
)
fastapi_response.headers.update(custom_headers)
return response
except HTTPException:

View file

@ -34,16 +34,12 @@ def get_instance_fn(value: str, config_file_path: Optional[str] = None) -> Any:
module_name = ".".join(parts[:-1])
instance_name = parts[-1]
# If config_file_path is provided, use it to determine the module spec and load the module
module_file_path = None
if config_file_path is not None:
directory = os.path.dirname(config_file_path)
module_file_path = os.path.join(directory, *module_name.split("."))
module_file_path += ".py"
# Check if the file exists before trying to load it
if not os.path.exists(module_file_path):
raise ImportError(f"Could not find module file {module_file_path}")
module_file_path = os.path.join(directory, *module_name.split(".")) + ".py"
if module_file_path is not None and os.path.exists(module_file_path):
spec = importlib.util.spec_from_file_location(module_name, module_file_path) # type: ignore
if spec is None:
raise ImportError(f"Could not find a module specification for {module_file_path}")
@ -52,7 +48,6 @@ def get_instance_fn(value: str, config_file_path: Optional[str] = None) -> Any:
raise ImportError(f"Could not find a module loader for {module_file_path}")
spec.loader.exec_module(module) # type: ignore
else:
# Dynamically import the module
module = importlib.import_module(module_name)
# Get the instance from the module

View file

@ -11,12 +11,14 @@ __all__ = ["ingest", "aingest", "query", "aquery"]
import asyncio
import contextvars
from contextlib import contextmanager
from functools import partial
from typing import (
TYPE_CHECKING,
Any,
Coroutine,
Dict,
Iterator,
List,
Optional,
Tuple,
@ -27,6 +29,9 @@ from typing import (
import httpx
import litellm
from litellm._internal_context import is_internal_call
from litellm.cost_calculator import vector_store_search_cost
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion
from litellm.rag.ingestion.bedrock_ingestion import BedrockRAGIngestion
from litellm.rag.ingestion.gemini_ingestion import GeminiRAGIngestion
@ -188,6 +193,25 @@ async def aingest(
)
@contextmanager
def _suppressed_sub_call_billing() -> Iterator[None]:
"""
Suppress a sub-call's own billing event so the parent aquery event bills it.
Every suppressed sub-call's cost must be folded into the parent event:
into the response's hidden response_cost on the non-streaming path, or via
the logging object's additional_response_cost on the streaming path (the
streamed cost is computed from assembled chunks after this pipeline
returns, so there is no response object to fold into here).
"""
previous = is_internal_call.get()
is_internal_call.set(True)
try:
yield
finally:
is_internal_call.set(previous)
async def _execute_query_pipeline(
model: str,
messages: List[Any],
@ -209,27 +233,46 @@ async def _execute_query_pipeline(
raise ValueError("No query found in messages for RAG query")
# 2. Search vector store
search_response = await litellm.vector_stores.asearch(
vector_store_id=retrieval_config["vector_store_id"],
query=query_text,
max_num_results=retrieval_config.get("top_k", 10),
custom_llm_provider=retrieval_config.get("custom_llm_provider", "openai"),
**kwargs,
)
with _suppressed_sub_call_billing():
search_response = await litellm.vector_stores.asearch(
vector_store_id=retrieval_config["vector_store_id"],
query=query_text,
max_num_results=retrieval_config.get("top_k", 10),
custom_llm_provider=retrieval_config.get("custom_llm_provider", "openai"),
**kwargs,
)
search_provider = retrieval_config.get("custom_llm_provider", "openai")
try:
search_cost = sum(
vector_store_search_cost(
model=search_provider if "/" in search_provider else None,
custom_llm_provider=search_provider,
response=search_response,
)
)
except Exception: # noqa: BLE001 - cost accounting must never break the query path
search_cost = 0.0
rerank_response = None
rerank_cost = 0.0
context_chunks = search_response.get("data", [])
# 3. Optional rerank
if rerank and rerank.get("enabled"):
documents = RAGQuery.extract_documents_from_search(search_response)
if documents:
rerank_response = await litellm.arerank(
model=rerank["model"],
query=query_text,
documents=documents,
top_n=rerank.get("top_n", 5),
)
with _suppressed_sub_call_billing():
rerank_response = await litellm.arerank(
model=rerank["model"],
query=query_text,
documents=documents,
top_n=rerank.get("top_n", 5),
)
rerank_hidden_params = getattr(rerank_response, "_hidden_params", None)
if isinstance(rerank_hidden_params, dict):
rerank_response_cost: float | None = rerank_hidden_params.get("response_cost")
rerank_cost = rerank_response_cost or 0.0
context_chunks = RAGQuery.get_top_chunks_from_rerank(search_response, rerank_response)
# 4. Build context message and call completion
@ -237,28 +280,40 @@ async def _execute_query_pipeline(
modified_messages = messages[:-1] + [context_message] + [messages[-1]]
# Use router if available to properly resolve virtual model names
if router is not None:
response = await router.acompletion(
model=model,
messages=modified_messages,
stream=stream,
**kwargs,
)
else:
response = await litellm.acompletion(
model=model,
messages=modified_messages,
stream=stream,
**kwargs,
)
with _suppressed_sub_call_billing():
if router is not None:
response = await router.acompletion(
model=model,
messages=modified_messages,
stream=stream,
**kwargs,
)
else:
response = await litellm.acompletion(
model=model,
messages=modified_messages,
stream=stream,
**kwargs,
)
# 5. Attach search results to response
sub_call_cost = search_cost + rerank_cost
if not stream and isinstance(response, ModelResponse):
response = RAGQuery.add_search_results_to_response(
response=response,
search_results=search_response,
rerank_results=rerank_response,
)
if sub_call_cost > 0:
hidden_params = getattr(response, "_hidden_params", None)
if isinstance(hidden_params, dict):
completion_response_cost: float | None = hidden_params.get("response_cost")
if completion_response_cost is not None:
hidden_params["response_cost"] = completion_response_cost + sub_call_cost
elif sub_call_cost > 0:
logging_obj: object = kwargs.get("litellm_logging_obj")
if isinstance(logging_obj, LiteLLMLoggingObj):
logging_obj.model_call_details["additional_response_cost"] = sub_call_cost
return response # type: ignore[return-value]

View file

@ -28,6 +28,7 @@ from litellm.types.utils import ModelResponse
from .config import (
DEFAULT_CODE_KEYWORDS,
DEFAULT_ESCALATION_KEYWORDS,
DEFAULT_REASONING_KEYWORDS,
DEFAULT_SIMPLE_KEYWORDS,
DEFAULT_TECHNICAL_KEYWORDS,
@ -173,6 +174,11 @@ class ComplexityRouter(CustomLogger):
self.config.custom_technical_keywords,
)
self.simple_keywords = self.config.simple_keywords or DEFAULT_SIMPLE_KEYWORDS
self.escalation_keywords = (
self.config.escalation_keywords
if self.config.escalation_keywords is not None
else DEFAULT_ESCALATION_KEYWORDS
)
# Lazily built on first semantic request and cached for reuse (route
# embeddings are static, only the prompt is embedded per request). The lock
@ -668,6 +674,53 @@ class ComplexityRouter(CustomLogger):
}
return best_model
def _escalation_triggered(self, user_message: str) -> bool:
"""Whether the prompt asks to escalate to a stronger model.
Matching is a case-sensitive substring test so the default "LITELLM ESCALATE"
only fires on the deliberate, shouted form and not on incidental lowercase
mentions of the word (e.g. "how do I escalate this ticket").
"""
if not self.escalation_keywords:
return False
return any(keyword in user_message for keyword in self.escalation_keywords)
def _tier_for_model(self, model: str) -> ComplexityTier | None:
"""Return the most-severe configured tier whose pool contains this model."""
pools = self._tier_pools()
matched = tuple(ComplexityTier(tier_name) for tier_name, models in pools.items() if model in models)
if not matched:
return None
return max(matched, key=TIER_SEVERITY_ORDER.index)
def _escalate_tier(self, tier: ComplexityTier) -> ComplexityTier:
"""Bump a tier one step up to the next-higher configured tier.
Returns the input tier unchanged when it is already the highest configured
tier, so escalation can never route below the model the user would otherwise
have received.
"""
configured = frozenset(self.config.tiers)
current_index = TIER_SEVERITY_ORDER.index(tier)
higher_tiers = tuple(
candidate for candidate in TIER_SEVERITY_ORDER[current_index + 1 :] if candidate.value in configured
)
return higher_tiers[0] if higher_tiers else tier
def _escalated_pin(self, pinned_model: str) -> str | None:
"""Bump a session's pinned model to the next-higher configured tier.
Returns None when the pin no longer maps to any configured tier, signalling
a full reclassification instead.
"""
pinned_tier = self._tier_for_model(pinned_model)
if pinned_tier is None:
return None
escalated_tier = self._escalate_tier(pinned_tier)
if escalated_tier == pinned_tier:
return pinned_model
return self.get_model_for_tier(escalated_tier)
def _lexical_tier_override(self, user_message: str) -> ComplexityTier | None:
"""When keyword_tier_rules match literally, the most-severe matched tier wins.
@ -910,29 +963,41 @@ class ComplexityRouter(CustomLogger):
if cache_key is not None:
pinned_model = await self.litellm_router_instance.cache.async_get_cache(key=cache_key)
if isinstance(pinned_model, str):
# Refresh the TTL on every hit so an active session doesn't lose its
# pin mid-conversation just because it outlives the original write.
await self.litellm_router_instance.cache.async_set_cache(
key=cache_key,
value=pinned_model,
ttl=self.config.session_affinity_ttl_seconds,
)
if self.config.adaptive:
from litellm.router_strategy.adaptive_router.config import (
ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY,
routed_model: str | None = pinned_model
if self.escalation_keywords:
resolved_messages = self._resolve_messages(messages, request_kwargs)
user_message = (
self._extract_user_message_and_system_prompt(resolved_messages)[0]
if resolved_messages
else None
)
if user_message is not None and self._escalation_triggered(user_message):
routed_model = self._escalated_pin(pinned_model)
if routed_model is not None:
# Refresh the TTL on every hit so an active session doesn't lose its
# pin mid-conversation just because it outlives the original write.
await self.litellm_router_instance.cache.async_set_cache(
key=cache_key,
value=routed_model,
ttl=self.config.session_affinity_ttl_seconds,
)
if self.config.adaptive:
from litellm.router_strategy.adaptive_router.config import (
ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY,
)
kwargs_metadata = request_kwargs.setdefault("metadata", {})
if isinstance(kwargs_metadata, dict):
kwargs_metadata[ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY] = pinned_model
verbose_router_logger.info(
f"ComplexityRouter: routing decision cause=session_affinity_pin, routed_model={pinned_model}"
)
has_original_messages = messages is not None and len(messages) > 0
return PreRoutingHookResponse(
model=pinned_model,
messages=messages if has_original_messages else None,
)
kwargs_metadata = request_kwargs.setdefault("metadata", {})
if isinstance(kwargs_metadata, dict):
kwargs_metadata[ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY] = routed_model
cause = "session_affinity_escalation" if routed_model != pinned_model else "session_affinity_pin"
verbose_router_logger.info(
f"ComplexityRouter: routing decision cause={cause}, routed_model={routed_model}"
)
has_original_messages = messages is not None and len(messages) > 0
return PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
)
response = await self._classify_and_route(
model=model,
@ -1004,13 +1069,17 @@ class ComplexityRouter(CustomLogger):
messages=messages if has_original_messages else None,
)
escalate = self._escalation_triggered(user_message)
override_tier = await self._resolve_keyword_tier_override(user_message, request_kwargs)
if override_tier is not None:
routed_model = await self._pick_model_for_tier(override_tier, messages, resolved_messages, request_kwargs)
cause = "semantic_keyword_match" if self.config.semantic_keyword_matching else "literal_keyword_match"
routed_tier = self._escalate_tier(override_tier) if escalate else override_tier
routed_model = await self._pick_model_for_tier(routed_tier, messages, resolved_messages, request_kwargs)
base_cause = "semantic_keyword_match" if self.config.semantic_keyword_matching else "literal_keyword_match"
cause = f"{base_cause}+escalation" if escalate else base_cause
verbose_router_logger.info(
f"ComplexityRouter: routing decision cause={cause}, "
f"tier={override_tier.value}, routed_model={routed_model}"
f"tier={routed_tier.value}, routed_model={routed_model}"
)
return PreRoutingHookResponse(
model=routed_model,
@ -1018,6 +1087,9 @@ class ComplexityRouter(CustomLogger):
)
tier, score, signals = await self.aclassify(user_message, system_prompt, request_kwargs)
if escalate:
tier = self._escalate_tier(tier)
signals = [*signals, "escalation"]
if self.config.adaptive:
routed_model = self._soft_floor_pick(tier, user_message, request_kwargs)
adaptive = self._ensure_adaptive_router()

View file

@ -162,6 +162,9 @@ DEFAULT_TECHNICAL_KEYWORDS: list[str] = [
# Note: "async", "kubernetes", "docker" are in DEFAULT_CODE_KEYWORDS
]
DEFAULT_ESCALATION_KEYWORDS: list[str] = ["LITELLM ESCALATE"]
DEFAULT_SIMPLE_KEYWORDS: list[str] = [
"what is",
"what's",
@ -339,6 +342,16 @@ class ComplexityRouterConfig(BaseModel):
),
)
escalation_keywords: list[str] | None = Field(
default=None,
description=(
"Case-sensitive phrases a user can include to force a bump to the next-higher "
"complexity tier when they aren't satisfied with results (they can force a stronger "
"model, but not choose which one). Defaults to ['LITELLM ESCALATE'] when unset; "
"set to an empty list to disable."
),
)
# Deterministic keyword -> tier overrides, evaluated before weighted scoring
keyword_tier_rules: list[KeywordTierRule] | None = Field(
default=None,
@ -400,6 +413,13 @@ class ComplexityRouterConfig(BaseModel):
coerced[key] = item
return coerced
@field_validator("escalation_keywords")
@classmethod
def _normalize_escalation_keywords(cls, value: list[str] | None) -> list[str] | None:
if value is None:
return None
return [stripped for keyword in value if (stripped := keyword.strip())]
@model_validator(mode="after")
def _validate_llm_classifier_config(self) -> "ComplexityRouterConfig":
if self.classifier_type == "llm" and self.classifier_llm_config is None:

View file

@ -529,6 +529,7 @@ class ChatCompletionDeltaToolCallChunk(TypedDict, total=False):
class ChatCompletionCachedContent(TypedDict):
type: Literal["ephemeral"]
ttl: NotRequired[Literal["5m", "1h"]]
class ChatCompletionThinkingBlock(TypedDict, total=False):

View file

@ -266,6 +266,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
"audio_transcription",
"responses",
"ocr",
"realtime",
]
]
tpm: Optional[int]
@ -402,6 +403,11 @@ class CallTypes(str, Enum):
vector_store_search = "vector_store_search"
avector_store_search = "avector_store_search"
ingest = "ingest"
aingest = "aingest"
query = "query"
aquery = "aquery"
#########################################################
# Container Call Types
#########################################################

View file

@ -3451,7 +3451,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2.2e-05,
"output_cost_per_token": 2.64e-06,
"supports_audio_input": true,
@ -3470,7 +3470,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 0.00022,
"output_cost_per_token": 2.2e-05,
"supports_audio_input": true,
@ -3489,7 +3489,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2.2e-05,
"supported_modalities": [
@ -4687,7 +4687,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supports_audio_input": true,
@ -4707,7 +4707,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -4739,7 +4739,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -4771,7 +4771,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supported_endpoints": [
@ -4832,7 +4832,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 0.0002,
"output_cost_per_token": 2e-05,
"supports_audio_input": true,
@ -4850,7 +4850,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2e-05,
"supported_modalities": [
@ -7922,7 +7922,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2.2e-05,
"output_cost_per_token": 2.64e-06,
"supports_audio_input": true,
@ -7941,7 +7941,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 0.00022,
"output_cost_per_token": 2.2e-05,
"supports_audio_input": true,
@ -7960,7 +7960,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2.2e-05,
"supported_modalities": [
@ -22169,7 +22169,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supports_audio_input": true,
@ -22188,7 +22188,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supports_audio_input": true,
@ -22282,7 +22282,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2e-05,
"supports_audio_input": true,
@ -22300,7 +22300,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2e-05,
"supports_audio_input": true,
@ -22318,7 +22318,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 8e-05,
"output_cost_per_token": 2e-05,
"supports_audio_input": true,
@ -24513,7 +24513,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -24545,7 +24545,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -24577,7 +24577,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -24610,7 +24610,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 32000,
"max_tokens": 32000,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 2.4e-05,
"regional_processing_uplift_multiplier_eu": 1.1,
@ -24645,7 +24645,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"regional_processing_uplift_multiplier_eu": 1.1,
@ -24678,7 +24678,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supported_endpoints": [
@ -24710,7 +24710,7 @@
"max_input_tokens": 32000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 1.6e-05,
"supported_endpoints": [
@ -43694,7 +43694,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supported_endpoints": [
@ -43727,7 +43727,7 @@
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"mode": "realtime",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"supported_endpoints": [

View file

@ -131,7 +131,7 @@ Quota Management - behavior features (entity- or config-driven caps and their ac
quota_management.<behavior>.<variant>.<assertion>
behavior : ratelimit | budget | spend_tracking
variant : <ratelimit> rpm | tpm | priority_generous | priority_strict
<budget> key | internal_user | end_user | organization | team_member | tag
<budget> key | internal_user | end_user | organization | team | team_member | tag
| model_max | soft | key_multi_window | team_multi_window
| fallback | spend_counter
<spend_tracking> chat_completions | stream | embeddings | cache_hit | key_rollup

View file

@ -352,6 +352,7 @@ quota_management.ratelimit.rpm.headers_report_remaining: {tier: P1, source: para
quota_management.ratelimit.priority_generous.picks_under_tpm: {tier: P1, source: 'dynamic_rate_limiter_v3.py:36-52', rationale: Generous mode (<80% sat) allows priority borrowing}
quota_management.ratelimit.priority_strict.picks_under_tpm: {tier: P1, source: 'dynamic_rate_limiter_v3.py:53-71', rationale: Strict mode (>=80% sat) enforces priority fairness}
quota_management.budget.key.blocks_over_limit: {tier: P0, source: proxy/auth/auth_checks.py, rationale: A key's max_budget blocks further paid calls once spend crosses it}
quota_management.budget.team.blocks_over_limit: {tier: P0, source: proxy/auth/auth_checks.py, rationale: "A team's max_budget blocks every key on the team once combined spend crosses it, including keys that spent nothing themselves"}
quota_management.budget.internal_user.blocks_over_limit: {tier: P1, source: proxy/auth/auth_checks.py, rationale: An internal user's max_budget governs personal keys}
quota_management.budget.end_user.blocks_over_limit: {tier: P1, source: proxy/auth/auth_checks.py, rationale: A customer (end-user) max_budget blocks calls attributed via user=}
quota_management.budget.organization.blocks_over_limit: {tier: P1, source: proxy/auth/auth_checks.py, rationale: An organization's max_budget blocks keys under its teams}

View file

@ -4,7 +4,7 @@ Each entity is an E2ECase (lifecycle.E2ECase) driven by run_case: init() creates
the budgeted entity + a key, run() drives spend until a `budget_exceeded` block,
teardown() deletes everything init() created (always runs, even on failure/skip).
Covers the entities with no prior live coverage - internal user, end-user,
organization, team member. See BUDGET_TEST_COVERAGE_MATRIX.md.
organization, team member - plus key and team. See BUDGET_TEST_COVERAGE_MATRIX.md.
A non-budget error fails hard (never a skip); if calls never get blocked, budget
enforcement is broken -> fail.
@ -18,16 +18,17 @@ import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import require_successful_call
from e2e_http import StreamingResponse, require_successful_call
from lifecycle import run_case
pytestmark = pytest.mark.e2e
def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") -> None:
"""Send paid calls until the entity's budget blocks one. Key/user/org/member
block within a couple calls off real-time reservation counters; the end-user
budget enforces off table spend that lands on the batch write, so it takes a
few more. A non-budget error fails hard (never a skip)."""
def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") -> StreamingResponse:
"""Send paid calls until the entity's budget blocks one; return the blocked
response so callers can assert on its shape. Key/user/org/member block within
a couple calls off real-time reservation counters; the end-user budget
enforces off table spend that lands on the batch write, so it takes a few
more. A non-budget error fails hard (never a skip)."""
for _ in range(40):
result = client.chat(
key,
@ -37,7 +38,7 @@ def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") ->
user=user or None,
)
if is_budget_block(result):
return
return result
require_successful_call(result)
time.sleep(2)
pytest.fail("budget never enforced within the call budget")
@ -69,10 +70,53 @@ class _BudgetCase:
class KeyBudgetCase(_BudgetCase):
"""A bare key (no team_id / user_id) carrying its own max_budget, so only the
key-level budget can be the thing that blocks. The refusal must be a 429
budget_exceeded; any other error already fails via _assert_budget_blocks."""
def init(self) -> None:
self.key = self.client.generate_key(max_budget=3e-6)
self._undo.append(lambda: self.client.delete_key(self.key))
def run(self) -> None:
blocked = _assert_budget_blocks(self.client, self.key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
class TeamBudgetCase(_BudgetCase):
"""An admin caps a whole team: two keys under a tiny-budget team, neither with
a key-level budget. Key A is driven until the team cap blocks it; key B's very
first call must then be refused too, proving the cap sits on the team, not the
key that spent. Both refusals must be 429 budget_exceeded."""
def init(self) -> None:
team_id = self.client.create_team(
alias=f"e2e-budget-team-{unique_marker()}", max_budget=3e-6
)
self._undo.append(lambda: self.client.delete_team(team_id))
self.key = self.client.generate_key(team_id=team_id)
self._undo.append(lambda: self.client.delete_key(self.key))
self._sibling_key = self.client.generate_key(team_id=team_id)
self._undo.append(lambda: self.client.delete_key(self._sibling_key))
def run(self) -> None:
blocked = _assert_budget_blocks(self.client, self.key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
sibling = self.client.chat(
self._sibling_key,
"claude-haiku-4-5",
f"spend {unique_marker()}",
max_tokens=16,
)
assert is_budget_block(sibling) and sibling.status_code == 429, (
f"a sibling key on the capped team must get the same 429 budget_exceeded, "
f"got {sibling.status_code}: {sibling.body[:200]}"
)
class InternalUserBudgetCase(_BudgetCase):
def init(self) -> None:
@ -138,6 +182,10 @@ def _case_id(case_cls: Type[_BudgetCase]) -> str:
KeyBudgetCase,
marks=pytest.mark.covers("quota_management.budget.key.blocks_over_limit"),
),
pytest.param(
TeamBudgetCase,
marks=pytest.mark.covers("quota_management.budget.team.blocks_over_limit"),
),
pytest.param(
InternalUserBudgetCase,
marks=pytest.mark.covers("quota_management.budget.internal_user.blocks_over_limit"),

View file

@ -172,7 +172,7 @@ def anthropic_messages():
"content": [
{
"type": "text",
"text": "Here is the full text of a complex legal agreement" * 400,
"text": "Here is the full text of a complex legal agreement" * 500,
"cache_control": {"type": "ephemeral"},
}
],

View file

@ -2,7 +2,9 @@ import copy
import datetime
import json
import os
import subprocess
import sys
import textwrap
import unittest
from typing import List, Optional, Tuple
from unittest.mock import ANY, MagicMock, Mock, patch
@ -1533,3 +1535,242 @@ class TestApplyToAnthropicMessagesRequest:
sys_blocks = sum(1 for b in (result_sys or []) if isinstance(b, dict) and b.get("cache_control") is not None)
total_blocks = sys_blocks + sum(AnthropicCacheControlHook._count_cache_control_blocks(m) for m in result_msgs)
assert total_blocks <= 4
class TestEnableAnthropicPromptCaching:
"""Auto-injected default breakpoints via litellm.enable_anthropic_prompt_caching."""
MESSAGES: List[AllMessageValues] = [
{"role": "system", "content": "a long system prompt"},
{"role": "user", "content": "first turn"},
{"role": "assistant", "content": "a reply"},
{"role": "user", "content": "latest turn"},
]
def _points(self, model="claude-sonnet-4-5", provider="anthropic", messages=None, system=None, tools=None):
return AnthropicCacheControlHook.get_default_injection_points(
messages=copy.deepcopy(self.MESSAGES) if messages is None else messages,
system=system,
model=model,
custom_llm_provider=provider,
tools=tools,
)
def test_disabled_by_default(self):
assert litellm.enable_anthropic_prompt_caching is False
assert self._points() == []
def test_injects_system_and_trailing_turn(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
assert self._points() == [
{"location": "message", "role": "system", "index": None, "control": {"type": "ephemeral"}},
{"location": "message", "role": None, "index": -1, "control": {"type": "ephemeral"}},
]
def test_bedrock_claude_is_injected(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
points = self._points(model="us.anthropic.claude-sonnet-4-5-20250929-v1:0", provider="bedrock")
assert [p["index"] for p in points] == [None, -1]
@pytest.mark.parametrize("model, provider", [("gpt-4o", "openai"), ("gemini-2.0-flash", "gemini")])
def test_non_anthropic_providers_never_injected(self, monkeypatch, model, provider):
"""These report supports_prompt_caching=True but never consume cache_control markers."""
from litellm.utils import supports_prompt_caching
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
assert supports_prompt_caching(model=model, custom_llm_provider=provider) is True
assert self._points(model=model, provider=provider) == []
def test_model_without_caching_support_not_injected(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
assert self._points(model="anthropic.claude-3-5-sonnet-20240620-v1:0", provider="bedrock") == []
def test_stands_down_when_client_sent_cache_control(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
messages = [
{"role": "system", "content": [{"type": "text", "text": "s", "cache_control": {"type": "ephemeral"}}]},
{"role": "user", "content": "latest turn"},
]
assert self._points(messages=messages) == []
def test_stands_down_when_system_block_has_cache_control(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
system = [{"type": "text", "text": "s", "cache_control": {"type": "ephemeral"}}]
assert self._points(messages=[{"role": "user", "content": "hi"}], system=system) == []
@staticmethod
def _tools(count: int, cached: bool) -> List[dict]:
tool: dict = {"type": "function", "function": {"name": "t", "description": "d", "parameters": {}}}
if cached:
tool["cache_control"] = {"type": "ephemeral"}
return [{**tool, "function": {**tool["function"], "name": f"t{i}"}} for i in range(count)]
def test_stands_down_when_only_tools_carry_cache_control(self, monkeypatch):
"""Caching just the tool definitions is a normal client pattern, and those
breakpoints count toward the provider's four-block limit. Three of them plus
our two would be five, which Anthropic rejects outright."""
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
assert self._points(tools=self._tools(3, cached=True)) == []
def test_injects_when_tools_carry_no_cache_control(self, monkeypatch):
"""Tools alone must not suppress injection; only client-marked ones do."""
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
assert [p["index"] for p in self._points(tools=self._tools(3, cached=False))] == [None, -1]
@pytest.mark.parametrize("tools", [None, []])
def test_absent_tools_do_not_suppress_injection(self, monkeypatch, tools):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
assert [p["index"] for p in self._points(tools=tools)] == [None, -1]
def test_seed_stands_down_when_only_tools_carry_cache_control(self, monkeypatch):
"""Same guard on the /chat/completions seeding path."""
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
params: dict = {}
AnthropicCacheControlHook.maybe_seed_default_injection_points(
non_default_params=params,
messages=copy.deepcopy(self.MESSAGES),
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
tools=self._tools(3, cached=True),
)
assert "cache_control_injection_points" not in params
def test_v1_messages_stands_down_when_only_tools_carry_cache_control(self, monkeypatch):
"""Same guard on the /v1/messages path, where tools reach the hook directly."""
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
messages = [{"role": "user", "content": [{"type": "text", "text": "hi"}]}]
result_msgs, result_sys = AnthropicCacheControlHook.maybe_inject_cache_control(
copy.deepcopy(messages),
"sys",
{},
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
tools=self._tools(3, cached=True),
)
assert result_sys == "sys"
assert result_msgs == messages
def test_default_ttl_is_anthropics_five_minute_cache(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
assert all(p["control"] == {"type": "ephemeral"} for p in self._points())
@pytest.mark.parametrize("ttl", ["5m", "1h"])
def test_ttl_override_applied(self, monkeypatch, ttl):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
monkeypatch.setattr(litellm, "anthropic_prompt_caching_ttl", ttl)
assert all(p["control"] == {"type": "ephemeral", "ttl": ttl} for p in self._points())
def test_seed_does_not_override_configured_points(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
configured = [{"location": "message", "role": "user", "index": 0}]
params = {"cache_control_injection_points": configured}
AnthropicCacheControlHook.maybe_seed_default_injection_points(
non_default_params=params,
messages=copy.deepcopy(self.MESSAGES),
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
)
assert params["cache_control_injection_points"] is configured
def test_seed_adds_defaults_when_enabled(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
params: dict = {}
AnthropicCacheControlHook.maybe_seed_default_injection_points(
non_default_params=params,
messages=copy.deepcopy(self.MESSAGES),
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
)
assert [p["index"] for p in params["cache_control_injection_points"]] == [None, -1]
def test_seed_is_noop_when_disabled(self):
params: dict = {}
AnthropicCacheControlHook.maybe_seed_default_injection_points(
non_default_params=params,
messages=copy.deepcopy(self.MESSAGES),
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
)
assert params == {}
def test_v1_messages_applies_defaults_end_to_end(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
messages = [
{"role": "user", "content": [{"type": "text", "text": "first"}]},
{"role": "assistant", "content": [{"type": "text", "text": "reply"}]},
{"role": "user", "content": [{"type": "text", "text": "latest"}]},
]
result_msgs, result_sys = AnthropicCacheControlHook.maybe_inject_cache_control(
messages,
"a system prompt",
{},
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
)
assert result_sys == [{"type": "text", "text": "a system prompt", "cache_control": {"type": "ephemeral"}}]
assert result_msgs[-1]["content"][-1]["cache_control"] == {"type": "ephemeral"}
assert "cache_control" not in result_msgs[0]["content"][-1]
def test_v1_messages_is_noop_when_disabled(self):
messages = [{"role": "user", "content": [{"type": "text", "text": "hi"}]}]
result_msgs, result_sys = AnthropicCacheControlHook.maybe_inject_cache_control(
messages,
"sys",
{},
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
)
assert result_sys == "sys"
assert result_msgs == messages
class TestAnthropicPromptCachingEnvVars:
"""Both settings are read from the environment at import, so an admin can enable
auto-caching without a config file. Each case re-imports litellm in a subprocess
so the env is read fresh without contaminating this process's module graph.
"""
@staticmethod
def _import_litellm_with_env(env_override: dict) -> Tuple[bool, Optional[str]]:
env = os.environ.copy()
env.pop("LITELLM_ENABLE_ANTHROPIC_PROMPT_CACHING", None)
env.pop("LITELLM_ANTHROPIC_PROMPT_CACHING_TTL", None)
env.update(env_override)
script = textwrap.dedent(
"""
import json, litellm
print(json.dumps([litellm.enable_anthropic_prompt_caching, litellm.anthropic_prompt_caching_ttl]))
"""
)
result = subprocess.run(
[sys.executable, "-c", script], capture_output=True, text=True, env=env, timeout=300
)
assert result.returncode == 0, result.stderr
enabled, ttl = json.loads(result.stdout.strip().splitlines()[-1])
return enabled, ttl
def test_unset_env_leaves_auto_caching_off(self):
assert self._import_litellm_with_env({}) == (False, None)
@pytest.mark.parametrize("value", ["true", "True", "TRUE"])
def test_env_enables_auto_caching_case_insensitively(self, value):
enabled, _ = self._import_litellm_with_env({"LITELLM_ENABLE_ANTHROPIC_PROMPT_CACHING": value})
assert enabled is True
@pytest.mark.parametrize("value", ["false", "0", "yes", ""])
def test_env_only_enables_on_true(self, value):
enabled, _ = self._import_litellm_with_env({"LITELLM_ENABLE_ANTHROPIC_PROMPT_CACHING": value})
assert enabled is False
@pytest.mark.parametrize("value", ["5m", "1h"])
def test_ttl_env_is_applied(self, value):
_, ttl = self._import_litellm_with_env({"LITELLM_ANTHROPIC_PROMPT_CACHING_TTL": value})
assert ttl == value
@pytest.mark.parametrize("value", ["10m", "1H", "3600", "ephemeral"])
def test_unsupported_ttl_env_falls_back_to_provider_default(self, value):
"""An unparseable TTL must fall back to Anthropic's 5m default, never reach the provider verbatim."""
_, ttl = self._import_litellm_with_env({"LITELLM_ANTHROPIC_PROMPT_CACHING_TTL": value})
assert ttl is None

View file

@ -1261,6 +1261,23 @@ class TestAnthropicThinkingSignatureSelfHeal:
)
assert is_anthropic_invalid_thinking_signature_error(raw) is True
def test_is_anthropic_invalid_thinking_signature_error_positive_bedrock(self):
from litellm.llms.anthropic.common_utils import (
is_anthropic_invalid_thinking_signature_error,
)
# Real user-reported Bedrock scenario
raw = '{"message":"messages.2.content.0.thinking.signature.str: Input should be a valid string"}'
assert is_anthropic_invalid_thinking_signature_error(raw) is True
def test_is_anthropic_invalid_thinking_signature_error_positive_vertex(self):
from litellm.llms.anthropic.common_utils import (
is_anthropic_invalid_thinking_signature_error,
)
raw = "messages.4.content.1.thinking.signature.str: Input should be a valid string"
assert is_anthropic_invalid_thinking_signature_error(raw) is True
def test_is_anthropic_invalid_thinking_signature_error_negative(self):
from litellm.llms.anthropic.common_utils import (
is_anthropic_invalid_thinking_signature_error,
@ -1271,6 +1288,11 @@ class TestAnthropicThinkingSignatureSelfHeal:
is_anthropic_invalid_thinking_signature_error("rate limit exceeded")
is False
)
assert (
is_anthropic_invalid_thinking_signature_error("invalid_request_error: model not found")
is False
)
assert is_anthropic_invalid_thinking_signature_error("thinking signature is malformed") is False
def test_strip_thinking_blocks_from_anthropic_messages(self):
from litellm.llms.anthropic.common_utils import (

View file

@ -0,0 +1,66 @@
import os
import sys
import pytest
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.llms.fireworks_ai.cost_calculator import cost_per_token
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
MODEL = "accounts/fireworks/models/glm-5p2"
INPUT_COST = 1.4e-06
CACHE_READ_COST = 2.6e-07
OUTPUT_COST = 4.4e-06
def _usage(prompt_tokens: int, cached_tokens: int, completion_tokens: int) -> Usage:
return Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=cached_tokens),
)
def test_cached_prompt_tokens_billed_at_cache_read_rate():
prompt_tokens = 7036
cached_tokens = 7020
completion_tokens = 8
prompt_cost, completion_cost = cost_per_token(
model=MODEL, usage=_usage(prompt_tokens, cached_tokens, completion_tokens)
)
expected_prompt_cost = (prompt_tokens - cached_tokens) * INPUT_COST + cached_tokens * CACHE_READ_COST
assert prompt_cost == pytest.approx(expected_prompt_cost)
assert completion_cost == pytest.approx(completion_tokens * OUTPUT_COST)
full_rate_cost = prompt_tokens * INPUT_COST
assert prompt_cost < full_rate_cost
def test_warm_call_cheaper_than_cold_call():
prompt_tokens = 7036
completion_tokens = 8
cold_prompt_cost, _ = cost_per_token(
model=MODEL, usage=_usage(prompt_tokens, 16, completion_tokens)
)
warm_prompt_cost, _ = cost_per_token(
model=MODEL, usage=_usage(prompt_tokens, 7020, completion_tokens)
)
assert warm_prompt_cost < cold_prompt_cost
def test_no_cached_tokens_matches_full_input_rate():
prompt_tokens = 100
completion_tokens = 10
prompt_cost, completion_cost = cost_per_token(
model=MODEL, usage=_usage(prompt_tokens, 0, completion_tokens)
)
assert prompt_cost == pytest.approx(prompt_tokens * INPUT_COST)
assert completion_cost == pytest.approx(completion_tokens * OUTPUT_COST)

View file

@ -22,6 +22,7 @@ from litellm.proxy.proxy_server import (
_scrub_db_overlay_remote_module_loads,
_scrub_guardrail_inner,
resolve_complexity_router_plugins,
resolve_routing_plugins,
)
from .conftest import normalize
@ -185,6 +186,75 @@ def test_resolve_complexity_router_plugins_rejects_synchronous_run_method(tmp_pa
)
# ---------------------------------------------------------------------------
# resolve_routing_plugins
# ---------------------------------------------------------------------------
def test_resolve_routing_plugins_resolves_dotted_paths(tmp_path):
plugin_file = tmp_path / "rs_plugin.py"
plugin_file.write_text(
"class _Plugin:\n"
" async def run(self, context):\n"
" return context\n"
"\n"
"rs_plugin_instance = _Plugin()\n"
)
resolved = resolve_routing_plugins(
plugin_paths=["rs_plugin.rs_plugin_instance"],
config_file_path=str(tmp_path / "config.yaml"),
source_label="router_settings.plugins",
)
assert len(resolved) == 1
assert type(resolved[0]).__name__ == "_Plugin"
def test_resolve_routing_plugins_passes_through_instances(tmp_path):
class _Plugin:
async def run(self, context):
return context
instance = _Plugin()
resolved = resolve_routing_plugins(
plugin_paths=[instance],
config_file_path=None,
source_label="router_settings.plugins",
)
assert resolved == [instance]
def test_resolve_routing_plugins_rejects_non_routing_plugin(tmp_path):
plugin_file = tmp_path / "bad_rs_plugin.py"
plugin_file.write_text("not_a_plugin = object()\n")
with pytest.raises(ValueError, match="router_settings.plugins"):
resolve_routing_plugins(
plugin_paths=["bad_rs_plugin.not_a_plugin"],
config_file_path=str(tmp_path / "config.yaml"),
source_label="router_settings.plugins",
)
def test_resolve_routing_plugins_rejects_synchronous_run(tmp_path):
plugin_file = tmp_path / "sync_rs_plugin.py"
plugin_file.write_text(
"class _SyncPlugin:\n"
" def run(self, context):\n"
" return context\n"
"\n"
"sync_plugin_instance = _SyncPlugin()\n"
)
with pytest.raises(ValueError, match="does not implement the RoutingPlugin interface"):
resolve_routing_plugins(
plugin_paths=["sync_rs_plugin.sync_plugin_instance"],
config_file_path=str(tmp_path / "config.yaml"),
source_label="router_settings.plugins",
)
# ---------------------------------------------------------------------------
# ProxyConfig.__init__
# ---------------------------------------------------------------------------
@ -793,6 +863,62 @@ async def test_ProxyConfig_load_config_minimal_yaml(tmp_path, monkeypatch):
}
@pytest.mark.asyncio
async def test_ProxyConfig_load_config_resolves_router_settings_plugins(tmp_path, monkeypatch):
"""Regression: router_settings.plugins dotted-path strings must be resolved to
live RoutingPlugin instances on the created Router. Previously they were passed
through as raw strings and only blew up at request time when the pipeline tried
to `await "some.string".run(context)`."""
plugin_file = tmp_path / "rs_plugin.py"
plugin_file.write_text(
"class _Plugin:\n"
" async def run(self, context):\n"
" return context\n"
"\n"
"rs_plugin_instance = _Plugin()\n"
)
f = tmp_path / "c.yaml"
f.write_text(
"model_list: []\n"
"general_settings: {}\n"
"litellm_settings: {}\n"
"router_settings:\n"
" plugins:\n"
" - rs_plugin.rs_plugin_instance\n"
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
router, _model_list, _general_settings = await ProxyConfig().load_config(
router=None, config_file_path=str(f)
)
assert len(router.routing_plugins) == 1
assert type(router.routing_plugins[0]).__name__ == "_Plugin"
@pytest.mark.asyncio
async def test_ProxyConfig_load_config_rejects_bad_router_settings_plugin(tmp_path, monkeypatch):
plugin_file = tmp_path / "bad_rs_plugin.py"
plugin_file.write_text("not_a_plugin = object()\n")
f = tmp_path / "c.yaml"
f.write_text(
"model_list: []\n"
"general_settings: {}\n"
"litellm_settings: {}\n"
"router_settings:\n"
" plugins:\n"
" - bad_rs_plugin.not_a_plugin\n"
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
with pytest.raises(ValueError, match="does not implement the RoutingPlugin interface"):
await ProxyConfig().load_config(router=None, config_file_path=str(f))
@pytest.mark.asyncio
async def test_ProxyConfig_load_config_wires_general_settings_url_validation(tmp_path, monkeypatch):
"""Regression for #26599: SSRF settings in general_settings must reach litellm globals."""

View file

@ -242,3 +242,88 @@ class TestRagIngestSSRFBlocked:
assert response.status_code != 400, (
f"Clean Bedrock ingest_options should not be rejected: {response.json()}"
)
def test_rag_query_returns_response_cost_header(client_internal_user):
"""
/v1/rag/query must surface the completion cost via the
x-litellm-response-cost response header, like /v1/chat/completions does.
"""
from litellm.types.utils import ModelResponse
mock_response = ModelResponse(
id="chatcmpl-test",
choices=[
{
"index": 0,
"message": {"role": "assistant", "content": "The codename is AZURE-FALCON-42."},
"finish_reason": "stop",
}
],
model="gpt-4o-mini",
usage={"prompt_tokens": 35, "completion_tokens": 14, "total_tokens": 49},
)
mock_response._hidden_params["response_cost"] = 3.45e-06
with patch(
"litellm.proxy.rag_endpoints.endpoints.litellm.aquery",
new_callable=AsyncMock,
return_value=mock_response,
), patch("litellm.vector_store_registry", None), patch(
"litellm.proxy.proxy_server.prisma_client", None
):
response = client_internal_user.post(
"/v1/rag/query",
json={
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "What is the codename?"}],
"retrieval_config": {
"vector_store_id": "vs_test_123",
"custom_llm_provider": "openai",
},
},
)
assert response.status_code == 200, response.json()
assert response.headers.get("x-litellm-response-cost") == "3.45e-06"
def test_rag_query_stream_returns_event_stream(client_internal_user):
"""
A stream=true /v1/rag/query must return an SSE response. Returning the raw
stream wrapper makes FastAPI try to serialize it, which raises and turns
every streaming RAG query into a 500; the stream then never drains, so its
single billing event (which carries the folded sub-call costs) never fires.
"""
import litellm as litellm_module
async def fake_aquery(**kwargs):
return await litellm_module.acompletion(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "What is the codename?"}],
mock_response="The codename is AZURE-FALCON-42.",
stream=True,
api_key="test-key",
)
with patch(
"litellm.proxy.rag_endpoints.endpoints.litellm.aquery",
new=AsyncMock(side_effect=fake_aquery),
), patch("litellm.vector_store_registry", None), patch("litellm.proxy.proxy_server.prisma_client", None):
response = client_internal_user.post(
"/v1/rag/query",
json={
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "What is the codename?"}],
"retrieval_config": {
"vector_store_id": "vs_test_123",
"custom_llm_provider": "openai",
},
"stream": True,
},
)
assert response.status_code == 200, response.text
assert response.headers.get("content-type", "").startswith("text/event-stream")
assert '"object":"chat.completion.chunk"' in response.text
assert "data: [DONE]" in response.text

View file

@ -64,6 +64,65 @@ def test_dotted_module_path_is_unaffected_by_gate():
assert result == "loaded"
def test_installed_package_resolved_when_local_file_absent(tmp_path, monkeypatch):
# Regression: with config_file_path set (startup load path) but no local
# module file next to it, get_instance_fn must fall back to importing the
# dotted name as an installed package. Previously it raised ImportError
# ("Could not find module file ..."), so plugins shipped as pip packages
# (e.g. router_settings/complexity_router plugins) could not be referenced.
pkg_dir = tmp_path / "site"
pkg_dir.mkdir()
(pkg_dir / "my_installed_plugin.py").write_text(
"class _P:\n"
" async def run(self, context):\n"
" return context\n"
"\n"
"instance = _P()\n"
)
monkeypatch.syspath_prepend(str(pkg_dir))
config_dir = tmp_path / "cfg"
config_dir.mkdir()
result = get_instance_fn(
value="my_installed_plugin.instance",
config_file_path=str(config_dir / "config.yaml"),
)
assert type(result).__name__ == "_P"
def test_local_module_file_wins_over_installed_package(tmp_path, monkeypatch):
# A local module file next to the config must still take precedence over an
# installed package of the same dotted name -- the fallback only kicks in
# when no local file exists.
pkg_dir = tmp_path / "site"
pkg_dir.mkdir()
(pkg_dir / "shadowed_mod.py").write_text("value = 'from-installed'\n")
monkeypatch.syspath_prepend(str(pkg_dir))
config_dir = tmp_path / "cfg"
config_dir.mkdir()
(config_dir / "shadowed_mod.py").write_text("value = 'from-local-file'\n")
result = get_instance_fn(
value="shadowed_mod.value",
config_file_path=str(config_dir / "config.yaml"),
)
assert result == "from-local-file"
def test_missing_module_everywhere_raises_import_error(tmp_path):
# Neither a local file nor an installed package: the fallback import must
# surface a real ImportError rather than silently succeeding.
config_dir = tmp_path / "cfg"
config_dir.mkdir()
with pytest.raises(ImportError):
get_instance_fn(
value="definitely_not_a_real_module_xyz.instance",
config_file_path=str(config_dir / "config.yaml"),
)
def test_pass_through_route_threads_config_file_path():
# ``create_pass_through_route`` must forward ``config_file_path`` so
# an operator with ``custom_handler: s3://...`` declared in

View file

View file

@ -0,0 +1,266 @@
"""
Tests for the RAG query pipeline in litellm/rag/main.py.
The RAG pipeline forwards its kwargs (including the parent litellm_logging_obj)
into @client-decorated sub-calls (vector store search, completion). Each logging
object allows exactly one async_success event, so if sub-calls are not marked as
internal, the vector store search consumes the slot first and the LLM
completion's usage/cost is never logged (spend tracking and budget enforcement
are bypassed). These tests pin the invariant that the single billing event for
aquery carries the completion response with real usage and cost.
"""
import asyncio
from unittest.mock import patch
import pytest
import litellm
from litellm._internal_context import is_internal_call
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.utils import CallTypes, ModelResponse
class RecordingLogger(CustomLogger):
def __init__(self):
super().__init__()
self.success_events = []
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
self.success_events.append({"kwargs": kwargs, "response_obj": response_obj})
@pytest.mark.asyncio
@pytest.mark.parametrize("use_router", [False, True])
async def test_aquery_single_billing_event_carries_completion_usage_and_cost(use_router):
"""
litellm.aquery must produce exactly one success event, and that event must
carry the LLM completion (a ModelResponse with non-zero usage and cost),
not the vector store search response. The proxy always passes a router, so
both the router and non-router completion branches are pinned.
"""
recording_logger = RecordingLogger()
original_callbacks = litellm.callbacks
litellm.callbacks = [recording_logger]
router_kwargs = {}
if use_router:
router_kwargs["router"] = litellm.Router(
model_list=[
{
"model_name": "gpt-4o-mini",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "test-key"},
}
]
)
try:
response = await litellm.aquery(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "What is the secret project codename?"}],
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
mock_response="The secret project codename is AZURE-FALCON-42.",
**router_kwargs,
)
assert isinstance(response, ModelResponse)
assert is_internal_call.get() is False
for _ in range(50):
if recording_logger.success_events:
break
await asyncio.sleep(0.1)
await asyncio.sleep(0.5)
finally:
litellm.callbacks = original_callbacks
assert len(recording_logger.success_events) == 1
event = recording_logger.success_events[0]
response_obj = event["response_obj"]
assert isinstance(response_obj, ModelResponse)
assert response_obj.usage.total_tokens > 0
standard_logging_object = event["kwargs"]["standard_logging_object"]
assert standard_logging_object["call_type"] == "aquery"
assert standard_logging_object["total_tokens"] > 0
assert standard_logging_object["prompt_tokens"] > 0
assert standard_logging_object["completion_tokens"] > 0
assert standard_logging_object["response_cost"] > 0
@pytest.mark.asyncio
async def test_aquery_response_hidden_params_carry_completion_cost():
"""
The aquery response must expose the completion's response_cost via hidden
params, so the proxy can return the x-litellm-response-cost header.
"""
response = await litellm.aquery(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hello"}],
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
mock_response="hi there",
)
assert isinstance(response, ModelResponse)
response_cost = response._hidden_params.get("response_cost")
assert response_cost is not None
assert response_cost > 0
@pytest.mark.asyncio
async def test_aquery_billed_cost_includes_priced_vector_store_search():
"""
When the vector store provider prices search calls (e.g. per-query cost),
that cost must be folded into the aquery billing instead of being dropped
with the suppressed sub-call event.
"""
recording_logger = RecordingLogger()
original_callbacks = litellm.callbacks
litellm.callbacks = [recording_logger]
try:
with patch("litellm.rag.main.vector_store_search_cost", return_value=(0.002, 0.0)):
response = await litellm.aquery(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hello"}],
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
mock_response="hi there",
)
for _ in range(50):
if recording_logger.success_events:
break
await asyncio.sleep(0.1)
await asyncio.sleep(0.5)
finally:
litellm.callbacks = original_callbacks
assert isinstance(response, ModelResponse)
total_cost = response._hidden_params.get("response_cost")
assert total_cost is not None
assert total_cost > 0.002
assert len(recording_logger.success_events) == 1
standard_logging_object = recording_logger.success_events[0]["kwargs"]["standard_logging_object"]
assert standard_logging_object["response_cost"] == total_cost
@pytest.mark.asyncio
async def test_aquery_with_rerank_bills_once_and_folds_rerank_cost():
"""
When rerank is enabled, its sub-call must run under the internal-call
context (no standalone billing event) and its cost must be folded into
the single aquery billing event.
"""
from litellm.types.rerank import RerankResponse
recording_logger = RecordingLogger()
original_callbacks = litellm.callbacks
litellm.callbacks = [recording_logger]
rerank_seen = {}
async def fake_arerank(**kwargs):
rerank_seen["internal"] = is_internal_call.get()
rerank_result = RerankResponse(id="rr_1", results=[{"index": 0, "relevance_score": 0.9}], meta={})
rerank_result._hidden_params["response_cost"] = 0.001
return rerank_result
try:
with patch("litellm.arerank", side_effect=fake_arerank):
response = await litellm.aquery(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hello"}],
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
rerank={"enabled": True, "model": "cohere/rerank-english-v3.0", "top_n": 1},
mock_response="hi there",
)
for _ in range(50):
if recording_logger.success_events:
break
await asyncio.sleep(0.1)
await asyncio.sleep(0.5)
finally:
litellm.callbacks = original_callbacks
assert rerank_seen["internal"] is True
assert is_internal_call.get() is False
assert isinstance(response, ModelResponse)
total_cost = response._hidden_params.get("response_cost")
assert total_cost is not None
assert total_cost > 0.001
assert len(recording_logger.success_events) == 1
standard_logging_object = recording_logger.success_events[0]["kwargs"]["standard_logging_object"]
assert standard_logging_object["call_type"] == "aquery"
assert standard_logging_object["response_cost"] == total_cost
@pytest.mark.asyncio
async def test_aquery_streaming_bills_sub_call_costs_into_final_event():
"""
On the streaming path the response cost is computed from the assembled
chunks after the pipeline returns, so there is no response object to fold
sub-call costs into. The pipeline must instead carry the accumulated
search and rerank cost through the logging object so the single streamed
billing event includes it; otherwise a caller passing stream=true incurs
priced vector search and rerank costs that never reach spend tracking.
"""
from litellm.types.rerank import RerankResponse
recording_logger = RecordingLogger()
original_callbacks = litellm.callbacks
litellm.callbacks = [recording_logger]
rerank_seen = {}
async def fake_arerank(**kwargs):
rerank_seen["internal"] = is_internal_call.get()
rerank_result = RerankResponse(id="rr_1", results=[{"index": 0, "relevance_score": 0.9}], meta={})
rerank_result._hidden_params["response_cost"] = 0.001
return rerank_result
try:
with (
patch("litellm.rag.main.vector_store_search_cost", return_value=(0.002, 0.0)),
patch("litellm.arerank", side_effect=fake_arerank),
):
response = await litellm.aquery(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hello"}],
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
rerank={"enabled": True, "model": "cohere/rerank-english-v3.0", "top_n": 1},
mock_response="hi there",
stream=True,
)
async for _ in response:
pass
for _ in range(50):
if recording_logger.success_events:
break
await asyncio.sleep(0.1)
await asyncio.sleep(0.5)
finally:
litellm.callbacks = original_callbacks
assert rerank_seen["internal"] is True
assert is_internal_call.get() is False
assert len(recording_logger.success_events) == 1
standard_logging_object = recording_logger.success_events[0]["kwargs"]["standard_logging_object"]
assert standard_logging_object["call_type"] == "aquery"
assert standard_logging_object["response_cost"] >= 0.003
def test_rag_call_types_are_registered():
"""
query/aquery/ingest/aingest are @client-decorated entry points, so their
function names must resolve to CallTypes members (deployment hooks and
call-type driven logic silently no-op for unregistered call types).
"""
assert CallTypes("query") is CallTypes.query
assert CallTypes("aquery") is CallTypes.aquery
assert CallTypes("ingest") is CallTypes.ingest
assert CallTypes("aingest") is CallTypes.aingest

View file

@ -3119,3 +3119,260 @@ class TestRoutingPlugins:
assert first.model == "gpt-4o-mini"
assert second.model == "gpt-4o-mini"
assert spy.call_count == 2
class TestEscalationKeywords:
"""Test user-triggered escalation: a keyword in the prompt bumps the resolved tier
one step higher so a user can force a stronger model when unhappy with results."""
@staticmethod
def _request_kwargs(session_id: str) -> Dict:
return {"metadata": {"session_id": session_id}}
def test_default_escalation_keyword(self, complexity_router):
assert complexity_router.escalation_keywords == ["LITELLM ESCALATE"]
def test_escalation_triggered_is_case_sensitive(self, complexity_router):
assert complexity_router._escalation_triggered("please LITELLM ESCALATE now") is True
assert complexity_router._escalation_triggered("please litellm escalate now") is False
assert complexity_router._escalation_triggered("how do I escalate this ticket") is False
def test_escalate_tier_bumps_one_step(self, complexity_router):
assert complexity_router._escalate_tier(ComplexityTier.SIMPLE) == ComplexityTier.MEDIUM
assert complexity_router._escalate_tier(ComplexityTier.MEDIUM) == ComplexityTier.COMPLEX
assert complexity_router._escalate_tier(ComplexityTier.COMPLEX) == ComplexityTier.REASONING
def test_escalate_tier_caps_at_highest_configured(self, complexity_router):
assert complexity_router._escalate_tier(ComplexityTier.REASONING) == ComplexityTier.REASONING
def test_escalate_tier_skips_unconfigured_intermediate(self, mock_router_instance):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={"tiers": {"SIMPLE": "gpt-4o-mini", "REASONING": "o1-preview"}},
)
assert router._escalate_tier(ComplexityTier.SIMPLE) == ComplexityTier.REASONING
def test_tier_for_model_returns_most_severe(self, mock_router_instance):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"tiers": {"SIMPLE": "shared", "COMPLEX": "shared", "REASONING": "top"}
},
)
assert router._tier_for_model("shared") == ComplexityTier.COMPLEX
assert router._tier_for_model("top") == ComplexityTier.REASONING
assert router._tier_for_model("unknown") is None
@pytest.mark.asyncio
async def test_escalation_bumps_classified_tier(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config=basic_config,
)
# Baseline: this prompt classifies SIMPLE.
baseline = await router.async_pre_routing_hook(
model="test-model", request_kwargs={}, messages=[{"role": "user", "content": "Hello there!"}]
)
assert baseline.model == "gpt-4o-mini"
escalated = await router.async_pre_routing_hook(
model="test-model",
request_kwargs={},
messages=[{"role": "user", "content": "LITELLM ESCALATE Hello there!"}],
)
assert escalated.model == "gpt-4o" # SIMPLE bumped to MEDIUM
@pytest.mark.asyncio
async def test_lowercase_keyword_does_not_escalate(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config=basic_config,
)
result = await router.async_pre_routing_hook(
model="test-model",
request_kwargs={},
messages=[{"role": "user", "content": "litellm escalate Hello there!"}],
)
assert result.model == "gpt-4o-mini" # not escalated
@pytest.mark.asyncio
async def test_custom_escalation_keyword(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={**basic_config, "escalation_keywords": ["MAKE IT BETTER"]},
)
# The default keyword no longer triggers once a custom list is supplied.
default = await router.async_pre_routing_hook(
model="test-model",
request_kwargs={},
messages=[{"role": "user", "content": "LITELLM ESCALATE Hello there!"}],
)
assert default.model == "gpt-4o-mini"
custom = await router.async_pre_routing_hook(
model="test-model",
request_kwargs={},
messages=[{"role": "user", "content": "MAKE IT BETTER Hello there!"}],
)
assert custom.model == "gpt-4o"
@pytest.mark.asyncio
async def test_empty_keyword_list_disables_escalation(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={**basic_config, "escalation_keywords": []},
)
result = await router.async_pre_routing_hook(
model="test-model",
request_kwargs={},
messages=[{"role": "user", "content": "LITELLM ESCALATE Hello there!"}],
)
assert result.model == "gpt-4o-mini"
@pytest.mark.asyncio
async def test_escalation_caps_at_highest_tier(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config=basic_config,
)
result = await router.async_pre_routing_hook(
model="test-model",
request_kwargs={},
messages=[
{
"role": "user",
"content": "LITELLM ESCALATE Let's think step by step and reason through this carefully.",
}
],
)
assert result.model == "o1-preview" # already REASONING, stays there
@pytest.mark.asyncio
async def test_escalation_bumps_keyword_tier_override(self, mock_router_instance, basic_config):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
**basic_config,
"keyword_tier_rules": [{"keywords": ["billing"], "tier": "SIMPLE"}],
},
)
baseline = await router.async_pre_routing_hook(
model="test-model", request_kwargs={}, messages=[{"role": "user", "content": "a billing question"}]
)
assert baseline.model == "gpt-4o-mini"
escalated = await router.async_pre_routing_hook(
model="test-model",
request_kwargs={},
messages=[{"role": "user", "content": "LITELLM ESCALATE a billing question"}],
)
assert escalated.model == "gpt-4o" # override SIMPLE bumped to MEDIUM
@pytest.mark.asyncio
async def test_escalation_overrides_session_pin_and_persists(self, mock_router_instance, basic_config):
"""Mid-session escalation bumps relative to the pinned model (never below it) and
the bumped model persists for later turns."""
mock_router_instance.cache = DualCache()
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={**basic_config, "session_affinity": True},
)
request_kwargs = self._request_kwargs("session-1")
first = await router.async_pre_routing_hook(
model="test-model", request_kwargs=request_kwargs, messages=[{"role": "user", "content": "Hello!"}]
)
assert first.model == "gpt-4o-mini" # pinned SIMPLE
with patch.object(router, "aclassify", wraps=router.aclassify) as spy_aclassify:
escalated = await router.async_pre_routing_hook(
model="test-model",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "LITELLM ESCALATE"}],
)
spy_aclassify.assert_not_called()
assert escalated.model == "gpt-4o" # bumped relative to the SIMPLE pin, not reclassified
# The bump persists: a later ordinary turn stays on the escalated model.
later = await router.async_pre_routing_hook(
model="test-model", request_kwargs=request_kwargs, messages=[{"role": "user", "content": "thanks"}]
)
assert later.model == "gpt-4o"
# Escalating again climbs one more tier.
again = await router.async_pre_routing_hook(
model="test-model",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "LITELLM ESCALATE still not good"}],
)
assert again.model == "claude-sonnet-4-20250514" # MEDIUM bumped to COMPLEX
def test_blank_escalation_keywords_are_stripped(self):
"""Blank/whitespace-only phrases are dropped so `"" in message` can't escalate
every request; surrounding whitespace on real phrases is trimmed."""
assert ComplexityRouterConfig(
tiers={"SIMPLE": "gpt-4o-mini", "MEDIUM": "gpt-4o"},
escalation_keywords=["", " "],
).escalation_keywords == []
assert ComplexityRouterConfig(
tiers={"SIMPLE": "gpt-4o-mini", "MEDIUM": "gpt-4o"},
escalation_keywords=[" LITELLM ESCALATE ", ""],
).escalation_keywords == ["LITELLM ESCALATE"]
@pytest.mark.asyncio
async def test_blank_escalation_keyword_does_not_escalate_everything(
self, mock_router_instance, basic_config
):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={**basic_config, "escalation_keywords": [""]},
)
assert router.escalation_keywords == []
result = await router.async_pre_routing_hook(
model="test-model",
request_kwargs={},
messages=[{"role": "user", "content": "Hello there!"}],
)
assert result.model == "gpt-4o-mini" # not escalated
def test_escalated_pin_stays_on_same_model_at_ceiling(self, mock_router_instance):
"""At the highest configured tier escalation keeps the exact pinned model, even
when that tier's pool has peers `get_model_for_tier` could randomly pick instead."""
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"tiers": {"SIMPLE": "gpt-4o-mini", "REASONING": ["o1-a", "o1-b", "o1-c"]}
},
)
for pinned in ("o1-a", "o1-b", "o1-c"):
assert router._escalated_pin(pinned) == pinned
@pytest.mark.asyncio
async def test_session_escalation_at_ceiling_keeps_multi_model_pin(self, mock_router_instance):
mock_router_instance.cache = DualCache()
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"tiers": {"SIMPLE": "gpt-4o-mini", "REASONING": ["o1-a", "o1-b", "o1-c"]},
"session_affinity": True,
},
)
cache_key = router._get_session_affinity_cache_key("session-top", {})
await mock_router_instance.cache.async_set_cache(key=cache_key, value="o1-b")
result = await router.async_pre_routing_hook(
model="test-model",
request_kwargs=self._request_kwargs("session-top"),
messages=[{"role": "user", "content": "LITELLM ESCALATE do better"}],
)
assert result.model == "o1-b" # unchanged: no random hop to o1-a / o1-c

View file

@ -0,0 +1,82 @@
import json
import typing
from pathlib import Path
import pytest
import litellm
from litellm.types.utils import ModelInfoBase
REALTIME_ONLY_GPT_MODELS = (
"azure/gpt-realtime-2025-08-28",
"azure/gpt-realtime-1.5-2026-02-23",
"azure/gpt-realtime-mini-2025-10-06",
"gpt-realtime",
"gpt-realtime-1.5",
"gpt-realtime-2",
"gpt-realtime-2.1",
"gpt-realtime-2.1-mini",
"gpt-realtime-mini",
"gpt-realtime-2025-08-28",
"gpt-realtime-mini-2025-10-06",
"gpt-realtime-mini-2025-12-15",
)
REALTIME_ONLY_GPT_MODELS_WITHOUT_ENDPOINTS = (
"azure/eu/gpt-4o-mini-realtime-preview-2024-12-17",
"azure/eu/gpt-4o-realtime-preview-2024-10-01",
"azure/eu/gpt-4o-realtime-preview-2024-12-17",
"azure/gpt-4o-mini-realtime-preview-2024-12-17",
"azure/gpt-4o-realtime-preview-2024-10-01",
"azure/gpt-4o-realtime-preview-2024-12-17",
"azure/us/gpt-4o-mini-realtime-preview-2024-12-17",
"azure/us/gpt-4o-realtime-preview-2024-10-01",
"azure/us/gpt-4o-realtime-preview-2024-12-17",
"gpt-4o-mini-realtime-preview",
"gpt-4o-mini-realtime-preview-2024-12-17",
"gpt-4o-realtime-preview",
"gpt-4o-realtime-preview-2024-12-17",
"gpt-4o-realtime-preview-2025-06-03",
)
ALL_REALTIME_ONLY_GPT_MODELS = REALTIME_ONLY_GPT_MODELS + REALTIME_ONLY_GPT_MODELS_WITHOUT_ENDPOINTS
def _load_cost_map() -> dict:
json_path = Path(__file__).parents[2] / "model_prices_and_context_window.json"
with open(json_path) as f:
return json.load(f)
def test_realtime_is_a_valid_mode_literal():
hints = typing.get_type_hints(ModelInfoBase, include_extras=False)
assert "realtime" in typing.get_args(hints["mode"])
@pytest.mark.parametrize("model", REALTIME_ONLY_GPT_MODELS)
def test_realtime_only_gpt_models_are_mode_realtime(model):
"""These models only serve /v1/realtime and are rejected by /v1/chat/completions
("This is not a chat model ..."), so they must not be tagged mode=chat."""
info = _load_cost_map()[model]
assert info["supported_endpoints"] == ["/v1/realtime"]
assert info["mode"] == "realtime"
@pytest.mark.parametrize("model", REALTIME_ONLY_GPT_MODELS_WITHOUT_ENDPOINTS)
def test_realtime_only_gpt_4o_models_are_mode_realtime(model):
"""gpt-4o(-mini)-realtime-preview are realtime-only and must not be mode=chat."""
assert _load_cost_map()[model]["mode"] == "realtime"
def test_get_model_info_reports_realtime_mode():
assert litellm.get_model_info("gpt-realtime-mini")["mode"] == "realtime"
def test_backup_matches_main_for_realtime_models():
repo_root = Path(__file__).parents[2]
with open(repo_root / "model_prices_and_context_window.json") as f:
main_cost = json.load(f)
with open(repo_root / "litellm" / "model_prices_and_context_window_backup.json") as f:
backup_cost = json.load(f)
for model in ALL_REALTIME_ONLY_GPT_MODELS:
assert backup_cost.get(model) == main_cost.get(model)

View file

@ -269,4 +269,22 @@ describe("ComplexityRouterConfig", () => {
);
expect(screen.getAllByText("This tier is required")).toHaveLength(1);
});
it("renders the escalation keywords section with current keywords when the handler is provided", () => {
renderWithProviders(
<ComplexityRouterConfig
{...baseProps}
escalationKeywords={["LITELLM ESCALATE"]}
onEscalationKeywordsChange={vi.fn()}
/>,
);
fireEvent.click(screen.getByText("Advanced: Escalation Keywords"));
expect(screen.getByText("Escalation Keywords")).toBeInTheDocument();
expect(screen.getByText("LITELLM ESCALATE")).toBeInTheDocument();
});
it("hides the escalation keywords section when no handler is provided", () => {
renderWithProviders(<ComplexityRouterConfig {...baseProps} />);
expect(screen.queryByText("Advanced: Escalation Keywords")).not.toBeInTheDocument();
});
});

View file

@ -4,6 +4,7 @@ import React from "react";
import { ModelGroup } from "@/components/llm_calls/fetch_models";
import AdaptiveRoutingConfig from "./AdaptiveRoutingConfig";
import ClassificationMethodConfig from "./ClassificationMethodConfig";
import EscalationKeywords from "./EscalationKeywords";
import KeywordTierRules, { KeywordTierRule } from "./KeywordTierRules";
import SemanticKeywordMatching from "./SemanticKeywordMatching";
@ -61,6 +62,8 @@ interface ComplexityRouterConfigProps {
onEmbeddingModelChange?: (model: string) => void;
matchThreshold?: number;
onMatchThresholdChange?: (threshold: number) => void;
escalationKeywords?: string[];
onEscalationKeywordsChange?: (keywords: string[]) => void;
showValidationErrors?: boolean;
}
@ -101,6 +104,8 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
onEmbeddingModelChange = () => {},
matchThreshold = 0.5,
onMatchThresholdChange = () => {},
escalationKeywords = [],
onEscalationKeywordsChange,
showValidationErrors = false,
}) => {
// Embedding models can't serve a chat-completion role, so they're excluded here.
@ -213,6 +218,19 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
),
children: <AdaptiveRoutingConfig value={value} onChange={onChange} />,
},
...(onEscalationKeywordsChange
? [
{
key: "escalation",
label: (
<Text strong style={{ color: "#374151" }}>
Advanced: Escalation Keywords
</Text>
),
children: <EscalationKeywords keywords={escalationKeywords} onChange={onEscalationKeywordsChange} />,
},
]
: []),
...(onKeywordTierRulesChange || onSemanticMatchingEnabledChange
? [
{

View file

@ -0,0 +1,45 @@
import { InfoCircleOutlined } from "@ant-design/icons";
import { Select as AntdSelect, Tooltip, Typography } from "antd";
import React from "react";
const { Text } = Typography;
export const DEFAULT_ESCALATION_KEYWORDS = ["LITELLM ESCALATE"];
interface EscalationKeywordsProps {
keywords: string[];
onChange: (keywords: string[]) => void;
}
const EscalationKeywords: React.FC<EscalationKeywordsProps> = ({ keywords, onChange }) => {
return (
<div className="w-full max-w-none">
<div className="flex items-center gap-2 mb-1">
<Typography.Title level={4} style={{ margin: 0 }}>
Escalation Keywords
</Typography.Title>
<Tooltip title="Case-sensitive phrases a user can include in their message to force a bump to the next-higher complexity tier when they aren't happy with results. They can force a stronger model, but not choose which one.">
<InfoCircleOutlined className="text-gray-400" />
</Tooltip>
</div>
<Text type="secondary" style={{ display: "block", marginBottom: 8, fontSize: 12 }}>
Optional: when a user message contains one of these phrases, the request is bumped one tier higher than it would
otherwise route to. Matching is case-sensitive, so &quot;LITELLM ESCALATE&quot; only fires on the exact, shouted
form. Leave empty to disable.
</Text>
<AntdSelect
mode="tags"
value={keywords}
onChange={onChange}
placeholder="e.g., LITELLM ESCALATE"
tokenSeparators={[","]}
open={false}
suffixIcon={null}
style={{ width: "100%" }}
allowClear
/>
</div>
);
};
export default EscalationKeywords;

View file

@ -14,6 +14,7 @@ import ComplexityRouterConfig, {
DEFAULT_TIER_DISTANCE_PENALTY,
} from "./ComplexityRouterConfig";
import { KeywordTierRule } from "./KeywordTierRules";
import { DEFAULT_ESCALATION_KEYWORDS } from "./EscalationKeywords";
import { DEFAULT_MATCH_THRESHOLD } from "./SemanticKeywordMatching";
import {
buildComplexityRouterConfig,
@ -52,6 +53,7 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
const [semanticMatchingEnabled, setSemanticMatchingEnabled] = useState<boolean>(false);
const [embeddingModel, setEmbeddingModel] = useState<string | undefined>(undefined);
const [matchThreshold, setMatchThreshold] = useState<number>(DEFAULT_MATCH_THRESHOLD);
const [escalationKeywords, setEscalationKeywords] = useState<string[]>(DEFAULT_ESCALATION_KEYWORDS);
const [showValidationErrors, setShowValidationErrors] = useState<boolean>(false);
// Semantic router config (existing)
@ -141,6 +143,7 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
semanticMatchingEnabled,
embeddingModel,
matchThreshold,
escalationKeywords,
adaptive,
adaptiveWeights,
tierDistancePenalty,
@ -316,6 +319,8 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
onEmbeddingModelChange={setEmbeddingModel}
matchThreshold={matchThreshold}
onMatchThresholdChange={setMatchThreshold}
escalationKeywords={escalationKeywords}
onEscalationKeywordsChange={setEscalationKeywords}
showValidationErrors={showValidationErrors}
/>
</div>

View file

@ -21,6 +21,7 @@ const baseParams: BuildComplexityRouterConfigParams = {
semanticMatchingEnabled: false,
embeddingModel: undefined,
matchThreshold: 0.5,
escalationKeywords: ["LITELLM ESCALATE"],
adaptive: false,
adaptiveWeights: { quality: 0.3, cost: 0.7 },
tierDistancePenalty: 0.5,
@ -28,9 +29,22 @@ const baseParams: BuildComplexityRouterConfigParams = {
};
describe("buildComplexityRouterConfig", () => {
it("emits only tiers and classifier_type when nothing else is configured", () => {
it("emits tiers, classifier_type, and escalation_keywords when nothing else is configured", () => {
const config = buildComplexityRouterConfig(baseParams);
expect(config).toEqual({ tiers, classifier_type: "heuristic" });
expect(config).toEqual({ tiers, classifier_type: "heuristic", escalation_keywords: ["LITELLM ESCALATE"] });
});
it("trims escalation keywords and drops blank entries", () => {
const config = buildComplexityRouterConfig({
...baseParams,
escalationKeywords: [" LITELLM ESCALATE ", "", " ", "MAKE IT BETTER"],
});
expect(config.escalation_keywords).toEqual(["LITELLM ESCALATE", "MAKE IT BETTER"]);
});
it("emits an empty escalation_keywords list so clearing the field disables escalation", () => {
const config = buildComplexityRouterConfig({ ...baseParams, escalationKeywords: [] });
expect(config.escalation_keywords).toEqual([]);
});
it("passes through a tier configured with more than one model as a pool", () => {

View file

@ -16,6 +16,7 @@ export interface BuildComplexityRouterConfigParams {
semanticMatchingEnabled: boolean;
embeddingModel: string | undefined;
matchThreshold: number;
escalationKeywords: string[];
adaptive: boolean;
adaptiveWeights: AdaptiveRouterWeights;
tierDistancePenalty: number;
@ -31,6 +32,7 @@ export interface ComplexityRouterConfigPayload {
semantic_keyword_matching?: boolean;
embedding_model?: string;
match_threshold?: number;
escalation_keywords?: string[];
adaptive?: boolean;
adaptive_weights?: AdaptiveRouterWeights;
tier_distance_penalty?: number;
@ -69,11 +71,13 @@ export const buildComplexityRouterConfig = ({
semanticMatchingEnabled,
embeddingModel,
matchThreshold,
escalationKeywords,
adaptive,
adaptiveWeights,
tierDistancePenalty,
adaptiveEligible,
}: BuildComplexityRouterConfigParams): ComplexityRouterConfigPayload => {
const cleanedEscalationKeywords = escalationKeywords.map((keyword) => keyword.trim()).filter(Boolean);
// Trim keywords and drop empty ones; drop any rule left with no keywords. Clicking
// "Add keyword rule" seeds a rule with an empty keywords list, so without this an
// unfilled row (common in the heuristic flow, where getSemanticConfigError doesn't run)
@ -88,6 +92,7 @@ export const buildComplexityRouterConfig = ({
...(classifierType === "llm" && classifierLlmConfig && { classifier_llm_config: classifierLlmConfig }),
...(customTechnicalKeywords.length > 0 && { custom_technical_keywords: customTechnicalKeywords }),
...(cleanedKeywordTierRules.length > 0 && { keyword_tier_rules: cleanedKeywordTierRules }),
escalation_keywords: cleanedEscalationKeywords,
...(semanticMatchingEnabled && {
semantic_keyword_matching: true,
embedding_model: embeddingModel,

View file

@ -21690,7 +21690,7 @@ export interface components {
* CallTypes
* @enum {string}
*/
CallTypes: "embedding" | "aembedding" | "completion" | "acompletion" | "atext_completion" | "text_completion" | "image_generation" | "aimage_generation" | "image_edit" | "aimage_edit" | "moderation" | "amoderation" | "atranscription" | "transcription" | "aspeech" | "speech" | "rerank" | "arerank" | "search" | "asearch" | "_arealtime" | "_aresponses_websocket" | "create_batch" | "acreate_batch" | "aretrieve_batch" | "retrieve_batch" | "acancel_batch" | "cancel_batch" | "pass_through_endpoint" | "anthropic_messages" | "aanthropic_messages" | "get_assistants" | "aget_assistants" | "create_assistants" | "acreate_assistants" | "delete_assistant" | "adelete_assistant" | "acreate_thread" | "create_thread" | "aget_thread" | "get_thread" | "a_add_message" | "add_message" | "aget_messages" | "get_messages" | "arun_thread" | "run_thread" | "arun_thread_stream" | "run_thread_stream" | "afile_retrieve" | "file_retrieve" | "afile_delete" | "file_delete" | "afile_list" | "file_list" | "acreate_file" | "create_file" | "afile_content" | "file_content" | "create_fine_tuning_job" | "acreate_fine_tuning_job" | "create_video" | "acreate_video" | "avideo_retrieve" | "video_retrieve" | "avideo_content" | "video_content" | "video_remix" | "avideo_remix" | "video_list" | "avideo_list" | "video_retrieve_job" | "avideo_retrieve_job" | "video_delete" | "avideo_delete" | "video_create_character" | "avideo_create_character" | "video_get_character" | "avideo_get_character" | "video_edit" | "avideo_edit" | "video_extension" | "avideo_extension" | "vector_store_file_create" | "avector_store_file_create" | "vector_store_file_list" | "avector_store_file_list" | "vector_store_file_retrieve" | "avector_store_file_retrieve" | "vector_store_file_content" | "avector_store_file_content" | "vector_store_file_update" | "avector_store_file_update" | "vector_store_file_delete" | "avector_store_file_delete" | "vector_store_create" | "avector_store_create" | "vector_store_search" | "avector_store_search" | "create_container" | "acreate_container" | "list_containers" | "alist_containers" | "retrieve_container" | "aretrieve_container" | "delete_container" | "adelete_container" | "list_container_files" | "alist_container_files" | "upload_container_file" | "aupload_container_file" | "create_sandbox" | "acreate_sandbox" | "delete_sandbox" | "adelete_sandbox" | "run_code" | "arun_code" | "code_interpreter_tool" | "acode_interpreter_tool" | "acancel_fine_tuning_job" | "cancel_fine_tuning_job" | "alist_fine_tuning_jobs" | "list_fine_tuning_jobs" | "aretrieve_fine_tuning_job" | "retrieve_fine_tuning_job" | "responses" | "aresponses" | "alist_input_items" | "llm_passthrough_route" | "allm_passthrough_route" | "generate_content" | "agenerate_content" | "generate_content_stream" | "agenerate_content_stream" | "ocr" | "aocr" | "call_mcp_tool" | "list_mcp_tools" | "asend_message" | "send_message" | "acreate_skill";
CallTypes: "embedding" | "aembedding" | "completion" | "acompletion" | "atext_completion" | "text_completion" | "image_generation" | "aimage_generation" | "image_edit" | "aimage_edit" | "moderation" | "amoderation" | "atranscription" | "transcription" | "aspeech" | "speech" | "rerank" | "arerank" | "search" | "asearch" | "_arealtime" | "_aresponses_websocket" | "create_batch" | "acreate_batch" | "aretrieve_batch" | "retrieve_batch" | "acancel_batch" | "cancel_batch" | "pass_through_endpoint" | "anthropic_messages" | "aanthropic_messages" | "get_assistants" | "aget_assistants" | "create_assistants" | "acreate_assistants" | "delete_assistant" | "adelete_assistant" | "acreate_thread" | "create_thread" | "aget_thread" | "get_thread" | "a_add_message" | "add_message" | "aget_messages" | "get_messages" | "arun_thread" | "run_thread" | "arun_thread_stream" | "run_thread_stream" | "afile_retrieve" | "file_retrieve" | "afile_delete" | "file_delete" | "afile_list" | "file_list" | "acreate_file" | "create_file" | "afile_content" | "file_content" | "create_fine_tuning_job" | "acreate_fine_tuning_job" | "create_video" | "acreate_video" | "avideo_retrieve" | "video_retrieve" | "avideo_content" | "video_content" | "video_remix" | "avideo_remix" | "video_list" | "avideo_list" | "video_retrieve_job" | "avideo_retrieve_job" | "video_delete" | "avideo_delete" | "video_create_character" | "avideo_create_character" | "video_get_character" | "avideo_get_character" | "video_edit" | "avideo_edit" | "video_extension" | "avideo_extension" | "vector_store_file_create" | "avector_store_file_create" | "vector_store_file_list" | "avector_store_file_list" | "vector_store_file_retrieve" | "avector_store_file_retrieve" | "vector_store_file_content" | "avector_store_file_content" | "vector_store_file_update" | "avector_store_file_update" | "vector_store_file_delete" | "avector_store_file_delete" | "vector_store_create" | "avector_store_create" | "vector_store_search" | "avector_store_search" | "ingest" | "aingest" | "query" | "aquery" | "create_container" | "acreate_container" | "list_containers" | "alist_containers" | "retrieve_container" | "aretrieve_container" | "delete_container" | "adelete_container" | "list_container_files" | "alist_container_files" | "upload_container_file" | "aupload_container_file" | "create_sandbox" | "acreate_sandbox" | "delete_sandbox" | "adelete_sandbox" | "run_code" | "arun_code" | "code_interpreter_tool" | "acode_interpreter_tool" | "acancel_fine_tuning_job" | "cancel_fine_tuning_job" | "alist_fine_tuning_jobs" | "list_fine_tuning_jobs" | "aretrieve_fine_tuning_job" | "retrieve_fine_tuning_job" | "responses" | "aresponses" | "alist_input_items" | "llm_passthrough_route" | "allm_passthrough_route" | "generate_content" | "agenerate_content" | "generate_content_stream" | "agenerate_content_stream" | "ocr" | "aocr" | "call_mcp_tool" | "list_mcp_tools" | "asend_message" | "send_message" | "acreate_skill";
/** CallbackDelete */
CallbackDelete: {
/** Callback Name */
@ -21870,6 +21870,11 @@ export interface components {
};
/** ChatCompletionCachedContent */
ChatCompletionCachedContent: {
/**
* Ttl
* @enum {string}
*/
ttl?: "5m" | "1h";
/**
* Type
* @constant