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:
Tin 2026-07-09 12:41:15 -07:00
parent 05f39bf942
commit 48124734a0
9 changed files with 367 additions and 144 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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