mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): compare the token identity decrypted and invalidate every per-user token store
Review follow-ups on the stale-token invalidation. The backend identity now decrypts client_id and client_secret before comparing: the stored values are NaCl-encrypted with a fresh nonce on every write, so comparing ciphertext flagged every routine save as a mint-relevant change and purged per-user tokens that were still valid. The identity also gains spec_path, the audience for OpenAPI servers, and parses credentials stored as a JSON string The purge now routes each (user, server) through the manager's invalidate_user_oauth_token_cache, which becomes the single invalidation point covering both the legacy per-user token cache and the v2 per-user OAuth token store; previously the purge evicted only the legacy cache while the revoke path evicted only the v2 store, so each path left the other cache serving a replaced token until its TTL. A credential row racing in between the find and the delete is now detected via the delete_many count and logged; its cache entry expires by TTL On the dashboard, CLEARED_ON_INVALIDATION and the staleness check move to types.tsx as the single shared implementation for both forms. The edit form's transport handler now rechecks the identity after its programmatic setFieldsValue calls, which antd does not report through onValuesChange, so a token no longer survives a transport switch that clears the mint target. The create form rebuilds formValues from the post-reset form state after an invalidation instead of publishing the pre-reset snapshot, so the tool preview can no longer refetch with the discarded DCR client. Both transport handlers now share the recheck, which also stops the create form from over-invalidating on an http to sse swap that keeps the same url and therefore the same audience
This commit is contained in:
parent
05f39bf942
commit
48124734a0
9 changed files with 367 additions and 144 deletions
|
|
@ -1070,48 +1070,82 @@ async def list_user_oauth_credentials(
|
|||
return results
|
||||
|
||||
|
||||
def _decrypted_credential_field(creds: Dict[str, Any], field: str) -> Any:
|
||||
"""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: Any) -> tuple[Any, ...]:
|
||||
"""The upstream-OAuth-token-determining fields of an MCP server: the resource/audience (url), 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)."""
|
||||
"""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)
|
||||
creds_dict: Dict[str, Any] = creds if isinstance(creds, dict) else {}
|
||||
if isinstance(creds, str):
|
||||
try:
|
||||
parsed: Any = json.loads(creds)
|
||||
except ValueError:
|
||||
parsed = None
|
||||
else:
|
||||
parsed = creds
|
||||
creds_dict: Dict[str, Any] = 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),
|
||||
creds_dict.get("client_id"),
|
||||
creds_dict.get("client_secret"),
|
||||
_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) -> int:
|
||||
"""Delete every stored per-user OAuth credential for a server and drop each from the per-user token
|
||||
cache, 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."""
|
||||
"""Delete every stored per-user OAuth credential 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. A row inserted between the find and the delete is removed from the DB but cannot be
|
||||
evicted from the caches (its user_id was never seen); that case is detected, logged, and bounded
|
||||
by the cache TTL."""
|
||||
repo = MCPUserCredentialsRepository(prisma_client)
|
||||
rows = await repo.table.find_many(where={"server_id": server_id})
|
||||
if not rows:
|
||||
return 0
|
||||
await repo.table.delete_many(where={"server_id": server_id})
|
||||
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
|
||||
mcp_per_user_token_cache,
|
||||
deleted_count = await repo.table.delete_many(where={"server_id": server_id})
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
for row in rows:
|
||||
try:
|
||||
await mcp_per_user_token_cache.delete(row.user_id, server_id)
|
||||
except Exception as exc: # noqa: BLE001 - cache drop is best-effort; the DB delete is authoritative
|
||||
verbose_proxy_logger.warning(
|
||||
"Failed to drop cached MCP OAuth token for user=%s server=%s: %s", row.user_id, server_id, exc
|
||||
)
|
||||
return len(rows)
|
||||
await global_mcp_server_manager.invalidate_user_oauth_token_cache(row.user_id, server_id)
|
||||
if deleted_count != len(rows):
|
||||
verbose_proxy_logger.warning(
|
||||
"MCP server %s: purge removed %d credential row(s) but %d were enumerated; "
|
||||
"row(s) raced in during the purge and their cached tokens will expire by TTL",
|
||||
server_id,
|
||||
deleted_count,
|
||||
len(rows),
|
||||
)
|
||||
return deleted_count
|
||||
|
||||
|
||||
async def refresh_user_oauth_token(
|
||||
|
|
|
|||
|
|
@ -57,7 +57,10 @@ 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 (
|
||||
mcp_per_user_token_cache,
|
||||
resolve_mcp_auth,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
|
||||
Error,
|
||||
Ok,
|
||||
|
|
@ -4053,10 +4056,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 +4070,7 @@ class MCPServerManager:
|
|||
verbose_logger.warning(
|
||||
"Failed to invalidate cached MCP OAuth token for user=%s server=%s: %s", user_id, server_id, exc
|
||||
)
|
||||
await mcp_per_user_token_cache.delete(user_id, server_id)
|
||||
|
||||
async def _resolve_oauth2_headers_for_tool_call(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -84,6 +84,7 @@ def _identity_server(**overrides):
|
|||
"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"},
|
||||
|
|
@ -113,29 +114,91 @@ def test_mcp_oauth_token_identity_stable_on_non_mint_fields(overrides):
|
|||
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)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_purge_user_oauth_credentials_for_server_deletes_rows_and_cache(monkeypatch):
|
||||
from litellm.proxy._experimental.mcp_server import oauth2_token_cache
|
||||
async def test_purge_user_oauth_credentials_for_server_invalidates_every_store(monkeypatch):
|
||||
"""The purge must route each (user, server) through the manager's shared invalidation, which is
|
||||
the single point covering both the legacy per-user token cache and the v2 per-user OAuth token
|
||||
store; evicting only one cache lets the other keep serving a token minted for the old 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
|
||||
|
||||
r1 = MagicMock(user_id="alice", server_id="srv-1")
|
||||
r2 = MagicMock(user_id="bob", server_id="srv-1")
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[r1, r2])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock()
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=2)
|
||||
|
||||
cache_deletes = []
|
||||
invalidations = []
|
||||
monkeypatch.setattr(
|
||||
oauth2_token_cache.mcp_per_user_token_cache,
|
||||
"delete",
|
||||
AsyncMock(side_effect=lambda uid, sid: cache_deletes.append((uid, sid))),
|
||||
mcp_server_manager.global_mcp_server_manager,
|
||||
"invalidate_user_oauth_token_cache",
|
||||
AsyncMock(side_effect=lambda uid, sid: invalidations.append((uid, sid))),
|
||||
)
|
||||
|
||||
purged = await purge_user_oauth_credentials_for_server(prisma, "srv-1")
|
||||
|
||||
assert purged == 2
|
||||
prisma.db.litellm_mcpusercredentials.delete_many.assert_awaited_once()
|
||||
assert set(cache_deletes) == {("alice", "srv-1"), ("bob", "srv-1")}
|
||||
assert set(invalidations) == {("alice", "srv-1"), ("bob", "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 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=[MagicMock(user_id="alice", server_id="srv-1")]
|
||||
)
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=2)
|
||||
monkeypatch.setattr(
|
||||
mcp_server_manager.global_mcp_server_manager,
|
||||
"invalidate_user_oauth_token_cache",
|
||||
AsyncMock(),
|
||||
)
|
||||
warning = MagicMock()
|
||||
monkeypatch.setattr(db_module.verbose_proxy_logger, "warning", warning)
|
||||
|
||||
purged = await purge_user_oauth_credentials_for_server(prisma, "srv-1")
|
||||
|
||||
assert purged == 2
|
||||
warning.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -225,9 +288,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:
|
||||
|
|
@ -310,9 +371,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")
|
||||
|
|
@ -354,18 +413,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")
|
||||
|
||||
|
|
@ -415,9 +470,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"]
|
||||
|
|
@ -445,19 +498,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(
|
||||
|
|
@ -485,14 +532,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
|
||||
|
|
@ -508,9 +551,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
|
||||
|
||||
|
||||
|
|
@ -527,12 +568,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():
|
||||
|
|
@ -554,9 +590,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()
|
||||
)
|
||||
|
|
@ -572,15 +606,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()
|
||||
)
|
||||
|
|
@ -599,9 +629,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()
|
||||
)
|
||||
|
|
@ -635,9 +663,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()
|
||||
)
|
||||
|
|
@ -654,9 +680,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 (
|
||||
|
|
@ -690,19 +714,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)
|
||||
|
|
@ -719,18 +737,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"]
|
||||
|
|
@ -758,9 +772,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(),
|
||||
|
|
@ -799,9 +811,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(),
|
||||
|
|
|
|||
|
|
@ -3328,8 +3328,34 @@ class TestMCPServerManager:
|
|||
assert store.invalidations == [("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."""
|
||||
async def test_invalidate_user_oauth_token_cache_drops_legacy_cache_too(self, monkeypatch):
|
||||
"""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."""
|
||||
from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module
|
||||
|
||||
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
|
||||
|
||||
legacy_deletes: list[tuple[str, str]] = []
|
||||
monkeypatch.setattr(
|
||||
manager_module.mcp_per_user_token_cache,
|
||||
"delete",
|
||||
AsyncMock(side_effect=lambda uid, sid: legacy_deletes.append((uid, sid))),
|
||||
)
|
||||
manager = MCPServerManager(per_user_oauth_token_store=_Store())
|
||||
await manager.invalidate_user_oauth_token_cache("alice", "srv-1")
|
||||
assert legacy_deletes == [("alice", "srv-1")]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalidate_user_oauth_token_cache_swallows_store_errors(self, monkeypatch):
|
||||
"""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."""
|
||||
from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module
|
||||
|
||||
class _Store:
|
||||
async def fetch(self, user_id: str, server_id: str):
|
||||
|
|
@ -3338,8 +3364,15 @@ class TestMCPServerManager:
|
|||
async def invalidate(self, user_id: str, server_id: str) -> None:
|
||||
raise RuntimeError("redis down")
|
||||
|
||||
legacy_deletes: list[tuple[str, str]] = []
|
||||
monkeypatch.setattr(
|
||||
manager_module.mcp_per_user_token_cache,
|
||||
"delete",
|
||||
AsyncMock(side_effect=lambda uid, sid: legacy_deletes.append((uid, sid))),
|
||||
)
|
||||
manager = MCPServerManager(per_user_oauth_token_store=_Store())
|
||||
await manager.invalidate_user_oauth_token_cache("alice", "srv-1")
|
||||
assert legacy_deletes == [("alice", "srv-1")]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_oauth2_headers_no_user_id(self):
|
||||
|
|
|
|||
|
|
@ -719,6 +719,56 @@ describe("CreateMCPServer", () => {
|
|||
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");
|
||||
expect(call[1]?.credentials?.access_token).not.toBe("stale-tok");
|
||||
}
|
||||
});
|
||||
|
||||
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",
|
||||
|
|
|
|||
|
|
@ -16,6 +16,8 @@ import {
|
|||
MCP_OAUTH2_FLOW_INTERACTIVE,
|
||||
isClientForwardedTokenMode,
|
||||
getOAuthAuthorizationIdentity,
|
||||
CLEARED_ON_INVALIDATION,
|
||||
isHeldOAuthTokenStale,
|
||||
} from "./types";
|
||||
import OAuthFormFields from "./OAuthFormFields";
|
||||
import TruePassthroughWarning from "./TruePassthroughWarning";
|
||||
|
|
@ -235,11 +237,9 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
|
|||
});
|
||||
|
||||
// Discard the held browser-authorized token and its tool preview when the authorization identity
|
||||
// changes (or the modal closes). For oauth2 the fetched token + DCR client also live in
|
||||
// form.credentials, and the discovered endpoints in authorization_url/token_url/registration_url, so
|
||||
// those form fields are reset too; whatever the admin just changed (passed via changedValues) is
|
||||
// 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.
|
||||
const CLEARED_ON_INVALIDATION = ["credentials", "authorization_url", "token_url", "registration_url"] as const;
|
||||
const clearHeldOAuthToken = (changedValues: Record<string, unknown> = {}) => {
|
||||
setOauthAccessToken(null);
|
||||
clearTools();
|
||||
|
|
@ -588,21 +588,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 =
|
||||
authorizedIdentity === undefined
|
||||
? transportValues
|
||||
: {
|
||||
...transportValues,
|
||||
credentials: undefined,
|
||||
authorization_url: undefined,
|
||||
token_url: undefined,
|
||||
registration_url: undefined,
|
||||
};
|
||||
|
||||
form.setFieldsValue(nextValues);
|
||||
if (authorizedIdentity !== undefined) {
|
||||
form.setFieldsValue(transportValues);
|
||||
if (isHeldOAuthTokenStale(form.getFieldsValue(true), authorizedIdentity)) {
|
||||
clearHeldOAuthToken();
|
||||
}
|
||||
setFormValues(form.getFieldsValue(true));
|
||||
};
|
||||
|
||||
// Generate options with existing groups and potential new group
|
||||
|
|
@ -672,9 +662,13 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
|
|||
const handleFormValuesChange = (changedValues: Record<string, unknown>, allValues: Record<string, unknown>) => {
|
||||
// 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.
|
||||
if (authorizedIdentity !== undefined && getOAuthAuthorizationIdentity(allValues) !== authorizedIdentity) {
|
||||
// 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);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -22,15 +22,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 +99,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 +460,74 @@ 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("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();
|
||||
|
|
|
|||
|
|
@ -6,6 +6,8 @@ import {
|
|||
AUTH_TYPE,
|
||||
isClientForwardedTokenMode,
|
||||
getOAuthAuthorizationIdentity,
|
||||
CLEARED_ON_INVALIDATION,
|
||||
isHeldOAuthTokenStale,
|
||||
OAUTH_FLOW,
|
||||
MCP_OAUTH2_FLOW_M2M,
|
||||
MCP_OAUTH2_FLOW_INTERACTIVE,
|
||||
|
|
@ -392,11 +394,12 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
// 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 form.credentials + the
|
||||
// discovered endpoint 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 (user input only), never programmatic resets.
|
||||
const CLEARED_ON_INVALIDATION = ["credentials", "authorization_url", "token_url", "registration_url"] as const;
|
||||
// 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) {
|
||||
|
|
@ -413,10 +416,7 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
};
|
||||
|
||||
const handleFormValuesChange = (changedValues: Record<string, unknown>) => {
|
||||
if (
|
||||
authorizedIdentityRef.current !== undefined &&
|
||||
getOAuthAuthorizationIdentity(form.getFieldsValue(true)) !== authorizedIdentityRef.current
|
||||
) {
|
||||
if (isHeldOAuthTokenStale(form.getFieldsValue(true), authorizedIdentityRef.current)) {
|
||||
clearHeldOAuthToken(changedValues);
|
||||
}
|
||||
};
|
||||
|
|
@ -539,6 +539,9 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
stdio_config: undefined,
|
||||
});
|
||||
}
|
||||
if (isHeldOAuthTokenStale(form.getFieldsValue(true), authorizedIdentityRef.current)) {
|
||||
clearHeldOAuthToken();
|
||||
}
|
||||
};
|
||||
|
||||
const handleSave = async (values: Record<string, any>) => {
|
||||
|
|
|
|||
|
|
@ -84,6 +84,21 @@ export const getOAuthAuthorizationIdentity = (values: Record<string, unknown>):
|
|||
return JSON.stringify(identity);
|
||||
};
|
||||
|
||||
// The form fields wiped when a held OAuth token is invalidated: the fetched token + DCR client live in
|
||||
// `credentials`, and the three endpoint fields were discovered by the authorize flow, so all of them are
|
||||
// stale together with the token. Shared by the create and edit forms so what gets wiped cannot drift.
|
||||
export const CLEARED_ON_INVALIDATION = ["credentials", "authorization_url", "token_url", "registration_url"] 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