From 812a2217ca6420f7ae36a140917b15a4a4ccbbb8 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 4 Jun 2026 22:22:28 -0700 Subject: [PATCH 1/8] [internal copy of #29511] feat(guardrails): add sensitive data routing to on-premise models (#29531) * feat(guardrails): add sensitive data routing to on-premise models When a guardrail detects sensitive data, route to an on-premise model instead of blocking or redacting. All subsequent requests in that session continue routing to the same model (sticky routing). New config options for guardrails: - on_sensitive_data: 'block' (default) or 'route' - sensitive_data_route_to_model: target model for rerouting - sticky_session_routing: persist routing for session (default: true) New exception SensitiveDataRouteException triggers rerouting when raised by guardrails. The proxy catches it, stores the routing decision in cache, and modifies the request's model field. New hook _PROXY_SensitiveDataRoutingHandler checks incoming requests against cached routing decisions and applies sticky routing. https://claude.ai/code/session_01SQd4isBa3UyouRoGVou9dK * fix: black formatting for custom_guardrail.py https://claude.ai/code/session_01SQd4isBa3UyouRoGVou9dK * test: improve test coverage for sensitive data routing feature Add additional tests for: - Cache key format and TTL constants - Session ID extraction from multiple locations - Custom guardrail initialization with routing config - Exception string representation and custom messages - Redis cache paths including fallback behavior - Edge cases in pre-call hook https://claude.ai/code/session_01SQd4isBa3UyouRoGVou9dK * fix: use correct GuardrailRaisedException parameters Replace invalid 'source' parameter with 'guardrail_name' to match the exception's actual signature. https://claude.ai/code/session_01SQd4isBa3UyouRoGVou9dK * test: move sensitive data routing tests to hooks directory Move test file to align with source code structure. https://claude.ai/code/session_01SQd4isBa3UyouRoGVou9dK * fix(guardrails): honor sticky_session_routing flag and scope session routing per API key Propagate sticky_session_routing through SensitiveDataRouteException so a guardrail configured with sticky_session_routing=False reroutes only the triggering request without persisting a session override. Scope the routing cache key to the requesting API key so sessions from different tenants cannot collide, and warn when sticky routing is requested but the hook is not registered. * refactor(guardrails): dedupe session-id extraction and drop redundant import Extract the shared session-id lookup into get_session_id_from_request_data so the sensitive-data routing hook and CustomGuardrail no longer keep two identical copies of the logic. Remove the redundant local import of GuardrailRaisedException in handle_sensitive_data_detection, and document that detection_info is surfaced in request metadata and logs so it must not carry raw sensitive values. * fix(guardrails): guard None user_api_key_dict in sensitive data route handler * fix(responses): send application/json Content-Type on responses DELETE OpenAI's responses DELETE endpoint now rejects requests that arrive without a Content-Type header, defaulting them to application/octet-stream and returning 'Unsupported content type: application/octet-stream'. The delete handler sent no body and therefore no Content-Type, so the request failed. Declare application/json on the delete request, matching the OpenAI SDK. * fix(guardrails): backfill in-memory cache after redis hit in sensitive data routing When _get_routed_model resolves a routing override from Redis it now also populates the local in-memory cache. Without the write-back, a non-writing instance that only ever reads from Redis would lose the sticky routing decision the moment Redis became unavailable, silently reverting sensitive sessions to the default model. * fix(guardrails): scope sticky sensitive-data routing to JWT principal Keyless auth (JWT and similar) has no api_key, so every such caller shared the "default" cache namespace. One authenticated user could reuse another user's session_id, trip the guardrail, and silently force the other user's subsequent requests onto the cached on-prem model for the TTL. Resolve the routing tenant from the api_key when present, otherwise from a stable principal built from the user/team/org identity, before reading or writing the session route. * fix(guardrails): require route target model when on_sensitive_data='route' * fix(guardrails): mark user_api_key_dict Optional in sensitive-data route handler * fix(guardrails): use remaining redis ttl for local backfill and str env default * fix(guardrails): graceful block when routing configured but no session_id handle_sensitive_data_detection promised to raise only SensitiveDataRouteException or GuardrailRaisedException, but when routing was configured and the request had no session_id it let a ValueError from raise_sensitive_data_route_exception propagate, surfacing as an HTTP 500 instead of a block. Fall back to a graceful block in that case so the documented contract holds. * fix(guardrails): run remaining guardrails after sensitive-data reroute Defer the SensitiveDataRouteException until every guardrail in the pre-call loop has run, so downstream security guardrails are no longer skipped when an earlier guardrail triggers routing. The first reroute wins and a later guardrail that blocks still propagates. Also normalize on_sensitive_data to lowercase like sibling on_* config fields so case-insensitive values are accepted. * fix(guardrails): classify sensitive-data reroute as guardrail intervention * fix(guardrails): record sensitive-data reroute as prometheus intervention not error * fix(guardrails): record service span for routing guardrail and move case-normalizer to base params Drop the early continue so a guardrail that signals sensitive-data routing still emits its PROXY_PRE_CALL service span like every other callback. Move the lowercase normalizer onto BaseLitellmParams so on_sensitive_data is normalized consistently when BaseLitellmParams is constructed directly, matching the cross-field route->model validator that already lives on the base. --- litellm/exceptions.py | 34 + litellm/integrations/custom_guardrail.py | 143 ++- litellm/integrations/prometheus.py | 2 +- litellm/llms/custom_httpx/llm_http_handler.py | 4 + litellm/proxy/hooks/__init__.py | 2 + litellm/proxy/hooks/sensitive_data_routing.py | 206 ++++ litellm/proxy/utils.py | 155 ++- litellm/types/guardrails.py | 71 +- .../integrations/test_custom_guardrail.py | 44 + .../custom_httpx/test_llm_http_handler.py | 40 + .../hooks/test_sensitive_data_routing.py | 1036 +++++++++++++++++ .../test_guardrails_case_normalization.py | 67 +- 12 files changed, 1745 insertions(+), 59 deletions(-) create mode 100644 litellm/proxy/hooks/sensitive_data_routing.py create mode 100644 tests/test_litellm/proxy/hooks/test_sensitive_data_routing.py 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/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/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/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/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/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/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/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/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/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" From df704d9016cd55ee11435a39c9ca465e4bf704ad Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 4 Jun 2026 22:46:08 -0700 Subject: [PATCH 2/8] fix(proxy/hooks): populate llm_provider on internal rate-limit errors (#27707) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat(proxy/hooks): add ProxyHTTPRateLimitError + provider resolver Introduces a small helper layer used by every proxy-side rate-limit hook so that the 429 they raise carries a populated llm_provider / model — instead of an empty exception.llm_provider that downstream loggers (Prometheus failure metric, observability callbacks) read as 'no provider attribution'. ProxyHTTPRateLimitError inherits from both fastapi.HTTPException (so the proxy server still renders it as a 429) and litellm.exceptions.RateLimitError (so isinstance checks and PrometheusLogger._get_exception_class_name pick up llm_provider). We deliberately don't call RateLimitError.__init__ — it constructs an httpx.Response we don't need and would just add failure surface; attribute parity is what downstream consumers care about. resolve_llm_provider_for_rate_limit() wraps litellm.get_llm_provider defensively. Internal limiter hooks fire from async_pre_call_hook — well before get_llm_provider runs anywhere else in the request lifecycle — so we have to call it ourselves at raise time. If the model is missing or unparseable (alias, router-only model) we fall back to llm_provider='litellm_proxy' rather than letting a second exception leak out and break the request path. Co-authored-by: Mateo Wang * fix(proxy/hooks): populate llm_provider on parallel-request 429s Both v1 and v3 parallel-request limiters fired bare HTTPException(429) from inside async_pre_call_hook. The downstream Prometheus failure metric reads exception.llm_provider via _get_exception_class_name — the empty value showed up as exception_class='HTTPException' and left model_id='None' on the time series. Threads requested_model through every raise site in: * parallel_request_limiter.py: - check_key_in_limits (the per-key/per-model/per-user/per-team/ per-customer over-limit path) - raise_rate_limit_error (zero-limit + global_max_parallel_requests paths) — now takes an optional requested_model kwarg * parallel_request_limiter_v3.py: - _handle_rate_limit_error (the OVER_LIMIT translator), called from both the should_rate_limit pre-check and the TPM reservation path Resolved via resolve_llm_provider_for_rate_limit so unknown / missing models silently fall back to llm_provider='litellm_proxy' instead of breaking the request path with a second exception. Co-authored-by: Mateo Wang * fix(proxy/hooks): populate llm_provider on dynamic-rate-limit 429s Same plumbing change as the parallel limiters, applied to both dynamic_rate_limiter (v1) and dynamic_rate_limiter_v3: * v1: TPM-zero and RPM-zero paths in async_pre_call_hook now resolve data['model'] -> (model, llm_provider) once and pass it into both raises. * v3: All three raise sites in _check_rate_limits — the model_saturation_check enforced raise, the priority_model enforced raise, and the fail-closed unknown-descriptor branch — now attribute the 429 to the actual provider. Falls back to llm_provider='litellm_proxy' when the model can't be resolved. Co-authored-by: Mateo Wang * fix(proxy/hooks): populate llm_provider on batch-rate-limit 429s batch_rate_limiter._raise_rate_limit_error now takes a requested_model kwarg threaded from data['model'] in _check_and_increment_batch_counters. The batch-creation 429 is what gets raised when the input file's tokens/requests count would push the per-key TPM/RPM window over its limit. Co-authored-by: Mateo Wang * fix(proxy/hooks): populate llm_provider on budget/iterations 429s Final batch of internal raise sites — the user/session-budget and max-iterations hooks. Same pattern: resolve data['model'] once at raise time, attach to ProxyHTTPRateLimitError so Prometheus and observability callbacks can attribute the 429. Hooks updated: * max_budget_limiter (per-user max_budget exceeded) * max_iterations_limiter (per-session agent iteration cap) * max_budget_per_session_limiter (per-session dollar cap) All three fall back to llm_provider='litellm_proxy' when data['model'] is missing or unparseable. Drops the now-unused HTTPException import from each module. Co-authored-by: Mateo Wang * test(proxy/hooks): pin provider field on internal rate-limit 429s Regression coverage for the 'provider field missing' bug across every proxy-side rate-limit hook + the helper layer: * ProxyHTTPRateLimitError class shape (HTTPException + RateLimitError, dict-detail stringification, None-provider normalization). * resolve_llm_provider_for_rate_limit happy paths (gpt-4o-mini, anthropic/..., bedrock/...) plus all three fallback branches (None, '', unknown name) plus a 'get_llm_provider raises' case that asserts we swallow the secondary exception. * For each limiter (parallel v1/v3, dynamic v1/v3, batch, max_budget, max_iterations, max_budget_per_session): assert the raised exception is a RateLimitError carrying the resolved model + llm_provider, and a sibling test that asserts the fallback path returns 'litellm_proxy' without leaking a second exception. * Two PrometheusLogger._get_exception_class_name pins so the Prometheus failure metric label flips from 'HTTPException' to 'Openai.ProxyHTTPRateLimitError' (or 'Litellm_proxy.*' on fallback) — that's what dashboards consume. Co-authored-by: Mateo Wang * perf(proxy/hooks): defer provider resolution to over-limit branches * fix: use error_message in raise_rate_limit_error to avoid literal 'None' in detail * Consolidate rate_limiter_utils imports in dynamic_rate_limiter * fix(proxy): set num_retries/max_retries on ProxyHTTPRateLimitError ProxyHTTPRateLimitError inherits from RateLimitError but did not call RateLimitError.__init__, so num_retries/max_retries were never set. When Starlette's HTTPException lacks __str__, MRO falls through to RateLimitError.__str__, which unconditionally reads these attributes and raises AttributeError during logging/traceback formatting. Initialize them to None defensively. * fix(mypy): silence base-class status_code conflict on ProxyHTTPRateLimitError HTTPException declares 'status_code: int' while openai.RateLimitError (via APIStatusError) declares 'status_code: Literal[429] = 429'. Mypy flags the multi-base override as [misc] in CI lint. The runtime semantics are fine (we set self.status_code in __init__), so silence the class-level annotation conflict with a targeted ignore. Co-authored-by: Mateo Wang --------- Co-authored-by: Cursor Agent Co-authored-by: Mateo Wang --- litellm/proxy/hooks/batch_rate_limiter.py | 14 +- litellm/proxy/hooks/dynamic_rate_limiter.py | 23 +- .../proxy/hooks/dynamic_rate_limiter_v3.py | 19 +- litellm/proxy/hooks/max_budget_limiter.py | 14 +- .../hooks/max_budget_per_session_limiter.py | 13 +- litellm/proxy/hooks/max_iterations_limiter.py | 13 +- .../proxy/hooks/parallel_request_limiter.py | 39 +- .../hooks/parallel_request_limiter_v3.py | 16 +- litellm/proxy/hooks/rate_limiter_utils.py | 96 +- .../test_proxy_rate_limit_provider_field.py | 968 ++++++++++++++++++ 10 files changed, 1186 insertions(+), 29 deletions(-) create mode 100644 tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py 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/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"])) From 2b7c97bff6d4af0d3fd71a9820d33957cd330548 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 5 Jun 2026 11:27:16 +0530 Subject: [PATCH 3/8] fix(vertex/anthropic): handle namespace tools and strip client_metadata for codex compatibility (#29489) * fix(vertex/anthropic): handle namespace tools and strip client_metadata for codex compatibility * fix(anthropic): cast nested namespace tools to fix mypy error, skip nameless flat tools --- litellm/llms/anthropic/chat/transformation.py | 35 ++++++- .../test_anthropic_chat_transformation.py | 98 +++++++++++++++++++ 2 files changed, 132 insertions(+), 1 deletion(-) diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 4f15d1b3cef..d9224281db3 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -918,7 +918,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 +2010,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/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 From 778a7f752dfd4f11d014d1dcb15dc76a3a27d234 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 4 Jun 2026 23:03:37 -0700 Subject: [PATCH 4/8] Support OAuth M2M for Databricks Apps A2A agents (#29586) * Add OAuth M2M support for A2A agents targeting Databricks Apps Databricks App endpoints reject static bearer tokens and require a short-lived OAuth token minted via the workspace OIDC token endpoint. A2A agents could previously only authenticate outbound with static_headers or client header passthrough, so Databricks App agents could not be registered. Agents configured with a databricks_oauth block in litellm_params now mint and cache a client_credentials token and attach it as the outbound Authorization header on both message/send and message/stream calls, overriding any statically configured Authorization. * Add tests covering Databricks App OAuth token error paths Cover the HTTP status error, transport error, non-object JSON body, and invalid expires_in fallback branches in the token cache so the failure handling is locked in by regression tests. * Harden Databricks App OAuth token cache Cap the cache TTL at the token's own lifetime so a token whose validity is shorter than the refresh buffer is never cached and served stale; include a digest of client_secret in the cache key so a rotated secret mints a fresh token instead of reusing the old one; and prune the per-key lock when its cached token is evicted so the lock map stays bounded by the live key set. * Clear per-key locks on Databricks OAuth cache flush * fix(a2a/databricks): mint OAuth token via Basic auth header, not unsupported auth= kwarg litellm's AsyncHTTPHandler.post (what get_async_httpx_client returns) has no auth parameter, so minting a Databricks App OAuth token raised "AsyncHTTPHandler.post() got an unexpected keyword argument 'auth'" before any network call ever left the proxy, breaking the feature end to end. The handler also calls raise_for_status() internally and re-raises a MaskedHTTPStatusError (a subclass of httpx.HTTPStatusError), so the explicit raise_for_status() after post() was dead code. Build the HTTP Basic Authorization header by hand and pass it via headers, which is what the Databricks workspace OIDC token endpoint documents for client authentication. The token-cache tests now model the real handler contract with create_autospec so the rejected auth= signature is enforced; the previous mocks accepted any kwargs and silently hid the bug. Co-authored-by: Mateo Wang * Prune Databricks OAuth lock on the short-lived-token path When expires_in is below the refresh buffer the token is intentionally not cached, so _remove_key never runs for that key and the per-key lock created by _get_lock leaked permanently. Drop the lock in that branch so _locks stays bounded by the live key set, and assert the cleanup in the short-lived-token test * Gate A2A Databricks OAuth on the databricks_oauth block at the call site Make the gating explicit where the header is applied so it is clear that only agents configured with a databricks_oauth block enter the OAuth path; every other agent is left untouched. Add a regression test asserting a non-Databricks agent never invokes the token resolver. --------- Co-authored-by: Cursor Agent Co-authored-by: Mateo Wang --- .../proxy/agent_endpoints/a2a_endpoints.py | 15 + .../proxy/agent_endpoints/databricks_oauth.py | 250 +++++++++ .../agent_endpoints/test_agent_headers.py | 92 +++- .../agent_endpoints/test_databricks_oauth.py | 496 ++++++++++++++++++ 4 files changed, 852 insertions(+), 1 deletion(-) create mode 100644 litellm/proxy/agent_endpoints/databricks_oauth.py create mode 100644 tests/test_litellm/proxy/agent_endpoints/test_databricks_oauth.py 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/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"} From 8259d6cd85ef6e7ba4d0db5b1937646d6f82a5e8 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 4 Jun 2026 23:30:05 -0700 Subject: [PATCH 5/8] fix: small CLAUDE.md nit (#29749) --- CLAUDE.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 From 1c741b91c0f174ccbb9243ebacb4dccc64773e51 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 5 Jun 2026 03:49:01 -0700 Subject: [PATCH 6/8] fix(anthropic): route Claude Opus 4.8 through adaptive thinking (#29702) * fix(anthropic): route Claude Opus 4.8 through adaptive thinking Opus 4.8 uses the same adaptive thinking contract as 4.6/4.7 (thinking.type=adaptive plus output_config.effort), but _is_adaptive_thinking_model only recognized 4.6/4.7 by name and otherwise leaned on the supports_adaptive_thinking cost-map flag. The Bedrock, Vertex, and Azure 4.8 entries don't carry that flag, so a bedrock/us.anthropic.claude-opus-4-8 request fell back to the legacy thinking.type=enabled shape and Bedrock rejected it with "thinking.type.enabled is not supported for this model". Add _is_claude_4_8_model and wire it in next to the existing 4.6/4.7 matchers in the adaptive-thinking detection, the effort=max gate, and the supported-params check, so every provider path treats 4.8 as adaptive regardless of whether its cost-map entry advertises the flag. * refactor(anthropic): drive Opus 4.8 adaptive thinking from the cost map Replace the _is_claude_4_8_model name matcher with cost-map data. Add supports_adaptive_thinking to every Opus 4.8 provider variant (Bedrock regional/global, Vertex, Azure) in both the root and bundled cost maps, and move the prefix-resolving capability lookup (_supports_model_capability) down to AnthropicModelInfo so _is_adaptive_thinking_model reads the flag through the bedrock/invoke/, bedrock/, and vertex_ai/ prefixes. The 4.6/4.7 name checks stay as a fallback since their provider entries don't carry the flag yet. A pure data fix is not enough on its own: _supports_factory doesn't strip the us.anthropic./invoke/ prefixes, so bedrock/invoke/us.anthropic.claude-opus-4-8 would still miss the flag without the resolver change. Add a cost-map guardrail test asserting every claude-opus-4-8 variant carries the flag, so a future variant added without it fails CI instead of silently sending the legacy thinking.type=enabled shape that the provider rejects. --- litellm/llms/anthropic/chat/transformation.py | 45 ------------ litellm/llms/anthropic/common_utils.py | 52 ++++++++++++-- ...odel_prices_and_context_window_backup.json | 8 +++ model_prices_and_context_window.json | 8 +++ .../test_reasoning_effort_translation.py | 51 +++++++++++++ .../anthropic/test_anthropic_common_utils.py | 72 +++++++++++++++++++ .../test_claude_opus_4_8_config.py | 21 ++++++ 7 files changed, 208 insertions(+), 49 deletions(-) diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index d9224281db3..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.""" 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/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/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/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/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}" From 3f79222350854f422ca6a52bf55e14faf0a94409 Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Fri, 5 Jun 2026 15:22:52 +0200 Subject: [PATCH 7/8] fix(proxy): persist oauth2_flow on MCP server registration (#29690) --- .../migration.sql | 2 + .../litellm_proxy_extras/schema.prisma | 1 + litellm/proxy/_types.py | 2 + litellm/proxy/schema.prisma | 1 + schema.prisma | 1 + .../test_mcp_management_endpoints.py | 52 +++++++++++++++++++ 6 files changed, 59 insertions(+) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260604120000_add_oauth2_flow_to_mcp_servers/migration.sql 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 c4754ef6117..6715f80464d 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -322,6 +322,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/_types.py b/litellm/proxy/_types.py index 6c36c813e76..6de8f0a3cdf 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1376,6 +1376,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 @@ -1449,6 +1450,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/schema.prisma b/litellm/proxy/schema.prisma index c4754ef6117..6715f80464d 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -322,6 +322,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/schema.prisma b/schema.prisma index c4754ef6117..6715f80464d 100644 --- a/schema.prisma +++ b/schema.prisma @@ -322,6 +322,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/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index a5c8320a5cc..7e044442e2b 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 @@ -3317,3 +3317,55 @@ def test_sanitize_mcp_server_for_non_admin_clears_credential_fields(): # server without exposing secrets. assert sanitized.server_id == server.server_id assert sanitized.alias == server.alias + + +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 + ) From ffd0e9fa7f49e2f040900630eedb2d5abdcd8375 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 5 Jun 2026 06:23:17 -0700 Subject: [PATCH 8/8] [internal copy of #27491] fix(realtime): Fix Realtime Audio Token Cost Tracking (#29722) * Normalize Realtime usage dict keys before ResponseAPIUsage transform * Test usage transform for Realtime versus tokens_details keys * Avoid usage_input dict in-place * Fix audio cost calculation * fix(responses): forward output audio_tokens into completion usage details Pass audio_tokens from output_tokens_details into CompletionTokensDetailsWrapper so cost can use output_cost_per_audio_token. Support dict output details like prompt path. Extend tests for Realtime and mixed completion audio. Co-authored-by: Cursor * Fix audio token usage formatting * style: Black-format Realtime usage and completion usage merge Resolve combine_usage_objects and responses/utils wrapping for CI black --check. Restore model_fields comments above completion_tokens_details merge loop. Co-authored-by: Cursor * Add test to cover combined usage objects * Fix merge conflict with test cases Removed unnecessary import statement and cleaned up assertions in test. * fix(cost_calculator): remove dead None guard in completion_tokens_details combiner --------- Co-authored-by: Liam McDonald Co-authored-by: Cursor --- litellm/cost_calculator.py | 9 ++- litellm/responses/utils.py | 15 +++++ .../responses/test_responses_utils.py | 59 ++++++++++++++++++- tests/test_litellm/test_cost_calculator.py | 45 ++++++++++++++ 4 files changed, 121 insertions(+), 7 deletions(-) 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/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/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_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 = [