From d0202ac364ea5e88fcf63d83634f47547677d685 Mon Sep 17 00:00:00 2001 From: joshua-berri Date: Thu, 8 Oct 2026 04:10:34 -0700 Subject: [PATCH] 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> --- .github/workflows/test-unit-proxy-db.yml | 10 + .../proxy/_experimental/mcp_server/catalog.py | 1 + litellm/proxy/_experimental/mcp_server/db.py | 5 + .../mcp_server/discoverable_endpoints.py | 113 ++++- .../mcp_server/mcp_server_manager.py | 79 +++- .../_experimental/mcp_server/oauth_utils.py | 59 ++- .../authz_code_refresher.py | 28 +- .../outbound_credentials/oauth_token_store.py | 1 + .../per_user_oauth_store.py | 2 + .../outbound_credentials/v2_token_store.py | 2 + litellm/proxy/_lazy_features.py | 1 + litellm/proxy/_types.py | 4 + litellm/proxy/proxy_server.py | 1 + .../types/mcp_server/mcp_server_manager.py | 2 + .../test_authz_code_refresher.py | 82 +++- .../mcp_server/test_db_credentials.py | 63 ++- .../mcp_server/test_discoverable_endpoints.py | 417 +++++++++++++++++- .../mcp_server/test_mcp_server_manager.py | 159 +++++++ tests/unit/proxy/test__lazy_features.py | 14 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 5 + 20 files changed, 997 insertions(+), 51 deletions(-) diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml index ba88464d652..3a989dc6651 100644 --- a/.github/workflows/test-unit-proxy-db.yml +++ b/.github/workflows/test-unit-proxy-db.yml @@ -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: >- diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py index d2d6291233e..d0383be61a8 100644 --- a/litellm/proxy/_experimental/mcp_server/catalog.py +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -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",))), diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index e2d26054c80..823f0ed675c 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -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 ) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index da6ce954b1d..0253719f410 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -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"} + ) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 701a8f74813..c42ab60c48b 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index e60c4ef1c30..276cf250b30 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -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: diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py index af1e82eab82..b926ca8a6a4 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/authz_code_refresher.py @@ -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, ) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py index 0ac4296f498..b9f09c266d5 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py index 2897c3e8e4a..d8094d5facb 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py @@ -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 {}), ) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/v2_token_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/v2_token_store.py index 0f18931d118..56ab8e6bbdc 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/v2_token_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/v2_token_store.py @@ -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, ) diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 517c0916e13..68188800438 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -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", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index bc27b2352a9..f009390576a 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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.", diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index b4b23f111d1..000ffa9c305 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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", diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 2f5d983cdc0..a7027102250 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py b/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py index df068b60338..350e7f1f910 100644 --- a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py +++ b/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_authz_code_refresher.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py index de86c1b62f4..10993280e89 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 23bdfa3eadf..4f22200145e 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -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}}) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 5bd4ff5d894..5731738f137 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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 diff --git a/tests/unit/proxy/test__lazy_features.py b/tests/unit/proxy/test__lazy_features.py index c3c5068b890..c8e34ad750c 100644 --- a/tests/unit/proxy/test__lazy_features.py +++ b/tests/unit/proxy/test__lazy_features.py @@ -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, ...] diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 1108d4a8743..bf863323e77 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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']).