diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index bf1f6fa90c1..b85ff2ce50e 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -20,7 +20,9 @@ from litellm.proxy._experimental.mcp_server.oauth_identity_binding import ( credential_binding_matches, 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, +) from litellm.proxy._types import ( LiteLLM_MCPServerTable, MCPApprovalStatus, @@ -105,7 +107,31 @@ _AUTH_FLOW_SCOPED_FIELDS: Final["frozenset[str]"] = frozenset( ) -def _blank_to_none(value: str | None) -> str | None: +_OAUTH_CLIENT_CREDENTIAL_FIELDS: Final = frozenset( + { + "client_id", + "client_secret", + "token_endpoint_auth_method", + "redirect_uris", + "dcr_issuer", + "dcr_server_url", + "access_token", + "refresh_token", + "expires_in", + "scope", + } +) + + +def stale_mcp_auth_fields(submitted: Mapping[str, object], previous_value: Callable[[str], object]) -> dict[str, None]: + return { + field: None + for field in _AUTH_FLOW_SCOPED_FIELDS + if field not in submitted or submitted[field] == previous_value(field) + } + + +def _blank_to_none(value: object) -> str | None: if not isinstance(value, str): return None return value.strip() or None @@ -137,6 +163,56 @@ _CLIENT_FORWARDED_AUTH_TYPES: Final["frozenset[str]"] = frozenset({"true_passthr _MINTED_TOKEN_CREDENTIAL_FIELDS: Final["frozenset[str]"] = frozenset({"access_token", "refresh_token", "expires_in"}) +def _bind_submitted_oauth_client( + credentials: Mapping[str, object], issuer: str | None, url: str | None +) -> dict[str, object]: + if not credentials.get("client_id") or credentials.get("dcr_issuer") or credentials.get("dcr_server_url"): + return dict(credentials) + return {**credentials, "dcr_issuer": issuer, "dcr_server_url": url} + + +def is_resubmitted_oauth_client(supplied: dict[str, object], existing: dict[str, object]) -> bool: + """Recognize the saved client without conflating same-ID clients from different issuers.""" + return bool( + supplied.get("client_id") + and _decrypted_credential_field(supplied, "client_id") == _decrypted_credential_field(existing, "client_id") + and all( + _decrypted_credential_field(supplied, field) == _decrypted_credential_field(existing, field) + for field in ("client_secret", "token_endpoint_auth_method", "dcr_issuer", "dcr_server_url") + if field in supplied + ) + ) + + +def oauth_credentials_for_upstream_edit( + credentials: Mapping[str, object], + previous_issuer: str | None, + previous_url: str | None, + *, + issuer_changed: bool, +) -> dict[str, object]: + """Bind an existing client on an explicit edit, so discovery can verify reuse at the new URL. + + No first-use backfill: unchanged legacy rows never enter this path. Without a known previous + issuer, or after a known issuer change, the old client cannot be carried to the new resource. + """ + registered_issuer: Final = _blank_to_none(credentials.get("dcr_issuer")) or previous_issuer + keep_client: Final = bool(registered_issuer and credentials.get("client_id") and not issuer_changed) + removed: Final = ( + _MINTED_TOKEN_CREDENTIAL_FIELDS | {"auth_value"} + if keep_client + else _OAUTH_CLIENT_CREDENTIAL_FIELDS | {"auth_value"} + ) + retained: Final = {key: value for key, value in credentials.items() if key not in removed} + if keep_client: + return { + **retained, + "dcr_issuer": registered_issuer, + "dcr_server_url": credentials.get("dcr_server_url") or previous_url, + } + return retained + + class _OAuthCredentialAccessToken(TypedDict): access_token: str @@ -386,7 +462,11 @@ def _prepare_mcp_server_data( if blob_value is not None and te_field not in data_dict: data_dict[te_field] = blob_value data_dict["credentials"] = encrypt_credentials(credentials=credentials, encryption_key=_get_salt_key()) - data_dict["credentials"] = safe_dumps(data_dict["credentials"]) + data_dict["credentials"] = safe_dumps( + _bind_submitted_oauth_client(data_dict["credentials"], data.issuer, data.url) + if not exclude_unset and data.auth_type == "oauth2" + else data_dict["credentials"] + ) # Serialize JSON fields from ``data_dict`` (not ``data``) so the # exclude_unset filter is respected. Reading back from ``data`` would @@ -1132,6 +1212,7 @@ async def _update_mcp_server_row( *, server_id: str, data_dict: Mapping[str, object], + expected_updated_at: datetime | None = None, ) -> "prisma_db_models.LiteLLM_MCPServerTable | McpIdentifierConflict | None": identifier_write: Final = any(field in data_dict for field in ("server_name", "alias")) protocol_write: Final = bool({"transport", "mcp_info"}.intersection(data_dict)) @@ -1144,6 +1225,11 @@ async def _update_mcp_server_row( if stored is None: return None _validate_mcp_protocol_write(stored, data_dict) + if expected_updated_at is not None: + changed: Final = await table.update_many( + where={"server_id": server_id, "updated_at": expected_updated_at}, data=data_dict + ) + return await table.find_unique(where={"server_id": server_id}) if changed else None return await table.update( where={"server_id": server_id}, data=data_dict, @@ -1180,6 +1266,7 @@ async def update_mcp_server( data: UpdateMCPServerRequest, touched_by: str, fields_set: set[str] | None = None, + expected_updated_at: datetime | None = None, ) -> LiteLLM_MCPServerTable | McpIdentifierConflict | None: """ Update a new mcp server record in the db @@ -1193,16 +1280,16 @@ async def update_mcp_server( # exclude_unset=True makes this a true partial update: fields the caller did # not provide are not written, so they keep their existing DB value instead # of being reset to a schema default (transport=sse, allow_all_keys=False...). - data_dict: Final = _prepare_mcp_server_data(data, exclude_unset=True, fields_set=fields_set) + prepared_data: Final = _prepare_mcp_server_data(data, exclude_unset=True, fields_set=fields_set) # Pre-fetch existing record once if we need it for auth_type, url, or credential logic existing = None - has_credentials: Final = "credentials" in data_dict and data_dict["credentials"] is not None + has_credentials: Final = "credentials" in prepared_data and prepared_data["credentials"] is not None # An explicit token-exchange column write (set or clear) also migrates the # legacy blob copies below, so the existing row is needed for those updates. - explicit_te_write: Final = bool(_TOKEN_EXCHANGE_COLUMN_FIELDS & data_dict.keys()) - url_provided: Final = "url" in data_dict and data_dict["url"] is not None - issuer_provided: Final = "issuer" in data_dict + explicit_te_write: Final = bool(_TOKEN_EXCHANGE_COLUMN_FIELDS & prepared_data.keys()) + url_provided: Final = "url" in prepared_data and prepared_data["url"] is not None + issuer_provided: Final = "issuer" in prepared_data if data.auth_type or has_credentials or explicit_te_write or url_provided or issuer_provided: existing = await _db_find_mcp_server_row(prisma_client, data.server_id) @@ -1213,29 +1300,60 @@ async def update_mcp_server( ) # A url change re-points the server at a potentially different upstream, so any discovered or # trust-on-first-use OAuth endpoints/issuer belong to the old upstream and must re-discover. - url_changed: Final = bool(url_provided and existing and existing.url != data_dict["url"]) + url_changed: Final = bool(url_provided and existing and existing.url != prepared_data["url"]) old_issuer: Final = _blank_to_none(getattr(existing, "issuer", None)) if existing else None issuer_changed: Final = bool( - issuer_provided and old_issuer is not None and _blank_to_none(data_dict.get("issuer")) != old_issuer + issuer_provided and existing is not None and _blank_to_none(prepared_data.get("issuer")) != old_issuer ) + oauth_upstream_edited: Final = bool(existing and existing.auth_type == "oauth2" and (url_changed or issuer_changed)) + existing_credentials: Final = _credentials_blob_to_mutable_dict((existing.credentials or {}) if existing else {}) + retained_credentials: Final = ( + oauth_credentials_for_upstream_edit( + existing_credentials, + old_issuer, + existing.url if existing else None, + issuer_changed=issuer_changed or auth_type_changed, + ) + if oauth_upstream_edited + else existing_credentials + ) + edited_credentials: Final = {"credentials": safe_dumps(retained_credentials)} if oauth_upstream_edited else {} + cleared_auth_fields: Final = ( + stale_mcp_auth_fields(prepared_data, lambda field: getattr(existing, field, None)) + if auth_type_changed or url_changed or issuer_changed + else {} + ) + supplied: Final = _credentials_blob_to_mutable_dict(prepared_data.get("credentials") or {}) + safe_submitted: Final = ( + oauth_credentials_for_upstream_edit( + supplied, + old_issuer, + existing.url if existing else None, + issuer_changed=issuer_changed or auth_type_changed, + ) + if oauth_upstream_edited and is_resubmitted_oauth_client(supplied, existing_credentials) + else supplied + ) + submitted_credentials: Final = ( + { + "credentials": safe_dumps( + _bind_submitted_oauth_client( + safe_submitted, + _blank_to_none(cleared_auth_fields.get("issuer", prepared_data.get("issuer", old_issuer))), + _blank_to_none(prepared_data.get("url", existing.url if existing else None)), + ) + ) + } + if has_credentials and (data.auth_type or (existing.auth_type if existing else None)) == "oauth2" + else {} + ) + data_dict: Final = {**edited_credentials, **prepared_data, **submitted_credentials, **cleared_auth_fields} + # Clear stale credentials when auth_type changes but no new credentials provided if auth_type_changed and "credentials" not in data_dict: data_dict["credentials"] = None - if auth_type_changed or url_changed or issuer_changed: - # Clear each auth-flow-scoped field that the caller either omitted (partial update) or - # resubmitted unchanged. The edit form re-sends every field, so a stale issuer/endpoint - # belonging to the old upstream would otherwise survive a url/auth_type change and win in the - # resolution merge; only a genuinely new submitted value is kept. - data_dict.update( - { - field: None - for field in _AUTH_FLOW_SCOPED_FIELDS - if field not in data_dict or data_dict[field] == getattr(existing, field, None) - } - ) - # An explicit column write that does not touch credentials must still migrate # the row's legacy blob copies: lift values for columns the caller left # untouched, strip every copy from the blob. Without this, clearing a column @@ -1263,11 +1381,10 @@ async def update_mcp_server( # within the client-forwarded class (true_passthrough ↔ oauth_delegate) keeps # the same declared app and so must merge, not replace. if not auth_type_changed: - existing_creds = _credentials_blob_to_mutable_dict(existing.credentials) new_creds: Final = _credentials_blob_to_mutable_dict(data_dict["credentials"]) # New values override existing; existing keys not in update are preserved. A client # rotation additionally drops the previous app's stale minted token keys. - merged: Final = _drop_stale_minted_on_client_rotation({**existing_creds, **new_creds}, new_creds) + merged: Final = _drop_stale_minted_on_client_rotation({**retained_credentials, **new_creds}, new_creds) # Migrate-on-write for legacy rows: token-exchange settings the # old blob shape carried move to their dedicated columns (unless # the caller set the column this update, or the row already has @@ -1298,6 +1415,7 @@ async def update_mcp_server( prisma_client, server_id=data.server_id, data_dict=data_dict, + expected_updated_at=expected_updated_at, ) if isinstance(updated_mcp_server, McpIdentifierConflict): diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 7a0f59c3c2b..342c650d2a6 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -12,7 +12,7 @@ from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse import httpx from fastapi import APIRouter, Depends, Form, HTTPException, Request from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Response -from pydantic import BaseModel, ConfigDict, Field, SecretStr, ValidationError +from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter, ValidationError from litellm._logging import verbose_logger from litellm.caching.in_memory_cache import InMemoryCache @@ -75,6 +75,7 @@ from litellm.proxy._experimental.mcp_server.oauth_utils import ( TOKEN_NO_CACHE_HEADERS, build_upstream_oauth2_token_request, get_request_base_url, + oauth_client_registration_matches, resolve_upstream_resource, validate_trusted_redirect_uri, well_known_root_suffix, @@ -200,6 +201,8 @@ def encode_state_with_base_url( dcr_client_id: str | None = None, dcr_client_secret: str | None = None, dcr_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None, + expected_issuer: str | None = None, + authorization_response_iss_parameter_supported: bool = False, oauth_nonce: str | None = None, ) -> str: """ @@ -225,6 +228,8 @@ def encode_state_with_base_url( response granted the minted client, sealed alongside the credentials so the exchange authenticates the way the upstream expects instead of falling back to the server row's configured method + expected_issuer: Issuer identifier of the authorization server this flow is being sent to, + sealed so /callback can hold the RFC 9207 ``iss`` of the response against it Returns: An encrypted string that encodes all values @@ -241,6 +246,8 @@ def encode_state_with_base_url( "dcr_client_id": dcr_client_id, "dcr_client_secret": dcr_client_secret, "dcr_token_endpoint_auth_method": dcr_token_endpoint_auth_method, + "expected_issuer": expected_issuer, + "authorization_response_iss_parameter_supported": authorization_response_iss_parameter_supported, } state_json: Final = json.dumps(state_data, sort_keys=True) encrypted_state: Final = encrypt_value_helper(state_json) @@ -807,7 +814,16 @@ def _dcr_bridge_relays_client_registration(mcp_server: MCPServer) -> bool: returns directly to the client's redirect URI without transiting the gateway. Gateway-side redirect trust and the ``/callback`` state relay therefore only apply to the short-circuit arm, where the upstream only knows the gateway's own callback.""" - return mcp_server.is_dcr_bridge and bool(mcp_server.effective_registration_url) and not mcp_server.client_id + return bool( + mcp_server.is_dcr_bridge + and mcp_server.effective_registration_url + and not ( + mcp_server.client_id + and oauth_client_registration_matches( + mcp_server.dcr_issuer, mcp_server.dcr_server_url, mcp_server.issuer, mcp_server.url + ) + ) + ) def _require_s256_pkce( @@ -932,6 +948,10 @@ async def authorize_with_server( ): _raise_if_not_oauth2(mcp_server) resolved_server: Final = await _server_with_oauth_endpoints(mcp_server, _register_flow_needed_endpoint) + if not oauth_client_registration_matches( + resolved_server.dcr_issuer, resolved_server.dcr_server_url, resolved_server.issuer, resolved_server.url + ): + raise HTTPException(status_code=400, detail="OAuth client belongs to a different issuer; register a new client") if resolved_server.effective_authorization_url is None: raise HTTPException( status_code=400, @@ -1003,6 +1023,8 @@ async def authorize_with_server( dcr_token_endpoint_auth_method=ephemeral_dcr_client.token_endpoint_auth_method if ephemeral_dcr_client else None, + expected_issuer=resolved_server.issuer, + authorization_response_iss_parameter_supported=resolved_server.authorization_response_iss_parameter_supported, ) relay_state: Final = secrets.token_urlsafe(_OAUTH_STATE_HANDLE_BYTES) @@ -1065,6 +1087,10 @@ async def exchange_token_with_server( raise HTTPException(status_code=400, detail="Unsupported grant_type") resolved_server: Final = await _server_with_oauth_endpoints(mcp_server, _token_flow_needed_endpoint) + if not oauth_client_registration_matches( + resolved_server.dcr_issuer, resolved_server.dcr_server_url, resolved_server.issuer, resolved_server.url + ): + raise HTTPException(status_code=400, detail="OAuth client belongs to a different issuer; register a new client") token_url: Final = resolved_server.effective_token_url if token_url is None: raise HTTPException( @@ -1364,6 +1390,8 @@ class _DcrClientRegistration(BaseModel): class _PersistedDcrCredentials(BaseModel): + dcr_issuer: str | None = None + dcr_server_url: str | None = None client_id: str | None = None client_secret: str | None = None token_endpoint_auth_method: str | None = None @@ -1410,19 +1438,28 @@ def _decrypt_persisted_dcr_credential(value: str | None, key: str) -> str | None def _apply_persisted_dcr_credentials(mcp_server: MCPServer, credentials: _PersistedDcrCredentials) -> bool: + if not oauth_client_registration_matches( + credentials.dcr_issuer, credentials.dcr_server_url, mcp_server.issuer, mcp_server.url + ): + return False client_id: Final = _decrypt_persisted_dcr_credential(credentials.client_id, "client_id") if not client_id: return False + mcp_server.dcr_issuer = credentials.dcr_issuer # rebind-ok: publish through the existing bool hydration contract + mcp_server.dcr_server_url = credentials.dcr_server_url # rebind-ok: retain binding in database-free shared state mcp_server.client_id = client_id mcp_server.client_secret = _decrypt_persisted_dcr_credential(credentials.client_secret, "client_secret") mcp_server.token_endpoint_auth_method = credentials.token_endpoint_auth_method return True -async def _load_store_dcr_credentials(mcp_server: MCPServer) -> _PersistedDcrCredentials | None: +async def _load_store_dcr_credentials( + mcp_server: MCPServer, *, raise_on_error: bool = False +) -> _PersistedDcrCredentials | None: """DCR client persisted in the server-scoped OAuth-client store for a config-declared server (which has no LiteLLM_MCPServerTable row). Returns None when the store has no usable client_id - or the DB is unreachable.""" + or the DB is unreachable. Registration writes request strict reads so a failed lookup cannot + be mistaken for an absent client.""" from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # avoids circular import get_mcp_server_oauth_client_credentials, ) @@ -1434,6 +1471,8 @@ async def _load_store_dcr_credentials(mcp_server: MCPServer) -> _PersistedDcrCre prisma_client=prisma_client, server_id=mcp_server.server_id ) except Exception as exc: # noqa: BLE001 # best-effort read; DB may be unreachable + if raise_on_error: + raise verbose_logger.debug( "register_client_with_server: failed to read stored DCR client for server_id=%s: %s", mcp_server.server_id, @@ -1463,6 +1502,8 @@ async def hydrate_config_server_dcr_client(mcp_server: MCPServer) -> bool: async def _resolve_persisted_dcr_client( mcp_server: MCPServer, + *, + raise_on_error: bool = False, ) -> tuple[Optional["LiteLLM_MCPServerTable"], _PersistedDcrCredentials | None]: """Resolve a server's persisted DCR client using the same two-level rule the write path uses, so read and write always agree. First, whether the server HAS a LiteLLM_MCPServerTable row: a row is @@ -1483,6 +1524,8 @@ async def _resolve_persisted_dcr_client( prisma_client = get_prisma_client_or_throw("Database not connected. Cannot read MCP OAuth client registration.") row: Final = await get_mcp_server(prisma_client=prisma_client, server_id=mcp_server.server_id) except Exception as exc: # noqa: BLE001 # best-effort read; DB may be unreachable + if raise_on_error: + raise verbose_logger.debug( "register_client_with_server: failed to read persisted DCR client for server_id=%s: %s", mcp_server.server_id, @@ -1491,12 +1534,14 @@ async def _resolve_persisted_dcr_client( return None, None if row is not None: + if row.url != mcp_server.url or (row.issuer and mcp_server.issuer and row.issuer != mcp_server.issuer): + return row, None credentials: Final = _get_persisted_dcr_credentials(row.credentials) if credentials is not None and credentials.client_id: return row, credentials return row, None if global_mcp_server_manager.is_config_declared_server(mcp_server.server_id): - return None, await _load_store_dcr_credentials(mcp_server) + return None, await _load_store_dcr_credentials(mcp_server, raise_on_error=raise_on_error) return None, None @@ -1519,20 +1564,24 @@ async def _reuse_persisted_dcr_client_if_available( if not _apply_persisted_dcr_credentials(mcp_server, credentials): return False - if persisted_mcp_server is not None: + await _refresh_persisted_dcr_server(persisted_mcp_server) + return bool(mcp_server.client_id) + + +async def _refresh_persisted_dcr_server(persisted_server: Optional["LiteLLM_MCPServerTable"]) -> None: + if persisted_server is not None and persisted_server.approval_status != "draft": from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # avoids circular import global_mcp_server_manager, ) try: - await global_mcp_server_manager.update_server(persisted_mcp_server) + await global_mcp_server_manager.update_server(persisted_server) except Exception as exc: # noqa: BLE001 # best-effort registry refresh verbose_logger.warning( "register_client_with_server: failed to refresh persisted DCR client registration for server_id=%s: %s", - mcp_server.server_id, + persisted_server.server_id, exc, ) - return bool(mcp_server.client_id) async def _persisted_dcr_redirect_uri_is_stale(mcp_server: MCPServer, current_redirect_uri: str) -> bool: @@ -1575,7 +1624,7 @@ async def _persist_dcr_client_registration( ``refresh_token`` grant has no client identity, so an expired access token forces a full re-authorization instead of a silent refresh. Mirrors the ``encrypt_credentials`` write that ``client_credentials`` and token exchange already use. Failures are logged, - never raised: registration still returns to the caller even when persistence fails. + never raised here: the registration endpoint rejects a failed persistence result. The client-forwarded token modes (``true_passthrough`` / ``oauth_delegate``) are skipped unconditionally: the caller holds the upstream token and the gateway must hold no OAuth @@ -1604,9 +1653,6 @@ async def _persist_dcr_client_registration( ) return "failed" - if await _reuse_persisted_dcr_client_if_available(mcp_server, current_redirect_uri=current_redirect_uri): - return "reused" - token_endpoint_auth_method: Final = ( "client_secret_basic" if registration.token_endpoint_auth_method == "client_secret_basic" else None ) @@ -1615,6 +1661,8 @@ async def _persist_dcr_client_registration( "client_secret": registration.client_secret, "token_endpoint_auth_method": token_endpoint_auth_method, "redirect_uris": [current_redirect_uri], + "dcr_issuer": mcp_server.issuer, + "dcr_server_url": mcp_server.url, } from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # avoids circular import @@ -1632,6 +1680,27 @@ async def _persist_dcr_client_registration( prisma_client: Final = get_prisma_client_or_throw( "Database not connected. Cannot persist MCP OAuth client registration." ) + except HTTPException: + # This getter only raises when no database is configured. The existing single-process + # temporary-session mode keeps registrations in memory; database write errors below fail. + _apply_persisted_dcr_credentials(mcp_server, _PersistedDcrCredentials.model_validate(credentials)) + return "persisted" + + try: + stored, latest_credentials = await _resolve_persisted_dcr_client(mcp_server, raise_on_error=True) + if stored is not None and ( + stored.url != mcp_server.url + or stored.auth_type != mcp_server.auth_type + or (stored.issuer and mcp_server.issuer and stored.issuer != mcp_server.issuer) + ): + return "failed" + if ( + latest_credentials is not None + and not _redirect_uri_not_registered(latest_credentials, current_redirect_uri) + and _apply_persisted_dcr_credentials(mcp_server, latest_credentials) + ): + await _refresh_persisted_dcr_server(stored) + return "reused" updated_row: Final = await update_mcp_server( prisma_client=prisma_client, data=( @@ -1649,19 +1718,25 @@ async def _persist_dcr_client_registration( ) ), touched_by="mcp_oauth_dcr", + expected_updated_at=stored.updated_at if stored else None, ) if updated_row is not None and not isinstance(updated_row, McpIdentifierConflict): - await global_mcp_server_manager.update_server(updated_row) + await _refresh_persisted_dcr_server(updated_row) + _apply_persisted_dcr_credentials(mcp_server, _PersistedDcrCredentials.model_validate(credentials)) return "persisted" + if stored is not None: + return ( + "reused" + if await _reuse_persisted_dcr_client_if_available(mcp_server, current_redirect_uri) + else "failed" + ) if global_mcp_server_manager.is_config_declared_server(mcp_server.server_id): await upsert_mcp_server_oauth_client_credentials( prisma_client=prisma_client, server_id=mcp_server.server_id, credentials=credentials, ) - mcp_server.client_id = registration.client_id - mcp_server.client_secret = registration.client_secret - mcp_server.token_endpoint_auth_method = token_endpoint_auth_method + _apply_persisted_dcr_credentials(mcp_server, _PersistedDcrCredentials.model_validate(credentials)) return "persisted" except Exception as exc: # noqa: BLE001 verbose_logger.warning( @@ -1859,20 +1934,27 @@ async def register_client_with_server( "redirect_uris": client_facing_redirect_uris, } - if mcp_server.client_id and not ( - persist_credentials - and mcp_server.registration_url - and await _persisted_dcr_redirect_uri_is_stale(mcp_server, current_redirect_uri) + resolved_server: Final = await _server_with_oauth_endpoints(mcp_server, _register_flow_needed_endpoint) + client_matches: Final = oauth_client_registration_matches( + resolved_server.dcr_issuer, resolved_server.dcr_server_url, resolved_server.issuer, resolved_server.url + ) + if ( + resolved_server.client_id + and client_matches + and not ( + persist_credentials + and resolved_server.registration_url + and await _persisted_dcr_redirect_uri_is_stale(resolved_server, current_redirect_uri) + ) ): return dummy_return if await _reuse_persisted_dcr_client_if_available( - mcp_server, + resolved_server, current_redirect_uri=current_redirect_uri if persist_credentials else None, ): return dummy_return - resolved_server: Final = await _server_with_oauth_endpoints(mcp_server, _register_flow_needed_endpoint) if resolved_server.effective_authorization_url is None: raise HTTPException( status_code=400, @@ -1908,19 +1990,32 @@ async def register_client_with_server( server_id=resolved_server.server_id, ) - token_response = response.json() - - if persist_credentials and not bridge_relay: - persistence_result = await _persist_dcr_client_registration( - resolved_server, token_response, current_redirect_uri - ) - if persistence_result == "reused": - return dummy_return - - if client_redirect_uris and not bridge_relay and isinstance(token_response, dict): - token_response = {**token_response, "redirect_uris": client_facing_redirect_uris} - - return JSONResponse(token_response) + token_response: Final = response.json() + persistence_result: Final = ( + await _persist_dcr_client_registration(resolved_server, token_response, current_redirect_uri) + if persist_credentials and not bridge_relay + else None + ) + if persistence_result == "reused": + return dummy_return + if persistence_result == "failed": + raise HTTPException(status_code=503, detail="OAuth client registration could not be saved; retry authorization") + bound_response: Final = ( + { + **token_response, + "dcr_issuer": resolved_server.issuer, + "dcr_server_url": resolved_server.url, + "dcr_redirect_uris": [current_redirect_uri], + } + if persistence_result == "persisted" + else token_response + ) + client_response: Final = ( + {**bound_response, "redirect_uris": client_facing_redirect_uris} + if client_redirect_uris and not bridge_relay and isinstance(bound_response, dict) + else bound_response + ) + return JSONResponse(client_response) @router.get("/authorize/mcp-session") @@ -2229,11 +2324,19 @@ def _render_oauth_error_html(error: str, description: str | None) -> HTMLRespons return HTMLResponse(body, status_code=400) +def _authorization_response_issuer_is_trusted(response_issuer: str | None, state_data: Mapping[str, object]) -> bool: + expected_issuer: Final = state_data.get("expected_issuer") + if response_issuer is None: + return state_data.get("authorization_response_iss_parameter_supported") is not True + return isinstance(expected_issuer, str) and bool(expected_issuer) and response_issuer == expected_issuer + + @router.get("/callback") async def callback( request: Request, code: str | None = None, state: str | None = None, + iss: str | None = None, error: str | None = None, error_description: str | None = None, error_uri: str | None = None, @@ -2244,7 +2347,9 @@ async def callback( - A successful authorization response (``code`` + ``state``), which is forwarded back to the validated client ``redirect_uri`` with the - original (un-wrapped) ``state``. + original (un-wrapped) ``state``, once the RFC 9207 ``iss`` (when the + authorization server sent one) matches the issuer /authorize sealed + into the state. - An error response (``error``[+``error_description``/``error_uri``]), per RFC 6749 §4.1.2.1. When ``state`` is present and decodes to a trusted ``redirect_uri``, the error params are propagated back to the client so @@ -2262,6 +2367,13 @@ async def callback( encoded_state = _resolve_encoded_oauth_state(request, state) try: state_data = decode_state_hash(encoded_state) + error_issuer_state: Final = TypeAdapter(dict[str, object]).validate_python(state_data) + if not _authorization_response_issuer_is_trusted(iss, error_issuer_state): + rejected_error: Final = _render_oauth_error_html( + "invalid_issuer", "Unexpected authorization issuer" + ) + _clear_oauth_state_cookie(rejected_error, request, state) + return rejected_error original_state = state_data.get("original_state") redirect_uri = _get_validated_client_redirect_uri(request, state_data) except Exception: @@ -2311,6 +2423,22 @@ async def callback( # states while permitting same-origin / allowlisted clients. redirect_uri = _get_validated_client_redirect_uri(request, state_data) + issuer_state: Final = TypeAdapter(dict[str, object]).validate_python(state_data) + if not _authorization_response_issuer_is_trusted(iss, issuer_state): + verbose_logger.warning( + "MCP /callback rejected an authorization response: RFC 9207 iss=%r does not match the " + "issuer this flow was sent to (%r)", + iss, + issuer_state.get("expected_issuer"), + ) + issuer_error_response: Final = _render_oauth_error_html( + "invalid_issuer", + "This authorization response came from a different identity provider than the one this " + "MCP server is configured to use.", + ) + _clear_oauth_state_cookie(issuer_error_response, request, state) + return issuer_error_response + # Interactive dcr_bridge oauth_delegate: the state carries the litellm user the authorize step # captured. Instead of forwarding the raw upstream code (which the client would present at the # token endpoint with no way to prove who signed in), seal the user and the upstream code into a diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index b8099c97861..abfb9dc56fe 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -719,13 +719,12 @@ def _normalized_authorize_endpoint(url: str) -> str: def _issuer_matches(claimed_issuer: object, configured_issuer: str) -> bool: """RFC 8414 §3.3 issuer equality between the metadata document's self-attested ``issuer`` and the - admin-configured issuer, tolerant only of URL-insignificant differences (scheme/host case, the - default port, a trailing slash). A non-string or empty claimed issuer never matches, so a + admin-configured issuer. A non-string or empty claimed issuer never matches, so a document that omits ``issuer`` fails closed under issuer-anchored discovery. """ if not isinstance(claimed_issuer, str) or not claimed_issuer: return False - return _normalized_authorize_endpoint(claimed_issuer) == _normalized_authorize_endpoint(configured_issuer) + return claimed_issuer == configured_issuer def _flow_endpoints_missing( @@ -856,6 +855,11 @@ def _carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_serv may_carry: Final = _endpoints_corroborate_authorization_url( previous_server.authorization_url, new_server.authorization_url ) + 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 + previous_server.authorization_response_iss_parameter_supported + ) if new_server.authorization_url is None and previous_server.authorization_url: new_server.authorization_url = previous_server.authorization_url if may_carry and new_server.token_url is None and previous_server.token_url: @@ -2110,13 +2114,20 @@ class MCPServerManager: if metadata is None: return server discovered_issuer: Final = metadata.discovered_issuer if not metadata.from_origin_fallback else None - resolved: Final = server.model_copy() - resolved.scopes = server.scopes or metadata.scopes - resolved.issuer = server.issuer or discovered_issuer - resolved.authorization_url = server.authorization_url or metadata.authorization_url - resolved.token_url = server.token_url or metadata.token_url - resolved.registration_url = server.registration_url or metadata.registration_url - return resolved + return server.model_copy( + update={ + "scopes": server.scopes or metadata.scopes, + "issuer": server.issuer or discovered_issuer, + "authorization_response_iss_parameter_supported": ( + metadata.authorization_response_iss_parameter_supported + if discovered_issuer is not None + else server.authorization_response_iss_parameter_supported + ), + "authorization_url": server.authorization_url or metadata.authorization_url, + "token_url": server.token_url or metadata.token_url, + "registration_url": server.registration_url or metadata.registration_url, + } + ) def _oauth_discovery_slot_is_current(self, server_id: str, generation: int) -> bool: slot: Final = self._oauth_discovery_slot(server_id) @@ -2641,6 +2652,11 @@ class MCPServerManager: scopes=resolved_scopes, configured_scopes=tuple(configured_scopes) if configured_scopes else None, issuer=effective_issuer, + authorization_response_iss_parameter_supported=( + gated_oauth_metadata.authorization_response_iss_parameter_supported + if gated_oauth_metadata + else False + ), issuer_is_anchored=use_issuer_anchor, authorization_url=resolved_authorization_url, token_url=resolved_token_url, @@ -3202,12 +3218,17 @@ class MCPServerManager: extra_headers=getattr(mcp_server, "extra_headers", None), static_headers=static_headers_dict, env_vars=env_vars_list, + dcr_issuer=credentials_dict.get("dcr_issuer") if credentials_dict else None, + dcr_server_url=credentials_dict.get("dcr_server_url") if credentials_dict else None, client_id=client_id_value or getattr(mcp_server, "client_id", None), client_secret=client_secret_value or getattr(mcp_server, "client_secret", None), oauth2_flow=self._explicit_oauth2_flow(getattr(mcp_server, "oauth2_flow", None)), scopes=resolved_scopes, configured_scopes=configured_scopes, issuer=effective_issuer, + authorization_response_iss_parameter_supported=( + gated_oauth_metadata.authorization_response_iss_parameter_supported if gated_oauth_metadata else False + ), issuer_is_anchored=use_issuer_anchor, authorization_url=manual_authorization_url or getattr(gated_oauth_metadata, "authorization_url", None), token_url=manual_token_url or getattr(gated_oauth_metadata, "token_url", None), @@ -5174,6 +5195,10 @@ class MCPServerManager: token_url=data.get("token_endpoint"), registration_url=data.get("registration_endpoint"), 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" + ) + is True, ) if any( diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index 39865a35ec6..3b5eda6cff2 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -658,6 +658,17 @@ def canonicalize_url_identity(url: str) -> str: return urlunparse((scheme, netloc, parsed.path.rstrip("/"), "", "", "")) +def oauth_client_registration_matches( + registered_issuer: str | None, + registered_url: str | None, + current_issuer: str | None, + current_url: str | None, +) -> bool: + if registered_issuer and current_issuer: + return registered_issuer == current_issuer + return not registered_url or registered_url == current_url + + def canonical_resource_uri(url: str) -> str | None: """Canonicalize an upstream MCP server URL into an RFC 8707 resource identifier. diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 8530a4d31b4..03408043f8e 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -31307,6 +31307,28 @@ ], "title": "Client Secret" }, + "dcr_issuer": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Dcr Issuer" + }, + "dcr_server_url": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Dcr Server Url" + }, "id_jag_resource": { "anyOf": [ { @@ -34023,6 +34045,28 @@ ], "title": "Client Secret" }, + "dcr_issuer": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Dcr Issuer" + }, + "dcr_server_url": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Dcr Server Url" + }, "id_jag_resource": { "anyOf": [ { @@ -35696,7 +35740,7 @@ }, "/callback": { "get": { - "description": "OAuth 2.0 authorization response handler for MCP loopback clients.\n\nAccepts either:\n\n- A successful authorization response (``code`` + ``state``), which is\n forwarded back to the validated client ``redirect_uri`` with the\n original (un-wrapped) ``state``.\n- An error response (``error``[+``error_description``/``error_uri``]), per\n RFC 6749 \u00a74.1.2.1. When ``state`` is present and decodes to a trusted\n ``redirect_uri``, the error params are propagated back to the client so\n its OAuth library can surface them. Otherwise we render an HTML error\n page so the user is not left on an opaque 422 / blank screen.", + "description": "OAuth 2.0 authorization response handler for MCP loopback clients.\n\nAccepts either:\n\n- A successful authorization response (``code`` + ``state``), which is\n forwarded back to the validated client ``redirect_uri`` with the\n original (un-wrapped) ``state``, once the RFC 9207 ``iss`` (when the\n authorization server sent one) matches the issuer /authorize sealed\n into the state.\n- An error response (``error``[+``error_description``/``error_uri``]), per\n RFC 6749 \u00a74.1.2.1. When ``state`` is present and decodes to a trusted\n ``redirect_uri``, the error params are propagated back to the client so\n its OAuth library can surface them. Otherwise we render an HTML error\n page so the user is not left on an opaque 422 / blank screen.", "operationId": "callback_callback_get", "parameters": [ { @@ -35731,6 +35775,22 @@ "title": "State" } }, + { + "in": "query", + "name": "iss", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Iss" + } + }, { "in": "query", "name": "error", @@ -37486,6 +37546,28 @@ ], "title": "Client Secret" }, + "dcr_issuer": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Dcr Issuer" + }, + "dcr_server_url": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Dcr Server Url" + }, "id_jag_resource": { "anyOf": [ { @@ -41480,6 +41562,28 @@ ], "title": "Client Secret" }, + "dcr_issuer": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Dcr Issuer" + }, + "dcr_server_url": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Dcr Server Url" + }, "id_jag_resource": { "anyOf": [ { diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index e1504172d9e..2aed0e2afb8 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -157,13 +157,16 @@ if MCP_AVAILABLE: get_user_env_vars, get_user_env_vars_bulk, get_user_oauth_credential, + is_resubmitted_oauth_client, list_server_user_credentials, list_user_oauth_credentials, mcp_oauth_token_identity, merge_user_env_vars, + oauth_credentials_for_upstream_edit, purge_user_oauth_credentials_for_server, reject_mcp_server, set_mcp_server_pinned_tools, + stale_mcp_auth_fields, store_user_credential, store_user_oauth_credential, update_mcp_server, @@ -870,6 +873,9 @@ if MCP_AVAILABLE: ("authentication_token", "auth_value"), ("client_id", "client_id"), ("client_secret", "client_secret"), + ("token_endpoint_auth_method", "token_endpoint_auth_method"), + ("dcr_issuer", "dcr_issuer"), + ("dcr_server_url", "dcr_server_url"), ("scopes", "scopes"), ("aws_access_key_id", "aws_access_key_id"), ("aws_secret_access_key", "aws_secret_access_key"), @@ -890,30 +896,73 @@ if MCP_AVAILABLE: value for key, value in as_dict.items() if key not in MCP_ADMIN_CONFIG_CREDENTIAL_KEYS and key != "scopes" ) + def _oauth_session_changes_upstream( + payload: NewMCPServerRequest, existing: MCPServer | LiteLLM_MCPServerTable + ) -> bool: + return existing.auth_type == MCPAuth.oauth2 and ( + payload.url != existing.url + or ("issuer" in payload.model_fields_set and (payload.issuer or None) != existing.issuer) + or (payload.auth_type is not None and payload.auth_type != existing.auth_type) + ) + def _inherit_credentials_from_existing_server( payload: NewMCPServerRequest, ) -> NewMCPServerRequest: - if not payload.server_id or _has_non_admin_config_credentials(payload.credentials): + if not payload.server_id: return payload existing_server: Final = global_mcp_server_manager.get_mcp_server_by_id(payload.server_id) if existing_server is None: return payload - - inherited_credentials: dict[str, object] = { + upstream_changed: Final = _oauth_session_changes_upstream(payload, existing_server) + issuer_changed: Final = ( + "issuer" in payload.model_fields_set and (payload.issuer or None) != existing_server.issuer + ) + cleared: Final = ( + stale_mcp_auth_fields( + payload.model_dump(exclude_unset=True), + lambda field: ( + getattr(existing_server, f"configured_{field}", None) or getattr(existing_server, field, None) + ), + ) + if upstream_changed + else {} + ) + resolved_payload: Final = ( + payload.model_copy(update={**cleared, "oauth2_flow": payload.oauth2_flow}) if upstream_changed else payload + ) + supplied: Final = dict(resolved_payload.credentials or {}) + existing_credentials: Final[dict[str, object]] = { credential_key: value for server_attr, credential_key in _INHERITED_CREDENTIAL_FIELDS if (value := getattr(existing_server, server_attr, None)) } + resubmitted_client: Final = is_resubmitted_oauth_client(supplied, existing_credentials) + if _has_non_admin_config_credentials(resolved_payload.credentials) and not ( + upstream_changed and resubmitted_client + ): + return resolved_payload + + bound_credentials: Final = ( + oauth_credentials_for_upstream_edit( + {**existing_credentials, **supplied}, + existing_server.issuer, + existing_server.url, + issuer_changed=issuer_changed + or (resolved_payload.auth_type or existing_server.auth_type) != existing_server.auth_type, + ) + if upstream_changed + else existing_credentials + ) # The gate above guarantees anything still supplied is admin config, which the admin just # typed, so it wins over the stored value. - inherited_credentials = {**inherited_credentials, **dict(payload.credentials or {})} + inherited_credentials: Final = bound_credentials if upstream_changed else {**bound_credentials, **supplied} if not inherited_credentials: - return payload + return resolved_payload.model_copy(update={"credentials": {}}) if upstream_changed else resolved_payload try: - return payload.model_copy(update={"credentials": inherited_credentials}) + return resolved_payload.model_copy(update={"credentials": inherited_credentials}) except AttributeError: pass @@ -937,15 +986,20 @@ if MCP_AVAILABLE: supplied: Final = payload.server_id if not supplied: return str(uuid.uuid4()) - if global_mcp_server_manager.get_mcp_server_by_id(supplied) is not None: - return supplied + registered: Final = global_mcp_server_manager.get_mcp_server_by_id(supplied) + if registered is not None: + return str(uuid.uuid4()) if _oauth_session_changes_upstream(payload, registered) else supplied prisma_client: Final = _get_prisma_client_or_none() if prisma_client is None: return supplied # A draft is another session's row, not a saved server, so re-supplying an id this # endpoint previously handed back must not let a later session adopt its configuration. existing: Final = await get_mcp_server(prisma_client, supplied) - if existing is None or existing.approval_status == MCPApprovalStatus.draft: + if ( + existing is None + or existing.approval_status == MCPApprovalStatus.draft + or _oauth_session_changes_upstream(payload, existing) + ): return str(uuid.uuid4()) return supplied diff --git a/litellm/types/mcp.py b/litellm/types/mcp.py index ce86e79a6f0..949506898b1 100644 --- a/litellm/types/mcp.py +++ b/litellm/types/mcp.py @@ -9,7 +9,7 @@ from urllib.parse import urlsplit import httpx from pydantic import BaseModel, ConfigDict, Field -from typing_extensions import TypedDict +from typing_extensions import ReadOnly, TypedDict from litellm.types.llms.base import HiddenParams @@ -161,6 +161,9 @@ MCPTokenEndpointAuthMethod = Literal["client_secret_basic", "client_secret_post" class MCPCredentials(TypedDict, total=False): + dcr_issuer: ReadOnly[str | None] + dcr_server_url: ReadOnly[str | None] + auth_value: str | None """ Authentication value diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 2ee19b3e59a..343fb731b06 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -38,6 +38,7 @@ class MCPOAuthMetadata(BaseModel): authorization_url: str | None = None token_url: str | None = None registration_url: str | 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 trust-on-first-use as the server's ``issuer`` when none is configured, so that later rebuilds @@ -118,6 +119,9 @@ class MCPServer(BaseModel): client_secret: str | None = None issuer: str | None = None issuer_is_anchored: bool = False + authorization_response_iss_parameter_supported: bool = False + dcr_issuer: str | None = None + dcr_server_url: str | None = None scopes: list[str] | None = None authorization_url: str | None = None token_url: str | None = None diff --git a/tests/e2e/mcp/test_mcp_oauth_happy_path_e2e.py b/tests/e2e/mcp/test_mcp_oauth_happy_path_e2e.py index c20b73c0d63..81470c21d51 100644 --- a/tests/e2e/mcp/test_mcp_oauth_happy_path_e2e.py +++ b/tests/e2e/mcp/test_mcp_oauth_happy_path_e2e.py @@ -38,6 +38,7 @@ pytestmark = [pytest.mark.e2e, pytest.mark.mcp_oauth_live, pytest.mark.provider_ class OAuthMetadata(BaseModel): + issuer: str authorization_endpoint: str token_endpoint: str registration_endpoint: str @@ -135,6 +136,7 @@ class TestMcpOauthHappyPath: auth_type="oauth2", oauth2_flow="authorization_code", per_server_oauth_discovery=route == "explicit_header_jwt", + issuer=metadata.issuer if metadata else None, authorization_url=metadata.authorization_endpoint if metadata else None, token_url=metadata.token_endpoint if metadata else None, registration_url=metadata.registration_endpoint if metadata else None, diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 90fe04ae896..ccb5cd35280 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -658,6 +658,7 @@ class McpServerCreateBody(BaseModel): auth_type: str | None = None oauth2_flow: Literal["client_credentials", "authorization_code"] | None = None per_server_oauth_discovery: bool | None = None + issuer: str | None = None authorization_url: str | None = None token_url: str | None = None registration_url: str | None = None 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 4e27ec134d4..6d71256ff18 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -11,6 +11,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import HTTPException +from litellm.proxy._types import LiteLLM_MCPServerTable from litellm.types.mcp import MCPAuth if TYPE_CHECKING: @@ -1365,6 +1366,8 @@ async def test_register_client_persists_dcr_client_identity(): mock_async_client = MagicMock() mock_async_client.post = AsyncMock(return_value=mock_response) + prisma = MagicMock() + prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=None) mock_update = AsyncMock(return_value=MagicMock()) mock_update_server = AsyncMock() @@ -1373,7 +1376,7 @@ async def test_register_client_persists_dcr_client_identity(): "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", return_value=mock_async_client, ), - patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()), + patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=prisma), patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=mock_update), patch.object(global_mcp_server_manager, "update_server", new=mock_update_server), ): @@ -1390,7 +1393,12 @@ async def test_register_client_persists_dcr_client_identity(): import json assert response.status_code == 200 - assert json.loads(response.body.decode("utf-8")) == mock_response.json.return_value + assert json.loads(response.body.decode("utf-8")) == { + **mock_response.json.return_value, + "dcr_issuer": None, + "dcr_server_url": None, + "dcr_redirect_uris": ["https://proxy.litellm.example/callback"], + } mock_update.assert_called_once() update_data = mock_update.call_args.kwargs["data"] @@ -1452,6 +1460,8 @@ async def _register_persistence_attempted_for_auth_type(auth_type: MCPAuth) -> b mock_async_client = MagicMock() mock_async_client.post = AsyncMock(return_value=mock_response) + prisma = MagicMock() + prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=None) mock_update = AsyncMock(return_value=MagicMock()) with ( @@ -1459,7 +1469,7 @@ async def _register_persistence_attempted_for_auth_type(auth_type: MCPAuth) -> b "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", return_value=mock_async_client, ), - patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()), + patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=prisma), patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=mock_update), patch.object(global_mcp_server_manager, "update_server", new=AsyncMock()), ): @@ -1473,7 +1483,12 @@ async def _register_persistence_attempted_for_auth_type(auth_type: MCPAuth) -> b persist_credentials=True, ) - assert json.loads(response.body.decode("utf-8")) == mock_response.json.return_value + expected_binding = { + "dcr_issuer": None, + "dcr_server_url": None, + "dcr_redirect_uris": ["https://proxy.litellm.example/callback"], + } if auth_type == MCPAuth.oauth2 else {} + assert json.loads(response.body.decode("utf-8")) == {**mock_response.json.return_value, **expected_binding} return mock_update.await_count > 0 @@ -1529,7 +1544,7 @@ async def test_register_client_persists_only_to_its_own_row_when_another_server_ ) sibling_row_with_client = MagicMock(server_id="server-a", url=shared_url) sibling_row_with_client.credentials = {"client_id": "client-a-do-not-adopt"} - own_row_without_client = MagicMock(server_id="server-b", url=shared_url) + own_row_without_client = LiteLLM_MCPServerTable.model_validate(fresh_server.model_dump(exclude_none=True)) own_row_without_client.credentials = {} rows_by_server_id = {"server-a": sibling_row_with_client, "server-b": own_row_without_client} @@ -1622,6 +1637,8 @@ async def test_register_client_does_not_clobber_token_url_when_absent(): mock_async_client = MagicMock() mock_async_client.post = AsyncMock(return_value=mock_response) + prisma = MagicMock() + prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=None) mock_update = AsyncMock(return_value=MagicMock()) mock_update_server = AsyncMock() @@ -1630,7 +1647,7 @@ async def test_register_client_does_not_clobber_token_url_when_absent(): "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", return_value=mock_async_client, ), - patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()), + patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=prisma), patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=mock_update), patch.object(global_mcp_server_manager, "update_server", new=mock_update_server), ): @@ -1685,7 +1702,7 @@ async def test_register_client_reuses_persisted_client_id_for_non_admin_when_reg mock_request.base_url = "https://proxy.litellm.example/" mock_request.headers = {} - persisted_server = MagicMock() + persisted_server = LiteLLM_MCPServerTable.model_validate(oauth2_server.model_dump(exclude_none=True)) persisted_server.credentials = {"client_id": "persisted-client"} mock_get_mcp_server = AsyncMock(return_value=persisted_server) mock_update_mcp_server = AsyncMock() @@ -1761,7 +1778,7 @@ async def test_register_client_reuse_refreshes_request_server_when_manager_updat mock_request.base_url = "https://proxy.litellm.example/" mock_request.headers = {} - persisted_server = MagicMock() + persisted_server = LiteLLM_MCPServerTable.model_validate(oauth2_server.model_dump(exclude_none=True)) persisted_server.credentials = { "client_id": "persisted-client", "client_secret": "persisted-secret", @@ -1803,7 +1820,8 @@ async def test_register_client_reuse_refreshes_request_server_when_manager_updat @pytest.mark.asyncio -async def test_register_client_returns_reused_client_when_concurrent_persist_wins(): +@pytest.mark.parametrize("late_winner", [False, True]) +async def test_register_client_returns_reused_client_when_concurrent_persist_wins(late_winner): try: from fastapi import Request @@ -1843,10 +1861,11 @@ async def test_register_client_returns_reused_client_when_concurrent_persist_win mock_async_client = MagicMock() mock_async_client.post = AsyncMock(return_value=mock_response) - persisted_server = MagicMock() + persisted_server = LiteLLM_MCPServerTable.model_validate(oauth2_server.model_dump(exclude_none=True)) persisted_server.credentials = {"client_id": "persisted-client"} - mock_get_mcp_server = AsyncMock(side_effect=[None, persisted_server]) - mock_update_mcp_server = AsyncMock() + empty_server = persisted_server.model_copy(update={"credentials": None}) + mock_get_mcp_server = AsyncMock(side_effect=[empty_server, empty_server, persisted_server] if late_winner else [empty_server, persisted_server]) + mock_update_mcp_server = AsyncMock(return_value=None) mock_update_server = AsyncMock() with ( @@ -1878,7 +1897,10 @@ async def test_register_client_returns_reused_client_when_concurrent_persist_win assert response["client_id"] == "remote_server" assert oauth2_server.client_id == "persisted-client" mock_async_client.post.assert_called_once() - mock_update_mcp_server.assert_not_called() + if late_winner: + mock_update_mcp_server.assert_awaited_once() + else: + mock_update_mcp_server.assert_not_called() mock_update_server.assert_called_once_with(persisted_server) @@ -1937,7 +1959,7 @@ async def test_register_client_re_registers_when_persisted_redirect_uri_no_longe mock_async_client = MagicMock() mock_async_client.post = AsyncMock(return_value=mock_response) - persisted_server = MagicMock() + persisted_server = LiteLLM_MCPServerTable.model_validate(oauth2_server.model_dump(exclude_none=True)) persisted_server.credentials = { "client_id": "stale-client", "client_secret": "stale-secret", @@ -2014,7 +2036,7 @@ async def test_register_client_grandfathers_persisted_client_without_recorded_re mock_async_client = MagicMock() mock_async_client.post = AsyncMock() - persisted_server = MagicMock() + persisted_server = LiteLLM_MCPServerTable.model_validate(oauth2_server.model_dump(exclude_none=True)) persisted_server.credentials = {"client_id": "legacy-client"} mock_get_mcp_server = AsyncMock(return_value=persisted_server) mock_update_mcp_server = AsyncMock() @@ -2072,7 +2094,7 @@ async def test_register_client_keeps_persisted_client_when_recorded_redirect_uri mock_async_client = MagicMock() mock_async_client.post = AsyncMock() - persisted_server = MagicMock() + persisted_server = LiteLLM_MCPServerTable.model_validate(oauth2_server.model_dump(exclude_none=True)) persisted_server.credentials = { "client_id": "kept-client", "redirect_uris": ["https://proxy.litellm.example/callback"], @@ -2137,7 +2159,7 @@ async def test_register_client_non_admin_reuses_persisted_client_despite_redirec mock_async_client = MagicMock() mock_async_client.post = AsyncMock() - persisted_server = MagicMock() + persisted_server = LiteLLM_MCPServerTable.model_validate(oauth2_server.model_dump(exclude_none=True)) persisted_server.credentials = { "client_id": "persisted-client", "redirect_uris": ["https://old-origin.example/callback"], @@ -4465,6 +4487,173 @@ async def test_callback_error_path_reads_cookie_and_clears_it(monkeypatch): assert cleared[cookie_name]["max-age"] == "0" +def _issuer_anchored_oauth_server(issuer: str | None = "https://idp.example.com"): + from litellm.proxy._types import MCPTransport + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + return MCPServer( + server_id="rfc9207_server", + name="rfc9207", + server_name="rfc9207", + alias="rfc9207", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="upstream-client-id", + issuer=issuer, + authorization_url="https://idp.example.com/oauth/authorize", + token_url="https://idp.example.com/oauth/token", + ) + + +async def _authorize_then_callback(server, iss, monkeypatch, expected_issuer_override=..., error=None): + """Run /authorize for ``server``, then feed the resulting flow back through /callback with the + RFC 9207 ``iss`` the authorization server supposedly returned. Returns (callback_response, + sealed_state_data).""" + from http.cookies import SimpleCookie + from urllib.parse import parse_qs, urlparse + + from fastapi import Request + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _oauth_state_cookie_name, + authorize_with_server, + callback, + decode_state_hash, + encode_state_with_base_url, + ) + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-for-LIT-5940") + client_redirect_uri = "http://127.0.0.1:6274/oauth/callback/debug" + + authorize_request = MagicMock(spec=Request) + authorize_request.base_url = "https://proxy.example.com/" + authorize_request.headers = {} + + authorize_response = await authorize_with_server( + request=authorize_request, + mcp_server=server, + client_id="upstream-client-id", + redirect_uri=client_redirect_uri, + state="client-original-state-9207", + code_challenge="challenge", + code_challenge_method="S256", + ) + handle = parse_qs(urlparse(authorize_response.headers["location"]).query)["state"][0] + jar = SimpleCookie() + jar.load(authorize_response.headers["set-cookie"]) + cookie_name = _oauth_state_cookie_name(handle) + sealed_state = jar[cookie_name].value + if expected_issuer_override is not ...: + # A state minted before the issuer was sealed into it: same shape, key absent. + sealed_state = encode_state_with_base_url( + base_url=client_redirect_uri, + original_state="client-original-state-9207", + client_redirect_uri=client_redirect_uri, + expected_issuer=expected_issuer_override, + ) + + callback_request = MagicMock(spec=Request) + callback_request.base_url = "https://proxy.example.com/" + callback_request.headers = {} + callback_request.cookies = {cookie_name: sealed_state} + + response = await callback( + request=callback_request, + code="upstream-auth-code" if error is None else None, + error=error, + state=handle, + iss=iss, + ) + return response, decode_state_hash(sealed_state) + + +@pytest.mark.asyncio +async def test_authorize_seals_the_issuer_and_callback_accepts_a_matching_rfc9207_iss(monkeypatch): + """The callback accepts the exact issuer sealed during authorization.""" + from urllib.parse import parse_qs, urlparse + + response, state_data = await _authorize_then_callback( + _issuer_anchored_oauth_server(), + iss="https://idp.example.com", + monkeypatch=monkeypatch, + ) + + assert state_data["expected_issuer"] == "https://idp.example.com" + assert response.status_code == 302 + query = parse_qs(urlparse(response.headers["location"]).query) + assert query["code"] == ["upstream-auth-code"] + assert query["state"] == ["client-original-state-9207"] + + +@pytest.mark.asyncio +async def test_callback_rejects_authorization_response_from_a_different_issuer(monkeypatch): + """LIT-5940 / RFC 9207 §2.4: an ``iss`` naming an authorization server we never sent the user to + is a mix-up attack, so the code must not reach the client's redirect_uri.""" + response, _ = await _authorize_then_callback( + _issuer_anchored_oauth_server(), + iss="https://attacker-idp.example.com", + monkeypatch=monkeypatch, + ) + + assert response.status_code == 400 + assert "location" not in response.headers + assert b"upstream-auth-code" not in response.body + assert b"invalid_issuer" in response.body + + +@pytest.mark.asyncio +@pytest.mark.parametrize("error", [None, "access_denied"]) +async def test_callback_rejects_supplied_issuer_without_expected_identity(monkeypatch, error): + response, state_data = await _authorize_then_callback( + _issuer_anchored_oauth_server(issuer=None), + iss="https://unknown-idp.example.com", + monkeypatch=monkeypatch, + error=error, + ) + assert state_data["expected_issuer"] is None + assert response.status_code == 400 + assert "location" not in response.headers + assert b"upstream-auth-code" not in response.body + + +@pytest.mark.asyncio +@pytest.mark.parametrize("issuer", [None, "https://idp.example.com"]) +async def test_callback_accepts_missing_unadvertised_issuer(monkeypatch, issuer): + response, _ = await _authorize_then_callback( + _issuer_anchored_oauth_server(issuer), iss=None, monkeypatch=monkeypatch, + ) + assert response.status_code == 302 + + +@pytest.mark.asyncio +async def test_callback_rejects_an_issuer_differing_only_outside_the_path(monkeypatch): + """A tenant a deployment encoded in a query string is part of that issuer's identity, so the + comparison must not canonicalize it away and let another tenant's response through.""" + response, state_data = await _authorize_then_callback( + _issuer_anchored_oauth_server(issuer="https://idp.example.com/?tenant=a"), + iss="https://idp.example.com/?tenant=b", + monkeypatch=monkeypatch, + ) + + assert state_data["expected_issuer"] == "https://idp.example.com/?tenant=a" + assert response.status_code == 400 + assert "location" not in response.headers + assert b"upstream-auth-code" not in response.body + + +@pytest.mark.asyncio +@pytest.mark.parametrize("iss,expected_status", [(None, 302), ("https://some-idp.example.com", 400)]) +async def test_callback_legacy_state_requires_absent_issuer(monkeypatch, iss, expected_status): + response, _ = await _authorize_then_callback( + _issuer_anchored_oauth_server(), + iss=iss, + monkeypatch=monkeypatch, + expected_issuer_override=None, + ) + assert response.status_code == expected_status + + @pytest.mark.asyncio async def test_oauth_authorize_includes_scopes_from_server_config(): """Test that authorize endpoint includes scopes from server configuration.""" @@ -9070,7 +9259,8 @@ async def test_persist_dcr_client_for_config_server_uses_side_store(): @pytest.mark.asyncio -async def test_hydrate_config_server_applies_stored_dcr_client(monkeypatch): +@pytest.mark.parametrize("legacy", [False, True]) +async def test_hydrate_config_server_applies_stored_dcr_client(monkeypatch, legacy): """On restart a config server's in-memory object has no client_id; hydration overlays the persisted DCR client from the server-scoped store, decrypting the encrypted-at-rest blob, so the refresh_token grant can authenticate as the registered client instead of re-authenticating.""" @@ -9090,12 +9280,14 @@ async def test_hydrate_config_server_applies_stored_dcr_client(monkeypatch): transport=MCPTransport.http, auth_type=MCPAuth.oauth2, client_id=None, + url="https://resource.example/mcp", + issuer="https://idp.example", ) monkeypatch.setattr(enc, "_get_salt_key", lambda: "salt-hydrate-key") stored_blob = safe_dumps( encrypt_credentials( - credentials={ + credentials={**({} if legacy else {"dcr_issuer": "https://idp.example", "dcr_server_url": "https://resource.example/mcp"}), "client_id": "stored-client", "client_secret": "stored-secret", "token_endpoint_auth_method": "client_secret_basic", @@ -9145,12 +9337,14 @@ async def test_reuse_config_server_reads_store_with_real_crypto(monkeypatch): transport=MCPTransport.http, auth_type=MCPAuth.oauth2, client_id=None, + url="https://resource.example/mcp", + issuer="https://idp.example", ) monkeypatch.setattr(enc, "_get_salt_key", lambda: "salt-reuse-key") blob = safe_dumps( encrypt_credentials( - credentials={"client_id": "stored-client", "client_secret": "sec", "redirect_uris": ["https://x/callback"]}, + credentials={"dcr_issuer": "https://idp.example", "dcr_server_url": "https://resource.example/mcp","client_id": "stored-client", "client_secret": "sec", "redirect_uris": ["https://x/callback"]}, encryption_key="salt-reuse-key", ) ) @@ -12759,3 +12953,232 @@ def test_invalidating_an_idle_server_leaves_no_generation_behind(): finally: for server_id in server_ids: discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server_id, None) + + +@pytest.mark.asyncio +async def test_callback_rejects_missing_advertised_issuer(monkeypatch): + server = _issuer_anchored_oauth_server().model_copy( + update={"authorization_response_iss_parameter_supported": True} + ) + response, _ = await _authorize_then_callback(server, iss=None, monkeypatch=monkeypatch) + assert response.status_code == 400 + assert "location" not in response.headers + assert b"upstream-auth-code" not in response.body + + +@pytest.mark.parametrize("issuer,url", [ + ("https://other.example", "https://resource.example/mcp"), + (None, "https://other.example/mcp"), +]) +def test_persisted_client_is_not_applied_to_another_upstream(issuer, url): + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _apply_persisted_dcr_credentials, _PersistedDcrCredentials, + ) + server = _issuer_anchored_oauth_server(issuer).model_copy(update={"url": url, "client_id": None}) + stored = _PersistedDcrCredentials.model_validate({ + "client_id": "old-client", "client_secret": "old-secret", + "dcr_issuer": "https://idp.example.com", "dcr_server_url": "https://resource.example/mcp", + }) + assert _apply_persisted_dcr_credentials(server, stored) is False + assert server.client_id is None + assert server.client_secret is None + + +@pytest.mark.asyncio +async def test_registration_does_not_write_client_after_upstream_edit(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "test-mcp-registration-salt") + from prisma import models + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _persist_dcr_client_registration + + server = _issuer_anchored_oauth_server().model_copy(update={"url": "https://resource.example/mcp", "client_id": None}) + current = models.LiteLLM_MCPServerTable.model_construct( + server_id=server.server_id, url="https://other.example/mcp", issuer="https://other.example", + auth_type="oauth2", transport="http", credentials=None, updated_at=datetime.now(timezone.utc), + ) + prisma = MagicMock() + table = AsyncMock() + table.find_unique.return_value = current + table.update.return_value = current + prisma.db.litellm_mcpservertable = table + monkeypatch.setattr(proxy_server, "prisma_client", prisma) + result = await _persist_dcr_client_registration(server, {"client_id": "old-issuer-client"}, "https://proxy.example/callback") + assert result == "failed" + table.update.assert_not_called() + table.update_many.assert_not_called() + assert server.client_id is None + + +@pytest.mark.parametrize("response_issuer", ["https://IDP.example.com", "https://idp.example.com/", "https://idp.example.com:443"]) +def test_issuer_validation_uses_exact_identifier(response_issuer): + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _authorization_response_issuer_is_trusted + assert _authorization_response_issuer_is_trusted(response_issuer, {"expected_issuer": "https://idp.example.com"}) is False + + +@pytest.mark.parametrize("issuer,url", [ + ("https://idp.example.com", "https://resource.example/mcp"), + ("https://idp.example.com", "https://other-resource.example/mcp"), +]) +def test_persisted_client_remains_usable_for_its_issuer(issuer, url, monkeypatch): + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _apply_persisted_dcr_credentials, _PersistedDcrCredentials + monkeypatch.setenv("LITELLM_SALT_KEY", "test-registration-binding") + server = _issuer_anchored_oauth_server(issuer).model_copy(update={"url": url, "client_id": None}) + credentials = _PersistedDcrCredentials( + client_id="registered-client", client_secret="registered-secret", dcr_issuer=issuer, + dcr_server_url="https://resource.example/mcp", + ) + assert _apply_persisted_dcr_credentials(server, credentials) is True + assert server.client_id == "registered-client" + assert server.client_secret == "registered-secret" + assert server.dcr_issuer == issuer + + +@pytest.mark.asyncio +async def test_callback_does_not_forward_error_from_another_issuer(monkeypatch): + response, _ = await _authorize_then_callback( + _issuer_anchored_oauth_server(), iss="https://other.example", monkeypatch=monkeypatch, error="access_denied", + ) + assert response.status_code == 400 + assert "location" not in response.headers + assert b"invalid_issuer" in response.body + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["authorize", "token"]) +async def test_bound_client_cannot_be_sent_to_a_different_issuer(monkeypatch, operation): + from fastapi import Request + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="bound-client", name="bound-client", transport="http", auth_type="oauth2", + url="https://new.example/mcp", issuer="https://new.example", client_id="old-client", + dcr_issuer="https://old.example", dcr_server_url="https://old.example/mcp", + ) + monkeypatch.setattr(endpoints, "_server_with_oauth_endpoints", AsyncMock(return_value=server)) + http_client = MagicMock() + monkeypatch.setattr(endpoints, "get_async_httpx_client", http_client) + request = ( + endpoints.authorize_with_server( + mcp_server=server, request=MagicMock(spec=Request), redirect_uri="http://localhost/callback", client_id="old-client", + ) + if operation == "authorize" + else endpoints.exchange_token_with_server( + mcp_server=server, request=MagicMock(spec=Request), grant_type="authorization_code", + code="test-code", redirect_uri="http://localhost/callback", client_id="old-client", + client_secret=None, code_verifier="verifier", + ) + ) + with pytest.raises(HTTPException) as error: + await request + assert error.value.status_code == 400 + assert "different issuer" in error.value.detail + http_client.assert_not_called() + + +@pytest.mark.asyncio +async def test_failed_registration_preserves_cached_credentials(monkeypatch): + from fastapi import Request + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + + server = _dcr_redirect_test_server("old-client").model_copy(update={ + "url": "https://new.example/mcp", "issuer": "https://new.example", + "dcr_issuer": "https://old.example", "client_secret": "old-secret", + }) + monkeypatch.setattr(endpoints, "_server_with_oauth_endpoints", AsyncMock(return_value=server)) + monkeypatch.setattr(endpoints, "_reuse_persisted_dcr_client_if_available", AsyncMock(return_value=False)) + response = MagicMock() + response.json.return_value = {"client_id": "new-client"} + register = AsyncMock(return_value=response) + monkeypatch.setattr(endpoints, "_post_dcr_registration", register) + persist = AsyncMock(return_value="failed") + monkeypatch.setattr(endpoints, "_persist_dcr_client_registration", persist) + request = MagicMock(spec=Request) + request.base_url = "https://gateway.example/" + request.headers = {} + with pytest.raises(HTTPException) as error: + await endpoints.register_client_with_server( + request=request, mcp_server=server, client_name="app", grant_types=["authorization_code"], + response_types=["code"], token_endpoint_auth_method="none", persist_credentials=True, + ) + assert error.value.status_code == 503 + assert "could not be saved" in error.value.detail + assert server.client_id == "old-client" + assert server.client_secret == "old-secret" + register.assert_awaited_once() + persist.assert_awaited_once_with(server, {"client_id": "new-client"}, "https://gateway.example/callback") + + +@pytest.mark.asyncio +async def test_registration_without_database_keeps_client_in_temporary_server(monkeypatch): + from fastapi import Request + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints + + monkeypatch.setattr(proxy_server, "prisma_client", None) + server = _dcr_redirect_test_server(None).model_copy(update={"issuer": "https://idp.example", "url": "https://resource.example/mcp"}) + monkeypatch.setattr(endpoints, "_server_with_oauth_endpoints", AsyncMock(return_value=server)) + upstream = MagicMock() + upstream.json.return_value = {"client_id": "temporary-client", "client_secret": "temporary-secret"} + monkeypatch.setattr(endpoints, "_post_dcr_registration", AsyncMock(return_value=upstream)) + request = MagicMock(spec=Request) + request.base_url = "https://gateway.example/" + request.headers = {} + response = await endpoints.register_client_with_server( + request=request, mcp_server=server, client_name="app", grant_types=["authorization_code"], + response_types=["code"], token_endpoint_auth_method="none", persist_credentials=True, + ) + assert response.status_code == 200 + assert json.loads(response.body)["client_id"] == "temporary-client" + assert server.client_id == "temporary-client" + assert server.client_secret == "temporary-secret" + assert server.dcr_issuer == server.issuer + assert server.dcr_server_url == server.url + + +@pytest.mark.asyncio +@pytest.mark.parametrize("config_store", [False, True]) +async def test_registration_does_not_overwrite_credentials_after_failed_identity_read(monkeypatch, config_store): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints, db + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + from litellm.proxy import utils + + server = _dcr_redirect_test_server(None) + monkeypatch.setattr(utils, "get_prisma_client_or_throw", lambda _: MagicMock()) + monkeypatch.setattr(db, "get_mcp_server", AsyncMock(return_value=None, side_effect=None if config_store else RuntimeError("unavailable"))) + monkeypatch.setattr(db, "get_mcp_server_oauth_client_credentials", AsyncMock(side_effect=RuntimeError("unavailable"))) + monkeypatch.setattr(global_mcp_server_manager, "is_config_declared_server", lambda _: config_store) + update = AsyncMock() + upsert = AsyncMock() + monkeypatch.setattr(db, "update_mcp_server", update) + monkeypatch.setattr(db, "upsert_mcp_server_oauth_client_credentials", upsert) + result = await endpoints._persist_dcr_client_registration(server, {"client_id": "new-client"}, "https://gateway.example/callback") + assert result == "failed" + assert server.client_id is None + update.assert_not_awaited() + upsert.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("winner_available", [False, True]) +async def test_registration_losing_conditional_write_reuses_only_a_matching_winner(monkeypatch, winner_available): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints, db + from litellm.proxy._types import LiteLLM_MCPServerTable + from litellm.proxy import utils + + server = _dcr_redirect_test_server(None).model_copy(update={"url": "https://upstream.example/mcp"}) + row = LiteLLM_MCPServerTable.model_validate(server.model_dump(exclude_none=True)) + winner = row.model_copy(update={"credentials": { + "client_id": "winner-client", "redirect_uris": ["https://gateway.example/callback"], + "dcr_server_url": server.url, + }}) if winner_available else row + monkeypatch.setattr(utils, "get_prisma_client_or_throw", lambda _: MagicMock()) + read = AsyncMock(side_effect=[row, winner]) + monkeypatch.setattr(db, "get_mcp_server", read) + monkeypatch.setattr(endpoints, "_refresh_persisted_dcr_server", AsyncMock()) + update = AsyncMock(return_value=None) + monkeypatch.setattr(db, "update_mcp_server", update) + result = await endpoints._persist_dcr_client_registration(server, {"client_id": "losing-client"}, "https://gateway.example/callback") + assert result == ("reused" if winner_available else "failed") + assert update.await_args.kwargs["expected_updated_at"] == row.updated_at + assert server.client_id == ("winner-client" if winner_available else None) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_partial_update.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_partial_update.py index a3c52dc16b7..122b22f6273 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_partial_update.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_partial_update.py @@ -17,6 +17,7 @@ from fastapi import HTTPException from litellm.proxy._experimental.mcp_server.db import ( create_mcp_server, + decrypt_credentials, set_mcp_server_pinned_tools, update_mcp_server, ) @@ -1181,3 +1182,186 @@ async def test_protocol_update_preserves_missing_server_without_writing(clear_al result = await update_mcp_server(prisma, payload, "admin") assert result is None table.update.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("changed,previous_issuer", [ + ({"url": "https://new.example/mcp"}, None), + ({"issuer": "https://new.example"}, "https://old.example"), + ({"url": "https://new.example/mcp", "auth_type": "api_key"}, "https://old.example"), +]) +async def test_upstream_identity_edit_drops_previous_oauth_client(changed, previous_issuer): + prisma = _mock_prisma() + existing = models.LiteLLM_MCPServerTable.model_construct( + server_id="test-server", transport="http", auth_type="oauth2", + url="https://old.example/mcp", issuer=previous_issuer, + credentials=json.dumps({"client_id": "old-client", "client_secret": "old-secret", "upstream_resource": "api://resource"}), + ) + prisma.db.litellm_mcpservertable.find_unique.return_value = existing + await update_mcp_server(prisma, UpdateMCPServerRequest(server_id="test-server", **changed), "test-user") + written = prisma.db.litellm_mcpservertable.update.call_args.kwargs["data"] + assert "credentials" in written + credentials = json.loads(written["credentials"]) + assert "client_id" not in credentials + assert "client_secret" not in credentials + assert credentials["upstream_resource"] == "api://resource" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("issuer", ["https://old.example/", "https://OLD.example:443"]) +async def test_distinct_issuer_identifier_edit_discards_registered_client(issuer): + prisma = _mock_prisma() + existing = models.LiteLLM_MCPServerTable.model_construct( + server_id="test-server", transport="http", auth_type="oauth2", + url="https://resource.example/mcp", issuer="https://old.example", + credentials=json.dumps({"client_id": "old-client", "dcr_issuer": "https://old.example"}), + ) + prisma.db.litellm_mcpservertable.find_unique.return_value = existing + await update_mcp_server(prisma, UpdateMCPServerRequest(server_id="test-server", issuer=issuer), "test-user") + written = prisma.db.litellm_mcpservertable.update.call_args.kwargs["data"] + assert "client_id" not in json.loads(written["credentials"]) + assert written["issuer"] == issuer + + +@pytest.mark.asyncio +@pytest.mark.parametrize("previous_issuer,binding", [("https://idp.example", None), (None, "https://idp.example")]) +@pytest.mark.parametrize("submitted_tokens", [None, {"access_token": "old-token", "refresh_token": "old-refresh", "expires_in": 3600}, {"access_token": "fresh-token", "refresh_token": "fresh-refresh", "expires_in": 3600}]) +async def test_url_edit_preserves_client_bound_to_previous_known_issuer_without_old_tokens(previous_issuer, binding, submitted_tokens): + prisma = _mock_prisma() + existing = models.LiteLLM_MCPServerTable.model_construct( + server_id="test-server", transport="http", auth_type="oauth2", + url="https://resource.example/mcp", issuer=previous_issuer, + credentials=json.dumps({"client_id": "static-client", "client_secret": "static-secret", "dcr_issuer": binding, + "access_token": "old-token", "refresh_token": "old-refresh", "expires_in": 3600, + "token_endpoint_auth_method": "client_secret_basic"}), + ) + prisma.db.litellm_mcpservertable.find_unique.return_value = existing + await update_mcp_server(prisma, UpdateMCPServerRequest(server_id="test-server", url=existing.url + "?v=2", **({"credentials": submitted_tokens} if submitted_tokens is not None else {})), "test-user") + written = prisma.db.litellm_mcpservertable.update.call_args.kwargs["data"] + credentials = json.loads(written["credentials"]) + assert credentials["client_id"] == "static-client" + assert credentials["client_secret"] == "static-secret" + assert credentials["dcr_issuer"] == "https://idp.example" + assert credentials["dcr_server_url"] == existing.url + assert "access_token" not in credentials + assert "refresh_token" not in credentials + assert "expires_in" not in credentials + assert credentials["token_endpoint_auth_method"] == "client_secret_basic" + assert written["issuer"] is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("changed", [0, 1]) +async def test_registration_write_requires_unchanged_server_revision(changed): + from datetime import datetime, timezone + from litellm.proxy._experimental.mcp_server.db import _update_mcp_server_row + + prisma = _mock_prisma() + table = prisma.db.litellm_mcpservertable + revision = datetime.now(timezone.utc) + table.update_many.return_value = changed + result = await _update_mcp_server_row( + prisma, server_id="test-server", data_dict={"credentials": "{}"}, expected_updated_at=revision, + ) + table.update_many.assert_awaited_once_with( + where={"server_id": "test-server", "updated_at": revision}, data={"credentials": "{}"}, + ) + table.update.assert_not_awaited() + assert (result is not None) is bool(changed) + if changed: + assert result.server_id == "test-server" + else: + table.find_unique.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("operation", ["create", "rotate"]) +async def test_explicit_oauth_client_write_binds_to_server_identity(operation): + prisma = _mock_prisma() + existing = models.LiteLLM_MCPServerTable.model_construct( + server_id="test-server", transport="http", auth_type="oauth2", + url="https://resource.example/mcp", issuer="https://idp.example", + credentials=json.dumps({"client_id": "old-client"}), + ) + prisma.db.litellm_mcpservertable.find_unique.return_value = existing + if operation == "create": + await create_mcp_server(prisma, NewMCPServerRequest( + server_name="bound-client", transport="http", auth_type="oauth2", + url=existing.url, issuer=existing.issuer, credentials={"client_id": "new-client"}, + ), "admin") + written = prisma.db.litellm_mcpservertable.create.call_args.kwargs["data"] + else: + await update_mcp_server(prisma, UpdateMCPServerRequest( + server_id=existing.server_id, credentials={"client_id": "new-client"}, + ), "admin") + written = prisma.db.litellm_mcpservertable.update.call_args.kwargs["data"] + credentials = json.loads(written["credentials"]) + assert credentials["dcr_issuer"] == existing.issuer + assert credentials["dcr_server_url"] == existing.url + + +@pytest.mark.asyncio +async def test_unrelated_edit_does_not_backfill_legacy_oauth_binding(): + prisma = _mock_prisma() + existing = models.LiteLLM_MCPServerTable.model_construct( + server_id="test-server", transport="http", auth_type="oauth2", + url="https://resource.example/mcp", issuer="https://idp.example", + credentials=json.dumps({"client_id": "legacy-client"}), + ) + prisma.db.litellm_mcpservertable.find_unique.return_value = existing + await update_mcp_server(prisma, UpdateMCPServerRequest( + server_id=existing.server_id, credentials={"scopes": ["tools.read"]}, + ), "admin") + written = prisma.db.litellm_mcpservertable.update.call_args.kwargs["data"] + credentials = json.loads(written["credentials"]) + assert credentials["client_id"] == "legacy-client" + assert "dcr_issuer" not in credentials + assert "dcr_server_url" not in credentials + + +@pytest.mark.asyncio +async def test_issuer_edit_does_not_rebind_resubmitted_saved_client(): + prisma = _mock_prisma() + existing = models.LiteLLM_MCPServerTable.model_construct( + server_id="test-server", transport="http", auth_type="oauth2", + url="https://old.example/mcp", issuer="https://old.example", + credentials=json.dumps({"client_id": "saved-client", "client_secret": "saved-secret"}), + ) + prisma.db.litellm_mcpservertable.find_unique.return_value = existing + await update_mcp_server(prisma, UpdateMCPServerRequest( + server_id=existing.server_id, url="https://new.example/mcp", issuer="https://new.example", + credentials={"client_id": "saved-client", "client_secret": "saved-secret"}, + ), "admin") + written = prisma.db.litellm_mcpservertable.update.call_args.kwargs["data"] + credentials = json.loads(written["credentials"]) + assert "client_id" not in credentials + assert "client_secret" not in credentials + + +@pytest.mark.asyncio +@pytest.mark.parametrize("replacement", [ + {"client_secret": "replacement-secret"}, + {"client_secret": "old-secret", "token_endpoint_auth_method": "client_secret_basic"}, + {"client_secret": None}, + {"dcr_issuer": "https://new.example", "dcr_server_url": "https://new.example/mcp"}, +]) +async def test_issuer_edit_preserves_replacement_with_same_client_id(replacement): + prisma = _mock_prisma() + existing = models.LiteLLM_MCPServerTable.model_construct( + server_id="test-server", transport="http", auth_type="oauth2", + url="https://old.example/mcp", issuer="https://old.example", + credentials=json.dumps({"client_id": "shared-client", "client_secret": "old-secret"}), + ) + prisma.db.litellm_mcpservertable.find_unique.return_value = existing + submitted = {"client_id": "shared-client", **replacement} + await update_mcp_server(prisma, UpdateMCPServerRequest( + server_id=existing.server_id, url="https://new.example/mcp", issuer="https://new.example", + credentials=submitted, + ), "admin") + written = prisma.db.litellm_mcpservertable.update.call_args.kwargs["data"] + credentials = decrypt_credentials(json.loads(written["credentials"])) + assert credentials["client_id"] == "shared-client" + assert credentials["dcr_issuer"] == "https://new.example" + assert credentials["dcr_server_url"] == "https://new.example/mcp" + assert credentials.get("client_secret") == replacement.get("client_secret") + assert credentials.get("token_endpoint_auth_method") == replacement.get("token_endpoint_auth_method") 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 e7602d7f825..123a5505953 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 @@ -9186,6 +9186,8 @@ class TestMCPServerTimestamps: authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", registration_url="https://idp.example.com/register", + issuer="https://idp.example.com", + authorization_response_iss_parameter_supported=True, ) blipped = MCPServer( @@ -9200,6 +9202,16 @@ class TestMCPServerTimestamps: assert blipped.authorization_url == "https://idp.example.com/authorize" assert blipped.token_url == "https://idp.example.com/token" assert blipped.registration_url == "https://idp.example.com/register" + assert blipped.issuer == "https://idp.example.com" + assert blipped.authorization_response_iss_parameter_supported is True + fallback = MCPServerManager._merge_discovered_oauth_metadata( + blipped, MCPOAuthMetadata(from_origin_fallback=True), + ) + assert fallback.authorization_response_iss_parameter_supported is True + refreshed = MCPServerManager._merge_discovered_oauth_metadata( + fallback, MCPOAuthMetadata(discovered_issuer="https://idp.example.com"), + ) + assert refreshed.authorization_response_iss_parameter_supported is False same_authorize = MCPServer( server_id="s1", @@ -9316,8 +9328,8 @@ class TestMCPServerTimestamps: from litellm.proxy._experimental.mcp_server.mcp_server_manager import _issuer_matches assert _issuer_matches("https://mcp.slack.com", "https://mcp.slack.com") - assert _issuer_matches("https://MCP.slack.com/", "https://mcp.slack.com") - assert _issuer_matches("https://mcp.slack.com:443", "https://mcp.slack.com") + assert not _issuer_matches("https://MCP.slack.com/", "https://mcp.slack.com") + assert not _issuer_matches("https://mcp.slack.com:443", "https://mcp.slack.com") assert _issuer_matches("https://login.example.com/tenant/v2.0", "https://login.example.com/tenant/v2.0") assert not _issuer_matches("https://attacker.example.com", "https://mcp.slack.com") assert not _issuer_matches("https://login.example.com/other/v2.0", "https://login.example.com/tenant/v2.0") diff --git a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py index 2cbf1d578b2..081a2d8ce73 100644 --- a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -1876,7 +1876,7 @@ class TestTemporaryMCPSessionEndpoints: url="https://temp.example.com", transport=MCPTransport.http, ) - existing_server = MagicMock() + existing_server = MagicMock(dcr_issuer=None, dcr_server_url=None, token_endpoint_auth_method=None) existing_server.authentication_token = "token-abc" existing_server.client_id = "client-123" existing_server.client_secret = "secret-xyz" @@ -1912,7 +1912,7 @@ class TestTemporaryMCPSessionEndpoints: @staticmethod def _inherit_with(payload_credentials, **server_overrides): - existing_server = MagicMock() + existing_server = MagicMock(dcr_issuer=None, dcr_server_url=None, token_endpoint_auth_method=None) existing_server.authentication_token = None existing_server.client_id = "client-123" existing_server.client_secret = "secret-xyz" @@ -2616,6 +2616,9 @@ class TestTemporaryMCPSessionEndpoints: user_id="admin-user", ) inherited_server = MagicMock( + dcr_issuer=None, + dcr_server_url=None, + token_endpoint_auth_method=None, authentication_token="token-abc", client_id="client-id", client_secret="client-secret", @@ -11145,3 +11148,177 @@ async def test_protocol_update_on_missing_server_preserves_not_found() -> None: user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), ) assert error.value.status_code == 404 + + +@pytest.mark.asyncio +async def test_changed_upstream_session_gets_an_isolated_id_without_saved_client(monkeypatch): + from litellm.proxy.management_endpoints import mcp_management_endpoints as management + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + saved = MCPServer( + server_id="saved-upstream", name="saved", transport="http", auth_type="oauth2", + url="https://old.example/mcp", issuer="https://old.example", + client_id="old-client", client_secret="old-secret", + ) + monkeypatch.setitem(management.global_mcp_server_manager.registry, saved.server_id, saved) + payload = NewMCPServerRequest( + server_id=saved.server_id, server_name="saved", transport="http", auth_type="oauth2", + oauth2_flow="authorization_code", url="https://new.example/mcp", issuer="https://new.example", + ) + staged = management._inherit_credentials_from_existing_server(payload) + assert not (staged.credentials or {}).get("client_id") + assert not (staged.credentials or {}).get("client_secret") + assert await management._resolve_session_server_id(staged) != saved.server_id + assert saved.client_id == "old-client" + + +def test_staged_url_edit_clears_resubmitted_issuer_and_endpoints(monkeypatch): + from litellm.proxy.management_endpoints import mcp_management_endpoints as management + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + saved = MCPServer( + server_id="saved-oauth", name="saved", transport="http", auth_type="oauth2", + url="https://old.example/mcp", issuer="https://old.example", + authorization_url="https://old.example/authorize", token_url="https://old.example/token", + registration_url="https://old.example/register", client_id="old-client", + ) + monkeypatch.setitem(management.global_mcp_server_manager.registry, saved.server_id, saved) + payload = NewMCPServerRequest( + server_id=saved.server_id, server_name="saved", transport="http", auth_type="oauth2", + oauth2_flow="authorization_code", url="https://new.example/mcp", issuer=saved.issuer, + authorization_url=saved.authorization_url, token_url=saved.token_url, registration_url=saved.registration_url, + ) + staged = management._inherit_credentials_from_existing_server(payload) + assert staged.issuer is None + assert staged.authorization_url is None + assert staged.token_url is None + assert staged.registration_url is None + assert staged.oauth2_flow == "authorization_code" + assert saved.issuer == "https://old.example" + + +@pytest.mark.asyncio +async def test_distinct_issuer_identifier_edit_isolates_session_client(monkeypatch): + from litellm.proxy.management_endpoints import mcp_management_endpoints as management + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + saved = MCPServer( + server_id="saved-oauth", name="saved", transport="http", auth_type="oauth2", + url="https://resource.example/mcp", issuer="https://idp.example", + client_id="saved-client", client_secret="saved-secret", + ) + monkeypatch.setitem(management.global_mcp_server_manager.registry, saved.server_id, saved) + payload = NewMCPServerRequest( + server_id=saved.server_id, server_name="saved", transport="http", auth_type="oauth2", + url=saved.url, issuer="https://IDP.example:443/", + ) + staged = management._inherit_credentials_from_existing_server(payload) + assert not (staged.credentials or {}).get("client_id") + assert not (staged.credentials or {}).get("client_secret") + assert await management._resolve_session_server_id(staged) != saved.server_id + + +@pytest.mark.asyncio +async def test_url_edit_stages_existing_client_with_previous_issuer_binding(monkeypatch): + from litellm.proxy.management_endpoints import mcp_management_endpoints as management + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + saved = MCPServer( + server_id="saved-oauth", name="saved", transport="http", auth_type="oauth2", + url="https://resource.example/mcp", issuer="https://idp.example", + client_id="static-client", client_secret="static-secret", authentication_token="old-token", + token_endpoint_auth_method="client_secret_basic", + ) + monkeypatch.setitem(management.global_mcp_server_manager.registry, saved.server_id, saved) + payload = NewMCPServerRequest( + server_id=saved.server_id, server_name="saved", transport="http", auth_type="oauth2", + url=saved.url + "?v=2", issuer=saved.issuer, + ) + staged = management._inherit_credentials_from_existing_server(payload) + assert staged.credentials["client_id"] == "static-client" + assert staged.credentials["client_secret"] == "static-secret" + assert staged.credentials["dcr_issuer"] == saved.issuer + assert staged.credentials["dcr_server_url"] == saved.url + assert staged.credentials["token_endpoint_auth_method"] == "client_secret_basic" + assert "auth_value" not in staged.credentials + assert staged.issuer is None + assert await management._resolve_session_server_id(staged) != saved.server_id + + +@pytest.mark.asyncio +@pytest.mark.parametrize("upstream_changed", [False, True]) +async def test_session_identity_checks_saved_row_when_registry_is_empty(monkeypatch, upstream_changed): + from litellm.proxy.management_endpoints import mcp_management_endpoints as management + + saved = generate_mock_mcp_server_db_record(server_id="saved-oauth-row") + saved.auth_type = "oauth2" + saved.url = "https://old.example/mcp" + saved.issuer = "https://old.example" + saved.approval_status = "approved" + monkeypatch.setattr(management.global_mcp_server_manager, "get_mcp_server_by_id", lambda _: None) + monkeypatch.setattr(management, "_get_prisma_client_or_none", lambda: MagicMock()) + monkeypatch.setattr(management, "get_mcp_server", AsyncMock(return_value=saved)) + payload = NewMCPServerRequest( + server_id=saved.server_id, auth_type="oauth2", transport="http", + url="https://new.example/mcp" if upstream_changed else saved.url, + ) + resolved = await management._resolve_session_server_id(payload) + assert (resolved != saved.server_id) is upstream_changed + + +def test_new_oauth_session_does_not_look_up_saved_credentials(monkeypatch): + from litellm.proxy.management_endpoints import mcp_management_endpoints as management + + lookup = MagicMock() + monkeypatch.setattr(management.global_mcp_server_manager, "get_mcp_server_by_id", lookup) + payload = NewMCPServerRequest(url="https://new.example/mcp", auth_type="oauth2", transport="http") + assert management._inherit_credentials_from_existing_server(payload) is payload + assert not payload.credentials + lookup.assert_not_called() + + +def test_edit_does_not_rebind_resubmitted_saved_client_to_new_issuer(monkeypatch): + from litellm.proxy.management_endpoints import mcp_management_endpoints as management + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + saved = MCPServer( + server_id="saved-static", name="saved", transport="http", auth_type="oauth2", + url="https://old.example/mcp", issuer="https://old.example", + client_id="saved-client", client_secret="saved-secret", + ) + monkeypatch.setitem(management.global_mcp_server_manager.registry, saved.server_id, saved) + payload = NewMCPServerRequest( + server_id=saved.server_id, transport="http", auth_type="oauth2", + url="https://new.example/mcp", issuer="https://new.example", + credentials={"client_id": saved.client_id, "client_secret": saved.client_secret}, + ) + staged = management._inherit_credentials_from_existing_server(payload) + assert not (staged.credentials or {}).get("client_id") + assert not (staged.credentials or {}).get("client_secret") + assert saved.client_id == "saved-client" + + +@pytest.mark.parametrize("replacement", [ + {"client_secret": "replacement-secret"}, + {"client_secret": "old-secret", "token_endpoint_auth_method": "client_secret_basic"}, + {"client_secret": None}, + {"dcr_issuer": "https://new.example", "dcr_server_url": "https://new.example/mcp"}, +]) +def test_staged_issuer_edit_preserves_replacement_with_same_client_id(monkeypatch, replacement): + from litellm.proxy.management_endpoints import mcp_management_endpoints as management + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + saved = MCPServer( + server_id="saved-static", name="saved", transport="http", auth_type="oauth2", + url="https://old.example/mcp", issuer="https://old.example", + client_id="shared-client", client_secret="old-secret", + ) + monkeypatch.setitem(management.global_mcp_server_manager.registry, saved.server_id, saved) + submitted = {"client_id": "shared-client", **replacement} + payload = NewMCPServerRequest( + server_id=saved.server_id, transport="http", auth_type="oauth2", + url="https://new.example/mcp", issuer="https://new.example", credentials=submitted, + ) + staged = management._inherit_credentials_from_existing_server(payload) + assert staged.credentials == submitted + assert saved.client_secret == "old-secret" diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx index 3cf2bcdf238..f6ad2e38658 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx @@ -28,7 +28,10 @@ const oauthHook = vi.hoisted(() => ({ tokenResponse: null as Record | null, reset: vi.fn(), onTokenReceived: null as - | ((token: Record | null, registeredClient?: { clientId?: string; clientSecret?: string }) => void) + | (( + token: Record | null, + registeredClient?: { client_id: string; client_secret?: string }, + ) => void) | null, getCredentials: null as (() => Record | undefined) | null, getTemporaryPayload: null as (() => Record | null) | null, @@ -37,7 +40,7 @@ vi.mock("@/hooks/useMcpOAuthFlow", () => ({ useMcpOAuthFlow: (opts: { onTokenReceived: ( token: Record | null, - registeredClient?: { clientId?: string; clientSecret?: string }, + registeredClient?: { client_id: string; client_secret?: string }, ) => void; getCredentials?: () => Record | undefined; getTemporaryPayload?: () => Record | null; @@ -628,7 +631,7 @@ describe("CreateMCPServer", () => { await act(async () => { oauthHook.onTokenReceived!( { access_token: "oauth2-minted-tok", refresh_token: "oauth2-minted-refresh", token_type: "Bearer" }, - { clientId: "dcr-minted-client", clientSecret: "dcr-minted-secret" }, + { client_id: "dcr-minted-client", client_secret: "dcr-minted-secret" }, ); }); @@ -675,7 +678,7 @@ describe("CreateMCPServer", () => { await act(async () => { oauthHook.onTokenReceived!( { access_token: "oauth2-tok", token_type: "Bearer" }, - { clientId: "dcr-client", clientSecret: "dcr-secret" }, + { client_id: "dcr-client", client_secret: "dcr-secret" }, ); }); @@ -702,7 +705,7 @@ describe("CreateMCPServer", () => { await act(async () => { oauthHook.onTokenReceived!( { access_token: "oauth2-tok", token_type: "Bearer" }, - { clientId: "leak-client", clientSecret: "leak-secret" }, + { client_id: "leak-client", client_secret: "leak-secret" }, ); }); // Ref is held while the modal is open. @@ -731,7 +734,7 @@ describe("CreateMCPServer", () => { await act(async () => { oauthHook.onTokenReceived!( { access_token: "oauth2-tok", token_type: "Bearer" }, - { clientId: "dcr-client", clientSecret: "dcr-secret" }, + { client_id: "dcr-client", client_secret: "dcr-secret" }, ); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx index 2d6e83aff9b..d53d11ddcfb 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx @@ -56,7 +56,7 @@ import EnvVarsSection from "./EnvVarsSection"; import { isAdminRole } from "@/utils/roles"; import { validateMCPServerUrl, validateMCPServerName } from "./utils"; import { toast } from "@/lib/toast"; -import { useMcpOAuthFlow } from "@/hooks/useMcpOAuthFlow"; +import { useMcpOAuthFlow, type McpDcrCredentials } from "@/hooks/useMcpOAuthFlow"; import { useTestMCPConnection } from "@/hooks/useTestMCPConnection"; import { MountedFormField, @@ -144,7 +144,7 @@ const CreateMCPServer: React.FC = ({ // it can never be collected as a client-forwarded server's declared app; injected into the payload // only on an oauth2 submit (where persisting the registered client is correct), and cleared on any // invalidation or modal close. An abandoned authorize leaves it null, which is the desired asymmetry. - const dcrClientRef = React.useRef<{ client_id: string; client_secret?: string } | null>(null); + const dcrClientRef = React.useRef(null); // Set when the upstream identity (url/endpoints) changed while a declared app is present, so the // section can warn that the saved app may not match the new upstream (the app is kept, not wiped). const [appMayNotMatchUpstream, setAppMayNotMatchUpstream] = useState(false); @@ -264,12 +264,7 @@ const CreateMCPServer: React.FC = ({ // The DCR-minted client is held in a ref, NOT written into form.credentials, so it can never be // collected as a client-forwarded server's declared app; it is injected into the payload only on // an oauth2 submit. An admin-typed client already lives in form.credentials and is left untouched. - dcrClientRef.current = registeredClient?.clientId - ? { - client_id: registeredClient.clientId, - ...(registeredClient.clientSecret && { client_secret: registeredClient.clientSecret }), - } - : null; + dcrClientRef.current = registeredClient ?? null; const current = (allFieldsValue(form).credentials as Record | undefined) ?? {}; const nextCredentials = { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/createServerPayload.ts b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/createServerPayload.ts index f45857fcc8b..30c4c502c42 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/createServerPayload.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/createServerPayload.ts @@ -29,7 +29,7 @@ export const AUTH_TYPES_REQUIRING_CREDENTIALS = [ export interface DcrClient { readonly client_id: string; - readonly client_secret?: string; + readonly client_secret?: string | null; } export interface CreateServerUiState { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.differential.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.differential.test.ts index 66786fac1c3..8a9bac582d6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.differential.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.differential.test.ts @@ -370,3 +370,31 @@ void isClientForwardedTokenMode; void normalizeEnvVars; void preservedAdminCredentials; export type { MCPServer }; + +it.each([ + [AUTH_TYPE.OAUTH2, OAUTH_FLOW.INTERACTIVE, true], + [AUTH_TYPE.OAUTH2, OAUTH_FLOW.M2M, false], + [AUTH_TYPE.TRUE_PASSTHROUGH, OAUTH_FLOW.INTERACTIVE, false], +])("applies a pending DCR client only to an interactive gateway OAuth save (%s, %s)", (authType, flow, useDcr) => { + const values: EditServerFormValues = { + auth_type: authType, + oauth_flow_type: flow, + transport: "http", + credentials: { client_id: "configured-client", client_secret: "configured-secret" }, + }; + const ui: EditServerUiState = { + ...baseUi, + dcrClient: { + client_id: "new-client", + client_secret: null, + dcr_issuer: "https://new.example", + dcr_server_url: "https://new.example/mcp", + }, + }; + const result = buildEditServerPayload(values, ui); + expect(result.kind).toBe("ok"); + if (result.kind !== "ok") throw new Error("Expected valid MCP server payload"); + expect(result.payload.credentials?.client_id).toBe(useDcr ? "new-client" : "configured-client"); + expect(result.payload.credentials?.client_secret).toBe(useDcr ? null : "configured-secret"); + expect(result.payload.credentials?.dcr_issuer).toBe(useDcr ? "https://new.example" : undefined); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.ts b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.ts index 3930f1edb9b..f166fa20d1f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/editServerPayload.ts @@ -1,3 +1,4 @@ +import type { McpDcrCredentials } from "@/hooks/useMcpOAuthFlow"; import { ADMIN_CONFIG_CREDENTIAL_KEYS, AUTH_TYPE, @@ -73,6 +74,7 @@ export interface EditServerPayload { } export interface EditServerUiState { + readonly dcrClient?: McpDcrCredentials | null; readonly mcpServer: MCPServer; readonly logoUrl: string | undefined; readonly costConfig: MCPServerCostInfo; @@ -213,6 +215,8 @@ const buildCredentials = (credentialValues: unknown): Readonly> | undefined; readonly includeCredentials: boolean; @@ -224,11 +228,16 @@ interface CredentialsEntryInput { // updates), so removal must be an explicit-null write: encrypt skips nulls and the merge overrides // the stored keys, returning the server to dynamic client registration. const resolveCredentialsEntry = ({ + dcrClient, + oauthFlow, authType, credentials, includeCredentials, removeStoredApp, }: CredentialsEntryInput): { readonly credentials?: Readonly> } => { + if (dcrClient && authType === AUTH_TYPE.OAUTH2 && oauthFlow !== OAUTH_FLOW.M2M) { + return { credentials: { ...credentials, ...dcrClient } }; + } if (removeStoredApp && isClientForwardedTokenMode(authType)) { return { credentials: { client_id: null, client_secret: null } }; } @@ -330,6 +339,8 @@ export const buildEditServerPayload = (values: EditServerFormValues, ui: EditSer const includeCredentials = restValues.auth_type && AUTH_TYPES_REQUIRING_CREDENTIALS.includes(restValues.auth_type); const credentialsEntryInput: CredentialsEntryInput = { + dcrClient: ui.dcrClient, + oauthFlow: restValues.oauth_flow_type, authType: restValues.auth_type, credentials: submitCredentials, includeCredentials: Boolean(includeCredentials), diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx index f79dd178571..91db38db389 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx @@ -16,21 +16,40 @@ vi.mock("@/components/networking", () => ({ })); const mockOauth: { + status: string; tokenResponse: any; getTemporaryPayload: (() => Record | null) | null; - onTokenReceived: ((token: Record | null) => void) | null; + onTokenReceived: + | (( + token: Record | null, + registeredClient?: { + client_id: string; + client_secret?: string; + dcr_issuer?: string; + dcr_server_url?: string; + }, + ) => void) + | null; reset: ReturnType; -} = { tokenResponse: null, getTemporaryPayload: null, onTokenReceived: null, reset: vi.fn() }; +} = { status: "idle", tokenResponse: null, getTemporaryPayload: null, onTokenReceived: null, reset: vi.fn() }; vi.mock("@/hooks/useMcpOAuthFlow", () => ({ useMcpOAuthFlow: (opts: { getTemporaryPayload?: () => Record | null; - onTokenReceived?: (token: Record | null) => void; + onTokenReceived?: ( + token: Record | null, + registeredClient?: { + client_id: string; + client_secret?: string; + dcr_issuer?: string; + dcr_server_url?: string; + }, + ) => void; }) => { mockOauth.getTemporaryPayload = opts?.getTemporaryPayload ?? null; mockOauth.onTokenReceived = opts?.onTokenReceived ?? null; return { startOAuthFlow: vi.fn(), - status: "idle", + status: mockOauth.status, error: null, tokenResponse: mockOauth.tokenResponse, reset: mockOauth.reset, @@ -547,6 +566,7 @@ describe("MCPServerEdit (auth type switch)", () => { describe("MCPServerEdit OAuth token invalidation", () => { beforeEach(() => { vi.clearAllMocks(); + mockOauth.status = "idle"; }); const renderOAuthEdit = () => @@ -560,6 +580,77 @@ describe("MCPServerEdit OAuth token invalidation", () => { />, ); + it.each(["authorizing", "exchanging"])("blocks Save while OAuth is %s", async (status) => { + mockOauth.status = status; + vi.mocked(networking.updateMCPServer).mockResolvedValue({ ...interactiveOAuthServer }); + const view = renderOAuthEdit(); + const save = screen.getAllByRole("button", { name: "Save Changes" })[0]; + expect(save).toBeDisabled(); + await act(async () => { + fireEvent.submit(screen.getByRole("form", { name: "Edit MCP server" })); + }); + expect(networking.updateMCPServer).not.toHaveBeenCalled(); + act(() => { + mockOauth.onTokenReceived?.( + { access_token: "new-token" }, + { + client_id: "new-client", + dcr_issuer: "https://new.example", + dcr_server_url: "https://new.example/mcp", + }, + ); + }); + mockOauth.status = "success"; + view.rerender( + , + ); + expect(screen.getAllByRole("button", { name: "Save Changes" })[0]).toBeEnabled(); + await act(async () => { + fireEvent.click(screen.getAllByRole("button", { name: "Save Changes" })[0]); + }); + await waitFor(() => expect(networking.updateMCPServer).toHaveBeenCalledTimes(1)); + expect(vi.mocked(networking.updateMCPServer).mock.calls[0][1].credentials).toMatchObject({ + client_id: "new-client", + dcr_issuer: "https://new.example", + }); + }); + + it.each(["Save Changes", "Cancel"])("keeps a newly registered client isolated until %s", async (action) => { + vi.mocked(networking.updateMCPServer).mockResolvedValue({ ...interactiveOAuthServer }); + renderOAuthEdit(); + act(() => { + mockOauth.onTokenReceived?.( + { access_token: "new-token" }, + { + client_id: "new-client", + client_secret: "new-secret", + dcr_issuer: "https://new.example", + dcr_server_url: "https://new.example/mcp", + }, + ); + }); + expect(networking.updateMCPServer).not.toHaveBeenCalled(); + await act(async () => { + fireEvent.click(screen.getAllByRole("button", { name: action })[0]); + }); + if (action === "Cancel") { + expect(networking.updateMCPServer).not.toHaveBeenCalled(); + return; + } + await waitFor(() => expect(networking.updateMCPServer).toHaveBeenCalledTimes(1)); + expect(vi.mocked(networking.updateMCPServer).mock.calls[0][1].credentials).toMatchObject({ + client_id: "new-client", + client_secret: "new-secret", + dcr_issuer: "https://new.example", + }); + }); + it("invalidates a session-authorized token when the transport switches to stdio", async () => { // Switching to stdio clears url/auth_type via programmatic form.setFieldsValue, which antd does // not report through onValuesChange; the explicit recheck in handleTransportChange must catch it. @@ -1687,6 +1778,51 @@ describe("MCPServerEdit (OAuth token persistence on save)", () => { expect(screen.getByText(/registered for the previous upstream/)).toBeInTheDocument(); }); + it("discards a canceled OAuth snapshot before saved server data loads", async () => { + setSecureItem( + EDIT_OAUTH_UI_STATE_KEY, + JSON.stringify({ + serverId: interactiveOAuthServer.server_id, + formValues: { ...interactiveOAuthServer, url: "https://new.example/mcp" }, + }), + ); + const props = { accessToken: "access-token", onCancel: vi.fn(), onSuccess: vi.fn(), availableAccessGroups: [] }; + const view = render(); + await act(async () => { + fireEvent.click(screen.getAllByRole("button", { name: "Cancel" })[0]); + }); + expect(props.onCancel).toHaveBeenCalledOnce(); + expect(window.sessionStorage.getItem(EDIT_OAUTH_UI_STATE_KEY)).toBeNull(); + expect(mockOauth.reset).toHaveBeenCalled(); + view.unmount(); + render(); + await waitFor(() => expect(screen.getByLabelText("MCP Server URL")).toHaveValue(interactiveOAuthServer.url)); + expect(networking.updateMCPServer).not.toHaveBeenCalled(); + }); + + it("restores the edited upstream after OAuth when saved server data loads later", async () => { + setSecureItem( + EDIT_OAUTH_UI_STATE_KEY, + JSON.stringify({ + serverId: interactiveOAuthServer.server_id, + formValues: { ...interactiveOAuthServer, url: "https://new.example/mcp", issuer: "https://new.example" }, + }), + ); + const props = { + accessToken: "access-token", + userID: "user-1", + onCancel: vi.fn(), + onSuccess: vi.fn(), + availableAccessGroups: [], + }; + const { rerender } = render( + , + ); + rerender(); + await waitFor(() => expect(screen.getByLabelText("MCP Server URL")).toHaveValue("https://new.example/mcp")); + expect(screen.getByLabelText("Issuer (optional)")).toHaveValue("https://new.example"); + }); + it("preserves a stored client_id on OAuth-resume restore even when the saved snapshot is token-only", async () => { // Post-redirect restore: the sessionStorage snapshot carries only a minted token (no client keys), // while the loaded server has a stored client_id. The restore must merge the server's declared app diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx index 3894cb6cd0b..1b1e4c27bde 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx @@ -55,7 +55,7 @@ import { EditServerFormValues, buildEditServerPayload, editPayloadErrorMessage } import { DUPLICATE_IDENTIFIER_MESSAGE, findDuplicateMcpServer, mcpSubmitErrorReason } from "./duplicateServerCheck"; import { toast } from "@/lib/toast"; import { getEditToolPreview } from "./editToolPreview"; -import { useMcpOAuthFlow } from "@/hooks/useMcpOAuthFlow"; +import { useMcpOAuthFlow, type McpDcrCredentials } from "@/hooks/useMcpOAuthFlow"; import { MountedFormField, MountedFormProvider, @@ -250,6 +250,7 @@ const MCPServerEdit: React.FC = ({ // in this edit session; undefined when none is held. If a mint-relevant field later diverges from it, // the held token (hook response + sessionStorage) is discarded so the admin must re-authorize. const authorizedIdentityRef = React.useRef(undefined); + const dcrClientRef = React.useRef(null); const { startOAuthFlow, @@ -300,7 +301,7 @@ const MCPServerEdit: React.FC = ({ env: values.env, }; }, - onTokenReceived: (token) => { + onTokenReceived: (token, registeredClient) => { if (!token?.access_token) { return; } @@ -319,6 +320,7 @@ const MCPServerEdit: React.FC = ({ return; } + dcrClientRef.current = registeredClient?.dcr_server_url ? registeredClient : null; const current = (allFieldsValue(form).credentials as Record | undefined) ?? {}; const nextCredentials = { ...(preservedAdminCredentials(current) ?? {}), @@ -393,6 +395,9 @@ const MCPServerEdit: React.FC = ({ if (!parsed || parsed.serverId !== mcpServer.server_id) { return; } + // The saved server may still be loading on the first render after the redirect. + // Consume this snapshot only after the matching server can restore it. + window.sessionStorage.removeItem(EDIT_OAUTH_UI_STATE_KEY); if (parsed.formValues) { // Rebuild credentials from the declared app in EITHER the loaded server or the saved snapshot, // then strip minted token material. Merging the two (server under snapshot) before stripping is @@ -426,7 +431,6 @@ const MCPServerEdit: React.FC = ({ } } catch (err) { console.error("Failed to restore MCP edit state", err); - } finally { window.sessionStorage.removeItem(EDIT_OAUTH_UI_STATE_KEY); } }, [form, mcpServer]); @@ -497,6 +501,7 @@ const MCPServerEdit: React.FC = ({ removeToken(mcpServer.server_id, userID); } setTools([]); + dcrClientRef.current = null; resetOAuthFlow(); // The admin-typed app is upstream-scoped config, not minted material, so it survives every // invalidation; only the held token is discarded. Token-shaped keys are excluded by the filter. @@ -720,7 +725,16 @@ const MCPServerEdit: React.FC = ({ return () => subscription.unsubscribe(); }, [form]); + const isOAuthPending = ["authorizing", "exchanging"].includes(oauthStatus); + + const handleCancel = () => { + window.sessionStorage.removeItem(EDIT_OAUTH_UI_STATE_KEY); + resetOAuthFlow(); + onCancel(); + }; + const submitForm = async () => { + if (isOAuthPending) return; const isValid = await form.trigger(mountedPaths(registry) as string[]); if (!isValid) { return; @@ -742,7 +756,8 @@ const MCPServerEdit: React.FC = ({ return; } try { - const built = buildEditServerPayload(values, { + const uiState = { + dcrClient: dcrClientRef.current, mcpServer, logoUrl, costConfig, @@ -752,7 +767,8 @@ const MCPServerEdit: React.FC = ({ toolNameToDisplayName, toolNameToDescription, removeStoredApp, - }); + }; + const built = buildEditServerPayload(values, uiState); if (built.kind !== "ok") { toast.fromError(editPayloadErrorMessage(built)); return; @@ -820,6 +836,7 @@ const MCPServerEdit: React.FC = ({
{ event.preventDefault(); void submitForm(); @@ -1308,10 +1325,12 @@ const MCPServerEdit: React.FC = ({
- - +
@@ -1323,10 +1342,12 @@ const MCPServerEdit: React.FC = ({
- - +
diff --git a/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.test.tsx b/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.test.tsx index 2e809cfd728..2dc404a9218 100644 --- a/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.test.tsx +++ b/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.test.tsx @@ -1,7 +1,7 @@ import { act, renderHook, waitFor } from "@testing-library/react"; import { beforeEach, describe, expect, it, vi } from "vitest"; import * as networking from "@/components/networking"; -import { setSecureItem } from "@/utils/secureStorage"; +import { getSecureItem, setSecureItem } from "@/utils/secureStorage"; import { useMcpOAuthFlow } from "./useMcpOAuthFlow"; vi.mock("@/components/networking", () => ({ @@ -73,7 +73,7 @@ describe("useMcpOAuthFlow reset", () => { await waitFor(() => expect(result.current.status).toBe("success")); expect(result.current.tokenResponse).toEqual(token); - expect(onTokenReceived).toHaveBeenCalledWith(token, expect.objectContaining({ clientId: "client-1" })); + expect(onTokenReceived).toHaveBeenCalledWith(token, expect.objectContaining({ client_id: "client-1" })); act(() => { result.current.reset(); @@ -84,6 +84,51 @@ describe("useMcpOAuthFlow reset", () => { expect(result.current.error).toBeNull(); }); + it("retains client binding from a redirect started before the state format changed", async () => { + seedCompletedRedirect(); + const client = { + client_id: "client-1", + client_secret: "registered-secret", + dcr_issuer: "https://issuer.example.com", + dcr_server_url: "https://server-1.example.com/mcp", + redirect_uris: ["https://app.example.com/ui/mcp/oauth/callback"], + }; + const state = JSON.parse(getSecureItem(FLOW_STATE_KEY)!); + setSecureItem( + FLOW_STATE_KEY, + JSON.stringify({ ...state, clientSecret: client.client_secret, dcrCredentials: client }), + ); + const token = { access_token: "registered-token" }; + vi.mocked(networking.exchangeMcpOAuthToken).mockResolvedValue(token); + const onTokenReceived = vi.fn(); + const { result } = renderFlow(onTokenReceived); + await waitFor(() => expect(result.current.status).toBe("success")); + expect(onTokenReceived).toHaveBeenCalledWith(token, client); + expect(networking.exchangeMcpOAuthToken).toHaveBeenCalledWith( + expect.objectContaining({ clientId: client.client_id, clientSecret: client.client_secret }), + ); + }); + + it("resumes a legacy server-managed flow without exposing a registered client", async () => { + seedCompletedRedirect(); + const state = JSON.parse(getSecureItem(FLOW_STATE_KEY)!); + delete state.clientId; + setSecureItem(FLOW_STATE_KEY, JSON.stringify(state)); + const token = { access_token: "server-managed-token" }; + vi.mocked(networking.exchangeMcpOAuthToken).mockResolvedValue(token); + const onTokenReceived = vi.fn(); + const { result } = renderFlow(onTokenReceived); + await waitFor(() => expect(result.current.status).toBe("success")); + expect(onTokenReceived).toHaveBeenCalledWith(token, undefined); + expect(networking.exchangeMcpOAuthToken).toHaveBeenCalledWith( + expect.objectContaining({ + serverId: "server-1", + clientId: undefined, + clientSecret: undefined, + }), + ); + }); + it("ignores an in-flight exchange result after reset", async () => { const token = { access_token: "stale-token" }; let resolveExchange: (value: typeof token) => void = () => undefined; @@ -137,7 +182,7 @@ describe("useMcpOAuthFlow reset", () => { rerender({ onTokenReceived: onTokenReceived2 }); await waitFor(() => - expect(onTokenReceived2).toHaveBeenCalledWith(token, expect.objectContaining({ clientId: "client-1" })), + expect(onTokenReceived2).toHaveBeenCalledWith(token, expect.objectContaining({ client_id: "client-1" })), ); }); @@ -163,8 +208,8 @@ describe("useMcpOAuthFlow reset", () => { await waitFor(() => expect(result.current.status).toBe("success")); expect(onTokenReceived).toHaveBeenCalledWith(token, { - clientId: "dcr-client-xyz", - clientSecret: "dcr-secret-abc", + client_id: "dcr-client-xyz", + client_secret: "dcr-secret-abc", }); }); @@ -196,13 +241,27 @@ describe("useMcpOAuthFlow reset", () => { ); }); - it("registers a fresh client when no client_id is present (new URL after the derived client is cleared)", async () => { - vi.mocked(networking.cacheTemporaryMcpServer).mockResolvedValue({ server_id: "server-2" }); - vi.mocked(networking.registerMcpOAuthClient).mockResolvedValue({ client_id: "fresh-client" }); - vi.mocked(networking.buildMcpOAuthAuthorizeUrl).mockReturnValue("https://idp.example.com/authorize"); + it.each([ + { method: "none", secret: undefined, issuer: "https://idp.example.com", bound: true }, + { method: "client_secret_basic", secret: "registered-secret", issuer: "https://idp.example.com", bound: true }, + { method: "none", secret: undefined, issuer: undefined, bound: true }, + { method: "none", secret: undefined, issuer: undefined, bound: false }, + ])( + "retains a fresh $method client with issuer $issuer and binding $bound", + async ({ method, secret, issuer, bound }) => { + vi.mocked(networking.cacheTemporaryMcpServer).mockResolvedValue({ server_id: "server-2" }); + const registration = { + client_id: "fresh-client", + client_secret: secret, + token_endpoint_auth_method: method, + dcr_issuer: issuer, + dcr_server_url: bound ? "https://server-2.example.com/mcp" : undefined, + dcr_redirect_uris: ["https://gateway.example.com/callback"], + }; + vi.mocked(networking.registerMcpOAuthClient).mockResolvedValue(registration); + vi.mocked(networking.buildMcpOAuthAuthorizeUrl).mockReturnValue("https://idp.example.com/authorize"); - const { result } = renderHook(() => - useMcpOAuthFlow({ + const options = { accessToken: "admin-token", getCredentials: () => ({}), getTemporaryPayload: () => ({ @@ -212,16 +271,45 @@ describe("useMcpOAuthFlow reset", () => { }), onTokenReceived: vi.fn(), flowSource: "create", - }), - ); + }; + const { result } = renderHook(() => useMcpOAuthFlow(options)); - await act(async () => { - await result.current.startOAuthFlow(); - }); + await act(async () => { + await result.current.startOAuthFlow(); + }); - expect(networking.registerMcpOAuthClient).toHaveBeenCalledTimes(1); - expect(networking.buildMcpOAuthAuthorizeUrl).toHaveBeenCalledWith( - expect.objectContaining({ clientId: "fresh-client" }), - ); - }); + expect(JSON.parse(getSecureItem(FLOW_STATE_KEY)!)).toEqual( + expect.objectContaining({ + client: { + client_id: "fresh-client", + client_secret: secret ?? null, + ...(bound + ? { + token_endpoint_auth_method: method === "client_secret_basic" ? method : null, + dcr_issuer: issuer ?? null, + dcr_server_url: "https://server-2.example.com/mcp", + redirect_uris: ["https://gateway.example.com/callback"], + } + : {}), + }, + }), + ); + const stored = JSON.parse(getSecureItem(FLOW_STATE_KEY)!); + vi.mocked(networking.exchangeMcpOAuthToken).mockResolvedValue({ access_token: "new-token" }); + setSecureItem(RESULT_KEY, JSON.stringify({ state: stored.state, code: "new-code" })); + const resumed = renderHook(() => useMcpOAuthFlow(options)); + await waitFor(() => expect(resumed.result.current.status).toBe("success")); + expect(options.onTokenReceived).toHaveBeenCalledWith({ access_token: "new-token" }, stored.client); + expect(networking.exchangeMcpOAuthToken).toHaveBeenCalledWith( + expect.objectContaining({ + clientId: "fresh-client", + clientSecret: secret, + }), + ); + expect(networking.registerMcpOAuthClient).toHaveBeenCalledTimes(1); + expect(networking.buildMcpOAuthAuthorizeUrl).toHaveBeenCalledWith( + expect.objectContaining({ clientId: "fresh-client" }), + ); + }, + ); }); diff --git a/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx b/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx index 2d6f0344eb9..25e8790e367 100644 --- a/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx +++ b/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx @@ -16,20 +16,48 @@ import { getSecureItem, setSecureItem } from "@/utils/secureStorage"; export type McpOAuthStatus = "idle" | "authorizing" | "exchanging" | "success" | "error"; +export interface McpDcrCredentials { + client_id: string; + client_secret?: string | null; + dcr_issuer?: string | null; + dcr_server_url?: string | null; + token_endpoint_auth_method?: string | null; + redirect_uris?: string[]; +} + +const getRegisteredOAuthClient = ( + registration: (Partial & { dcr_redirect_uris?: string[] }) | undefined, +): McpDcrCredentials | undefined => + registration?.client_id + ? { + client_id: registration.client_id, + client_secret: registration.client_secret ?? null, + ...(registration.dcr_server_url && { + dcr_issuer: registration.dcr_issuer ?? null, + dcr_server_url: registration.dcr_server_url, + token_endpoint_auth_method: + registration.token_endpoint_auth_method === "client_secret_basic" ? "client_secret_basic" : null, + redirect_uris: registration.dcr_redirect_uris, + }), + } + : undefined; + +const oauthClientRequest = (client: McpDcrCredentials | undefined) => ({ + clientId: client?.client_id, + clientSecret: client?.client_secret ?? undefined, +}); + interface UseMcpOAuthFlowOptions { accessToken: string | null; getCredentials: () => | { client_id?: string; - client_secret?: string; + client_secret?: string | null; scopes?: string[]; } | undefined; getTemporaryPayload: () => Record | null; - onTokenReceived: ( - tokenResponse: Record, - registeredClient?: { clientId?: string; clientSecret?: string }, - ) => void; + onTokenReceived: (tokenResponse: Record, registeredClient?: McpDcrCredentials) => void; onBeforeRedirect?: () => void; // Distinguishes which form started the flow (e.g. "create" vs "edit"). Both forms // mount this hook with shared storage keys, so the return handler only processes a @@ -69,6 +97,8 @@ export const useMcpOAuthFlow = ({ codeVerifier: string; clientId?: string; clientSecret?: string; + client?: McpDcrCredentials; + dcrCredentials?: McpDcrCredentials; serverId: string; redirectUri: string; flowSource?: string; @@ -147,7 +177,7 @@ export const useMcpOAuthFlow = ({ throw new Error("Temporary MCP server identifier missing. Please retry."); } - let registeredClient: { clientId?: string; clientSecret?: string } = {}; + let registeredClient: McpDcrCredentials | undefined; const hasPreconfiguredCredentials = Boolean(temporaryPayload.credentials?.client_id); if (!hasPreconfiguredCredentials) { @@ -162,24 +192,22 @@ export const useMcpOAuthFlow = ({ // rejects the registration and the admin authorize dead-ends. redirect_uris: [callbackUrl()], }); - registeredClient = { - clientId: registration?.client_id, - clientSecret: registration?.client_secret, - }; + registeredClient = getRegisteredOAuthClient(registration); } const verifier = generateCodeVerifier(); const challenge = await generateCodeChallenge(verifier); const state = crypto.randomUUID(); - const clientId = registeredClient.clientId || credentials.client_id; + const client = registeredClient ?? getRegisteredOAuthClient(credentials); const scopeString = Array.isArray(credentials.scopes) ? credentials.scopes.filter((s) => s && s.trim().length > 0).join(" ") : undefined; + const { clientId } = oauthClientRequest(client); const authorizeUrl = buildMcpOAuthAuthorizeUrl({ serverId, - clientId: clientId, + clientId, redirectUri: callbackUrl(), state, codeChallenge: challenge, @@ -189,8 +217,7 @@ export const useMcpOAuthFlow = ({ const flowState: StoredFlowState = { state, codeVerifier: verifier, - clientId, - clientSecret: registeredClient.clientSecret || credentials.client_secret, + client, serverId, redirectUri: callbackUrl(), flowSource, @@ -311,12 +338,15 @@ export const useMcpOAuthFlow = ({ throw new Error("Authorization code missing in callback."); } + const client = + flowState.client ?? + flowState.dcrCredentials ?? + getRegisteredOAuthClient({ client_id: flowState.clientId, client_secret: flowState.clientSecret }); setStatus("exchanging"); const token = await exchangeMcpOAuthToken({ serverId: flowState.serverId, code: payload.code, - clientId: flowState.clientId, - clientSecret: flowState.clientSecret, + ...oauthClientRequest(client), codeVerifier: flowState.codeVerifier, redirectUri: flowState.redirectUri, accessToken, @@ -326,7 +356,7 @@ export const useMcpOAuthFlow = ({ return; } - onTokenReceived(token, { clientId: flowState.clientId, clientSecret: flowState.clientSecret }); + onTokenReceived(token, client); setTokenResponse(token); setStatus("success"); setError(null); diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index ae8b99448b0..73336ca1823 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -2273,7 +2273,9 @@ export interface paths { * * - A successful authorization response (``code`` + ``state``), which is * forwarded back to the validated client ``redirect_uri`` with the - * original (un-wrapped) ``state``. + * original (un-wrapped) ``state``, once the RFC 9207 ``iss`` (when the + * authorization server sent one) matches the issuer /authorize sealed + * into the state. * - An error response (``error``[+``error_description``/``error_uri``]), per * RFC 6749 §4.1.2.1. When ``state`` is present and decodes to a trusted * ``redirect_uri``, the error params are propagated back to the client so @@ -37147,6 +37149,10 @@ export interface components { client_private_key_id?: string | null; /** Client Secret */ client_secret?: string | null; + /** Dcr Issuer */ + dcr_issuer?: string | null; + /** Dcr Server Url */ + dcr_server_url?: string | null; /** Id Jag Resource */ id_jag_resource?: string | null; /** Id Jag Resource Token Endpoint */ @@ -53354,6 +53360,7 @@ export interface operations { query?: { code?: string | null; state?: string | null; + iss?: string | null; error?: string | null; error_description?: string | null; error_uri?: string | null;