diff --git a/CLAUDE.md b/CLAUDE.md index d2d8601a9c9..02a9630b486 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -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 diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260604120000_add_oauth2_flow_to_mcp_servers/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260604120000_add_oauth2_flow_to_mcp_servers/migration.sql new file mode 100644 index 00000000000..fee6926d963 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260604120000_add_oauth2_flow_to_mcp_servers/migration.sql @@ -0,0 +1,2 @@ +-- AlterTable +ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "oauth2_flow" TEXT; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index cbc9f7fd5a8..34cf76c4c86 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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) diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 9a4b158b622..88029615ba8 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -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, diff --git a/litellm/exceptions.py b/litellm/exceptions.py index 17f5b43c273..15f6030d4a3 100644 --- a/litellm/exceptions.py +++ b/litellm/exceptions.py @@ -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) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 6d0d73e033d..fc5f0429b63 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -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 diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 9fc09807369..648fe671140 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -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" """ diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 4f15d1b3cef..7949e150c23 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -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, diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index 31131d722ab..3f002d73cbc 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -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) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 5c502c56ffe..eedab7fc36c 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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, diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index fbe6c097202..4eec27f29bd 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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, diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index eb5a6c1a0b6..39945d5ebe2 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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 diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index 9f1403d4328..7b2f75e1cff 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -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 diff --git a/litellm/proxy/agent_endpoints/databricks_oauth.py b/litellm/proxy/agent_endpoints/databricks_oauth.py new file mode 100644 index 00000000000..1c1f5a2b4c4 --- /dev/null +++ b/litellm/proxy/agent_endpoints/databricks_oauth.py @@ -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 "}`` 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}"} diff --git a/litellm/proxy/hooks/__init__.py b/litellm/proxy/hooks/__init__.py index 34505427d79..0db661fb508 100644 --- a/litellm/proxy/hooks/__init__.py +++ b/litellm/proxy/hooks/__init__.py @@ -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 ## diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 435b6eea45b..df485411d76 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -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( diff --git a/litellm/proxy/hooks/dynamic_rate_limiter.py b/litellm/proxy/hooks/dynamic_rate_limiter.py index f1c1d487cc1..57cd538507e 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter.py @@ -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 diff --git a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py index 861083e7dfa..bfc6e2c2f72 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py @@ -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 diff --git a/litellm/proxy/hooks/max_budget_limiter.py b/litellm/proxy/hooks/max_budget_limiter.py index 9a7e5117945..658d7995631 100644 --- a/litellm/proxy/hooks/max_budget_limiter.py +++ b/litellm/proxy/hooks/max_budget_limiter.py @@ -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: diff --git a/litellm/proxy/hooks/max_budget_per_session_limiter.py b/litellm/proxy/hooks/max_budget_per_session_limiter.py index 59fb101f557..0b63465c4a5 100644 --- a/litellm/proxy/hooks/max_budget_per_session_limiter.py +++ b/litellm/proxy/hooks/max_budget_per_session_limiter.py @@ -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 diff --git a/litellm/proxy/hooks/max_iterations_limiter.py b/litellm/proxy/hooks/max_iterations_limiter.py index df9a298ca03..d5bc669c928 100644 --- a/litellm/proxy/hooks/max_iterations_limiter.py +++ b/litellm/proxy/hooks/max_iterations_limiter.py @@ -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( diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index 43c5fc68723..c6324c3e3a3 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -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: diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 4343747d104..9fdb146b19d 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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( diff --git a/litellm/proxy/hooks/rate_limiter_utils.py b/litellm/proxy/hooks/rate_limiter_utils.py index 927bac0de58..0ba3df448e5 100644 --- a/litellm/proxy/hooks/rate_limiter_utils.py +++ b/litellm/proxy/hooks/rate_limiter_utils.py @@ -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] diff --git a/litellm/proxy/hooks/sensitive_data_routing.py b/litellm/proxy/hooks/sensitive_data_routing.py new file mode 100644 index 00000000000..0a907b1d71c --- /dev/null +++ b/litellm/proxy/hooks/sensitive_data_routing.py @@ -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 diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index cbc9f7fd5a8..34cf76c4c86 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 8bd50a50a38..e77e24c9e71 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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] diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index 46a2894bd10..60badb57d2a 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -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( diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 25c0bcabb4a..0d81e25592d 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -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): diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 4c227656e5f..c50b3def651 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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, diff --git a/schema.prisma b/schema.prisma index cbc9f7fd5a8..34cf76c4c86 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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) diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index 4956fccb8db..f0bc7b8ebed 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -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" diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py index d501ae0f79a..4c330312930 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py @@ -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 diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_reasoning_effort_translation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_reasoning_effort_translation.py index 09601a65811..71755e6da3f 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_reasoning_effort_translation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_reasoning_effort_translation.py @@ -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", [ diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py index d34b6ffc831..a09e55d4ed7 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py @@ -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 diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index 083f4a97ab6..279e9730e69 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -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) diff --git a/tests/test_litellm/proxy/agent_endpoints/test_agent_headers.py b/tests/test_litellm/proxy/agent_endpoints/test_agent_headers.py index 93ba9dc922c..fa530e0975a 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_agent_headers.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_headers.py @@ -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 # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/agent_endpoints/test_databricks_oauth.py b/tests/test_litellm/proxy/agent_endpoints/test_databricks_oauth.py new file mode 100644 index 00000000000..39f51ee0401 --- /dev/null +++ b/tests/test_litellm/proxy/agent_endpoints/test_databricks_oauth.py @@ -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"} diff --git a/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py b/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py new file mode 100644 index 00000000000..8c74919df19 --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py @@ -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"])) diff --git a/tests/test_litellm/proxy/hooks/test_sensitive_data_routing.py b/tests/test_litellm/proxy/hooks/test_sensitive_data_routing.py new file mode 100644 index 00000000000..78d2c3af0f3 --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_sensitive_data_routing.py @@ -0,0 +1,1036 @@ +""" +Tests for Sensitive Data Routing feature. + +This feature allows guardrails to route requests to a different model +(typically on-premise) when sensitive data is detected, instead of blocking. +All subsequent requests in the same session are routed to the same model. +""" + +import asyncio +from typing import Any, Dict, Optional +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from litellm.caching.caching import DualCache +from litellm.exceptions import SensitiveDataRouteException +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + get_session_id_from_request_data, +) +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.hooks.sensitive_data_routing import ( + _PROXY_SensitiveDataRoutingHandler, + SENSITIVE_ROUTING_CACHE_PREFIX, + DEFAULT_SENSITIVE_ROUTING_TTL, +) + + +class MockInternalUsageCache: + def __init__(self): + self._cache: Dict[str, Any] = {} + self._ttls: Dict[str, int] = {} + self.dual_cache = MagicMock() + self.dual_cache.redis_cache = None + + async def async_get_cache(self, key: str, **kwargs) -> Optional[Any]: + return self._cache.get(key) + + async def async_set_cache(self, key: str, value: Any, ttl: int = 3600, **kwargs): + self._cache[key] = value + self._ttls[key] = ttl + + +class TestSensitiveDataRoutingHandler: + @pytest.fixture + def handler(self): + cache = MockInternalUsageCache() + return _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + + @pytest.fixture + def user_api_key_dict(self): + return UserAPIKeyAuth(api_key="test-key") + + @pytest.mark.asyncio + async def test_set_session_routing(self, handler): + key = UserAPIKeyAuth(api_key="hashed-key") + await handler.set_session_routing( + session_id="test-session-123", + model="on-premise-model", + user_api_key_dict=key, + guardrail_name="test-guardrail", + ) + + routed_model = await handler._get_routed_model("test-session-123", key) + assert routed_model == "on-premise-model" + + def test_get_session_id_from_metadata(self): + data = {"metadata": {"session_id": "session-from-metadata"}} + session_id = get_session_id_from_request_data(data) + assert session_id == "session-from-metadata" + + def test_get_session_id_from_litellm_metadata(self): + data = {"litellm_metadata": {"session_id": "session-from-litellm-metadata"}} + session_id = get_session_id_from_request_data(data) + assert session_id == "session-from-litellm-metadata" + + def test_get_session_id_from_litellm_session_id(self): + data = {"litellm_session_id": "session-direct"} + session_id = get_session_id_from_request_data(data) + assert session_id == "session-direct" + + @pytest.mark.asyncio + async def test_pre_call_hook_no_session(self, handler, user_api_key_dict): + data = {"model": "gpt-4"} + result = await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data, + call_type="completion", + ) + assert result is None + assert data["model"] == "gpt-4" + + @pytest.mark.asyncio + async def test_pre_call_hook_with_routing_override( + self, handler, user_api_key_dict + ): + await handler.set_session_routing( + session_id="routed-session", + model="on-premise-model", + user_api_key_dict=user_api_key_dict, + ) + + data = { + "model": "gpt-4", + "metadata": {"session_id": "routed-session"}, + } + result = await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data, + call_type="completion", + ) + + assert result is not None + assert result["model"] == "on-premise-model" + assert result["metadata"]["sensitive_data_routing_applied"] is True + assert result["metadata"]["sensitive_data_routing_original_model"] == "gpt-4" + + @pytest.mark.asyncio + async def test_pre_call_hook_no_override_needed(self, handler, user_api_key_dict): + data = { + "model": "gpt-4", + "metadata": {"session_id": "no-override-session"}, + } + result = await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data, + call_type="completion", + ) + assert result is None + assert data["model"] == "gpt-4" + + +class TestSensitiveDataRouteException: + def test_exception_creation(self): + exc = SensitiveDataRouteException( + route_to_model="on-premise-model", + session_id="test-session", + guardrail_name="test-guardrail", + detection_info={"detected_entities": ["SSN", "CREDIT_CARD"]}, + ) + + assert exc.route_to_model == "on-premise-model" + assert exc.session_id == "test-session" + assert exc.guardrail_name == "test-guardrail" + assert "SSN" in exc.detection_info["detected_entities"] + + +class TestCustomGuardrailSensitiveDataRouting: + def test_should_route_on_sensitive_data_false_by_default(self): + guardrail = CustomGuardrail(guardrail_name="test") + assert guardrail.should_route_on_sensitive_data() is False + + def test_should_route_on_sensitive_data_true(self): + guardrail = CustomGuardrail( + guardrail_name="test", + on_sensitive_data="route", + sensitive_data_route_to_model="on-premise-model", + ) + assert guardrail.should_route_on_sensitive_data() is True + + def test_should_route_on_sensitive_data_missing_model(self): + guardrail = CustomGuardrail( + guardrail_name="test", + on_sensitive_data="route", + ) + assert guardrail.should_route_on_sensitive_data() is False + + def test_raise_sensitive_data_route_exception(self): + guardrail = CustomGuardrail( + guardrail_name="test", + on_sensitive_data="route", + sensitive_data_route_to_model="on-premise-model", + ) + + request_data = {"model": "gpt-4", "metadata": {"session_id": "test-session"}} + + with pytest.raises(SensitiveDataRouteException) as exc_info: + guardrail.raise_sensitive_data_route_exception( + route_to_model="on-premise-model", + request_data=request_data, + detection_info={"type": "PII"}, + ) + + assert exc_info.value.route_to_model == "on-premise-model" + assert exc_info.value.session_id == "test-session" + + def test_raise_exception_carries_sticky_flag_false(self): + guardrail = CustomGuardrail( + guardrail_name="test", + on_sensitive_data="route", + sensitive_data_route_to_model="on-premise-model", + sticky_session_routing=False, + ) + + request_data = {"metadata": {"session_id": "test-session"}} + + with pytest.raises(SensitiveDataRouteException) as exc_info: + guardrail.raise_sensitive_data_route_exception( + route_to_model="on-premise-model", + request_data=request_data, + ) + + assert exc_info.value.sticky_session_routing is False + + def test_raise_exception_carries_sticky_flag_default_true(self): + guardrail = CustomGuardrail( + guardrail_name="test", + on_sensitive_data="route", + sensitive_data_route_to_model="on-premise-model", + ) + + request_data = {"metadata": {"session_id": "test-session"}} + + with pytest.raises(SensitiveDataRouteException) as exc_info: + guardrail.raise_sensitive_data_route_exception( + route_to_model="on-premise-model", + request_data=request_data, + ) + + assert exc_info.value.sticky_session_routing is True + + def test_raise_sensitive_data_route_exception_missing_session(self): + guardrail = CustomGuardrail(guardrail_name="test") + + request_data = {"model": "gpt-4"} + + with pytest.raises(ValueError) as exc_info: + guardrail.raise_sensitive_data_route_exception( + route_to_model="on-premise-model", + request_data=request_data, + ) + + assert "session_id" in str(exc_info.value) + + def test_handle_sensitive_data_detection_route(self): + guardrail = CustomGuardrail( + guardrail_name="test", + on_sensitive_data="route", + sensitive_data_route_to_model="on-premise-model", + ) + + request_data = {"model": "gpt-4", "metadata": {"session_id": "test-session"}} + + with pytest.raises(SensitiveDataRouteException) as exc_info: + guardrail.handle_sensitive_data_detection( + request_data=request_data, + detection_info={"type": "PII"}, + ) + + assert exc_info.value.route_to_model == "on-premise-model" + + def test_handle_sensitive_data_detection_block(self): + from litellm.exceptions import GuardrailRaisedException + + guardrail = CustomGuardrail(guardrail_name="test") + + request_data = {"model": "gpt-4", "metadata": {"session_id": "test-session"}} + + with pytest.raises(GuardrailRaisedException): + guardrail.handle_sensitive_data_detection( + request_data=request_data, + ) + + def test_handle_sensitive_data_detection_route_no_session_falls_back_to_block(self): + from litellm.exceptions import GuardrailRaisedException + + guardrail = CustomGuardrail( + guardrail_name="test", + on_sensitive_data="route", + sensitive_data_route_to_model="on-premise-model", + ) + + request_data = {"model": "gpt-4"} + + with pytest.raises(GuardrailRaisedException) as exc_info: + guardrail.handle_sensitive_data_detection( + request_data=request_data, + detection_info={"type": "PII"}, + ) + + assert "session_id" in str(exc_info.value) + + +class TestStickySessionRouting: + @pytest.fixture + def handler(self): + cache = MockInternalUsageCache() + return _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + + @pytest.fixture + def user_api_key_dict(self): + return UserAPIKeyAuth(api_key="test-key") + + @pytest.mark.asyncio + async def test_sticky_routing_persists(self, handler, user_api_key_dict): + session_id = "sticky-session" + await handler.set_session_routing( + session_id=session_id, + model="on-premise-model", + user_api_key_dict=user_api_key_dict, + ) + + for i in range(5): + data = { + "model": f"gpt-{i}", + "metadata": {"session_id": session_id}, + } + result = await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data, + call_type="completion", + ) + + assert result is not None + assert result["model"] == "on-premise-model" + assert ( + result["metadata"]["sensitive_data_routing_original_model"] + == f"gpt-{i}" + ) + + @pytest.mark.asyncio + async def test_different_sessions_independent(self, handler, user_api_key_dict): + await handler.set_session_routing( + session_id="session-a", + model="on-premise-model-a", + user_api_key_dict=user_api_key_dict, + ) + await handler.set_session_routing( + session_id="session-b", + model="on-premise-model-b", + user_api_key_dict=user_api_key_dict, + ) + + data_a = {"model": "gpt-4", "metadata": {"session_id": "session-a"}} + data_b = {"model": "gpt-4", "metadata": {"session_id": "session-b"}} + data_c = {"model": "gpt-4", "metadata": {"session_id": "session-c"}} + + result_a = await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data_a, + call_type="completion", + ) + result_b = await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data_b, + call_type="completion", + ) + result_c = await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data_c, + call_type="completion", + ) + + assert result_a["model"] == "on-premise-model-a" + assert result_b["model"] == "on-premise-model-b" + assert result_c is None + + @pytest.mark.asyncio + async def test_routing_is_isolated_per_api_key(self, handler): + shared_session = "shared-session-id" + await handler.set_session_routing( + session_id=shared_session, + model="on-premise-model", + user_api_key_dict=UserAPIKeyAuth(api_key="tenant-a"), + ) + + data_for_tenant_b = { + "model": "gpt-4", + "metadata": {"session_id": shared_session}, + } + result = await handler.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="tenant-b"), + cache=DualCache(), + data=data_for_tenant_b, + call_type="completion", + ) + assert result is None + assert data_for_tenant_b["model"] == "gpt-4" + + data_for_tenant_a = { + "model": "gpt-4", + "metadata": {"session_id": shared_session}, + } + result = await handler.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="tenant-a"), + cache=DualCache(), + data=data_for_tenant_a, + call_type="completion", + ) + assert result is not None + assert result["model"] == "on-premise-model" + + +class TestCacheKeyAndTTL: + def test_cache_prefix_constant(self): + assert SENSITIVE_ROUTING_CACHE_PREFIX == "sensitive_route" + + def test_default_ttl_constant(self): + assert DEFAULT_SENSITIVE_ROUTING_TTL == 3600 + + def test_make_cache_key_format(self): + cache = MockInternalUsageCache() + handler = _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + key = handler._make_cache_key("test-session-123", "hashed-key") + assert key == "{sensitive_route:hashed-key:test-session-123}:model" + + def test_make_cache_key_is_tenant_scoped(self): + cache = MockInternalUsageCache() + handler = _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + key_a = handler._make_cache_key("shared-session", "key-a") + key_b = handler._make_cache_key("shared-session", "key-b") + assert key_a != key_b + + def test_resolve_tenant_prefers_api_key(self): + tenant = _PROXY_SensitiveDataRoutingHandler._resolve_tenant( + UserAPIKeyAuth(api_key="hashed-key", user_id="alice") + ) + assert tenant == "hashed-key" + + def test_resolve_tenant_falls_back_to_jwt_principal(self): + tenant = _PROXY_SensitiveDataRoutingHandler._resolve_tenant( + UserAPIKeyAuth(api_key=None, user_id="alice", team_id="t1", org_id="o1") + ) + assert tenant == "user:alice|team:t1|org:o1" + + def test_resolve_tenant_distinguishes_keyless_principals(self): + tenant_a = _PROXY_SensitiveDataRoutingHandler._resolve_tenant( + UserAPIKeyAuth(api_key=None, user_id="alice") + ) + tenant_b = _PROXY_SensitiveDataRoutingHandler._resolve_tenant( + UserAPIKeyAuth(api_key=None, user_id="bob") + ) + assert tenant_a != tenant_b + + def test_resolve_tenant_defaults_when_anonymous(self): + assert _PROXY_SensitiveDataRoutingHandler._resolve_tenant(None) == "default" + assert ( + _PROXY_SensitiveDataRoutingHandler._resolve_tenant( + UserAPIKeyAuth(api_key=None) + ) + == "default" + ) + + +class TestCustomGuardrailSessionIdExtraction: + def test_get_session_id_from_litellm_session_id(self): + guardrail = CustomGuardrail(guardrail_name="test") + request_data = {"litellm_session_id": "session-direct-123"} + session_id = guardrail._get_session_id_from_request_data(request_data) + assert session_id == "session-direct-123" + + def test_get_session_id_from_metadata(self): + guardrail = CustomGuardrail(guardrail_name="test") + request_data = {"metadata": {"session_id": "session-metadata-456"}} + session_id = guardrail._get_session_id_from_request_data(request_data) + assert session_id == "session-metadata-456" + + def test_get_session_id_from_litellm_metadata(self): + guardrail = CustomGuardrail(guardrail_name="test") + request_data = {"litellm_metadata": {"session_id": "session-litellm-meta-789"}} + session_id = guardrail._get_session_id_from_request_data(request_data) + assert session_id == "session-litellm-meta-789" + + def test_get_session_id_returns_none_when_missing(self): + guardrail = CustomGuardrail(guardrail_name="test") + request_data = {"model": "gpt-4"} + session_id = guardrail._get_session_id_from_request_data(request_data) + assert session_id is None + + def test_get_session_id_priority_litellm_session_id_first(self): + guardrail = CustomGuardrail(guardrail_name="test") + request_data = { + "litellm_session_id": "priority-session", + "metadata": {"session_id": "should-not-use"}, + "litellm_metadata": {"session_id": "also-not-this"}, + } + session_id = guardrail._get_session_id_from_request_data(request_data) + assert session_id == "priority-session" + + def test_get_session_id_converts_to_string(self): + guardrail = CustomGuardrail(guardrail_name="test") + request_data = {"litellm_session_id": 12345} + session_id = guardrail._get_session_id_from_request_data(request_data) + assert session_id == "12345" + assert isinstance(session_id, str) + + +class TestCustomGuardrailInit: + def test_init_with_routing_config(self): + guardrail = CustomGuardrail( + guardrail_name="test-guardrail", + on_sensitive_data="route", + sensitive_data_route_to_model="on-premise-model", + sticky_session_routing=True, + ) + assert guardrail.on_sensitive_data == "route" + assert guardrail.sensitive_data_route_to_model == "on-premise-model" + assert guardrail.sticky_session_routing is True + + def test_init_default_values(self): + guardrail = CustomGuardrail(guardrail_name="test") + assert guardrail.on_sensitive_data is None + assert guardrail.sensitive_data_route_to_model is None + assert guardrail.sticky_session_routing is True + + +class TestSensitiveDataRouteExceptionStr: + def test_exception_str_representation(self): + exc = SensitiveDataRouteException( + route_to_model="on-premise-model", + session_id="test-session", + guardrail_name="pii-detector", + ) + assert ( + str(exc) + == "Sensitive data detected by pii-detector. Routing to model: on-premise-model" + ) + + def test_exception_custom_message(self): + exc = SensitiveDataRouteException( + route_to_model="on-premise-model", + session_id="test-session", + guardrail_name="pii-detector", + message="Custom error message", + ) + assert str(exc) == "Custom error message" + assert exc.message == "Custom error message" + + +class TestRedisCache: + @pytest.fixture + def handler_with_redis(self): + cache = MockInternalUsageCache() + mock_redis = AsyncMock() + cache.dual_cache.redis_cache = mock_redis + return _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + + @pytest.mark.asyncio + async def test_get_routed_model_from_redis(self, handler_with_redis): + handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_cache = AsyncMock( + return_value="redis-model" + ) + result = await handler_with_redis._get_routed_model( + "session-123", UserAPIKeyAuth(api_key="hashed-key") + ) + assert result == "redis-model" + + @pytest.mark.asyncio + async def test_get_routed_model_backfills_in_memory_after_redis_hit( + self, handler_with_redis + ): + cache_key = "{sensitive_route:hashed-key:session-123}:model" + key = UserAPIKeyAuth(api_key="hashed-key") + handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_cache = AsyncMock( + return_value="on-premise-model" + ) + handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_ttl = ( + AsyncMock(return_value=120) + ) + + first = await handler_with_redis._get_routed_model("session-123", key) + assert first == "on-premise-model" + assert handler_with_redis.internal_usage_cache._cache[cache_key] == ( + "on-premise-model" + ) + + handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_cache = AsyncMock( + side_effect=Exception("Redis went down") + ) + second = await handler_with_redis._get_routed_model("session-123", key) + assert second == "on-premise-model" + + @pytest.mark.asyncio + async def test_backfill_uses_remaining_redis_ttl(self, handler_with_redis): + cache_key = "{sensitive_route:hashed-key:session-123}:model" + key = UserAPIKeyAuth(api_key="hashed-key") + handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_cache = AsyncMock( + return_value="on-premise-model" + ) + handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_ttl = ( + AsyncMock(return_value=42) + ) + + await handler_with_redis._get_routed_model("session-123", key) + + assert handler_with_redis.internal_usage_cache._ttls[cache_key] == 42 + + @pytest.mark.asyncio + async def test_backfill_falls_back_to_full_ttl_when_redis_ttl_missing( + self, handler_with_redis + ): + cache_key = "{sensitive_route:hashed-key:session-123}:model" + key = UserAPIKeyAuth(api_key="hashed-key") + handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_cache = AsyncMock( + return_value="on-premise-model" + ) + handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_ttl = ( + AsyncMock(return_value=None) + ) + + await handler_with_redis._get_routed_model("session-123", key) + + assert ( + handler_with_redis.internal_usage_cache._ttls[cache_key] + == handler_with_redis.ttl + ) + + @pytest.mark.asyncio + async def test_get_routed_model_redis_fallback_on_error(self, handler_with_redis): + handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_get_cache = AsyncMock( + side_effect=Exception("Redis connection error") + ) + handler_with_redis.internal_usage_cache._cache[ + "{sensitive_route:hashed-key:session-123}:model" + ] = "fallback-model" + result = await handler_with_redis._get_routed_model( + "session-123", UserAPIKeyAuth(api_key="hashed-key") + ) + assert result == "fallback-model" + + @pytest.mark.asyncio + async def test_set_session_routing_with_redis(self, handler_with_redis): + handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_set_cache = ( + AsyncMock() + ) + await handler_with_redis.set_session_routing( + session_id="session-456", + model="on-premise-model", + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + guardrail_name="test-guardrail", + ) + handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_set_cache.assert_called_once() + + @pytest.mark.asyncio + async def test_set_session_routing_redis_fallback_on_error( + self, handler_with_redis + ): + handler_with_redis.internal_usage_cache.dual_cache.redis_cache.async_set_cache = AsyncMock( + side_effect=Exception("Redis connection error") + ) + await handler_with_redis.set_session_routing( + session_id="session-789", + model="on-premise-model", + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + ) + cache_key = "{sensitive_route:hashed-key:session-789}:model" + assert ( + handler_with_redis.internal_usage_cache._cache[cache_key] + == "on-premise-model" + ) + + +class TestPreCallHookEdgeCases: + @pytest.fixture + def handler(self): + cache = MockInternalUsageCache() + return _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + + @pytest.fixture + def user_api_key_dict(self): + return UserAPIKeyAuth(api_key="test-key") + + @pytest.mark.asyncio + async def test_pre_call_hook_same_model_no_change(self, handler, user_api_key_dict): + await handler.set_session_routing( + session_id="same-model-session", + model="gpt-4", + user_api_key_dict=user_api_key_dict, + ) + data = { + "model": "gpt-4", + "metadata": {"session_id": "same-model-session"}, + } + result = await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data, + call_type="completion", + ) + assert result is None + + +class TestHandleSensitiveDataDetectionWithRouting: + def test_handle_sensitive_data_detection_full_flow(self): + guardrail = CustomGuardrail( + guardrail_name="pii-guardrail", + on_sensitive_data="route", + sensitive_data_route_to_model="on-premise-model", + ) + + request_data = { + "model": "gpt-4", + "metadata": {"session_id": "flow-test-session"}, + "messages": [{"role": "user", "content": "My SSN is 123-45-6789"}], + } + + with pytest.raises(SensitiveDataRouteException) as exc_info: + guardrail.handle_sensitive_data_detection( + request_data=request_data, + detection_info={"detected_entities": ["SSN"]}, + ) + + exc = exc_info.value + assert exc.route_to_model == "on-premise-model" + assert exc.session_id == "flow-test-session" + assert exc.guardrail_name == "pii-guardrail" + assert exc.detection_info == {"detected_entities": ["SSN"]} + + +class TestProxyHandleSensitiveDataRouteException: + @pytest.fixture + def proxy_logging(self): + from litellm.proxy.utils import ProxyLogging + + return ProxyLogging(user_api_key_cache=DualCache()) + + @pytest.fixture + def routing_hook(self): + cache = MockInternalUsageCache() + return _PROXY_SensitiveDataRoutingHandler(internal_usage_cache=cache) + + @pytest.mark.asyncio + async def test_sticky_routing_persists_override(self, proxy_logging, routing_hook): + proxy_logging.proxy_hook_mapping["sensitive_data_routing"] = routing_hook + exc = SensitiveDataRouteException( + route_to_model="on-premise-model", + session_id="sess-sticky", + guardrail_name="pii", + sticky_session_routing=True, + ) + data = {"model": "gpt-4", "metadata": {"session_id": "sess-sticky"}} + + result = await proxy_logging._handle_sensitive_data_route_exception( + exc, data, UserAPIKeyAuth(api_key="tenant-a") + ) + + assert result["model"] == "on-premise-model" + assert ( + await routing_hook._get_routed_model( + "sess-sticky", UserAPIKeyAuth(api_key="tenant-a") + ) + == "on-premise-model" + ) + + @pytest.mark.asyncio + async def test_non_sticky_routing_does_not_persist_override( + self, proxy_logging, routing_hook + ): + proxy_logging.proxy_hook_mapping["sensitive_data_routing"] = routing_hook + exc = SensitiveDataRouteException( + route_to_model="on-premise-model", + session_id="sess-non-sticky", + guardrail_name="pii", + sticky_session_routing=False, + ) + data = {"model": "gpt-4", "metadata": {"session_id": "sess-non-sticky"}} + + result = await proxy_logging._handle_sensitive_data_route_exception( + exc, data, UserAPIKeyAuth(api_key="tenant-a") + ) + + assert result["model"] == "on-premise-model" + assert ( + await routing_hook._get_routed_model( + "sess-non-sticky", UserAPIKeyAuth(api_key="tenant-a") + ) + is None + ) + + @pytest.mark.asyncio + async def test_sticky_routing_handles_none_user_api_key_dict( + self, proxy_logging, routing_hook + ): + proxy_logging.proxy_hook_mapping["sensitive_data_routing"] = routing_hook + exc = SensitiveDataRouteException( + route_to_model="on-premise-model", + session_id="sess-no-key", + guardrail_name="pii", + sticky_session_routing=True, + ) + data = {"model": "gpt-4", "metadata": {"session_id": "sess-no-key"}} + + result = await proxy_logging._handle_sensitive_data_route_exception( + exc, data, None + ) + + assert result["model"] == "on-premise-model" + assert ( + await routing_hook._get_routed_model("sess-no-key", None) + == "on-premise-model" + ) + + @pytest.mark.asyncio + async def test_sticky_routing_scopes_jwt_users_by_principal( + self, proxy_logging, routing_hook + ): + proxy_logging.proxy_hook_mapping["sensitive_data_routing"] = routing_hook + exc = SensitiveDataRouteException( + route_to_model="on-premise-model", + session_id="shared-jwt-session", + guardrail_name="pii", + sticky_session_routing=True, + ) + attacker = UserAPIKeyAuth(api_key=None, user_id="attacker", team_id="team-x") + await proxy_logging._handle_sensitive_data_route_exception( + exc, + {"model": "gpt-4", "metadata": {"session_id": "shared-jwt-session"}}, + attacker, + ) + + victim = UserAPIKeyAuth(api_key=None, user_id="victim", team_id="team-y") + victim_data = { + "model": "gpt-4", + "metadata": {"session_id": "shared-jwt-session"}, + } + result = await routing_hook.async_pre_call_hook( + user_api_key_dict=victim, + cache=DualCache(), + data=victim_data, + call_type="completion", + ) + assert result is None + assert victim_data["model"] == "gpt-4" + + attacker_data = { + "model": "gpt-4", + "metadata": {"session_id": "shared-jwt-session"}, + } + result = await routing_hook.async_pre_call_hook( + user_api_key_dict=attacker, + cache=DualCache(), + data=attacker_data, + call_type="completion", + ) + assert result is not None + assert result["model"] == "on-premise-model" + + @pytest.mark.asyncio + async def test_sticky_routing_warns_when_hook_not_registered(self, proxy_logging): + exc = SensitiveDataRouteException( + route_to_model="on-premise-model", + session_id="sess-no-hook", + sticky_session_routing=True, + ) + data = {"model": "gpt-4", "metadata": {"session_id": "sess-no-hook"}} + + with patch("litellm.proxy.utils.verbose_proxy_logger.warning") as mock_warning: + result = await proxy_logging._handle_sensitive_data_route_exception( + exc, data, UserAPIKeyAuth(api_key="tenant-a") + ) + + assert result["model"] == "on-premise-model" + mock_warning.assert_called_once() + + +class _RoutingGuardrail(CustomGuardrail): + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + self.handle_sensitive_data_detection(request_data=data) + + +class _RecordingGuardrail(CustomGuardrail): + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.ran = False + + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + self.ran = True + return None + + +class _BlockingGuardrail(CustomGuardrail): + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.ran = False + + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + from litellm.exceptions import GuardrailRaisedException + + self.ran = True + raise GuardrailRaisedException( + message="blocked", guardrail_name=self.guardrail_name + ) + + +class TestPreCallHookDeferredRouting: + """Guardrails after the one that triggers routing must still run.""" + + @pytest.fixture + def proxy_logging(self): + from litellm.proxy.utils import ProxyLogging + + return ProxyLogging(user_api_key_cache=DualCache()) + + @pytest.fixture(autouse=True) + def restore_callbacks(self): + import litellm + + original = litellm.callbacks + litellm.callbacks = [] + yield + litellm.callbacks = original + + @pytest.mark.asyncio + async def test_later_guardrail_runs_and_routing_applied(self, proxy_logging): + import litellm + + router = _RoutingGuardrail( + guardrail_name="router", + default_on=True, + event_hook="pre_call", + on_sensitive_data="route", + sensitive_data_route_to_model="on-prem-model", + sticky_session_routing=False, + ) + recorder = _RecordingGuardrail( + guardrail_name="recorder", + default_on=True, + event_hook="pre_call", + ) + litellm.callbacks = [router, recorder] + + data = {"model": "gpt-4", "metadata": {"session_id": "sess-defer"}} + result = await proxy_logging.pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="tenant-a"), + data=data, + call_type="completion", + ) + + assert recorder.ran is True + assert result["model"] == "on-prem-model" + assert result["metadata"]["sensitive_data_routing_applied"] is True + + @pytest.mark.asyncio + async def test_later_blocking_guardrail_overrides_routing(self, proxy_logging): + import litellm + from litellm.exceptions import GuardrailRaisedException + + router = _RoutingGuardrail( + guardrail_name="router", + default_on=True, + event_hook="pre_call", + on_sensitive_data="route", + sensitive_data_route_to_model="on-prem-model", + sticky_session_routing=False, + ) + blocker = _BlockingGuardrail( + guardrail_name="blocker", + default_on=True, + event_hook="pre_call", + ) + litellm.callbacks = [router, blocker] + + data = {"model": "gpt-4", "metadata": {"session_id": "sess-block"}} + with pytest.raises(GuardrailRaisedException): + await proxy_logging.pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="tenant-a"), + data=data, + call_type="completion", + ) + + assert blocker.ran is True + + @pytest.mark.asyncio + async def test_routing_guardrail_records_service_span(self, proxy_logging): + import litellm + from litellm.types.services import ServiceTypes + + class _SlowRoutingGuardrail(CustomGuardrail): + async def async_pre_call_hook( + self, user_api_key_dict, cache, data, call_type + ): + await asyncio.sleep(0.02) + self.handle_sensitive_data_detection(request_data=data) + + router = _SlowRoutingGuardrail( + guardrail_name="router", + default_on=True, + event_hook="pre_call", + on_sensitive_data="route", + sensitive_data_route_to_model="on-prem-model", + sticky_session_routing=False, + ) + litellm.callbacks = [router] + + recorded = AsyncMock() + proxy_logging.service_logging_obj.async_service_success_hook = recorded + + data = {"model": "gpt-4", "metadata": {"session_id": "sess-span"}} + result = await proxy_logging.pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="tenant-a"), + data=data, + call_type="completion", + ) + + assert result["model"] == "on-prem-model" + recorded.assert_called_once() + assert recorded.call_args.kwargs["call_type"] == "_SlowRoutingGuardrail" + assert recorded.call_args.kwargs["service"] == ServiceTypes.PROXY_PRE_CALL + + @pytest.mark.asyncio + async def test_routing_recorded_as_intervention_not_prometheus_error( + self, proxy_logging + ): + import litellm + from litellm.integrations.prometheus import PrometheusLogger + + router = _RoutingGuardrail( + guardrail_name="router", + default_on=True, + event_hook="pre_call", + on_sensitive_data="route", + sensitive_data_route_to_model="on-prem-model", + sticky_session_routing=False, + ) + prom = MagicMock(spec=PrometheusLogger) + litellm.callbacks = [router, prom] + + data = {"model": "gpt-4", "metadata": {"session_id": "sess-prom"}} + result = await proxy_logging.pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="tenant-a"), + data=data, + call_type="completion", + ) + + assert result["model"] == "on-prem-model" + prom._record_guardrail_metrics.assert_called_once() + metrics_kwargs = prom._record_guardrail_metrics.call_args.kwargs + assert metrics_kwargs["status"] == "intervened" + assert metrics_kwargs["error_type"] is None diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index daa3033b231..e33b4b3d03d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -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 + ) diff --git a/tests/test_litellm/responses/test_responses_utils.py b/tests/test_litellm/responses/test_responses_utils.py index 60b84f0e0a8..bd441321507 100644 --- a/tests/test_litellm/responses/test_responses_utils.py +++ b/tests/test_litellm/responses/test_responses_utils.py @@ -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: diff --git a/tests/test_litellm/test_claude_opus_4_8_config.py b/tests/test_litellm/test_claude_opus_4_8_config.py index 0ea4026e165..32f7d249e05 100644 --- a/tests/test_litellm/test_claude_opus_4_8_config.py +++ b/tests/test_litellm/test_claude_opus_4_8_config.py @@ -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}" diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 3d45a3409d8..82a4a60bf82 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -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 = [ diff --git a/tests/test_litellm/types/test_guardrails_case_normalization.py b/tests/test_litellm/types/test_guardrails_case_normalization.py index e1e03fe6b88..3e7a573ea8e 100644 --- a/tests/test_litellm/types/test_guardrails_case_normalization.py +++ b/tests/test_litellm/types/test_guardrails_case_normalization.py @@ -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"