adopt MCP SDK OAuth discovery utils and token models for upstream auth

Replace the hand-rolled pieces of upstream OAuth in the MCP gateway with
their official SDK equivalents:

- WWW-Authenticate parsing now uses mcp.client.auth.utils
  extract_field_from_www_auth instead of a custom regex
- protected-resource and authorization-server discovery URL candidates come
  from build_protected_resource_metadata_discovery_urls and
  build_oauth_authorization_server_metadata_discovery_urls (SEP-985 /
  RFC 8414 ordering); the gateway keeps its legacy extras (root endpoints
  for path-ful issuers, bare-issuer fetch, Azure Entra fallback) appended
  after the spec-ordered list
- discovery payloads validate against mcp.shared.auth ProtectedResourceMetadata
  and OAuthMetadata first, with the previous lenient extraction kept as a
  fallback for IdPs that serve partial documents
- token-endpoint responses (client_credentials and RFC 8693 token exchange)
  parse via mcp.shared.auth.OAuthToken instead of hand-rolled dict checks
- the RFC 8693 exchange request now sends an explicit requested_token_type
  per spec section 2.1
- fetch_upstream_oauth_protected_resource now probes the RFC 9728 section 3.1
  path-suffix URL before the host root, matching spec ordering (previously
  reversed)
This commit is contained in:
mateo-berri 2026-06-10 16:45:20 +00:00
parent 67a8c0f27b
commit 0c55b3537a
No known key found for this signature in database
6 changed files with 145 additions and 146 deletions

View file

