diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 1d7c0eaa116..82365be60ca 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -15,7 +15,6 @@ from typing import Any, Callable, Dict, List, Literal, Optional, Set, Tuple, Uni from urllib.parse import urlparse import anyio - from fastapi import HTTPException from httpx import HTTPStatusError from mcp import ReadResourceResult, Resource @@ -72,7 +71,9 @@ try: from mcp.shared.tool_name_validation import ( validate_tool_name, # pyright: ignore[reportAssignmentType] ) - from mcp.shared.tool_name_validation import SEP_986_URL + from mcp.shared.tool_name_validation import ( + SEP_986_URL, + ) except ImportError: from pydantic import BaseModel @@ -332,6 +333,15 @@ class MCPServerManager: available_on_public_internet=bool( server_config.get("available_on_public_internet", False) ), + # Token Exchange (OBO) fields + token_exchange_endpoint=server_config.get( + "token_exchange_endpoint", None + ), + audience=server_config.get("audience", None), + subject_token_type=server_config.get( + "subject_token_type", + "urn:ietf:params:oauth:token-type:access_token", + ), ) self.config_mcp_servers[server_id] = new_server @@ -636,6 +646,21 @@ class MCPServerManager: getattr(mcp_server, "available_on_public_internet", False) ), updated_at=getattr(mcp_server, "updated_at", None), + # Token Exchange (OBO) fields — read from credentials JSON blob + token_exchange_endpoint=( + credentials_dict.get("token_exchange_endpoint") + if credentials_dict + else None + ), + audience=( + credentials_dict.get("audience") if credentials_dict else None + ), + subject_token_type=( + credentials_dict.get("subject_token_type") + if credentials_dict + else None + ) + or "urn:ietf:params:oauth:token-type:access_token", ) return new_server @@ -832,6 +857,27 @@ class MCPServerManager: ######################################################### # Methods that call the upstream MCP servers ######################################################### + @staticmethod + def _extract_bearer_token( + oauth2_headers: Optional[Dict[str, str]], + raw_headers: Optional[Dict[str, str]], + ) -> Optional[str]: + """Extract the bare Bearer token from oauth2_headers or raw_headers. + + Returns the token string without the ``Bearer `` prefix, or ``None`` + if no Authorization header is found. + """ + auth_value: Optional[str] = None + if oauth2_headers and "Authorization" in oauth2_headers: + auth_value = oauth2_headers["Authorization"] + elif raw_headers: + # raw_headers may have lowercase keys depending on the ASGI server + normalized = {k.lower(): v for k, v in raw_headers.items()} + auth_value = normalized.get("authorization") + if auth_value and auth_value.startswith("Bearer "): + return auth_value[len("Bearer "):] + return auth_value + def _build_stdio_env( self, server: MCPServer, @@ -865,25 +911,30 @@ class MCPServerManager: mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None, extra_headers: Optional[Dict[str, str]] = None, stdio_env: Optional[Dict[str, str]] = None, + subject_token: Optional[str] = None, ) -> MCPClient: """ 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 + 2. OAuth2 Token Exchange (OBO) — exchange user token for scoped token + 3. OAuth2 client_credentials token — auto-fetched and cached + 4. ``server.authentication_token`` — static token from config/DB Args: 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. + subject_token: Optional user JWT for token exchange (OBO) flow. Returns: Configured MCP client instance. """ - auth_value = await resolve_mcp_auth(server, mcp_auth_header) + auth_value = await resolve_mcp_auth( + server, mcp_auth_header, subject_token=subject_token + ) transport = server.transport or MCPTransport.sse @@ -1942,9 +1993,12 @@ class MCPServerManager: if server_auth_header is None: server_auth_header = mcp_auth_header - # oauth2 headers + # Extract subject token for OAuth2 Token Exchange (OBO) flow + subject_token: Optional[str] = None extra_headers: Optional[Dict[str, str]] = None - if mcp_server.auth_type == MCPAuth.oauth2: + if mcp_server.auth_type == MCPAuth.oauth2_token_exchange: + subject_token = self._extract_bearer_token(oauth2_headers, raw_headers) + elif mcp_server.auth_type == MCPAuth.oauth2: extra_headers = oauth2_headers if mcp_server.extra_headers and raw_headers: @@ -1974,6 +2028,7 @@ class MCPServerManager: mcp_auth_header=server_auth_header, extra_headers=extra_headers, stdio_env=stdio_env, + subject_token=subject_token, ) call_tool_params = MCPCallToolRequestParams(