resolve_mcp_auth

This commit is contained in:
Ishaan Jaffer 2026-02-09 14:49:16 -08:00
parent 3c25cda0b6
commit 2579169ccb
2 changed files with 33 additions and 17 deletions

View file

@ -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,

View file

@ -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