mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge pull request #32652 from BerriAI/litellm_mcp_stale_token_invalidation
fix(mcp): invalidate a browser-authorized upstream token when a mint-relevant field changes
This commit is contained in:
commit
eec948dcb8
14 changed files with 1166 additions and 137 deletions
|
|
@ -3,7 +3,7 @@ import binascii
|
|||
import hashlib
|
||||
import json
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import TYPE_CHECKING, Any, Dict, Iterable, List, Optional, Set, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Awaitable, Callable, Dict, Iterable, List, Optional, Set, Union, cast
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
|
|
@ -558,7 +558,11 @@ async def delete_mcp_server_from_virtualkey():
|
|||
pass
|
||||
|
||||
|
||||
async def delete_mcp_server(prisma_client: PrismaClient, server_id: str) -> Optional[LiteLLM_MCPServerTable]:
|
||||
async def delete_mcp_server(
|
||||
prisma_client: PrismaClient,
|
||||
server_id: str,
|
||||
invalidate_token_cache: Optional[Callable[[str, str], Awaitable[None]]] = None,
|
||||
) -> Optional[LiteLLM_MCPServerTable]:
|
||||
"""
|
||||
Delete the mcp server from the db by server_id
|
||||
|
||||
|
|
@ -569,6 +573,12 @@ async def delete_mcp_server(prisma_client: PrismaClient, server_id: str) -> Opti
|
|||
caller-visible error. Each table is cleaned independently so a failure on one
|
||||
still attempts the other.
|
||||
|
||||
Each enumerated credential row's user also gets their cached per-user token
|
||||
invalidated (legacy cache + v2 store, via invalidate_token_cache, defaulting
|
||||
to the manager's shared invalidation): the caches are keyed by
|
||||
(user_id, server_id), so without this a re-created server reusing the same
|
||||
server_id would serve tokens minted for the deleted server until TTL.
|
||||
|
||||
Returns the deleted mcp server record if it exists, otherwise None
|
||||
"""
|
||||
deleted_server = await MCPServerRepository(prisma_client).table.delete(
|
||||
|
|
@ -577,6 +587,18 @@ async def delete_mcp_server(prisma_client: PrismaClient, server_id: str) -> Opti
|
|||
},
|
||||
)
|
||||
if deleted_server is not None:
|
||||
credential_user_ids: List[str] = []
|
||||
try:
|
||||
credential_rows = await prisma_client.db.litellm_mcpusercredentials.find_many(
|
||||
where={"server_id": server_id}
|
||||
)
|
||||
credential_user_ids = [row.user_id for row in credential_rows]
|
||||
except Exception as e: # noqa: BLE001 - enumeration is best-effort; cached tokens expire by TTL
|
||||
verbose_proxy_logger.warning(
|
||||
"MCP server %s deleted but per-user credential enumeration failed; cached tokens expire by TTL: %s",
|
||||
server_id,
|
||||
e,
|
||||
)
|
||||
for model, label in (
|
||||
(prisma_client.db.litellm_mcpusercredentials, "credential"),
|
||||
(prisma_client.db.litellm_mcpuserenvvars, "env var"),
|
||||
|
|
@ -591,6 +613,15 @@ async def delete_mcp_server(prisma_client: PrismaClient, server_id: str) -> Opti
|
|||
label,
|
||||
e,
|
||||
)
|
||||
if credential_user_ids:
|
||||
if invalidate_token_cache is None:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
invalidate_token_cache = global_mcp_server_manager.invalidate_user_oauth_token_cache
|
||||
for user_id in credential_user_ids:
|
||||
await invalidate_token_cache(user_id, server_id)
|
||||
return deleted_server
|
||||
|
||||
|
||||
|
|
@ -1070,6 +1101,103 @@ async def list_user_oauth_credentials(
|
|||
return results
|
||||
|
||||
|
||||
def _decrypted_credential_field(creds: Dict[str, object], field: str) -> object:
|
||||
"""Return one credential field decrypted with the global salt key; non-string and legacy
|
||||
plaintext values come back unchanged (decrypt_value_helper returns the original on failure)."""
|
||||
value = creds.get(field)
|
||||
if not isinstance(value, str):
|
||||
return value
|
||||
return decrypt_value_helper(
|
||||
value=value,
|
||||
key=field,
|
||||
exception_type="debug",
|
||||
return_original_value=True,
|
||||
)
|
||||
|
||||
|
||||
def mcp_oauth_token_identity(server: object) -> tuple[object, ...]:
|
||||
"""The upstream-OAuth-token-determining fields of an MCP server: the resource/audience (url, or
|
||||
spec_path for OpenAPI servers), the OAuth mode/grant (auth_type, oauth2_flow), the
|
||||
authorization-server endpoints, and the OAuth client + scopes. Mirrors the dashboard's
|
||||
getOAuthAuthorizationIdentity. When any of these change on a server update, previously stored
|
||||
per-user tokens were minted for the old identity and are stale. Excludes transport and
|
||||
delegate_auth_to_upstream, which do not affect what token is minted (RFC 8707/8693).
|
||||
|
||||
client_id/client_secret are compared decrypted: stored values are NaCl-encrypted with a fresh
|
||||
nonce on every write, so comparing ciphertext would flag every routine save as an identity
|
||||
change and purge tokens that are still valid."""
|
||||
creds = getattr(server, "credentials", None)
|
||||
if isinstance(creds, str):
|
||||
try:
|
||||
parsed: object = json.loads(creds)
|
||||
except ValueError:
|
||||
parsed = None
|
||||
else:
|
||||
parsed = creds
|
||||
creds_dict: Dict[str, object] = parsed if isinstance(parsed, dict) else {}
|
||||
return (
|
||||
getattr(server, "url", None),
|
||||
getattr(server, "spec_path", None),
|
||||
getattr(server, "auth_type", None),
|
||||
getattr(server, "oauth2_flow", None),
|
||||
getattr(server, "authorization_url", None),
|
||||
getattr(server, "token_url", None),
|
||||
getattr(server, "registration_url", None),
|
||||
_decrypted_credential_field(creds_dict, "client_id"),
|
||||
_decrypted_credential_field(creds_dict, "client_secret"),
|
||||
creds_dict.get("scopes"),
|
||||
)
|
||||
|
||||
|
||||
async def purge_user_oauth_credentials_for_server(
|
||||
prisma_client: PrismaClient,
|
||||
server_id: str,
|
||||
invalidate_token_cache: Optional[Callable[[str, str], Awaitable[None]]] = None,
|
||||
) -> int:
|
||||
"""Delete every stored per-user OAuth token for a server and invalidate each user's cached
|
||||
token everywhere it can be served from (the legacy per-user token cache and the v2 per-user OAuth
|
||||
token store), so no user keeps a token minted for a superseded configuration. Called when a server
|
||||
update changes a mint-relevant field (see mcp_oauth_token_identity). Returns the number of rows
|
||||
removed.
|
||||
|
||||
LiteLLM_MCPUserCredentials also stores BYOK API keys in the same column; only rows whose payload
|
||||
decodes as an OAuth2 credential (see _decode_oauth_payload) are deleted, because a config change
|
||||
only invalidates minted tokens, never a user's own stored key. Rows are therefore deleted per
|
||||
(user_id, server_id) pair rather than by a blanket server_id filter. An OAuth row inserted while
|
||||
the purge runs for a user not yet enumerated survives; a re-auth completing in the window for an
|
||||
already-enumerated user is deleted along with the stale row (the pair delete cannot tell them
|
||||
apart), which costs that user one extra re-auth and nothing else.
|
||||
|
||||
invalidate_token_cache is injectable for tests; it defaults to the manager's shared
|
||||
invalidate_user_oauth_token_cache, the single invalidation point for per-user tokens."""
|
||||
repo = MCPUserCredentialsRepository(prisma_client)
|
||||
rows = await repo.table.find_many(where={"server_id": server_id})
|
||||
oauth_rows = [row for row in rows if _decode_oauth_payload(row.credential_b64) is not None]
|
||||
if not oauth_rows:
|
||||
return 0
|
||||
deleted_count = await repo.table.delete_many(
|
||||
where={"server_id": server_id, "user_id": {"in": [row.user_id for row in oauth_rows]}}
|
||||
)
|
||||
if invalidate_token_cache is None:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
invalidate_token_cache = global_mcp_server_manager.invalidate_user_oauth_token_cache
|
||||
|
||||
for row in oauth_rows:
|
||||
await invalidate_token_cache(row.user_id, server_id)
|
||||
if deleted_count != len(oauth_rows):
|
||||
verbose_proxy_logger.warning(
|
||||
"MCP server %s: purge removed %d OAuth credential row(s) but %d were enumerated; "
|
||||
"row(s) were deleted concurrently during the purge",
|
||||
server_id,
|
||||
deleted_count,
|
||||
len(oauth_rows),
|
||||
)
|
||||
return deleted_count
|
||||
|
||||
|
||||
async def refresh_user_oauth_token(
|
||||
prisma_client: PrismaClient,
|
||||
user_id: str,
|
||||
|
|
|
|||
|
|
@ -57,7 +57,11 @@ from litellm.proxy._experimental.mcp_server.elicitation_handler import (
|
|||
from litellm.proxy._experimental.mcp_server.sampling_handler import (
|
||||
MCP_SAMPLING_AVAILABLE,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import resolve_mcp_auth
|
||||
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
|
||||
MCPPerUserTokenCache,
|
||||
mcp_per_user_token_cache,
|
||||
resolve_mcp_auth,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
|
||||
Error,
|
||||
Ok,
|
||||
|
|
@ -799,10 +803,12 @@ class MCPServerManager:
|
|||
self,
|
||||
cred_provider: Optional[UpstreamCredentialProvider] = None,
|
||||
per_user_oauth_token_store: Optional[InvalidatableOAuthTokenStore] = None,
|
||||
per_user_token_cache: Optional[MCPPerUserTokenCache] = None,
|
||||
):
|
||||
self._per_user_oauth_token_store = per_user_oauth_token_store or LazyPerUserOAuthTokenStore(
|
||||
self.get_mcp_server_by_id
|
||||
)
|
||||
self._per_user_token_cache = per_user_token_cache or mcp_per_user_token_cache
|
||||
self._cred_provider = cred_provider or UpstreamCredentialProvider(
|
||||
oauth_token_store=self._per_user_oauth_token_store,
|
||||
token_exchanger=build_token_exchanger(),
|
||||
|
|
@ -4053,10 +4059,13 @@ class MCPServerManager:
|
|||
return await self._cred_provider.has_user_token(to_subject(user_api_key_auth, None), spec)
|
||||
|
||||
async def invalidate_user_oauth_token_cache(self, user_id: str, server_id: str) -> None:
|
||||
"""Drop the v2 chain's cached token for ``(user_id, server_id)`` after the credential row
|
||||
changes (re-auth, revoke), so the next resolve reads the new row instead of serving the
|
||||
replaced token until its cache TTL. Best-effort: a cache-drop failure is logged, never
|
||||
raised, because the DB write already succeeded and the TTL remains the backstop.
|
||||
"""Drop every cached token for ``(user_id, server_id)`` after the credential row changes
|
||||
(re-auth, revoke, config-change purge): the v2 chain's cache and the legacy per-user token
|
||||
cache, so the next resolve reads the new row instead of serving the replaced token until its
|
||||
cache TTL, whichever path resolves it. This is the single invalidation point for per-user
|
||||
OAuth tokens; callers must not evict individual caches directly. Best-effort: a cache-drop
|
||||
failure is logged, never raised, because the DB write already succeeded and the TTL remains
|
||||
the backstop.
|
||||
"""
|
||||
try:
|
||||
await self._per_user_oauth_token_store.invalidate(user_id, server_id)
|
||||
|
|
@ -4064,6 +4073,12 @@ class MCPServerManager:
|
|||
verbose_logger.warning(
|
||||
"Failed to invalidate cached MCP OAuth token for user=%s server=%s: %s", user_id, server_id, exc
|
||||
)
|
||||
try:
|
||||
await self._per_user_token_cache.delete(user_id, server_id)
|
||||
except Exception as exc: # noqa: BLE001 - cache drop is best-effort; TTL is the backstop
|
||||
verbose_logger.warning(
|
||||
"Failed to drop legacy cached MCP OAuth token for user=%s server=%s: %s", user_id, server_id, exc
|
||||
)
|
||||
|
||||
async def _resolve_oauth2_headers_for_tool_call(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -125,7 +125,9 @@ if MCP_AVAILABLE:
|
|||
get_user_env_vars_bulk,
|
||||
get_user_oauth_credential,
|
||||
list_user_oauth_credentials,
|
||||
mcp_oauth_token_identity,
|
||||
merge_user_env_vars,
|
||||
purge_user_oauth_credentials_for_server,
|
||||
reject_mcp_server,
|
||||
store_user_credential,
|
||||
store_user_oauth_credential,
|
||||
|
|
@ -2318,6 +2320,19 @@ if MCP_AVAILABLE:
|
|||
},
|
||||
)
|
||||
|
||||
# Snapshot the pre-update identity so we can detect a mint-relevant change below. The read is
|
||||
# advisory (it only feeds the stale-token purge decision), so a failure skips the purge with a
|
||||
# warning instead of failing the edit, whose primary job is the update itself.
|
||||
try:
|
||||
old_server_record = await get_mcp_server(prisma_client, payload.server_id)
|
||||
except Exception as exc: # noqa: BLE001 - advisory read; invalidation is best-effort end-to-end
|
||||
verbose_logger.warning(
|
||||
"MCP server %s: could not snapshot the pre-update record; skipping the stale-token check: %s",
|
||||
payload.server_id,
|
||||
exc,
|
||||
)
|
||||
old_server_record = None
|
||||
|
||||
# try to update the mcp server
|
||||
mcp_server_record_updated = await update_mcp_server(
|
||||
prisma_client,
|
||||
|
|
@ -2336,6 +2351,30 @@ if MCP_AVAILABLE:
|
|||
# Ensure registry is up to date by reloading from database
|
||||
await global_mcp_server_manager.reload_servers_from_database()
|
||||
|
||||
# If a field that determines which upstream OAuth token gets minted changed (url/audience, OAuth
|
||||
# mode/grant, authorization-server endpoints, or the OAuth client + scopes), every stored per-user
|
||||
# token was minted for the old configuration and is stale. Purge them (DB + cache) so the next
|
||||
# tool call re-authorizes instead of forwarding a token for a resource/AS/client that no longer
|
||||
# matches. Best-effort: a purge failure must not fail the update, whose primary job already
|
||||
# succeeded.
|
||||
if old_server_record is not None and mcp_oauth_token_identity(old_server_record) != mcp_oauth_token_identity(
|
||||
mcp_server_record_updated
|
||||
):
|
||||
try:
|
||||
purged = await purge_user_oauth_credentials_for_server(prisma_client, payload.server_id)
|
||||
if purged:
|
||||
verbose_logger.info(
|
||||
"MCP server %s: purged %d stale per-user OAuth token(s) after a mint-relevant config change",
|
||||
payload.server_id,
|
||||
purged,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 - purge is best-effort; the server update already succeeded
|
||||
verbose_logger.warning(
|
||||
"MCP server %s: failed to purge stale per-user OAuth tokens after config change: %s",
|
||||
payload.server_id,
|
||||
exc,
|
||||
)
|
||||
|
||||
# TODO: Enterprise: Finish audit log trail
|
||||
if litellm.store_audit_logs:
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -2533,9 +2533,10 @@ class StandardLoggingMCPToolCall(TypedDict, total=False):
|
|||
|
||||
mcp_server_resource: Optional[str]
|
||||
"""
|
||||
The upstream MCP server resource identifier (scheme + host + path) the tool call was
|
||||
forwarded to. Redacted for logging: userinfo, query string, and fragment are stripped so an
|
||||
upstream URL carrying an embedded token or secret query parameter never reaches log metadata.
|
||||
The origin (scheme + host + port) of the upstream MCP server the tool call was forwarded
|
||||
to. Redacted for logging: userinfo, the path, the query string, and the fragment are all
|
||||
stripped, because hosted MCP servers routinely embed the credential in the URL path and
|
||||
this value is readable by callers via request logs.
|
||||
Records which upstream received a relayed request; never a credential.
|
||||
"""
|
||||
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ keeps a plain-base64 fallback on read so existing rows continue to work.
|
|||
import base64
|
||||
import json
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
|
@ -63,6 +64,264 @@ def _legacy_row(payload: str):
|
|||
return row
|
||||
|
||||
|
||||
def _identity_server(**overrides):
|
||||
base = dict(
|
||||
url="https://up.example.com/mcp",
|
||||
auth_type="oauth2",
|
||||
oauth2_flow="authorization_code",
|
||||
authorization_url="https://idp.example.com/authorize",
|
||||
token_url="https://idp.example.com/token",
|
||||
registration_url="https://idp.example.com/register",
|
||||
credentials={"client_id": "cid", "client_secret": "csec", "scopes": ["a"]},
|
||||
server_name="srv",
|
||||
description="d",
|
||||
)
|
||||
base.update(overrides)
|
||||
return SimpleNamespace(**base)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"overrides",
|
||||
[
|
||||
{"url": "https://other.example.com/mcp"},
|
||||
{"spec_path": "https://up.example.com/openapi.json"},
|
||||
{"auth_type": "oauth_delegate"},
|
||||
{"oauth2_flow": "client_credentials"},
|
||||
{"authorization_url": "https://other.example.com/authorize"},
|
||||
{"token_url": "https://other.example.com/token"},
|
||||
{"registration_url": "https://other.example.com/register"},
|
||||
{"credentials": {"client_id": "new", "client_secret": "csec", "scopes": ["a"]}},
|
||||
{"credentials": {"client_id": "cid", "client_secret": "rotated", "scopes": ["a"]}},
|
||||
{"credentials": {"client_id": "cid", "client_secret": "csec", "scopes": ["b"]}},
|
||||
],
|
||||
)
|
||||
def test_mcp_oauth_token_identity_changes_on_mint_relevant_fields(overrides):
|
||||
from litellm.proxy._experimental.mcp_server.db import mcp_oauth_token_identity
|
||||
|
||||
assert mcp_oauth_token_identity(_identity_server()) != mcp_oauth_token_identity(_identity_server(**overrides))
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"overrides",
|
||||
[
|
||||
{"server_name": "renamed"},
|
||||
{"description": "changed"},
|
||||
],
|
||||
)
|
||||
def test_mcp_oauth_token_identity_stable_on_non_mint_fields(overrides):
|
||||
from litellm.proxy._experimental.mcp_server.db import mcp_oauth_token_identity
|
||||
|
||||
assert mcp_oauth_token_identity(_identity_server()) == mcp_oauth_token_identity(_identity_server(**overrides))
|
||||
|
||||
|
||||
def _encrypted_creds_json(client_id: str = "cid", client_secret: str = "csec") -> str:
|
||||
from litellm.proxy._experimental.mcp_server.db import encrypt_credentials
|
||||
|
||||
encrypted = encrypt_credentials(
|
||||
credentials={"client_id": client_id, "client_secret": client_secret, "scopes": ["a"]},
|
||||
encryption_key=None,
|
||||
)
|
||||
return json.dumps(encrypted)
|
||||
|
||||
|
||||
def test_mcp_oauth_token_identity_stable_across_reencryption():
|
||||
"""Stored client_id/client_secret are NaCl-encrypted with a fresh nonce on every write, so two
|
||||
saves of the SAME plaintext produce different ciphertext. The identity must compare decrypted
|
||||
values; comparing ciphertext would flag every routine save as a mint-relevant change and purge
|
||||
per-user tokens that are still valid."""
|
||||
from litellm.proxy._experimental.mcp_server.db import mcp_oauth_token_identity
|
||||
|
||||
first = _encrypted_creds_json()
|
||||
second = _encrypted_creds_json()
|
||||
assert first != second
|
||||
|
||||
assert mcp_oauth_token_identity(_identity_server(credentials=first)) == mcp_oauth_token_identity(
|
||||
_identity_server(credentials=second)
|
||||
)
|
||||
|
||||
|
||||
def test_mcp_oauth_token_identity_detects_change_under_encryption():
|
||||
from litellm.proxy._experimental.mcp_server.db import mcp_oauth_token_identity
|
||||
|
||||
unchanged = _identity_server(credentials=_encrypted_creds_json())
|
||||
changed = _identity_server(credentials=_encrypted_creds_json(client_id="other"))
|
||||
assert mcp_oauth_token_identity(unchanged) != mcp_oauth_token_identity(changed)
|
||||
|
||||
|
||||
def _oauth_row(user_id: str, server_id: str = "srv-1"):
|
||||
"""A stored per-user OAuth token row (payload tagged type=oauth2, legacy plain-base64 encoding)."""
|
||||
row = _legacy_row(json.dumps({"type": "oauth2", "access_token": "tok-" + user_id}))
|
||||
row.user_id = user_id
|
||||
row.server_id = server_id
|
||||
return row
|
||||
|
||||
|
||||
def _byok_row(user_id: str, server_id: str = "srv-1"):
|
||||
"""A stored BYOK API key row: the same column, but the payload is a plain string, not OAuth JSON."""
|
||||
row = _legacy_row("sk-byok-" + user_id)
|
||||
row.user_id = user_id
|
||||
row.server_id = server_id
|
||||
return row
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_purge_user_oauth_credentials_for_server_invalidates_each_user():
|
||||
"""The purge must route each (user, server) row through the invalidator exactly once."""
|
||||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_oauth_row("alice"), _oauth_row("bob")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=2)
|
||||
|
||||
invalidations = []
|
||||
|
||||
async def record_invalidation(user_id: str, server_id: str) -> None:
|
||||
invalidations.append((user_id, server_id))
|
||||
|
||||
purged = await purge_user_oauth_credentials_for_server(prisma, "srv-1", invalidate_token_cache=record_invalidation)
|
||||
|
||||
assert purged == 2
|
||||
prisma.db.litellm_mcpusercredentials.delete_many.assert_awaited_once_with(
|
||||
where={"server_id": "srv-1", "user_id": {"in": ["alice", "bob"]}}
|
||||
)
|
||||
assert set(invalidations) == {("alice", "srv-1"), ("bob", "srv-1")}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_purge_user_oauth_credentials_for_server_spares_byok_rows():
|
||||
"""Regression: the purge used to delete_many on server_id alone, wiping BYOK API keys that share
|
||||
the LiteLLM_MCPUserCredentials table. Only rows holding an OAuth2 payload may be deleted (one
|
||||
batched query filtered to their user_ids), and only their users' token caches invalidated."""
|
||||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_byok_row("carol"), _oauth_row("alice")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=1)
|
||||
|
||||
invalidations = []
|
||||
|
||||
async def record_invalidation(user_id: str, server_id: str) -> None:
|
||||
invalidations.append((user_id, server_id))
|
||||
|
||||
purged = await purge_user_oauth_credentials_for_server(prisma, "srv-1", invalidate_token_cache=record_invalidation)
|
||||
|
||||
assert purged == 1
|
||||
prisma.db.litellm_mcpusercredentials.delete_many.assert_awaited_once_with(
|
||||
where={"server_id": "srv-1", "user_id": {"in": ["alice"]}}
|
||||
)
|
||||
assert invalidations == [("alice", "srv-1")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_purge_user_oauth_credentials_for_server_all_byok_is_noop():
|
||||
"""An api_key (BYOK-only) server whose identity tuple changes (e.g. its url) must purge nothing."""
|
||||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_byok_row("carol"), _byok_row("dave")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock()
|
||||
|
||||
purged = await purge_user_oauth_credentials_for_server(prisma, "srv-1")
|
||||
|
||||
assert purged == 0
|
||||
prisma.db.litellm_mcpusercredentials.delete_many.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_purge_user_oauth_credentials_for_server_defaults_to_manager_invalidator(monkeypatch):
|
||||
"""When no invalidator is injected, the purge must resolve to the manager's shared
|
||||
invalidate_user_oauth_token_cache, the single point covering both the legacy per-user token cache
|
||||
and the v2 per-user OAuth token store; a wrong or no-op default silently leaves every cache
|
||||
serving tokens minted for the superseded config."""
|
||||
from litellm.proxy._experimental.mcp_server import mcp_server_manager
|
||||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_oauth_row("alice")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=1)
|
||||
|
||||
shared_invalidator = AsyncMock()
|
||||
monkeypatch.setattr(
|
||||
mcp_server_manager.global_mcp_server_manager,
|
||||
"invalidate_user_oauth_token_cache",
|
||||
shared_invalidator,
|
||||
)
|
||||
|
||||
purged = await purge_user_oauth_credentials_for_server(prisma, "srv-1")
|
||||
|
||||
assert purged == 1
|
||||
shared_invalidator.assert_awaited_once_with("alice", "srv-1")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_purge_user_oauth_credentials_for_server_logs_raced_rows(monkeypatch):
|
||||
from litellm.proxy._experimental.mcp_server import db as db_module
|
||||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_oauth_row("alice")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=0)
|
||||
warning = MagicMock()
|
||||
monkeypatch.setattr(db_module.verbose_proxy_logger, "warning", warning)
|
||||
|
||||
purged = await purge_user_oauth_credentials_for_server(prisma, "srv-1", invalidate_token_cache=AsyncMock())
|
||||
|
||||
assert purged == 0
|
||||
warning.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_mcp_server_invalidates_cached_tokens_for_enumerated_users():
|
||||
"""Deleting a server must invalidate each enumerated user's cached per-user token: the caches are
|
||||
keyed by (user_id, server_id), so a re-created server reusing the same server_id would otherwise
|
||||
serve tokens minted for the deleted server until TTL."""
|
||||
from litellm.proxy._experimental.mcp_server.db import delete_mcp_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=MagicMock(server_id="srv-1"))
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_oauth_row("alice"), _byok_row("bob")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=2)
|
||||
prisma.db.litellm_mcpuserenvvars.delete_many = AsyncMock(return_value=0)
|
||||
|
||||
invalidations = []
|
||||
|
||||
async def record_invalidation(user_id: str, server_id: str) -> None:
|
||||
invalidations.append((user_id, server_id))
|
||||
|
||||
deleted = await delete_mcp_server(prisma, "srv-1", invalidate_token_cache=record_invalidation)
|
||||
|
||||
assert deleted is not None
|
||||
assert set(invalidations) == {("alice", "srv-1"), ("bob", "srv-1")}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_mcp_server_returns_none_without_cleanup_when_server_missing():
|
||||
from litellm.proxy._experimental.mcp_server.db import delete_mcp_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=None)
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock()
|
||||
|
||||
deleted = await delete_mcp_server(prisma, "srv-1", invalidate_token_cache=AsyncMock())
|
||||
|
||||
assert deleted is None
|
||||
prisma.db.litellm_mcpusercredentials.find_many.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_purge_user_oauth_credentials_for_server_noop_when_empty():
|
||||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock()
|
||||
|
||||
purged = await purge_user_oauth_credentials_for_server(prisma, "srv-1")
|
||||
|
||||
assert purged == 0
|
||||
prisma.db.litellm_mcpusercredentials.delete_many.assert_not_awaited()
|
||||
|
||||
|
||||
def _stored_value(prisma) -> str:
|
||||
"""Pull the credential_b64 value passed to the most recent upsert call."""
|
||||
call = prisma.db.litellm_mcpusercredentials.upsert.call_args
|
||||
|
|
@ -136,9 +395,7 @@ async def test_store_user_oauth_credential_does_not_persist_plaintext():
|
|||
access_token = "ya29.a0AfH6SMBverysecretaccesstoken"
|
||||
prisma = _make_prisma_with_existing(row=None)
|
||||
|
||||
await store_user_oauth_credential(
|
||||
prisma, "alice", "srv-1", access_token, refresh_token="rfr-xyz"
|
||||
)
|
||||
await store_user_oauth_credential(prisma, "alice", "srv-1", access_token, refresh_token="rfr-xyz")
|
||||
|
||||
stored = _stored_value(prisma)
|
||||
try:
|
||||
|
|
@ -221,9 +478,7 @@ async def test_byok_guard_rejects_overwriting_encrypted_byok():
|
|||
|
||||
encrypted_row = MagicMock()
|
||||
encrypted_row.credential_b64 = _stored_value(prisma)
|
||||
prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(
|
||||
return_value=encrypted_row
|
||||
)
|
||||
prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=encrypted_row)
|
||||
|
||||
with pytest.raises(ValueError, match="could not be verified as an OAuth2"):
|
||||
await store_user_oauth_credential(prisma, "alice", "srv-1", "tok")
|
||||
|
|
@ -265,18 +520,14 @@ async def test_list_oauth_credentials_filters_byok_and_returns_payloads():
|
|||
"connected_at": "2024-01-01T00:00:00Z",
|
||||
}
|
||||
legacy_row = MagicMock()
|
||||
legacy_row.credential_b64 = base64.urlsafe_b64encode(
|
||||
json.dumps(legacy_payload).encode()
|
||||
).decode()
|
||||
legacy_row.credential_b64 = base64.urlsafe_b64encode(json.dumps(legacy_payload).encode()).decode()
|
||||
legacy_row.server_id = "srv-legacy"
|
||||
|
||||
byok_row = MagicMock()
|
||||
byok_row.credential_b64 = base64.urlsafe_b64encode(b"plain-byok-key").decode()
|
||||
byok_row.server_id = "srv-byok"
|
||||
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(
|
||||
return_value=[encrypted_row, legacy_row, byok_row]
|
||||
)
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[encrypted_row, legacy_row, byok_row])
|
||||
|
||||
results = await list_user_oauth_credentials(prisma, "alice")
|
||||
|
||||
|
|
@ -326,9 +577,7 @@ async def test_rotate_re_encrypts_byok_with_new_key(monkeypatch):
|
|||
prisma.db.litellm_mcpusercredentials.update = AsyncMock()
|
||||
|
||||
new_master_key = "rotated-salt-key-9999-9999-9999-9999"
|
||||
await rotate_mcp_user_credentials_master_key(
|
||||
prisma_client=prisma, new_master_key=new_master_key
|
||||
)
|
||||
await rotate_mcp_user_credentials_master_key(prisma_client=prisma, new_master_key=new_master_key)
|
||||
|
||||
update_call = prisma.db.litellm_mcpusercredentials.update.call_args
|
||||
new_stored = update_call.kwargs["data"]["credential_b64"]
|
||||
|
|
@ -356,19 +605,13 @@ async def test_rotate_migrates_legacy_plaintext_rows(monkeypatch):
|
|||
legacy_row.user_id = "alice"
|
||||
legacy_row.server_id = "srv-legacy"
|
||||
legacy_row.credential_b64 = base64.urlsafe_b64encode(b"legacy-plain").decode()
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(
|
||||
return_value=[legacy_row]
|
||||
)
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[legacy_row])
|
||||
prisma.db.litellm_mcpusercredentials.update = AsyncMock()
|
||||
|
||||
new_key = "another-rotation-key-aaaa-bbbb-cccc-dddd"
|
||||
await rotate_mcp_user_credentials_master_key(
|
||||
prisma_client=prisma, new_master_key=new_key
|
||||
)
|
||||
await rotate_mcp_user_credentials_master_key(prisma_client=prisma, new_master_key=new_key)
|
||||
|
||||
new_stored = prisma.db.litellm_mcpusercredentials.update.call_args.kwargs["data"][
|
||||
"credential_b64"
|
||||
]
|
||||
new_stored = prisma.db.litellm_mcpusercredentials.update.call_args.kwargs["data"]["credential_b64"]
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", new_key)
|
||||
assert (
|
||||
decrypt_value_helper(
|
||||
|
|
@ -396,14 +639,10 @@ async def test_rotate_skips_undecodable_rows():
|
|||
good_row.server_id = "srv-ok"
|
||||
good_row.credential_b64 = base64.urlsafe_b64encode(b"good-byok").decode()
|
||||
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(
|
||||
return_value=[bad_row, good_row]
|
||||
)
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[bad_row, good_row])
|
||||
prisma.db.litellm_mcpusercredentials.update = AsyncMock()
|
||||
|
||||
await rotate_mcp_user_credentials_master_key(
|
||||
prisma_client=prisma, new_master_key="new-key-xxxx"
|
||||
)
|
||||
await rotate_mcp_user_credentials_master_key(prisma_client=prisma, new_master_key="new-key-xxxx")
|
||||
|
||||
# Only one update call — the good row.
|
||||
assert prisma.db.litellm_mcpusercredentials.update.call_count == 1
|
||||
|
|
@ -419,9 +658,7 @@ def _oauth_cred(access_token="at-live", refresh_token=None, expires_in_seconds=N
|
|||
if refresh_token is not None:
|
||||
cred["refresh_token"] = refresh_token
|
||||
if expires_in_seconds is not None:
|
||||
cred["expires_at"] = (
|
||||
datetime.now(timezone.utc) + timedelta(seconds=expires_in_seconds)
|
||||
).isoformat()
|
||||
cred["expires_at"] = (datetime.now(timezone.utc) + timedelta(seconds=expires_in_seconds)).isoformat()
|
||||
return cred
|
||||
|
||||
|
||||
|
|
@ -438,12 +675,7 @@ def test_expiry_buffer_treats_soon_to_expire_as_expired():
|
|||
cred = _oauth_cred(expires_in_seconds=30)
|
||||
assert is_oauth_credential_expired(cred, buffer_seconds=60) is True
|
||||
# A token comfortably beyond the buffer stays valid.
|
||||
assert (
|
||||
is_oauth_credential_expired(
|
||||
_oauth_cred(expires_in_seconds=600), buffer_seconds=60
|
||||
)
|
||||
is False
|
||||
)
|
||||
assert is_oauth_credential_expired(_oauth_cred(expires_in_seconds=600), buffer_seconds=60) is False
|
||||
|
||||
|
||||
def test_expiry_past_is_expired_regardless_of_buffer():
|
||||
|
|
@ -465,9 +697,7 @@ async def test_resolve_returns_valid_token_without_refreshing(monkeypatch):
|
|||
refresh = AsyncMock()
|
||||
monkeypatch.setattr(db_mod, "refresh_user_oauth_token", refresh)
|
||||
|
||||
cred = _oauth_cred(
|
||||
access_token="at-live", refresh_token="rt-1", expires_in_seconds=600
|
||||
)
|
||||
cred = _oauth_cred(access_token="at-live", refresh_token="rt-1", expires_in_seconds=600)
|
||||
result = await resolve_valid_user_oauth_token(
|
||||
user_id="alice", server=MagicMock(), cred=cred, prisma_client=MagicMock()
|
||||
)
|
||||
|
|
@ -483,15 +713,11 @@ async def test_resolve_refreshes_expired_token_with_refresh_token(monkeypatch):
|
|||
# new token rather than returning None (which left the UI tool list empty).
|
||||
import litellm.proxy._experimental.mcp_server.db as db_mod
|
||||
|
||||
refreshed = _oauth_cred(
|
||||
access_token="at-fresh", refresh_token="rt-2", expires_in_seconds=3600
|
||||
)
|
||||
refreshed = _oauth_cred(access_token="at-fresh", refresh_token="rt-2", expires_in_seconds=3600)
|
||||
refresh = AsyncMock(return_value=refreshed)
|
||||
monkeypatch.setattr(db_mod, "refresh_user_oauth_token", refresh)
|
||||
|
||||
expired = _oauth_cred(
|
||||
access_token="at-dead", refresh_token="rt-1", expires_in_seconds=-5
|
||||
)
|
||||
expired = _oauth_cred(access_token="at-dead", refresh_token="rt-1", expires_in_seconds=-5)
|
||||
result = await resolve_valid_user_oauth_token(
|
||||
user_id="alice", server=MagicMock(), cred=expired, prisma_client=MagicMock()
|
||||
)
|
||||
|
|
@ -510,9 +736,7 @@ async def test_resolve_refreshes_token_expiring_within_buffer(monkeypatch):
|
|||
refresh = AsyncMock(return_value=refreshed)
|
||||
monkeypatch.setattr(db_mod, "refresh_user_oauth_token", refresh)
|
||||
|
||||
soon = _oauth_cred(
|
||||
access_token="at-soon", refresh_token="rt-1", expires_in_seconds=30
|
||||
)
|
||||
soon = _oauth_cred(access_token="at-soon", refresh_token="rt-1", expires_in_seconds=30)
|
||||
result = await resolve_valid_user_oauth_token(
|
||||
user_id="alice", server=MagicMock(), cred=soon, prisma_client=MagicMock()
|
||||
)
|
||||
|
|
@ -546,9 +770,7 @@ async def test_resolve_returns_none_when_refresh_fails(monkeypatch):
|
|||
refresh = AsyncMock(return_value=None)
|
||||
monkeypatch.setattr(db_mod, "refresh_user_oauth_token", refresh)
|
||||
|
||||
expired = _oauth_cred(
|
||||
access_token="at-dead", refresh_token="rt-1", expires_in_seconds=-5
|
||||
)
|
||||
expired = _oauth_cred(access_token="at-dead", refresh_token="rt-1", expires_in_seconds=-5)
|
||||
result = await resolve_valid_user_oauth_token(
|
||||
user_id="alice", server=MagicMock(), cred=expired, prisma_client=MagicMock()
|
||||
)
|
||||
|
|
@ -565,9 +787,7 @@ async def test_resolve_returns_none_for_missing_credential(monkeypatch):
|
|||
monkeypatch.setattr(db_mod, "refresh_user_oauth_token", refresh)
|
||||
|
||||
assert (
|
||||
await resolve_valid_user_oauth_token(
|
||||
user_id="alice", server=MagicMock(), cred=None, prisma_client=MagicMock()
|
||||
)
|
||||
await resolve_valid_user_oauth_token(user_id="alice", server=MagicMock(), cred=None, prisma_client=MagicMock())
|
||||
is None
|
||||
)
|
||||
assert (
|
||||
|
|
@ -601,19 +821,13 @@ async def test_rotate_user_env_vars_re_encrypts_with_new_key(monkeypatch):
|
|||
encrypted_old = encrypt_value_helper(json.dumps(values))
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpuserenvvars.find_many = AsyncMock(
|
||||
return_value=[_env_var_row(encrypted_old)]
|
||||
)
|
||||
prisma.db.litellm_mcpuserenvvars.find_many = AsyncMock(return_value=[_env_var_row(encrypted_old)])
|
||||
prisma.db.litellm_mcpuserenvvars.update = AsyncMock()
|
||||
|
||||
new_master_key = "rotated-env-key-1111-2222-3333-4444"
|
||||
await rotate_mcp_user_env_vars_master_key(
|
||||
prisma_client=prisma, new_master_key=new_master_key
|
||||
)
|
||||
await rotate_mcp_user_env_vars_master_key(prisma_client=prisma, new_master_key=new_master_key)
|
||||
|
||||
new_stored = prisma.db.litellm_mcpuserenvvars.update.call_args.kwargs["data"][
|
||||
"values_b64"
|
||||
]
|
||||
new_stored = prisma.db.litellm_mcpuserenvvars.update.call_args.kwargs["data"]["values_b64"]
|
||||
assert new_stored != encrypted_old, "rotation must produce different ciphertext"
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", new_master_key)
|
||||
|
|
@ -630,18 +844,14 @@ async def test_rotate_user_env_vars_re_encrypts_with_new_key(monkeypatch):
|
|||
async def test_rotate_user_env_vars_skips_undecryptable_rows():
|
||||
# A corrupt row must be skipped (not overwritten) so recoverable data is
|
||||
# preserved and one bad row does not abort the rest of the rotation.
|
||||
good = _env_var_row(
|
||||
encrypt_value_helper(json.dumps({"A": "1"})), server_id="srv-ok"
|
||||
)
|
||||
good = _env_var_row(encrypt_value_helper(json.dumps({"A": "1"})), server_id="srv-ok")
|
||||
bad = _env_var_row("!!! not encrypted !!!", server_id="srv-corrupt")
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpuserenvvars.find_many = AsyncMock(return_value=[bad, good])
|
||||
prisma.db.litellm_mcpuserenvvars.update = AsyncMock()
|
||||
|
||||
await rotate_mcp_user_env_vars_master_key(
|
||||
prisma_client=prisma, new_master_key="new-key-xxxx"
|
||||
)
|
||||
await rotate_mcp_user_env_vars_master_key(prisma_client=prisma, new_master_key="new-key-xxxx")
|
||||
|
||||
assert prisma.db.litellm_mcpuserenvvars.update.call_count == 1
|
||||
where = prisma.db.litellm_mcpuserenvvars.update.call_args.kwargs["where"]
|
||||
|
|
@ -669,9 +879,7 @@ async def test_refresh_user_oauth_token_uses_client_secret_basic(monkeypatch):
|
|||
|
||||
monkeypatch.setattr(db_mod, "get_async_httpx_client", lambda **kwargs: mock_client)
|
||||
monkeypatch.setattr(db_mod, "store_user_oauth_credential", AsyncMock())
|
||||
monkeypatch.setattr(
|
||||
db_mod, "get_user_oauth_credential", AsyncMock(return_value={"access_token": "new-at"})
|
||||
)
|
||||
monkeypatch.setattr(db_mod, "get_user_oauth_credential", AsyncMock(return_value={"access_token": "new-at"}))
|
||||
|
||||
result = await db_mod.refresh_user_oauth_token(
|
||||
prisma_client=MagicMock(),
|
||||
|
|
@ -710,9 +918,7 @@ async def test_refresh_user_oauth_token_defaults_to_client_secret_post(monkeypat
|
|||
|
||||
monkeypatch.setattr(db_mod, "get_async_httpx_client", lambda **kwargs: mock_client)
|
||||
monkeypatch.setattr(db_mod, "store_user_oauth_credential", AsyncMock())
|
||||
monkeypatch.setattr(
|
||||
db_mod, "get_user_oauth_credential", AsyncMock(return_value={"access_token": "new-at"})
|
||||
)
|
||||
monkeypatch.setattr(db_mod, "get_user_oauth_credential", AsyncMock(return_value={"access_token": "new-at"}))
|
||||
|
||||
await db_mod.refresh_user_oauth_token(
|
||||
prisma_client=MagicMock(),
|
||||
|
|
|
|||
|
|
@ -6869,6 +6869,7 @@ async def test_call_tool_with_legacy_db_m2m_server_resolves_oauth2_flow():
|
|||
(None, None),
|
||||
("", None),
|
||||
("not a url", None),
|
||||
("http://[::1", None),
|
||||
],
|
||||
)
|
||||
def test_redact_mcp_resource_url_strips_credentials(url, expected):
|
||||
|
|
|
|||
|
|
@ -3327,9 +3327,35 @@ class TestMCPServerManager:
|
|||
await manager.invalidate_user_oauth_token_cache("alice", "srv-1")
|
||||
assert store.invalidations == [("alice", "srv-1")]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalidate_user_oauth_token_cache_drops_legacy_cache_too(self):
|
||||
"""A per-user token can be served from the legacy per-user token cache as well as the v2
|
||||
store; the shared invalidation must evict both, or the path not evicted keeps serving a
|
||||
token minted for a replaced credential row until its TTL."""
|
||||
|
||||
class _Store:
|
||||
async def fetch(self, user_id: str, server_id: str):
|
||||
return None
|
||||
|
||||
async def invalidate(self, user_id: str, server_id: str) -> None:
|
||||
return None
|
||||
|
||||
class _LegacyCache:
|
||||
def __init__(self) -> None:
|
||||
self.deletes: list[tuple[str, str]] = []
|
||||
|
||||
async def delete(self, user_id: str, server_id: str) -> None:
|
||||
self.deletes.append((user_id, server_id))
|
||||
|
||||
legacy_cache = _LegacyCache()
|
||||
manager = MCPServerManager(per_user_oauth_token_store=_Store(), per_user_token_cache=legacy_cache)
|
||||
await manager.invalidate_user_oauth_token_cache("alice", "srv-1")
|
||||
assert legacy_cache.deletes == [("alice", "srv-1")]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalidate_user_oauth_token_cache_swallows_store_errors(self):
|
||||
"""A cache-drop failure must not fail the credential write that triggered it."""
|
||||
"""A cache-drop failure must not fail the credential write that triggered it, and the
|
||||
legacy cache must still be evicted after the v2 store drop fails."""
|
||||
|
||||
class _Store:
|
||||
async def fetch(self, user_id: str, server_id: str):
|
||||
|
|
@ -3338,7 +3364,35 @@ class TestMCPServerManager:
|
|||
async def invalidate(self, user_id: str, server_id: str) -> None:
|
||||
raise RuntimeError("redis down")
|
||||
|
||||
manager = MCPServerManager(per_user_oauth_token_store=_Store())
|
||||
class _LegacyCache:
|
||||
def __init__(self) -> None:
|
||||
self.deletes: list[tuple[str, str]] = []
|
||||
|
||||
async def delete(self, user_id: str, server_id: str) -> None:
|
||||
self.deletes.append((user_id, server_id))
|
||||
|
||||
legacy_cache = _LegacyCache()
|
||||
manager = MCPServerManager(per_user_oauth_token_store=_Store(), per_user_token_cache=legacy_cache)
|
||||
await manager.invalidate_user_oauth_token_cache("alice", "srv-1")
|
||||
assert legacy_cache.deletes == [("alice", "srv-1")]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalidate_user_oauth_token_cache_swallows_legacy_cache_errors(self):
|
||||
"""The legacy cache drop is best-effort like the v2 drop: a failure must be logged, never
|
||||
raised into the credential write that triggered the invalidation."""
|
||||
|
||||
class _Store:
|
||||
async def fetch(self, user_id: str, server_id: str):
|
||||
return None
|
||||
|
||||
async def invalidate(self, user_id: str, server_id: str) -> None:
|
||||
return None
|
||||
|
||||
class _RaisingLegacyCache:
|
||||
async def delete(self, user_id: str, server_id: str) -> None:
|
||||
raise RuntimeError("redis down")
|
||||
|
||||
manager = MCPServerManager(per_user_oauth_token_store=_Store(), per_user_token_cache=_RaisingLegacyCache())
|
||||
await manager.invalidate_user_oauth_token_cache("alice", "srv-1")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -5134,3 +5134,93 @@ def test_stamp_oauth2_flow_ignores_non_oauth2():
|
|||
payload = _oauth2_create_payload(auth_type="none")
|
||||
mgmt_endpoints.stamp_omitted_oauth2_flow(payload)
|
||||
assert payload.oauth2_flow is None
|
||||
|
||||
|
||||
async def _run_edit(old_record, updated_record, purge_mock=None):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import edit_mcp_server
|
||||
|
||||
server_id = updated_record.server_id
|
||||
with (
|
||||
patch("litellm.proxy.management_endpoints.mcp_management_endpoints.MCP_AVAILABLE", True),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server",
|
||||
AsyncMock(side_effect=old_record)
|
||||
if isinstance(old_record, Exception)
|
||||
else AsyncMock(return_value=old_record),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.update_mcp_server",
|
||||
AsyncMock(return_value=updated_record),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.validate_and_normalize_mcp_server_payload",
|
||||
autospec=True,
|
||||
),
|
||||
patch("litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager") as mock_manager,
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.purge_user_oauth_credentials_for_server",
|
||||
purge_mock if purge_mock is not None else AsyncMock(return_value=1),
|
||||
) as mock_purge,
|
||||
):
|
||||
mock_manager.update_server = AsyncMock()
|
||||
mock_manager.reload_servers_from_database = AsyncMock()
|
||||
payload = UpdateMCPServerRequest(server_id=server_id, alias=updated_record.alias, url=updated_record.url)
|
||||
user_auth = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
result = await edit_mcp_server(payload=payload, user_api_key_dict=user_auth)
|
||||
return result, mock_purge
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_edit_mcp_server_purges_user_tokens_on_mint_relevant_change():
|
||||
server_id = str(uuid.uuid4())
|
||||
old = generate_mock_mcp_server_db_record(server_id=server_id, url="https://old.example.com/mcp")
|
||||
updated = generate_mock_mcp_server_db_record(server_id=server_id, url="https://new.example.com/mcp")
|
||||
|
||||
result, mock_purge = await _run_edit(old, updated)
|
||||
|
||||
assert result.server_id == server_id
|
||||
mock_purge.assert_awaited_once()
|
||||
assert mock_purge.await_args.args[1] == server_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_edit_mcp_server_skips_purge_when_identity_unchanged():
|
||||
server_id = str(uuid.uuid4())
|
||||
old = generate_mock_mcp_server_db_record(server_id=server_id, alias="Before")
|
||||
updated = generate_mock_mcp_server_db_record(server_id=server_id, alias="After")
|
||||
|
||||
result, mock_purge = await _run_edit(old, updated)
|
||||
|
||||
assert result.server_id == server_id
|
||||
mock_purge.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_edit_mcp_server_purge_failure_does_not_fail_the_edit():
|
||||
"""The purge is best-effort: a purge exception after a successful update must be swallowed and
|
||||
logged, never turned into an error response for an edit whose primary job already succeeded."""
|
||||
server_id = str(uuid.uuid4())
|
||||
old = generate_mock_mcp_server_db_record(server_id=server_id, url="https://old.example.com/mcp")
|
||||
updated = generate_mock_mcp_server_db_record(server_id=server_id, url="https://new.example.com/mcp")
|
||||
|
||||
result, mock_purge = await _run_edit(old, updated, purge_mock=AsyncMock(side_effect=RuntimeError("db down")))
|
||||
|
||||
assert result.server_id == server_id
|
||||
mock_purge.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_edit_mcp_server_snapshot_failure_skips_purge_but_edit_succeeds():
|
||||
"""The pre-update snapshot read is advisory (it only feeds the purge decision); a read failure
|
||||
must skip the stale-token check with a warning, never fail the edit itself."""
|
||||
server_id = str(uuid.uuid4())
|
||||
updated = generate_mock_mcp_server_db_record(server_id=server_id, url="https://new.example.com/mcp")
|
||||
|
||||
result, mock_purge = await _run_edit(RuntimeError("db read failed"), updated)
|
||||
|
||||
assert result.server_id == server_id
|
||||
mock_purge.assert_not_awaited()
|
||||
|
|
|
|||
|
|
@ -374,6 +374,68 @@ describe("CreateMCPServer", () => {
|
|||
expect(credentials.access_token).toBeUndefined();
|
||||
});
|
||||
|
||||
it.each([
|
||||
["true_passthrough", "True Passthrough (no LiteLLM auth)"],
|
||||
["oauth_delegate", "OAuth Delegate (client-supplied upstream token)"],
|
||||
])("persists only tool config on create for %s; the token stays browser-held", async (_authType, optionLabel) => {
|
||||
oauthHook.tokenResponse = { access_token: "upstream-tok", token_type: "Bearer" };
|
||||
await selectHttpTransport();
|
||||
|
||||
const user = userEvent.setup({ delay: null });
|
||||
await user.type(getServerNameInput(), "CF_Server");
|
||||
await user.type(screen.getByPlaceholderText("https://your-mcp-server.com"), "https://example.com/mcp");
|
||||
|
||||
await selectAntOption("Authentication", optionLabel);
|
||||
|
||||
await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy());
|
||||
await act(async () => {
|
||||
oauthHook.onTokenReceived!({ access_token: "upstream-tok", token_type: "Bearer" }, undefined);
|
||||
});
|
||||
|
||||
fireEvent.click(screen.getByRole("button", { name: "Disable all tools" }));
|
||||
|
||||
// Previewing and configuring must stay stateless: nothing is persisted anywhere (server row,
|
||||
// per-user DB credential, sessionStorage) until the admin submits.
|
||||
expect(networking.createMCPServer).not.toHaveBeenCalled();
|
||||
expect(networking.storeMCPOAuthUserCredential).not.toHaveBeenCalled();
|
||||
expect(setToken).not.toHaveBeenCalled();
|
||||
|
||||
const createdServer = {
|
||||
server_id: "new-cf-server",
|
||||
server_name: "CF_Server",
|
||||
alias: "CF_Server",
|
||||
url: "https://example.com/mcp",
|
||||
transport: "http",
|
||||
auth_type: _authType,
|
||||
created_at: "2024-01-01T00:00:00Z",
|
||||
created_by: "user-1",
|
||||
updated_at: "2024-01-01T00:00:00Z",
|
||||
updated_by: "user-1",
|
||||
};
|
||||
vi.mocked(networking.createMCPServer).mockResolvedValue(createdServer);
|
||||
|
||||
const submitButton = screen.getByRole("button", { name: "Add MCP Server" });
|
||||
await act(async () => {
|
||||
fireEvent.click(submitButton);
|
||||
});
|
||||
|
||||
await waitFor(() => expect(networking.createMCPServer).toHaveBeenCalledTimes(1));
|
||||
const [, payload] = vi.mocked(networking.createMCPServer).mock.calls[0];
|
||||
|
||||
// Only the tool configuration persists on the server row; the upstream token appears nowhere
|
||||
// in the create payload and no per-user DB credential is written. The token is committed to
|
||||
// sessionStorage only, keyed to the created server.
|
||||
expect(payload.allowed_tools).toEqual([]);
|
||||
expect(payload.credentials).toBeUndefined();
|
||||
expect(JSON.stringify(payload)).not.toContain("upstream-tok");
|
||||
expect(networking.storeMCPOAuthUserCredential).not.toHaveBeenCalled();
|
||||
expect(setToken).toHaveBeenCalledWith(
|
||||
"new-cf-server",
|
||||
expect.objectContaining({ access_token: "upstream-tok" }),
|
||||
undefined,
|
||||
);
|
||||
});
|
||||
|
||||
it("should not show auth value field when None auth type is selected", async () => {
|
||||
await selectHttpTransport();
|
||||
|
||||
|
|
@ -681,6 +743,93 @@ describe("CreateMCPServer", () => {
|
|||
// Asserted in setupOAuthInteractive
|
||||
});
|
||||
|
||||
it("invalidates the held token when the auth mode changes after Authorize & Fetch", async () => {
|
||||
await setupOAuthInteractive();
|
||||
const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com");
|
||||
await act(async () => {
|
||||
fireEvent.change(urlInput, { target: { value: "https://a.example.com/mcp" } });
|
||||
});
|
||||
act(() => {
|
||||
oauthHook.onTokenReceived?.({ access_token: "tok-a" }, { clientId: "client-a", clientSecret: "secret-a" });
|
||||
});
|
||||
oauthHook.reset.mockClear();
|
||||
|
||||
// Switching the Authentication mode changes the OAuth identity, so the held token is discarded.
|
||||
await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)");
|
||||
|
||||
await waitFor(() => expect(oauthHook.reset).toHaveBeenCalled());
|
||||
});
|
||||
|
||||
it("does NOT invalidate the held token when a non-mint field (server name) changes", async () => {
|
||||
await setupOAuthInteractive();
|
||||
const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com");
|
||||
await act(async () => {
|
||||
fireEvent.change(urlInput, { target: { value: "https://a.example.com/mcp" } });
|
||||
});
|
||||
act(() => {
|
||||
oauthHook.onTokenReceived?.({ access_token: "tok-a" }, { clientId: "client-a", clientSecret: "secret-a" });
|
||||
});
|
||||
oauthHook.reset.mockClear();
|
||||
|
||||
const nameInput = document.getElementById("server_name") as HTMLInputElement;
|
||||
await act(async () => {
|
||||
fireEvent.change(nameInput, { target: { value: "Renamed_Server" } });
|
||||
});
|
||||
|
||||
// server_name is not part of the OAuth identity, so the held token must survive the edit.
|
||||
await waitFor(() => expect(screen.getAllByRole("button", { name: "Add MCP Server" }).length).toBeGreaterThan(0));
|
||||
expect(oauthHook.reset).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("does not refetch the tool preview with a discarded token after invalidation", async () => {
|
||||
// Regression: handleFormValuesChange used to publish the pre-reset antd snapshot into
|
||||
// formValues after clearHeldOAuthToken, so useTestMCPConnection kept the discarded OAuth
|
||||
// material (the DCR client minted for the old identity) and sent it on the next tool-preview
|
||||
// request.
|
||||
await setupOAuthInteractive();
|
||||
const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com");
|
||||
await act(async () => {
|
||||
fireEvent.change(urlInput, { target: { value: "https://a.example.com/mcp" } });
|
||||
});
|
||||
act(() => {
|
||||
oauthHook.onTokenReceived?.({ access_token: "stale-tok" }, { clientId: "client-a", clientSecret: "secret-a" });
|
||||
});
|
||||
const nameInput = document.getElementById("server_name") as HTMLInputElement;
|
||||
await act(async () => {
|
||||
fireEvent.change(nameInput, { target: { value: "Sync_FormValues" } });
|
||||
});
|
||||
vi.mocked(networking.testMCPToolsListRequest).mockClear();
|
||||
|
||||
await selectAntOption("Authentication", "API Key");
|
||||
|
||||
await waitFor(() => expect(vi.mocked(networking.testMCPToolsListRequest)).toHaveBeenCalled());
|
||||
for (const call of vi.mocked(networking.testMCPToolsListRequest).mock.calls) {
|
||||
expect(call[1]?.credentials?.client_id).not.toBe("client-a");
|
||||
expect(call[1]?.credentials?.client_secret).not.toBe("secret-a");
|
||||
}
|
||||
});
|
||||
|
||||
it("keeps the held token on an http to sse switch with the same url", async () => {
|
||||
// Same url means the same resource/audience (RFC 8707): the minted token is still valid, so a
|
||||
// pure transport swap between the two MCP wire protocols must not force a re-authorize.
|
||||
await setupOAuthInteractive();
|
||||
const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com");
|
||||
await act(async () => {
|
||||
fireEvent.change(urlInput, { target: { value: "https://a.example.com/mcp" } });
|
||||
});
|
||||
act(() => {
|
||||
oauthHook.onTokenReceived?.({ access_token: "tok-a" }, { clientId: "client-a", clientSecret: "secret-a" });
|
||||
});
|
||||
oauthHook.reset.mockClear();
|
||||
|
||||
await selectAntOption("Transport Type", "Server-Sent Events (SSE)");
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByPlaceholderText("https://your-mcp-server.com")).toBeInTheDocument();
|
||||
});
|
||||
expect(oauthHook.reset).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("includes token_validation in payload when token_validation_json is filled with valid JSON", async () => {
|
||||
vi.mocked(networking.createMCPServer).mockResolvedValue({
|
||||
server_id: "new-server-oauth",
|
||||
|
|
|
|||
|
|
@ -15,6 +15,9 @@ import {
|
|||
MCP_OAUTH2_FLOW_M2M,
|
||||
MCP_OAUTH2_FLOW_INTERACTIVE,
|
||||
isClientForwardedTokenMode,
|
||||
getOAuthAuthorizationIdentity,
|
||||
CLEARED_ON_INVALIDATION,
|
||||
isHeldOAuthTokenStale,
|
||||
} from "./types";
|
||||
import OAuthFormFields from "./OAuthFormFields";
|
||||
import TruePassthroughWarning from "./TruePassthroughWarning";
|
||||
|
|
@ -99,7 +102,10 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
|
|||
const [oauthAccessToken, setOauthAccessToken] = useState<string | null>(null);
|
||||
const [logoUrl, setLogoUrl] = useState<string | undefined>(undefined);
|
||||
const [oauthDocsUrl, setOauthDocsUrl] = useState<string | null>(null);
|
||||
const [authorizedUrl, setAuthorizedUrl] = useState<string | undefined>(undefined);
|
||||
// The OAuth authorization identity (see getOAuthAuthorizationIdentity) captured at the moment a token
|
||||
// was fetched; undefined when no valid token is held. If any mint-relevant field diverges from this,
|
||||
// the held token is stale and is discarded so the admin must re-authorize.
|
||||
const [authorizedIdentity, setAuthorizedIdentity] = useState<string | undefined>(undefined);
|
||||
|
||||
// Single hook call shared by MCPConnectionStatus and MCPToolConfiguration to avoid duplicate requests.
|
||||
const {
|
||||
|
|
@ -125,12 +131,6 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
|
|||
const isAwsSigV4AuthType = authType === AUTH_TYPE.AWS_SIGV4;
|
||||
const isM2MFlow = isOAuthAuthType && formValues.oauth_flow_type === OAUTH_FLOW.M2M;
|
||||
|
||||
const getOAuthAuthorizationTarget = (values: Record<string, unknown>): string | undefined => {
|
||||
const transport = values.transport || transportType;
|
||||
const target = transport === TRANSPORT.OPENAPI ? values.spec_path : values.url;
|
||||
return typeof target === "string" ? target : undefined;
|
||||
};
|
||||
|
||||
const persistCreateUiState = () => {
|
||||
if (typeof window === "undefined") {
|
||||
return;
|
||||
|
|
@ -207,6 +207,7 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
|
|||
// and committed to sessionStorage on submit; it must never be written into form.credentials,
|
||||
// which would persist it as server-level credentials on the created server row. Mirrors the
|
||||
// edit form's onTokenReceived early return.
|
||||
setAuthorizedIdentity(getOAuthAuthorizationIdentity(form.getFieldsValue(true)));
|
||||
NotificationsManager.success(
|
||||
"Token held for this browser session. Tools can now be previewed and configured; nothing will be saved to LiteLLM.",
|
||||
);
|
||||
|
|
@ -223,7 +224,9 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
|
|||
};
|
||||
|
||||
form.setFieldsValue({ credentials });
|
||||
setAuthorizedUrl(getOAuthAuthorizationTarget(form.getFieldsValue(true)));
|
||||
// Capture the identity AFTER writing the DCR'd credentials so the held token is not spuriously
|
||||
// invalidated by its own credential write.
|
||||
setAuthorizedIdentity(getOAuthAuthorizationIdentity(form.getFieldsValue(true)));
|
||||
|
||||
NotificationsManager.success(
|
||||
"OAuth authorization successful! Please click 'Create MCP Server' to save the configuration.",
|
||||
|
|
@ -233,13 +236,23 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
|
|||
flowSource: "create",
|
||||
});
|
||||
|
||||
const clearAuthorizedOAuthState = (values: Record<string, unknown>) => {
|
||||
form.resetFields(["credentials", "authorization_url", "token_url", "registration_url"]);
|
||||
form.setFieldsValue(values);
|
||||
// Discard the held browser-authorized token and its tool preview when the authorization identity
|
||||
// changes (or the modal closes). The CLEARED_ON_INVALIDATION form fields (shared with the edit form
|
||||
// via types.tsx) are reset too; whatever the admin just changed (passed via changedValues) is
|
||||
// re-applied so the invalidation never wipes their in-flight edit. Admin-typed endpoint fields are
|
||||
// left alone (see CLEARED_ON_INVALIDATION).
|
||||
const clearHeldOAuthToken = (changedValues: Record<string, unknown> = {}) => {
|
||||
setOauthAccessToken(null);
|
||||
clearTools();
|
||||
resetOAuthFlow();
|
||||
setAuthorizedUrl(undefined);
|
||||
setAuthorizedIdentity(undefined);
|
||||
form.resetFields([...CLEARED_ON_INVALIDATION]);
|
||||
const preserved = Object.fromEntries(
|
||||
CLEARED_ON_INVALIDATION.filter((key) => key in changedValues).map((key) => [key, changedValues[key]]),
|
||||
);
|
||||
if (Object.keys(preserved).length > 0) {
|
||||
form.setFieldsValue(preserved);
|
||||
}
|
||||
};
|
||||
|
||||
React.useEffect(() => {
|
||||
|
|
@ -576,22 +589,11 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
|
|||
? { url: undefined, command: undefined, args: undefined, env: undefined }
|
||||
: { spec_path: undefined, command: undefined, args: undefined, env: undefined };
|
||||
|
||||
const nextValues =
|
||||
authorizedUrl === undefined
|
||||
? transportValues
|
||||
: {
|
||||
...transportValues,
|
||||
credentials: undefined,
|
||||
authorization_url: undefined,
|
||||
token_url: undefined,
|
||||
registration_url: undefined,
|
||||
};
|
||||
|
||||
if (authorizedUrl !== undefined) {
|
||||
clearAuthorizedOAuthState(nextValues);
|
||||
} else {
|
||||
form.setFieldsValue(nextValues);
|
||||
form.setFieldsValue(transportValues);
|
||||
if (isHeldOAuthTokenStale(form.getFieldsValue(true), authorizedIdentity)) {
|
||||
clearHeldOAuthToken();
|
||||
}
|
||||
setFormValues(form.getFieldsValue(true));
|
||||
};
|
||||
|
||||
// Generate options with existing groups and potential new group
|
||||
|
|
@ -652,27 +654,21 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
|
|||
setOauthAccessToken(null);
|
||||
clearTools();
|
||||
resetOAuthFlow();
|
||||
setAuthorizedUrl(undefined);
|
||||
setAuthorizedIdentity(undefined);
|
||||
}
|
||||
}, [isModalVisible, form, clearTools, resetOAuthFlow]);
|
||||
|
||||
const isAdmin = isAdminRole(userRole);
|
||||
|
||||
const handleFormValuesChange = (changedValues: Record<string, unknown>, allValues: Record<string, unknown>) => {
|
||||
const changedAuthorizationTarget = "url" in changedValues || "spec_path" in changedValues;
|
||||
if (
|
||||
changedAuthorizationTarget &&
|
||||
authorizedUrl !== undefined &&
|
||||
getOAuthAuthorizationTarget(allValues) !== authorizedUrl
|
||||
) {
|
||||
const invalidated = {
|
||||
credentials: undefined,
|
||||
authorization_url: changedValues.authorization_url,
|
||||
token_url: changedValues.token_url,
|
||||
registration_url: changedValues.registration_url,
|
||||
};
|
||||
clearAuthorizedOAuthState(invalidated);
|
||||
setFormValues({ ...allValues, ...invalidated });
|
||||
// Any change to a mint-relevant field (url, auth_type, oauth_flow_type, client creds/scopes, or the
|
||||
// authorization/token/registration endpoints — see getOAuthAuthorizationIdentity) makes a held token
|
||||
// stale, so discard it and force a fresh authorize. When that happens, formValues must be rebuilt
|
||||
// from the form's post-reset state, not the pre-reset allValues snapshot: the snapshot still holds
|
||||
// the discarded token in credentials, and useTestMCPConnection reads formValues for tool preview.
|
||||
if (isHeldOAuthTokenStale(allValues, authorizedIdentity)) {
|
||||
clearHeldOAuthToken(changedValues);
|
||||
setFormValues({ ...form.getFieldsValue(true), ...changedValues });
|
||||
return;
|
||||
}
|
||||
setFormValues(allValues);
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ vi.mock("../networking", () => ({
|
|||
updateMCPServer: vi.fn(),
|
||||
listMCPTools: vi.fn().mockResolvedValue({ tools: [], error: null }),
|
||||
storeMCPOAuthUserCredential: vi.fn().mockResolvedValue({}),
|
||||
testMCPToolsListRequest: vi.fn().mockResolvedValue({ tools: [], error: null }),
|
||||
}));
|
||||
|
||||
vi.mock("../molecules/notifications_manager", () => ({
|
||||
|
|
@ -22,15 +23,22 @@ vi.mock("../molecules/notifications_manager", () => ({
|
|||
const mockOauth: {
|
||||
tokenResponse: any;
|
||||
getTemporaryPayload: (() => Record<string, unknown> | null) | null;
|
||||
} = { tokenResponse: null, getTemporaryPayload: null };
|
||||
onTokenReceived: ((token: Record<string, unknown> | null) => void) | null;
|
||||
reset: ReturnType<typeof vi.fn>;
|
||||
} = { tokenResponse: null, getTemporaryPayload: null, onTokenReceived: null, reset: vi.fn() };
|
||||
vi.mock("@/hooks/useMcpOAuthFlow", () => ({
|
||||
useMcpOAuthFlow: (opts: { getTemporaryPayload?: () => Record<string, unknown> | null }) => {
|
||||
useMcpOAuthFlow: (opts: {
|
||||
getTemporaryPayload?: () => Record<string, unknown> | null;
|
||||
onTokenReceived?: (token: Record<string, unknown> | null) => void;
|
||||
}) => {
|
||||
mockOauth.getTemporaryPayload = opts?.getTemporaryPayload ?? null;
|
||||
mockOauth.onTokenReceived = opts?.onTokenReceived ?? null;
|
||||
return {
|
||||
startOAuthFlow: vi.fn(),
|
||||
status: "idle",
|
||||
error: null,
|
||||
tokenResponse: mockOauth.tokenResponse,
|
||||
reset: mockOauth.reset,
|
||||
};
|
||||
},
|
||||
}));
|
||||
|
|
@ -92,10 +100,12 @@ vi.mock("./mcp_tool_configuration", () => ({
|
|||
const mockGetToken = vi.fn();
|
||||
const mockIsTokenValid = vi.fn();
|
||||
const mockSetToken = vi.fn();
|
||||
const mockRemoveToken = vi.fn();
|
||||
vi.mock("@/utils/mcpTokenStore", () => ({
|
||||
getToken: (...args: any[]) => mockGetToken(...args),
|
||||
isTokenValid: (...args: any[]) => mockIsTokenValid(...args),
|
||||
setToken: (...args: any[]) => mockSetToken(...args),
|
||||
removeToken: (...args: unknown[]) => mockRemoveToken(...args),
|
||||
}));
|
||||
|
||||
// ── fixtures ──────────────────────────────────────────────────────────────────
|
||||
|
|
@ -451,6 +461,159 @@ describe("MCPServerEdit (auth type switch)", () => {
|
|||
});
|
||||
});
|
||||
|
||||
describe("MCPServerEdit OAuth token invalidation", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
|
||||
const renderOAuthEdit = () =>
|
||||
render(
|
||||
<MCPServerEdit
|
||||
mcpServer={{ ...interactiveOAuthServer }}
|
||||
accessToken="access-token"
|
||||
onCancel={vi.fn()}
|
||||
onSuccess={vi.fn()}
|
||||
availableAccessGroups={[]}
|
||||
/>,
|
||||
);
|
||||
|
||||
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.
|
||||
// Regression: the token used to survive this switch (sessionStorage + hook state kept the old
|
||||
// token minted for the http url).
|
||||
renderOAuthEdit();
|
||||
|
||||
act(() => {
|
||||
mockOauth.onTokenReceived?.({ access_token: "tok-1" });
|
||||
});
|
||||
mockOauth.reset.mockClear();
|
||||
|
||||
await selectAntOption("Transport Type", "Standard Input/Output (stdio)");
|
||||
|
||||
await waitFor(() => expect(mockOauth.reset).toHaveBeenCalled());
|
||||
expect(mockRemoveToken).toHaveBeenCalledWith("oauth_server_1", undefined);
|
||||
});
|
||||
|
||||
it("invalidates a session-authorized token when the server URL changes", async () => {
|
||||
renderOAuthEdit();
|
||||
|
||||
act(() => {
|
||||
mockOauth.onTokenReceived?.({ access_token: "tok-1" });
|
||||
});
|
||||
mockOauth.reset.mockClear();
|
||||
|
||||
const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com");
|
||||
await act(async () => {
|
||||
fireEvent.change(urlInput, { target: { value: "https://other.example.com/mcp" } });
|
||||
});
|
||||
|
||||
await waitFor(() => expect(mockOauth.reset).toHaveBeenCalled());
|
||||
expect(mockRemoveToken).toHaveBeenCalledWith("oauth_server_1", undefined);
|
||||
});
|
||||
|
||||
it("previews tools with a staged interactive OAuth token before it is saved", async () => {
|
||||
// Regression: for authorization_code the fetch went by server_id only, relying on the stored DB
|
||||
// credential, so a token authorized in this edit session gave an empty preview until the admin
|
||||
// saved; the create form previews the identical state via the config-based preview endpoint.
|
||||
mockOauth.tokenResponse = { access_token: "staged-obo-tok" };
|
||||
|
||||
renderOAuthEdit();
|
||||
|
||||
await waitFor(() => {
|
||||
expect(vi.mocked(networking.testMCPToolsListRequest)).toHaveBeenCalledWith(
|
||||
"access-token",
|
||||
// oauth2_flow must be explicit: the preview endpoint infers client_credentials from
|
||||
// inherited client_id/client_secret/token_url and would strip the staged bearer.
|
||||
expect.objectContaining({
|
||||
server_id: "oauth_server_1",
|
||||
url: "https://example.com/mcp",
|
||||
oauth2_flow: "authorization_code",
|
||||
}),
|
||||
"staged-obo-tok",
|
||||
);
|
||||
});
|
||||
expect(networking.listMCPTools).not.toHaveBeenCalled();
|
||||
// Previewing must stay stateless: the staged token is committed only by an explicit Save
|
||||
// (storeMCPOAuthUserCredential for authorization_code, setToken for the client-forwarded modes).
|
||||
expect(networking.storeMCPOAuthUserCredential).not.toHaveBeenCalled();
|
||||
expect(mockSetToken).not.toHaveBeenCalled();
|
||||
expect(networking.updateMCPServer).not.toHaveBeenCalled();
|
||||
mockOauth.tokenResponse = null;
|
||||
});
|
||||
|
||||
it("previews an OpenAPI server's staged token against its spec_path", async () => {
|
||||
mockOauth.tokenResponse = { access_token: "staged-obo-tok" };
|
||||
|
||||
render(
|
||||
<MCPServerEdit
|
||||
mcpServer={{
|
||||
...interactiveOAuthServer,
|
||||
transport: "openapi",
|
||||
url: null,
|
||||
spec_path: "https://example.com/openapi.json",
|
||||
}}
|
||||
accessToken="access-token"
|
||||
onCancel={vi.fn()}
|
||||
onSuccess={vi.fn()}
|
||||
availableAccessGroups={[]}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(vi.mocked(networking.testMCPToolsListRequest)).toHaveBeenCalledWith(
|
||||
"access-token",
|
||||
expect.objectContaining({ spec_path: "https://example.com/openapi.json" }),
|
||||
"staged-obo-tok",
|
||||
);
|
||||
});
|
||||
mockOauth.tokenResponse = null;
|
||||
});
|
||||
|
||||
it("keeps the admin's in-flight endpoint edits when the token is invalidated", async () => {
|
||||
// Regression: invalidation used to form.resetFields the endpoint fields; with the edit Form's
|
||||
// initialValues that silently reverted an admin-corrected token_url back to the saved (wrong)
|
||||
// value while still looking plausible. Only credentials (the minted material) may be wiped.
|
||||
renderOAuthEdit();
|
||||
|
||||
const tokenUrlInput = screen.getByPlaceholderText("https://example.com/oauth/token");
|
||||
await act(async () => {
|
||||
fireEvent.change(tokenUrlInput, { target: { value: "https://corrected.example.com/token" } });
|
||||
});
|
||||
|
||||
act(() => {
|
||||
mockOauth.onTokenReceived?.({ access_token: "tok-1" });
|
||||
});
|
||||
mockOauth.reset.mockClear();
|
||||
|
||||
const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com");
|
||||
await act(async () => {
|
||||
fireEvent.change(urlInput, { target: { value: "https://moved.example.com/mcp" } });
|
||||
});
|
||||
|
||||
await waitFor(() => expect(mockOauth.reset).toHaveBeenCalled());
|
||||
expect((screen.getByPlaceholderText("https://example.com/oauth/token") as HTMLInputElement).value).toBe(
|
||||
"https://corrected.example.com/token",
|
||||
);
|
||||
});
|
||||
|
||||
it("keeps a session-authorized token on an http to sse switch with the same url", async () => {
|
||||
// Same url means the same resource/audience (RFC 8707): the minted token is still valid, so a
|
||||
// pure transport swap between the two MCP wire protocols must not force a re-authorize.
|
||||
renderOAuthEdit();
|
||||
|
||||
act(() => {
|
||||
mockOauth.onTokenReceived?.({ access_token: "tok-1" });
|
||||
});
|
||||
mockOauth.reset.mockClear();
|
||||
|
||||
await selectAntOption("Transport Type", "Server-Sent Events (SSE)");
|
||||
|
||||
expect(mockOauth.reset).not.toHaveBeenCalled();
|
||||
expect(mockRemoveToken).not.toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
||||
describe("MCPServerEdit (tool allowlist)", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
|
|
@ -1210,6 +1373,7 @@ describe("MCPServerEdit (OAuth token persistence on save)", () => {
|
|||
expect(networking.storeMCPOAuthUserCredential).not.toHaveBeenCalled();
|
||||
const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0];
|
||||
expect(payload.credentials).toBeUndefined();
|
||||
expect(JSON.stringify(payload)).not.toContain("cf-tok");
|
||||
},
|
||||
);
|
||||
|
||||
|
|
|
|||
|
|
@ -5,6 +5,9 @@ import { Button, TabGroup, TabList, Tab, TabPanels, TabPanel } from "@tremor/rea
|
|||
import {
|
||||
AUTH_TYPE,
|
||||
isClientForwardedTokenMode,
|
||||
getOAuthAuthorizationIdentity,
|
||||
CLEARED_ON_INVALIDATION,
|
||||
isHeldOAuthTokenStale,
|
||||
OAUTH_FLOW,
|
||||
MCP_OAUTH2_FLOW_M2M,
|
||||
MCP_OAUTH2_FLOW_INTERACTIVE,
|
||||
|
|
@ -14,8 +17,8 @@ import {
|
|||
getMcpOAuthMode,
|
||||
oauth2FlowToFormValue,
|
||||
} from "./types";
|
||||
import { updateMCPServer, listMCPTools, storeMCPOAuthUserCredential } from "../networking";
|
||||
import { getToken, isTokenValid, setToken } from "@/utils/mcpTokenStore";
|
||||
import { updateMCPServer, listMCPTools, storeMCPOAuthUserCredential, testMCPToolsListRequest } from "../networking";
|
||||
import { getToken, isTokenValid, removeToken, setToken } from "@/utils/mcpTokenStore";
|
||||
import { buildMcpPassthroughAuthHeader } from "@/utils/mcpHeaderUtils";
|
||||
import MCPServerCostConfig from "./mcp_server_cost_config";
|
||||
import MCPPermissionManagement from "./MCPPermissionManagement";
|
||||
|
|
@ -136,11 +139,17 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
// that read only mcpServer.auth_type go stale the moment the admin switches modes in the form.
|
||||
const getEffectiveAuthType = () => form.getFieldValue("auth_type") ?? mcpServer.auth_type;
|
||||
|
||||
// The OAuth authorization identity (see getOAuthAuthorizationIdentity) captured when a token is fetched
|
||||
// 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 {
|
||||
startOAuthFlow,
|
||||
status: oauthStatus,
|
||||
error: oauthError,
|
||||
tokenResponse: oauthTokenResponse,
|
||||
reset: resetOAuthFlow,
|
||||
} = useMcpOAuthFlow({
|
||||
accessToken,
|
||||
getCredentials: () => form.getFieldValue("credentials"),
|
||||
|
|
@ -183,6 +192,7 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
return;
|
||||
}
|
||||
|
||||
authorizedIdentityRef.current = getOAuthAuthorizationIdentity(form.getFieldsValue(true));
|
||||
if (isClientForwardedTokenMode(getEffectiveAuthType())) {
|
||||
const browserHeldToken = {
|
||||
access_token: token.access_token,
|
||||
|
|
@ -205,6 +215,8 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
};
|
||||
|
||||
form.setFieldsValue({ credentials });
|
||||
// Re-capture after writing credentials so the token is not invalidated by its own credential write.
|
||||
authorizedIdentityRef.current = getOAuthAuthorizationIdentity(form.getFieldsValue(true));
|
||||
|
||||
NotificationsManager.success(
|
||||
"OAuth authorization successful! Please click 'Update MCP Server' to save the credentials.",
|
||||
|
|
@ -378,6 +390,91 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [mcpServer, accessToken, userID, oauthTokenResponse?.access_token]);
|
||||
|
||||
// Invalidate a token authorized in this edit session once any mint-relevant field diverges from the
|
||||
// identity it was minted against (url, auth_type, oauth_flow_type, client creds/scopes, or the
|
||||
// authorization/token/registration endpoints — see getOAuthAuthorizationIdentity). Discards the hook
|
||||
// token (resetOAuthFlow, which re-runs fetchTools to prompt a fresh authorize), the sessionStorage
|
||||
// token (removeToken, browser-held modes), and the fetched token/DCR client in the shared
|
||||
// CLEARED_ON_INVALIDATION form fields; the admin's in-flight edit is re-applied so it is never wiped.
|
||||
// Only fires when a token was actually authorized here (ref set), so a token already valid for the
|
||||
// saved server on mount is left untouched. Driven from onValuesChange for user input, plus an explicit
|
||||
// recheck after programmatic setFieldsValue paths (handleTransportChange), which antd does not report
|
||||
// through onValuesChange.
|
||||
const clearHeldOAuthToken = (changedValues: Record<string, unknown> = {}) => {
|
||||
authorizedIdentityRef.current = undefined;
|
||||
if (mcpServer.server_id) {
|
||||
removeToken(mcpServer.server_id, userID);
|
||||
}
|
||||
setTools([]);
|
||||
resetOAuthFlow();
|
||||
form.resetFields([...CLEARED_ON_INVALIDATION]);
|
||||
const preserved = Object.fromEntries(
|
||||
CLEARED_ON_INVALIDATION.filter((key) => key in changedValues).map((key) => [key, changedValues[key]]),
|
||||
);
|
||||
if (Object.keys(preserved).length > 0) {
|
||||
form.setFieldsValue(preserved);
|
||||
}
|
||||
};
|
||||
|
||||
const handleFormValuesChange = (changedValues: Record<string, unknown>) => {
|
||||
if (isHeldOAuthTokenStale(form.getFieldsValue(true), authorizedIdentityRef.current)) {
|
||||
clearHeldOAuthToken(changedValues);
|
||||
}
|
||||
};
|
||||
|
||||
// A token authorized in this edit session for interactive OAuth (authorization_code) is only
|
||||
// committed to the DB on save, so a plain by-server_id listing cannot use it and the preview would
|
||||
// stay empty until the admin saves; the create form previews the identical state through the
|
||||
// config-based preview endpoint, which takes the staged token explicitly. Returns false when there
|
||||
// is no staged interactive token so fetchTools falls through to the by-server_id listing.
|
||||
const previewWithStagedInteractiveToken = async (
|
||||
isPassthrough: boolean,
|
||||
isBrowserHeldTokenMode: boolean,
|
||||
): Promise<boolean> => {
|
||||
const stagedToken =
|
||||
!isPassthrough && !isBrowserHeldTokenMode && getEffectiveAuthType() === AUTH_TYPE.OAUTH2
|
||||
? oauthTokenResponse?.access_token
|
||||
: undefined;
|
||||
if (!stagedToken) {
|
||||
return false;
|
||||
}
|
||||
setIsLoadingTools(true);
|
||||
setToolsError(null);
|
||||
try {
|
||||
const values = form.getFieldsValue(true);
|
||||
const rawTransport = values.transport || mcpServer.transport;
|
||||
// oauth2_flow must be explicit: the preview endpoint infers client_credentials from the
|
||||
// inherited client_id/client_secret/token_url (common once DCR or discovery filled them) and
|
||||
// would strip the staged bearer to preview as M2M. spec_path keeps OpenAPI servers on the
|
||||
// spec-based preview path, mirroring the create form's config.
|
||||
const previewConfig = {
|
||||
server_id: mcpServer.server_id,
|
||||
server_name: values.server_name || mcpServer.server_name || mcpServer.alias,
|
||||
url: values.url || mcpServer.url,
|
||||
spec_path: values.spec_path || mcpServer.spec_path,
|
||||
transport: rawTransport === TRANSPORT.OPENAPI ? TRANSPORT.HTTP : rawTransport,
|
||||
auth_type: AUTH_TYPE.OAUTH2,
|
||||
oauth2_flow: MCP_OAUTH2_FLOW_INTERACTIVE,
|
||||
authorization_url: values.authorization_url,
|
||||
token_url: values.token_url,
|
||||
registration_url: values.registration_url,
|
||||
};
|
||||
const toolsResponse = await testMCPToolsListRequest(accessToken, previewConfig, stagedToken);
|
||||
if (toolsResponse.tools && !toolsResponse.error) {
|
||||
setTools(toolsResponse.tools);
|
||||
} else {
|
||||
setTools([]);
|
||||
setToolsError(toolsResponse.message || "Failed to load tools");
|
||||
}
|
||||
} catch (error) {
|
||||
setTools([]);
|
||||
setToolsError(error instanceof Error ? error.message : "Failed to load tools");
|
||||
} finally {
|
||||
setIsLoadingTools(false);
|
||||
}
|
||||
return true;
|
||||
};
|
||||
|
||||
const fetchTools = async () => {
|
||||
if (!accessToken || !mcpServer.server_id) return;
|
||||
|
||||
|
|
@ -393,6 +490,10 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
delegate_auth_to_upstream: mcpServer.delegate_auth_to_upstream,
|
||||
}) === "passthrough";
|
||||
const isBrowserHeldTokenMode = isClientForwardedTokenMode(getEffectiveAuthType());
|
||||
|
||||
if (await previewWithStagedInteractiveToken(isPassthrough, isBrowserHeldTokenMode)) {
|
||||
return;
|
||||
}
|
||||
if (isPassthrough || isBrowserHeldTokenMode) {
|
||||
const token =
|
||||
oauthTokenResponse?.access_token ??
|
||||
|
|
@ -496,6 +597,9 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
stdio_config: undefined,
|
||||
});
|
||||
}
|
||||
if (isHeldOAuthTokenStale(form.getFieldsValue(true), authorizedIdentityRef.current)) {
|
||||
clearHeldOAuthToken();
|
||||
}
|
||||
};
|
||||
|
||||
const handleSave = async (values: Record<string, any>) => {
|
||||
|
|
@ -805,7 +909,13 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
</TabList>
|
||||
<TabPanels className="mt-6">
|
||||
<TabPanel>
|
||||
<Form form={form} onFinish={handleSave} initialValues={initialValues} layout="vertical">
|
||||
<Form
|
||||
form={form}
|
||||
onFinish={handleSave}
|
||||
onValuesChange={handleFormValuesChange}
|
||||
initialValues={initialValues}
|
||||
layout="vertical"
|
||||
>
|
||||
<Form.Item
|
||||
label="MCP Server Name"
|
||||
name="server_name"
|
||||
|
|
|
|||
|
|
@ -7,9 +7,36 @@ import {
|
|||
handleTransport,
|
||||
handleAuth,
|
||||
getMcpOAuthMode,
|
||||
getOAuthAuthorizationIdentity,
|
||||
isHeldOAuthTokenStale,
|
||||
oauth2FlowToFormValue,
|
||||
} from "./types";
|
||||
|
||||
describe("getOAuthAuthorizationIdentity", () => {
|
||||
// Regression: the identity used to pick the audience from spec_path only when values.transport was
|
||||
// OPENAPI, but the create form keeps transport in component state, so values.transport was absent and
|
||||
// spec_path edits on OpenAPI servers never invalidated a held token.
|
||||
it("changes when spec_path changes even when transport is absent from form values", () => {
|
||||
const authorized = { auth_type: AUTH_TYPE.OAUTH2, spec_path: "https://a.example.com/openapi.json" };
|
||||
const edited = { auth_type: AUTH_TYPE.OAUTH2, spec_path: "https://b.example.com/openapi.json" };
|
||||
expect(getOAuthAuthorizationIdentity(edited)).not.toBe(getOAuthAuthorizationIdentity(authorized));
|
||||
expect(isHeldOAuthTokenStale(edited, getOAuthAuthorizationIdentity(authorized))).toBe(true);
|
||||
});
|
||||
|
||||
it("changes when url changes", () => {
|
||||
const authorized = { auth_type: AUTH_TYPE.OAUTH2, url: "https://a.example.com/mcp" };
|
||||
const edited = { auth_type: AUTH_TYPE.OAUTH2, url: "https://b.example.com/mcp" };
|
||||
expect(getOAuthAuthorizationIdentity(edited)).not.toBe(getOAuthAuthorizationIdentity(authorized));
|
||||
});
|
||||
|
||||
it("is stable across non-mint fields", () => {
|
||||
const authorized = { auth_type: AUTH_TYPE.OAUTH2, url: "https://a.example.com/mcp", server_name: "one" };
|
||||
const renamed = { auth_type: AUTH_TYPE.OAUTH2, url: "https://a.example.com/mcp", server_name: "two" };
|
||||
expect(getOAuthAuthorizationIdentity(renamed)).toBe(getOAuthAuthorizationIdentity(authorized));
|
||||
expect(isHeldOAuthTokenStale(renamed, getOAuthAuthorizationIdentity(authorized))).toBe(false);
|
||||
});
|
||||
});
|
||||
|
||||
describe("handleTransport", () => {
|
||||
it("should default to SSE when transport is null", () => {
|
||||
expect(handleTransport(null)).toBe(TRANSPORT.SSE);
|
||||
|
|
|
|||
|
|
@ -57,6 +57,55 @@ export const OAUTH_FLOW = {
|
|||
M2M: "m2m",
|
||||
};
|
||||
|
||||
// The fields that determine which upstream OAuth token "Authorize & Fetch" mints: the resource/audience
|
||||
// (url, or spec_path for OpenAPI servers), the OAuth mode/grant (auth_type, oauth_flow_type), the OAuth
|
||||
// client and requested scope (credentials.client_id / client_secret / scopes), and the authorization-server
|
||||
// endpoints (authorization_url / token_url / registration_url). Grounded in RFC 8707 / RFC 8693 and the MCP
|
||||
// auth spec: an access token is bound to exactly this tuple (resource/audience + scope + client + issuer), so
|
||||
// a previously authorized token is stale if and only if this identity changes and must be re-minted.
|
||||
// url and spec_path are compared independently rather than selected by transport: the create form keeps
|
||||
// transport in component state, not in form values, so a transport-conditional target would silently pin the
|
||||
// audience to a missing url and never fire for spec_path edits on OpenAPI servers. Mirrors the backend's
|
||||
// mcp_oauth_token_identity. Deliberately EXCLUDES: transport itself (http<->sse on the same url is the same
|
||||
// audience; a switch to/from OpenAPI shows up as url/spec_path changes because each form clears the field the
|
||||
// new transport does not use), delegate_auth_to_upstream (a downstream-usage toggle that is never sent to the
|
||||
// authorize request), and all metadata/RBAC/routing fields. Shared by the create and edit forms so their
|
||||
// invalidation logic cannot drift.
|
||||
export const getOAuthAuthorizationIdentity = (values: Record<string, unknown>): string => {
|
||||
const credentials = (values.credentials ?? {}) as Record<string, unknown>;
|
||||
const identity = {
|
||||
url: typeof values.url === "string" ? values.url : null,
|
||||
spec_path: typeof values.spec_path === "string" ? values.spec_path : null,
|
||||
auth_type: values.auth_type ?? null,
|
||||
oauth_flow_type: values.oauth_flow_type ?? null,
|
||||
client_id: credentials.client_id ?? null,
|
||||
client_secret: credentials.client_secret ?? null,
|
||||
scopes: credentials.scopes ?? null,
|
||||
authorization_url: values.authorization_url ?? null,
|
||||
token_url: values.token_url ?? null,
|
||||
registration_url: values.registration_url ?? null,
|
||||
};
|
||||
return JSON.stringify(identity);
|
||||
};
|
||||
|
||||
// The form fields wiped when a held OAuth token is invalidated: only `credentials`, which holds the
|
||||
// minted material (the fetched token + DCR client). The authorization/token/registration endpoint
|
||||
// fields are deliberately NOT wiped: nothing programmatic ever writes them (upstream discovery happens
|
||||
// backend-side), so they only ever hold admin input, and resetting them would wipe it (create) or
|
||||
// silently revert it to the saved record (edit, whose Form has initialValues). Shared by the create and
|
||||
// edit forms so what gets wiped cannot drift.
|
||||
export const CLEARED_ON_INVALIDATION = ["credentials"] as const;
|
||||
|
||||
// True when a token was authorized in this session (authorizedIdentity recorded at mint time) and the
|
||||
// form's current identity no longer matches it. Every invalidation decision in both forms goes through
|
||||
// this single check: onValuesChange for user edits, and an explicit recheck after any programmatic
|
||||
// form.setFieldsValue (antd does not fire onValuesChange for those), so a missed event path cannot let a
|
||||
// stale token survive.
|
||||
export const isHeldOAuthTokenStale = (
|
||||
values: Record<string, unknown>,
|
||||
authorizedIdentity: string | undefined,
|
||||
): boolean => authorizedIdentity !== undefined && getOAuthAuthorizationIdentity(values) !== authorizedIdentity;
|
||||
|
||||
// Backend value of `oauth2_flow` that marks a machine-to-machine server. Distinct
|
||||
// from the UI-local OAUTH_FLOW.M2M ("m2m"); this is what the API actually returns.
|
||||
export const MCP_OAUTH2_FLOW_M2M = "client_credentials";
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue