mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(mcp): extract subject token and wire OBO into MCP tool calls
Adds _extract_bearer_token() helper, updates _create_mcp_client() and _call_regular_mcp_tool() to extract the user's JWT and pass it as the subject token for OBO exchange. Also updates load_servers_from_config() and build_mcp_server_from_table() to read token exchange config fields. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
a2b010b2a5
commit
df081c02d6
1 changed files with 62 additions and 7 deletions
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue