diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index e7174f943b2..532aea249bf 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -36,6 +36,7 @@ from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, ) +from litellm.proxy._experimental.mcp_server.oauth2_token_cache import resolve_mcp_auth from litellm.proxy._experimental.mcp_server.utils import ( MCP_TOOL_PREFIX_SEPARATOR, add_server_prefix_to_name, @@ -833,7 +834,7 @@ class MCPServerManager: return resolved_env - def _create_mcp_client( + async def _create_mcp_client( self, server: MCPServer, mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None, @@ -843,13 +844,22 @@ class MCPServerManager: """ Create an MCPClient instance for the given server. + Auth resolution (single place for all auth logic): + 1. ``mcp_auth_header`` — per-request/per-user override + 2. OAuth2 client_credentials token — auto-fetched and cached + 3. ``server.authentication_token`` — static token from config/DB + Args: - server (MCPServer): The server configuration - mcp_auth_header: MCP auth header to be passed to the MCP server. This is optional and will be used if provided. + server: The server configuration. + mcp_auth_header: Optional per-request auth override. + extra_headers: Additional headers to forward. + stdio_env: Environment variables for stdio transport. Returns: - MCPClient: Configured MCP client instance + Configured MCP client instance. """ + auth_value = await resolve_mcp_auth(server, mcp_auth_header) + transport = server.transport or MCPTransport.sse # Handle stdio transport @@ -868,7 +878,7 @@ class MCPServerManager: server_url="", # Not used for stdio transport_type=transport, auth_type=server.auth_type, - auth_value=mcp_auth_header or server.authentication_token, + auth_value=auth_value, timeout=60.0, stdio_config=stdio_config, extra_headers=extra_headers, @@ -880,7 +890,7 @@ class MCPServerManager: server_url=server_url, transport_type=transport, auth_type=server.auth_type, - auth_value=mcp_auth_header or server.authentication_token, + auth_value=auth_value, timeout=60.0, extra_headers=extra_headers, ) @@ -920,7 +930,7 @@ class MCPServerManager: stdio_env = self._build_stdio_env(server, raw_headers) - client = self._create_mcp_client( + client = await self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -980,7 +990,7 @@ class MCPServerManager: stdio_env = self._build_stdio_env(server, raw_headers) - client = self._create_mcp_client( + client = await self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -1024,7 +1034,7 @@ class MCPServerManager: stdio_env = self._build_stdio_env(server, raw_headers) - client = self._create_mcp_client( + client = await self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -1068,7 +1078,7 @@ class MCPServerManager: stdio_env = self._build_stdio_env(server, raw_headers) - client = self._create_mcp_client( + client = await self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -1109,7 +1119,7 @@ class MCPServerManager: stdio_env = self._build_stdio_env(server, raw_headers) - client = self._create_mcp_client( + client = await self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -1139,7 +1149,7 @@ class MCPServerManager: stdio_env = self._build_stdio_env(server, raw_headers) - client = self._create_mcp_client( + client = await self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -1943,7 +1953,7 @@ class MCPServerManager: stdio_env = self._build_stdio_env(mcp_server, raw_headers) - client = self._create_mcp_client( + client = await self._create_mcp_client( server=mcp_server, mcp_auth_header=server_auth_header, extra_headers=extra_headers, @@ -2119,8 +2129,8 @@ class MCPServerManager: Note: This now handles prefixed tool names """ for server in self.get_registry().values(): - if server.auth_type == MCPAuth.oauth2: - # Skip OAuth2 servers for now as they may require user-specific tokens + if server.needs_user_oauth_token: + # Skip OAuth2 servers that rely on user-provided tokens continue tools = await self._get_tools_from_server(server) for tool in tools: @@ -2414,7 +2424,7 @@ class MCPServerManager: should_skip_health_check = False # Skip if auth_type is oauth2 - if server.auth_type == MCPAuth.oauth2: + if server.needs_user_oauth_token: should_skip_health_check = True # Skip if auth_type is not none and authentication_token is missing elif ( @@ -2429,7 +2439,7 @@ class MCPServerManager: if server.static_headers: extra_headers.update(server.static_headers) - client = self._create_mcp_client( + client = await self._create_mcp_client( server=server, mcp_auth_header=None, extra_headers=extra_headers, diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 3e6fb739821..2cd385c5bf6 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -4,6 +4,7 @@ from typing import Any, Dict, List, Optional from pydantic import BaseModel, ConfigDict from litellm.proxy._types import MCPAuthType, MCPTransportType +from litellm.types.mcp import MCPAuth # MCPInfo now allows arbitrary additional fields for custom metadata MCPInfo = Dict[str, Any] @@ -59,3 +60,8 @@ class MCPServer(BaseModel): def has_client_credentials(self) -> bool: """True if this server has OAuth2 client_credentials config (client_id, client_secret, token_url).""" return bool(self.client_id and self.client_secret and self.token_url) + + @property + def needs_user_oauth_token(self) -> bool: + """True if this is an OAuth2 server that relies on per-user tokens (no client_credentials).""" + return self.auth_type == MCPAuth.oauth2 and not self.has_client_credentials