From 582d324a768b0d9f5125eeeb13e5a92578bce4ce Mon Sep 17 00:00:00 2001 From: Jay Prajapati <79649559+jayy-77@users.noreply.github.com> Date: Mon, 26 Jan 2026 12:29:50 +0530 Subject: [PATCH 1/7] fix(proxy): support slashes in google generateContent model names (#19737) * fix(proxy): support slashes in google route params * fix(proxy): extract google model ids with slashes * test(proxy): cover google model ids with slashes --- litellm/proxy/_types.py | 12 +++++----- litellm/proxy/auth/auth_utils.py | 20 ++++++++++++++-- litellm/proxy/auth/route_checks.py | 10 +++++++- .../proxy/auth/test_auth_utils.py | 23 +++++++++++++++++++ .../proxy/auth/test_route_checks.py | 4 ++++ 5 files changed, 60 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0516a4aaa66..60f968d9be1 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -425,12 +425,12 @@ class LiteLLMRoutes(enum.Enum): ] google_routes = [ - "/v1beta/models/{model_name}:countTokens", - "/v1beta/models/{model_name}:generateContent", - "/v1beta/models/{model_name}:streamGenerateContent", - "/models/{model_name}:countTokens", - "/models/{model_name}:generateContent", - "/models/{model_name}:streamGenerateContent", + "/v1beta/models/{model_name:path}:countTokens", + "/v1beta/models/{model_name:path}:generateContent", + "/v1beta/models/{model_name:path}:streamGenerateContent", + "/models/{model_name:path}:countTokens", + "/models/{model_name:path}:generateContent", + "/models/{model_name:path}:streamGenerateContent", # Google Interactions API "/interactions", "/v1beta/interactions", diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 9b9a988c07a..2bd84a1d98d 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -758,11 +758,27 @@ def get_model_from_request( if match: model = match.group(1) + # If still not found, extract model from Google generateContent-style routes. + # These routes put the model in the path and allow "/" inside the model id. + # Examples: + # - /v1beta/models/gemini-2.0-flash:generateContent + # - /v1beta/models/bedrock/claude-sonnet-3.7:generateContent + # - /models/custom/ns/model:streamGenerateContent + if model is None and not route.lower().startswith("/vertex"): + google_match = re.search(r"/(?:v1beta|beta)/models/([^:]+):", route) + if google_match: + model = google_match.group(1) + + if model is None and not route.lower().startswith("/vertex"): + google_match = re.search(r"^/models/([^:]+):", route) + if google_match: + model = google_match.group(1) + # If still not found, extract from Vertex AI passthrough route # Pattern: /vertex_ai/.../models/{model_id}:* # Example: /vertex_ai/v1/.../models/gemini-1.5-pro:generateContent - if model is None and "/vertex" in route.lower(): - vertex_match = re.search(r"/models/([^/:]+)", route) + if model is None and route.lower().startswith("/vertex"): + vertex_match = re.search(r"/models/([^:]+)", route) if vertex_match: model = vertex_match.group(1) diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index 96a70b8016f..fa63ff01b3a 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -392,7 +392,15 @@ class RouteChecks: # Ensure route is a string before attempting regex matching if not isinstance(route, str): return False - pattern = re.sub(r"\{[^}]+\}", r"[^/]+", pattern) + + def _placeholder_to_regex(match: re.Match) -> str: + placeholder = match.group(0).strip("{}") + if placeholder.endswith(":path"): + # allow "/" in the placeholder value, but don't eat the route suffix after ":" + return r"[^:]+" + return r"[^/]+" + + pattern = re.sub(r"\{[^}]+\}", _placeholder_to_regex, pattern) # Anchor the pattern to match the entire string pattern = f"^{pattern}$" if re.match(pattern, route): diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 62f9cc33b64..82920ce1d80 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -8,6 +8,7 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_utils import ( _get_customer_id_from_standard_headers, get_end_user_id_from_request_body, + get_model_from_request, get_key_model_rpm_limit, get_key_model_tpm_limit, ) @@ -186,3 +187,25 @@ class TestGetEndUserIdFromRequestBodyWithStandardHeaders: request_body=request_body, request_headers=headers ) assert result == "body-user" + + +def test_get_model_from_request_supports_google_model_names_with_slashes(): + assert ( + get_model_from_request( + request_data={}, + route="/v1beta/models/bedrock/claude-sonnet-3.7:generateContent", + ) + == "bedrock/claude-sonnet-3.7" + ) + assert ( + get_model_from_request( + request_data={}, + route="/models/hosted_vllm/gpt-oss-20b:generateContent", + ) + == "hosted_vllm/gpt-oss-20b" + ) + + +def test_get_model_from_request_vertex_passthrough_still_works(): + route = "/vertex_ai/v1/projects/p/locations/l/publishers/google/models/gemini-1.5-pro:generateContent" + assert get_model_from_request(request_data={}, route=route) == "gemini-1.5-pro" diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index b8084906fa5..a745ac3de13 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -161,9 +161,11 @@ def test_virtual_key_llm_api_route_includes_passthrough_prefix(route): [ "/v1beta/models/gemini-2.5-flash:countTokens", "/v1beta/models/gemini-2.0-flash:generateContent", + "/v1beta/models/bedrock/claude-sonnet-3.7:generateContent", "/v1beta/models/gemini-1.5-pro:streamGenerateContent", "/models/gemini-2.5-flash:countTokens", "/models/gemini-2.0-flash:generateContent", + "/models/bedrock/claude-sonnet-3.7:generateContent", "/models/gemini-1.5-pro:streamGenerateContent", ], ) @@ -187,9 +189,11 @@ def test_virtual_key_llm_api_routes_allows_google_routes(route): "/v1beta/models/google-gemini-2-5-pro-code-reviewer-k8s:generateContent", "/v1beta/models/gemini-2.5-flash-exp:countTokens", "/v1beta/models/custom-model-name-123:streamGenerateContent", + "/v1beta/models/bedrock/claude-sonnet-3.7:generateContent", "/models/google-gemini-2-5-pro-code-reviewer-k8s:generateContent", "/models/gemini-2.5-flash-exp:countTokens", "/models/custom-model-name-123:streamGenerateContent", + "/models/bedrock/claude-sonnet-3.7:generateContent", ], ) def test_google_routes_with_dynamic_model_names_recognized_as_llm_api_route(route): From 5666c725ce8a7b197ea4271fadc77a7a0736d89e Mon Sep 17 00:00:00 2001 From: jquinter Date: Mon, 26 Jan 2026 04:00:53 -0300 Subject: [PATCH 2/7] Fix/non standard mcp url pattern (#19738) * fix(mcp): Add standard MCP URL pattern support for OAuth discovery (#17272) OAuth discovery endpoints now support both URL patterns: - Standard MCP pattern: /mcp/{server_name} (new) - Legacy LiteLLM pattern: /{server_name}/mcp (backward compatible) The standard pattern is required by MCP-compliant clients like mcp-inspector and VSCode Copilot, which expect resource URLs following the /mcp/{server_name} convention per RFC 9728. Changes: - Add _build_oauth_protected_resource_response() helper - Add oauth_protected_resource_mcp_standard() endpoint - Add oauth_authorization_server_mcp_standard() endpoint - Keep legacy endpoints for backward compatibility - Add tests for both URL patterns Fixes #17272 * fix(mcp): Add standard MCP URL pattern support for OAuth discovery (#17272) OAuth discovery endpoints now support both URL patterns: - Standard MCP pattern: /mcp/{server_name} (new) - Legacy LiteLLM pattern: /{server_name}/mcp (backward compatible) The standard pattern is required by MCP-compliant clients like mcp-inspector and VSCode Copilot, which expect resource URLs following the /mcp/{server_name} convention per RFC 9728. Changes: - Add _build_oauth_protected_resource_response() helper - Add oauth_protected_resource_mcp_standard() endpoint - Add oauth_authorization_server_mcp_standard() endpoint - Keep legacy endpoints for backward compatibility - Add tests for both URL patterns Fixes #17272 * Test was relocated * refactor(mcp): Extract helper methods from run_with_session to fix PLR0915 Split the large run_with_session method (55 statements) into smaller helper methods to satisfy ruff's PLR0915 rule (max 50 statements): - _create_transport_context(): Creates transport based on type - _execute_session_operation(): Handles session lifecycle Also changed cleanup exception handling from Exception to BaseException to properly catch asyncio.CancelledError (which is a BaseException subclass in Python 3.8+). Co-Authored-By: Claude Opus 4.5 * test(mcp): Fix flaky test by mocking health_check_server The test_mcp_server_manager_config_integration_with_database test was making real network calls to fake URLs which caused timeouts and CancelledError exceptions. Fixed by mocking health_check_server to return a proper LiteLLM_MCPServerTable object instead of making network calls. * test(mcp): Fix skip condition to properly detect claude model names The skip condition for missing API keys was checking for "anthropic" in the model name, but the test uses "claude-haiku-4-5" which doesn't match. Updated to check for both "anthropic" and "claude" model patterns. Also added skip condition for OpenAI models when OPENAI_API_KEY is not set. Co-Authored-By: Claude Opus 4.5 * test(mcp): Fix skip condition to properly detect claude model names The skip condition for missing API keys was checking for "anthropic" in the model name, but the test uses "claude-haiku-4-5" which doesn't match. Updated to check for both "anthropic" and "claude" model patterns. Also added skip condition for OpenAI models when OPENAI_API_KEY is not set. --------- Co-authored-by: Claude Opus 4.5 --- litellm/experimental_mcp_client/client.py | 159 +-- .../mcp_server/discoverable_endpoints.py | 172 ++- .../mcp_server/test_discoverable_endpoints.py | 1108 +++++++++++++++++ .../mcp_tests/test_aresponses_api_with_mcp.py | 6 +- tests/mcp_tests/test_mcp_server.py | 21 +- 5 files changed, 1361 insertions(+), 105 deletions(-) create mode 100644 tests/litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 5ad2dd54853..ab576d49f3c 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -4,7 +4,7 @@ LiteLLM Proxy uses this MCP Client to connnect to other MCP servers. import asyncio import base64 -from typing import Awaitable, Callable, Dict, List, Optional, TypeVar, Union +from typing import Any, Awaitable, Callable, Dict, List, Optional, Tuple, TypeVar, Union import httpx from mcp import ClientSession, ReadResourceResult, Resource, StdioServerParameters @@ -74,97 +74,102 @@ class MCPClient: if auth_value: self.update_auth_value(auth_value) + def _create_transport_context( + self, + ) -> Tuple[Any, Optional[httpx.AsyncClient]]: + """ + Create the appropriate transport context based on transport type. + + Returns: + Tuple of (transport_context, http_client). + http_client is only set for HTTP transport and needs cleanup. + """ + http_client: Optional[httpx.AsyncClient] = None + + if self.transport_type == MCPTransport.stdio: + if not self.stdio_config: + raise ValueError("stdio_config is required for stdio transport") + server_params = StdioServerParameters( + command=self.stdio_config.get("command", ""), + args=self.stdio_config.get("args", []), + env=self.stdio_config.get("env", {}), + ) + return stdio_client(server_params), None + + if self.transport_type == MCPTransport.sse: + headers = self._get_auth_headers() + httpx_client_factory = self._create_httpx_client_factory() + return sse_client( + url=self.server_url, + timeout=self.timeout, + headers=headers, + httpx_client_factory=httpx_client_factory, + ), None + + # HTTP transport (default) + headers = self._get_auth_headers() + httpx_client_factory = self._create_httpx_client_factory() + verbose_logger.debug( + "litellm headers for streamable_http_client: %s", headers + ) + http_client = httpx_client_factory( + headers=headers, + timeout=httpx.Timeout(self.timeout), + ) + transport_ctx = streamable_http_client( + url=self.server_url, + http_client=http_client, + ) + return transport_ctx, http_client + + async def _execute_session_operation( + self, + transport_ctx: Any, + operation: Callable[[ClientSession], Awaitable[TSessionResult]], + ) -> TSessionResult: + """ + Execute an operation within a transport and session context. + + Handles entering/exiting contexts and running the operation. + """ + transport = await transport_ctx.__aenter__() + try: + read_stream, write_stream = transport[0], transport[1] + session_ctx = ClientSession(read_stream, write_stream) + session = await session_ctx.__aenter__() + try: + await session.initialize() + return await operation(session) + finally: + try: + await session_ctx.__aexit__(None, None, None) + except BaseException as e: + verbose_logger.debug(f"Error during session context exit: {e}") + finally: + try: + await transport_ctx.__aexit__(None, None, None) + except BaseException as e: + verbose_logger.debug(f"Error during transport context exit: {e}") + async def run_with_session( self, operation: Callable[[ClientSession], Awaitable[TSessionResult]] ) -> TSessionResult: """Open a session, run the provided coroutine, and clean up.""" - transport_ctx = None http_client: Optional[httpx.AsyncClient] = None - transport = None - session_ctx = None - try: - if self.transport_type == MCPTransport.stdio: - if not self.stdio_config: - raise ValueError("stdio_config is required for stdio transport") - - server_params = StdioServerParameters( - command=self.stdio_config.get("command", ""), - args=self.stdio_config.get("args", []), - env=self.stdio_config.get("env", {}), - ) - transport_ctx = stdio_client(server_params) - elif self.transport_type == MCPTransport.sse: - headers = self._get_auth_headers() - httpx_client_factory = self._create_httpx_client_factory() - transport_ctx = sse_client( - url=self.server_url, - timeout=self.timeout, - headers=headers, - httpx_client_factory=httpx_client_factory, - ) - else: - headers = self._get_auth_headers() - httpx_client_factory = self._create_httpx_client_factory() - verbose_logger.debug( - "litellm headers for streamable_http_client: %s", headers - ) - http_client = httpx_client_factory( - headers=headers, - timeout=httpx.Timeout(self.timeout), - ) - transport_ctx = streamable_http_client( - url=self.server_url, - http_client=http_client, - ) - - if transport_ctx is None: - raise RuntimeError("Failed to create transport context") - - # Enter transport context - transport = await transport_ctx.__aenter__() - try: - read_stream, write_stream = transport[0], transport[1] - session_ctx = ClientSession(read_stream, write_stream) - - # Enter session context - session = await session_ctx.__aenter__() - try: - await session.initialize() - result = await operation(session) - return result - finally: - # Ensure session context is properly exited - if session_ctx is not None: - try: - await session_ctx.__aexit__(None, None, None) - except Exception as e: - verbose_logger.debug( - f"Error during session context exit: {e}" - ) - finally: - # Ensure transport context is properly exited - if transport_ctx is not None: - try: - await transport_ctx.__aexit__(None, None, None) - except Exception as e: - verbose_logger.debug( - f"Error during transport context exit: {e}" - ) + transport_ctx, http_client = self._create_transport_context() + return await self._execute_session_operation(transport_ctx, operation) except Exception: verbose_logger.warning( "MCP client run_with_session failed for %s", self.server_url or "stdio" ) raise finally: - # Always clean up http_client if it was created if http_client is not None: try: await http_client.aclose() - except Exception as e: - verbose_logger.debug( - f"Error during http_client cleanup: {e}" - ) + except BaseException as e: + verbose_logger.debug(f"Error during http_client cleanup: {e}") def update_auth_value(self, mcp_auth_value: Union[str, Dict[str, str]]): """ diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index ded591a8f53..56feff548ad 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -387,25 +387,57 @@ async def callback(code: str, state: str): 1. Try resource_metadata from WWW-Authenticate header (if present) 2. Fall back to path-based well-known URI: /.well-known/oauth-protected-resource/{path} ( - If the resource identifier value contains a path or query component, any terminating slash (/) - following the host component MUST be removed before inserting /.well-known/ and the well-known - URI path suffix between the host component and the path(include root path) and/or query components. + If the resource identifier value contains a path or query component, any terminating slash (/) + following the host component MUST be removed before inserting /.well-known/ and the well-known + URI path suffix between the host component and the path(include root path) and/or query components. https://datatracker.ietf.org/doc/html/rfc9728#section-3.1) 3. Fall back to root-based well-known URI: /.well-known/oauth-protected-resource + + Dual Pattern Support: + - Standard MCP pattern: /mcp/{server_name} (recommended, used by mcp-inspector, VSCode Copilot) + - LiteLLM legacy pattern: /{server_name}/mcp (backward compatibility) + + The resource URL returned matches the pattern used in the discovery request. """ -@router.get(f"/.well-known/oauth-protected-resource{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}/mcp") -@router.get("/.well-known/oauth-protected-resource") -async def oauth_protected_resource_mcp( - request: Request, mcp_server_name: Optional[str] = None -): + + +def _build_oauth_protected_resource_response( + request: Request, + mcp_server_name: Optional[str], + use_standard_pattern: bool, +) -> dict: + """ + Build OAuth protected resource response with the appropriate URL pattern. + + Args: + request: FastAPI Request object + mcp_server_name: Name of the MCP server + use_standard_pattern: If True, use /mcp/{server_name} pattern; + if False, use /{server_name}/mcp pattern + + Returns: + OAuth protected resource metadata dict + """ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, ) - # Get the correct base URL considering X-Forwarded-* headers + request_base_url = get_request_base_url(request) mcp_server: Optional[MCPServer] = None if mcp_server_name: mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name) + + # Build resource URL based on the pattern + if mcp_server_name: + if use_standard_pattern: + # Standard MCP pattern: /mcp/{server_name} + resource_url = f"{request_base_url}/mcp/{mcp_server_name}" + else: + # LiteLLM legacy pattern: /{server_name}/mcp + resource_url = f"{request_base_url}/{mcp_server_name}/mcp" + else: + resource_url = f"{request_base_url}/mcp" + return { "authorization_servers": [ ( @@ -414,14 +446,55 @@ async def oauth_protected_resource_mcp( else f"{request_base_url}" ) ], - "resource": ( - f"{request_base_url}/{mcp_server_name}/mcp" - if mcp_server_name - else f"{request_base_url}/mcp" - ), # this is what Claude will call + "resource": resource_url, "scopes_supported": mcp_server.scopes if mcp_server else [], } + +# Standard MCP pattern: /.well-known/oauth-protected-resource/mcp/{server_name} +# This is the pattern expected by standard MCP clients (mcp-inspector, VSCode Copilot) +@router.get(f"/.well-known/oauth-protected-resource{'' if get_server_root_path() == '/' else get_server_root_path()}/mcp/{{mcp_server_name}}") +async def oauth_protected_resource_mcp_standard( + request: Request, mcp_server_name: str +): + """ + OAuth protected resource discovery endpoint using standard MCP URL pattern. + + Standard pattern: /mcp/{server_name} + Discovery path: /.well-known/oauth-protected-resource/mcp/{server_name} + + This endpoint is compliant with MCP specification and works with standard + MCP clients like mcp-inspector and VSCode Copilot. + """ + return _build_oauth_protected_resource_response( + request=request, + mcp_server_name=mcp_server_name, + use_standard_pattern=True, + ) + + +# LiteLLM legacy pattern: /.well-known/oauth-protected-resource/{server_name}/mcp +# Kept for backward compatibility with existing deployments +@router.get(f"/.well-known/oauth-protected-resource{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}/mcp") +@router.get("/.well-known/oauth-protected-resource") +async def oauth_protected_resource_mcp( + request: Request, mcp_server_name: Optional[str] = None +): + """ + OAuth protected resource discovery endpoint using LiteLLM legacy URL pattern. + + Legacy pattern: /{server_name}/mcp + Discovery path: /.well-known/oauth-protected-resource/{server_name}/mcp + + This endpoint is kept for backward compatibility. New integrations should + use the standard MCP pattern (/mcp/{server_name}) instead. + """ + return _build_oauth_protected_resource_response( + request=request, + mcp_server_name=mcp_server_name, + use_standard_pattern=False, + ) + """ https://datatracker.ietf.org/doc/html/rfc8414#section-3.1 RFC 8414: Path-aware OAuth discovery @@ -430,15 +503,26 @@ async def oauth_protected_resource_mcp( the well-known URI suffix between the host component and the path(include root path) component. """ -@router.get(f"/.well-known/oauth-authorization-server{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}") -@router.get("/.well-known/oauth-authorization-server") -async def oauth_authorization_server_mcp( - request: Request, mcp_server_name: Optional[str] = None -): + + +def _build_oauth_authorization_server_response( + request: Request, + mcp_server_name: Optional[str], +) -> dict: + """ + Build OAuth authorization server metadata response. + + Args: + request: FastAPI Request object + mcp_server_name: Name of the MCP server + + Returns: + OAuth authorization server metadata dict + """ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, ) - # Get the correct base URL considering X-Forwarded-* headers + request_base_url = get_request_base_url(request) authorization_endpoint = ( @@ -470,18 +554,58 @@ async def oauth_authorization_server_mcp( } +# Standard MCP pattern: /.well-known/oauth-authorization-server/mcp/{server_name} +@router.get(f"/.well-known/oauth-authorization-server{'' if get_server_root_path() == '/' else get_server_root_path()}/mcp/{{mcp_server_name}}") +async def oauth_authorization_server_mcp_standard( + request: Request, mcp_server_name: str +): + """ + OAuth authorization server discovery endpoint using standard MCP URL pattern. + + Standard pattern: /mcp/{server_name} + Discovery path: /.well-known/oauth-authorization-server/mcp/{server_name} + """ + return _build_oauth_authorization_server_response( + request=request, + mcp_server_name=mcp_server_name, + ) + + +# LiteLLM legacy pattern and root endpoint +@router.get(f"/.well-known/oauth-authorization-server{'' if get_server_root_path() == '/' else get_server_root_path()}/{{mcp_server_name}}") +@router.get("/.well-known/oauth-authorization-server") +async def oauth_authorization_server_mcp( + request: Request, mcp_server_name: Optional[str] = None +): + """ + OAuth authorization server discovery endpoint. + + Supports both legacy pattern (/{server_name}) and root endpoint. + """ + return _build_oauth_authorization_server_response( + request=request, + mcp_server_name=mcp_server_name, + ) + + # Alias for standard OpenID discovery @router.get("/.well-known/openid-configuration") async def openid_configuration(request: Request): return await oauth_authorization_server_mcp(request) +# Additional legacy pattern support @router.get("/.well-known/oauth-authorization-server/{mcp_server_name}/mcp") -@router.get("/.well-known/oauth-authorization-server") -async def oauth_authorization_server_root( - request: Request, mcp_server_name: Optional[str] = None +async def oauth_authorization_server_legacy( + request: Request, mcp_server_name: str ): - return await oauth_authorization_server_mcp(request, mcp_server_name) + """ + OAuth authorization server discovery for legacy /{server_name}/mcp pattern. + """ + return _build_oauth_authorization_server_response( + request=request, + mcp_server_name=mcp_server_name, + ) @router.post("/{mcp_server_name}/register") diff --git a/tests/litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py new file mode 100644 index 00000000000..ac2c945baa6 --- /dev/null +++ b/tests/litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -0,0 +1,1108 @@ +"""Tests for MCP OAuth discoverable endpoints""" +import pytest +from unittest.mock import AsyncMock, MagicMock, patch + + +@pytest.mark.asyncio +async def test_authorize_endpoint_includes_response_type(): + """Test that authorize endpoint includes response_type=code parameter (fixes #15684)""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + authorize, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Clear registry + global_mcp_server_manager.registry.clear() + + # Create mock OAuth2 server + oauth2_server = MCPServer( + server_id="test_oauth_server", + name="test_oauth", + server_name="test_oauth", + alias="test_oauth", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="test_client_id", + client_secret="test_client_secret", + authorization_url="https://provider.com/oauth/authorize", + token_url="https://provider.com/oauth/token", + scopes=["read", "write"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + # Mock request + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://litellm.example.com/" + mock_request.headers = {} + + # Mock the encryption functions to avoid needing a signing key + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper" + ) as mock_encrypt: + mock_encrypt.return_value = "mocked_encrypted_state" + + # Call authorize endpoint + response = await authorize( + request=mock_request, + client_id="test_client_id", + mcp_server_name="test_oauth", + redirect_uri="https://client.example.com/callback", + state="test_state", + ) + + # Verify response is a redirect + assert response.status_code == 307 # FastAPI RedirectResponse default + + # Verify response_type is in the redirect URL + assert "response_type=code" in response.headers["location"] + assert "https://provider.com/oauth/authorize" in response.headers["location"] + assert "client_id=test_client_id" in response.headers["location"] + assert "scope=read+write" in response.headers["location"] + + +@pytest.mark.asyncio +async def test_authorize_endpoint_forwards_pkce_parameters(): + """Test that authorize endpoint forwards PKCE parameters (code_challenge and code_challenge_method)""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + authorize, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Clear registry + global_mcp_server_manager.registry.clear() + + # Create mock OAuth2 server (simulating Google OAuth) + oauth2_server = MCPServer( + server_id="google_mcp", + name="google_mcp", + server_name="google_mcp", + alias="google_mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="669428968603-test.apps.googleusercontent.com", + client_secret="GOCSPX-test_secret", + authorization_url="https://accounts.google.com/o/oauth2/v2/auth", + token_url="https://oauth2.googleapis.com/token", + scopes=["https://www.googleapis.com/auth/drive", "openid", "email"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + # Mock request + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://litellm-proxy.example.com/" + mock_request.headers = {} + + # Mock the encryption function + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper" + ) as mock_encrypt: + mock_encrypt.return_value = "mocked_encrypted_state_with_pkce" + + # Call authorize endpoint with PKCE parameters + response = await authorize( + request=mock_request, + client_id="669428968603-test.apps.googleusercontent.com", + mcp_server_name="google_mcp", + redirect_uri="http://localhost:60108/callback", + state="test_client_state", + code_challenge="x6YH_qgwbvOzbsHDuL1sW9gYkR9-gObUiIB5RkPwxDk", + code_challenge_method="S256", + ) + + # Verify response is a redirect + assert response.status_code == 307 + + # Verify PKCE parameters are included in the redirect URL + location = response.headers["location"] + assert "https://accounts.google.com/o/oauth2/v2/auth" in location + assert "code_challenge=x6YH_qgwbvOzbsHDuL1sW9gYkR9-gObUiIB5RkPwxDk" in location + assert "code_challenge_method=S256" in location + assert "client_id=669428968603-test.apps.googleusercontent.com" in location + assert "response_type=code" in location + + +@pytest.mark.asyncio +async def test_token_endpoint_forwards_code_verifier(): + """Test that token endpoint forwards code_verifier for PKCE flow""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + token_endpoint, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + import httpx + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Clear registry + global_mcp_server_manager.registry.clear() + + # Create mock OAuth2 server + oauth2_server = MCPServer( + server_id="google_mcp", + name="google_mcp", + server_name="google_mcp", + alias="google_mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="669428968603-test.apps.googleusercontent.com", + client_secret="GOCSPX-test_secret", + authorization_url="https://accounts.google.com/o/oauth2/v2/auth", + token_url="https://oauth2.googleapis.com/token", + scopes=["https://www.googleapis.com/auth/drive", "openid", "email"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + # Mock request + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://litellm-proxy.example.com/" + mock_request.headers = {} + + # Mock httpx client response + mock_response = MagicMock() + mock_response.json.return_value = { + "access_token": "ya29.test_access_token", + "token_type": "Bearer", + "expires_in": 3599, + "scope": "openid email https://www.googleapis.com/auth/drive", + } + mock_response.raise_for_status = MagicMock() + + # Mock the async httpx client with AsyncMock for async methods + from unittest.mock import AsyncMock + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client" + ) as mock_get_client: + mock_async_client = MagicMock() + # Use AsyncMock for the async post method + mock_async_client.post = AsyncMock(return_value=mock_response) + mock_get_client.return_value = mock_async_client + + # Call token endpoint with code_verifier + response = await token_endpoint( + request=mock_request, + grant_type="authorization_code", + code="4/test_authorization_code", + redirect_uri="http://localhost:60108/callback", + client_id="669428968603-test.apps.googleusercontent.com", + mcp_server_name="google_mcp", + client_secret="GOCSPX-test_secret", + code_verifier="test_code_verifier_from_client", + ) + + # Verify that the token endpoint was called with code_verifier + mock_async_client.post.assert_called_once() + call_args = mock_async_client.post.call_args + + # Check the data parameter includes code_verifier + assert call_args[1]["data"]["code_verifier"] == "test_code_verifier_from_client" + assert call_args[1]["data"]["code"] == "4/test_authorization_code" + assert ( + call_args[1]["data"]["client_id"] + == "669428968603-test.apps.googleusercontent.com" + ) + assert call_args[1]["data"]["client_secret"] == "GOCSPX-test_secret" + assert call_args[1]["data"]["grant_type"] == "authorization_code" + + # Verify response + response_data = response.body + import json + + token_data = json.loads(response_data) + assert token_data["access_token"] == "ya29.test_access_token" + assert token_data["token_type"] == "Bearer" + + +@pytest.mark.asyncio +async def test_register_client_without_mcp_server_name_returns_dummy(): + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + register_client, + ) + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://proxy.litellm.example/" + mock_request.headers = {} + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + new=AsyncMock(return_value={}), + ): + result = await register_client(request=mock_request) + + assert result == { + "client_id": "dummy_client", + "client_secret": "dummy", + "redirect_uris": ["https://proxy.litellm.example/callback"], + } + + +@pytest.mark.asyncio +async def test_register_client_returns_existing_server_credentials(): + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + register_client, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + global_mcp_server_manager.registry.clear() + oauth2_server = MCPServer( + server_id="stored_server", + name="stored_server", + server_name="stored_server", + alias="stored_server", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="existing-client", + client_secret="existing-secret", + authorization_url="https://provider.example/oauth/authorize", + token_url="https://provider.example/oauth/token", + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://proxy.litellm.example/" + mock_request.headers = {} + + try: + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + new=AsyncMock(return_value={}), + ): + result = await register_client( + request=mock_request, mcp_server_name=oauth2_server.server_name + ) + finally: + global_mcp_server_manager.registry.clear() + + assert result == { + "client_id": "stored_server", + "client_secret": "dummy", + "redirect_uris": ["https://proxy.litellm.example/callback"], + } + + +@pytest.mark.asyncio +async def test_register_client_remote_registration_success(): + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + register_client, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + global_mcp_server_manager.registry.clear() + oauth2_server = MCPServer( + server_id="remote_server", + name="remote_server", + server_name="remote_server", + alias="remote_server", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id=None, + client_secret=None, + authorization_url="https://provider.example/oauth/authorize", + token_url="https://provider.example/oauth/token", + registration_url="https://provider.example/oauth/register", + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://proxy.litellm.example/" + mock_request.headers = {} + + request_payload = { + "client_name": "Litellm Proxy", + "grant_types": ["authorization_code", "refresh_token"], + "response_types": ["code"], + "token_endpoint_auth_method": "client_secret_post", + } + + mock_response = MagicMock() + mock_response.json.return_value = { + "client_id": "generated-client", + "client_secret": "generated-secret", + } + mock_response.raise_for_status = MagicMock() + mock_async_client = MagicMock() + mock_async_client.post = AsyncMock(return_value=mock_response) + + try: + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + new=AsyncMock(return_value=request_payload), + ), patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=mock_async_client, + ): + response = await register_client( + request=mock_request, mcp_server_name=oauth2_server.server_name + ) + finally: + global_mcp_server_manager.registry.clear() + + import json + + assert response.status_code == 200 + payload = json.loads(response.body.decode("utf-8")) + assert payload == mock_response.json.return_value + + mock_async_client.post.assert_called_once() + call_args = mock_async_client.post.call_args + assert call_args.args[0] == oauth2_server.registration_url + assert call_args.kwargs["headers"] == { + "Content-Type": "application/json", + "Accept": "application/json", + } + assert call_args.kwargs["json"]["redirect_uris"] == [ + "https://proxy.litellm.example/callback" + ] + assert call_args.kwargs["json"]["grant_types"] == request_payload["grant_types"] + assert ( + call_args.kwargs["json"]["token_endpoint_auth_method"] + == request_payload["token_endpoint_auth_method"] + ) + + +@pytest.mark.asyncio +async def test_authorize_endpoint_respects_x_forwarded_proto(): + """Test that authorize endpoint uses X-Forwarded-Proto header to construct correct redirect_uri""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + authorize, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Clear registry + global_mcp_server_manager.registry.clear() + + # Create mock OAuth2 server + oauth2_server = MCPServer( + server_id="test_oauth_server", + name="test_oauth", + server_name="test_oauth", + alias="test_oauth", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="test_client_id", + client_secret="test_client_secret", + authorization_url="https://provider.com/oauth/authorize", + token_url="https://provider.com/oauth/token", + scopes=["read", "write"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + # Mock request with http base_url but X-Forwarded-Proto: https + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://litellm.example.com/" # HTTP + mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + + # Mock the encryption functions + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper" + ) as mock_encrypt: + mock_encrypt.return_value = "mocked_encrypted_state" + + # Call authorize endpoint + response = await authorize( + request=mock_request, + client_id="test_client_id", + mcp_server_name="test_oauth", + redirect_uri="https://client.example.com/callback", + state="test_state", + ) + + # Verify redirect URL uses HTTPS in the redirect_uri parameter + location = response.headers["location"] + + # The redirect_uri parameter sent to the OAuth provider should use HTTPS + assert ( + "redirect_uri=https%3A%2F%2Flitellm.example.com%2Fcallback" in location + or "redirect_uri=https://litellm.example.com/callback" in location + ) + + +@pytest.mark.asyncio +async def test_token_endpoint_respects_x_forwarded_proto(): + """Test that token endpoint uses X-Forwarded-Proto header for redirect_uri""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + token_endpoint, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Clear registry + global_mcp_server_manager.registry.clear() + + # Create mock OAuth2 server + oauth2_server = MCPServer( + server_id="google_mcp", + name="google_mcp", + server_name="google_mcp", + alias="google_mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="test_client_id", + client_secret="test_secret", + authorization_url="https://accounts.google.com/o/oauth2/v2/auth", + token_url="https://oauth2.googleapis.com/token", + scopes=["openid", "email"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + # Mock request with http base_url but X-Forwarded-Proto: https + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://litellm-proxy.example.com/" # HTTP + mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + + # Mock httpx client response + mock_response = MagicMock() + mock_response.json.return_value = { + "access_token": "test_token", + "token_type": "Bearer", + "expires_in": 3599, + } + mock_response.raise_for_status = MagicMock() + + # Mock the async httpx client + mock_async_client = MagicMock() + mock_async_client.post = AsyncMock(return_value=mock_response) + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client" + ) as mock_get_client: + mock_get_client.return_value = mock_async_client + + # Call token endpoint + response = await token_endpoint( + request=mock_request, + grant_type="authorization_code", + code="test_code", + redirect_uri="http://localhost:60108/callback", + client_id="test_client_id", + mcp_server_name="google_mcp", + client_secret="test_secret", + ) + + # Verify that the redirect_uri sent to the provider uses HTTPS + call_args = mock_async_client.post.call_args + assert ( + call_args[1]["data"]["redirect_uri"] + == "https://litellm-proxy.example.com/callback" + ) + + +@pytest.mark.asyncio +async def test_oauth_protected_resource_standard_pattern(): + """Test that oauth_protected_resource_mcp_standard returns standard MCP URL pattern (/mcp/{server_name})""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + oauth_protected_resource_mcp_standard, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Clear registry + global_mcp_server_manager.registry.clear() + + # Create mock OAuth2 server + oauth2_server = MCPServer( + server_id="test_server", + name="test_server", + server_name="test_server", + alias="test_server", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="test_client_id", + client_secret="test_client_secret", + authorization_url="https://provider.com/oauth/authorize", + token_url="https://provider.com/oauth/token", + scopes=["read", "write"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + # Mock request + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://litellm.example.com/" + mock_request.headers = {} + + # Call the standard pattern endpoint + response = await oauth_protected_resource_mcp_standard( + request=mock_request, + mcp_server_name="test_server", + ) + + # Verify response uses standard MCP pattern: /mcp/{server_name} + assert response["resource"] == "https://litellm.example.com/mcp/test_server" + assert response["authorization_servers"][0] == "https://litellm.example.com/test_server" + assert response["scopes_supported"] == oauth2_server.scopes + + +@pytest.mark.asyncio +async def test_oauth_protected_resource_legacy_pattern(): + """Test that oauth_protected_resource_mcp returns legacy URL pattern (/{server_name}/mcp)""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + oauth_protected_resource_mcp, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Clear registry + global_mcp_server_manager.registry.clear() + + # Create mock OAuth2 server + oauth2_server = MCPServer( + server_id="test_server", + name="test_server", + server_name="test_server", + alias="test_server", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="test_client_id", + client_secret="test_client_secret", + authorization_url="https://provider.com/oauth/authorize", + token_url="https://provider.com/oauth/token", + scopes=["read", "write"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + # Mock request + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://litellm.example.com/" + mock_request.headers = {} + + # Call the legacy pattern endpoint + response = await oauth_protected_resource_mcp( + request=mock_request, + mcp_server_name="test_server", + ) + + # Verify response uses legacy pattern: /{server_name}/mcp + assert response["resource"] == "https://litellm.example.com/test_server/mcp" + assert response["authorization_servers"][0] == "https://litellm.example.com/test_server" + assert response["scopes_supported"] == oauth2_server.scopes + + +@pytest.mark.asyncio +async def test_oauth_protected_resource_respects_x_forwarded_proto(): + """Test that oauth_protected_resource_mcp uses X-Forwarded-Proto for URLs""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + oauth_protected_resource_mcp, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + # Clear registry + global_mcp_server_manager.registry.clear() + + # Create mock OAuth2 server + oauth2_server = MCPServer( + server_id="test_oauth_server", + name="test_oauth", + server_name="test_oauth", + alias="test_oauth", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="test_client_id", + client_secret="test_client_secret", + authorization_url="https://provider.com/oauth/authorize", + token_url="https://provider.com/oauth/token", + scopes=["read", "write"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + # Mock request with http base_url but X-Forwarded-Proto: https + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://litellm.example.com/" # HTTP + mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + + # Call the endpoint + response = await oauth_protected_resource_mcp( + request=mock_request, + mcp_server_name="test_oauth", + ) + + # Verify response uses HTTPS URLs + assert response["authorization_servers"][0].startswith( + "https://litellm.example.com/" + ) + assert response["scopes_supported"] == oauth2_server.scopes + + +@pytest.mark.asyncio +async def test_oauth_authorization_server_respects_x_forwarded_proto(): + """Test that oauth_authorization_server_mcp uses X-Forwarded-Proto for URLs""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + oauth_authorization_server_mcp, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + # Clear registry + global_mcp_server_manager.registry.clear() + + # Create mock OAuth2 server + oauth2_server = MCPServer( + server_id="test_oauth_server", + name="test_oauth", + server_name="test_oauth", + alias="test_oauth", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="test_client_id", + client_secret="test_client_secret", + authorization_url="https://provider.com/oauth/authorize", + token_url="https://provider.com/oauth/token", + scopes=["read", "write"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + # Mock request with http base_url but X-Forwarded-Proto: https + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://litellm.example.com/" # HTTP + mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + + # Call the endpoint + response = await oauth_authorization_server_mcp( + request=mock_request, + mcp_server_name="test_oauth", + ) + + # Verify response uses HTTPS URLs + assert response["authorization_endpoint"].startswith("https://litellm.example.com/") + assert response["token_endpoint"].startswith("https://litellm.example.com/") + assert response["registration_endpoint"].startswith("https://litellm.example.com/") + assert response["grant_types_supported"] == ["authorization_code", "refresh_token"] + assert response["scopes_supported"] == oauth2_server.scopes + + +@pytest.mark.asyncio +async def test_register_client_respects_x_forwarded_proto(): + """Test that register_client uses X-Forwarded-Proto for redirect_uris""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + register_client, + ) + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Mock request with http base_url but X-Forwarded-Proto: https + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://proxy.litellm.example/" # HTTP + mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", + new=AsyncMock(return_value={}), + ): + result = await register_client(request=mock_request) + + # Verify the redirect_uris use HTTPS + assert result == { + "client_id": "dummy_client", + "client_secret": "dummy", + "redirect_uris": ["https://proxy.litellm.example/callback"], + } + + +@pytest.mark.asyncio +async def test_authorize_endpoint_respects_x_forwarded_host(): + """Test that authorize endpoint uses X-Forwarded-Host and X-Forwarded-Proto to construct correct redirect_uri""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + authorize, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Clear registry + global_mcp_server_manager.registry.clear() + + # Create mock OAuth2 server + oauth2_server = MCPServer( + server_id="test_oauth_server", + name="test_oauth", + server_name="test_oauth", + alias="test_oauth", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="test_client_id", + client_secret="test_client_secret", + authorization_url="https://provider.com/oauth/authorize", + token_url="https://provider.com/oauth/token", + scopes=["read", "write"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + # Mock request simulating nginx proxy: + # Internal: http://localhost:8888/github/mcp + # External: https://proxy.example.com/github/mcp + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://localhost:8888/github/mcp" + mock_request.headers = { + "X-Forwarded-Proto": "https", + "X-Forwarded-Host": "proxy.example.com", + } + + # Mock the encryption functions + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper" + ) as mock_encrypt: + mock_encrypt.return_value = "mocked_encrypted_state" + + # Call authorize endpoint + response = await authorize( + request=mock_request, + client_id="test_client_id", + mcp_server_name="test_oauth", + redirect_uri="https://client.example.com/callback", + state="test_state", + ) + + # Verify redirect URL uses the forwarded host and scheme + location = response.headers["location"] + + # The redirect_uri parameter should use the external URL + assert ( + "redirect_uri=https%3A%2F%2Fproxy.example.com%2Fgithub%2Fmcp%2Fcallback" + in location + or "redirect_uri=https://proxy.example.com/github/mcp/callback" in location + ) + + +@pytest.mark.asyncio +async def test_token_endpoint_respects_x_forwarded_host(): + """Test that token endpoint uses X-Forwarded-Host and X-Forwarded-Proto for redirect_uri""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + token_endpoint, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Clear registry + global_mcp_server_manager.registry.clear() + + # Create mock OAuth2 server + oauth2_server = MCPServer( + server_id="google_mcp", + name="google_mcp", + server_name="google_mcp", + alias="google_mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="test_client_id", + client_secret="test_secret", + authorization_url="https://accounts.google.com/o/oauth2/v2/auth", + token_url="https://oauth2.googleapis.com/token", + scopes=["openid", "email"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + # Mock request simulating nginx proxy without port in host + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://localhost:8888/github/mcp" + mock_request.headers = { + "X-Forwarded-Proto": "https", + "X-Forwarded-Host": "proxy.example.com", + } + + # Mock httpx client response + mock_response = MagicMock() + mock_response.json.return_value = { + "access_token": "test_token", + "token_type": "Bearer", + "expires_in": 3599, + } + mock_response.raise_for_status = MagicMock() + + # Mock the async httpx client + mock_async_client = MagicMock() + mock_async_client.post = AsyncMock(return_value=mock_response) + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client" + ) as mock_get_client: + mock_get_client.return_value = mock_async_client + + # Call token endpoint + response = await token_endpoint( + request=mock_request, + grant_type="authorization_code", + code="test_code", + redirect_uri="http://localhost:60108/callback", + client_id="test_client_id", + mcp_server_name="google_mcp", + client_secret="test_secret", + ) + + # Verify that the redirect_uri sent to the provider uses the external URL + call_args = mock_async_client.post.call_args + assert ( + call_args[1]["data"]["redirect_uri"] + == "https://proxy.example.com/github/mcp/callback" + ) + + +@pytest.mark.parametrize( + "base_url,x_forwarded_proto,x_forwarded_host,x_forwarded_port,expected_url", + [ + # Case 1: No forwarded headers - use original URL as-is (no trailing slash) + ( + "http://localhost:4000/", + None, + None, + None, + "http://localhost:4000", + ), + # Case 2: Only X-Forwarded-Proto - change scheme only + ( + "http://localhost:4000/", + "https", + None, + None, + "https://localhost:4000", + ), + # Case 3: X-Forwarded-Proto + X-Forwarded-Host - change scheme and host + ( + "http://localhost:4000/", + "https", + "proxy.example.com", + None, + "https://proxy.example.com", + ), + # Case 4: X-Forwarded-Host with port included in host header + ( + "http://localhost:4000/", + "https", + "proxy.example.com:8080", + None, + "https://proxy.example.com:8080", + ), + # Case 5: X-Forwarded-Host + X-Forwarded-Port as separate headers + ( + "http://localhost:4000/", + "https", + "proxy.example.com", + "8443", + "https://proxy.example.com:8443", + ), + # Case 6: Only X-Forwarded-Host without proto - use original scheme + ( + "http://localhost:4000/", + None, + "proxy.example.com", + None, + "http://proxy.example.com", + ), + # Case 7: Only X-Forwarded-Port without host - preserves original port if present + # (This is safer behavior - X-Forwarded-Port alone is unusual) + ( + "http://localhost:4000/", + None, + None, + "8443", + "http://localhost:4000", # Original port preserved when already present + ), + # Case 8: Complex internal URL with path (path is preserved) + ( + "http://localhost:8888/github/mcp", + "https", + "proxy.example.com", + None, + "https://proxy.example.com/github/mcp", + ), + # Case 9: IPv6 address in X-Forwarded-Host (should not treat :: as port separator) + ( + "http://localhost:4000/", + "https", + "[2001:db8::1]", + None, + "https://[2001:db8::1]", + ), + # Case 10: IPv6 address with port + ( + "http://localhost:4000/", + "https", + "[2001:db8::1]:8080", + None, + "https://[2001:db8::1]:8080", + ), + # Case 11: X-Forwarded-Host already has port, X-Forwarded-Port also provided (host wins) + ( + "http://localhost:4000/", + "https", + "proxy.example.com:9000", + "8443", + "https://proxy.example.com:9000", + ), + # Case 12: Standard proxy setup (most common case) + ( + "http://127.0.0.1:8888/", + "https", + "chatproxy.company.com", + None, + "https://chatproxy.company.com", + ), + # Case 13: Internal URL already has port, X-Forwarded-Port does NOT override + # (safer behavior - preserves original port when X-Forwarded-Host not provided) + ( + "http://localhost:4000/", + None, + None, + "443", + "http://localhost:4000", # Original port preserved + ), + # Case 14: Original URL with existing port in netloc, X-Forwarded-Host replaces it + ( + "http://internal.local:8888/", + "https", + "external.com", + None, + "https://external.com", + ), + ], +) +def test_get_request_base_url_comprehensive( + base_url, x_forwarded_proto, x_forwarded_host, x_forwarded_port, expected_url +): + """Comprehensive test for get_request_base_url with various header combinations""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + get_request_base_url, + ) + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + # Create mock request + mock_request = MagicMock(spec=Request) + mock_request.base_url = base_url + + # Build headers dict + headers = {} + if x_forwarded_proto: + headers["X-Forwarded-Proto"] = x_forwarded_proto + if x_forwarded_host: + headers["X-Forwarded-Host"] = x_forwarded_host + if x_forwarded_port: + headers["X-Forwarded-Port"] = x_forwarded_port + + # Mock headers.get() to return our test values + def mock_get(header_name, default=None): + return headers.get(header_name, default) + + mock_request.headers.get = mock_get + + # Test the function + result = get_request_base_url(mock_request) + + # Verify result + assert result == expected_url, ( + f"Expected '{expected_url}' but got '{result}'\n" + f"Input: base_url={base_url}, " + f"X-Forwarded-Proto={x_forwarded_proto}, " + f"X-Forwarded-Host={x_forwarded_host}, " + f"X-Forwarded-Port={x_forwarded_port}" + ) diff --git a/tests/mcp_tests/test_aresponses_api_with_mcp.py b/tests/mcp_tests/test_aresponses_api_with_mcp.py index 6c8f51201d8..a7bbfef14af 100644 --- a/tests/mcp_tests/test_aresponses_api_with_mcp.py +++ b/tests/mcp_tests/test_aresponses_api_with_mcp.py @@ -683,9 +683,11 @@ async def test_streaming_responses_api_with_mcp_tools( Return the user the result of request 2 """ - # Skip test if ANTHROPIC_API_KEY is not set for anthropic models - if "anthropic" in model.lower() and not os.getenv("ANTHROPIC_API_KEY"): + # Skip test if required API keys are not set + if ("anthropic" in model.lower() or "claude" in model.lower()) and not os.getenv("ANTHROPIC_API_KEY"): pytest.skip("ANTHROPIC_API_KEY not set, skipping anthropic model test") + if ("gpt" in model.lower() or "openai" in model.lower()) and not os.getenv("OPENAI_API_KEY"): + pytest.skip("OPENAI_API_KEY not set, skipping openai model test") from unittest.mock import AsyncMock, patch diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index e8a1231c6fb..5785ff750fd 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -1105,9 +1105,26 @@ async def test_mcp_server_manager_config_integration_with_database(): test_manager.get_allowed_mcp_servers = mock_get_allowed_servers - # Test the method (this tests our second fix) - import asyncio + # Mock health_check_server to avoid real network calls that timeout + async def mock_health_check(server_id: str, mcp_auth_header=None): + server = test_manager.get_mcp_server_by_id(server_id) + if not server: + return None + return LiteLLM_MCPServerTable( + server_id=server_id, + server_name=server.name, + url=server.url, + transport=server.transport, + description=server.mcp_info.get("description") if server.mcp_info else None, + mcp_access_groups=server.access_groups, + status="healthy", + last_health_check=datetime.datetime.now(), + mcp_info=server.mcp_info, + ) + test_manager.health_check_server = mock_health_check + + # Test the method (this tests our second fix) servers_list = await test_manager.get_all_mcp_servers_with_health_and_teams( user_api_key_auth=mock_user_auth ) From a65ac2af1372db1f728a141bf185a02adb249d8c Mon Sep 17 00:00:00 2001 From: Harshit Jain <48647625+Harshit28j@users.noreply.github.com> Date: Mon, 26 Jan 2026 12:33:23 +0530 Subject: [PATCH 3/7] add callbacks and labels to prometheus (#19708) --- litellm/types/integrations/prometheus.py | 24 ++++++ .../test_prometheus_missing_metrics.py | 77 +++++++++++++++++++ 2 files changed, 101 insertions(+) create mode 100644 tests/test_litellm/integrations/test_prometheus_missing_metrics.py diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index fd9b722287d..74d7cbdfaea 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -150,6 +150,7 @@ class UserAPIKeyLabelNames(Enum): FALLBACK_MODEL = "fallback_model" ROUTE = "route" MODEL_GROUP = "model_group" + CALLBACK_NAME = "callback_name" DEFINED_PROMETHEUS_METRICS = Literal[ @@ -196,6 +197,9 @@ DEFINED_PROMETHEUS_METRICS = Literal[ "litellm_cache_hits_metric", "litellm_cache_misses_metric", "litellm_cached_tokens_metric", + "litellm_remaining_api_key_requests_for_model", + "litellm_remaining_api_key_tokens_for_model", + "litellm_callback_logging_failures_metric", ] @@ -436,6 +440,26 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.USER.value, ] + litellm_user_budget_remaining_hours_metric = [ + UserAPIKeyLabelNames.USER.value, + ] + + litellm_remaining_api_key_requests_for_model = [ + UserAPIKeyLabelNames.API_KEY_HASH.value, + UserAPIKeyLabelNames.API_KEY_ALIAS.value, + UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value, + ] + + litellm_remaining_api_key_tokens_for_model = [ + UserAPIKeyLabelNames.API_KEY_HASH.value, + UserAPIKeyLabelNames.API_KEY_ALIAS.value, + UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value, + ] + + litellm_callback_logging_failures_metric = [ + UserAPIKeyLabelNames.CALLBACK_NAME.value, + ] + # Add deployment metrics litellm_deployment_failure_responses = [ UserAPIKeyLabelNames.REQUESTED_MODEL.value, diff --git a/tests/test_litellm/integrations/test_prometheus_missing_metrics.py b/tests/test_litellm/integrations/test_prometheus_missing_metrics.py new file mode 100644 index 00000000000..7fcfb21ed4c --- /dev/null +++ b/tests/test_litellm/integrations/test_prometheus_missing_metrics.py @@ -0,0 +1,77 @@ +""" +Unit tests for the new Prometheus metrics that were previously missing from validation. + +Tests for: +- litellm_remaining_api_key_requests_for_model +- litellm_remaining_api_key_tokens_for_model +- litellm_callback_logging_failures_metric +""" +from typing import get_args +from litellm.types.integrations.prometheus import ( + DEFINED_PROMETHEUS_METRICS, + PrometheusMetricLabels, + UserAPIKeyLabelNames, +) + + +def test_new_metrics_in_defined_metrics(): + """ + Test that the new metrics are present in DEFINED_PROMETHEUS_METRICS. + """ + defined_metrics = get_args(DEFINED_PROMETHEUS_METRICS) + + new_metrics = [ + "litellm_remaining_api_key_requests_for_model", + "litellm_remaining_api_key_tokens_for_model", + "litellm_callback_logging_failures_metric", + ] + + for metric in new_metrics: + assert ( + metric in defined_metrics + ), f"{metric} should be in DEFINED_PROMETHEUS_METRICS" + + +def test_new_metrics_have_correct_labels(): + """ + Test that the new metrics have the correct labels defined. + """ + # Test API Key limits metrics labels + api_key_metrics = [ + "litellm_remaining_api_key_requests_for_model", + "litellm_remaining_api_key_tokens_for_model", + ] + + expected_api_key_labels = [ + UserAPIKeyLabelNames.API_KEY_HASH.value, + UserAPIKeyLabelNames.API_KEY_ALIAS.value, + UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value, + ] + + for metric in api_key_metrics: + labels = PrometheusMetricLabels.get_labels(metric) + for expected_label in expected_api_key_labels: + assert ( + expected_label in labels + ), f"{metric} should have label {expected_label}" + + # Test Callback failure metric labels + callback_metric = "litellm_callback_logging_failures_metric" + callback_labels = PrometheusMetricLabels.get_labels(callback_metric) + + assert ( + UserAPIKeyLabelNames.CALLBACK_NAME.value in callback_labels + ), f"{callback_metric} should have label {UserAPIKeyLabelNames.CALLBACK_NAME.value}" + + +def test_callback_name_label_definition(): + """ + Test that CALLBACK_NAME is defined correctly in UserAPIKeyLabelNames. + """ + assert UserAPIKeyLabelNames.CALLBACK_NAME.value == "callback_name" + + +if __name__ == "__main__": + test_new_metrics_in_defined_metrics() + test_new_metrics_have_correct_labels() + test_callback_name_label_definition() From 344ea3d9f26592fd50b894556d7bb92df45a5c51 Mon Sep 17 00:00:00 2001 From: Harshit Jain <48647625+Harshit28j@users.noreply.github.com> Date: Mon, 26 Jan 2026 12:37:19 +0530 Subject: [PATCH 4/7] feat: add clientip and user agent in metrics (#19717) * feat: add clientip and user agent in metrics * fix: lint errors * Add model id and other req labels --------- Co-authored-by: Krish Dholakia --- litellm/integrations/prometheus.py | 118 ++++++---- litellm/litellm_core_utils/litellm_logging.py | 12 +- litellm/proxy/litellm_pre_call_utils.py | 42 +++- litellm/types/integrations/prometheus.py | 62 ++++++ litellm/types/utils.py | 25 ++- .../test_prometheus_client_ip_user_agent.py | 203 ++++++++++++++++++ .../integrations/test_prometheus_labels.py | 140 ++++++++---- 7 files changed, 492 insertions(+), 110 deletions(-) create mode 100644 tests/test_litellm/integrations/test_prometheus_client_ip_user_agent.py diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index bafb0d88c82..abae36bf55b 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -229,14 +229,18 @@ class PrometheusLogger(CustomLogger): self.litellm_remaining_api_key_requests_for_model = self._gauge_factory( "litellm_remaining_api_key_requests_for_model", "Remaining Requests API Key can make for model (model based rpm limit on key)", - labelnames=["hashed_api_key", "api_key_alias", "model"], + labelnames=self.get_labels_for_metric( + "litellm_remaining_api_key_requests_for_model" + ), ) # Remaining MODEL TPM limit for API Key self.litellm_remaining_api_key_tokens_for_model = self._gauge_factory( "litellm_remaining_api_key_tokens_for_model", "Remaining Tokens API Key can make for model (model based tpm limit on key)", - labelnames=["hashed_api_key", "api_key_alias", "model"], + labelnames=self.get_labels_for_metric( + "litellm_remaining_api_key_tokens_for_model" + ), ) ######################################## @@ -373,15 +377,9 @@ class PrometheusLogger(CustomLogger): self.litellm_llm_api_failed_requests_metric = self._counter_factory( name="litellm_llm_api_failed_requests_metric", documentation="deprecated - use litellm_proxy_failed_requests_metric", - labelnames=[ - "end_user", - "hashed_api_key", - "api_key_alias", - "model", - "team", - "team_alias", - "user", - ], + labelnames=self.get_labels_for_metric( + "litellm_llm_api_failed_requests_metric" + ), ) self.litellm_requests_metric = self._counter_factory( @@ -954,6 +952,8 @@ class PrometheusLogger(CustomLogger): route=standard_logging_payload["metadata"].get( "user_api_key_request_route" ), + client_ip=standard_logging_payload["metadata"].get("requester_ip_address"), + user_agent=standard_logging_payload["metadata"].get("user_agent"), ) if ( @@ -1011,6 +1011,7 @@ class PrometheusLogger(CustomLogger): user_api_key_alias=user_api_key_alias, kwargs=kwargs, metadata=_metadata, + model_id=enum_values.model_id, ) # set latency metrics @@ -1245,6 +1246,7 @@ class PrometheusLogger(CustomLogger): user_api_key_alias: Optional[str], kwargs: dict, metadata: dict, + model_id: Optional[str] = None, ): from litellm.proxy.common_utils.callback_utils import ( get_model_group_from_litellm_kwargs, @@ -1266,11 +1268,11 @@ class PrometheusLogger(CustomLogger): ) self.litellm_remaining_api_key_requests_for_model.labels( - user_api_key, user_api_key_alias, model_group + user_api_key, user_api_key_alias, model_group, model_id ).set(remaining_requests) self.litellm_remaining_api_key_tokens_for_model.labels( - user_api_key, user_api_key_alias, model_group + user_api_key, user_api_key_alias, model_group, model_id ).set(remaining_tokens) def _set_latency_metrics( @@ -1365,14 +1367,14 @@ class PrometheusLogger(CustomLogger): standard_logging_payload: StandardLoggingPayload = kwargs.get( "standard_logging_object", {} ) - + if self._should_skip_metrics_for_invalid_key( kwargs=kwargs, standard_logging_payload=standard_logging_payload ): return - + model = kwargs.get("model", "") - + litellm_params = kwargs.get("litellm_params", {}) or {} get_end_user_id_for_cost_tracking = _get_cached_end_user_id_for_cost_tracking() @@ -1396,6 +1398,7 @@ class PrometheusLogger(CustomLogger): user_api_team, user_api_team_alias, user_id, + standard_logging_payload.get("model_id", ""), ).inc() self.set_llm_deployment_failure_metrics(kwargs) except Exception as e: @@ -1413,49 +1416,57 @@ class PrometheusLogger(CustomLogger): ) -> Optional[int]: """ Extract HTTP status code from various input formats for validation. - + This is a centralized helper to extract status code from different callback function signatures. Handles both ProxyException (uses 'code') and standard exceptions (uses 'status_code'). - + Args: kwargs: Dictionary potentially containing 'exception' key enum_values: Object with 'status_code' attribute exception: Exception object to extract status code from directly - + Returns: Status code as integer if found, None otherwise """ status_code = None - + # Try from enum_values first (most common in our callbacks) - if enum_values and hasattr(enum_values, "status_code") and enum_values.status_code: + if ( + enum_values + and hasattr(enum_values, "status_code") + and enum_values.status_code + ): try: status_code = int(enum_values.status_code) except (ValueError, TypeError): pass - + if not status_code and exception: # ProxyException uses 'code' attribute, other exceptions may use 'status_code' - status_code = getattr(exception, "status_code", None) or getattr(exception, "code", None) + status_code = getattr(exception, "status_code", None) or getattr( + exception, "code", None + ) if status_code is not None: try: status_code = int(status_code) except (ValueError, TypeError): status_code = None - + if not status_code and kwargs: exception_in_kwargs = kwargs.get("exception") if exception_in_kwargs: - status_code = getattr(exception_in_kwargs, "status_code", None) or getattr(exception_in_kwargs, "code", None) + status_code = getattr( + exception_in_kwargs, "status_code", None + ) or getattr(exception_in_kwargs, "code", None) if status_code is not None: try: status_code = int(status_code) except (ValueError, TypeError): status_code = None - + return status_code - + def _is_invalid_api_key_request( self, status_code: Optional[int], @@ -1463,23 +1474,23 @@ class PrometheusLogger(CustomLogger): ) -> bool: """ Determine if a request has an invalid API key based on status code and exception. - + This method prevents invalid authentication attempts from being recorded in Prometheus metrics. A 401 status code is the definitive indicator of authentication failure. Additionally, we check exception messages for authentication error patterns to catch cases where the exception hasn't been converted to a ProxyException yet. - + Args: status_code: HTTP status code (401 indicates authentication error) exception: Exception object to check for auth-related error messages - + Returns: True if the request has an invalid API key and metrics should be skipped, False otherwise """ if status_code == 401: return True - + # Handle cases where AssertionError is raised before conversion to ProxyException if exception is not None: exception_str = str(exception).lower() @@ -1492,9 +1503,9 @@ class PrometheusLogger(CustomLogger): ] if any(pattern in exception_str for pattern in auth_error_patterns): return True - + return False - + def _should_skip_metrics_for_invalid_key( self, kwargs: Optional[dict] = None, @@ -1505,18 +1516,18 @@ class PrometheusLogger(CustomLogger): ) -> bool: """ Determine if Prometheus metrics should be skipped for invalid API key requests. - + This is a centralized validation method that extracts status code and exception information from various callback function signatures and determines if the request represents an invalid API key attempt that should be filtered from metrics. - + Args: kwargs: Dictionary potentially containing exception and other data user_api_key_dict: User API key authentication object (currently unused) enum_values: Object with status_code attribute standard_logging_payload: Standard logging payload dictionary exception: Exception object to check directly - + Returns: True if metrics should be skipped (invalid key detected), False otherwise """ @@ -1525,17 +1536,17 @@ class PrometheusLogger(CustomLogger): enum_values=enum_values, exception=exception, ) - + if exception is None and kwargs: exception = kwargs.get("exception") - + if self._is_invalid_api_key_request(status_code, exception=exception): verbose_logger.debug( "Skipping Prometheus metrics for invalid API key request: " f"status_code={status_code}, exception={type(exception).__name__ if exception else None}" ) return True - + return False async def async_post_call_failure_hook( @@ -1576,6 +1587,10 @@ class PrometheusLogger(CustomLogger): litellm_params=request_data, proxy_server_request=request_data.get("proxy_server_request", {}), ) + _metadata = request_data.get("metadata", {}) or {} + model_id = _metadata.get("model_info", {}).get("id") or request_data.get( + "model_info", {} + ).get("id") enum_values = UserAPIKeyLabelValues( end_user=user_api_key_dict.end_user_id, user=user_api_key_dict.user_id, @@ -1590,6 +1605,9 @@ class PrometheusLogger(CustomLogger): exception_class=self._get_exception_class_name(original_exception), tags=_tags, route=user_api_key_dict.request_route, + client_ip=_metadata.get("requester_ip_address"), + user_agent=_metadata.get("user_agent"), + model_id=model_id, ) _labels = prometheus_label_factory( supported_enum_labels=self.get_labels_for_metric( @@ -1629,6 +1647,7 @@ class PrometheusLogger(CustomLogger): ): return + _metadata = data.get("metadata", {}) or {} enum_values = UserAPIKeyLabelValues( end_user=user_api_key_dict.end_user_id, hashed_api_key=user_api_key_dict.api_key, @@ -1644,6 +1663,8 @@ class PrometheusLogger(CustomLogger): litellm_params=data, proxy_server_request=data.get("proxy_server_request", {}), ), + client_ip=_metadata.get("requester_ip_address"), + user_agent=_metadata.get("user_agent"), ) _labels = prometheus_label_factory( supported_enum_labels=self.get_labels_for_metric( @@ -1684,7 +1705,7 @@ class PrometheusLogger(CustomLogger): exception = request_kwargs.get("exception", None) llm_provider = _litellm_params.get("custom_llm_provider", None) - + if self._should_skip_metrics_for_invalid_key( kwargs=request_kwargs, standard_logging_payload=standard_logging_payload, @@ -1716,6 +1737,10 @@ class PrometheusLogger(CustomLogger): "user_api_key_team_alias" ], tags=standard_logging_payload.get("request_tags", []), + client_ip=standard_logging_payload["metadata"].get( + "requester_ip_address" + ), + user_agent=standard_logging_payload["metadata"].get("user_agent"), ) """ @@ -2263,7 +2288,10 @@ class PrometheusLogger(CustomLogger): async def fetch_keys( page_size: int, page: int - ) -> Tuple[List[Union[str, UserAPIKeyAuth, LiteLLM_DeletedVerificationToken]], Optional[int]]: + ) -> Tuple[ + List[Union[str, UserAPIKeyAuth, LiteLLM_DeletedVerificationToken]], + Optional[int], + ]: key_list_response = await _list_key_helper( prisma_client=prisma_client, page=page, @@ -2379,12 +2407,16 @@ class PrometheusLogger(CustomLogger): # Get total user count total_users = await prisma_client.db.litellm_usertable.count() self.litellm_total_users_metric.set(total_users) - verbose_logger.debug(f"Prometheus: set litellm_total_users to {total_users}") + verbose_logger.debug( + f"Prometheus: set litellm_total_users to {total_users}" + ) # Get total team count total_teams = await prisma_client.db.litellm_teamtable.count() self.litellm_teams_count_metric.set(total_teams) - verbose_logger.debug(f"Prometheus: set litellm_teams_count to {total_teams}") + verbose_logger.debug( + f"Prometheus: set litellm_teams_count to {total_teams}" + ) except Exception as e: verbose_logger.exception( f"Error initializing user/team count metrics: {str(e)}" diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index fadeeffa9cc..e5412a650b7 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -335,7 +335,9 @@ class Logging(LiteLLMLoggingBaseClass): self.start_time = start_time # log the call start time self.call_type = call_type self.litellm_call_id = litellm_call_id - self.litellm_trace_id: str = litellm_trace_id if litellm_trace_id else str(uuid.uuid4()) + self.litellm_trace_id: str = ( + litellm_trace_id if litellm_trace_id else str(uuid.uuid4()) + ) self.function_id = function_id self.streaming_chunks: List[Any] = [] # for generating complete stream response self.sync_streaming_chunks: List[ @@ -544,7 +546,10 @@ class Logging(LiteLLMLoggingBaseClass): if "stream_options" in additional_params: self.stream_options = additional_params["stream_options"] ## check if custom pricing set ## - if any(litellm_params.get(key) is not None for key in _CUSTOM_PRICING_KEYS & litellm_params.keys()): + if any( + litellm_params.get(key) is not None + for key in _CUSTOM_PRICING_KEYS & litellm_params.keys() + ): self.custom_pricing = True if "custom_llm_provider" in self.model_call_details: @@ -4453,6 +4458,7 @@ class StandardLoggingPayloadSetup: user_api_key_request_route=None, spend_logs_metadata=None, requester_ip_address=None, + user_agent=None, requester_metadata=None, prompt_management_metadata=prompt_management_metadata, applied_guardrails=applied_guardrails, @@ -5138,6 +5144,7 @@ def get_standard_logging_object_payload( model_group=_model_group, model_id=_model_id, requester_ip_address=clean_metadata.get("requester_ip_address", None), + user_agent=clean_metadata.get("user_agent", None), messages=StandardLoggingPayloadSetup.append_system_prompt_messages( kwargs=kwargs, messages=kwargs.get("messages") ), @@ -5203,6 +5210,7 @@ def get_standard_logging_metadata( user_api_key_team_alias=None, spend_logs_metadata=None, requester_ip_address=None, + user_agent=None, requester_metadata=None, user_api_key_end_user_id=None, prompt_management_metadata=None, diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 064538ef3bf..02dcb25c82f 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -846,7 +846,9 @@ async def add_litellm_data_to_request( # noqa: PLR0915 # Add headers to metadata for guardrails to access (fixes #17477) # Guardrails use metadata["headers"] to access request headers (e.g., User-Agent) - if _metadata_variable_name in data and isinstance(data[_metadata_variable_name], dict): + if _metadata_variable_name in data and isinstance( + data[_metadata_variable_name], dict + ): data[_metadata_variable_name]["headers"] = _headers # check for forwardable headers @@ -1002,7 +1004,9 @@ async def add_litellm_data_to_request( # noqa: PLR0915 # User spend, budget - used by prometheus.py # Follow same pattern as team and API key budgets - data[_metadata_variable_name]["user_api_key_user_spend"] = user_api_key_dict.user_spend + data[_metadata_variable_name][ + "user_api_key_user_spend" + ] = user_api_key_dict.user_spend data[_metadata_variable_name][ "user_api_key_user_max_budget" ] = user_api_key_dict.user_max_budget @@ -1029,8 +1033,8 @@ async def add_litellm_data_to_request( # noqa: PLR0915 ## [Enterprise Only] # Add User-IP Address requester_ip_address = "" - if premium_user is True: - # Only set the IP Address for Enterprise Users + if True: # Always set the IP Address if available + # logic for tracking IP Address # logic for tracking IP Address if ( @@ -1050,6 +1054,16 @@ async def add_litellm_data_to_request( # noqa: PLR0915 requester_ip_address = request.client.host data[_metadata_variable_name]["requester_ip_address"] = requester_ip_address + # Add User-Agent + user_agent = "" + if ( + request is not None + and hasattr(request, "headers") + and "user-agent" in request.headers + ): + user_agent = request.headers["user-agent"] + data[_metadata_variable_name]["user_agent"] = user_agent + # Check if using tag based routing tags = LiteLLMProxyRequestSetup.add_request_tag_to_metadata( llm_router=llm_router, @@ -1532,7 +1546,9 @@ def add_guardrails_from_policy_engine( f"policy_count={len(registry.get_all_policies())}" ) if not registry.is_initialized(): - verbose_proxy_logger.debug("Policy engine not initialized, skipping policy matching") + verbose_proxy_logger.debug( + "Policy engine not initialized, skipping policy matching" + ) return # Build context from request @@ -1550,13 +1566,17 @@ def add_guardrails_from_policy_engine( # Get matching policies via attachments matching_policy_names = PolicyMatcher.get_matching_policies(context=context) - verbose_proxy_logger.debug(f"Policy engine: matched policies via attachments: {matching_policy_names}") + verbose_proxy_logger.debug( + f"Policy engine: matched policies via attachments: {matching_policy_names}" + ) # Combine attachment-based policies with dynamic request body policies all_policy_names = set(matching_policy_names) if request_body_policies and isinstance(request_body_policies, list): all_policy_names.update(request_body_policies) - verbose_proxy_logger.debug(f"Policy engine: added dynamic policies from request body: {request_body_policies}") + verbose_proxy_logger.debug( + f"Policy engine: added dynamic policies from request body: {request_body_policies}" + ) if not all_policy_names: return @@ -1567,7 +1587,9 @@ def add_guardrails_from_policy_engine( context=context, ) - verbose_proxy_logger.debug(f"Policy engine: applied policies (conditions matched): {applied_policy_names}") + verbose_proxy_logger.debug( + f"Policy engine: applied policies (conditions matched): {applied_policy_names}" + ) # Track applied policies in metadata for response headers for policy_name in applied_policy_names: @@ -1578,7 +1600,9 @@ def add_guardrails_from_policy_engine( # Resolve guardrails from matching policies resolved_guardrails = PolicyResolver.resolve_guardrails_for_context(context=context) - verbose_proxy_logger.debug(f"Policy engine: resolved guardrails: {resolved_guardrails}") + verbose_proxy_logger.debug( + f"Policy engine: resolved guardrails: {resolved_guardrails}" + ) if not resolved_guardrails: return diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index 74d7cbdfaea..ea9c9bd325d 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -150,6 +150,8 @@ class UserAPIKeyLabelNames(Enum): FALLBACK_MODEL = "fallback_model" ROUTE = "route" MODEL_GROUP = "model_group" + CLIENT_IP = "client_ip" + USER_AGENT = "user_agent" CALLBACK_NAME = "callback_name" @@ -199,6 +201,7 @@ DEFINED_PROMETHEUS_METRICS = Literal[ "litellm_cached_tokens_metric", "litellm_remaining_api_key_requests_for_model", "litellm_remaining_api_key_tokens_for_model", + "litellm_llm_api_failed_requests_metric", "litellm_callback_logging_failures_metric", ] @@ -213,6 +216,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.REQUESTED_MODEL.value, UserAPIKeyLabelNames.END_USER.value, UserAPIKeyLabelNames.USER.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_llm_api_time_to_first_token_metric = [ @@ -221,6 +225,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.API_KEY_ALIAS.value, UserAPIKeyLabelNames.TEAM.value, UserAPIKeyLabelNames.TEAM_ALIAS.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_request_total_latency_metric = [ @@ -232,6 +237,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.TEAM_ALIAS.value, UserAPIKeyLabelNames.USER.value, UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_request_queue_time_seconds = [ @@ -243,6 +249,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.TEAM_ALIAS.value, UserAPIKeyLabelNames.USER.value, UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] # Guardrail metrics - these use custom labels (guardrail_name, status, error_type, hook_type) @@ -262,6 +269,9 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.STATUS_CODE.value, UserAPIKeyLabelNames.USER_EMAIL.value, UserAPIKeyLabelNames.ROUTE.value, + UserAPIKeyLabelNames.CLIENT_IP.value, + UserAPIKeyLabelNames.USER_AGENT.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_proxy_failed_requests_metric = [ @@ -276,6 +286,9 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.EXCEPTION_STATUS.value, UserAPIKeyLabelNames.EXCEPTION_CLASS.value, UserAPIKeyLabelNames.ROUTE.value, + UserAPIKeyLabelNames.CLIENT_IP.value, + UserAPIKeyLabelNames.USER_AGENT.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_deployment_latency_per_output_token = [ @@ -296,6 +309,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.v2_LITELLM_MODEL_NAME.value, UserAPIKeyLabelNames.API_KEY_HASH.value, UserAPIKeyLabelNames.API_KEY_ALIAS.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_remaining_requests_metric = [ @@ -305,6 +319,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.v2_LITELLM_MODEL_NAME.value, UserAPIKeyLabelNames.API_KEY_HASH.value, UserAPIKeyLabelNames.API_KEY_ALIAS.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_remaining_tokens_metric = [ @@ -314,6 +329,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.v2_LITELLM_MODEL_NAME.value, UserAPIKeyLabelNames.API_KEY_HASH.value, UserAPIKeyLabelNames.API_KEY_ALIAS.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_requests_metric = [ @@ -325,6 +341,9 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.TEAM_ALIAS.value, UserAPIKeyLabelNames.USER.value, UserAPIKeyLabelNames.USER_EMAIL.value, + UserAPIKeyLabelNames.CLIENT_IP.value, + UserAPIKeyLabelNames.USER_AGENT.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_spend_metric = [ @@ -336,6 +355,9 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.TEAM_ALIAS.value, UserAPIKeyLabelNames.USER.value, UserAPIKeyLabelNames.USER_EMAIL.value, + UserAPIKeyLabelNames.CLIENT_IP.value, + UserAPIKeyLabelNames.USER_AGENT.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_input_tokens_metric = [ @@ -348,6 +370,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.USER.value, UserAPIKeyLabelNames.USER_EMAIL.value, UserAPIKeyLabelNames.REQUESTED_MODEL.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_total_tokens_metric = [ @@ -360,6 +383,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.USER.value, UserAPIKeyLabelNames.USER_EMAIL.value, UserAPIKeyLabelNames.REQUESTED_MODEL.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_output_tokens_metric = [ @@ -372,6 +396,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.USER.value, UserAPIKeyLabelNames.USER_EMAIL.value, UserAPIKeyLabelNames.REQUESTED_MODEL.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_deployment_state = [ @@ -398,6 +423,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.TEAM_ALIAS.value, UserAPIKeyLabelNames.EXCEPTION_STATUS.value, UserAPIKeyLabelNames.EXCEPTION_CLASS.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_deployment_failed_fallbacks = litellm_deployment_successful_fallbacks @@ -473,6 +499,8 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.API_KEY_ALIAS.value, UserAPIKeyLabelNames.TEAM.value, UserAPIKeyLabelNames.TEAM_ALIAS.value, + UserAPIKeyLabelNames.CLIENT_IP.value, + UserAPIKeyLabelNames.USER_AGENT.value, ] litellm_deployment_total_requests = [ @@ -485,10 +513,37 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.API_KEY_ALIAS.value, UserAPIKeyLabelNames.TEAM.value, UserAPIKeyLabelNames.TEAM_ALIAS.value, + UserAPIKeyLabelNames.CLIENT_IP.value, + UserAPIKeyLabelNames.USER_AGENT.value, ] litellm_deployment_success_responses = litellm_deployment_total_requests + litellm_remaining_api_key_requests_for_model = [ + UserAPIKeyLabelNames.API_KEY_HASH.value, + UserAPIKeyLabelNames.API_KEY_ALIAS.value, + UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value, + UserAPIKeyLabelNames.MODEL_ID.value, + ] + + litellm_remaining_api_key_tokens_for_model = [ + UserAPIKeyLabelNames.API_KEY_HASH.value, + UserAPIKeyLabelNames.API_KEY_ALIAS.value, + UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value, + UserAPIKeyLabelNames.MODEL_ID.value, + ] + + litellm_llm_api_failed_requests_metric = [ + UserAPIKeyLabelNames.END_USER.value, + UserAPIKeyLabelNames.API_KEY_HASH.value, + UserAPIKeyLabelNames.API_KEY_ALIAS.value, + UserAPIKeyLabelNames.v1_LITELLM_MODEL_NAME.value, + UserAPIKeyLabelNames.TEAM.value, + UserAPIKeyLabelNames.TEAM_ALIAS.value, + UserAPIKeyLabelNames.USER.value, + UserAPIKeyLabelNames.MODEL_ID.value, + ] + # Buffer monitoring metrics - these typically don't need additional labels litellm_pod_lock_manager_size: List[str] = [] @@ -509,6 +564,7 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.TEAM_ALIAS.value, UserAPIKeyLabelNames.END_USER.value, UserAPIKeyLabelNames.USER.value, + UserAPIKeyLabelNames.MODEL_ID.value, ] litellm_cache_hits_metric = _cache_metric_labels @@ -601,6 +657,12 @@ class UserAPIKeyLabelValues(BaseModel): route: Annotated[ Optional[str], Field(..., alias=UserAPIKeyLabelNames.ROUTE.value) ] = None + client_ip: Annotated[ + Optional[str], Field(..., alias=UserAPIKeyLabelNames.CLIENT_IP.value) + ] = None + user_agent: Annotated[ + Optional[str], Field(..., alias=UserAPIKeyLabelNames.USER_AGENT.value) + ] = None class PrometheusMetricsConfig(BaseModel): diff --git a/litellm/types/utils.py b/litellm/types/utils.py index cd797dd1e54..2ac5443b3bc 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3,25 +3,26 @@ import time from enum import Enum from typing import TYPE_CHECKING, Any, Dict, List, Literal, Mapping, Optional, Union -from aiohttp import FormData from openai._models import BaseModel as OpenAIObject -from openai.types.audio.transcription_create_params import FileTypes # type: ignore -from openai.types.chat.chat_completion import ChatCompletion +from openai.types.audio.transcription_create_params import FileTypes as FileTypes # type: ignore +from openai.types.chat.chat_completion import ChatCompletion as ChatCompletion from openai.types.completion_usage import ( CompletionTokensDetails, CompletionUsage, PromptTokensDetails, ) from openai.types.moderation import ( - Categories, - CategoryAppliedInputTypes, - CategoryScores, + Categories as Categories, + CategoryAppliedInputTypes as CategoryAppliedInputTypes, + CategoryScores as CategoryScores, +) +from openai.types.moderation_create_response import ( + Moderation as Moderation, + ModerationCreateResponse as ModerationCreateResponse, ) -from openai.types.moderation_create_response import Moderation, ModerationCreateResponse from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, model_validator -from typing_extensions import Callable, Dict, Required, TypedDict, override +from typing_extensions import Required, TypedDict -import litellm from litellm._uuid import uuid from litellm.types.llms.base import ( BaseLiteLLMOpenAIResponseObject, @@ -52,7 +53,7 @@ from .llms.openai import ( ResponsesAPIResponse, WebSearchOptions, ) -from .rerank import RerankResponse +from .rerank import RerankResponse as RerankResponse if TYPE_CHECKING: from .vector_stores import VectorStoreSearchResponse @@ -1411,7 +1412,7 @@ class Usage(SafeAttributeModel, CompletionUsage): prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None """Breakdown of tokens used in the prompt.""" - def __init__( + def __init__( # noqa: PLR0915 self, prompt_tokens: Optional[int] = None, completion_tokens: Optional[int] = None, @@ -2501,6 +2502,7 @@ class StandardLoggingMetadata(StandardLoggingUserAPIKeyMetadata): dict ] # special param to log k,v pairs to spendlogs for a call requester_ip_address: Optional[str] + user_agent: Optional[str] requester_metadata: Optional[dict] requester_custom_headers: Optional[ Dict[str, str] @@ -2686,6 +2688,7 @@ class StandardLoggingPayload(TypedDict): request_tags: list end_user: Optional[str] requester_ip_address: Optional[str] + user_agent: Optional[str] messages: Optional[Union[str, list, dict]] response: Optional[Union[str, list, dict]] error_str: Optional[str] diff --git a/tests/test_litellm/integrations/test_prometheus_client_ip_user_agent.py b/tests/test_litellm/integrations/test_prometheus_client_ip_user_agent.py new file mode 100644 index 00000000000..4a9fa3de5fd --- /dev/null +++ b/tests/test_litellm/integrations/test_prometheus_client_ip_user_agent.py @@ -0,0 +1,203 @@ +import pytest +from unittest.mock import MagicMock, patch +from litellm.integrations.prometheus import PrometheusLogger +from litellm.types.integrations.prometheus import ( + UserAPIKeyLabelValues, +) +from litellm.proxy._types import UserAPIKeyAuth + + +@pytest.mark.asyncio +async def test_async_post_call_failure_hook_includes_client_ip_user_agent(): + """ + Test that async_post_call_failure_hook includes client_ip and user_agent in UserAPIKeyLabelValues + """ + # Mocking + # Mocking + with patch( + "litellm.integrations.prometheus.PrometheusLogger.__init__", return_value=None + ): + logger = PrometheusLogger() + # Initialize attributes manually as __init__ is mocked + logger.litellm_proxy_failed_requests_metric = MagicMock() + logger.litellm_proxy_total_requests_metric = MagicMock() + logger.get_labels_for_metric = MagicMock( + return_value=["client_ip", "user_agent"] + ) + + request_data = { + "model": "gpt-4", + "metadata": { + "requester_ip_address": "127.0.0.1", + "user_agent": "test-agent", + }, + } + user_api_key_dict = UserAPIKeyAuth(token="test_token") + original_exception = Exception("Test exception") + + # Mock prometheus_label_factory to inspect arguments + with patch( + "litellm.integrations.prometheus.prometheus_label_factory" + ) as mock_label_factory: + mock_label_factory.return_value = {} + + await logger.async_post_call_failure_hook( + request_data=request_data, + original_exception=original_exception, + user_api_key_dict=user_api_key_dict, + ) + + # Verification + assert mock_label_factory.call_count >= 1 + + # Check calls + calls = mock_label_factory.call_args_list + found = False + for call in calls: + kwargs = call.kwargs + enum_values = kwargs.get("enum_values") + if isinstance(enum_values, UserAPIKeyLabelValues): + if ( + enum_values.client_ip == "127.0.0.1" + and enum_values.user_agent == "test-agent" + ): + found = True + break + + assert ( + found + ), "UserAPIKeyLabelValues should contain client_ip='127.0.0.1' and user_agent='test-agent'" + + +@pytest.mark.asyncio +async def test_async_post_call_success_hook_includes_client_ip_user_agent(): + """ + Test that async_post_call_success_hook includes client_ip and user_agent in UserAPIKeyLabelValues + """ + # Mocking + # Mocking + with patch( + "litellm.integrations.prometheus.PrometheusLogger.__init__", return_value=None + ): + logger = PrometheusLogger() + logger.litellm_proxy_total_requests_metric = MagicMock() + logger.get_labels_for_metric = MagicMock( + return_value=["client_ip", "user_agent"] + ) + + data = { + "model": "gpt-4", + "metadata": { + "requester_ip_address": "192.168.1.1", + "user_agent": "success-agent", + }, + } + user_api_key_dict = UserAPIKeyAuth(token="test_token") + response = MagicMock() + + # Mock prometheus_label_factory to inspect arguments + with patch( + "litellm.integrations.prometheus.prometheus_label_factory" + ) as mock_label_factory: + mock_label_factory.return_value = {} + + await logger.async_post_call_success_hook( + data=data, + user_api_key_dict=user_api_key_dict, + response=response, + ) + + # Verification + assert mock_label_factory.call_count >= 1 + + # Check calls + calls = mock_label_factory.call_args_list + found = False + for call in calls: + kwargs = call.kwargs + enum_values = kwargs.get("enum_values") + if isinstance(enum_values, UserAPIKeyLabelValues): + if ( + enum_values.client_ip == "192.168.1.1" + and enum_values.user_agent == "success-agent" + ): + found = True + break + + assert ( + found + ), "UserAPIKeyLabelValues should contain client_ip='192.168.1.1' and user_agent='success-agent'" + + +def test_set_llm_deployment_failure_metrics_includes_client_ip_user_agent(): + """ + Test that set_llm_deployment_failure_metrics includes client_ip and user_agent in UserAPIKeyLabelValues + """ + # Mocking + # Mocking + with patch( + "litellm.integrations.prometheus.PrometheusLogger.__init__", return_value=None + ): + logger = PrometheusLogger() + logger.litellm_deployment_failure_responses = MagicMock() + logger.litellm_deployment_total_requests = MagicMock() + logger.get_labels_for_metric = MagicMock( + return_value=["client_ip", "user_agent"] + ) + logger.set_deployment_partial_outage = MagicMock() + + request_kwargs = { + "model": "gpt-4", + "standard_logging_object": { + "metadata": { + "requester_ip_address": "10.0.0.1", + "user_agent": "failure-deployment", + "user_api_key_team_id": "team_1", + "user_api_key_team_alias": "team_alias_1", + "user_api_key_alias": "key_alias_1", + }, + "model_group": "group_1", + "api_base": "http://api.base", + "model_id": "model_1", + }, + "litellm_params": {}, + "exception": Exception("Deployment failure"), + } + + # Mock prometheus_label_factory to inspect arguments + with patch( + "litellm.integrations.prometheus.prometheus_label_factory" + ) as mock_label_factory: + mock_label_factory.return_value = {} + + logger.set_llm_deployment_failure_metrics(request_kwargs=request_kwargs) + + # Verification + assert mock_label_factory.call_count >= 1 + + # Check calls + calls = mock_label_factory.call_args_list + found = False + for call in calls: + kwargs = call.kwargs + enum_values = kwargs.get("enum_values") + if isinstance(enum_values, UserAPIKeyLabelValues): + if ( + enum_values.client_ip == "10.0.0.1" + and enum_values.user_agent == "failure-deployment" + ): + found = True + break + + assert ( + found + ), "UserAPIKeyLabelValues should contain client_ip='10.0.0.1' and user_agent='failure-deployment'" + + +if __name__ == "__main__": + import asyncio + + asyncio.run(test_async_post_call_failure_hook_includes_client_ip_user_agent()) + asyncio.run(test_async_post_call_success_hook_includes_client_ip_user_agent()) + test_set_llm_deployment_failure_metrics_includes_client_ip_user_agent() + print("✅ All client_ip and user_agent tests passed!") diff --git a/tests/test_litellm/integrations/test_prometheus_labels.py b/tests/test_litellm/integrations/test_prometheus_labels.py index c0b863ef6ee..a83bc1df1e1 100644 --- a/tests/test_litellm/integrations/test_prometheus_labels.py +++ b/tests/test_litellm/integrations/test_prometheus_labels.py @@ -26,15 +26,49 @@ def test_user_email_in_required_metrics(): "litellm_input_tokens_metric", "litellm_output_tokens_metric", "litellm_requests_metric", - "litellm_spend_metric" + "litellm_spend_metric", ] for metric_name in metrics_with_user_email: labels = PrometheusMetricLabels.get_labels(metric_name) - assert user_email_label in labels, f"Metric {metric_name} should contain user_email label" + assert ( + user_email_label in labels + ), f"Metric {metric_name} should contain user_email label" print(f"✅ {metric_name} contains user_email label") +def test_model_id_in_required_metrics(): + """ + Test that model_id label is present in all the metrics that should have it + """ + model_id_label = UserAPIKeyLabelNames.MODEL_ID.value + + # Metrics that should have model_id + metrics_with_model_id = [ + "litellm_proxy_total_requests_metric", + "litellm_proxy_failed_requests_metric", + "litellm_input_tokens_metric", + "litellm_output_tokens_metric", + "litellm_requests_metric", + "litellm_spend_metric", + "litellm_llm_api_latency_metric", + "litellm_remaining_requests_metric", + "litellm_deployment_successful_fallbacks", + "litellm_cache_hits_metric", + "litellm_cache_misses_metric", + "litellm_remaining_api_key_requests_for_model", + "litellm_remaining_api_key_tokens_for_model", + "litellm_llm_api_failed_requests_metric", + ] + + for metric_name in metrics_with_model_id: + labels = PrometheusMetricLabels.get_labels(metric_name) + assert ( + model_id_label in labels + ), f"Metric {metric_name} should contain model_id label" + print(f"✅ {metric_name} contains model_id label") + + def test_user_email_label_exists(): """Test that the USER_EMAIL label is properly defined""" assert UserAPIKeyLabelNames.USER_EMAIL.value == "user_email" @@ -52,12 +86,14 @@ def test_prometheus_metric_labels_structure(): "litellm_proxy_failed_requests_metric", "litellm_input_tokens_metric", "litellm_output_tokens_metric", - "litellm_spend_metric" + "litellm_spend_metric", ] for metric_name in test_metrics: # Check metric is in DEFINED_PROMETHEUS_METRICS - assert metric_name in get_args(DEFINED_PROMETHEUS_METRICS), f"{metric_name} should be in DEFINED_PROMETHEUS_METRICS" + assert metric_name in get_args( + DEFINED_PROMETHEUS_METRICS + ), f"{metric_name} should be in DEFINED_PROMETHEUS_METRICS" # Check labels can be retrieved labels = PrometheusMetricLabels.get_labels(metric_name) @@ -74,11 +110,11 @@ def test_route_normalization_for_responses_api(): """ Test that route normalization prevents high cardinality in Prometheus metrics for the /v1/responses/{response_id} endpoint. - + Issue: https://github.com/BerriAI/litellm/issues/XXXX Each unique response ID was creating a separate metric line, causing the /metrics endpoint to grow to ~30MB and take ~40 seconds to respond. - + Fix: Routes are normalized to collapse dynamic IDs into placeholders. """ from litellm.proxy.auth.auth_utils import normalize_request_route @@ -91,43 +127,53 @@ def test_route_normalization_for_responses_api(): ("/v1/responses/resp_abc123", "/v1/responses/{response_id}"), ("/v1/responses/litellm_poll_xyz", "/v1/responses/{response_id}"), ] - + for original, expected in responses_routes: normalized = normalize_request_route(original) - assert normalized == expected, \ - f"Failed: {original} -> {normalized} (expected {expected})" - + assert ( + normalized == expected + ), f"Failed: {original} -> {normalized} (expected {expected})" + # Verify cardinality reduction - unique_normalized = set(normalize_request_route(route) for route, _ in responses_routes) - assert len(unique_normalized) == 1, \ - f"Expected 1 unique normalized route, got {len(unique_normalized)}: {unique_normalized}" - - print(f"✅ Responses API routes: {len(responses_routes)} different IDs normalized to 1 metric label") - + unique_normalized = set( + normalize_request_route(route) for route, _ in responses_routes + ) + assert ( + len(unique_normalized) == 1 + ), f"Expected 1 unique normalized route, got {len(unique_normalized)}: {unique_normalized}" + + print( + f"✅ Responses API routes: {len(responses_routes)} different IDs normalized to 1 metric label" + ) + def test_route_normalization_for_sub_routes(): """Test that sub-routes like /cancel and /input_items are normalized correctly""" from litellm.proxy.auth.auth_utils import normalize_request_route - + sub_routes = [ ("/v1/responses/id1/cancel", "/v1/responses/{response_id}/cancel"), ("/v1/responses/id2/cancel", "/v1/responses/{response_id}/cancel"), ("/v1/responses/id3/input_items", "/v1/responses/{response_id}/input_items"), - ("/openai/v1/responses/id4/input_items", "/openai/v1/responses/{response_id}/input_items"), + ( + "/openai/v1/responses/id4/input_items", + "/openai/v1/responses/{response_id}/input_items", + ), ] - + for original, expected in sub_routes: normalized = normalize_request_route(original) - assert normalized == expected, \ - f"Failed: {original} -> {normalized} (expected {expected})" - + assert ( + normalized == expected + ), f"Failed: {original} -> {normalized} (expected {expected})" + print("✅ Sub-routes normalized correctly") def test_route_normalization_preserves_static_routes(): """Test that static routes are not affected by normalization""" from litellm.proxy.auth.auth_utils import normalize_request_route - + static_routes = [ "/chat/completions", "/v1/chat/completions", @@ -137,46 +183,47 @@ def test_route_normalization_preserves_static_routes(): "/v1/models", "/v1/responses", # List endpoint without ID ] - + for route in static_routes: normalized = normalize_request_route(route) - assert normalized == route, \ - f"Static route should not be modified: {route} -> {normalized}" - + assert ( + normalized == route + ), f"Static route should not be modified: {route} -> {normalized}" + print(f"✅ {len(static_routes)} static routes preserved") def test_route_normalization_other_dynamic_apis(): """Test normalization for other OpenAI-compatible APIs with dynamic IDs""" from litellm.proxy.auth.auth_utils import normalize_request_route - + test_cases = [ # Threads API ("/v1/threads/thread_123", "/v1/threads/{thread_id}"), ("/v1/threads/thread_abc/messages", "/v1/threads/{thread_id}/messages"), - ("/v1/threads/thread_abc/runs/run_123", "/v1/threads/{thread_id}/runs/{run_id}"), - + ( + "/v1/threads/thread_abc/runs/run_123", + "/v1/threads/{thread_id}/runs/{run_id}", + ), # Vector Stores API ("/v1/vector_stores/vs_123", "/v1/vector_stores/{vector_store_id}"), ("/v1/vector_stores/vs_123/files", "/v1/vector_stores/{vector_store_id}/files"), - # Assistants API ("/v1/assistants/asst_123", "/v1/assistants/{assistant_id}"), - # Files API ("/v1/files/file_123", "/v1/files/{file_id}"), ("/v1/files/file_123/content", "/v1/files/{file_id}/content"), - # Batches API ("/v1/batches/batch_123", "/v1/batches/{batch_id}"), ("/v1/batches/batch_123/cancel", "/v1/batches/{batch_id}/cancel"), ] - + for original, expected in test_cases: normalized = normalize_request_route(original) - assert normalized == expected, \ - f"Failed: {original} -> {normalized} (expected {expected})" - + assert ( + normalized == expected + ), f"Failed: {original} -> {normalized} (expected {expected})" + print(f"✅ {len(test_cases)} other API routes normalized correctly") @@ -195,26 +242,29 @@ def test_prometheus_metrics_use_normalized_routes(): # Create a mock PrometheusLogger prometheus_logger = MagicMock() - prometheus_logger.get_labels_for_metric = PrometheusLogger.get_labels_for_metric.__get__(prometheus_logger) - + prometheus_logger.get_labels_for_metric = ( + PrometheusLogger.get_labels_for_metric.__get__(prometheus_logger) + ) + # Test with a normalized route enum_values = UserAPIKeyLabelValues( route="/v1/responses/{response_id}", # Normalized route status_code="200", requested_model="gpt-4", ) - + labels = prometheus_label_factory( supported_enum_labels=prometheus_logger.get_labels_for_metric( metric_name="litellm_proxy_total_requests_metric" ), enum_values=enum_values, ) - + # Verify the route is normalized in labels - assert labels["route"] == "/v1/responses/{response_id}", \ - f"Expected normalized route in labels, got: {labels.get('route')}" - + assert ( + labels["route"] == "/v1/responses/{response_id}" + ), f"Expected normalized route in labels, got: {labels.get('route')}" + print("✅ Prometheus metrics use normalized routes in labels") @@ -227,4 +277,4 @@ if __name__ == "__main__": test_route_normalization_preserves_static_routes() test_route_normalization_other_dynamic_apis() test_prometheus_metrics_use_normalized_routes() - print("\n✅ All prometheus label tests passed!") \ No newline at end of file + print("\n✅ All prometheus label tests passed!") From 79603b9c3a1c16ddae78ad00609e5268a5b07a9e Mon Sep 17 00:00:00 2001 From: Harshit Jain <48647625+Harshit28j@users.noreply.github.com> Date: Mon, 26 Jan 2026 12:38:16 +0530 Subject: [PATCH 5/7] fix: optimize logo fetching and resolve mcp import blockers (#19719) --- litellm/experimental_mcp_client/client.py | 6 +- .../mcp_server/mcp_server_manager.py | 31 +++++-- .../proxy/_experimental/mcp_server/server.py | 6 +- .../mcp_management_endpoints.py | 18 +++- litellm/proxy/proxy_server.py | 40 ++++++--- tests/proxy_unit_tests/test_get_image.py | 89 +++++++++++++++++++ 6 files changed, 163 insertions(+), 27 deletions(-) create mode 100644 tests/proxy_unit_tests/test_get_image.py diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index ab576d49f3c..50ec5cab429 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -10,7 +10,11 @@ import httpx from mcp import ClientSession, ReadResourceResult, Resource, StdioServerParameters from mcp.client.sse import sse_client from mcp.client.stdio import stdio_client -from mcp.client.streamable_http import streamable_http_client + +try: + from mcp.client.streamable_http import streamable_http_client # type: ignore +except ImportError: + streamable_http_client = None from mcp.types import CallToolRequestParams as MCPCallToolRequestParams from mcp.types import CallToolResult as MCPCallToolResult from mcp.types import ( diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index e0217cd9e00..7fceec005e6 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -63,7 +63,20 @@ from litellm.types.mcp_server.mcp_server_manager import ( MCPOAuthMetadata, MCPServer, ) -from mcp.shared.tool_name_validation import SEP_986_URL, validate_tool_name + +try: + from mcp.shared.tool_name_validation import SEP_986_URL, validate_tool_name # type: ignore +except ImportError: + SEP_986_URL = "https://github.com/modelcontextprotocol/protocol/blob/main/proposals/0001-tool-name-validation.md" + + def validate_tool_name(name: str): + from pydantic import BaseModel + + class MockResult(BaseModel): + is_valid: bool = True + warnings: list = [] + + return MockResult() # Probe includes characters on both sides of the separator to mimic real prefixed tool names. @@ -90,7 +103,9 @@ def _warn_on_server_name_fields( if result.is_valid: return - warning_text = "; ".join(result.warnings) if result.warnings else "Validation failed" + warning_text = ( + "; ".join(result.warnings) if result.warnings else "Validation failed" + ) verbose_logger.warning( "MCP server '%s' has invalid %s '%s': %s", server_id, @@ -103,7 +118,6 @@ def _warn_on_server_name_fields( _warn("server_name", server_name) - def _deserialize_json_dict(data: Any) -> Optional[Dict[str, str]]: """ Deserialize optional JSON mappings stored in the database. @@ -391,10 +405,13 @@ class MCPServerManager: # Note: `extra_headers` on MCPServer is a List[str] of header names to forward # from the client request (not available in this OpenAPI tool generation step). # `static_headers` is a dict of concrete headers to always send. - headers = merge_mcp_headers( - extra_headers=headers, - static_headers=server.static_headers, - ) or {} + headers = ( + merge_mcp_headers( + extra_headers=headers, + static_headers=server.static_headers, + ) + or {} + ) verbose_logger.debug( f"Using headers for OpenAPI tools (excluding sensitive values): " diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 03652ae155e..bd7870a0fa0 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -74,7 +74,11 @@ if MCP_AVAILABLE: AuthContextMiddleware, auth_context_var, ) - from mcp.server.streamable_http_manager import StreamableHTTPSessionManager + + try: + from mcp.server.streamable_http_manager import StreamableHTTPSessionManager + except ImportError: + StreamableHTTPSessionManager = None # type: ignore from mcp.types import ( CallToolResult, EmbeddedResource, diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index a15f47d13bd..83d7f3fde4c 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -56,7 +56,19 @@ except ImportError as e: MCP_AVAILABLE = False if MCP_AVAILABLE: - from mcp.shared.tool_name_validation import validate_tool_name + try: + from mcp.shared.tool_name_validation import validate_tool_name # type: ignore + except ImportError: + + def validate_tool_name(name: str): + from pydantic import BaseModel + + class MockResult(BaseModel): + is_valid: bool = True + warnings: list = [] + + return MockResult() + from litellm.proxy._experimental.mcp_server.db import ( create_mcp_server, delete_mcp_server, @@ -122,9 +134,7 @@ if MCP_AVAILABLE: ) if validation_result.warnings: error_messages_text = ( - error_messages_text - + "\n" - + "\n".join(validation_result.warnings) + error_messages_text + "\n" + "\n".join(validation_result.warnings) ) raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 0aa8ccff1d4..bcd51c36632 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -9644,7 +9644,7 @@ def get_logo_url(): @app.get("/get_image", include_in_schema=False) -def get_image(): +async def get_image(): """Get logo to show on admin UI""" # get current_dir @@ -9663,25 +9663,37 @@ def get_image(): if is_non_root and not os.path.exists(default_logo): default_logo = default_site_logo + cache_dir = assets_dir if is_non_root else current_dir + cache_path = os.path.join(cache_dir, "cached_logo.jpg") + + # [OPTIMIZATION] Check if the cached image exists first + if os.path.exists(cache_path): + return FileResponse(cache_path, media_type="image/jpeg") + logo_path = os.getenv("UI_LOGO_PATH", default_logo) verbose_proxy_logger.debug("Reading logo from path: %s", logo_path) # Check if the logo path is an HTTP/HTTPS URL if logo_path.startswith(("http://", "https://")): - # Download the image and cache it - client = HTTPHandler() - response = client.get(logo_path) - if response.status_code == 200: - # Save the image to a local file - cache_dir = assets_dir if is_non_root else current_dir - cache_path = os.path.join(cache_dir, "cached_logo.jpg") - with open(cache_path, "wb") as f: - f.write(response.content) + try: + # Download the image and cache it + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - # Return the cached image as a FileResponse - return FileResponse(cache_path, media_type="image/jpeg") - else: - # Handle the case when the image cannot be downloaded + async_client = AsyncHTTPHandler(timeout=5.0) + response = await async_client.get(logo_path) + if response.status_code == 200: + # Save the image to a local file + with open(cache_path, "wb") as f: + f.write(response.content) + + # Return the cached image as a FileResponse + return FileResponse(cache_path, media_type="image/jpeg") + else: + # Handle the case when the image cannot be downloaded + return FileResponse(default_logo, media_type="image/jpeg") + except Exception as e: + # Handle any exceptions during the download (e.g., timeout, connection error) + verbose_proxy_logger.debug(f"Error downloading logo from {logo_path}: {e}") return FileResponse(default_logo, media_type="image/jpeg") else: # Return the local image file if the logo path is not an HTTP/HTTPS URL diff --git a/tests/proxy_unit_tests/test_get_image.py b/tests/proxy_unit_tests/test_get_image.py new file mode 100644 index 00000000000..ad8c2672754 --- /dev/null +++ b/tests/proxy_unit_tests/test_get_image.py @@ -0,0 +1,89 @@ +import os +import sys +from unittest import mock + +# Standard path insertion +sys.path.insert(0, os.path.abspath("../..")) + +import pytest +import httpx +from litellm.proxy.proxy_server import app + + +@pytest.mark.asyncio +async def test_get_image_error_handling(): + """ + Test that get_image handles network errors gracefully and doesn't hang. + """ + # Set an unreachable URL + os.environ["UI_LOGO_PATH"] = "http://invalid-url-12345.com/logo.jpg" + + # Clear cache + parent_dir = os.path.dirname( + os.path.dirname( + app.__file__ + if hasattr(app, "__file__") + else "litellm/proxy/proxy_server.py" + ) + ) + cache_path = os.path.join(parent_dir, "proxy", "cached_logo.jpg") + if os.path.exists(cache_path): + os.remove(cache_path) + + # Mock AsyncHTTPHandler to simulate a timeout or connection error + with mock.patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get" + ) as mock_get: + mock_get.side_effect = httpx.ConnectError("Network is unreachable") + + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), base_url="http://testserver" + ) as ac: + response = await ac.get("/get_image") + + assert response.status_code == 200 + assert response.headers["content-type"] == "image/jpeg" + + +@pytest.mark.asyncio +async def test_get_image_cache_logic(): + """ + Test that once cached, get_image doesn't hit the network. + """ + os.environ["UI_LOGO_PATH"] = "http://example.com/logo.jpg" + + # Clear cache + parent_dir = os.path.dirname( + os.path.dirname( + app.__file__ + if hasattr(app, "__file__") + else "litellm/proxy/proxy_server.py" + ) + ) + cache_path = os.path.join(parent_dir, "proxy", "cached_logo.jpg") + if os.path.exists(cache_path): + os.remove(cache_path) + + # Mock response + mock_response = mock.Mock() + mock_response.status_code = 200 + mock_response.content = b"fake image data" + + with mock.patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get" + ) as mock_get: + mock_get.return_value = mock_response + + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), base_url="http://testserver" + ) as ac: + # First call - should hit download logic + response1 = await ac.get("/get_image") + assert response1.status_code == 200 + assert mock_get.call_count == 1 + + # Second call - should hit cache + response2 = await ac.get("/get_image") + assert response2.status_code == 200 + # If cache works, mock_get shouldn't be called again + assert mock_get.call_count == 1 From 87acdef8990837895035fe2233ec2da0fef338aa Mon Sep 17 00:00:00 2001 From: Harshit Jain <48647625+Harshit28j@users.noreply.github.com> Date: Mon, 26 Jan 2026 12:41:33 +0530 Subject: [PATCH 6/7] feat: tpm-rpm limit in prometheus metrics (#19725) Co-authored-by: Krish Dholakia --- litellm/integrations/prometheus.py | 65 +++++++++++++++++++ .../litellm_core_utils/get_litellm_params.py | 9 ++- litellm/main.py | 38 +++++------ litellm/types/integrations/prometheus.py | 11 ++++ 4 files changed, 102 insertions(+), 21 deletions(-) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index abae36bf55b..6772520bd52 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -316,6 +316,18 @@ class PrometheusLogger(CustomLogger): labelnames=self.get_labels_for_metric("litellm_deployment_state"), ) + self.litellm_deployment_tpm_limit = self._gauge_factory( + "litellm_deployment_tpm_limit", + "Deployment TPM limit found in config", + labelnames=self.get_labels_for_metric("litellm_deployment_tpm_limit"), + ) + + self.litellm_deployment_rpm_limit = self._gauge_factory( + "litellm_deployment_rpm_limit", + "Deployment RPM limit found in config", + labelnames=self.get_labels_for_metric("litellm_deployment_rpm_limit"), + ) + self.litellm_deployment_cooled_down = self._counter_factory( "litellm_deployment_cooled_down", "LLM Deployment Analytics - Number of times a deployment has been cooled down by LiteLLM load balancing logic. exception_status is the status of the exception that caused the deployment to be cooled down", @@ -1778,6 +1790,49 @@ class PrometheusLogger(CustomLogger): ) ) + def _set_deployment_tpm_rpm_limit_metrics( + self, + model_info: dict, + litellm_params: dict, + litellm_model_name: Optional[str], + model_id: Optional[str], + api_base: Optional[str], + llm_provider: Optional[str], + ): + """ + Set the deployment TPM and RPM limits metrics + """ + tpm = model_info.get("tpm") or litellm_params.get("tpm") + rpm = model_info.get("rpm") or litellm_params.get("rpm") + + if tpm is not None: + _labels = prometheus_label_factory( + supported_enum_labels=self.get_labels_for_metric( + metric_name="litellm_deployment_tpm_limit" + ), + enum_values=UserAPIKeyLabelValues( + litellm_model_name=litellm_model_name, + model_id=model_id, + api_base=api_base, + api_provider=llm_provider, + ), + ) + self.litellm_deployment_tpm_limit.labels(**_labels).set(tpm) + + if rpm is not None: + _labels = prometheus_label_factory( + supported_enum_labels=self.get_labels_for_metric( + metric_name="litellm_deployment_rpm_limit" + ), + enum_values=UserAPIKeyLabelValues( + litellm_model_name=litellm_model_name, + model_id=model_id, + api_base=api_base, + api_provider=llm_provider, + ), + ) + self.litellm_deployment_rpm_limit.labels(**_labels).set(rpm) + def set_llm_deployment_success_metrics( self, request_kwargs: dict, @@ -1811,6 +1866,16 @@ class PrometheusLogger(CustomLogger): _model_info = _metadata.get("model_info") or {} model_id = _model_info.get("id", None) + if _model_info or _litellm_params: + self._set_deployment_tpm_rpm_limit_metrics( + model_info=_model_info, + litellm_params=_litellm_params, + litellm_model_name=litellm_model_name, + model_id=model_id, + api_base=api_base, + llm_provider=llm_provider, + ) + remaining_requests: Optional[int] = None remaining_tokens: Optional[int] = None if additional_headers := standard_logging_payload["hidden_params"][ diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index e290101f8bf..060e98fd49f 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -93,8 +93,11 @@ def get_litellm_params( "text_completion": text_completion, "azure_ad_token_provider": azure_ad_token_provider, "user_continue_message": user_continue_message, - "base_model": base_model or ( - _get_base_model_from_litellm_call_metadata(metadata=metadata) if metadata else None + "base_model": base_model + or ( + _get_base_model_from_litellm_call_metadata(metadata=metadata) + if metadata + else None ), "litellm_trace_id": litellm_trace_id, "litellm_session_id": litellm_session_id, @@ -139,5 +142,7 @@ def get_litellm_params( "aws_sts_endpoint": kwargs.get("aws_sts_endpoint"), "aws_external_id": kwargs.get("aws_external_id"), "aws_bedrock_runtime_endpoint": kwargs.get("aws_bedrock_runtime_endpoint"), + "tpm": kwargs.get("tpm"), + "rpm": kwargs.get("rpm"), } return litellm_params diff --git a/litellm/main.py b/litellm/main.py index ce84c8988e0..a4bcfdec81b 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -148,7 +148,7 @@ from litellm.utils import ( validate_and_fix_openai_messages, validate_and_fix_openai_tools, validate_chat_completion_tool_choice, - validate_openai_optional_params + validate_openai_optional_params, ) from ._logging import verbose_logger @@ -368,7 +368,7 @@ class AsyncCompletions: @tracer.wrap() @client -async def acompletion( # noqa: PLR0915 +async def acompletion( # noqa: PLR0915 model: str, # Optional OpenAI params: see https://platform.openai.com/docs/api-reference/chat/create messages: List = [], @@ -603,12 +603,11 @@ async def acompletion( # noqa: PLR0915 if timeout is not None and isinstance(timeout, (int, float)): timeout_value = float(timeout) init_response = await asyncio.wait_for( - loop.run_in_executor(None, func_with_context), - timeout=timeout_value + loop.run_in_executor(None, func_with_context), timeout=timeout_value ) else: init_response = await loop.run_in_executor(None, func_with_context) - + if isinstance(init_response, dict) or isinstance( init_response, ModelResponse ): ## CACHING SCENARIO @@ -640,6 +639,7 @@ async def acompletion( # noqa: PLR0915 except asyncio.TimeoutError: custom_llm_provider = custom_llm_provider or "openai" from litellm.exceptions import Timeout + raise Timeout( message=f"Request timed out after {timeout} seconds", model=model, @@ -1118,7 +1118,6 @@ def completion( # type: ignore # noqa: PLR0915 # validate optional params stop = validate_openai_optional_params(stop=stop) - ######### unpacking kwargs ##################### args = locals() @@ -1135,7 +1134,9 @@ def completion( # type: ignore # noqa: PLR0915 # Check if MCP tools are present (following responses pattern) # Cast tools to Optional[Iterable[ToolParam]] for type checking tools_for_mcp = cast(Optional[Iterable[ToolParam]], tools) - if LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(tools=tools_for_mcp): + if LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway( + tools=tools_for_mcp + ): # Return coroutine - acompletion will await it # completion() can return a coroutine when MCP tools are present, which acompletion() awaits return acompletion_with_mcp( # type: ignore[return-value] @@ -1536,6 +1537,8 @@ def completion( # type: ignore # noqa: PLR0915 max_retries=max_retries, timeout=timeout, litellm_request_debug=kwargs.get("litellm_request_debug", False), + tpm=kwargs.get("tpm"), + rpm=kwargs.get("rpm"), ) cast(LiteLLMLoggingObj, logging).update_environment_variables( model=model, @@ -2361,11 +2364,7 @@ def completion( # type: ignore # noqa: PLR0915 input=messages, api_key=api_key, original_response=response ) elif custom_llm_provider == "minimax": - api_key = ( - api_key - or get_secret_str("MINIMAX_API_KEY") - or litellm.api_key - ) + api_key = api_key or get_secret_str("MINIMAX_API_KEY") or litellm.api_key api_base = ( api_base @@ -2413,7 +2412,9 @@ def completion( # type: ignore # noqa: PLR0915 or custom_llm_provider == "wandb" or custom_llm_provider == "clarifai" or custom_llm_provider in litellm.openai_compatible_providers - or JSONProviderRegistry.exists(custom_llm_provider) # JSON-configured providers + or JSONProviderRegistry.exists( + custom_llm_provider + ) # JSON-configured providers or "ft:gpt-3.5-turbo" in model # finetune gpt-3.5-turbo ): # allow user to make an openai call with a custom base # note: if a user sets a custom base - we should ensure this works @@ -4724,7 +4725,7 @@ def embedding( # noqa: PLR0915 if headers is not None and headers != {}: optional_params["extra_headers"] = headers - + if encoding_format is not None: optional_params["encoding_format"] = encoding_format else: @@ -6759,9 +6760,7 @@ def speech( # noqa: PLR0915 if text_to_speech_provider_config is None: text_to_speech_provider_config = MinimaxTextToSpeechConfig() - minimax_config = cast( - MinimaxTextToSpeechConfig, text_to_speech_provider_config - ) + minimax_config = cast(MinimaxTextToSpeechConfig, text_to_speech_provider_config) if api_base is not None: litellm_params_dict["api_base"] = api_base @@ -6901,7 +6900,7 @@ async def ahealth_check( custom_llm_provider_from_params = model_params.get("custom_llm_provider", None) api_base_from_params = model_params.get("api_base", None) api_key_from_params = model_params.get("api_key", None) - + model, custom_llm_provider, _, _ = get_llm_provider( model=model, custom_llm_provider=custom_llm_provider_from_params, @@ -7275,8 +7274,9 @@ def __getattr__(name: str) -> Any: _encoding = tiktoken.get_encoding("cl100k_base") # Cache it in the module's __dict__ for subsequent accesses import sys + sys.modules[__name__].__dict__["encoding"] = _encoding global _encoding_cache _encoding_cache = _encoding return _encoding - raise AttributeError(f"module {__name__!r} has no attribute {name!r}") \ No newline at end of file + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index ea9c9bd325d..ee49ba1a19c 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -199,6 +199,8 @@ DEFINED_PROMETHEUS_METRICS = Literal[ "litellm_cache_hits_metric", "litellm_cache_misses_metric", "litellm_cached_tokens_metric", + "litellm_deployment_tpm_limit", + "litellm_deployment_rpm_limit", "litellm_remaining_api_key_requests_for_model", "litellm_remaining_api_key_tokens_for_model", "litellm_llm_api_failed_requests_metric", @@ -406,6 +408,15 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.API_PROVIDER.value, ] + litellm_deployment_tpm_limit = [ + UserAPIKeyLabelNames.v2_LITELLM_MODEL_NAME.value, + UserAPIKeyLabelNames.MODEL_ID.value, + UserAPIKeyLabelNames.API_BASE.value, + UserAPIKeyLabelNames.API_PROVIDER.value, + ] + + litellm_deployment_rpm_limit = litellm_deployment_tpm_limit + litellm_deployment_cooled_down = [ UserAPIKeyLabelNames.v2_LITELLM_MODEL_NAME.value, UserAPIKeyLabelNames.MODEL_ID.value, From aa8134fee9d85f2f2f023fff79cc41b17312c4d0 Mon Sep 17 00:00:00 2001 From: Tamir Kiviti <95572081+tamirkiviti13@users.noreply.github.com> Date: Mon, 26 Jan 2026 09:13:46 +0200 Subject: [PATCH 7/7] add timeout to onyx guardrail (#19731) * add timeout to onyx guardrail * add tests --- .../docs/proxy/guardrails/onyx_security.md | 3 + .../guardrails/guardrail_hooks/onyx/onyx.py | 7 +- .../proxy/guardrails/guardrail_hooks/onyx.py | 5 + .../guardrails/guardrail_hooks/test_onyx.py | 302 +++++++++++++++++- 4 files changed, 313 insertions(+), 4 deletions(-) diff --git a/docs/my-website/docs/proxy/guardrails/onyx_security.md b/docs/my-website/docs/proxy/guardrails/onyx_security.md index 85b0ba9f830..d240902eb52 100644 --- a/docs/my-website/docs/proxy/guardrails/onyx_security.md +++ b/docs/my-website/docs/proxy/guardrails/onyx_security.md @@ -128,6 +128,7 @@ guardrails: mode: ["pre_call", "post_call", "during_call"] # Run at multiple stages api_key: os.environ/ONYX_API_KEY api_base: os.environ/ONYX_API_BASE + timeout: 10.0 # Optional, defaults to 10 seconds ``` ### Required Parameters @@ -137,6 +138,7 @@ guardrails: ### Optional Parameters - **`api_base`**: Onyx API base URL (defaults to `https://ai-guard.onyx.security`) +- **`timeout`**: Request timeout in seconds (defaults to `10.0`) ## Environment Variables @@ -145,4 +147,5 @@ You can set these environment variables instead of hardcoding values in your con ```shell export ONYX_API_KEY="your-api-key-here" export ONYX_API_BASE="https://ai-guard.onyx.security" # Optional +export ONYX_TIMEOUT=10 # Optional, timeout in seconds ``` diff --git a/litellm/proxy/guardrails/guardrail_hooks/onyx/onyx.py b/litellm/proxy/guardrails/guardrail_hooks/onyx/onyx.py index 5f57cab1db4..3598dbe741e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/onyx/onyx.py +++ b/litellm/proxy/guardrails/guardrail_hooks/onyx/onyx.py @@ -8,6 +8,7 @@ import os import uuid from typing import TYPE_CHECKING, Any, Literal, Optional, Type +import httpx from fastapi import HTTPException from litellm._logging import verbose_proxy_logger @@ -25,10 +26,12 @@ if TYPE_CHECKING: class OnyxGuardrail(CustomGuardrail): def __init__( - self, api_base: Optional[str] = None, api_key: Optional[str] = None, **kwargs + self, api_base: Optional[str] = None, api_key: Optional[str] = None, timeout: Optional[float] = 10.0, **kwargs ): + timeout = timeout or int(os.getenv("ONYX_TIMEOUT", 10.0)) self.async_handler = get_async_httpx_client( - llm_provider=httpxSpecialProvider.GuardrailCallback + llm_provider=httpxSpecialProvider.GuardrailCallback, + params={"timeout": httpx.Timeout(timeout=timeout, connect=5.0)}, ) self.api_base = api_base or os.getenv( "ONYX_API_BASE", diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/onyx.py b/litellm/types/proxy/guardrails/guardrail_hooks/onyx.py index aa5b9d7a3fc..42d7e94829f 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/onyx.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/onyx.py @@ -16,6 +16,11 @@ class OnyxGuardrailConfigModel(GuardrailConfigModel): description="The API key for the Onyx Guard server. If not provided, the `ONYX_API_KEY` environment variable is checked.", ) + timeout: Optional[float] = Field( + default=None, + description="The timeout for the Onyx Guard server in seconds. If not provided, the `ONYX_TIMEOUT` environment variable is checked.", + ) + @staticmethod def ui_friendly_name() -> str: return "Onyx Guardrail" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_onyx.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_onyx.py index 9ede649f392..fb7480d263c 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_onyx.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_onyx.py @@ -3,6 +3,7 @@ import sys import uuid from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest from fastapi import HTTPException from httpx import Request, Response @@ -47,20 +48,129 @@ def test_onyx_guard_config(): del os.environ["ONYX_API_KEY"] +def test_onyx_guard_with_custom_timeout_from_kwargs(): + """Test Onyx guard instantiation with custom timeout passed via kwargs.""" + # Set environment variables for testing + os.environ["ONYX_API_BASE"] = "https://test.onyx.security" + os.environ["ONYX_API_KEY"] = "test-api-key" + + with patch( + "litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client" + ) as mock_get_client: + mock_get_client.return_value = MagicMock() + + # Simulate how guardrail is instantiated from config with timeout + guardrail = OnyxGuardrail( + guardrail_name="onyx-guard-custom-timeout", + event_hook="pre_call", + default_on=True, + timeout=45.0, + ) + + # Verify the client was initialized with custom timeout + mock_get_client.assert_called() + call_kwargs = mock_get_client.call_args.kwargs + timeout_param = call_kwargs["params"]["timeout"] + assert timeout_param.read == 45.0 + assert timeout_param.connect == 5.0 + + # Clean up + if "ONYX_API_BASE" in os.environ: + del os.environ["ONYX_API_BASE"] + if "ONYX_API_KEY" in os.environ: + del os.environ["ONYX_API_KEY"] + + +def test_onyx_guard_with_timeout_none_uses_env_var(): + """Test Onyx guard with timeout=None uses ONYX_TIMEOUT env var. + + When timeout=None is passed (as it would be from config model with default None), + the ONYX_TIMEOUT environment variable should be used. + """ + # Set environment variables for testing + os.environ["ONYX_API_BASE"] = "https://test.onyx.security" + os.environ["ONYX_API_KEY"] = "test-api-key" + os.environ["ONYX_TIMEOUT"] = "60" + + with patch( + "litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client" + ) as mock_get_client: + mock_get_client.return_value = MagicMock() + + # Pass timeout=None to simulate config model behavior + guardrail = OnyxGuardrail( + guardrail_name="onyx-guard-env-timeout", + event_hook="pre_call", + default_on=True, + timeout=None, # This triggers env var lookup + ) + + # Verify the client was initialized with timeout from env var + mock_get_client.assert_called() + call_kwargs = mock_get_client.call_args.kwargs + timeout_param = call_kwargs["params"]["timeout"] + assert timeout_param.read == 60.0 + assert timeout_param.connect == 5.0 + + # Clean up + if "ONYX_API_BASE" in os.environ: + del os.environ["ONYX_API_BASE"] + if "ONYX_API_KEY" in os.environ: + del os.environ["ONYX_API_KEY"] + if "ONYX_TIMEOUT" in os.environ: + del os.environ["ONYX_TIMEOUT"] + + +def test_onyx_guard_with_timeout_none_defaults_to_10(): + """Test Onyx guard with timeout=None and no env var defaults to 10 seconds.""" + # Set environment variables for testing + os.environ["ONYX_API_BASE"] = "https://test.onyx.security" + os.environ["ONYX_API_KEY"] = "test-api-key" + # Ensure ONYX_TIMEOUT is not set + if "ONYX_TIMEOUT" in os.environ: + del os.environ["ONYX_TIMEOUT"] + + with patch( + "litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client" + ) as mock_get_client: + mock_get_client.return_value = MagicMock() + + # Pass timeout=None with no env var - should default to 10.0 + guardrail = OnyxGuardrail( + guardrail_name="onyx-guard-default-timeout", + event_hook="pre_call", + default_on=True, + timeout=None, + ) + + # Verify the client was initialized with default timeout of 10.0 + mock_get_client.assert_called() + call_kwargs = mock_get_client.call_args.kwargs + timeout_param = call_kwargs["params"]["timeout"] + assert timeout_param.read == 10.0 + assert timeout_param.connect == 5.0 + + # Clean up + if "ONYX_API_BASE" in os.environ: + del os.environ["ONYX_API_BASE"] + if "ONYX_API_KEY" in os.environ: + del os.environ["ONYX_API_KEY"] + + class TestOnyxGuardrail: """Test suite for Onyx Security Guardrail integration.""" def setup_method(self): """Setup test environment.""" # Clean up any existing environment variables - for key in ["ONYX_API_BASE", "ONYX_API_KEY"]: + for key in ["ONYX_API_BASE", "ONYX_API_KEY", "ONYX_TIMEOUT"]: if key in os.environ: del os.environ[key] def teardown_method(self): """Clean up test environment.""" # Clean up any environment variables set during tests - for key in ["ONYX_API_BASE", "ONYX_API_KEY"]: + for key in ["ONYX_API_BASE", "ONYX_API_KEY", "ONYX_TIMEOUT"]: if key in os.environ: del os.environ[key] @@ -103,6 +213,95 @@ class TestOnyxGuardrail: ): OnyxGuardrail(guardrail_name="test-guard", event_hook="pre_call") + def test_initialization_with_default_timeout(self): + """Test that default timeout is 10.0 seconds.""" + os.environ["ONYX_API_KEY"] = "test-api-key" + + with patch( + "litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client" + ) as mock_get_client: + mock_get_client.return_value = MagicMock() + guardrail = OnyxGuardrail( + guardrail_name="test-guard", event_hook="pre_call", default_on=True + ) + + # Verify the client was initialized with correct timeout + mock_get_client.assert_called_once() + call_kwargs = mock_get_client.call_args.kwargs + timeout_param = call_kwargs["params"]["timeout"] + assert timeout_param.read == 10.0 + assert timeout_param.connect == 5.0 + + def test_initialization_with_custom_timeout_parameter(self): + """Test initialization with custom timeout parameter.""" + os.environ["ONYX_API_KEY"] = "test-api-key" + + with patch( + "litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client" + ) as mock_get_client: + mock_get_client.return_value = MagicMock() + guardrail = OnyxGuardrail( + guardrail_name="test-guard", + event_hook="pre_call", + default_on=True, + timeout=30.0, + ) + + # Verify the client was initialized with custom timeout + mock_get_client.assert_called_once() + call_kwargs = mock_get_client.call_args.kwargs + timeout_param = call_kwargs["params"]["timeout"] + assert timeout_param.read == 30.0 + assert timeout_param.connect == 5.0 + + def test_initialization_with_timeout_from_env_var(self): + """Test initialization with timeout from ONYX_TIMEOUT environment variable. + + Note: The env var is only used when timeout=None is explicitly passed, + since the default parameter value is 10.0 (not None). + """ + os.environ["ONYX_API_KEY"] = "test-api-key" + os.environ["ONYX_TIMEOUT"] = "25" + + with patch( + "litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client" + ) as mock_get_client: + mock_get_client.return_value = MagicMock() + # Must pass timeout=None explicitly to trigger env var lookup + guardrail = OnyxGuardrail( + guardrail_name="test-guard", event_hook="pre_call", default_on=True, timeout=None + ) + + # Verify the client was initialized with timeout from env var + mock_get_client.assert_called_once() + call_kwargs = mock_get_client.call_args.kwargs + timeout_param = call_kwargs["params"]["timeout"] + assert timeout_param.read == 25.0 + assert timeout_param.connect == 5.0 + + def test_initialization_timeout_parameter_overrides_env_var(self): + """Test that timeout parameter overrides ONYX_TIMEOUT environment variable.""" + os.environ["ONYX_API_KEY"] = "test-api-key" + os.environ["ONYX_TIMEOUT"] = "25" + + with patch( + "litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client" + ) as mock_get_client: + mock_get_client.return_value = MagicMock() + guardrail = OnyxGuardrail( + guardrail_name="test-guard", + event_hook="pre_call", + default_on=True, + timeout=15.0, + ) + + # Verify the client was initialized with parameter timeout (not env var) + mock_get_client.assert_called_once() + call_kwargs = mock_get_client.call_args.kwargs + timeout_param = call_kwargs["params"]["timeout"] + assert timeout_param.read == 15.0 + assert timeout_param.connect == 5.0 + @pytest.mark.asyncio async def test_apply_guardrail_request_no_violations(self): """Test apply_guardrail for request with no violations detected.""" @@ -388,6 +587,105 @@ class TestOnyxGuardrail: assert result == inputs + @pytest.mark.asyncio + async def test_apply_guardrail_timeout_error_handling(self): + """Test handling of timeout errors in apply_guardrail (graceful degradation).""" + # Set required API key + os.environ["ONYX_API_KEY"] = "test-api-key" + + guardrail = OnyxGuardrail( + guardrail_name="test-guard", event_hook="pre_call", default_on=True, timeout=1.0 + ) + + inputs = GenericGuardrailAPIInputs() + + request_data = { + "proxy_server_request": { + "messages": [{"role": "user", "content": "Test message"}], + "model": "gpt-3.5-turbo", + } + } + + # Test httpx timeout error + with patch.object( + guardrail.async_handler, "post", side_effect=httpx.TimeoutException("Request timed out") + ): + # Should return original inputs on timeout (graceful degradation) + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + + assert result == inputs + + @pytest.mark.asyncio + async def test_apply_guardrail_read_timeout_error_handling(self): + """Test handling of read timeout errors in apply_guardrail.""" + # Set required API key + os.environ["ONYX_API_KEY"] = "test-api-key" + + guardrail = OnyxGuardrail( + guardrail_name="test-guard", event_hook="pre_call", default_on=True, timeout=5.0 + ) + + inputs = GenericGuardrailAPIInputs() + + request_data = { + "proxy_server_request": { + "messages": [{"role": "user", "content": "Test message"}], + "model": "gpt-3.5-turbo", + } + } + + # Test httpx ReadTimeout error + with patch.object( + guardrail.async_handler, "post", side_effect=httpx.ReadTimeout("Read timed out") + ): + # Should return original inputs on timeout (graceful degradation) + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + + assert result == inputs + + @pytest.mark.asyncio + async def test_apply_guardrail_connect_timeout_error_handling(self): + """Test handling of connect timeout errors in apply_guardrail.""" + # Set required API key + os.environ["ONYX_API_KEY"] = "test-api-key" + + guardrail = OnyxGuardrail( + guardrail_name="test-guard", event_hook="pre_call", default_on=True, timeout=5.0 + ) + + inputs = GenericGuardrailAPIInputs() + + request_data = { + "proxy_server_request": { + "messages": [{"role": "user", "content": "Test message"}], + "model": "gpt-3.5-turbo", + } + } + + # Test httpx ConnectTimeout error + with patch.object( + guardrail.async_handler, "post", side_effect=httpx.ConnectTimeout("Connect timed out") + ): + # Should return original inputs on timeout (graceful degradation) + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + + assert result == inputs + @pytest.mark.asyncio async def test_apply_guardrail_no_logging_obj(self): """Test apply_guardrail without logging object (uses UUID)."""