fix(mcp): bind OAuth clients to their upstream issuer (#37777)

* feat(mcp): advertise the SDK's latest spec revision and validate the RFC 9207 iss

MCPSpecVersion stopped at 2025-06-18 while the pinned SDK negotiates 2025-11-25, and the version LiteLLM puts on its own outbound initialize was a hardcoded historical member. Add the missing revision, name the highest revision we speak once, and pin it to the SDK's LATEST_PROTOCOL_VERSION with a test so the two cannot drift apart silently.

/authorize now seals the issuer it sent the user to into the OAuth state, and /callback holds the authorization response's RFC 9207 iss against it, refusing to forward a code that came back from an authorization server we never sent the user to. An absent iss, an unanchored server row and a state minted before the seal all keep their current behavior.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): keep params, query and fragment significant in issuer comparison

The shared canonicalizer drops all three, so two issuers differing only outside the path compared equal and a response from another tenant's authorization server would have continued through the flow.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): refresh generated API snapshots

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(mcp): cover OAuth client isolation and lint checks

* fix(mcp): preserve registered clients in the existing save payload

* test(mcp): cover optional OAuth registration metadata

* fix(mcp): preserve compatible OAuth registrations across edits

* fix(mcp): retain OAuth state through pending authorization

* fix(mcp): guard pending OAuth at form submission

* fix(mcp): discard canceled OAuth edit snapshots

* test(mcp): preserve complete OAuth registration assertions

* refactor(mcp): construct OAuth credential updates without mutation

* fix(mcp): simplify issuer binding and reject unverifiable callbacks

* fix(mcp): preserve replacement clients and pending redirect bindings

* fix(mcp): retain clients with replacement authentication methods

* fix(mcp): preserve cached clients and pin manual OAuth issuers

---------

Co-authored-by: yucheng <yucheng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-05 10:25:47 -07:00 • committed by GitHub
parent 2e81db03b9
commit 2d82915084
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
24 changed files with 1739 additions and 174 deletions

View file

@ -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):

View file

@ -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

View file

