Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_mcp_server_env_vars

# Conflicts:
#	tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py
This commit is contained in:
mateo-berri 2026-06-05 15:37:18 +00:00
commit aa78a90a58
No known key found for this signature in database
45 changed files with 4304 additions and 147 deletions

View file

@ -42,7 +42,7 @@ When you must use real LLM models to, for example, write e2e tests, write a QA r
If you're an internal contributor, when creating a new PR, the typical flow is to branch off litellm_internal_staging and create a branch prefixed with litellm_. Do not create a branch prefixed with claude/ and generally do not have / in your branch names
Do not add `Co-Authored-By: Claude` or any Claude attribution to commit messages. Never use a `claude/` prefix or put a `/` in a branch name. Do not add "Generated with Claude Code" (or any similar attribution) to PR descriptions. Do not create a new PR/branch off the existing PR to fix/add something that is related and could've just been committed directly to the existing PR's branch
Do not add `Co-Authored-By: Claude` or any Claude attribution to commit messages. Never use a `claude/` prefix or put a `/` in a branch name. Do not add "Generated with Claude Code" (or any similar attribution) to PR descriptions or comments. Do not create a new PR/branch off the existing PR to fix/add something that is related and could've just been committed directly to the existing PR's branch
When working on a PR, keep the PR description in sync with new commits being made

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "oauth2_flow" TEXT;

View file

