mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
resolve_mcp_auth
This commit is contained in:
parent
3c25cda0b6
commit
2579169ccb
2 changed files with 33 additions and 17 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue