feat(mcp): support upstream OAuth client metadata identities (#45231)

* feat(mcp): support upstream OAuth client metadata identities

* fix(mcp): load client metadata on cold OAuth requests

* fix(mcp): discover client metadata with configured OAuth endpoints

* fix(mcp): preserve configured endpoints when discovery fails

* fix(mcp): retain CIMD identities across refresh and catalog reloads

* fix(mcp): reuse saved CIMD identity for token endpoint refresh

* fix(mcp): keep dynamic client registration ahead of CIMD unless the deployment opts in

A provider that advertises both a registration endpoint and client ID
metadata documents now gets the registration flow the gateway used
before, so a gateway on a private network keeps working against it;
the metadata document identity applies when the provider offers no
registration, or when general_settings.mcp_prefer_client_id_metadata_document
is true. The saved CIMD refresh identity compares tokens as bytes so a
non-ASCII refresh token cannot crash the token route.

* chore(ui): regenerate dashboard API types for the new MCP general setting

---------

Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
joshua-berri 2026-10-08 04:10:34 -07:00 • committed by GitHub
parent 3a3b8cad80
commit d0202ac364
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
20 changed files with 997 additions and 51 deletions

View file

@ -141,6 +141,16 @@ jobs:
dist: loadscope
timeout: 15
- test-group: mcp-oauth
test-path: >-
tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py
tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py
tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py
tests/unit/proxy/_experimental/mcp_server/outbound_credentials
workers: 4
dist: loadscope
timeout: 15
# ---- logging: split into 2 shards ----
- test-group: custom-logging
test-path: >-

View file

@ -72,6 +72,7 @@ def _configuration_identity(server: MCPServer) -> str:
"token_url",
"registration_url",
"authorization_response_iss_parameter_supported",
"client_id_metadata_document_supported",
)
)
| (frozenset() if server.issuer_is_anchored else frozenset(("issuer",))),

View file

@ -223,6 +223,7 @@ class _OAuthCredentialAccessToken(TypedDict):
class OAuthCredentialPayload(_OAuthCredentialAccessToken, total=False):
identity_binding_proof: ReadOnly[str]
cimd_client_id: ReadOnly[str]
type: str
refresh_token: str
expires_at: str
@ -1765,6 +1766,7 @@ async def store_user_oauth_credential(
scopes: list[str] | None = None,
skip_byok_guard: bool = False,
identity_binding_proof: str | None = None,
cimd_client_id: str | None = None,
) -> None:
"""Persist an OAuth2 access token for a user+server pair.
@ -1782,6 +1784,7 @@ async def store_user_oauth_credential(
"access_token": access_token,
"connected_at": datetime.now(timezone.utc).isoformat(),
**({"identity_binding_proof": identity_binding_proof} if identity_binding_proof else {}),
**({"cimd_client_id": cimd_client_id} if cimd_client_id else {}),
}
if refresh_token:
payload["refresh_token"] = refresh_token
@ -2063,6 +2066,7 @@ async def refresh_user_oauth_token(
auth_method=getattr(server, "token_endpoint_auth_method", None),
client_id=client_id,
client_secret=client_secret,
cimd_client_id=cred.get("cimd_client_id"),
)
token_data: Final[dict[str, str]] = {
"grant_type": "refresh_token",
@ -2132,6 +2136,7 @@ async def refresh_user_oauth_token(
expires_in=expires_in,
scopes=scopes,
identity_binding_proof=binding_proof,
cimd_client_id=cred.get("cimd_client_id"),
skip_byok_guard=True, # Row is already OAuth2; skip the extra find_unique check
)

View file

@ -79,9 +79,13 @@ from litellm.proxy._experimental.mcp_server.oauth_identity_binding import (
enforce_oauth_identity_binding,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
CIMD_METADATA_PATH,
TOKEN_NO_CACHE_HEADERS,
build_upstream_oauth2_token_request,
get_cimd_client_id,
get_cimd_document_url,
get_request_base_url,
needs_cimd_discovery,
oauth_client_registration_matches,
resolve_upstream_resource,
validate_trusted_redirect_uri,
@ -638,6 +642,7 @@ async def _store_per_user_token_server_side(
user_id: str,
token_response: dict[str, Any],
identity_binding_proof: str | None = None,
cimd_client_id: str | None = None,
) -> None:
"""Persist the OAuth token server-side and warm the Redis cache.
@ -680,6 +685,7 @@ async def _store_per_user_token_server_side(
expires_in=expires_in,
scopes=scopes,
identity_binding_proof=identity_binding_proof,
**({"cimd_client_id": cimd_client_id} if cimd_client_id is not None else {}),
)
verbose_logger.info(
"_store_per_user_token_server_side: stored token for user=%s server=%s",
@ -778,20 +784,20 @@ async def _server_with_oauth_endpoints(
mcp_server: MCPServer,
needed_endpoint: Callable[[MCPServer], str | None],
) -> MCPServer:
"""Join deferred OAuth discovery only when the endpoint this caller needs is still missing.
"""Join deferred discovery for missing endpoints or unknown public-client metadata.
Admin-entered endpoints live on ``configured_*`` after an anchored issuer empties the
resolved fields. A caller whose needed endpoint already resolves never awaits discovery
and cannot 503 over a leftover pin. A server still missing it joins the deferred task;
no slot is a no-op and the caller 400s.
Capability discovery is optional when manual endpoints already resolve; the manager
preserves those endpoints if discovery fails. No discovery slot remains a no-op.
"""
if needed_endpoint(mcp_server) is not None:
if needed_endpoint(mcp_server) is not None and not needs_cimd_discovery(mcp_server):
return mcp_server
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # circular import with mcp_server_manager at module load
global_mcp_server_manager,
)
return await global_mcp_server_manager.ensure_oauth_metadata_discovered(mcp_server)
if needed_endpoint(mcp_server) is None:
return await global_mcp_server_manager.ensure_oauth_metadata_discovered(mcp_server)
return await global_mcp_server_manager.ensure_oauth_metadata_discovered(mcp_server, needed_endpoint=needed_endpoint)
def _raise_unless_oauth2_discovery_server(
@ -979,7 +985,8 @@ async def authorize_with_server(
binding: Final = resolved_server.oauth_identity_binding
enforce_binding: Final = binding is not None and binding.mode == "enforce"
if enforce_binding:
cimd_client_id: Final = get_cimd_client_id(resolved_server)
if enforce_binding or cimd_client_id:
_require_s256_pkce(code_challenge, code_challenge_method)
if resolved_server.is_dcr_bridge:
@ -1043,7 +1050,7 @@ async def authorize_with_server(
relay_state: Final = secrets.token_urlsafe(_OAUTH_STATE_HANDLE_BYTES)
params: Final = {
"client_id": resolved_server.client_id if resolved_server.client_id else client_id,
"client_id": resolved_server.client_id or cimd_client_id or client_id,
"redirect_uri": f"{request_base_url}/callback",
"state": relay_state,
"response_type": response_type or "code",
@ -1080,7 +1087,37 @@ def _token_credential_source(mcp_server: MCPServer) -> CredentialSource:
"""Mirrors the resolved-client rule in :func:`exchange_token_with_server`: when the server has a
stored client_id the gateway presents its own credentials upstream, so a credential rejection is
the operator's fault, not the caller's."""
return "gateway_stored" if mcp_server.client_id else "caller_supplied"
return "gateway_stored" if mcp_server.client_id or get_cimd_client_id(mcp_server) else "caller_supplied"
async def _saved_cimd_refresh_client_id(
server: MCPServer, user_id: str | None, refresh_token: str | None
) -> str | None:
"""Reuse only the client identity bound to this caller's presented refresh grant."""
if (
not user_id
or not refresh_token
or not server.needs_user_oauth_token
or server.auth_type != MCPAuth.oauth2
or server.client_id
or server.client_secret
):
return None
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server.db import get_user_oauth_credential
if proxy_server.prisma_client is None:
return None
try:
credential: Final = await get_user_oauth_credential(proxy_server.prisma_client, user_id, server.server_id)
except Exception: # noqa: BLE001 # optional storage must not prevent a caller-owned OAuth exchange
return None
if credential is None:
return None
stored_refresh: Final = credential.get("refresh_token")
if not stored_refresh or not secrets.compare_digest(stored_refresh.encode(), refresh_token.encode()):
return None
return credential.get("cimd_client_id")
async def exchange_token_with_server(
@ -1117,6 +1154,17 @@ async def exchange_token_with_server(
),
)
request_user_id: Final = (
await extract_user_id_from_request(request)
if resolved_server.needs_user_oauth_token or resolved_server.oauth_identity_binding is not None
else None
)
cimd_client_id: Final = (
await _saved_cimd_refresh_client_id(resolved_server, request_user_id, refresh_token)
if grant_type == "refresh_token"
else None
) or get_cimd_client_id(resolved_server)
# The id, secret, and token-endpoint auth method must come from the same source. When the
# server-side client_id wins, falling back to the caller's secret pairs the persisted client
# with a foreign secret; the register short-circuit hands clients a placeholder secret
@ -1138,16 +1186,11 @@ async def exchange_token_with_server(
auth_method=resolved_auth_method,
client_id=resolved_client_id,
client_secret=resolved_client_secret,
cimd_client_id=cimd_client_id,
)
except TokenEndpointAuthConfigError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
request_user_id: Final = (
await extract_user_id_from_request(request)
if resolved_server.needs_user_oauth_token or resolved_server.oauth_identity_binding is not None
else None
)
bridge_identity: _BridgeAuthorizationCode | None = None
bridge_mint_ready: _BridgeMintReady | None = None
bridge_upstream_refresh: SecretStr | None = None
@ -1265,7 +1308,7 @@ async def exchange_token_with_server(
except httpx.HTTPStatusError as exc:
fault: Final = classify_upstream_token_rejection(
exc.response,
credential_source=_token_credential_source(resolved_server),
credential_source="gateway_stored" if cimd_client_id else _token_credential_source(resolved_server),
log_context=resolved_server.server_id,
)
upstream_rejected_bridge_refresh: Final = (
@ -1336,6 +1379,7 @@ async def exchange_token_with_server(
user_id=user_id,
token_response=token_response,
identity_binding_proof=binding_proof,
**({"cimd_client_id": cimd_client_id} if cimd_client_id is not None else {}),
)
else:
verbose_logger.warning(
@ -1991,8 +2035,21 @@ async def register_client_with_server(
),
)
cimd_client_id: Final = get_cimd_client_id(resolved_server)
if cimd_client_id:
return {
"client_id": cimd_client_id,
"token_endpoint_auth_method": "none",
"redirect_uris": client_facing_redirect_uris,
}
registration_url: Final = resolved_server.effective_registration_url
if registration_url is None:
if resolved_server.client_id_metadata_document_supported and resolved_server.is_gateway_managed_oauth2:
raise HTTPException(
status_code=400,
detail="CIMD requires a stable HTTPS PROXY_BASE_URL and public-client authentication; "
"configure these or provide a pre-registered OAuth client",
)
return dummy_return
bridge_relay: Final = _dcr_bridge_relays_client_registration(resolved_server)
@ -3162,3 +3219,25 @@ async def register_client(request: Request, mcp_server_name: str | None = None):
client_redirect_uris=client_redirect_uris,
client_application_type=client_application_type,
)
@router.get(CIMD_METADATA_PATH, include_in_schema=False)
async def oauth_client_metadata() -> JSONResponse:
from mcp.shared.auth import OAuthClientInformationFull
from pydantic import AnyUrl
document_url: Final = get_cimd_document_url()
if document_url is None:
raise HTTPException(status_code=404, detail="CIMD requires a configured HTTPS PROXY_BASE_URL")
base_url: Final = document_url.removesuffix(CIMD_METADATA_PATH)
metadata: Final = OAuthClientInformationFull(
client_id=document_url,
client_name="LiteLLM MCP Gateway",
redirect_uris=[AnyUrl(f"{base_url}/callback")],
token_endpoint_auth_method="none",
grant_types=["authorization_code", "refresh_token"],
response_types=["code"],
)
return JSONResponse(
metadata.model_dump(mode="json", exclude_none=True), headers={"Cache-Control": "public, max-age=300"}
)

View file

@ -105,6 +105,7 @@ from litellm.proxy._experimental.mcp_server.oauth_utils import ( # noqa: F401
_redact_mcp_resource_url, # pyright: ignore[reportPrivateUsage,reportUnusedImport] # backwards-compatible package export
canonicalize_url_identity,
get_byok_www_authenticate,
needs_cimd_discovery,
redact_mcp_resource_url,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
@ -452,6 +453,7 @@ class _AuthorizationServerMetadataPayload(TypedDict, total=False):
authorization_endpoint: str
token_endpoint: str
registration_endpoint: str
client_id_metadata_document_supported: ReadOnly[object]
scopes_supported: Sequence[str]
grant_types_supported: Sequence[str]
token_endpoint_auth_methods_supported: Sequence[str]
@ -767,14 +769,16 @@ def _flow_endpoints_missing(
return authorization_url is None or token_url is None
def oauth_endpoints_unresolved(server: MCPServer) -> bool:
"""``_flow_endpoints_missing`` over a built registry entry, for the reload fast-path check.
def oauth_endpoints_unresolved(server: MCPServer, *, include_client_metadata: bool = True) -> bool:
"""Whether endpoint or eligible client metadata discovery is still pending.
The flow comes from ``effective_oauth2_flow``, the one column-first, shape-fallback judge every
flow decision uses, not from the raw column: a legacy row the startup backfill deliberately left
unstamped (the ambiguous M2M shape) serves M2M at request time, and reading the bare column here
would classify it as interactive-missing-endpoints and re-run discovery on every reload.
"""
if include_client_metadata and needs_cimd_discovery(server):
return True
if (
server.auth_type == MCPAuth.oauth2_token_exchange
and server.token_exchange_profile == "entra_obo"
@ -868,6 +872,8 @@ def carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_serve
may_carry: Final = _endpoints_corroborate_authorization_url(
previous_server.authorization_url, new_server.authorization_url
)
if may_carry and new_server.client_id_metadata_document_supported is None:
new_server.client_id_metadata_document_supported = previous_server.client_id_metadata_document_supported
if may_carry and new_server.issuer is None:
new_server.issuer = previous_server.issuer
new_server.authorization_response_iss_parameter_supported = ( # rebind-ok: publish on the existing rebuild object
@ -908,7 +914,7 @@ def _restrict_discovery_to_corroborated_authorization_server(
return metadata
if _endpoints_corroborate_authorization_url(metadata.authorization_url, manual_authorization_url):
return metadata
if not metadata.token_url and not metadata.registration_url:
if not metadata.token_url and not metadata.registration_url and not metadata.client_id_metadata_document_supported:
return metadata
bridge_note: Final = (
" The discovered registration_url is rejected with it, so this dcr_bridge server stays on the"
@ -926,7 +932,9 @@ def _restrict_discovery_to_corroborated_authorization_server(
_normalized_authorize_endpoint(manual_authorization_url),
bridge_note,
)
return metadata.model_copy(update={"token_url": None, "registration_url": None})
return metadata.model_copy(
update={"token_url": None, "registration_url": None, "client_id_metadata_document_supported": False}
)
def _redacted_origin_list(urls: Sequence[str]) -> str:
@ -2060,6 +2068,7 @@ class MCPServerManager:
update={
"scopes": server.scopes or metadata.scopes,
"issuer": server.issuer or discovered_issuer,
"client_id_metadata_document_supported": metadata.client_id_metadata_document_supported,
"authorization_response_iss_parameter_supported": (
metadata.authorization_response_iss_parameter_supported
if discovered_issuer is not None
@ -2224,12 +2233,27 @@ class MCPServerManager:
if should_defer != has_slot:
self._set_oauth_discovery_deferred(server.server_id, should_defer)
async def ensure_oauth_metadata_discovered(self, server: MCPServer, *, _retry_stale: bool = True) -> MCPServer:
async def ensure_oauth_metadata_discovered(
self,
server: MCPServer,
*,
needed_endpoint: Callable[[MCPServer], str | None] | None = None,
_retry_stale: bool = True,
) -> MCPServer:
return await self.catalog.resolve_oauth_metadata(
server, lambda selected: self._ensure_oauth_metadata_discovered(selected, _retry_stale=_retry_stale)
server,
lambda selected: self._ensure_oauth_metadata_discovered(
selected, needed_endpoint=needed_endpoint, _retry_stale=_retry_stale
),
)
async def _ensure_oauth_metadata_discovered(self, server: MCPServer, *, _retry_stale: bool = True) -> MCPServer:
async def _ensure_oauth_metadata_discovered(
self,
server: MCPServer,
*,
needed_endpoint: Callable[[MCPServer], str | None] | None = None,
_retry_stale: bool = True,
) -> MCPServer:
"""Join the bounded discovery task and return the resolved server.
Concurrent callers share one task per server. A failed attempt remains
@ -2237,9 +2261,12 @@ class MCPServerManager:
Args:
server: The MCP server whose OAuth metadata must be resolved.
needed_endpoint: A caller-specific endpoint that may remain usable when
optional capability discovery fails.
Returns:
The resolved server; the registered server when no discovery is
The resolved server; a configured caller endpoint remains usable on
optional capability-discovery failure. The registered server when no discovery is
pending, or when discovery failed for a client-forwarded-token
server, whose session consumes no discovered endpoint.
@ -2257,18 +2284,26 @@ class MCPServerManager:
outcome: Final = await asyncio.shield(task)
except asyncio.CancelledError:
if task.cancelled() and not self._oauth_discovery_slot_is_current(server.server_id, generation):
return await self._rejoin_oauth_metadata_discovery(server, retry_stale=_retry_stale)
return await self._rejoin_oauth_metadata_discovery(
server, needed_endpoint=needed_endpoint, retry_stale=_retry_stale
)
raise
match outcome:
case _OAuthDiscoveryResolved(resolved_server):
self.catalog.assert_current(resolved_server)
return resolved_server
case _OAuthDiscoveryStale():
return await self._rejoin_oauth_metadata_discovery(server, retry_stale=_retry_stale)
return await self._rejoin_oauth_metadata_discovery(
server, needed_endpoint=needed_endpoint, retry_stale=_retry_stale
)
case _OAuthDiscoveryFailed(timed_out=timed_out):
current: Final = self._registered_server(server)
self.catalog.assert_current(current)
if current.is_client_forwarded_token:
if (
current.is_client_forwarded_token
or not oauth_endpoints_unresolved(current, include_client_metadata=False)
or (needed_endpoint is not None and needed_endpoint(current) is not None)
):
return current
server_ref: Final = current.alias or current.server_name or current.name or current.server_id
reason: Final = "timed out" if timed_out else "returned incomplete metadata"
@ -2279,11 +2314,19 @@ class MCPServerManager:
return assert_never(outcome)
async def _rejoin_oauth_metadata_discovery(self, server: MCPServer, *, retry_stale: bool) -> MCPServer:
async def _rejoin_oauth_metadata_discovery(
self, server: MCPServer, *, needed_endpoint: Callable[[MCPServer], str | None] | None = None, retry_stale: bool
) -> MCPServer:
if retry_stale:
return await self.ensure_oauth_metadata_discovered(server, _retry_stale=False)
return await self.ensure_oauth_metadata_discovered(
server, needed_endpoint=needed_endpoint, _retry_stale=False
)
current: Final = self._registered_server(server)
if not oauth_endpoints_unresolved(current) or current.is_client_forwarded_token:
if (
not oauth_endpoints_unresolved(current, include_client_metadata=False)
or current.is_client_forwarded_token
or (needed_endpoint is not None and needed_endpoint(current))
):
return current
raise HTTPException(status_code=503, detail="OAuth metadata discovery changed repeatedly; retry shortly")
@ -2616,6 +2659,9 @@ class MCPServerManager:
scopes=resolved_scopes,
configured_scopes=tuple(configured_scopes) if configured_scopes else None,
issuer=effective_issuer,
client_id_metadata_document_supported=(
gated_oauth_metadata.client_id_metadata_document_supported if gated_oauth_metadata else None
),
authorization_response_iss_parameter_supported=(
gated_oauth_metadata.authorization_response_iss_parameter_supported
if gated_oauth_metadata
@ -3194,6 +3240,9 @@ class MCPServerManager:
scopes=resolved_scopes,
configured_scopes=configured_scopes,
issuer=effective_issuer,
client_id_metadata_document_supported=(
gated_oauth_metadata.client_id_metadata_document_supported if gated_oauth_metadata else None
),
authorization_response_iss_parameter_supported=(
gated_oauth_metadata.authorization_response_iss_parameter_supported if gated_oauth_metadata else False
),
@ -5304,6 +5353,7 @@ class MCPServerManager:
authorization_url=data.get("authorization_endpoint"),
token_url=data.get("token_endpoint"),
registration_url=data.get("registration_endpoint"),
client_id_metadata_document_supported=data.get("client_id_metadata_document_supported") is True,
discovered_issuer=claimed_issuer if isinstance(claimed_issuer, str) and claimed_issuer else None,
authorization_response_iss_parameter_supported=data.get(
"authorization_response_iss_parameter_supported"
@ -5336,6 +5386,7 @@ class MCPServerManager:
return MCPOAuthMetadata(
authorization_url=f"{base}/oauth2/v2.0/authorize",
token_url=f"{base}/oauth2/v2.0/token",
client_id_metadata_document_supported=False,
)
@staticmethod

View file

@ -130,6 +130,57 @@ def _resolve_proxy_base_url_env() -> str | None:
return None
CIMD_METADATA_PATH: Final = "/oauth/client-metadata.json"
def get_cimd_document_url() -> str | None:
try:
configured: Final = _resolve_proxy_base_url_env()
parsed: Final = urlparse(configured or "")
except ValueError:
return None
if parsed.scheme != "https" or not parsed.hostname or parsed.username is not None or parsed.password is not None:
return None
return f"{configured}{CIMD_METADATA_PATH}"
def _can_use_cimd(server: "MCPServer") -> bool:
return (
server.is_gateway_managed_oauth2
and server.needs_user_oauth_token
and not server.client_id
and not server.client_secret
and server.token_endpoint_auth_method != "client_secret_basic"
)
def needs_cimd_discovery(server: "MCPServer") -> bool:
"""Resolve unknown client metadata support even when OAuth endpoints are configured."""
return (
getattr(server, "client_id_metadata_document_supported", False) is None
and _can_use_cimd(server)
and get_cimd_document_url() is not None
)
def _deployment_prefers_cimd() -> bool:
from litellm.proxy.proxy_server import general_settings
return general_settings.get("mcp_prefer_client_id_metadata_document") is True
def _dynamic_registration_takes_precedence(server: "MCPServer") -> bool:
return server.effective_registration_url is not None and not _deployment_prefers_cimd()
def get_cimd_client_id(server: "MCPServer") -> str | None:
if getattr(server, "client_id_metadata_document_supported", False) is not True or not _can_use_cimd(server):
return None
if _dynamic_registration_takes_precedence(server):
return None
return get_cimd_document_url()
BYOK_RESOURCE_METADATA_PATH: Final = "/v1/mcp/oauth/protected-resource"
@ -748,6 +799,7 @@ def build_upstream_oauth2_token_request(
auth_method: object,
client_id: str | None,
client_secret: str | None,
cimd_client_id: str | None = None,
) -> TokenEndpointClientAuth:
"""Client auth plus the RFC 8707 ``resource`` for one upstream plain-OAuth2 token request.
@ -757,10 +809,11 @@ def build_upstream_oauth2_token_request(
authenticate as the caller's own client rather than the server's; ``resource`` always comes from
the server, so no leg can choose or forget it.
"""
selected_cimd_id: Final = cimd_client_id or get_cimd_client_id(mcp_server)
client_auth: Final = build_token_endpoint_client_auth(
auth_method=normalize_token_endpoint_auth_method(auth_method),
client_id=client_id,
client_secret=client_secret,
auth_method=None if selected_cimd_id else normalize_token_endpoint_auth_method(auth_method),
client_id=selected_cimd_id or client_id,
client_secret=None if selected_cimd_id else client_secret,
)
resource: Final = resolve_upstream_resource(mcp_server)
if not resource:

View file

@ -25,7 +25,7 @@ from litellm.proxy._experimental.mcp_server.oauth_identity_binding import (
RefreshTokenPresented,
enforce_oauth_identity_binding,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import build_upstream_oauth2_token_request
from litellm.proxy._experimental.mcp_server.oauth_utils import build_upstream_oauth2_token_request, get_cimd_client_id
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
@ -47,6 +47,7 @@ class CredentialPersist(Protocol):
expires_in: int | None,
scopes: tuple[str, ...] | None,
identity_binding_proof: str | None = None,
cimd_client_id: str | None = None,
) -> None: ...
@ -112,12 +113,14 @@ class AuthorizationCodeRefresher:
if not token_url:
return None
cimd_client_id: Final = token.cimd_client_id or get_cimd_client_id(server)
try:
token_request: Final = build_upstream_oauth2_token_request(
server,
auth_method=server.token_endpoint_auth_method,
client_id=server.client_id,
client_secret=server.client_secret,
cimd_client_id=cimd_client_id,
)
except TokenEndpointAuthConfigError as exc:
verbose_logger.warning("MCP OAuth refresh misconfigured for server %s: %s", server_id, exc)
@ -155,22 +158,21 @@ class AuthorizationCodeRefresher:
expires_in: Final = _parse_expires_in(body.get("expires_in"))
scopes: Final = _parse_scopes(body.get("scope")) or token.scopes
if binding_proof is not None:
await self._persist(
user_id,
server_id,
access_token,
new_refresh,
expires_in,
scopes or None,
identity_binding_proof=binding_proof,
)
else:
await self._persist(user_id, server_id, access_token, new_refresh, expires_in, scopes or None)
await self._persist(
user_id,
server_id,
access_token,
new_refresh,
expires_in,
scopes or None,
**({"identity_binding_proof": binding_proof} if binding_proof is not None else {}),
**({"cimd_client_id": cimd_client_id} if cimd_client_id is not None else {}),
)
return OAuthToken(
access_token=access_token,
expires_at=self._clock() + expires_in if expires_in is not None else None,
refresh_token=new_refresh,
scopes=scopes,
identity_binding_proof=binding_proof,
cimd_client_id=cimd_client_id,
)

View file

@ -42,6 +42,7 @@ class OAuthToken:
refresh_token: str | None = None
scopes: tuple[str, ...] = ()
identity_binding_proof: str | None = None
cimd_client_id: str | None = None
def __repr__(self) -> str:
has_refresh: Final = self.refresh_token is not None

View file

@ -72,6 +72,7 @@ async def _persist_credential(
expires_in: int | None,
scopes: tuple[str, ...] | None,
identity_binding_proof: str | None = None,
cimd_client_id: str | None = None,
) -> None:
from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415
store_user_oauth_credential,
@ -90,6 +91,7 @@ async def _persist_credential(
scopes=list(scopes) if scopes else None,
skip_byok_guard=True,
identity_binding_proof=identity_binding_proof,
**({"cimd_client_id": cimd_client_id} if cimd_client_id is not None else {}),
)

View file

@ -47,12 +47,14 @@ def _to_oauth_token(payload: Mapping[str, object]) -> OAuthToken | None:
refresh_token: Final = payload.get("refresh_token")
expires_at: Final = payload.get("expires_at")
binding_proof: Final = payload.get("identity_binding_proof")
cimd_client_id: Final = payload.get("cimd_client_id")
return OAuthToken(
access_token=access_token,
expires_at=_iso_to_epoch(expires_at) if isinstance(expires_at, str) else None,
refresh_token=refresh_token if isinstance(refresh_token, str) else None,
scopes=_to_scopes(payload.get("scopes")),
identity_binding_proof=binding_proof if isinstance(binding_proof, str) else None,
cimd_client_id=cimd_client_id if isinstance(cimd_client_id, str) else None,
)

View file

@ -166,6 +166,7 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = (
name="mcp_discoverable",
module_path="litellm.proxy._experimental.mcp_server.discoverable_endpoints",
path_prefixes=(
"/oauth/client-metadata.json",
"/.well-known/oauth-",
"/.well-known/openid-configuration",
"/.well-known/jwks.json",

View file

@ -3185,6 +3185,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
ge=1,
description="Number of trusted reverse proxies/load balancers in front of the gateway that append to X-Forwarded-For. When set (and mcp_trusted_proxy_ranges validates the direct peer), the client IP for MCP access control is read this many entries from the right of the chain instead of the spoofable leftmost value, defeating append-style X-Forwarded-For forgery.",
)
mcp_prefer_client_id_metadata_document: bool | None = Field(
None,
description="When true, a gateway-managed OAuth2 MCP server whose authorization server advertises Client ID Metadata Document support identifies itself with the gateway's public metadata document URL even when that authorization server also offers dynamic client registration. Requires a public HTTPS PROXY_BASE_URL the authorization server can fetch. Default false: dynamic client registration is used whenever the authorization server offers it, and the metadata document only when it does not.",
)
trusted_proxy_ranges: list[str] | None = Field(
None,
description="CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler, and whose X-Forwarded-For is used to attribute Admin UI sign-in attempts to a source address. Set it to an empty list when clients connect directly, so the peer address is the source. Left unset, or containing an entry that is not an address or CIDR range, the per-source sign-in limit is off.",

View file

@ -18538,6 +18538,7 @@ _GENERAL_SETTINGS_CONFIG_LIST_FIELD_TYPES: Final[Mapping[str, str]] = MappingPro
"mcp_client_id_header": "String",
"mcp_trusted_proxy_ranges": "List",
"mcp_xff_num_trusted_hops": "Integer",
"mcp_prefer_client_id_metadata_document": "Boolean",
"always_include_stream_usage": "Boolean",
"forward_client_headers_to_llm_api": "Boolean",
"mcp_required_fields": "List",

View file

@ -39,6 +39,7 @@ class MCPOAuthMetadata(LiteLLMBaseModel):
authorization_url: str | None = None
token_url: str | None = None
registration_url: str | None = None
client_id_metadata_document_supported: bool | None = None
authorization_response_iss_parameter_supported: bool = False
discovered_issuer: str | None = None
"""The ``issuer`` the authorization-server metadata document self-attests (RFC 8414). Persisted
@ -120,6 +121,7 @@ class MCPServer(LiteLLMBaseModel):
client_secret: str | None = None
issuer: str | None = None
issuer_is_anchored: bool = False
client_id_metadata_document_supported: bool | None = None
authorization_response_iss_parameter_supported: bool = False
dcr_issuer: str | None = None
dcr_server_url: str | None = None

View file

@ -52,7 +52,7 @@ def _endpoint(body, sink=None):
def _recording_persist(sink):
async def persist(
user_id, server_id, access_token, refresh_token, expires_in, scopes
user_id, server_id, access_token, refresh_token, expires_in, scopes, cimd_client_id=None
):
sink.append(
(user_id, server_id, access_token, refresh_token, expires_in, scopes)
@ -325,3 +325,83 @@ async def test_verified_refresh_preserves_binding_proof_in_storage():
assert token.refresh_token == "rotated"
assert token.identity_binding_proof == "verified-binding"
assert persist.await_args.kwargs["identity_binding_proof"] == "verified-binding"
@pytest.mark.asyncio
async def test_cimd_refresh_on_fresh_replica_preserves_user_and_client_identity(monkeypatch):
from litellm.types.mcp_server.mcp_server_manager import MCPServer
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
server = MCPServer.model_validate(
{
"server_id": "srv",
"name": "srv",
"server_name": "srv",
"transport": "http",
"url": "https://mcp.example.com/mcp",
"auth_type": "oauth2",
"oauth2_flow": "authorization_code",
"token_url": "https://idp.example.com/token",
"client_id_metadata_document_supported": True,
}
)
posted = []
persisted = []
refreshed = await _refresher(
server=server,
body={"access_token": "new-at", "expires_in": 3600},
post_sink=posted,
persist_sink=persisted,
).refresh("alice", "srv", OAuthToken(access_token="old-at", refresh_token="old-rt"))
assert refreshed is not None
assert refreshed.access_token == "new-at"
assert posted == [
(
"https://idp.example.com/token",
{
"grant_type": "refresh_token",
"refresh_token": "old-rt",
"client_id": "https://gateway.example.com/oauth/client-metadata.json",
},
{},
)
]
assert persisted == [("alice", "srv", "new-at", "old-rt", 3600, None)]
assert server.client_id is None
@pytest.mark.asyncio
async def test_saved_cimd_identity_refreshes_without_discovery(monkeypatch):
from litellm.proxy._experimental.mcp_server.outbound_credentials.v2_token_store import V2PerUserTokenStore
from litellm.types.mcp_server.mcp_server_manager import MCPServer
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
identity = "https://gateway.example.com/oauth/client-metadata.json"
stored = {"access_token": "old", "refresh_token": "old-rt", "cimd_client_id": identity}
async def read_credential(user_id, server_id):
assert (user_id, server_id) == ("alice", "srv")
return stored
async def persist(user_id, server_id, access_token, refresh_token, expires_in, scopes, **metadata):
assert (user_id, server_id, access_token) == ("alice", "srv", "new")
assert metadata["cimd_client_id"] == identity
async def post(url, form, headers):
assert url == "https://idp.example.com/token"
assert form["client_id"] == identity
assert "client_secret" not in form
assert "Authorization" not in headers
return {"access_token": "new", "refresh_token": "rotated"}
server = MCPServer(
server_id="srv", name="srv", transport="http", auth_type="oauth2", oauth2_flow="authorization_code",
authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token",
)
token = await V2PerUserTokenStore(read_credential).fetch("alice", "srv")
assert token is not None
refreshed = await AuthorizationCodeRefresher(lambda _: server, post, persist).refresh("alice", "srv", token)
assert refreshed is not None
assert refreshed.access_token == "new"
assert refreshed.refresh_token == "rotated"
assert refreshed.cimd_client_id == identity

View file

@ -708,7 +708,8 @@ async def test_store_user_oauth_credential_does_not_persist_plaintext():
@pytest.mark.asyncio
async def test_oauth_round_trip_returns_payload():
@pytest.mark.parametrize("cimd_client_id", [None, "https://gateway.example.com/oauth/client-metadata.json"])
async def test_oauth_round_trip_returns_payload(cimd_client_id):
access_token = "ya29.a0AfH6SMBverysecretaccesstoken"
prisma = _make_prisma_with_existing(row=None)
await store_user_oauth_credential(
@ -719,6 +720,7 @@ async def test_oauth_round_trip_returns_payload():
refresh_token="rfr-xyz",
scopes=["a", "b"],
identity_binding_proof="verified-proof",
cimd_client_id=cimd_client_id,
)
stored = _stored_value(prisma)
@ -734,6 +736,7 @@ async def test_oauth_round_trip_returns_payload():
assert result["refresh_token"] == "rfr-xyz"
assert result["scopes"] == ["a", "b"]
assert result["identity_binding_proof"] == "verified-proof"
assert result.get("cimd_client_id") == cimd_client_id
@pytest.mark.asyncio
@ -1795,3 +1798,61 @@ async def test_unverified_legacy_cache_cannot_bypass_enforcement(monkeypatch):
await mcp_per_user_token_cache.set("alice", "srv", "bob", 60)
assert await module.resolve_user_oauth_access_token("alice", server) is None
assert await mcp_per_user_token_cache.get("alice", "srv") is None
@pytest.mark.asyncio
@pytest.mark.parametrize("owner", ["native", "legacy"])
@pytest.mark.parametrize("capability", [None, False])
async def test_saved_cimd_grant_refreshes_from_encrypted_storage_on_a_fresh_replica(monkeypatch, respx_mock, owner, capability):
from urllib.parse import parse_qs
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import db as db_module
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _store_per_user_token_server_side
from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import (
LazyPerUserOAuthTokenStore,
)
from litellm.types.mcp_server.mcp_server_manager import MCPServer
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
identity = "https://gateway.example.com/oauth/client-metadata.json"
prisma = _make_prisma_with_existing(row=None)
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
server = MCPServer(
server_id="saved-cimd", name="saved_cimd", transport="http", auth_type="oauth2",
client_id_metadata_document_supported=capability,
oauth2_flow="authorization_code", authorization_url="https://idp.example.com/authorize",
token_url="https://idp.example.com/token",
)
await _store_per_user_token_server_side(
server, "alice", {"access_token": "old", "refresh_token": "old-refresh", "expires_in": -10},
cimd_client_id=identity,
)
async def read_row(**kwargs):
return SimpleNamespace(credential_b64=_stored_value(prisma), user_id="alice", server_id=server.server_id)
prisma.db.litellm_mcpusercredentials.find_unique.side_effect = read_row
saved = await get_user_oauth_credential(prisma, "alice", server.server_id)
assert saved is not None and saved["cimd_client_id"] == identity
assert identity not in _stored_value(prisma)
monkeypatch.setenv("PROXY_BASE_URL", "https://renamed-gateway.example.com")
upstream = respx_mock.post("https://idp.example.com/token").respond(
200, json={"access_token": "fresh", "refresh_token": "rotated", "expires_in": 3600}
)
if owner == "native":
refreshed = await LazyPerUserOAuthTokenStore(lambda _: server).fetch("alice", server.server_id)
assert refreshed is not None and refreshed.access_token == "fresh"
assert refreshed.cimd_client_id == identity
else:
legacy = await db_module.refresh_user_oauth_token(prisma, "alice", server, saved)
assert legacy is not None and legacy["access_token"] == "fresh"
persisted = await get_user_oauth_credential(prisma, "alice", server.server_id)
assert persisted is not None
assert persisted["cimd_client_id"] == identity
assert persisted["refresh_token"] == "rotated"
assert upstream.call_count == 1
assert parse_qs(upstream.calls[0].request.content.decode())["client_id"] == [identity]
assert "client_secret" not in parse_qs(upstream.calls[0].request.content.decode())
assert server.client_id is None and server.client_id_metadata_document_supported is capability

View file

@ -20,9 +20,9 @@ if TYPE_CHECKING:
import httpx
from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey
from fastapi import APIRouter
from respx import MockRouter
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@ -12568,10 +12568,12 @@ async def test_oauth_write_denial_does_not_erase_identity_binding(
@pytest.mark.asyncio
@pytest.mark.parametrize("admin_only", [False, True])
@pytest.mark.parametrize("cimd", [False, True])
async def test_signed_oauth_callback_honors_credential_write_policy(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
monkeypatch: pytest.MonkeyPatch,
admin_only: bool,
cimd: bool,
) -> None:
import httpx
import litellm
@ -12586,7 +12588,8 @@ async def test_signed_oauth_callback_honors_credential_write_policy(
server: Final = MCPServer(
server_id="signed-server", name="signed-server", transport=MCPTransport.http,
auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id="client",
auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id=None if cimd else "client",
client_id_metadata_document_supported=cimd,
token_url="https://upstream.example.test/token",
)
monkeypatch.setattr(proxy_server, "general_settings", {
@ -12594,6 +12597,7 @@ async def test_signed_oauth_callback_honors_credential_write_policy(
"admin_only_routes": [f"/v1/mcp/server/{server.server_id}/oauth-user-credential"] if admin_only else [],
})
monkeypatch.setenv("LITELLM_SALT_KEY", "signed-oauth-test-salt")
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
manager: Final = MagicMock()
manager.get_allowed_mcp_servers = AsyncMock(return_value=[server.server_id])
manager.invalidate_user_oauth_token_cache = AsyncMock()
@ -12630,6 +12634,16 @@ async def test_signed_oauth_callback_honors_credential_write_policy(
"user_id": "jwt-owner", "server_id": server.server_id,
}
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
payload: Final = json.loads(decrypt_value_helper(
table.upsert.call_args.kwargs["data"]["create"]["credential_b64"], "credential_b64"
))
if cimd:
assert payload["cimd_client_id"] == "https://gateway.example.com/oauth/client-metadata.json"
else:
assert "cimd_client_id" not in payload
@pytest.mark.asyncio
@pytest.mark.parametrize("allowed", [False, True])
@ -13362,3 +13376,402 @@ async def test_register_application_type_keeps_no_registration_endpoint_fallback
"redirect_uris": ["https://gateway.example/callback"],
}
assert len(upstream.calls) == 0
def _cimd_oauth_server():
from litellm.types.mcp_server.mcp_server_manager import MCPServer
return MCPServer.model_validate(
{
"server_id": "cimd-server",
"name": "cimd-server",
"server_name": "cimd-server",
"url": "https://mcp.example.com/mcp",
"transport": "http",
"auth_type": "oauth2",
"oauth2_flow": "authorization_code",
"authorization_url": "https://idp.example.com/authorize",
"token_url": "https://idp.example.com/token",
"client_id_metadata_document_supported": True,
}
)
def _cimd_request():
from starlette.requests import Request
return Request(
{
"type": "http",
"method": "GET",
"scheme": "https",
"path": "/",
"root_path": "",
"query_string": b"",
"headers": [],
"server": ("gateway.example.com", 443),
"client": ("127.0.0.1", 10000),
}
)
@pytest.mark.asyncio
async def test_cimd_registration_returns_https_identity_without_dcr(monkeypatch, respx_mock):
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
response = await endpoints.register_client_with_server(
_cimd_request(), _cimd_oauth_server(), "Gateway", None, None, None
)
body = json.loads(response.body) if hasattr(response, "body") else response
assert body["client_id"] == "https://gateway.example.com/oauth/client-metadata.json"
assert "client_secret" not in body
assert len(respx_mock.calls) == 0
@pytest.mark.asyncio
async def test_cimd_authorization_uses_metadata_identity_and_s256(monkeypatch):
from urllib.parse import parse_qs, urlparse
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
monkeypatch.setenv("LITELLM_SALT_KEY", "cimd-test-state-signing-key")
response = await endpoints.authorize_with_server(
_cimd_request(),
_cimd_oauth_server(),
"placeholder",
"https://gateway.example.com/ui/",
code_challenge="a" * 43,
code_challenge_method="S256",
)
params = parse_qs(urlparse(response.headers["location"]).query)
assert params["client_id"] == ["https://gateway.example.com/oauth/client-metadata.json"]
assert params["redirect_uri"] == ["https://gateway.example.com/callback"]
assert params["code_challenge_method"] == ["S256"]
@pytest.mark.asyncio
async def test_cimd_authorization_rejects_missing_pkce(monkeypatch):
from fastapi import HTTPException
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
monkeypatch.setenv("LITELLM_SALT_KEY", "cimd-test-state-signing-key")
with pytest.raises(HTTPException) as exc:
await endpoints.authorize_with_server(
_cimd_request(), _cimd_oauth_server(), "placeholder", "https://gateway.example.com/ui/"
)
assert exc.value.status_code == 400
def test_cimd_refresh_request_uses_same_identity_without_caller_secret(monkeypatch):
from litellm.proxy._experimental.mcp_server.oauth_utils import build_upstream_oauth2_token_request
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
request = build_upstream_oauth2_token_request(
_cimd_oauth_server(), auth_method=None, client_id="placeholder", client_secret="dummy"
)
assert request.body["client_id"] == "https://gateway.example.com/oauth/client-metadata.json"
assert "client_secret" not in request.body
assert "Authorization" not in request.headers
@pytest.mark.parametrize(
"base", [None, "http://gateway.example.com", "invalid", "https://user:secret@gateway.example.com"]
)
@pytest.mark.asyncio
async def test_cimd_without_stable_https_origin_reports_actionable_error(monkeypatch, base, respx_mock):
from fastapi import HTTPException
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
monkeypatch.delenv("PROXY_BASE_URL", raising=False)
if base is not None:
monkeypatch.setenv("PROXY_BASE_URL", base)
with pytest.raises(HTTPException) as exc:
await endpoints.register_client_with_server(_cimd_request(), _cimd_oauth_server(), "Gateway", None, None, None)
assert exc.value.status_code == 400
assert "HTTPS PROXY_BASE_URL" in str(exc.value.detail)
assert len(respx_mock.calls) == 0
@pytest.mark.parametrize("base", [None, "http://gateway.example.com"])
@pytest.mark.asyncio
async def test_cimd_without_https_origin_falls_back_to_available_dcr(monkeypatch, base, respx_mock):
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
monkeypatch.delenv("PROXY_BASE_URL", raising=False)
if base is not None:
monkeypatch.setenv("PROXY_BASE_URL", base)
server = _cimd_oauth_server().model_copy(update={"registration_url": "https://idp.example.com/register"})
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
post = respx_mock.post("https://idp.example.com/register").respond(201, json={"client_id": "registered-client"})
response = await endpoints.register_client_with_server(_cimd_request(), server, "Gateway", None, None, None)
assert json.loads(response.body)["client_id"] == "registered-client"
assert post.call_count == 1
@pytest.mark.asyncio
async def test_cimd_yields_to_dynamic_registration_when_the_authorization_server_offers_both(monkeypatch, respx_mock):
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
from litellm.proxy._experimental.mcp_server.oauth_utils import get_cimd_client_id
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
server = _cimd_oauth_server().model_copy(update={"registration_url": "https://idp.example.com/register"})
post = respx_mock.post("https://idp.example.com/register").respond(201, json={"client_id": "registered-client"})
response = await endpoints.register_client_with_server(_cimd_request(), server, "Gateway", None, None, None)
assert json.loads(response.body)["client_id"] == "registered-client"
assert post.call_count == 1
assert get_cimd_client_id(server) is None
@pytest.mark.asyncio
async def test_cimd_is_preferred_over_dynamic_registration_when_the_deployment_opts_in(monkeypatch, respx_mock):
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
monkeypatch.setattr(proxy_server, "general_settings", {"mcp_prefer_client_id_metadata_document": True})
server = _cimd_oauth_server().model_copy(update={"registration_url": "https://idp.example.com/register"})
post = respx_mock.post("https://idp.example.com/register").respond(201, json={"client_id": "registered-client"})
response = await endpoints.register_client_with_server(_cimd_request(), server, "Gateway", None, None, None)
body = json.loads(response.body) if hasattr(response, "body") else response
assert body["client_id"] == "https://gateway.example.com/oauth/client-metadata.json"
assert "client_secret" not in body
assert post.call_count == 0
@pytest.mark.parametrize(
"updates",
[
{"client_id": "static-client", "client_secret": "static-secret"},
{"client_id": "persisted-dcr-client", "dcr_issuer": "https://idp.example.com"},
{"client_id_metadata_document_supported": False},
{"auth_type": "oauth_delegate", "dcr_bridge": True},
{"auth_type": "true_passthrough", "dcr_bridge": True},
{"delegate_auth_to_upstream": True},
{"oauth2_flow": "client_credentials"},
{"client_secret": "configured-secret"},
{"token_endpoint_auth_method": "client_secret_basic"},
],
)
def test_cimd_preserves_existing_identity_and_other_auth_modes(monkeypatch, updates):
from litellm.proxy._experimental.mcp_server.oauth_utils import get_cimd_client_id
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
assert get_cimd_client_id(_cimd_oauth_server().model_copy(update=updates)) is None
@pytest.mark.asyncio
async def test_cimd_document_is_public_and_binds_configured_origin(monkeypatch):
import httpx
from fastapi import FastAPI
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com/proxy")
app = FastAPI()
app.include_router(endpoints.router)
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="https://attacker.example") as client:
response = await client.get("/oauth/client-metadata.json", headers={"X-Forwarded-Host": "attacker.example"})
assert response.status_code == 200
assert response.json()["client_id"] == "https://gateway.example.com/proxy/oauth/client-metadata.json"
assert response.json()["redirect_uris"] == ["https://gateway.example.com/proxy/callback"]
assert response.json()["token_endpoint_auth_method"] == "none"
assert "client_secret" not in response.json()
assert response.headers["cache-control"] == "public, max-age=300"
@pytest.mark.asyncio
@pytest.mark.parametrize("base", [None, "https://[invalid"])
async def test_cimd_document_is_unavailable_without_configured_https_origin(monkeypatch, base):
from fastapi import HTTPException
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
monkeypatch.delenv("PROXY_BASE_URL", raising=False)
if base is not None:
monkeypatch.setenv("PROXY_BASE_URL", base)
with pytest.raises(HTTPException) as exc:
await endpoints.oauth_client_metadata()
assert exc.value.status_code == 404
@pytest.mark.asyncio
async def test_cimd_upstream_client_rejection_is_gateway_fault(monkeypatch, respx_mock):
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
upstream = respx_mock.post("https://idp.example.com/token").respond(
401, json={"error": "invalid_client", "error_description": "provider-private-detail"}
)
response = await endpoints.exchange_token_with_server(
request=_cimd_request(),
mcp_server=_cimd_oauth_server(),
grant_type="authorization_code",
code="code",
redirect_uri="https://gateway.example.com/callback",
client_id="caller-placeholder",
client_secret="dummy",
code_verifier="verifier",
)
assert upstream.call_count == 1
assert response.status_code == 502
body = json.loads(response.body)
assert body["error"] == "server_error"
assert "provider-private-detail" not in body["error_description"]
@pytest.mark.asyncio
@pytest.mark.parametrize("flow", ["refresh", "authorize"])
async def test_optional_cimd_discovery_preserves_the_callers_configured_endpoint(
monkeypatch: pytest.MonkeyPatch, respx_mock: "MockRouter", flow: str
) -> None:
from urllib.parse import parse_qs
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
monkeypatch.setenv("LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP", "0")
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
monkeypatch.setenv("LITELLM_SALT_KEY", "cimd-test-state-signing-key")
token: Final = respx_mock.post("https://idp.example.com/token").respond(
200, json={"access_token": "upstream-access", "token_type": "Bearer", "expires_in": 3600}
)
discovery: Final = respx_mock.route().respond(503)
manager: Final = manager_module.MCPServerManager()
monkeypatch.setattr(manager_module, "global_mcp_server_manager", manager)
await manager.load_servers_from_config({"manual": {
"url": "https://mcp.example.com/mcp", "transport": "http", "auth_type": "oauth2",
"oauth2_flow": "authorization_code",
**({"token_url": "https://idp.example.com/token"} if flow == "refresh"
else {"authorization_url": "https://idp.example.com/authorize"}),
}})
server: Final = next(iter(manager.config_mcp_servers.values()))
async with manager.catalog.operation():
if flow == "refresh":
response: Final = await endpoints.exchange_token_with_server(
request=_cimd_request(), mcp_server=server, grant_type="refresh_token",
refresh_token="existing-refresh", client_id="existing-client",
code=None, redirect_uri=None, client_secret=None, code_verifier=None,
)
assert response.status_code == 200
assert json.loads(response.body)["access_token"] == "upstream-access"
assert parse_qs(token.calls[0].request.content.decode())["refresh_token"] == ["existing-refresh"]
else:
redirect: Final = await endpoints.authorize_with_server(
_cimd_request(), server, "existing-client", "https://gateway.example.com/ui/",
code_challenge="a" * 43, code_challenge_method="S256",
)
assert redirect.status_code == 307
assert redirect.headers["location"].startswith("https://idp.example.com/authorize?")
assert discovery.call_count > 0
@pytest.mark.asyncio
@pytest.mark.parametrize("origin", [None, "https://renamed-gateway.example.com"])
async def test_token_route_preserves_saved_cimd_grant_after_origin_change(
jwt_oauth_identity: tuple["JWTHandler", "RSAPrivateKey"],
monkeypatch: pytest.MonkeyPatch,
respx_mock: "MockRouter",
origin: str | None,
) -> None:
from types import SimpleNamespace
from urllib.parse import parse_qs
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import db, mcp_server_manager
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
monkeypatch.setenv("LITELLM_SALT_KEY", "saved-cimd-route-test-salt")
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
server: Final = _cimd_oauth_server()
mcp_server_manager.global_mcp_server_manager.registry[server.server_id] = server
identity: Final = "https://gateway.example.com/oauth/client-metadata.json"
table: Final = proxy_server.prisma_client.db.litellm_mcpusercredentials
table.find_unique = AsyncMock(return_value=None)
table.upsert = AsyncMock()
await endpoints._store_per_user_token_server_side(
server, "jwt-owner", {"access_token": "expired", "refresh_token": "saved-refresh", "expires_in": -1},
cimd_client_id=identity,
)
async def saved_row(**kwargs):
return SimpleNamespace(credential_b64=table.upsert.call_args.kwargs["data"]["create"]["credential_b64"])
table.find_unique.side_effect = saved_row
if origin is None:
monkeypatch.delenv("PROXY_BASE_URL")
else:
monkeypatch.setenv("PROXY_BASE_URL", origin)
_, key = jwt_oauth_identity
upstream: Final = respx_mock.post(server.token_url).respond(
200, json={"access_token": "fresh", "refresh_token": "rotated", "expires_in": 3600}
)
response: Final = await endpoints.exchange_token_with_server(
_token_request({"Authorization": f"Bearer {_oauth_identity_jwt(key, scope='litellm_proxy_admin')}"}),
server, "refresh_token", None, None, identity, None, None, refresh_token="saved-refresh",
)
assert response.status_code == 200
assert parse_qs(upstream.calls[0].request.content.decode())["client_id"] == [identity]
persisted: Final = await db.get_user_oauth_credential(proxy_server.prisma_client, "jwt-owner", server.server_id)
assert persisted is not None and persisted["cimd_client_id"] == identity
assert persisted["refresh_token"] == "rotated"
assert table.upsert.await_count == 2
@pytest.mark.asyncio
@pytest.mark.parametrize("state", [
"matching", "unicode_grant", "foreign_refresh", "missing_refresh", "missing_grant", "database_missing",
"database_outage", "static_client", "anonymous",
])
async def test_saved_cimd_refresh_identity_is_bound_to_the_callers_stored_grant(
monkeypatch: pytest.MonkeyPatch, state: str,
) -> None:
from types import SimpleNamespace
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import db
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
monkeypatch.setenv("LITELLM_SALT_KEY", "saved-cimd-owner-test-salt")
server: Final = _cimd_oauth_server()
identity: Final = "https://original-gateway.example.com/oauth/client-metadata.json"
grant: Final = "alice-refresh-\u00e9" if state == "unicode_grant" else "alice-refresh"
database: Final = MagicMock()
table: Final = database.db.litellm_mcpusercredentials
table.find_unique = AsyncMock(return_value=None)
table.upsert = AsyncMock()
monkeypatch.setattr(proxy_server, "prisma_client", database)
await db.store_user_oauth_credential(
database, "alice", server.server_id, "expired",
refresh_token=None if state == "missing_refresh" else grant, cimd_client_id=identity,
)
table.find_unique.reset_mock()
table.find_unique.return_value = (
None if state == "missing_grant" else SimpleNamespace(
credential_b64=table.upsert.call_args.kwargs["data"]["create"]["credential_b64"]
)
)
if state == "database_missing":
monkeypatch.setattr(proxy_server, "prisma_client", None)
if state == "database_outage":
table.find_unique.side_effect = RuntimeError("database unavailable")
if state == "static_client":
server.client_id = "configured-client"
resolved: Final = await endpoints._saved_cimd_refresh_client_id(
server, None if state == "anonymous" else "alice",
"foreign-refresh" if state == "foreign_refresh" else grant,
)
assert resolved == (identity if state in ("matching", "unicode_grant") else None)
if state in ("anonymous", "static_client", "database_missing"):
table.find_unique.assert_not_awaited()
else:
table.find_unique.assert_awaited_once_with(where={"user_id_server_id": {"user_id": "alice", "server_id": server.server_id}})

View file

@ -1461,6 +1461,7 @@ class TestMCPServerManager:
manager = MCPServerManager()
metadata = MCPOAuthMetadata(
client_id_metadata_document_supported=True,
authorization_url="https://attacker.example.com/authorize",
token_url="https://attacker.example.com/token",
scopes=["read", "admin"],
@ -1478,6 +1479,8 @@ class TestMCPServerManager:
assert server.token_url is None
assert server.scopes == ["read", "admin"]
assert server.client_id_metadata_document_supported is False
@pytest.mark.asyncio
async def test_load_servers_from_config_fills_token_url_when_metadata_corroborates_manual_authorization_url(self):
"""Corroborated metadata keeps the self-heal on the config path: when the discovered document
@ -1487,6 +1490,7 @@ class TestMCPServerManager:
manager = MCPServerManager()
metadata = MCPOAuthMetadata(
client_id_metadata_document_supported=True,
authorization_url="https://idp.example.com/authorize",
token_url="https://idp.example.com/token",
scopes=["read", "admin"],
@ -1503,6 +1507,8 @@ class TestMCPServerManager:
assert server.token_url == "https://idp.example.com/token"
assert server.scopes == ["read", "admin"]
assert server.client_id_metadata_document_supported is True
@pytest.mark.asyncio
@pytest.mark.parametrize("blank_authorization_url", ["", " "])
async def test_load_servers_from_config_blank_authorization_url_is_not_a_pin(self, blank_authorization_url):
@ -2535,6 +2541,7 @@ class TestMCPServerManager:
)
metadata = MCPOAuthMetadata(
client_id_metadata_document_supported=True,
authorization_url="https://idp.example.com/authorize",
token_url="https://idp.example.com/token",
registration_url="https://idp.example.com/register",
@ -2548,6 +2555,8 @@ class TestMCPServerManager:
assert built.registration_url == "https://idp.example.com/register"
assert built.scopes == ["read"]
assert built.client_id_metadata_document_supported is True
@pytest.mark.asyncio
async def test_build_from_table_uses_issuer_anchored_endpoints_when_issuer_configured(self):
"""When an admin configures an issuer, the build takes its endpoints from the issuer-anchored
@ -4507,6 +4516,7 @@ class TestMCPServerManager:
"authorization_endpoint": "https://idp.example.com/authorize",
"token_endpoint": "https://idp.example.com/token",
"scopes_supported": ["read", "write"],
"client_id_metadata_document_supported": True,
},
)
mock_client = MagicMock()
@ -4522,6 +4532,8 @@ class TestMCPServerManager:
assert result.token_url == "https://idp.example.com/token"
assert result.scopes == ["read", "write"]
assert result.client_id_metadata_document_supported is True
@pytest.mark.asyncio
async def test_fetch_single_authorization_server_metadata_rejects_issuer_mismatch(self):
"""RFC 8414 §3.3 fail-closed: a document self-attesting a DIFFERENT issuer than the one it was
@ -19108,3 +19120,150 @@ def test_discovery_keys_bind_static_auth_to_caller_and_configuration() -> None:
manager._discovery_key(updated, first, None, None, None, None),
)
assert len(set(keys)) == 3
@pytest.mark.parametrize("advertised", [False, True])
def test_untrusted_metadata_cannot_enable_cimd_without_matching_authorization_endpoint(advertised):
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_restrict_discovery_to_corroborated_authorization_server,
)
metadata = MCPOAuthMetadata(scopes=["read"], client_id_metadata_document_supported=advertised)
result = _restrict_discovery_to_corroborated_authorization_server(
metadata, "https://trusted.example.com/authorize", "server", False
)
assert result is not None
assert result.client_id_metadata_document_supported is False
assert result.scopes == ["read"]
@pytest.mark.asyncio
@pytest.mark.parametrize("source", ["config", "database"])
@pytest.mark.parametrize("advertised", [True, False])
@pytest.mark.parametrize("startup", [True, False])
async def test_manual_oauth_endpoints_discover_client_metadata_once(
source: str, advertised: bool, startup: bool, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch
) -> None:
from starlette.requests import Request
from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module
from litellm.proxy._experimental.mcp_server.oauth_utils import get_cimd_client_id
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
monkeypatch.setenv("LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP", "1" if startup else "0")
await _mock_oauth_discovery(respx_mock, monkeypatch, server_url="https://up.example.com/mcp", scopes=["read"])
metadata: Final = respx_mock.get("https://up.example.com/.well-known/oauth-authorization-server").respond(
json={
"issuer": "https://up.example.com",
"authorization_endpoint": "https://up.example.com/authorize",
"token_endpoint": "https://up.example.com/token",
"client_id_metadata_document_supported": advertised,
}
)
manager: Final = MCPServerManager()
monkeypatch.setattr(manager_module, "global_mcp_server_manager", manager)
configured: Final[MCPServer]
if source == "config":
await manager.load_servers_from_config({"manual": {
"url": "https://up.example.com/mcp", "transport": "http", "auth_type": "oauth2",
"oauth2_flow": "authorization_code", "authorization_url": "https://up.example.com/authorize",
"token_url": "https://up.example.com/token", "scopes": ["read"],
}})
configured = next(iter(manager.config_mcp_servers.values()))
else:
row: Final = LiteLLM_MCPServerTable(
server_id="manual", alias="manual", url="https://up.example.com/mcp", transport=MCPTransport.http,
auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code",
authorization_url="https://up.example.com/authorize", token_url="https://up.example.com/token",
credentials={"scopes": ["read"]}, created_at=datetime.now(), updated_at=datetime.now(),
)
built: Final = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
manager.registry[built.server_id] = built
configured = built
async with manager.catalog.operation():
response: Final = await discoverable_endpoints.register_client_with_server(
Request({"type": "http", "scheme": "https", "server": ("gateway.example.com", 443),
"path": "/register", "root_path": "", "headers": [], "query_string": b""}),
configured, "Gateway", None, None, None,
)
body: Final = json.loads(response.body) if hasattr(response, "body") else response
assert body["client_id"] == (
"https://gateway.example.com/oauth/client-metadata.json" if advertised else configured.server_name
)
resolved: Final = await manager.ensure_oauth_metadata_discovered(configured)
again: Final = await manager.ensure_oauth_metadata_discovered(resolved)
assert metadata.call_count == 1
assert resolved.client_id_metadata_document_supported is advertised
assert again == resolved
assert get_cimd_client_id(resolved) == (
"https://gateway.example.com/oauth/client-metadata.json" if advertised else None
)
assert resolved.effective_authorization_url == "https://up.example.com/authorize"
assert resolved.effective_token_url == "https://up.example.com/token"
@pytest.mark.asyncio
async def test_optional_client_metadata_discovery_failure_preserves_manual_endpoints(
respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
monkeypatch.setenv("LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP", "0")
await _mock_oauth_discovery(respx_mock, monkeypatch, server_url="https://up.example.com/mcp", scopes=["read"])
respx_mock.get("https://up.example.com/.well-known/oauth-authorization-server").respond(503)
respx_mock.route().respond(404)
manager: Final = MCPServerManager()
await manager.load_servers_from_config({"manual": {
"url": "https://up.example.com/mcp", "transport": "http", "auth_type": "oauth2",
"oauth2_flow": "authorization_code", "authorization_url": "https://up.example.com/authorize",
"token_url": "https://up.example.com/token", "scopes": ["read"],
}})
configured: Final = next(iter(manager.config_mcp_servers.values()))
resolved: Final = await manager.ensure_oauth_metadata_discovered(configured)
attempts: Final = len(respx_mock.calls)
retry: Final = await manager.ensure_oauth_metadata_discovered(resolved)
assert attempts > 0
assert len(respx_mock.calls) == attempts
assert retry == resolved
assert resolved.client_id_metadata_document_supported is None
assert resolved.effective_authorization_url == "https://up.example.com/authorize"
assert resolved.effective_token_url == "https://up.example.com/token"
assert manager.oauth_discovery_slot(resolved.server_id) is not None
@pytest.mark.parametrize("capability", [True, False])
@pytest.mark.parametrize("rebuild", ["same", "repointed", "fresh_discovery", "anchored"])
def test_oauth_rebuild_retains_only_corroborated_cimd_capability(capability: bool, rebuild: str) -> None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import carry_forward_resolved_oauth_endpoints
previous: Final = MCPServer(
server_id="cimd-rebuild", name="cimd_rebuild", url="https://mcp.example.com/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.oauth2,
authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token",
client_id_metadata_document_supported=capability,
)
rebuilt: Final = previous.model_copy(update={
"client_id_metadata_document_supported": not capability if rebuild == "fresh_discovery" else None,
"authorization_url": "https://changed.example.com/authorize" if rebuild == "repointed" else previous.authorization_url,
"token_url": None,
"issuer": "https://idp.example.com" if rebuild == "anchored" else None,
"issuer_is_anchored": rebuild == "anchored",
})
carry_forward_resolved_oauth_endpoints(rebuilt, previous)
expected: Final = not capability if rebuild == "fresh_discovery" else None if rebuild in ("repointed", "anchored") else capability
assert rebuilt.client_id_metadata_document_supported is expected
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["authorization_url", "token_url"])
async def test_repeated_stale_discovery_uses_current_callers_endpoint(endpoint: str) -> None:
manager: Final = MCPServerManager()
original: Final = MCPServer(
server_id="partial-replacement", name="replacement", url="https://old.example.com/mcp",
transport=MCPTransport.http, auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code",
)
replacement: Final = original.model_copy(update={endpoint: "https://new.example.com/oauth"})
manager.registry[original.server_id] = replacement
resolved: Final = await manager._rejoin_oauth_metadata_discovery(
original, needed_endpoint=lambda server: getattr(server, endpoint), retry_stale=False,
)
assert resolved is replacement

View file

@ -21,6 +21,20 @@ FLAG: Final = "LITELLM_DISABLE_LAZY_ROUTES"
WARMUP_PATH: Final = "/lazy/warm/{name}"
def test_cimd_metadata_is_available_before_any_oauth_request(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv(FLAG, "false")
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com")
app: Final = FastAPI()
attach_lazy_features(app)
with TestClient(app) as client:
response: Final = client.get("/oauth/client-metadata.json")
assert response.status_code == 200
assert response.json()["client_id"] == "https://gateway.example.com/oauth/client-metadata.json"
assert response.json()["redirect_uris"] == ["https://gateway.example.com/callback"]
class _Operation(BaseModel):
tags: tuple[str, ...]

View file

@ -30314,6 +30314,11 @@ export interface components {
* @description Custom CIDR ranges that define internal/private networks for MCP access control. When set, only these ranges are treated as internal. Defaults to RFC 1918 private ranges (10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16, 127.0.0.0/8).
*/
mcp_internal_ip_ranges?: string[] | null;
/**
* Mcp Prefer Client Id Metadata Document
* @description When true, a gateway-managed OAuth2 MCP server whose authorization server advertises Client ID Metadata Document support identifies itself with the gateway's public metadata document URL even when that authorization server also offers dynamic client registration. Requires a public HTTPS PROXY_BASE_URL the authorization server can fetch. Default false: dynamic client registration is used whenever the authorization server offers it, and the metadata document only when it does not.
*/
mcp_prefer_client_id_metadata_document?: boolean | null;
/**
* Mcp Required Fields
* @description List of MCP server fields that must be filled in for a submission to pass standards checks (e.g. ['description', 'source_url', 'alias']).