@ -327,6 +327,7 @@ model LiteLLM_MCPServerTable {
authorization_url String?
token_url String?
registration_url String?
oauth2_flow String?
allow_all_keys Boolean @default(false)
available_on_public_internet Boolean @default(true)
delegate_auth_to_upstream Boolean @default(false)

View file

@ -2425,12 +2425,11 @@ class BaseTokenUsageProcessor:
if not attr.startswith("_") and not callable(
getattr(usage.completion_tokens_details, attr)
):
current_val = getattr(
combined.completion_tokens_details, attr, 0
current_val = (
getattr(combined.completion_tokens_details, attr, 0) or 0
)
new_val = getattr(usage.completion_tokens_details, attr, 0)
if new_val is not None and current_val is not None:
new_val = getattr(usage.completion_tokens_details, attr, 0) or 0
if isinstance(new_val, (int, float)):
setattr(
combined.completion_tokens_details,
attr,

View file

@ -1062,3 +1062,37 @@ class GuardrailInterventionNormalStringError(
def __repr__(self):
return self.__str__()
class SensitiveDataRouteException(Exception):
"""
Exception raised when a guardrail detects sensitive data and wants to reroute the request.
Instead of blocking the request, this exception signals that the request should be
routed to a different model (typically an on-premise model for data privacy).
The proxy catches this exception and:
1. Reroutes the current request to the specified model
2. When sticky_session_routing is True, stores the routing decision in session
cache so all subsequent requests in the same session are routed to the same model
"""
def __init__(
self,
route_to_model: str,
session_id: str,
guardrail_name: Optional[str] = None,
detection_info: Optional[Dict[str, Any]] = None,
message: Optional[str] = None,
sticky_session_routing: bool = True,
):
self.route_to_model = route_to_model
self.session_id = session_id
self.guardrail_name = guardrail_name
self.detection_info = detection_info or {}
self.sticky_session_routing = sticky_session_routing
self.message = (
message
or f"Sensitive data detected by {guardrail_name}. Routing to model: {route_to_model}"
)
super().__init__(self.message)

View file

@ -47,9 +47,29 @@ from litellm.exceptions import (
BlockedPiiEntityError,
GuardrailRaisedException,
ModifyResponseException,
SensitiveDataRouteException,
)
def get_session_id_from_request_data(request_data: Dict[str, Any]) -> Optional[str]:
"""Extract session_id from request data (litellm_session_id or metadata)."""
session_id = request_data.get("litellm_session_id")
if session_id:
return str(session_id)
metadata = request_data.get("metadata") or {}
session_id = metadata.get("session_id")
if session_id:
return str(session_id)
litellm_metadata = request_data.get("litellm_metadata") or {}
session_id = litellm_metadata.get("session_id")
if session_id:
return str(session_id)
return None
class CustomGuardrail(CustomLogger):
# If True, during_call runs async_moderation_hook instead of the unified apply_guardrail path.
use_native_during_call_hook: ClassVar[bool] = False
@ -68,6 +88,9 @@ class CustomGuardrail(CustomLogger):
end_session_after_n_fails: Optional[int] = None,
on_violation: Optional[str] = None,
realtime_violation_message: Optional[str] = None,
on_sensitive_data: Optional[str] = None,
sensitive_data_route_to_model: Optional[str] = None,
sticky_session_routing: bool = True,
**kwargs,
):
"""
@ -83,6 +106,9 @@ class CustomGuardrail(CustomLogger):
end_session_after_n_fails: For /v1/realtime sessions, end the session after this many violations
on_violation: For /v1/realtime sessions, 'warn' or 'end_session'
realtime_violation_message: Message the bot speaks aloud when a /v1/realtime guardrail fires
on_sensitive_data: Action when sensitive data is detected. 'block' (default) or 'route'
sensitive_data_route_to_model: Model to route to when on_sensitive_data='route'
sticky_session_routing: When True, all subsequent requests in the session use the same model
"""
self.guardrail_name = guardrail_name
self.supported_event_hooks = supported_event_hooks
@ -96,6 +122,11 @@ class CustomGuardrail(CustomLogger):
self.end_session_after_n_fails: Optional[int] = end_session_after_n_fails
self.on_violation: Optional[str] = on_violation
self.realtime_violation_message: Optional[str] = realtime_violation_message
self.on_sensitive_data: Optional[str] = on_sensitive_data
self.sensitive_data_route_to_model: Optional[str] = (
sensitive_data_route_to_model
)
self.sticky_session_routing: bool = sticky_session_routing
if supported_event_hooks:
## validate event_hook is in supported_event_hooks
@ -167,6 +198,108 @@ class CustomGuardrail(CustomLogger):
detection_info=detection_info,
)
def raise_sensitive_data_route_exception(
self,
route_to_model: str,
request_data: Dict[str, Any],
detection_info: Optional[Dict[str, Any]] = None,
) -> None:
"""
Raise an exception to reroute the request to a different model.
Use this when sensitive data is detected and the guardrail is configured
to route to an on-premise model instead of blocking.
The exception will reroute this request to the specified model. When
sticky_session_routing is enabled (the default), it also stores the
routing decision so subsequent requests in this session reuse the model.
Args:
route_to_model: The model to route this request (and session) to
request_data: The original request data dictionary
detection_info: Optional non-sensitive detection metadata (e.g. matched
entity types, rule ids, scores). This is surfaced in request metadata
and logs, so it must not contain the raw detected sensitive values.
Raises:
SensitiveDataRouteException: Always raises to trigger rerouting
"""
session_id = self._get_session_id_from_request_data(request_data)
if not session_id:
raise ValueError(
"Cannot route sensitive data without a session_id. "
"Ensure the request includes a session_id in metadata or headers."
)
raise SensitiveDataRouteException(
route_to_model=route_to_model,
session_id=session_id,
guardrail_name=self.guardrail_name,
detection_info=detection_info,
sticky_session_routing=self.sticky_session_routing,
)
def _get_session_id_from_request_data(
self, request_data: Dict[str, Any]
) -> Optional[str]:
"""Extract session_id from request data."""
return get_session_id_from_request_data(request_data)
def should_route_on_sensitive_data(self) -> bool:
"""
Returns True if this guardrail is configured to route requests
to a different model when sensitive data is detected.
"""
return (
self.on_sensitive_data == "route"
and self.sensitive_data_route_to_model is not None
)
def handle_sensitive_data_detection(
self,
request_data: Dict[str, Any],
detection_info: Optional[Dict[str, Any]] = None,
) -> None:
"""
Handle sensitive data detection based on guardrail configuration.
If on_sensitive_data='route', raises SensitiveDataRouteException to reroute.
Otherwise, raises GuardrailRaisedException to block. When routing is
configured but the request carries no session_id, routing is not possible
so the request falls back to a graceful block.
Args:
request_data: The request data dictionary
detection_info: Optional non-sensitive detection metadata. When routing,
this is surfaced in request metadata and logs, so it must not contain
the raw detected sensitive values.
Raises:
SensitiveDataRouteException: When configured to route and a session_id is present
GuardrailRaisedException: When configured to block, or when routing is
configured but no session_id is available
"""
if self.should_route_on_sensitive_data():
try:
self.raise_sensitive_data_route_exception(
route_to_model=self.sensitive_data_route_to_model, # type: ignore
request_data=request_data,
detection_info=detection_info,
)
except ValueError:
raise GuardrailRaisedException(
message=(
f"Sensitive data detected by {self.guardrail_name} "
"(routing skipped: request has no session_id)"
),
guardrail_name=self.guardrail_name,
)
else:
raise GuardrailRaisedException(
message=f"Sensitive data detected by {self.guardrail_name}",
guardrail_name=self.guardrail_name,
)
@staticmethod
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
"""
@ -753,12 +886,20 @@ class CustomGuardrail(CustomLogger):
Guardrails signal intentional blocks by raising:
- GuardrailRaisedException (generic guardrail API, tool permission)
- BlockedPiiEntityError (Presidio PII detection)
- SensitiveDataRouteException (sensitive-data reroute to on-premise model)
- HTTPException with status 400 (content policy violation)
- ModifyResponseException (passthrough mode violation)
"""
if isinstance(e, ModifyResponseException):
return True
if isinstance(e, (GuardrailRaisedException, BlockedPiiEntityError)):
if isinstance(
e,
(
GuardrailRaisedException,
BlockedPiiEntityError,
SensitiveDataRouteException,
),
):
return True
if (
HTTPException is not None

View file

@ -2690,7 +2690,7 @@ class PrometheusLogger(CustomLogger):
Args:
guardrail_name: Name of the guardrail
latency_seconds: Execution latency in seconds
status: "success" or "error"
status: "success", "error", or "intervened"
error_type: Type of error if any, None otherwise
hook_type: "pre_call", "during_call", or "post_call"
"""

View file

@ -81,7 +81,6 @@ from litellm.types.utils import (
from litellm.utils import (
ModelResponse,
Usage,
_supports_factory,
add_dummy_tool,
any_assistant_message_has_thinking_blocks,
get_max_tokens,
@ -337,50 +336,6 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
v in model_lower for v in ("opus-4-7", "opus_4_7", "opus-4.7", "opus_4.7")
)
@staticmethod
def _supports_model_capability(model: str, key: str) -> bool:
"""Check a boolean capability ``key`` in the model map.
Strips bedrock/vertex prefixes so a provider-routed Claude still
resolves to the Anthropic model-map entry.
"""
try:
if _supports_factory(
model=model,
custom_llm_provider="anthropic",
key=key,
):
return True
except Exception:
pass
candidates = [model]
for prefix in (
"bedrock/converse/",
"bedrock/invoke/",
"bedrock/",
"vertex_ai/",
):
if model.startswith(prefix):
candidates.append(model[len(prefix) :])
try:
from litellm.llms.bedrock.common_utils import BedrockModelInfo
base = BedrockModelInfo.get_base_model(model)
if base:
candidates.append(base)
candidates.append(f"bedrock/{base}")
except Exception:
pass
try:
for cand in candidates:
if cand in litellm.model_cost and (
litellm.model_cost[cand].get(key) is True
):
return True
except Exception:
pass
return False
@staticmethod
def _supports_effort_level(model: str, level: str) -> bool:
"""Check ``supports_{level}_reasoning_effort`` in the model map."""
@ -918,7 +873,39 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
anthropic_tools = []
mcp_servers = []
for tool in tools:
if "input_schema" in tool: # assume in anthropic format
if tool.get("type") == "namespace":
# Namespace is a grouping container (e.g. codex's multi_agent_v1).
# Extract its nested tools and map them individually.
for nested in tool.get("tools") or []:
if "input_schema" in nested:
# Already in Anthropic format.
anthropic_tools.append(nested)
elif "function" not in nested and "name" in nested:
# Flat format: {type, name, description, parameters, ...}.
# Normalize to OpenAI-wrapped format before mapping.
wrapped = cast(
ChatCompletionToolParam,
{
"type": nested.get("type", "function"),
"function": {
k: v for k, v in nested.items() if k != "type"
},
},
)
nested_tool, nested_mcp = self._map_tool_helper(wrapped)
if nested_tool is not None:
anthropic_tools.append(nested_tool)
if nested_mcp is not None:
mcp_servers.append(nested_mcp)
elif "function" in nested:
nested_tool, nested_mcp = self._map_tool_helper(
cast(ChatCompletionToolParam, nested)
)
if nested_tool is not None:
anthropic_tools.append(nested_tool)
if nested_mcp is not None:
mcp_servers.append(nested_mcp)
elif "input_schema" in tool: # assume in anthropic format
anthropic_tools.append(tool)
else: # assume openai tool call
new_tool, mcp_server_tool = self._map_tool_helper(tool)
@ -1978,6 +1965,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
# Remove internal LiteLLM parameters that should not be sent to Anthropic API
optional_params.pop("is_vertex_request", None)
optional_params.pop("client_metadata", None)
data = {
"model": model,

View file

@ -272,19 +272,63 @@ class AnthropicModelInfo(BaseLLMModelInfo):
)
@staticmethod
def _is_adaptive_thinking_model(model: str) -> bool:
"""Claude 4.6+ models use adaptive thinking with ``output_config.effort``."""
def _supports_model_capability(model: str, key: str) -> bool:
"""Check a boolean capability ``key`` in the model map.
Strips bedrock/vertex prefixes so a provider-routed Claude still
resolves to the Anthropic model-map entry.
"""
from litellm.utils import _supports_factory
try:
if _supports_factory(
model=model,
custom_llm_provider=None,
key="supports_adaptive_thinking",
custom_llm_provider="anthropic",
key=key,
):
return True
except Exception:
pass
candidates = [model]
for prefix in (
"bedrock/converse/",
"bedrock/invoke/",
"bedrock/",
"vertex_ai/",
):
if model.startswith(prefix):
candidates.append(model[len(prefix) :])
try:
from litellm.llms.bedrock.common_utils import BedrockModelInfo
base = BedrockModelInfo.get_base_model(model)
if base:
candidates.append(base)
candidates.append(f"bedrock/{base}")
except Exception:
pass
try:
for cand in candidates:
if cand in litellm.model_cost and (
litellm.model_cost[cand].get(key) is True
):
return True
except Exception:
pass
return False
@staticmethod
def _is_adaptive_thinking_model(model: str) -> bool:
"""Claude 4.6+ models use adaptive thinking with ``output_config.effort``.
Driven by the ``supports_adaptive_thinking`` flag in the model map; the
4.6/4.7 name checks remain only as a fallback for provider-routed ids
whose map entries predate the flag.
"""
if AnthropicModelInfo._supports_model_capability(
model, "supports_adaptive_thinking"
):
return True
return AnthropicModelInfo._is_claude_4_6_model(
model
) or AnthropicModelInfo._is_claude_4_7_model(model)

View file

@ -2586,6 +2586,8 @@ class BaseLLMHTTPHandler:
headers=headers,
)
headers.setdefault("Content-Type", "application/json")
## LOGGING
logging_obj.pre_call(
input=input,
@ -2676,6 +2678,8 @@ class BaseLLMHTTPHandler:
headers=headers,
)
headers.setdefault("Content-Type", "application/json")
## LOGGING
logging_obj.pre_call(
input=input,

View file

@ -1319,6 +1319,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1350,6 +1351,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1381,6 +1383,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1412,6 +1415,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1443,6 +1447,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -2194,6 +2199,7 @@
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -33927,6 +33933,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -33955,6 +33962,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,

View file

@ -1406,6 +1406,7 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
authorization_url: Optional[str] = None
token_url: Optional[str] = None
registration_url: Optional[str] = None
oauth2_flow: Optional[Literal["client_credentials", "authorization_code"]] = None
allow_all_keys: bool = False
available_on_public_internet: bool = True
delegate_auth_to_upstream: bool = False
@ -1480,6 +1481,7 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
authorization_url: Optional[str] = None
token_url: Optional[str] = None
registration_url: Optional[str] = None
oauth2_flow: Optional[Literal["client_credentials", "authorization_code"]] = None
allow_all_keys: bool = False
available_on_public_internet: bool = True
delegate_auth_to_upstream: bool = False

View file

@ -15,6 +15,10 @@ from fastapi.responses import JSONResponse, StreamingResponse
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.url_utils import SSRFError, validate_url
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints.databricks_oauth import (
DATABRICKS_OAUTH_PARAM,
resolve_databricks_app_auth_header,
)
from litellm.proxy.agent_endpoints.utils import merge_agent_headers
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.types.utils import all_litellm_params
@ -677,6 +681,17 @@ async def invoke_agent_a2a( # noqa: PLR0915
static_headers=static_headers or None,
)
# Databricks App endpoints require a short-lived OAuth M2M token rather
# than a static bearer. Only agents explicitly configured with a
# ``databricks_oauth`` block get one; every other agent is left untouched.
if litellm_params.get(DATABRICKS_OAUTH_PARAM):
databricks_auth = await resolve_databricks_app_auth_header(litellm_params)
if databricks_auth:
agent_extra_headers = {
**(agent_extra_headers or {}),
**databricks_auth,
}
# Merge agent-level guardrails into data so post_call_success_hook and
# _handle_stream_message both pick them up. A2A agents use model
# a2a_agent/*, which is not an llm_router deployment, so

View file

@ -0,0 +1,250 @@
"""
OAuth M2M (client_credentials) support for A2A agents that target Databricks
App endpoints.
Databricks Apps reject static bearer tokens; they require a short-lived OAuth
access token minted from the workspace OIDC token endpoint. When an agent is
registered with a ``databricks_oauth`` block in its ``litellm_params``, LiteLLM
fetches that token via the client_credentials grant, caches it until shortly
before expiry, and attaches it as the outbound ``Authorization`` header on every
call the proxy makes to the agent.
Config example::
agents:
- agent_name: my-databricks-app
agent_card_params:
url: https://my-app-1234.aws.databricksapps.com
litellm_params:
databricks_oauth:
client_id: os.environ/DATABRICKS_CLIENT_ID
client_secret: os.environ/DATABRICKS_CLIENT_SECRET
workspace_url: https://dbc-abc123.cloud.databricks.com
"""
import asyncio
import base64
import hashlib
from dataclasses import dataclass
from typing import Any, Dict, Optional, Tuple
import httpx
from litellm._logging import verbose_logger
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.custom_http import httpxSpecialProvider
DATABRICKS_OAUTH_PARAM = "databricks_oauth"
_DEFAULT_SCOPE = "all-apis"
_TOKEN_EXPIRY_BUFFER_SECONDS = 60
_DEFAULT_TTL_SECONDS = 3600
def _resolve_secret(value: Any) -> Optional[str]:
"""Resolve a config value, expanding ``os.environ/`` references."""
if not isinstance(value, str):
return None
if value.startswith("os.environ/"):
return get_secret_str(value)
return value
def _token_url_from_workspace(workspace_url: str) -> str:
"""Build the workspace OIDC token endpoint from a workspace URL."""
base = workspace_url.strip().rstrip("/")
if base.endswith("/serving-endpoints"):
base = base[: -len("/serving-endpoints")]
return f"{base}/oidc/v1/token"
@dataclass(frozen=True)
class DatabricksAppOAuthConfig:
client_id: str
client_secret: str
token_url: str
scope: str
@property
def cache_key(self) -> str:
# Include a digest of the secret so a rotated client_secret yields a new
# key and forces a fresh token instead of serving the stale one.
secret_digest = hashlib.sha256(self.client_secret.encode()).hexdigest()[:16]
return f"{self.token_url}|{self.client_id}|{self.scope}|{secret_digest}"
def parse_databricks_oauth_config(
litellm_params: Optional[Dict[str, Any]],
) -> Optional[DatabricksAppOAuthConfig]:
"""Build a Databricks App OAuth config from an agent's ``litellm_params``.
Returns ``None`` when the agent has no ``databricks_oauth`` block. Raises
``ValueError`` when the block is present but incomplete, so misconfiguration
surfaces loudly instead of silently sending an unauthenticated request.
"""
if not litellm_params:
return None
raw = litellm_params.get(DATABRICKS_OAUTH_PARAM)
if raw is None:
return None
if not isinstance(raw, dict):
raise ValueError(
f"'{DATABRICKS_OAUTH_PARAM}' must be a mapping of OAuth settings, "
f"got {type(raw).__name__}"
)
client_id = _resolve_secret(raw.get("client_id"))
client_secret = _resolve_secret(raw.get("client_secret"))
workspace_url = _resolve_secret(raw.get("workspace_url"))
missing = [
name
for name, value in (
("client_id", client_id),
("client_secret", client_secret),
("workspace_url", workspace_url),
)
if not value
]
if missing:
raise ValueError(
f"Databricks App OAuth config is missing required field(s): "
f"{', '.join(missing)}"
)
scope = _resolve_secret(raw.get("scope")) or _DEFAULT_SCOPE
return DatabricksAppOAuthConfig(
client_id=client_id, # type: ignore[arg-type]
client_secret=client_secret, # type: ignore[arg-type]
token_url=_token_url_from_workspace(workspace_url), # type: ignore[arg-type]
scope=scope,
)
class DatabricksAppOAuthTokenCache(InMemoryCache):
"""In-memory cache for Databricks App OAuth client_credentials tokens.
Keyed by token endpoint + client_id + scope so distinct agents and service
principals never share a token. A per-key ``asyncio.Lock`` collapses
concurrent fetches into a single token request.
"""
def __init__(self) -> None:
super().__init__(default_ttl=_DEFAULT_TTL_SECONDS)
self._locks: Dict[str, asyncio.Lock] = {}
def _get_lock(self, cache_key: str) -> asyncio.Lock:
return self._locks.setdefault(cache_key, asyncio.Lock())
def _remove_key(self, key: str) -> None:
# Drop the per-key lock alongside the cached token so ``_locks`` stays
# bounded by the live key set rather than growing for every key ever seen.
super()._remove_key(key)
self._locks.pop(key, None)
def flush_cache(self) -> None:
super().flush_cache()
self._locks.clear()
async def async_get_token(self, config: DatabricksAppOAuthConfig) -> str:
cache_key = config.cache_key
cached = self.get_cache(cache_key)
if cached is not None:
return cached
async with self._get_lock(cache_key):
cached = self.get_cache(cache_key)
if cached is not None:
return cached
token, ttl = await self._fetch_token(config)
# ttl == 0 means the token's own lifetime is shorter than the
# refresh buffer; skip caching so we never hand out a stale token,
# and drop the lock we just created since no cached entry will ever
# trigger _remove_key to clean it up.
if ttl > 0:
self.set_cache(cache_key, token, ttl=ttl)
else:
self._locks.pop(cache_key, None)
return token
async def _fetch_token(self, config: DatabricksAppOAuthConfig) -> Tuple[str, int]:
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.A2A)
verbose_logger.debug(
"Fetching Databricks App OAuth token from %s", config.token_url
)
basic_auth = base64.b64encode(
f"{config.client_id}:{config.client_secret}".encode()
).decode()
try:
response = await client.post(
config.token_url,
data={
"grant_type": "client_credentials",
"scope": config.scope,
},
headers={
"Authorization": f"Basic {basic_auth}",
"Content-Type": "application/x-www-form-urlencoded",
},
)
except httpx.HTTPStatusError as exc:
raise ValueError(
"Databricks App OAuth token request failed with status "
f"{exc.response.status_code}"
) from exc
except httpx.HTTPError as exc:
raise ValueError(
f"Databricks App OAuth token request failed: {exc}"
) from exc
body = response.json()
if not isinstance(body, dict):
raise ValueError(
"Databricks App OAuth token response returned non-object JSON "
f"(got {type(body).__name__})"
)
access_token = body.get("access_token")
if not access_token:
raise ValueError(
"Databricks App OAuth token response missing 'access_token'"
)
raw_expires_in = body.get("expires_in")
try:
expires_in = (
int(raw_expires_in)
if raw_expires_in is not None
else _DEFAULT_TTL_SECONDS
)
except (TypeError, ValueError):
expires_in = _DEFAULT_TTL_SECONDS
ttl = max(expires_in - _TOKEN_EXPIRY_BUFFER_SECONDS, 0)
return access_token, ttl
databricks_app_oauth_token_cache = DatabricksAppOAuthTokenCache()
async def resolve_databricks_app_auth_header(
litellm_params: Optional[Dict[str, Any]],
) -> Optional[Dict[str, str]]:
"""Return ``{"Authorization": "Bearer <token>"}`` for a Databricks App agent.
Returns ``None`` when the agent is not configured for Databricks App OAuth.
"""
config = parse_databricks_oauth_config(litellm_params)
if config is None:
return None
token = await databricks_app_oauth_token_cache.async_get_token(config)
return {"Authorization": f"Bearer {token}"}

View file

@ -10,6 +10,7 @@ from .max_iterations_limiter import _PROXY_MaxIterationsHandler
from .parallel_request_limiter import _PROXY_MaxParallelRequestsHandler
from .parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3
from .responses_id_security import ResponsesIDSecurity
from .sensitive_data_routing import _PROXY_SensitiveDataRoutingHandler
# List of all available hooks that can be enabled.
# Defined before the enterprise import below so that any module re-imported
@ -23,6 +24,7 @@ PROXY_HOOKS = {
"litellm_skills": SkillsInjectionHook,
"max_iterations_limiter": _PROXY_MaxIterationsHandler,
"max_budget_per_session_limiter": _PROXY_MaxBudgetPerSessionHandler,
"sensitive_data_routing": _PROXY_SensitiveDataRoutingHandler,
}
## FEATURE FLAG HOOKS ##

View file

@ -32,6 +32,10 @@ from litellm.batches.batch_utils import (
)
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import SpecialModelNames, UserAPIKeyAuth
from litellm.proxy.hooks.rate_limiter_utils import (
ProxyHTTPRateLimitError,
resolve_llm_provider_for_rate_limit,
)
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
@ -375,6 +379,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
descriptors: List["RateLimitDescriptor"],
batch_usage: BatchFileUsage,
limit_type: str,
requested_model: Optional[str] = None,
) -> None:
"""Raise HTTPException for rate limit exceeded."""
from datetime import datetime
@ -419,7 +424,10 @@ class _PROXY_BatchRateLimiter(CustomLogger):
f"Limit resets at: {reset_time_formatted}"
)
raise HTTPException(
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
requested_model
)
raise ProxyHTTPRateLimitError(
status_code=429,
detail=detail,
headers={
@ -427,6 +435,8 @@ class _PROXY_BatchRateLimiter(CustomLogger):
"rate_limit_type": limit_type,
"reset_at": reset_time_formatted,
},
model=resolved_model,
llm_provider=llm_provider,
)
async def _check_and_increment_batch_counters(
@ -470,6 +480,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
)
if rate_limit_response["overall_code"] == "OVER_LIMIT":
requested_model = data.get("model") if data else None
for status in rate_limit_response["statuses"]:
if status["code"] == "OVER_LIMIT":
self._raise_rate_limit_error(
@ -477,6 +488,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
descriptors,
batch_usage,
status["rate_limit_type"],
requested_model=requested_model,
)
async def count_input_file_usage(

View file

@ -6,20 +6,21 @@ import asyncio
import os
from typing import List, Optional, Tuple, Union
from fastapi import HTTPException
import litellm
from litellm import ModelResponse, Router
from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.hooks.rate_limiter_utils import (
ProxyHTTPRateLimitError,
convert_priority_to_percent,
resolve_llm_provider_for_rate_limit,
)
from litellm.types.router import ModelGroupInfo
from litellm.types.utils import CallTypesLiteral
from litellm.utils import get_utc_datetime
from .rate_limiter_utils import convert_priority_to_percent
class DynamicRateLimiterCache:
"""
@ -218,7 +219,10 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger):
)
### CHECK TPM ###
if available_tpm is not None and available_tpm == 0:
raise HTTPException(
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
data.get("model")
)
raise ProxyHTTPRateLimitError(
status_code=429,
detail={
"error": "Key={} over available TPM={}. Model TPM={}, Active keys={}".format(
@ -228,10 +232,15 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger):
active_projects,
)
},
model=resolved_model,
llm_provider=llm_provider,
)
### CHECK RPM ###
elif available_rpm is not None and available_rpm == 0:
raise HTTPException(
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
data.get("model")
)
raise ProxyHTTPRateLimitError(
status_code=429,
detail={
"error": "Key={} over available RPM={}. Model RPM={}, Active keys={}".format(
@ -241,6 +250,8 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger):
active_projects,
)
},
model=resolved_model,
llm_provider=llm_provider,
)
elif available_rpm is not None or available_tpm is not None:
## UPDATE CACHE WITH ACTIVE PROJECT

View file