@ -14,6 +14,8 @@ import weakref
from typing import TYPE_CHECKING, Dict, Tuple
import httpx
from mcp.shared.auth import OAuthToken
from pydantic import ValidationError
from litellm._logging import verbose_logger
from litellm.caching.in_memory_cache import InMemoryCache
@ -35,6 +37,21 @@ TOKEN_EXCHANGE_GRANT_TYPE = "urn:ietf:params:oauth:grant-type:token-exchange"
DEFAULT_SUBJECT_TOKEN_TYPE = "urn:ietf:params:oauth:token-type:access_token"
def parse_oauth_token_response(response: httpx.Response, server_id: str) -> OAuthToken:
"""Validate a token-endpoint response per RFC 6749 §5.1 using the SDK model.
Raises ``ValueError`` with a server-scoped message on malformed payloads so
callers surface a clear configuration error instead of a pydantic trace.
"""
try:
return OAuthToken.model_validate(response.json())
except ValidationError as exc:
raise ValueError(
f"OAuth2 token response for MCP server '{server_id}' is not a valid "
f"RFC 6749 token payload: {exc.error_count()} validation error(s)"
) from exc
class TokenExchangeHandler:
"""Handles OAuth 2.0 Token Exchange (RFC 8693) for MCP servers.
@ -120,6 +137,8 @@ class TokenExchangeHandler:
"subject_token": subject_token,
"subject_token_type": server.subject_token_type
or DEFAULT_SUBJECT_TOKEN_TYPE,
# RFC 8693 §2.1: explicit requested_token_type (the spec default)
"requested_token_type": "urn:ietf:params:oauth:token-type:access_token",
"client_id": server.client_id,
"client_secret": server.client_secret,
}
@ -150,30 +169,12 @@ class TokenExchangeHandler:
f"failed with status {exc.response.status_code}"
) from exc
body = response.json()
if not isinstance(body, dict):
raise ValueError(
f"Token exchange response for MCP server '{server.server_id}' "
f"returned non-object JSON (got {type(body).__name__})"
)
access_token = body.get("access_token")
if not access_token:
raise ValueError(
f"Token exchange response for MCP server '{server.server_id}' "
f"missing 'access_token'"
)
raw_expires_in = body.get("expires_in")
try:
expires_in = (
int(raw_expires_in)
if raw_expires_in is not None
else MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL
)
except (TypeError, ValueError):
expires_in = MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL
token = parse_oauth_token_response(response, server.server_id)
expires_in = (
token.expires_in
if token.expires_in is not None
else MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL
)
ttl = max(
expires_in - MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS,
MCP_OAUTH2_TOKEN_CACHE_MIN_TTL,
@ -184,7 +185,7 @@ class TokenExchangeHandler:
server.server_id,
expires_in,
)
return access_token, ttl
return token.access_token, ttl
def invalidate(self, subject_token: str, server_id: str) -> None:
"""Remove a cached exchanged token (e.g. after a 401)."""

View file

@ -844,9 +844,9 @@ async def fetch_upstream_oauth_protected_resource(
"""Fetch the upstream MCP server's ``.well-known/oauth-protected-resource``
metadata for a pass-through server.
Tries host-only first, then falls back to the RFC 9728 §3.1 path-suffix
form (e.g. ``https://host/.well-known/oauth-protected-resource/mcp``) to
cover upstreams that scope metadata per resource path.
Tries the RFC 9728 §3.1 path-suffix form first (e.g.
``https://host/.well-known/oauth-protected-resource/mcp``) so upstreams
that scope metadata per resource path win, then falls back to host-root.
Responses are cached in-process for ~5 minutes keyed on
``(server_id, resource_url)`` so we do not hammer the IdP.
@ -876,14 +876,14 @@ async def fetch_upstream_oauth_protected_resource(
if cached is not None and cached[0] > now:
return cached[1]
host_base = f"{upstream.scheme}://{upstream.netloc}"
candidates = [f"{host_base}/.well-known/oauth-protected-resource"]
# RFC 9728 §3.1 path fallback
if upstream.path and upstream.path not in ("", "/"):
candidates.append(
f"{host_base}/.well-known/oauth-protected-resource"
f"{upstream.path.rstrip('/')}"
)
from mcp.client.auth.utils import (
build_protected_resource_metadata_discovery_urls,
)
# RFC 9728 §3.1 ordering: path-suffix form first, host-root fallback.
candidates = build_protected_resource_metadata_discovery_urls(
None, mcp_server.url
)
async_client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.Oauth2Check

View file

@ -20,6 +20,12 @@ import anyio
from fastapi import HTTPException
from httpx import HTTPStatusError
from mcp import ReadResourceResult, Resource
from mcp.client.auth.utils import (
build_oauth_authorization_server_metadata_discovery_urls,
build_protected_resource_metadata_discovery_urls,
extract_field_from_www_auth,
)
from mcp.shared.auth import OAuthMetadata, ProtectedResourceMetadata
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
from mcp.types import (
CallToolResult,
@ -29,7 +35,7 @@ from mcp.types import (
ResourceTemplate,
)
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl
from pydantic import AnyUrl, ValidationError
import litellm
from litellm._logging import verbose_logger
@ -2407,15 +2413,15 @@ class MCPServerManager:
exc,
)
header_value: Optional[str] = None
resource_metadata_url: Optional[str] = None
scopes: Optional[List[str]] = None
if exc.response is not None:
header_value = exc.response.headers.get(
"WWW-Authenticate"
) or exc.response.headers.get("www-authenticate")
resource_metadata_url, scopes = self._parse_www_authenticate_header(
header_value
)
resource_metadata_url = extract_field_from_www_auth(
exc.response, "resource_metadata"
)
scopes = self._extract_scopes(
extract_field_from_www_auth(exc.response, "scope")
)
authorization_servers = []
resource_scopes = None
@ -2461,29 +2467,6 @@ class MCPServerManager:
)
return None
def _parse_www_authenticate_header(
self, header_value: Optional[str]
) -> Tuple[Optional[str], Optional[List[str]]]:
if not header_value:
return None, None
_, _, params_section = header_value.partition(" ")
params_section = params_section or header_value
param_pattern = re.compile(r"([a-zA-Z0-9_]+)\s*=\s*\"?([^\",]+)\"?")
params: Dict[str, str] = {
match.group(1).lower(): match.group(2).strip()
for match in param_pattern.finditer(params_section)
}
resource_metadata_url = params.get("resource_metadata")
scope_value = params.get("scope")
scopes_list = [s for s in (scope_value.split() if scope_value else []) if s]
scopes = scopes_list or None
return resource_metadata_url, scopes
async def _fetch_oauth_metadata_from_resource(
self, resource_metadata_url: str, server_url: str
) -> Tuple[List[str], Optional[List[str]]]:
@ -2513,19 +2496,32 @@ class MCPServerManager:
)
return [], None
raw_servers = data.get("authorization_servers")
if isinstance(raw_servers, list):
authorization_servers = [
entry
for entry in raw_servers
if isinstance(entry, str) and entry.strip() != ""
]
else:
authorization_servers = []
if not isinstance(data, dict):
return [], None
scopes = self._extract_scopes(
data.get("scopes_supported") or data.get("scopes")
)
try:
metadata = ProtectedResourceMetadata.model_validate(data)
authorization_servers = [
str(entry) for entry in metadata.authorization_servers
]
scopes = self._extract_scopes(metadata.scopes_supported)
except ValidationError:
# Lenient fallback: real upstreams routinely omit RFC 9728
# required fields (e.g. ``resource``) while still advertising
# usable authorization_servers.
raw_servers = data.get("authorization_servers")
if isinstance(raw_servers, list):
authorization_servers = [
entry
for entry in raw_servers
if isinstance(entry, str) and entry.strip() != ""
]
else:
authorization_servers = []
scopes = self._extract_scopes(data.get("scopes_supported"))
if scopes is None:
scopes = self._extract_scopes(data.get("scopes"))
return authorization_servers, scopes
@ -2540,16 +2536,7 @@ class MCPServerManager:
if not parsed.scheme or not parsed.netloc:
return [], None
base = f"{parsed.scheme}://{parsed.netloc}"
path = parsed.path or ""
path = path.strip("/")
candidate_urls: List[str] = []
if path:
candidate_urls.append(f"{base}/.well-known/oauth-protected-resource/{path}")
candidate_urls.append(f"{base}/.well-known/oauth-protected-resource")
for url in candidate_urls:
for url in build_protected_resource_metadata_discovery_urls(None, server_url):
(
authorization_servers,
scopes,
@ -2581,21 +2568,20 @@ class MCPServerManager:
if not parsed.scheme or not parsed.netloc:
return None
candidate_urls = build_oauth_authorization_server_metadata_discovery_urls(
issuer_url, server_url
)
# Legacy fallbacks beyond the SDK's RFC 8414 ordered list: root
# endpoints for path-ful issuers, and the bare issuer URL (some IdPs
# serve their OIDC document directly at the issuer).
base = f"{parsed.scheme}://{parsed.netloc}"
path = (parsed.path or "").strip("/")
candidate_urls: List[str] = []
if path:
candidate_urls.append(
f"{base}/.well-known/oauth-authorization-server/{path}"
)
candidate_urls.append(f"{base}/.well-known/openid-configuration/{path}")
candidate_urls.append(
f"{issuer_url.rstrip('/')}/.well-known/openid-configuration"
)
candidate_urls.append(f"{base}/.well-known/oauth-authorization-server")
candidate_urls.append(f"{base}/.well-known/openid-configuration")
candidate_urls.append(issuer_url.rstrip("/"))
for legacy_url in (
f"{base}/.well-known/oauth-authorization-server",
f"{base}/.well-known/openid-configuration",
issuer_url.rstrip("/"),
):
if legacy_url not in candidate_urls:
candidate_urls.append(legacy_url)
for url in candidate_urls:
try:
@ -2619,25 +2605,51 @@ class MCPServerManager:
)
continue
scopes = self._extract_scopes(data.get("scopes_supported"))
metadata = self._oauth_metadata_from_asm_payload(data)
if metadata is not None:
return metadata
return self._build_azure_authorization_server_metadata(parsed)
def _oauth_metadata_from_asm_payload(self, data: Any) -> Optional[MCPOAuthMetadata]:
"""Map an RFC 8414 / OIDC discovery payload onto MCPOAuthMetadata.
Validates with the SDK's strict OAuthMetadata model first; falls back
to lenient field extraction for IdPs that serve partial documents
(e.g. missing ``issuer``).
"""
if not isinstance(data, dict):
return None
try:
asm = OAuthMetadata.model_validate(data)
metadata = MCPOAuthMetadata(
scopes=scopes,
scopes=self._extract_scopes(asm.scopes_supported),
authorization_url=str(asm.authorization_endpoint),
token_url=str(asm.token_endpoint),
registration_url=(
str(asm.registration_endpoint)
if asm.registration_endpoint
else None
),
)
except ValidationError:
metadata = MCPOAuthMetadata(
scopes=self._extract_scopes(data.get("scopes_supported")),
authorization_url=data.get("authorization_endpoint"),
token_url=data.get("token_endpoint"),
registration_url=data.get("registration_endpoint"),
)
if any(
[
metadata.scopes,
metadata.authorization_url,
metadata.token_url,
metadata.registration_url,
]
):
return metadata
return self._build_azure_authorization_server_metadata(parsed)
if any(
[
metadata.scopes,
metadata.authorization_url,
metadata.token_url,
metadata.registration_url,
]
):
return metadata
return None
@staticmethod
def _build_azure_authorization_server_metadata(

View file

@ -27,6 +27,9 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
encrypt_value_helper,
)
from litellm.proxy._experimental.mcp_server.auth import token_exchange
from litellm.proxy._experimental.mcp_server.auth.token_exchange import (
parse_oauth_token_response as _parse_oauth_token_response,
)
from litellm.types.llms.custom_http import httpxSpecialProvider
if TYPE_CHECKING:
@ -125,32 +128,12 @@ class MCPOAuth2TokenCache(InMemoryCache):
f"failed with status {exc.response.status_code}"
) from exc
body = response.json()
if not isinstance(body, dict):
raise ValueError(
f"OAuth2 token response for MCP server '{server.server_id}' "
f"returned non-object JSON (got {type(body).__name__})"
)
access_token = body.get("access_token")
if not access_token:
raise ValueError(
f"OAuth2 token response for MCP server '{server.server_id}' "
f"missing 'access_token'"
)
# Safely parse expires_in — providers may return null or non-numeric values
raw_expires_in = body.get("expires_in")
try:
expires_in = (
int(raw_expires_in)
if raw_expires_in is not None
else MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL
)
except (TypeError, ValueError):
expires_in = MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL
token = _parse_oauth_token_response(response, server.server_id)
expires_in = (
token.expires_in
if token.expires_in is not None
else MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL
)
ttl = max(
expires_in - MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS,
MCP_OAUTH2_TOKEN_CACHE_MIN_TTL,
@ -161,7 +144,7 @@ class MCPOAuth2TokenCache(InMemoryCache):
server.server_id,
expires_in,
)
return access_token, ttl
return token.access_token, ttl
def invalidate(self, server_id: str) -> None:
"""Remove a cached token (e.g. after a 401)."""

View file

@ -79,6 +79,9 @@ async def test_exchange_token_success():
assert data["grant_type"] == TOKEN_EXCHANGE_GRANT_TYPE
assert data["subject_token"] == "user-jwt-xyz"
assert data["subject_token_type"] == "urn:ietf:params:oauth:token-type:access_token"
assert (
data["requested_token_type"] == "urn:ietf:params:oauth:token-type:access_token"
)
assert data["audience"] == "api://mcp-server"
assert data["scope"] == "mcp.tools.read mcp.tools.execute"
assert data["client_id"] == "litellm-client-id"
@ -258,7 +261,7 @@ async def test_exchange_token_missing_access_token():
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
),
pytest.raises(ValueError, match="missing 'access_token'"),
pytest.raises(ValueError, match="not a valid RFC 6749 token payload"),
):
await handler.exchange_token("jwt", server)

View file

@ -170,6 +170,6 @@ async def test_non_dict_response_raises_value_error():
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client",
return_value=mock_client,
),
pytest.raises(ValueError, match="non-object JSON"),
pytest.raises(ValueError, match="not a valid RFC 6749 token payload"),
):
await resolve_mcp_auth(server)