@ -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(

View file

@ -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.

View file

@ -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": [
{

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -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

View file

@ -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)

View file

@ -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")

View file

@ -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")

View file

@ -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"

View file

@ -28,7 +28,10 @@ const oauthHook = vi.hoisted(() => ({
tokenResponse: null as Record<string, unknown> | null,
reset: vi.fn(),
onTokenReceived: null as
| ((token: Record<string, unknown> | null, registeredClient?: { clientId?: string; clientSecret?: string }) => void)
| ((
token: Record<string, unknown> | null,
registeredClient?: { client_id: string; client_secret?: string },
) => void)
| null,
getCredentials: null as (() => Record<string, unknown> | undefined) | null,
getTemporaryPayload: null as (() => Record<string, unknown> | null) | null,
@ -37,7 +40,7 @@ vi.mock("@/hooks/useMcpOAuthFlow", () => ({
useMcpOAuthFlow: (opts: {
onTokenReceived: (
token: Record<string, unknown> | null,
registeredClient?: { clientId?: string; clientSecret?: string },
registeredClient?: { client_id: string; client_secret?: string },
) => void;
getCredentials?: () => Record<string, unknown> | undefined;
getTemporaryPayload?: () => Record<string, unknown> | 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" },
);
});

View file

@ -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<CreateMCPServerProps> = ({
// 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<McpDcrCredentials | null>(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<CreateMCPServerProps> = ({
// 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<string, unknown> | undefined) ?? {};
const nextCredentials = {

View file

@ -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 {

View file

@ -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);
});

View file

@ -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<Record<string, un
};
interface CredentialsEntryInput {
readonly dcrClient?: McpDcrCredentials | null;
readonly oauthFlow?: string;
readonly authType: string | undefined;
readonly credentials: Readonly<Record<string, unknown>> | 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<Record<string, unknown>> } => {
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),

View file

@ -16,21 +16,40 @@ vi.mock("@/components/networking", () => ({
}));
const mockOauth: {
status: string;
tokenResponse: any;
getTemporaryPayload: (() => Record<string, unknown> | null) | null;
onTokenReceived: ((token: Record<string, unknown> | null) => void) | null;
onTokenReceived:
| ((
token: Record<string, unknown> | null,
registeredClient?: {
client_id: string;
client_secret?: string;
dcr_issuer?: string;
dcr_server_url?: string;
},
) => void)
| null;
reset: ReturnType<typeof vi.fn>;
} = { 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<string, unknown> | null;
onTokenReceived?: (token: Record<string, unknown> | null) => void;
onTokenReceived?: (
token: Record<string, unknown> | 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(
<MCPServerEdit
mcpServer={{ ...interactiveOAuthServer }}
accessToken="access-token"
onCancel={vi.fn()}
onSuccess={vi.fn()}
availableAccessGroups={[]}
/>,
);
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(<MCPServerEdit {...props} mcpServer={{ ...interactiveOAuthServer, server_id: "", url: "" }} />);
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(<MCPServerEdit {...props} mcpServer={interactiveOAuthServer} />);
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(
<MCPServerEdit {...props} mcpServer={{ ...interactiveOAuthServer, server_id: "", url: "" }} />,
);
rerender(<MCPServerEdit {...props} mcpServer={interactiveOAuthServer} />);
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

View file

@ -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<MCPServerEditProps> = ({
// 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<string | undefined>(undefined);
const dcrClientRef = React.useRef<McpDcrCredentials | null>(null);
const {
startOAuthFlow,
@ -300,7 +301,7 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
env: values.env,
};
},
onTokenReceived: (token) => {
onTokenReceived: (token, registeredClient) => {
if (!token?.access_token) {
return;
}
@ -319,6 +320,7 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
return;
}
dcrClientRef.current = registeredClient?.dcr_server_url ? registeredClient : null;
const current = (allFieldsValue(form).credentials as Record<string, unknown> | undefined) ?? {};
const nextCredentials = {
...(preservedAdminCredentials(current) ?? {}),
@ -393,6 +395,9 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
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<MCPServerEditProps> = ({
}
} 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<MCPServerEditProps> = ({
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<MCPServerEditProps> = ({
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<MCPServerEditProps> = ({
return;
}
try {
const built = buildEditServerPayload(values, {
const uiState = {
dcrClient: dcrClientRef.current,
mcpServer,
logoUrl,
costConfig,
@ -752,7 +767,8 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
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<MCPServerEditProps> = ({
<FormProvider {...form}>
<MountedFormProvider value={{ control: form.control, registry }}>
<form
aria-label="Edit MCP server"
onSubmit={(event) => {
event.preventDefault();
void submitForm();
@ -1308,10 +1325,12 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
</div>
<div className="flex justify-end gap-2">
<Button variant="outline" onClick={onCancel}>
<Button variant="outline" onClick={handleCancel}>
Cancel
</Button>
<Button type="submit">Save Changes</Button>
<Button type="submit" disabled={isOAuthPending}>
Save Changes
</Button>
</div>
</form>
</MountedFormProvider>
@ -1323,10 +1342,12 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
<MCPServerCostConfig value={costConfig} onChange={setCostConfig} tools={tools} disabled={isLoadingTools} />
<div className="flex justify-end gap-2">
<Button variant="outline" onClick={onCancel}>
<Button variant="outline" onClick={handleCancel}>
Cancel
</Button>
<Button onClick={() => void submitForm()}>Save Changes</Button>
<Button onClick={() => void submitForm()} disabled={isOAuthPending}>
Save Changes
</Button>
</div>
</div>
</TabsContent>

View file

@ -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" }),
);
},
);
});

View file

@ -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<McpDcrCredentials> & { 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<string, any> | null;
onTokenReceived: (
tokenResponse: Record<string, any>,
registeredClient?: { clientId?: string; clientSecret?: string },
) => void;
onTokenReceived: (tokenResponse: Record<string, any>, 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);

View file

@ -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;