@ -19,7 +19,11 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import (
RateLimitDescriptorRateLimitObject,
_PROXY_MaxParallelRequestsHandler_v3,
)
from litellm.proxy.hooks.rate_limiter_utils import convert_priority_to_percent
from litellm.proxy.hooks.rate_limiter_utils import (
ProxyHTTPRateLimitError,
convert_priority_to_percent,
resolve_llm_provider_for_rate_limit,
)
from litellm.proxy.utils import InternalUsageCache
from litellm.types.router import ModelGroupInfo
from litellm.types.utils import CallTypesLiteral
@ -487,12 +491,13 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
)
if atomic_response["overall_code"] == "OVER_LIMIT":
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(model)
for status in atomic_response["statuses"]:
if status["code"] != "OVER_LIMIT":
continue
descriptor_key = status["descriptor_key"]
if descriptor_key == "model_saturation_check":
raise HTTPException(
raise ProxyHTTPRateLimitError(
status_code=429,
detail={
"error": f"Model capacity reached for {model}. "
@ -507,13 +512,15 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
"rate_limit_type": str(status["rate_limit_type"]),
"x-litellm-priority": priority or "default",
},
model=resolved_model,
llm_provider=llm_provider,
)
if descriptor_key == "priority_model":
verbose_proxy_logger.debug(
f"Enforcing priority limits for {model}, saturation: {saturation:.1%}, "
f"priority: {priority}"
)
raise HTTPException(
raise ProxyHTTPRateLimitError(
status_code=429,
detail={
"error": f"Priority-based rate limit exceeded. "
@ -531,6 +538,8 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
"x-litellm-priority": priority or "default",
"x-litellm-saturation": f"{saturation:.2%}",
},
model=resolved_model,
llm_provider=llm_provider,
)
# Fail-closed guard: overall_code says OVER_LIMIT but no status
@ -547,7 +556,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
f"Dynamic rate limiter: OVER_LIMIT response with unknown "
f"descriptor_key(s) — refusing request. response={atomic_response}"
)
raise HTTPException(
raise ProxyHTTPRateLimitError(
status_code=429,
detail={
"error": "Rate limit exceeded",
@ -562,6 +571,8 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
"retry-after": str(self.v3_limiter.window_size),
"x-litellm-priority": priority or "default",
},
model=resolved_model,
llm_provider=llm_provider,
)
# If priority is NOT enforced (saturation below threshold) but

View file

@ -5,6 +5,10 @@ from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.hooks.rate_limiter_utils import (
ProxyHTTPRateLimitError,
resolve_llm_provider_for_rate_limit,
)
class _PROXY_MaxBudgetLimiter(CustomLogger):
@ -63,7 +67,15 @@ class _PROXY_MaxBudgetLimiter(CustomLogger):
# CHECK IF REQUEST ALLOWED
if curr_spend >= max_budget:
raise HTTPException(status_code=429, detail="Max budget limit reached.")
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
data.get("model") if data else None
)
raise ProxyHTTPRateLimitError(
status_code=429,
detail="Max budget limit reached.",
model=resolved_model,
llm_provider=llm_provider,
)
except HTTPException as e:
raise e
except Exception as e:

View file

@ -17,12 +17,14 @@ Follows the same pattern as max_iterations_limiter.py.
import os
from typing import TYPE_CHECKING, Any, Optional, Union
from fastapi import HTTPException
from litellm import DualCache
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.hooks.rate_limiter_utils import (
ProxyHTTPRateLimitError,
resolve_llm_provider_for_rate_limit,
)
if TYPE_CHECKING:
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
@ -112,13 +114,18 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
)
if current_spend >= max_budget:
raise HTTPException(
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
data.get("model") if data else None
)
raise ProxyHTTPRateLimitError(
status_code=429,
detail=(
f"Session budget exceeded for session {session_id}. "
f"Current spend: ${current_spend:.4f}, "
f"max_budget_per_session: ${max_budget:.2f}."
),
model=resolved_model,
llm_provider=llm_provider,
)
return None

View file

@ -13,12 +13,14 @@ Follows the same pattern as parallel_request_limiter_v3.py.
import os
from typing import TYPE_CHECKING, Any, Optional, Union
from fastapi import HTTPException
from litellm import DualCache
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.hooks.rate_limiter_utils import (
ProxyHTTPRateLimitError,
resolve_llm_provider_for_rate_limit,
)
if TYPE_CHECKING:
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
@ -116,12 +118,17 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
current_count = await self._increment_and_get(cache_key)
if current_count > max_iterations:
raise HTTPException(
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
data.get("model") if data else None
)
raise ProxyHTTPRateLimitError(
status_code=429,
detail=(
f"Max iterations exceeded for session {session_id}. "
f"Current count: {current_count}, max_iterations: {max_iterations}."
),
model=resolved_model,
llm_provider=llm_provider,
)
verbose_proxy_logger.debug(

View file

@ -17,6 +17,10 @@ from litellm.proxy.auth.auth_utils import (
get_key_model_rpm_limit,
get_key_model_tpm_limit,
)
from litellm.proxy.hooks.rate_limiter_utils import (
ProxyHTTPRateLimitError,
resolve_llm_provider_for_rate_limit,
)
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
@ -73,7 +77,8 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
if max_parallel_requests == 0 or tpm_limit == 0 or rpm_limit == 0:
# base case
raise self.raise_rate_limit_error(
additional_details=f"{CommonProxyErrors.max_parallel_request_limit_reached.value}. Hit limit for {rate_limit_type}. Current limits: max_parallel_requests: {max_parallel_requests}, tpm_limit: {tpm_limit}, rpm_limit: {rpm_limit}"
additional_details=f"{CommonProxyErrors.max_parallel_request_limit_reached.value}. Hit limit for {rate_limit_type}. Current limits: max_parallel_requests: {max_parallel_requests}, tpm_limit: {tpm_limit}, rpm_limit: {rpm_limit}",
requested_model=data.get("model") if data else None,
)
new_val = {
"current_requests": 1,
@ -95,10 +100,16 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
values_to_update_in_cache.append((request_count_api_key, new_val))
else:
raise HTTPException(
requested_model = data.get("model") if data else None
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
requested_model
)
raise ProxyHTTPRateLimitError(
status_code=429,
detail=f"LiteLLM Rate Limit Handler for rate limit type = {rate_limit_type}. {CommonProxyErrors.max_parallel_request_limit_reached.value}. current rpm: {current['current_rpm']}, rpm limit: {rpm_limit}, current tpm: {current['current_tpm']}, tpm limit: {tpm_limit}, current max_parallel_requests: {current['current_requests']}, max_parallel_requests: {max_parallel_requests}",
headers={"retry-after": str(self.time_to_next_minute())},
model=resolved_model,
llm_provider=llm_provider,
)
await self.internal_usage_cache.async_batch_set_cache(
@ -122,18 +133,31 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
return seconds_to_next_minute
def raise_rate_limit_error(
self, additional_details: Optional[str] = None
self,
additional_details: Optional[str] = None,
requested_model: Optional[str] = None,
) -> HTTPException:
"""
Raise an HTTPException with a 429 status code and a retry-after header
Raise an HTTPException with a 429 status code and a retry-after header.
``requested_model`` is resolved via :func:`get_llm_provider` so the
raised exception carries ``llm_provider`` for downstream loggers
(Prometheus failure metric, observability callbacks). Falls back to
``llm_provider="litellm_proxy"`` when the model is missing or
unparseable — see ``resolve_llm_provider_for_rate_limit``.
"""
error_message = "Max parallel request limit reached"
if additional_details is not None:
error_message = error_message + " " + additional_details
raise HTTPException(
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
requested_model
)
raise ProxyHTTPRateLimitError(
status_code=429,
detail=f"Max parallel request limit reached {additional_details}",
detail=error_message,
headers={"retry-after": str(self.time_to_next_minute())},
model=resolved_model,
llm_provider=llm_provider,
)
async def get_all_cache_objects(
@ -225,7 +249,8 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
# if above -> raise error
if current_global_requests >= global_max_parallel_requests:
return self.raise_rate_limit_error(
additional_details=f"Hit Global Limit: Limit={global_max_parallel_requests}, current: {current_global_requests}"
additional_details=f"Hit Global Limit: Limit={global_max_parallel_requests}, current: {current_global_requests}",
requested_model=data.get("model") if data else None,
)
# if below -> increment
else:

View file

@ -23,8 +23,6 @@ from typing import (
cast,
)
from fastapi import HTTPException
from litellm import DualCache
from litellm._logging import verbose_proxy_logger
from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE
@ -34,6 +32,10 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import get_model_rate_limit_from_metadata
from litellm.proxy.hooks.rate_limiter_utils import (
ProxyHTTPRateLimitError,
resolve_llm_provider_for_rate_limit,
)
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject
from litellm.types.utils import CallTypes, ModelResponse, Usage
@ -1967,6 +1969,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
self,
response: RateLimitResponse,
descriptors: List[RateLimitDescriptor],
requested_model: Optional[str] = None,
) -> None:
"""Handle rate limit exceeded error by raising HTTPException."""
for status in response["statuses"]:
@ -1999,7 +2002,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
f"Limit resets at: {reset_time_formatted}"
)
raise HTTPException(
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
requested_model
)
raise ProxyHTTPRateLimitError(
status_code=429,
detail=detail,
headers={
@ -2007,6 +2013,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
"rate_limit_type": str(status["rate_limit_type"]),
"reset_at": reset_time_formatted,
},
model=resolved_model,
llm_provider=llm_provider,
)
async def async_pre_call_hook(
@ -2115,6 +2123,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
self._handle_rate_limit_error(
response=response,
descriptors=descriptors,
requested_model=requested_model,
)
else:
# add descriptors to request headers
@ -2188,6 +2197,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
self._handle_rate_limit_error(
response=tpm_response,
descriptors=descriptors,
requested_model=requested_model,
)
else:
self._stash_value_in_metadata_channels(

View file

@ -2,11 +2,105 @@
Shared utility functions for rate limiter hooks.
"""
from typing import Optional, Union
from typing import Any, Optional, Tuple, Union
from fastapi import HTTPException
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.exceptions import RateLimitError
from litellm.types.router import ModelGroupInfo
from litellm.types.utils import PriorityReservationDict
PROXY_LLM_PROVIDER_FALLBACK = "litellm_proxy"
def resolve_llm_provider_for_rate_limit(
model: Optional[str],
) -> Tuple[str, str]:
"""
Resolve ``(model, llm_provider)`` for a request being rejected by an
internal proxy-side rate-limit hook.
These hooks fire from ``async_pre_call_hook`` — well before
:func:`litellm.get_llm_provider` is invoked anywhere else in the request
lifecycle — so the raised 429 would otherwise have an empty
``llm_provider`` field, making the resulting Prometheus
``litellm_proxy_failed_requests_metric`` show up with
``exception_class="RateLimitError"`` and no provider attribution.
Wrapped defensively: if ``model`` is missing, malformed, or
``get_llm_provider`` raises (unknown alias, router-only model, etc.) we
fall back to ``("", "litellm_proxy")`` so we never break the request path
by piling a second exception on top of the rate-limit one we're trying to
raise.
"""
if not model:
return "", PROXY_LLM_PROVIDER_FALLBACK
try:
resolved_model, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=model,
)
return (
resolved_model or model,
custom_llm_provider or PROXY_LLM_PROVIDER_FALLBACK,
)
except Exception as e:
verbose_proxy_logger.debug(
"rate_limiter_utils.resolve_llm_provider_for_rate_limit: "
"could not resolve provider for model=%s, falling back to %s. err=%s",
model,
PROXY_LLM_PROVIDER_FALLBACK,
str(e),
)
return model, PROXY_LLM_PROVIDER_FALLBACK
class ProxyHTTPRateLimitError(HTTPException, RateLimitError): # type: ignore[misc]
"""
HTTPException raised by proxy-side rate-limit hooks that *also* exposes
``model`` and ``llm_provider`` attributes.
Why both base classes:
- The proxy server's exception handler keys off ``HTTPException`` to render
a 429 response, so we must remain an ``HTTPException``.
- Downstream loggers (Prometheus ``async_post_call_failure_hook``,
structured logging, observability callbacks) read ``exception.llm_provider``
via :meth:`litellm.integrations.prometheus.PrometheusLogger._get_exception_class_name`
and ``isinstance(exc, RateLimitError)`` for category routing. Inheriting
from :class:`litellm.exceptions.RateLimitError` keeps that wiring intact.
We intentionally do not call ``RateLimitError.__init__`` (which constructs
an httpx.Response) — it isn't needed here and just adds failure surface.
Attribute parity is what downstream consumers rely on.
"""
def __init__(
self,
status_code: int,
detail: Any = None,
headers: Optional[dict] = None,
*,
model: str = "",
llm_provider: str = PROXY_LLM_PROVIDER_FALLBACK,
) -> None:
HTTPException.__init__(
self, status_code=status_code, detail=detail, headers=headers
)
self.status_code = status_code
self.model = model or ""
self.llm_provider = llm_provider or PROXY_LLM_PROVIDER_FALLBACK
# `message` is what RateLimitError.__str__ would print and what some
# observability callbacks log. Keep it human-readable.
self.message = detail if isinstance(detail, str) else str(detail)
# `RateLimitError.__str__` (resolved via MRO since Starlette's
# HTTPException doesn't define `__str__`) unconditionally reads
# these attributes. Set them so `str(exc)` doesn't raise
# AttributeError from logging/traceback paths.
self.num_retries: Optional[int] = None
self.max_retries: Optional[int] = None
def convert_priority_to_percent(
value: Union[float, PriorityReservationDict], model_info: Optional[ModelGroupInfo]

View file

@ -0,0 +1,206 @@
"""
Sensitive Data Routing Hook for LiteLLM Proxy.
When a guardrail detects sensitive data and is configured with on_sensitive_data='route',
this hook manages:
1. Storing the routing decision (session_id -> model) in cache
2. Checking incoming requests for existing routing overrides
3. Applying sticky routing so all subsequent requests in a session go to the same model
Works across multiple proxy instances via DualCache (in-memory + Redis).
"""
import os
from typing import TYPE_CHECKING, Any, Optional, Union
from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_guardrail import get_session_id_from_request_data
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
if TYPE_CHECKING:
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
InternalUsageCache = _InternalUsageCache
else:
InternalUsageCache = Any
SENSITIVE_ROUTING_CACHE_PREFIX = "sensitive_route"
DEFAULT_SENSITIVE_ROUTING_TTL = 3600
class _PROXY_SensitiveDataRoutingHandler(CustomLogger):
"""
Pre-call hook that checks for existing sensitive data routing overrides
and applies them to incoming requests.
This hook runs early in the pre-call chain and modifies the request's
model field if a routing override exists for the session.
"""
def __init__(self, internal_usage_cache: InternalUsageCache):
self.internal_usage_cache = internal_usage_cache
self.ttl = int(
os.getenv(
"LITELLM_SENSITIVE_ROUTING_TTL",
str(DEFAULT_SENSITIVE_ROUTING_TTL),
)
)
def _make_cache_key(self, session_id: str, tenant: str) -> str:
return f"{{{SENSITIVE_ROUTING_CACHE_PREFIX}:{tenant}:{session_id}}}:model"
@staticmethod
def _resolve_tenant(user_api_key_dict: Optional[UserAPIKeyAuth]) -> str:
"""
Identify the authenticated principal the routing override belongs to.
API-key auth is scoped by the hashed key. JWT (and other keyless) auth
has no api_key, so fall back to a stable identity claim. Without this,
every keyless caller would share the ``default`` namespace and could read
or overwrite another principal's session routing.
"""
if user_api_key_dict is None:
return "default"
if user_api_key_dict.api_key:
return user_api_key_dict.api_key
principal = [
f"{label}:{value}"
for label, value in (
("user", user_api_key_dict.user_id),
("team", user_api_key_dict.team_id),
("org", user_api_key_dict.org_id),
)
if value
]
return "|".join(principal) if principal else "default"
async def _get_routed_model(
self, session_id: str, user_api_key_dict: Optional[UserAPIKeyAuth]
) -> Optional[str]:
"""Get the model this session should be routed to, if any."""
cache_key = self._make_cache_key(
session_id, self._resolve_tenant(user_api_key_dict)
)
if self.internal_usage_cache.dual_cache.redis_cache is not None:
try:
result = await self.internal_usage_cache.dual_cache.redis_cache.async_get_cache(
key=cache_key
)
if result is not None:
routed_model = str(result)
remaining_ttl = await self.internal_usage_cache.dual_cache.redis_cache.async_get_ttl(
key=cache_key
)
await self.internal_usage_cache.async_set_cache(
key=cache_key,
value=routed_model,
ttl=remaining_ttl if remaining_ttl is not None else self.ttl,
litellm_parent_otel_span=None,
local_only=True,
)
return routed_model
except Exception as e:
verbose_proxy_logger.warning(
"SensitiveDataRoutingHandler: Redis GET failed, falling back to in-memory: %s",
str(e),
)
result = await self.internal_usage_cache.async_get_cache(
key=cache_key,
litellm_parent_otel_span=None,
local_only=True,
)
if result is not None:
return str(result)
return None
async def set_session_routing(
self,
session_id: str,
model: str,
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
guardrail_name: Optional[str] = None,
) -> None:
"""
Store a routing override for a session.
Called by guardrails when they detect sensitive data and want to
route the session to a specific model. The override is scoped to the
requesting principal so sessions from different tenants cannot collide.
"""
cache_key = self._make_cache_key(
session_id, self._resolve_tenant(user_api_key_dict)
)
verbose_proxy_logger.info(
"SensitiveDataRoutingHandler: Setting session routing session_id=%s model=%s guardrail=%s ttl=%s",
session_id,
model,
guardrail_name,
self.ttl,
)
if self.internal_usage_cache.dual_cache.redis_cache is not None:
try:
await self.internal_usage_cache.dual_cache.redis_cache.async_set_cache(
key=cache_key,
value=model,
ttl=self.ttl,
)
except Exception as e:
verbose_proxy_logger.warning(
"SensitiveDataRoutingHandler: Redis SET failed, falling back to in-memory: %s",
str(e),
)
await self.internal_usage_cache.async_set_cache(
key=cache_key,
value=model,
ttl=self.ttl,
litellm_parent_otel_span=None,
local_only=True,
)
async def async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
cache: DualCache,
data: dict,
call_type: str,
) -> Optional[Union[Exception, str, dict]]:
"""
Before each LLM call, check if this session has a routing override.
If so, modify the request's model field.
"""
session_id = get_session_id_from_request_data(data)
if session_id is None:
return None
routed_model = await self._get_routed_model(session_id, user_api_key_dict)
if routed_model is None:
return None
original_model = data.get("model")
if original_model == routed_model:
return None
verbose_proxy_logger.info(
"SensitiveDataRoutingHandler: Applying session routing override "
"session_id=%s original_model=%s routed_model=%s",
session_id,
original_model,
routed_model,
)
data["model"] = routed_model
metadata = data.get("metadata") or {}
metadata["sensitive_data_routing_applied"] = True
metadata["sensitive_data_routing_original_model"] = original_model
data["metadata"] = metadata
return data

View file

@ -327,6 +327,7 @@ model LiteLLM_MCPServerTable {
authorization_url String?
token_url String?
registration_url String?
oauth2_flow String?
allow_all_keys Boolean @default(false)
available_on_public_internet Boolean @default(true)
delegate_auth_to_upstream Boolean @default(false)

View file

@ -89,11 +89,14 @@ from litellm._logging import _redact_string, verbose_proxy_logger
from litellm._service_logger import ServiceLogging, ServiceTypes
from litellm.caching.caching import DualCache, RedisCache
from litellm.caching.dual_cache import LimitedSizeOrderedDict
from litellm.exceptions import RejectedRequestError
from litellm.exceptions import RejectedRequestError, SensitiveDataRouteException
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
ModifyResponseException,
)
from litellm.proxy.hooks.sensitive_data_routing import (
_PROXY_SensitiveDataRoutingHandler,
)
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
from litellm.integrations.SlackAlerting.utils import _add_langfuse_trace_id_to_alert
@ -1151,6 +1154,9 @@ class ProxyLogging:
response=response, data=data, call_type=call_type
)
except SensitiveDataRouteException:
status = "intervened"
raise
except Exception as e:
status = "error"
error_type = type(e).__name__
@ -1460,47 +1466,57 @@ class ProxyLogging:
self._process_guardrail_metadata(data)
return data
deferred_route_exc: Optional[SensitiveDataRouteException] = None
for _callback in caps.resolved_callbacks:
start_time = time.time()
if isinstance(_callback, CustomGuardrail) and data is not None:
# Skip guardrails managed by a pipeline
if (
_callback.guardrail_name
and _callback.guardrail_name in pipeline_managed
):
continue
try:
if isinstance(_callback, CustomGuardrail) and data is not None:
# Skip guardrails managed by a pipeline
if (
_callback.guardrail_name
and _callback.guardrail_name in pipeline_managed
):
continue
result = await self._process_guardrail_callback(
callback=_callback,
data=data, # type: ignore
user_api_key_dict=user_api_key_dict,
call_type=call_type,
event_type=GuardrailEventHooks.pre_call,
)
if result is None:
continue
data = result
elif (
_callback is not None
and isinstance(_callback, CustomLogger)
and "async_pre_call_hook" in vars(_callback.__class__)
and _callback.__class__.async_pre_call_hook
!= CustomLogger.async_pre_call_hook
):
if call_type == "call_mcp_tool" and user_api_key_dict is None:
continue
response = await _callback.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=self.call_details["user_api_key_cache"],
data=data, # type: ignore
call_type=call_type, # type: ignore
)
if response is not None:
data = await self.process_pre_call_hook_response(
response=response, data=data, call_type=call_type
result = await self._process_guardrail_callback(
callback=_callback,
data=data, # type: ignore
user_api_key_dict=user_api_key_dict,
call_type=call_type,
event_type=GuardrailEventHooks.pre_call,
)
if result is None:
continue
data = result
elif (
_callback is not None
and isinstance(_callback, CustomLogger)
and "async_pre_call_hook" in vars(_callback.__class__)
and _callback.__class__.async_pre_call_hook
!= CustomLogger.async_pre_call_hook
):
if call_type == "call_mcp_tool" and user_api_key_dict is None:
continue
response = await _callback.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=self.call_details["user_api_key_cache"],
data=data, # type: ignore
call_type=call_type, # type: ignore
)
if response is not None:
data = await self.process_pre_call_hook_response(
response=response, data=data, call_type=call_type
)
except SensitiveDataRouteException as e:
# Defer the reroute until remaining guardrails have run so later
# security checks are not skipped; the first reroute wins and a
# later guardrail that blocks still propagates. Fall through to the
# service-span recording below so the triggering guardrail is still
# timed like every other callback.
if deferred_route_exc is None:
deferred_route_exc = e
end_time = time.time()
duration = end_time - start_time
@ -1516,13 +1532,76 @@ class ProxyLogging:
end_time=end_time,
)
if deferred_route_exc is not None and data is not None:
data = await self._handle_sensitive_data_route_exception(
deferred_route_exc, data, user_api_key_dict
)
if data is not None:
self._process_guardrail_metadata(data)
return data
except SensitiveDataRouteException as e:
data = await self._handle_sensitive_data_route_exception(
e, data, user_api_key_dict
)
if data is not None:
self._process_guardrail_metadata(data)
return data
except Exception as e:
raise e
async def _handle_sensitive_data_route_exception(
self,
exc: SensitiveDataRouteException,
data: Optional[dict],
user_api_key_dict: Optional[UserAPIKeyAuth],
) -> Optional[dict]:
"""
Handle SensitiveDataRouteException by rerouting the current request to
the target model and, when sticky_session_routing is enabled, persisting
the session override so subsequent requests reuse the same model.
"""
if data is None:
return None
verbose_proxy_logger.info(
"SensitiveDataRouteException caught: session_id=%s route_to_model=%s guardrail=%s sticky=%s",
exc.session_id,
exc.route_to_model,
exc.guardrail_name,
exc.sticky_session_routing,
)
if exc.sticky_session_routing:
sensitive_routing_hook = self.get_proxy_hook("sensitive_data_routing")
if isinstance(sensitive_routing_hook, _PROXY_SensitiveDataRoutingHandler):
await sensitive_routing_hook.set_session_routing(
session_id=exc.session_id,
model=exc.route_to_model,
user_api_key_dict=user_api_key_dict,
guardrail_name=exc.guardrail_name,
)
else:
verbose_proxy_logger.warning(
"SensitiveDataRouteException requested sticky routing for session_id=%s "
"but the 'sensitive_data_routing' hook is not registered. Only this request "
"will be rerouted; subsequent requests will not be sticky.",
exc.session_id,
)
original_model = data.get("model")
data["model"] = exc.route_to_model
metadata = data.get("metadata") or {}
metadata["sensitive_data_routing_applied"] = True
metadata["sensitive_data_routing_original_model"] = original_model
metadata["sensitive_data_routing_guardrail"] = exc.guardrail_name
metadata["sensitive_data_routing_detection_info"] = exc.detection_info
data["metadata"] = metadata
return data
@staticmethod
async def _run_guardrail_task_with_enrichment(
callback: Any, coro: Awaitable[Any]

View file

@ -1006,6 +1006,20 @@ class ResponseAPILoggingUtils:
)
response_api_usage: ResponseAPIUsage
if isinstance(usage_input, dict):
usage_input = dict(usage_input) # shallow copy; avoid mutating caller
# Realtime *_token_details → *_tokens_details when unset.
if (
usage_input.get("input_tokens_details") is None
and "input_token_details" in usage_input
):
usage_input["input_tokens_details"] = usage_input["input_token_details"]
if (
usage_input.get("output_tokens_details") is None
and "output_token_details" in usage_input
):
usage_input["output_tokens_details"] = usage_input[
"output_token_details"
]
total_tokens = usage_input.get("total_tokens")
if total_tokens is None:
input_tokens = usage_input.get("input_tokens")
@ -1050,6 +1064,7 @@ class ResponseAPILoggingUtils:
),
image_tokens=getattr(output_tokens_details, "image_tokens", None),
text_tokens=getattr(output_tokens_details, "text_tokens", None),
audio_tokens=getattr(output_tokens_details, "audio_tokens", None),
)
chat_usage = Usage(

View file

@ -2,7 +2,7 @@ from datetime import datetime
from enum import Enum
from typing import Any, Dict, List, Literal, Optional, Union
from pydantic import BaseModel, ConfigDict, Field, field_validator
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from typing_extensions import Required, TypedDict
from litellm.types.proxy.guardrails.guardrail_hooks.akto import (
@ -771,6 +771,58 @@ class BaseLitellmParams(
),
)
on_sensitive_data: Optional[Literal["block", "route"]] = Field(
default=None,
description=(
"Action to take when sensitive data is detected. "
"'block' raises an exception (default behavior). "
"'route' reroutes the request to the model specified in sensitive_data_route_to_model."
),
)
sensitive_data_route_to_model: Optional[str] = Field(
default=None,
description=(
"Model to route requests to when sensitive data is detected and on_sensitive_data='route'. "
"This is typically an on-premise model for data privacy. "
"The routing decision persists for the entire session."
),
)
sticky_session_routing: Optional[bool] = Field(
default=True,
description=(
"When True (default), after sensitive data is detected and routed, all subsequent "
"requests in the same session will continue routing to the same model."
),
)
@field_validator(
"mode",
"default_action",
"on_disallowed_action",
"unreachable_fallback",
"on_sensitive_data",
mode="before",
check_fields=False,
)
@classmethod
def normalize_lowercase(cls, v):
"""Normalize string and list fields to lowercase for ALL guardrail types."""
if isinstance(v, str):
return v.lower()
if isinstance(v, list):
return [x.lower() if isinstance(x, str) else x for x in v]
return v
@model_validator(mode="after")
def validate_sensitive_data_routing(self) -> "BaseLitellmParams":
if self.on_sensitive_data == "route" and not self.sensitive_data_route_to_model:
raise ValueError(
"sensitive_data_route_to_model must be set when on_sensitive_data='route'"
)
return self
model_config = ConfigDict(extra="allow", protected_namespaces=())
@ -811,23 +863,6 @@ class LitellmParams(
description="When to apply the guardrail (pre_call, post_call, during_call, logging_only)"
)
@field_validator(
"mode",
"default_action",
"on_disallowed_action",
"unreachable_fallback",
mode="before",
check_fields=False,
)
@classmethod
def normalize_lowercase(cls, v):
"""Normalize string and list fields to lowercase for ALL guardrail types."""
if isinstance(v, str):
return v.lower()
if isinstance(v, list):
return [x.lower() if isinstance(x, str) else x for x in v]
return v
@field_validator("timeout", mode="before", check_fields=False)
@classmethod
def coerce_timeout(cls, v):

View file

@ -1319,6 +1319,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1350,6 +1351,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1381,6 +1383,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1412,6 +1415,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1443,6 +1447,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -2194,6 +2199,7 @@
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -33962,6 +33968,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -33990,6 +33997,7 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,

View file

@ -327,6 +327,7 @@ model LiteLLM_MCPServerTable {
authorization_url String?
token_url String?
registration_url String?
oauth2_flow String?
allow_all_keys Boolean @default(false)
available_on_public_internet Boolean @default(true)
delegate_auth_to_upstream Boolean @default(false)

View file

@ -1249,3 +1249,47 @@ class TestCustomGuardrailSpendLogMatchRedaction:
slg = request_data["metadata"]["standard_logging_guardrail_information"][0]
assert slg["guardrail_response"]["filters"][0]["regex"] == "[REDACTED]"
assert raw["filters"][0]["regex"] == r"\d{3}-\d{2}-\d{4}"
class TestGuardrailInterventionClassification:
"""A routing decision is a deliberate guardrail intervention, not a failure."""
def test_sensitive_data_route_exception_is_intervention(self):
from litellm.exceptions import SensitiveDataRouteException
exc = SensitiveDataRouteException(
route_to_model="on-prem-model",
session_id="sess-1",
guardrail_name="pii-rail",
)
assert CustomGuardrail._is_guardrail_intervention(exc) is True
@pytest.mark.asyncio
async def test_routing_logged_as_intervened_not_failed(self):
from litellm.exceptions import SensitiveDataRouteException
from litellm.integrations.custom_guardrail import log_guardrail_information
from litellm.types.guardrails import GuardrailEventHooks
class RoutingGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="pii-rail",
event_hook=GuardrailEventHooks.pre_call,
)
@log_guardrail_information
async def async_pre_call_hook(self, data, **kwargs):
raise SensitiveDataRouteException(
route_to_model="on-prem-model",
session_id="sess-1",
guardrail_name=self.guardrail_name,
)
guardrail = RoutingGuardrail()
request_data: dict = {"metadata": {}}
with pytest.raises(SensitiveDataRouteException):
await guardrail.async_pre_call_hook(data=request_data)
slg = request_data["metadata"]["standard_logging_guardrail_information"][0]
assert slg["guardrail_status"] == "guardrail_intervened"

View file

@ -5090,3 +5090,101 @@ def test_map_tool_helper_collision_prefers_definitions_over_components_schemas()
# Cross-namespace ref *also* resolves to the `definitions` body because
# ``unpack_defs`` keys by last path segment -- documented limitation.
assert transformed["input_schema"]["properties"]["from_components"] == expected
def test_namespace_tool_flat_nested_tools_are_extracted():
"""Codex sends nested tools in flat format {type, name, description, parameters} with no 'function' wrapper.
These must be normalized and mapped without raising KeyError: 'function'."""
config = AnthropicConfig()
tools = [
{
"type": "namespace",
"name": "multi_agent_v1",
"tools": [
{
"type": "function",
"name": "close_agent",
"description": "Close an agent.",
"strict": False,
"parameters": {
"type": "object",
"properties": {"target": {"type": "string"}},
"required": ["target"],
"additionalProperties": False,
},
},
],
}
]
anthropic_tools, _ = config._map_tools(tools)
assert len(anthropic_tools) == 1
assert anthropic_tools[0]["name"] == "close_agent"
def test_namespace_tool_nested_tools_are_extracted():
"""Codex sends type='namespace' wrapping nested tools in Anthropic format.
The namespace container must be dropped and its nested tools extracted individually.
"""
config = AnthropicConfig()
tools = [
{
"type": "namespace",
"name": "multi_agent_v1",
"description": "Tools for spawning and managing sub-agents.",
"tools": [
{
"name": "close_agent",
"type": "custom",
"description": "Close an agent.",
"input_schema": {
"type": "object",
"properties": {"target": {"type": "string"}},
"required": ["target"],
},
},
{
"name": "resume_agent",
"type": "custom",
"description": "Resume a closed agent.",
"input_schema": {
"type": "object",
"properties": {"id": {"type": "string"}},
"required": ["id"],
},
},
],
},
{
"type": "function",
"function": {
"name": "exec_command",
"description": "Run a command.",
"parameters": {
"type": "object",
"properties": {"cmd": {"type": "string"}},
"required": ["cmd"],
},
},
},
]
anthropic_tools, mcp_servers = config._map_tools(tools)
names = [t["name"] for t in anthropic_tools]
assert "close_agent" in names
assert "resume_agent" in names
assert "exec_command" in names
assert "multi_agent_v1" not in names
assert len(anthropic_tools) == 3
assert mcp_servers == []
def test_client_metadata_stripped_from_anthropic_request():
"""client_metadata passed by codex must not reach the Anthropic (or Vertex Anthropic) payload."""
config = AnthropicConfig()
result = config.transform_request(
model="claude-3-5-haiku-20241022",
messages=[{"role": "user", "content": "hello"}],
optional_params={"max_tokens": 10, "client_metadata": {"originator": "codex"}},
litellm_params={},
headers={},
)
assert "client_metadata" not in result

View file

@ -2,6 +2,7 @@
import pytest
import litellm
from litellm.llms.anthropic.common_utils import AnthropicError
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
@ -11,6 +12,22 @@ from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_tran
)
@pytest.fixture
def local_model_cost_map(monkeypatch):
"""Force the bundled backup cost map so Opus 4.8 adaptive detection (driven
by the ``supports_adaptive_thinking`` flag) doesn't depend on the
network-fetched ``main`` copy, which lacks the flag until this branch merges."""
original = litellm.model_cost
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
litellm.model_cost = litellm.get_model_cost_map(url="")
litellm.get_model_info.cache_clear()
try:
yield
finally:
litellm.model_cost = original
litellm.get_model_info.cache_clear()
@pytest.mark.parametrize(
"reasoning_effort,expected_effort",
[
@ -298,6 +315,40 @@ def test_legacy_thinking_high_budget_keeps_xhigh_when_supported():
assert result.get("output_config") == {"effort": "xhigh"}
@pytest.mark.parametrize(
"model",
[
"claude-opus-4-8",
"bedrock/us.anthropic.claude-opus-4-8",
"bedrock/invoke/us.anthropic.claude-opus-4-8",
],
)
def test_legacy_thinking_translates_to_adaptive_for_opus_48(
model, local_model_cost_map
):
"""Regression for issue #29188: Opus 4.8 requires adaptive thinking, but the
legacy ``thinking.type='enabled'`` shape was passed through unchanged for
Bedrock 4.8 (its cost-map entry lacked ``supports_adaptive_thinking`` and the
lookup didn't strip the provider prefix), so Bedrock rejected the request. The
reporter's reproducer used ``budget_tokens=24000``, the ``xhigh`` bucket."""
config = AnthropicMessagesConfig()
optional_params = {
"max_tokens": 100,
"thinking": {"type": "enabled", "budget_tokens": 24000},
}
result = config.transform_anthropic_messages_request(
model=model,
messages=[{"role": "user", "content": "ping"}],
anthropic_messages_optional_request_params=optional_params,
litellm_params={},
headers={},
)
assert result.get("thinking") == {"type": "adaptive"}
assert result.get("output_config") == {"effort": "xhigh"}
@pytest.mark.parametrize(
"budget_tokens,expected_effort",
[

View file

@ -14,6 +14,8 @@ import os
import sys
from unittest.mock import patch
import pytest
sys.path.insert(
0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../.."))
)
@ -1378,3 +1380,73 @@ class TestAnthropicThinkingSignatureSelfHeal:
config.transform_anthropic_messages_request_on_http_error(err, data)
assert "thinking" not in data
assert data["messages"] == []
@pytest.fixture
def local_model_cost_map(monkeypatch):
"""Force the bundled backup cost map so detection doesn't depend on the
network-fetched ``main`` copy (which lacks this branch's flags until merge)."""
import litellm
original = litellm.model_cost
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
litellm.model_cost = litellm.get_model_cost_map(url="")
litellm.get_model_info.cache_clear()
try:
yield
finally:
litellm.model_cost = original
litellm.get_model_info.cache_clear()
class TestClaudeOpus48AdaptiveThinking:
"""Opus 4.8 requires adaptive thinking (``thinking.type='adaptive'`` +
``output_config.effort``). Detection is driven by the
``supports_adaptive_thinking`` cost-map flag, resolved through provider
prefixes. Before the fix the Bedrock entries lacked the flag and the lookup
didn't strip the ``us.anthropic.``/``invoke/`` prefixes, so a
``bedrock/us.anthropic.claude-opus-4-8`` call sent the legacy
``thinking.type='enabled'`` shape and Bedrock rejected it (issue #29188)."""
@pytest.mark.parametrize(
"model",
[
"claude-opus-4-8",
"anthropic/claude-opus-4-8",
"anthropic.claude-opus-4-8",
"bedrock/us.anthropic.claude-opus-4-8",
"bedrock/invoke/us.anthropic.claude-opus-4-8",
"bedrock/eu.anthropic.claude-opus-4-8",
"vertex_ai/claude-opus-4-8",
"azure_ai/claude-opus-4-8",
],
)
def test_adaptive_thinking_detected_for_opus_4_8(self, local_model_cost_map, model):
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
assert AnthropicModelInfo._is_adaptive_thinking_model(model) is True
def test_resolver_reads_flag_through_bedrock_invoke_prefix(
self, local_model_cost_map
):
"""The resolver fix: ``bedrock/invoke/...`` resolves to the flagged
Bedrock entry. Pure ``_supports_factory`` without prefix-stripping
returns False here, which is why the data-only fix alone was not enough."""
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
assert (
AnthropicModelInfo._supports_model_capability(
"bedrock/invoke/us.anthropic.claude-opus-4-8",
"supports_adaptive_thinking",
)
is True
)
@pytest.mark.parametrize(
"model",
["claude-opus-4-5", "claude-3-7-sonnet", "claude-3-5-haiku-20241022"],
)
def test_non_adaptive_models_not_detected(self, local_model_cost_map, model):
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
assert AnthropicModelInfo._is_adaptive_thinking_model(model) is False

View file

@ -564,6 +564,46 @@ def test_sync_delete_responses_omits_body_for_azure():
)
def _content_type(headers: dict) -> str:
for key, value in headers.items():
if key.lower() == "content-type":
return value
return ""
def test_async_delete_responses_sets_json_content_type():
"""OpenAI rejects a responses DELETE with no Content-Type by treating it as
application/octet-stream. The handler must declare application/json."""
captured: dict = {}
fake_async_delete, _ = _build_delete_response_mock(captured)
async def run():
with patch.object(AsyncHTTPHandler, "delete", new=fake_async_delete):
await litellm.adelete_responses(
response_id="resp_xyz",
custom_llm_provider="openai",
api_key="test-key",
)
asyncio.run(run())
assert _content_type(captured["headers"]) == "application/json"
def test_sync_delete_responses_sets_json_content_type():
captured: dict = {}
_, fake_sync_delete = _build_delete_response_mock(captured)
with patch.object(HTTPHandler, "delete", new=fake_sync_delete):
litellm.delete_responses(
response_id="resp_xyz",
custom_llm_provider="openai",
api_key="test-key",
)
assert _content_type(captured["headers"]) == "application/json"
# ---------------------------------------------------------------------------
# Parity tests: request-body is serialized once and reused for the wire.
# (_async_post_anthropic_messages_with_http_error_retry)

View file

@ -14,7 +14,6 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
# ---------------------------------------------------------------------------
# Helper: build a minimal mock agent
# ---------------------------------------------------------------------------
@ -307,6 +306,97 @@ async def test_convention_unrelated_prefix_not_forwarded():
assert headers is None
# ---------------------------------------------------------------------------
# Databricks App OAuth M2M injection
# ---------------------------------------------------------------------------
def _mock_databricks_token_client(access_token="dbx-oauth-token"):
response = MagicMock()
response.raise_for_status = MagicMock()
response.json = MagicMock(
return_value={"access_token": access_token, "expires_in": 3600}
)
client = MagicMock()
client.post = AsyncMock(return_value=response)
return client
@pytest.mark.asyncio
async def test_databricks_oauth_header_injected():
"""A databricks_oauth block mints an outbound Bearer Authorization header."""
from litellm.proxy.agent_endpoints import databricks_oauth
databricks_oauth.databricks_app_oauth_token_cache.flush_cache()
mock_agent = _make_mock_agent()
mock_agent.litellm_params = {
"databricks_oauth": {
"client_id": "cid",
"client_secret": "secret",
"workspace_url": "https://dbc.cloud.databricks.com",
}
}
mock_request = _make_mock_request()
with patch(
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
return_value=_mock_databricks_token_client("minted-token"),
):
mock_asend = await _invoke(mock_agent, mock_request, None)
headers = mock_asend.call_args.kwargs.get("agent_extra_headers")
assert headers is not None
assert headers.get("Authorization") == "Bearer minted-token"
@pytest.mark.asyncio
async def test_databricks_oauth_overrides_static_authorization():
"""The minted OAuth token wins over a statically configured Authorization."""
from litellm.proxy.agent_endpoints import databricks_oauth
databricks_oauth.databricks_app_oauth_token_cache.flush_cache()
mock_agent = _make_mock_agent(static_headers={"Authorization": "Bearer static-pat"})
mock_agent.litellm_params = {
"databricks_oauth": {
"client_id": "cid",
"client_secret": "secret",
"workspace_url": "https://dbc.cloud.databricks.com",
}
}
mock_request = _make_mock_request()
with patch(
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
return_value=_mock_databricks_token_client("oauth-wins"),
):
mock_asend = await _invoke(mock_agent, mock_request, None)
headers = mock_asend.call_args.kwargs.get("agent_extra_headers")
assert headers is not None
assert headers.get("Authorization") == "Bearer oauth-wins"
@pytest.mark.asyncio
async def test_non_databricks_agent_skips_oauth_resolution():
"""Agents without a databricks_oauth block never enter the OAuth path."""
mock_agent = _make_mock_agent(static_headers={"x-custom": "v"})
mock_agent.litellm_params = {"require_trace_id_on_calls_to_agent": False}
mock_request = _make_mock_request()
with patch(
"litellm.proxy.agent_endpoints.a2a_endpoints.resolve_databricks_app_auth_header",
new_callable=AsyncMock,
) as mock_resolve:
mock_asend = await _invoke(mock_agent, mock_request, None)
mock_resolve.assert_not_called()
headers = mock_asend.call_args.kwargs.get("agent_extra_headers")
assert headers == {"x-custom": "v"}
assert "Authorization" not in headers
# ---------------------------------------------------------------------------
# Direct unit test for the merge utility
# ---------------------------------------------------------------------------

View file

@ -0,0 +1,496 @@
"""
Unit tests for Databricks App OAuth M2M support for A2A agents.
Covers config parsing (including os.environ/ resolution and validation),
workspace token-URL construction, client_credentials token fetching, caching
with expiry buffering, and the public ``resolve_databricks_app_auth_header``
helper.
"""
import base64
from unittest.mock import MagicMock, create_autospec, patch
import httpx
import pytest
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy.agent_endpoints.databricks_oauth import (
DatabricksAppOAuthConfig,
DatabricksAppOAuthTokenCache,
parse_databricks_oauth_config,
resolve_databricks_app_auth_header,
)
def _expected_basic_auth(client_id: str, client_secret: str) -> str:
token = base64.b64encode(f"{client_id}:{client_secret}".encode()).decode()
return f"Basic {token}"
def _mock_http_handler(access_token="tok-abc", expires_in=3600, post_error=None):
"""Return a mock that mirrors litellm's ``AsyncHTTPHandler`` contract.
Two properties of the real handler matter for these tests and were the
source of a runtime bug the original suite missed:
1. ``post`` does not accept an ``auth`` kwarg. ``create_autospec`` enforces
the real signature, so reintroducing HTTP Basic via ``auth=`` fails with
``TypeError`` instead of silently passing.
2. ``post`` calls ``raise_for_status`` internally and raises
``httpx.HTTPStatusError`` itself on non-2xx; callers never inspect the
returned response's status. Error-path tests therefore raise from
``post`` rather than from ``response.raise_for_status``.
"""
handler = create_autospec(AsyncHTTPHandler, instance=True)
if post_error is not None:
handler.post.side_effect = post_error
else:
response = MagicMock()
response.json = MagicMock(
return_value={"access_token": access_token, "expires_in": expires_in}
)
handler.post.return_value = response
return handler
# ---------------------------------------------------------------------------
# Config parsing
# ---------------------------------------------------------------------------
def test_parse_returns_none_without_block():
assert parse_databricks_oauth_config(None) is None
assert parse_databricks_oauth_config({}) is None
assert parse_databricks_oauth_config({"other": "value"}) is None
def test_parse_builds_config_and_token_url():
config = parse_databricks_oauth_config(
{
"databricks_oauth": {
"client_id": "cid",
"client_secret": "secret",
"workspace_url": "https://dbc-abc.cloud.databricks.com",
}
}
)
assert config == DatabricksAppOAuthConfig(
client_id="cid",
client_secret="secret",
token_url="https://dbc-abc.cloud.databricks.com/oidc/v1/token",
scope="all-apis",
)
def test_parse_strips_serving_endpoints_and_trailing_slash():
config = parse_databricks_oauth_config(
{
"databricks_oauth": {
"client_id": "cid",
"client_secret": "secret",
"workspace_url": "https://dbc-abc.cloud.databricks.com/serving-endpoints/",
}
}
)
assert config is not None
assert config.token_url == "https://dbc-abc.cloud.databricks.com/oidc/v1/token"
def test_parse_custom_scope():
config = parse_databricks_oauth_config(
{
"databricks_oauth": {
"client_id": "cid",
"client_secret": "secret",
"workspace_url": "https://dbc-abc.cloud.databricks.com",
"scope": "custom-scope",
}
}
)
assert config is not None
assert config.scope == "custom-scope"
@pytest.mark.parametrize(
"missing_field", ["client_id", "client_secret", "workspace_url"]
)
def test_parse_raises_on_missing_field(missing_field):
block = {
"client_id": "cid",
"client_secret": "secret",
"workspace_url": "https://dbc-abc.cloud.databricks.com",
}
block.pop(missing_field)
with pytest.raises(ValueError, match=missing_field):
parse_databricks_oauth_config({"databricks_oauth": block})
def test_parse_raises_on_non_mapping_block():
with pytest.raises(ValueError, match="mapping"):
parse_databricks_oauth_config({"databricks_oauth": "not-a-dict"})
def test_parse_resolves_os_environ_references(monkeypatch):
monkeypatch.setenv("MY_DBX_CLIENT_ID", "env-cid")
monkeypatch.setenv("MY_DBX_SECRET", "env-secret")
config = parse_databricks_oauth_config(
{
"databricks_oauth": {
"client_id": "os.environ/MY_DBX_CLIENT_ID",
"client_secret": "os.environ/MY_DBX_SECRET",
"workspace_url": "https://dbc-abc.cloud.databricks.com",
}
}
)
assert config is not None
assert config.client_id == "env-cid"
assert config.client_secret == "env-secret"
# ---------------------------------------------------------------------------
# Token fetching + caching
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_fetch_token_posts_client_credentials_with_basic_auth():
cache = DatabricksAppOAuthTokenCache()
config = DatabricksAppOAuthConfig(
client_id="cid",
client_secret="secret",
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
scope="all-apis",
)
client = _mock_http_handler(access_token="tok-1")
with patch(
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
return_value=client,
):
token = await cache.async_get_token(config)
assert token == "tok-1"
client.post.assert_awaited_once()
call = client.post.call_args
assert call.args[0] == config.token_url
assert call.kwargs["data"] == {
"grant_type": "client_credentials",
"scope": "all-apis",
}
# Databricks authenticates the client with HTTP Basic; it must be sent as a
# header because litellm's AsyncHTTPHandler.post has no ``auth`` parameter.
assert call.kwargs["headers"]["Authorization"] == _expected_basic_auth(
"cid", "secret"
)
assert "auth" not in call.kwargs
@pytest.mark.asyncio
async def test_token_is_cached_across_calls():
cache = DatabricksAppOAuthTokenCache()
config = DatabricksAppOAuthConfig(
client_id="cid",
client_secret="secret",
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
scope="all-apis",
)
client = _mock_http_handler(access_token="tok-cached")
with patch(
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
return_value=client,
):
first = await cache.async_get_token(config)
second = await cache.async_get_token(config)
assert first == second == "tok-cached"
client.post.assert_awaited_once()
@pytest.mark.asyncio
async def test_distinct_clients_do_not_share_token():
cache = DatabricksAppOAuthTokenCache()
config_a = DatabricksAppOAuthConfig(
client_id="cid-a",
client_secret="secret",
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
scope="all-apis",
)
config_b = DatabricksAppOAuthConfig(
client_id="cid-b",
client_secret="secret",
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
scope="all-apis",
)
clients = [_mock_http_handler("tok-a"), _mock_http_handler("tok-b")]
def _next_client(*args, **kwargs):
return clients.pop(0)
with patch(
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
side_effect=_next_client,
):
token_a = await cache.async_get_token(config_a)
token_b = await cache.async_get_token(config_b)
assert token_a == "tok-a"
assert token_b == "tok-b"
@pytest.mark.asyncio
async def test_ttl_applies_expiry_buffer():
cache = DatabricksAppOAuthTokenCache()
config = DatabricksAppOAuthConfig(
client_id="cid",
client_secret="secret",
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
scope="all-apis",
)
client = _mock_http_handler(access_token="tok", expires_in=600)
captured = {}
real_set = cache.set_cache
def _spy_set(key, value, **kwargs):
captured["ttl"] = kwargs.get("ttl")
return real_set(key, value, **kwargs)
with (
patch(
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
return_value=client,
),
patch.object(cache, "set_cache", side_effect=_spy_set),
):
await cache.async_get_token(config)
assert captured["ttl"] == 600 - 60
@pytest.mark.asyncio
async def test_missing_access_token_raises():
cache = DatabricksAppOAuthTokenCache()
config = DatabricksAppOAuthConfig(
client_id="cid",
client_secret="secret",
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
scope="all-apis",
)
client = _mock_http_handler()
client.post.return_value.json.return_value = {"not_a_token": "x"}
with patch(
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
return_value=client,
):
with pytest.raises(ValueError, match="access_token"):
await cache.async_get_token(config)
def _config():
return DatabricksAppOAuthConfig(
client_id="cid",
client_secret="secret",
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
scope="all-apis",
)
@pytest.mark.asyncio
async def test_http_status_error_raises_value_error():
cache = DatabricksAppOAuthTokenCache()
request = httpx.Request("POST", _config().token_url)
error_response = httpx.Response(status_code=401, request=request)
client = _mock_http_handler(
post_error=httpx.HTTPStatusError(
"unauthorized", request=request, response=error_response
)
)
with patch(
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
return_value=client,
):
with pytest.raises(ValueError, match="status 401"):
await cache.async_get_token(_config())
@pytest.mark.asyncio
async def test_transport_error_raises_value_error():
cache = DatabricksAppOAuthTokenCache()
client = _mock_http_handler(post_error=httpx.ConnectError("boom"))
with patch(
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
return_value=client,
):
with pytest.raises(ValueError, match="token request failed"):
await cache.async_get_token(_config())
@pytest.mark.asyncio
async def test_non_object_json_body_raises():
cache = DatabricksAppOAuthTokenCache()
client = _mock_http_handler()
client.post.return_value.json.return_value = ["not", "an", "object"]
with patch(
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
return_value=client,
):
with pytest.raises(ValueError, match="non-object JSON"):
await cache.async_get_token(_config())
@pytest.mark.asyncio
@pytest.mark.parametrize("expires_in", [None, "not-a-number"])
async def test_invalid_expires_in_falls_back_to_default_ttl(expires_in):
cache = DatabricksAppOAuthTokenCache()
client = _mock_http_handler(expires_in=expires_in)
captured = {}
real_set = cache.set_cache
def _spy_set(key, value, **kwargs):
captured["ttl"] = kwargs.get("ttl")
return real_set(key, value, **kwargs)
with (
patch(
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
return_value=client,
),
patch.object(cache, "set_cache", side_effect=_spy_set),
):
await cache.async_get_token(_config())
# default TTL (3600) minus the 60s expiry buffer
assert captured["ttl"] == 3600 - 60
@pytest.mark.asyncio
async def test_short_lived_token_not_cached():
"""A token whose lifetime is below the refresh buffer is never cached and
leaves no per-key lock behind."""
cache = DatabricksAppOAuthTokenCache()
config = _config()
client = _mock_http_handler(access_token="short", expires_in=30)
with patch(
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
return_value=client,
):
await cache.async_get_token(config)
await cache.async_get_token(config)
assert cache.get_cache(config.cache_key) is None
assert config.cache_key not in cache._locks
assert client.post.await_count == 2
@pytest.mark.asyncio
async def test_rotated_secret_forces_new_token():
"""Rotating client_secret changes the cache key so a fresh token is minted."""
cache = DatabricksAppOAuthTokenCache()
old = DatabricksAppOAuthConfig(
client_id="cid",
client_secret="old-secret",
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
scope="all-apis",
)
rotated = DatabricksAppOAuthConfig(
client_id="cid",
client_secret="new-secret",
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
scope="all-apis",
)
assert old.cache_key != rotated.cache_key
clients = [_mock_http_handler("old-token"), _mock_http_handler("new-token")]
with patch(
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
side_effect=lambda *a, **k: clients.pop(0),
):
assert await cache.async_get_token(old) == "old-token"
assert await cache.async_get_token(rotated) == "new-token"
@pytest.mark.asyncio
async def test_lock_pruned_when_token_evicted():
"""The per-key lock is removed when its cached token is deleted/evicted."""
cache = DatabricksAppOAuthTokenCache()
config = _config()
client = _mock_http_handler("tok")
with patch(
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
return_value=client,
):
await cache.async_get_token(config)
assert config.cache_key in cache._locks
cache.delete_cache(config.cache_key)
assert config.cache_key not in cache._locks
@pytest.mark.asyncio
async def test_flush_cache_clears_locks():
"""flush_cache drops the per-key locks alongside the cached tokens."""
cache = DatabricksAppOAuthTokenCache()
config = _config()
client = _mock_http_handler("tok")
with patch(
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
return_value=client,
):
await cache.async_get_token(config)
assert config.cache_key in cache._locks
cache.flush_cache()
assert cache._locks == {}
assert cache.get_cache(config.cache_key) is None
# ---------------------------------------------------------------------------
# Public helper
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_resolve_returns_none_when_not_configured():
assert await resolve_databricks_app_auth_header(None) is None
assert await resolve_databricks_app_auth_header({"foo": "bar"}) is None
@pytest.mark.asyncio
async def test_resolve_returns_bearer_header():
from litellm.proxy.agent_endpoints.databricks_oauth import (
databricks_app_oauth_token_cache,
)
databricks_app_oauth_token_cache.flush_cache()
litellm_params = {
"databricks_oauth": {
"client_id": "resolve-cid",
"client_secret": "secret",
"workspace_url": "https://resolve.cloud.databricks.com",
}
}
client = _mock_http_handler(access_token="resolved-token")
with patch(
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
return_value=client,
):
header = await resolve_databricks_app_auth_header(litellm_params)
assert header == {"Authorization": "Bearer resolved-token"}

View file

@ -0,0 +1,968 @@
"""
Regression tests for the "provider field missing" bug on proxy-side
rate-limit errors.
Background
----------
The proxy's internal rate-limit hooks (parallel_request_limiter,
parallel_request_limiter_v3, dynamic_rate_limiter, dynamic_rate_limiter_v3,
batch_rate_limiter, max_budget_limiter, max_iterations_limiter,
max_budget_per_session_limiter) all fire from ``async_pre_call_hook`` —
*before* :func:`litellm.get_llm_provider` runs anywhere else in the request
lifecycle.
Until now, those hooks raised a bare ``HTTPException(429, ...)`` which carries
no ``llm_provider`` / ``model`` attribute. Downstream:
- The Prometheus ``litellm_proxy_failed_requests_metric`` reads
``exception.llm_provider`` via ``_get_exception_class_name`` — it came back
empty, so dashboards showed ``exception_class="HTTPException"`` with no
provider attribution.
- Observability callbacks that ``isinstance(e, RateLimitError)`` for
category routing missed these entirely.
The fix wraps every internal raise site in
:class:`ProxyHTTPRateLimitError` (an ``HTTPException`` *and* a
``litellm.RateLimitError``), and resolves ``model`` / ``llm_provider`` from
``data["model"]`` via :func:`get_llm_provider`. When the model is missing or
unparseable we fall back to ``llm_provider="litellm_proxy"`` so we never break
the request path with a second exception.
These tests pin both the happy path (provider correctly resolved) and the
fallback path (unknown model, missing model) for every limiter.
"""
import sys
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
import litellm
from litellm.caching.caching import DualCache
from litellm.exceptions import RateLimitError
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.hooks.batch_rate_limiter import (
BatchFileUsage,
_PROXY_BatchRateLimiter,
)
from litellm.proxy.hooks.dynamic_rate_limiter import _PROXY_DynamicRateLimitHandler
from litellm.proxy.hooks.dynamic_rate_limiter_v3 import (
_PROXY_DynamicRateLimitHandlerV3,
)
from litellm.proxy.hooks.max_budget_limiter import _PROXY_MaxBudgetLimiter
from litellm.proxy.hooks.max_budget_per_session_limiter import (
_PROXY_MaxBudgetPerSessionHandler,
)
from litellm.proxy.hooks.max_iterations_limiter import _PROXY_MaxIterationsHandler
from litellm.proxy.hooks.parallel_request_limiter import (
_PROXY_MaxParallelRequestsHandler,
)
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3,
)
from litellm.proxy.hooks.rate_limiter_utils import (
PROXY_LLM_PROVIDER_FALLBACK,
ProxyHTTPRateLimitError,
resolve_llm_provider_for_rate_limit,
)
from litellm.proxy.utils import InternalUsageCache
from litellm.types.agents import AgentResponse
# ---------------------------------------------------------------------------
# Helper class itself
# ---------------------------------------------------------------------------
class TestProxyHTTPRateLimitErrorClass:
"""Pin the dual ``HTTPException`` + ``RateLimitError`` shape."""
def test_is_both_http_exception_and_rate_limit_error(self):
e = ProxyHTTPRateLimitError(
status_code=429,
detail="boom",
model="gpt-4o-mini",
llm_provider="openai",
)
# FastAPI handler keys off HTTPException to render the 429.
assert isinstance(e, HTTPException)
# Prometheus / observability key off RateLimitError + .llm_provider.
assert isinstance(e, RateLimitError)
assert e.status_code == 429
assert e.model == "gpt-4o-mini"
assert e.llm_provider == "openai"
assert e.message == "boom"
assert e.detail == "boom"
def test_dict_detail_is_stringified_for_message(self):
# Some hooks pass a dict detail (e.g. dynamic_rate_limiter v1) — the
# `message` attr (read by RateLimitError.__str__ and observability
# callbacks) must still be a string.
e = ProxyHTTPRateLimitError(
status_code=429,
detail={"error": "over rpm"},
model="claude-3-5-sonnet",
llm_provider="anthropic",
)
assert isinstance(e.message, str)
assert "over rpm" in e.message
def test_defaults_to_litellm_proxy_provider(self):
e = ProxyHTTPRateLimitError(status_code=429, detail="x")
assert e.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
assert e.model == ""
def test_none_provider_normalized_to_fallback(self):
e = ProxyHTTPRateLimitError(
status_code=429,
detail="x",
model=None, # type: ignore[arg-type]
llm_provider=None, # type: ignore[arg-type]
)
assert e.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
assert e.model == ""
class TestResolveLLMProviderForRateLimit:
@pytest.mark.parametrize(
"model, expected_provider",
[
("gpt-4o-mini", "openai"),
("anthropic/claude-3-5-sonnet", "anthropic"),
("bedrock/meta.llama3-1-70b-instruct-v1:0", "bedrock"),
],
)
def test_known_models_resolve_provider(self, model, expected_provider):
resolved_model, provider = resolve_llm_provider_for_rate_limit(model)
assert provider == expected_provider
assert resolved_model # non-empty
@pytest.mark.parametrize("model", [None, "", "totally-not-a-real-model-name"])
def test_missing_or_unknown_model_falls_back(self, model):
# Must never raise — the resolver wraps `get_llm_provider` defensively
# because raising here would mask the rate-limit error we're trying
# to surface to the user.
resolved_model, provider = resolve_llm_provider_for_rate_limit(model)
assert provider == PROXY_LLM_PROVIDER_FALLBACK
# Resolver returns the input model verbatim on the unknown branch so
# the `.model` attribute is never silently swapped to a different one.
if not model:
assert resolved_model == ""
else:
assert resolved_model == model
def test_get_llm_provider_raising_is_swallowed(self):
# If get_llm_provider itself blows up (unexpected error), we still
# fall back rather than letting the secondary exception escape.
with patch.object(
litellm,
"get_llm_provider",
side_effect=RuntimeError("boom"),
):
resolved_model, provider = resolve_llm_provider_for_rate_limit("anything")
assert provider == PROXY_LLM_PROVIDER_FALLBACK
assert resolved_model == "anything"
# ---------------------------------------------------------------------------
# parallel_request_limiter v1
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_parallel_request_limiter_v1_populates_provider_when_at_rpm_limit():
"""
Trip the per-key RPM cap and assert the raised exception carries
``model`` / ``llm_provider`` resolved from ``data["model"]``.
"""
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-rl-test",
max_parallel_requests=10,
rpm_limit=1,
tpm_limit=10,
)
data = {"model": "gpt-4o-mini"}
# First request consumes the budget.
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data=data,
call_type="completion",
)
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data=data,
call_type="completion",
)
exc = exc_info.value
assert exc.status_code == 429
assert isinstance(exc, RateLimitError)
assert exc.llm_provider == "openai"
assert exc.model == "gpt-4o-mini"
@pytest.mark.asyncio
async def test_parallel_request_limiter_v1_zero_limit_path_populates_provider():
"""
When tpm_limit / rpm_limit is 0 the limiter takes the
``raise_rate_limit_error`` path. That path receives ``requested_model``
via the call-site change and must pass it through.
"""
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-rl-zero",
max_parallel_requests=0,
rpm_limit=10,
tpm_limit=10,
)
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={"model": "anthropic/claude-3-5-sonnet"},
call_type="completion",
)
exc = exc_info.value
assert exc.status_code == 429
assert isinstance(exc, RateLimitError)
assert exc.llm_provider == "anthropic"
assert exc.model == "claude-3-5-sonnet"
@pytest.mark.asyncio
async def test_parallel_request_limiter_v1_global_limit_populates_provider():
"""global_max_parallel_requests path also threads the model through."""
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-global")
# Pre-fill the global counter so the next call exceeds it.
await handler.internal_usage_cache.async_set_cache(
key="global_max_parallel_requests",
value=5,
local_only=True,
litellm_parent_otel_span=None,
)
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={
"model": "bedrock/meta.llama3-1-70b-instruct-v1:0",
"metadata": {"global_max_parallel_requests": 1},
},
call_type="completion",
)
exc = exc_info.value
assert exc.status_code == 429
assert exc.llm_provider == "bedrock"
assert exc.model == "meta.llama3-1-70b-instruct-v1:0"
@pytest.mark.asyncio
async def test_parallel_request_limiter_v1_unknown_model_falls_back():
"""
When ``data["model"]`` is unparseable, the resolver falls back to
``litellm_proxy`` — and crucially does *not* leak a secondary exception.
"""
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-rl-unknown",
max_parallel_requests=10,
rpm_limit=1,
tpm_limit=10,
)
data = {"model": "totally-not-a-real-model"}
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data=data,
call_type="completion",
)
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data=data,
call_type="completion",
)
exc = exc_info.value
assert exc.status_code == 429
assert exc.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
# Resolver returns the input verbatim so we don't silently relabel the
# model in the user-facing 429 detail.
assert exc.model == "totally-not-a-real-model"
@pytest.mark.asyncio
async def test_parallel_request_limiter_v1_missing_model_falls_back():
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-rl-no-model",
max_parallel_requests=10,
rpm_limit=1,
tpm_limit=10,
)
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={},
call_type="completion",
)
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={},
call_type="completion",
)
exc = exc_info.value
assert exc.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
assert exc.model == ""
# ---------------------------------------------------------------------------
# parallel_request_limiter v3
# ---------------------------------------------------------------------------
def _v3_over_limit_response(rate_limit_type: str = "rpm") -> dict:
return {
"overall_code": "OVER_LIMIT",
"statuses": [
{
"code": "OVER_LIMIT",
"descriptor_key": "key",
"current_limit": 1,
"limit_remaining": -1,
"rate_limit_type": rate_limit_type,
}
],
}
@pytest.mark.asyncio
@pytest.mark.parametrize(
"model, expected_provider",
[
("gpt-4o-mini", "openai"),
("anthropic/claude-3-5-sonnet", "anthropic"),
],
)
async def test_parallel_request_limiter_v3_populates_provider(model, expected_provider):
handler = _PROXY_MaxParallelRequestsHandler_v3(
internal_usage_cache=InternalUsageCache(DualCache())
)
descriptors = [{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 1}}]
over = _v3_over_limit_response()
with pytest.raises(HTTPException) as exc_info:
handler._handle_rate_limit_error(
response=over,
descriptors=descriptors,
requested_model=model,
)
exc = exc_info.value
assert exc.status_code == 429
assert isinstance(exc, RateLimitError)
assert exc.llm_provider == expected_provider
# v3 may strip the "anthropic/" prefix in the resolved model — accept
# either; we only care that the provider field is correct and the model
# is non-empty.
assert exc.model
@pytest.mark.asyncio
async def test_parallel_request_limiter_v3_unknown_model_falls_back():
handler = _PROXY_MaxParallelRequestsHandler_v3(
internal_usage_cache=InternalUsageCache(DualCache())
)
descriptors = [{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 1}}]
with pytest.raises(HTTPException) as exc_info:
handler._handle_rate_limit_error(
response=_v3_over_limit_response(),
descriptors=descriptors,
requested_model="totally-bogus",
)
assert exc_info.value.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
assert exc_info.value.model == "totally-bogus"
@pytest.mark.asyncio
async def test_parallel_request_limiter_v3_missing_model_falls_back():
handler = _PROXY_MaxParallelRequestsHandler_v3(
internal_usage_cache=InternalUsageCache(DualCache())
)
descriptors = [{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 1}}]
with pytest.raises(HTTPException) as exc_info:
handler._handle_rate_limit_error(
response=_v3_over_limit_response(),
descriptors=descriptors,
requested_model=None,
)
assert exc_info.value.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
assert exc_info.value.model == ""
# ---------------------------------------------------------------------------
# dynamic_rate_limiter v1
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_dynamic_rate_limiter_v1_tpm_zero_populates_provider():
handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache())
handler.check_available_usage = AsyncMock(return_value=(0, 5, 100, 5, 1))
user_api_key_dict = UserAPIKeyAuth(api_key="sk-dyn")
user_api_key_dict.metadata = {}
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={"model": "gpt-4o-mini"},
call_type="completion",
)
exc = exc_info.value
assert exc.status_code == 429
assert isinstance(exc, RateLimitError)
assert exc.llm_provider == "openai"
assert exc.model == "gpt-4o-mini"
@pytest.mark.asyncio
async def test_dynamic_rate_limiter_v1_rpm_zero_populates_provider():
handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache())
handler.check_available_usage = AsyncMock(return_value=(5, 0, 5, 100, 1))
user_api_key_dict = UserAPIKeyAuth(api_key="sk-dyn")
user_api_key_dict.metadata = {}
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={"model": "anthropic/claude-3-5-sonnet"},
call_type="completion",
)
exc = exc_info.value
assert exc.llm_provider == "anthropic"
assert exc.model == "claude-3-5-sonnet"
@pytest.mark.asyncio
async def test_dynamic_rate_limiter_v1_unknown_model_falls_back():
handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache())
handler.check_available_usage = AsyncMock(return_value=(0, 5, 100, 5, 1))
user_api_key_dict = UserAPIKeyAuth(api_key="sk-dyn")
user_api_key_dict.metadata = {}
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={"model": "no-such-model"},
call_type="completion",
)
assert exc_info.value.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
assert exc_info.value.model == "no-such-model"
# ---------------------------------------------------------------------------
# dynamic_rate_limiter v3 — exercise just the raise path via the helper, not
# the full Redis/Lua stack.
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_dynamic_rate_limiter_v3_model_capacity_path_populates_provider():
"""
The v3 dynamic limiter has three raise sites: model_saturation_check,
priority_model, and the fail-closed unknown-descriptor branch. We patch
the atomic increment to short-circuit straight into the model_saturation
path — that's the most common production trip — and confirm the
raised exception carries provider info.
"""
from litellm.types.router import ModelGroupInfo
handler = _PROXY_DynamicRateLimitHandlerV3(internal_usage_cache=DualCache())
handler.v3_limiter.atomic_check_and_increment_by_n = AsyncMock(
return_value={
"overall_code": "OVER_LIMIT",
"statuses": [
{
"code": "OVER_LIMIT",
"descriptor_key": "model_saturation_check",
"current_limit": 100,
"limit_remaining": 0,
"rate_limit_type": "rpm",
}
],
}
)
handler._create_priority_based_descriptors = MagicMock(return_value=[])
handler._create_model_tracking_descriptor = MagicMock(
return_value={
"key": "model_saturation_check",
"value": "gpt-4o-mini",
"rate_limit": {"requests_per_unit": 100},
}
)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-dyn-v3")
user_api_key_dict.metadata = {}
model_info = ModelGroupInfo(model_group="gpt-4o-mini", providers=["openai"])
with pytest.raises(HTTPException) as exc_info:
await handler._check_rate_limits(
model="gpt-4o-mini",
model_group_info=model_info,
user_api_key_dict=user_api_key_dict,
priority="default",
saturation=1.0,
data={"model": "gpt-4o-mini"},
)
exc = exc_info.value
assert exc.status_code == 429
assert isinstance(exc, RateLimitError)
assert exc.llm_provider == "openai"
assert exc.model == "gpt-4o-mini"
@pytest.mark.asyncio
async def test_dynamic_rate_limiter_v3_unknown_descriptor_path_populates_provider():
"""Fail-closed unknown-descriptor branch must still attribute provider."""
from litellm.types.router import ModelGroupInfo
handler = _PROXY_DynamicRateLimitHandlerV3(internal_usage_cache=DualCache())
handler.v3_limiter.atomic_check_and_increment_by_n = AsyncMock(
return_value={
"overall_code": "OVER_LIMIT",
"statuses": [
{
"code": "OVER_LIMIT",
"descriptor_key": "something_we_dont_handle",
"current_limit": 1,
"limit_remaining": 0,
"rate_limit_type": "rpm",
}
],
}
)
handler._create_priority_based_descriptors = MagicMock(return_value=[])
handler._create_model_tracking_descriptor = MagicMock(
return_value={
"key": "model_saturation_check",
"value": "gpt-4o-mini",
"rate_limit": {"requests_per_unit": 1},
}
)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-dyn-v3-unknown")
user_api_key_dict.metadata = {}
model_info = ModelGroupInfo(model_group="gpt-4o-mini", providers=["openai"])
with pytest.raises(HTTPException) as exc_info:
await handler._check_rate_limits(
model="gpt-4o-mini",
model_group_info=model_info,
user_api_key_dict=user_api_key_dict,
priority="default",
saturation=1.0,
data={"model": "gpt-4o-mini"},
)
assert exc_info.value.llm_provider == "openai"
# ---------------------------------------------------------------------------
# batch_rate_limiter
# ---------------------------------------------------------------------------
def _batch_over_limit_response() -> dict:
return {
"overall_code": "OVER_LIMIT",
"statuses": [
{
"code": "OVER_LIMIT",
"descriptor_key": "key",
"current_limit": 10,
"limit_remaining": -5,
"rate_limit_type": "requests",
}
],
}
@pytest.mark.asyncio
async def test_batch_rate_limiter_populates_provider():
"""
batch_rate_limiter trips when the file's request/token count exceeds the
remaining window. The raise must thread `data["model"]` through the
helper.
"""
parallel_limiter = MagicMock()
parallel_limiter.window_size = 60
parallel_limiter._create_rate_limit_descriptors = MagicMock(
return_value=[
{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 10}}
]
)
parallel_limiter.atomic_check_and_increment_by_n = AsyncMock(
return_value=_batch_over_limit_response()
)
handler = _PROXY_BatchRateLimiter(
internal_usage_cache=InternalUsageCache(DualCache()),
parallel_request_limiter=parallel_limiter,
)
with pytest.raises(HTTPException) as exc_info:
await handler._check_and_increment_batch_counters(
user_api_key_dict=UserAPIKeyAuth(api_key="sk-batch"),
data={"model": "gpt-4o-mini"},
batch_usage=BatchFileUsage(total_tokens=100, request_count=15),
)
exc = exc_info.value
assert exc.status_code == 429
assert isinstance(exc, RateLimitError)
assert exc.llm_provider == "openai"
assert exc.model == "gpt-4o-mini"
@pytest.mark.asyncio
async def test_batch_rate_limiter_unknown_model_falls_back():
parallel_limiter = MagicMock()
parallel_limiter.window_size = 60
parallel_limiter._create_rate_limit_descriptors = MagicMock(
return_value=[
{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 10}}
]
)
parallel_limiter.atomic_check_and_increment_by_n = AsyncMock(
return_value=_batch_over_limit_response()
)
handler = _PROXY_BatchRateLimiter(
internal_usage_cache=InternalUsageCache(DualCache()),
parallel_request_limiter=parallel_limiter,
)
with pytest.raises(HTTPException) as exc_info:
await handler._check_and_increment_batch_counters(
user_api_key_dict=UserAPIKeyAuth(api_key="sk-batch"),
data={"model": "fake-model-xyz"},
batch_usage=BatchFileUsage(total_tokens=100, request_count=15),
)
assert exc_info.value.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
# ---------------------------------------------------------------------------
# max_budget_limiter
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_max_budget_limiter_populates_provider():
handler = _PROXY_MaxBudgetLimiter()
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-budget",
user_id="user-1",
user_max_budget=10.0,
)
with patch(
"litellm.proxy.proxy_server.get_current_spend",
new=AsyncMock(return_value=10.0),
):
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={"model": "gpt-4o-mini"},
call_type="completion",
)
exc = exc_info.value
assert exc.status_code == 429
assert isinstance(exc, RateLimitError)
assert exc.llm_provider == "openai"
assert exc.model == "gpt-4o-mini"
@pytest.mark.asyncio
async def test_max_budget_limiter_no_model_falls_back():
handler = _PROXY_MaxBudgetLimiter()
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-budget",
user_id="user-1",
user_max_budget=10.0,
)
with patch(
"litellm.proxy.proxy_server.get_current_spend",
new=AsyncMock(return_value=10.0),
):
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={},
call_type="completion",
)
assert exc_info.value.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
assert exc_info.value.model == ""
# ---------------------------------------------------------------------------
# max_iterations_limiter
# ---------------------------------------------------------------------------
def _make_iter_agent(max_iterations: int) -> AgentResponse:
return AgentResponse(
agent_id="agent-iter",
agent_name="iter-agent",
litellm_params={"max_iterations": max_iterations},
agent_card_params={"name": "iter-agent", "version": "1.0.0"},
)
@pytest.mark.asyncio
async def test_max_iterations_limiter_populates_provider():
local_cache = DualCache()
handler = _PROXY_MaxIterationsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-iter", agent_id="agent-iter")
with patch(
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = _make_iter_agent(max_iterations=1)
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={
"model": "gpt-4o-mini",
"metadata": {"session_id": "session-iter-1"},
},
call_type="completion",
)
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={
"model": "gpt-4o-mini",
"metadata": {"session_id": "session-iter-1"},
},
call_type="completion",
)
exc = exc_info.value
assert exc.status_code == 429
assert isinstance(exc, RateLimitError)
assert exc.llm_provider == "openai"
assert exc.model == "gpt-4o-mini"
@pytest.mark.asyncio
async def test_max_iterations_limiter_unknown_model_falls_back():
local_cache = DualCache()
handler = _PROXY_MaxIterationsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-iter", agent_id="agent-iter")
with patch(
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = _make_iter_agent(max_iterations=1)
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={
"model": "no-such-model",
"metadata": {"session_id": "session-iter-2"},
},
call_type="completion",
)
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={
"model": "no-such-model",
"metadata": {"session_id": "session-iter-2"},
},
call_type="completion",
)
assert exc_info.value.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
# ---------------------------------------------------------------------------
# max_budget_per_session_limiter
# ---------------------------------------------------------------------------
def _make_session_budget_agent(max_budget: float) -> AgentResponse:
return AgentResponse(
agent_id="agent-session-budget",
agent_name="session-budget-agent",
litellm_params={"max_budget_per_session": max_budget},
agent_card_params={"name": "session-budget-agent", "version": "1.0.0"},
)
@pytest.mark.asyncio
async def test_max_budget_per_session_limiter_populates_provider():
handler = _PROXY_MaxBudgetPerSessionHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-session-budget", agent_id="agent-session-budget"
)
with patch(
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = _make_session_budget_agent(
max_budget=1.0
)
with patch.object(
handler, "_get_current_spend", new=AsyncMock(return_value=5.0)
):
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={
"model": "anthropic/claude-3-5-sonnet",
"metadata": {"session_id": "session-budget-1"},
},
call_type="completion",
)
exc = exc_info.value
assert exc.status_code == 429
assert isinstance(exc, RateLimitError)
assert exc.llm_provider == "anthropic"
@pytest.mark.asyncio
async def test_max_budget_per_session_limiter_unknown_model_falls_back():
handler = _PROXY_MaxBudgetPerSessionHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-session-budget", agent_id="agent-session-budget"
)
with patch(
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
) as mock_registry:
mock_registry.get_agent_by_id.return_value = _make_session_budget_agent(
max_budget=1.0
)
with patch.object(
handler, "_get_current_spend", new=AsyncMock(return_value=5.0)
):
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=DualCache(),
data={
"model": "no-such-model",
"metadata": {"session_id": "session-budget-2"},
},
call_type="completion",
)
assert exc_info.value.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
# ---------------------------------------------------------------------------
# Prometheus integration: failure metric reads exception.llm_provider
# via _get_exception_class_name. With the fix, this returns
# "Openai.RateLimitError" instead of plain "HTTPException" for proxy-side
# 429s on a known model. Pin that contract — that's what dashboards see.
# ---------------------------------------------------------------------------
def test_prometheus_exception_class_name_includes_provider():
from litellm.integrations.prometheus import PrometheusLogger
exc = ProxyHTTPRateLimitError(
status_code=429,
detail="over limit",
model="gpt-4o-mini",
llm_provider="openai",
)
name = PrometheusLogger._get_exception_class_name(exc)
# Format is "{Provider.}{ClassName}" per `_get_exception_class_name`.
assert name.startswith("Openai.")
# And specifically: it ends in our exception class. (We don't pin the
# full string to avoid coupling the test to PR #27687's parallel rename.)
assert name.endswith("ProxyHTTPRateLimitError")
def test_prometheus_exception_class_name_falls_back_when_no_model():
from litellm.integrations.prometheus import PrometheusLogger
exc = ProxyHTTPRateLimitError(status_code=429, detail="over limit")
name = PrometheusLogger._get_exception_class_name(exc)
# `litellm_proxy` -> `Litellm_proxy.` (capitalize first char only).
assert name.startswith("Litellm_proxy.")
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-vv", "-x"]))

File diff suppressed because it is too large Load diff

View file

@ -4398,3 +4398,55 @@ class TestMCPUserEnvVarsAccessControl:
)
assert exc.value.status_code == 403
get_mcp_server_mock.assert_not_awaited()
def test_oauth2_flow_accepted_on_create_request():
"""NewMCPServerRequest carries oauth2_flow through to the persisted dict."""
from litellm.proxy._experimental.mcp_server.db import _prepare_mcp_server_data
payload = NewMCPServerRequest(
server_name="m2m-server",
url="https://example.com/mcp",
transport="http",
auth_type="oauth2",
token_url="https://idp.example.com/oauth/token",
oauth2_flow="client_credentials",
)
data_dict = _prepare_mcp_server_data(payload)
assert data_dict["oauth2_flow"] == "client_credentials"
def test_oauth2_flow_round_trips_on_update_and_response_models():
"""oauth2_flow survives UpdateMCPServerRequest and the LiteLLM_MCPServerTable
response model. Before the fix these models dropped the field (no attribute),
which is why a persisted value never round-tripped."""
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
UpdateMCPServerRequest,
)
update = UpdateMCPServerRequest(
server_id="srv-1", oauth2_flow="client_credentials"
)
assert update.oauth2_flow == "client_credentials"
row = LiteLLM_MCPServerTable(
server_id="srv-1",
transport="http",
oauth2_flow="client_credentials",
)
assert row.oauth2_flow == "client_credentials"
def test_oauth2_flow_defaults_to_none_when_omitted():
"""Omitting oauth2_flow is valid and resolves to None (runtime infers it)."""
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
UpdateMCPServerRequest,
)
assert UpdateMCPServerRequest(server_id="srv-1").oauth2_flow is None
assert (
LiteLLM_MCPServerTable(server_id="srv-1", transport="http").oauth2_flow
is None
)

View file

@ -327,7 +327,8 @@ class TestResponseAPILoggingUtils:
"output_tokens_details": {
"reasoning_tokens": 30,
"image_tokens": 100,
"text_tokens": 70,
"text_tokens": 50,
"audio_tokens": 20,
},
}
@ -346,7 +347,61 @@ class TestResponseAPILoggingUtils:
assert result.completion_tokens_details is not None
assert result.completion_tokens_details.reasoning_tokens == 30
assert result.completion_tokens_details.image_tokens == 100
assert result.completion_tokens_details.text_tokens == 70
assert result.completion_tokens_details.text_tokens == 50
assert result.completion_tokens_details.audio_tokens == 20
def test_transform_response_api_usage_with_realtime_keys(self):
"""Realtime input_token_details / output_token_details normalize for Usage."""
usage = {
"input_tokens": 10,
"output_tokens": 20,
"total_tokens": 30,
"input_token_details": {
"text_tokens": 8,
"audio_tokens": 2,
"cached_tokens": 0,
},
"output_token_details": {
"text_tokens": 12,
"audio_tokens": 8,
},
}
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
)
assert result.prompt_tokens_details is not None
assert result.prompt_tokens_details.text_tokens == 8
assert result.prompt_tokens_details.audio_tokens == 2
assert result.completion_tokens_details is not None
assert result.completion_tokens_details.text_tokens == 12
assert result.completion_tokens_details.audio_tokens == 8
def test_transform_response_api_usage_tokens_details_keep_values(self):
"""Keeps input_tokens_details / output_tokens_details when singular keys are also present."""
usage = {
"input_tokens": 10,
"output_tokens": 20,
"total_tokens": 30,
"input_tokens_details": {"text_tokens": 10},
"output_tokens_details": {"text_tokens": 20},
"input_token_details": {"text_tokens": 1, "audio_tokens": 99},
"output_token_details": {"text_tokens": 2, "audio_tokens": 98},
}
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
)
assert result.prompt_tokens_details is not None
assert result.prompt_tokens_details.text_tokens == 10
assert result.prompt_tokens_details.audio_tokens is None
assert result.completion_tokens_details is not None
assert result.completion_tokens_details.text_tokens == 20
assert result.completion_tokens_details.audio_tokens is None
class TestResponsesAPIProviderSpecificParams:

View file

@ -182,3 +182,24 @@ def test_opus_4_8_provider_resolves_via_model_info(local_model_cost_map):
assert info["litellm_provider"] == "anthropic"
assert info["max_input_tokens"] == 1000000
assert info["max_output_tokens"] == 128000
@pytest.mark.parametrize(
"cost_map",
[_load_root_cost_map(), GetModelCostMap.load_local_model_cost_map()],
ids=["root", "bundled_backup"],
)
def test_opus_4_8_all_variants_carry_adaptive_thinking_flag(cost_map):
"""Every Opus 4.8 entry must advertise ``supports_adaptive_thinking``.
Adaptive-thinking detection is cost-map driven, so a single variant missing
the flag silently sends the legacy ``thinking.type='enabled'`` shape and the
provider 400s (issue #29188, which the Bedrock/Vertex/Azure variants hit
because only the bare ``claude-opus-4-8`` entry carried the flag). This guards
against a future variant being added without it."""
variants = [k for k in cost_map if "claude-opus-4-8" in k]
assert variants, "no claude-opus-4-8 entries found in cost map"
missing = [
k for k in variants if cost_map[k].get("supports_adaptive_thinking") is not True
]
assert not missing, f"missing supports_adaptive_thinking: {missing}"

View file

@ -385,6 +385,51 @@ def test_handle_realtime_stream_cost_calculation():
)
assert cost == 0.0 # No usage, no cost
def test_realtime_stream_combines_text_and_audio_token_details():
"""Realtime response.done usage with input_token_details / output_token_details."""
from litellm.cost_calculator import RealtimeAPITokenUsageProcessor
results: OpenAIRealtimeStreamList = [
{"type": "session.created", "session": {"model": "gpt-4o-realtime-preview"}},
{
"type": "response.done",
"response": {
"usage": {
"input_tokens": 10,
"output_tokens": 20,
"total_tokens": 30,
"input_token_details": {"text_tokens": 8, "audio_tokens": 2},
"output_token_details": {"text_tokens": 12, "audio_tokens": 8},
}
},
},
{
"type": "response.done",
"response": {
"usage": {
"input_tokens": 5,
"output_tokens": 15,
"total_tokens": 20,
"input_token_details": {"text_tokens": 3, "audio_tokens": 2},
"output_token_details": {"text_tokens": 5, "audio_tokens": 10},
}
},
},
]
combined = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(
results=results,
)
assert combined.prompt_tokens_details is not None
assert combined.prompt_tokens_details.text_tokens == 11
assert combined.prompt_tokens_details.audio_tokens == 4
assert combined.completion_tokens_details is not None
assert combined.completion_tokens_details.text_tokens == 17
assert combined.completion_tokens_details.audio_tokens == 18
def test_realtime_logging_object_allows_null_transcript_in_conversation_item_added():
results: OpenAIRealtimeStreamList = [

View file

@ -3,7 +3,9 @@ Test case normalization in LitellmParams for all guardrail types
"""
import pytest
from litellm.types.guardrails import LitellmParams
from pydantic import ValidationError
from litellm.types.guardrails import BaseLitellmParams, LitellmParams
class TestLitellmParamsCaseNormalization:
@ -89,3 +91,66 @@ class TestLitellmParamsCaseNormalization:
)
assert params.on_disallowed_action in ["block", "rewrite"]
assert params.on_disallowed_action.islower()
class TestSensitiveDataRoutingValidation:
"""on_sensitive_data='route' requires a target model to be set"""
def test_route_with_target_model_is_valid(self):
params = LitellmParams(
guardrail="presidio",
mode="pre_call",
on_sensitive_data="route",
sensitive_data_route_to_model="on-prem-model",
)
assert params.on_sensitive_data == "route"
assert params.sensitive_data_route_to_model == "on-prem-model"
def test_route_without_target_model_raises(self):
with pytest.raises(ValidationError, match="sensitive_data_route_to_model"):
LitellmParams(
guardrail="presidio",
mode="pre_call",
on_sensitive_data="route",
)
def test_base_params_route_without_target_model_raises(self):
with pytest.raises(ValidationError, match="sensitive_data_route_to_model"):
BaseLitellmParams(on_sensitive_data="route")
def test_base_params_normalize_on_sensitive_data_case(self):
params = BaseLitellmParams(
on_sensitive_data="Route",
sensitive_data_route_to_model="on-prem-model",
)
assert params.on_sensitive_data == "route"
def test_base_params_capitalized_route_without_target_model_raises(self):
with pytest.raises(ValidationError, match="sensitive_data_route_to_model"):
BaseLitellmParams(on_sensitive_data="ROUTE")
def test_block_without_target_model_is_valid(self):
params = LitellmParams(
guardrail="presidio",
mode="pre_call",
on_sensitive_data="block",
)
assert params.on_sensitive_data == "block"
assert params.sensitive_data_route_to_model is None
def test_on_sensitive_data_is_case_normalized(self):
params = LitellmParams(
guardrail="presidio",
mode="pre_call",
on_sensitive_data="Route",
sensitive_data_route_to_model="on-prem-model",
)
assert params.on_sensitive_data == "route"
def test_on_sensitive_data_uppercase_block_normalized(self):
params = LitellmParams(
guardrail="presidio",
mode="pre_call",
on_sensitive_data="BLOCK",
)
assert params.on_sensitive_data == "block"