fix(mcp): persist config.yaml DCR clients in a server-scoped store

Config.yaml-declared OAuth2 MCP servers using Dynamic Client Registration have no LiteLLM_MCPServerTable row, so the DCR persist path called update_mcp_server, which returns None for a missing row, then update_server(None), which dereferenced .approval_status and raised AttributeError. The exception was swallowed to a warning while /register still returned 200, so the minted client was never stored and every access-token expiry forced a full re-authorization

Persist the acquired DCR client (client_id, client_secret, token_endpoint_auth_method, redirect_uris, encrypted at rest) in a dedicated LiteLLM_MCPServerOAuthClient store keyed by server_id when the server has no row, overlay it onto the in-memory config server so the refresh_token grant can authenticate within the process, and rehydrate it when the registry syncs from the database (which runs after the DB connects, unlike config load) so restarts and other pods pick it up. The store is encrypted at rest and is re-encrypted by the master-key rotation path alongside the server rows, through a shared helper so the two sites cannot diverge. The DB-backed server path is unchanged, and guarding the None return removes the swallowed-crash footgun

Resolves the config.yaml DCR persistence regression introduced in v1.92.0 by #31912
This commit is contained in:
Tin Chi Lo 2026-07-17 12:10:12 -07:00
parent e5a9f3f5d7
commit 99b85a3f2c
11 changed files with 679 additions and 51 deletions

View file

@ -0,0 +1,9 @@
-- CreateTable
CREATE TABLE IF NOT EXISTS "LiteLLM_MCPServerOAuthClient" (
"server_id" TEXT NOT NULL,
"credentials" JSONB,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
CONSTRAINT "LiteLLM_MCPServerOAuthClient_pkey" PRIMARY KEY ("server_id")
);

View file

@ -396,6 +396,13 @@ model LiteLLM_MCPUserEnvVars {
@@index([server_id])
}
model LiteLLM_MCPServerOAuthClient {
server_id String @id
credentials Json?
created_at DateTime @default(now()) @map("created_at")
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
}
// Generate Tokens for Proxy
model LiteLLM_VerificationToken {
token String @id

View file

@ -33,6 +33,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
from litellm.proxy.utils import PrismaClient
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
from litellm.repositories.table_repositories import (
MCPServerOAuthClientRepository,
MCPServerRepository,
MCPUserCredentialsRepository,
)
@ -639,6 +640,7 @@ async def delete_mcp_server(
for model, label in (
(prisma_client.db.litellm_mcpusercredentials, "credential"),
(prisma_client.db.litellm_mcpuserenvvars, "env var"),
(prisma_client.db.litellm_mcpserveroauthclient, "OAuth client"),
):
try:
await model.delete_many(where={"server_id": server_id})
@ -823,26 +825,66 @@ async def update_mcp_server(
return updated_mcp_server
async def rotate_mcp_server_credentials_master_key(prisma_client: PrismaClient, touched_by: str, new_master_key: str):
async def get_mcp_server_oauth_client_credentials(prisma_client: PrismaClient, server_id: str) -> object | None:
"""Read the persisted (encrypted) DCR OAuth client blob for a server from the
server-scoped store, or None. Config.yaml-declared servers have no
LiteLLM_MCPServerTable row, so their dynamically registered client lives here keyed
by server_id. The returned value is the raw credentials blob for
``_get_persisted_dcr_credentials`` to parse."""
row = await MCPServerOAuthClientRepository(prisma_client).table.find_unique(where={"server_id": server_id})
if row is None:
return None
return row.credentials
async def upsert_mcp_server_oauth_client_credentials(
prisma_client: PrismaClient, server_id: str, credentials: MCPCredentials
) -> None:
"""Persist a server's dynamically registered OAuth client (RFC 7591 DCR) in the
server-scoped store keyed by server_id, independent of any LiteLLM_MCPServerTable row.
client_id/client_secret are encrypted at rest with the same salt key used for the
server row's credentials blob, so ``_apply_persisted_dcr_credentials`` decrypts them the
same way regardless of which store a server's client came from."""
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
encrypted = encrypt_credentials(credentials=dict(credentials), encryption_key=_get_salt_key())
blob = safe_dumps(encrypted)
await MCPServerOAuthClientRepository(prisma_client).table.upsert(
where={"server_id": server_id},
data={
"create": {"server_id": server_id, "credentials": blob},
"update": {"credentials": blob},
},
)
def _reencrypt_mcp_credentials_blob(credentials: object, new_master_key: str) -> str | None:
"""Decrypt an at-rest MCP credentials blob with the current key and re-encrypt it under
new_master_key, returning the serialized blob or None when there is nothing to rotate. Shared by
every table that stores an encrypted MCP credentials blob so a master-key rotation covers them
uniformly and cannot silently skip one."""
if not credentials:
return None
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps # noqa: PLC0415 # avoids circular import
creds_dict = json.loads(credentials) if isinstance(credentials, str) else dict(credentials)
decrypted = decrypt_credentials(credentials=cast(MCPCredentials, creds_dict))
encrypted = encrypt_credentials(credentials=decrypted, encryption_key=new_master_key)
return safe_dumps(encrypted)
async def rotate_mcp_server_credentials_master_key(prisma_client: PrismaClient, touched_by: str, new_master_key: str):
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps # noqa: PLC0415 # avoids circular import
mcp_servers = await MCPServerRepository(prisma_client).table.find_many()
updated = 0
for mcp_server in mcp_servers:
update_data: Dict[str, Any] = {}
credentials = mcp_server.credentials
if credentials:
# Decrypt with current key first, then re-encrypt with new key
decrypted_credentials = decrypt_credentials(
credentials=cast(MCPCredentials, dict(credentials)),
)
encrypted_credentials = encrypt_credentials(
credentials=decrypted_credentials,
encryption_key=new_master_key,
)
update_data["credentials"] = safe_dumps(encrypted_credentials)
rotated_credentials = _reencrypt_mcp_credentials_blob(mcp_server.credentials, new_master_key)
if rotated_credentials is not None:
update_data["credentials"] = rotated_credentials
rotated_env_vars = _reencrypt_global_env_var_values(mcp_server.env_vars, new_master_key)
if rotated_env_vars is not None:
@ -857,9 +899,23 @@ async def rotate_mcp_server_credentials_master_key(prisma_client: PrismaClient,
data=update_data,
)
updated += 1
oauth_clients = await MCPServerOAuthClientRepository(prisma_client).table.find_many()
oauth_updated = 0
for oauth_client in oauth_clients:
rotated_credentials = _reencrypt_mcp_credentials_blob(oauth_client.credentials, new_master_key)
if rotated_credentials is None:
continue
await MCPServerOAuthClientRepository(prisma_client).table.update(
where={"server_id": oauth_client.server_id},
data={"credentials": rotated_credentials},
)
oauth_updated += 1
verbose_proxy_logger.info(
"rotate_mcp_server_credentials_master_key: rotated %d MCP server row(s)",
"rotate_mcp_server_credentials_master_key: rotated %d MCP server row(s) and %d OAuth-client row(s)",
updated,
oauth_updated,
)

View file

@ -971,43 +971,93 @@ def _apply_persisted_dcr_credentials(mcp_server: MCPServer, credentials: _Persis
return True
async def _get_persisted_mcp_server_with_dcr_client_id(
mcp_server: MCPServer,
) -> Optional[tuple["LiteLLM_MCPServerTable", _PersistedDcrCredentials]]:
from litellm.proxy._experimental.mcp_server.db import get_mcp_server # noqa: PLC0415
from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415
async def _load_store_dcr_credentials(mcp_server: MCPServer) -> _PersistedDcrCredentials | None:
"""DCR client persisted in the server-scoped OAuth-client store for a config-declared server
(which has no LiteLLM_MCPServerTable row). Returns None when the store has no usable client_id
or the DB is unreachable."""
from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # avoids circular import
get_mcp_server_oauth_client_credentials,
)
from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415 # avoids circular import
try:
prisma_client = get_prisma_client_or_throw("Database not connected. Cannot read MCP OAuth client registration.")
persisted_mcp_server = await get_mcp_server(
prisma_client=prisma_client,
server_id=mcp_server.server_id,
blob = await get_mcp_server_oauth_client_credentials(
prisma_client=prisma_client, server_id=mcp_server.server_id
)
except Exception as exc: # noqa: BLE001
except Exception as exc: # noqa: BLE001 # best-effort read; DB may be unreachable
verbose_logger.debug(
"register_client_with_server: failed to read persisted DCR client registration for server_id=%s: %s",
"register_client_with_server: failed to read stored DCR client for server_id=%s: %s",
mcp_server.server_id,
exc,
)
return None
if persisted_mcp_server is None:
return None
credentials = _get_persisted_dcr_credentials(persisted_mcp_server.credentials)
credentials = _get_persisted_dcr_credentials(blob)
if credentials is None or not credentials.client_id:
return None
return credentials
return persisted_mcp_server, credentials
async def hydrate_config_server_dcr_client(mcp_server: MCPServer) -> bool:
"""Overlay a config-declared server's persisted DCR client onto its in-memory object so token
refresh can authenticate. Config.yaml servers have no LiteLLM_MCPServerTable row, so their
minted client lives in the server-scoped store; without this overlay the in-memory server
carries no client_id after a restart. An explicit client_id set in config.yaml wins and is never
overwritten by a persisted store client."""
if mcp_server.client_id:
return False
credentials = await _load_store_dcr_credentials(mcp_server)
if credentials is None:
return False
return _apply_persisted_dcr_credentials(mcp_server, credentials)
async def _resolve_persisted_dcr_client(
mcp_server: MCPServer,
) -> tuple[Optional["LiteLLM_MCPServerTable"], _PersistedDcrCredentials | None]:
"""Resolve a server's persisted DCR client using the same two-level rule the write path uses, so
read and write always agree. First, whether the server HAS a LiteLLM_MCPServerTable row: a row is
always resolved to that row and the store is never consulted for a server that has a row, so a
caller-chosen server_id colliding with a config-declared server cannot inherit that config
server's client, and a row that exists but carries no usable client_id yields (row, None) rather
than a store fallback. Second, among rowless servers: a config-declared server keeps its client in
the server-scoped store, while a rowless non-config server is a throwaway temp/session server with
no persisted client. Returns (row_or_None, credentials_or_None); the row is only needed by the
reuse path to refresh the registry for a DB-declared server."""
from litellm.proxy._experimental.mcp_server.db import get_mcp_server # noqa: PLC0415 # avoids circular import
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # avoids circular import
global_mcp_server_manager,
)
from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415 # avoids circular import
try:
prisma_client = get_prisma_client_or_throw("Database not connected. Cannot read MCP OAuth client registration.")
row = await get_mcp_server(prisma_client=prisma_client, server_id=mcp_server.server_id)
except Exception as exc: # noqa: BLE001 # best-effort read; DB may be unreachable
verbose_logger.debug(
"register_client_with_server: failed to read persisted DCR client for server_id=%s: %s",
mcp_server.server_id,
exc,
)
return None, None
if row is not None:
credentials = _get_persisted_dcr_credentials(row.credentials)
if credentials is not None and credentials.client_id:
return row, credentials
return row, None
if global_mcp_server_manager.is_config_declared_server(mcp_server.server_id):
return None, await _load_store_dcr_credentials(mcp_server)
return None, None
async def _reuse_persisted_dcr_client_if_available(
mcp_server: MCPServer, current_redirect_uri: Optional[str] = None
) -> bool:
persisted = await _get_persisted_mcp_server_with_dcr_client_id(mcp_server)
if persisted is None:
persisted_mcp_server, credentials = await _resolve_persisted_dcr_client(mcp_server)
if credentials is None:
return False
persisted_mcp_server, credentials = persisted
if current_redirect_uri is not None and _redirect_uri_not_registered(credentials, current_redirect_uri):
verbose_logger.debug(
"register_client_with_server: not reusing persisted DCR client for server_id=%s; its registered "
@ -1021,18 +1071,19 @@ async def _reuse_persisted_dcr_client_if_available(
if not _apply_persisted_dcr_credentials(mcp_server, credentials):
return False
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
global_mcp_server_manager,
)
try:
await global_mcp_server_manager.update_server(persisted_mcp_server)
except Exception as exc: # noqa: BLE001
verbose_logger.warning(
"register_client_with_server: failed to refresh persisted DCR client registration for server_id=%s: %s",
mcp_server.server_id,
exc,
if persisted_mcp_server is not None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # avoids circular import
global_mcp_server_manager,
)
try:
await global_mcp_server_manager.update_server(persisted_mcp_server)
except Exception as exc: # noqa: BLE001 # best-effort registry refresh
verbose_logger.warning(
"register_client_with_server: failed to refresh persisted DCR client registration for server_id=%s: %s",
mcp_server.server_id,
exc,
)
return bool(mcp_server.client_id)
@ -1044,10 +1095,9 @@ async def _persisted_dcr_redirect_uri_is_stale(mcp_server: MCPServer, current_re
otherwise short-circuits registration before any redirect check can run. Servers
without a persisted DCR recording (admin-configured client_id, or registered before
redirect_uris were recorded) are never reported stale."""
persisted = await _get_persisted_mcp_server_with_dcr_client_id(mcp_server)
if persisted is None:
_, credentials = await _resolve_persisted_dcr_client(mcp_server)
if credentials is None:
return False
_, credentials = persisted
if not _redirect_uri_not_registered(credentials, current_redirect_uri):
return False
verbose_logger.warning(
@ -1067,7 +1117,10 @@ DcrRegistrationPersistenceResult = Literal["persisted", "reused", "skipped", "fa
async def _persist_dcr_client_registration(
mcp_server: MCPServer, registration_response: object, current_redirect_uri: str
) -> DcrRegistrationPersistenceResult:
"""Persist the dynamically registered OAuth client (RFC 7591) onto the MCP server row.
"""Persist the dynamically registered OAuth client (RFC 7591) to its single home: the server's
``LiteLLM_MCPServerTable`` row when it has one, otherwise the server-scoped store when the server
is config-declared. A rowless server that is not config-declared is a throwaway temp/session
server, so its client is overlaid in memory only and not persisted.
The interactive authorization_code flow mints a ``client_id`` via Dynamic Client
Registration that discovery cannot re-derive; without persisting it the autonomous
@ -1106,16 +1159,20 @@ async def _persist_dcr_client_registration(
if await _reuse_persisted_dcr_client_if_available(mcp_server, current_redirect_uri=current_redirect_uri):
return "reused"
token_endpoint_auth_method = (
"client_secret_basic" if registration.token_endpoint_auth_method == "client_secret_basic" else None
)
credentials: MCPCredentials = {
"client_id": registration.client_id,
"client_secret": registration.client_secret,
"token_endpoint_auth_method": (
"client_secret_basic" if registration.token_endpoint_auth_method == "client_secret_basic" else None
),
"token_endpoint_auth_method": token_endpoint_auth_method,
"redirect_uris": [current_redirect_uri],
}
from litellm.proxy._experimental.mcp_server.db import update_mcp_server # noqa: PLC0415
from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # avoids circular import
update_mcp_server,
upsert_mcp_server_oauth_client_credentials,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
global_mcp_server_manager,
)
@ -1136,7 +1193,18 @@ async def _persist_dcr_client_registration(
),
touched_by="mcp_oauth_dcr",
)
await global_mcp_server_manager.update_server(updated_row)
if updated_row is not None:
await global_mcp_server_manager.update_server(updated_row)
return "persisted"
if global_mcp_server_manager.is_config_declared_server(mcp_server.server_id):
await upsert_mcp_server_oauth_client_credentials(
prisma_client=prisma_client,
server_id=mcp_server.server_id,
credentials=credentials,
)
mcp_server.client_id = registration.client_id
mcp_server.client_secret = registration.client_secret
mcp_server.token_endpoint_auth_method = token_endpoint_auth_method
return "persisted"
except Exception as exc: # noqa: BLE001
verbose_logger.warning(

View file

@ -1127,6 +1127,14 @@ class MCPServerManager:
"""
return self.config_mcp_servers | self.registry
def is_config_declared_server(self, server_id: str) -> bool:
"""True when server_id was declared in config.yaml (present in the in-memory config map).
Config servers are rowless and persistent, so their DCR client belongs in the server-scoped
store; a rowless server that is NOT config-declared is a throwaway temp/session server whose
client must not be persisted. This never overrides the row-existence check: a server that has
a LiteLLM_MCPServerTable row is always resolved to that row first."""
return server_id in self.config_mcp_servers
async def load_servers_from_config(
self,
mcp_servers_config: dict[str, Any],
@ -1367,8 +1375,32 @@ class MCPServerManager:
verbose_logger.debug(f"Loaded MCP Servers: {json.dumps(self.config_mcp_servers, indent=4, default=str)}")
await self._hydrate_config_servers_dcr_clients()
self.initialize_tool_name_to_mcp_server_name_mapping()
async def _hydrate_config_servers_dcr_clients(self) -> None:
"""Overlay each config-declared server's persisted DCR client (from the server-scoped
store) onto its in-memory object so token refresh authenticates after a restart. A
best-effort no-op when the DB is unreachable at config-load time."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( # noqa: PLC0415 # circular import
hydrate_config_server_dcr_client,
)
for server in self.config_mcp_servers.values():
try:
if await hydrate_config_server_dcr_client(server):
verbose_logger.debug(
"hydrated persisted DCR client onto config MCP server server_id=%s",
server.server_id,
)
except Exception as exc: # noqa: BLE001 # best-effort hydration; never fail config load
verbose_logger.debug(
"load_servers_from_config: failed to hydrate DCR client for server_id=%s: %s",
server.server_id,
exc,
)
async def _register_openapi_tools(self, spec_path: str, server: MCPServer, base_url: str):
"""
Register tools from an OpenAPI specification for a given server.
@ -4968,6 +5000,8 @@ class MCPServerManager:
verbose_logger.debug("MCP registry refreshed (%s servers in registry)", len(registered_registry))
await self._hydrate_config_servers_dcr_clients()
def get_mcp_servers_from_ids(self, server_ids: list[str]) -> list[MCPServer]:
servers = []
registry = self.get_registry()

View file

@ -396,6 +396,13 @@ model LiteLLM_MCPUserEnvVars {
@@index([server_id])
}
model LiteLLM_MCPServerOAuthClient {
server_id String @id
credentials Json?
created_at DateTime @default(now()) @map("created_at")
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
}
// Generate Tokens for Proxy
model LiteLLM_VerificationToken {
token String @id

View file

@ -77,6 +77,10 @@ class MCPUserCredentialsRepository(PrismaTableRepository):
table_name = "litellm_mcpusercredentials"
class MCPServerOAuthClientRepository(PrismaTableRepository):
table_name = "litellm_mcpserveroauthclient"
class PromptRepository(PrismaTableRepository):
table_name = "litellm_prompttable"

View file

@ -396,6 +396,13 @@ model LiteLLM_MCPUserEnvVars {
@@index([server_id])
}
model LiteLLM_MCPServerOAuthClient {
server_id String @id
credentials Json?
created_at DateTime @default(now()) @map("created_at")
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
}
// Generate Tokens for Proxy
model LiteLLM_VerificationToken {
token String @id

View file

@ -978,3 +978,67 @@ def test_prepare_mcp_server_data_update_carries_token_exchange_columns():
assert data["audience"] == "https://upstream.example.com"
assert data["subject_token_type"] == "urn:ietf:params:oauth:token-type:jwt"
assert data["token_exchange_profile"] == "entra_obo"
@pytest.mark.asyncio
async def test_master_key_rotation_reencrypts_oauth_client_store(monkeypatch):
"""The server-scoped DCR client store (LiteLLM_MCPServerOAuthClient) is encrypted at rest, so a
master-key rotation must re-encrypt it alongside the server rows. Skipping it leaves
config-declared DCR clients under the retired key, where they decrypt back to ciphertext and
force a full re-authorization."""
import litellm.proxy.common_utils.encrypt_decrypt_utils as enc
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._experimental.mcp_server.db import (
decrypt_credentials,
encrypt_credentials,
rotate_mcp_server_credentials_master_key,
)
key_old, key_new = "salt-old-key", "salt-new-key"
blob_old = safe_dumps(
encrypt_credentials(
credentials={"client_id": "cid-123", "client_secret": "sec-456"},
encryption_key=key_old,
)
)
monkeypatch.setattr(enc, "_get_salt_key", lambda: key_old)
prisma = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(
return_value=[SimpleNamespace(server_id="config_faros", credentials=blob_old)]
)
store_update = AsyncMock()
prisma.db.litellm_mcpserveroauthclient.update = store_update
await rotate_mcp_server_credentials_master_key(prisma, touched_by="test", new_master_key=key_new)
store_update.assert_awaited_once()
assert store_update.await_args.kwargs["where"] == {"server_id": "config_faros"}
rotated_blob = store_update.await_args.kwargs["data"]["credentials"]
monkeypatch.setattr(enc, "_get_salt_key", lambda: key_new)
recovered = decrypt_credentials(credentials=json.loads(rotated_blob))
assert recovered["client_id"] == "cid-123"
assert recovered["client_secret"] == "sec-456"
@pytest.mark.asyncio
async def test_delete_mcp_server_cleans_oauth_client_store():
"""Deleting a server must remove its server-scoped DCR client store entry alongside the per-user
credential and env-var rows, or a re-created server reusing the same server_id would inherit the
deleted server's OAuth client."""
from litellm.proxy._experimental.mcp_server.db import delete_mcp_server
prisma = MagicMock()
prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=SimpleNamespace(server_id="s1"))
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[])
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock()
prisma.db.litellm_mcpuserenvvars.delete_many = AsyncMock()
prisma.db.litellm_mcpserveroauthclient.delete_many = AsyncMock()
await delete_mcp_server(prisma, "s1", invalidate_token_cache=AsyncMock())
prisma.db.litellm_mcpserveroauthclient.delete_many.assert_awaited_once_with(where={"server_id": "s1"})

View file

@ -7130,3 +7130,373 @@ async def test_token_exchange_unreadable_body_still_renders_oauth_fault():
assert response.status_code == 502
body = json.loads(response.body)
assert body == {"error": "server_error", "error_description": "upstream token endpoint returned HTTP 400"}
@pytest.mark.asyncio
async def test_persist_dcr_client_for_config_server_uses_side_store():
"""A config.yaml-declared OAuth2 DCR server has no LiteLLM_MCPServerTable row, so
update_mcp_server returns None. The minted client must then persist to the server-scoped
OAuth-client store keyed by server_id (never a shadow server row), overlay onto the in-memory
server so refresh can authenticate this process, and never call update_server(None) (which
previously raised AttributeError on .approval_status, was swallowed, and reported a 200 that
persisted nothing)."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_persist_dcr_client_registration,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
config_server = MCPServer(
server_id="config_faros",
name="config_faros",
server_name="config_faros",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
client_id=None,
client_secret=None,
authorization_url="https://provider.example/oauth/authorize",
token_url="https://provider.example/oauth/token",
registration_url="https://provider.example/oauth/register",
)
mock_upsert = AsyncMock()
mock_update_server = AsyncMock()
with (
patch.object(global_mcp_server_manager, "is_config_declared_server", return_value=True),
patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
patch(
"litellm.proxy._experimental.mcp_server.db.update_mcp_server",
new=AsyncMock(return_value=None),
),
patch(
"litellm.proxy._experimental.mcp_server.db.get_mcp_server",
new=AsyncMock(return_value=None),
),
patch(
"litellm.proxy._experimental.mcp_server.db.get_mcp_server_oauth_client_credentials",
new=AsyncMock(return_value=None),
),
patch(
"litellm.proxy._experimental.mcp_server.db.upsert_mcp_server_oauth_client_credentials",
new=mock_upsert,
),
patch.object(global_mcp_server_manager, "update_server", new=mock_update_server),
):
result = await _persist_dcr_client_registration(
mcp_server=config_server,
registration_response={
"client_id": "minted-client",
"client_secret": "minted-secret",
"token_endpoint_auth_method": "client_secret_basic",
},
current_redirect_uri="https://proxy.litellm.example/callback",
)
assert result == "persisted"
mock_upsert.assert_called_once()
assert mock_upsert.call_args.kwargs["server_id"] == "config_faros"
stored = mock_upsert.call_args.kwargs["credentials"]
assert stored["client_id"] == "minted-client"
assert stored["client_secret"] == "minted-secret"
assert stored["token_endpoint_auth_method"] == "client_secret_basic"
assert stored["redirect_uris"] == ["https://proxy.litellm.example/callback"]
assert config_server.client_id == "minted-client"
assert config_server.client_secret == "minted-secret"
assert config_server.token_endpoint_auth_method == "client_secret_basic"
mock_update_server.assert_not_called()
@pytest.mark.asyncio
async def test_hydrate_config_server_applies_stored_dcr_client(monkeypatch):
"""On restart a config server's in-memory object has no client_id; hydration overlays the
persisted DCR client from the server-scoped store, decrypting the encrypted-at-rest blob, so the
refresh_token grant can authenticate as the registered client instead of re-authenticating."""
import litellm.proxy.common_utils.encrypt_decrypt_utils as enc
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._experimental.mcp_server.db import encrypt_credentials
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
hydrate_config_server_dcr_client,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(
server_id="config_faros",
name="config_faros",
server_name="config_faros",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
client_id=None,
)
monkeypatch.setattr(enc, "_get_salt_key", lambda: "salt-hydrate-key")
stored_blob = safe_dumps(
encrypt_credentials(
credentials={
"client_id": "stored-client",
"client_secret": "stored-secret",
"token_endpoint_auth_method": "client_secret_basic",
"redirect_uris": ["https://proxy.litellm.example/callback"],
},
encryption_key="salt-hydrate-key",
)
)
assert "stored-client" not in stored_blob and "stored-secret" not in stored_blob
with (
patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
patch(
"litellm.proxy._experimental.mcp_server.db.get_mcp_server_oauth_client_credentials",
new=AsyncMock(return_value=stored_blob),
),
):
applied = await hydrate_config_server_dcr_client(server)
assert applied is True
assert server.client_id == "stored-client"
assert server.client_secret == "stored-secret"
assert server.token_endpoint_auth_method == "client_secret_basic"
@pytest.mark.asyncio
async def test_reuse_config_server_reads_store_with_real_crypto(monkeypatch):
"""A config-declared server (rowless) keeps its DCR client in the store, so the reuse read
resolves it from the store and decrypts the encrypted-at-rest client, mirroring the write path so
a re-authorize reuses the client instead of re-minting one."""
import litellm.proxy.common_utils.encrypt_decrypt_utils as enc
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._experimental.mcp_server.db import encrypt_credentials
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_reuse_persisted_dcr_client_if_available,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(
server_id="config_faros",
name="config_faros",
server_name="config_faros",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
client_id=None,
)
monkeypatch.setattr(enc, "_get_salt_key", lambda: "salt-reuse-key")
blob = safe_dumps(
encrypt_credentials(
credentials={"client_id": "stored-client", "client_secret": "sec", "redirect_uris": ["https://x/callback"]},
encryption_key="salt-reuse-key",
)
)
assert "stored-client" not in blob
store_lookup = AsyncMock(return_value=blob)
with (
patch.object(global_mcp_server_manager, "is_config_declared_server", return_value=True),
patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
patch("litellm.proxy._experimental.mcp_server.db.get_mcp_server", new=AsyncMock(return_value=None)),
patch(
"litellm.proxy._experimental.mcp_server.db.get_mcp_server_oauth_client_credentials",
new=store_lookup,
),
):
result = await _reuse_persisted_dcr_client_if_available(server, current_redirect_uri="https://x/callback")
assert result is True
assert server.client_id == "stored-client"
store_lookup.assert_awaited_once()
@pytest.mark.asyncio
async def test_temp_server_is_not_persisted_to_store():
"""A rowless server that is NOT config-declared (a throwaway /server/oauth/session server) must
not leave a permanent store row on persist, and the read must never consult the store for it. Its
minted client is overlaid in memory for the session only."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_persist_dcr_client_registration,
_reuse_persisted_dcr_client_if_available,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
temp = MCPServer(
server_id="temp-uuid",
name="temp",
server_name="temp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
client_id=None,
authorization_url="https://p.example/authorize",
token_url="https://p.example/token",
registration_url="https://p.example/register",
)
upsert = AsyncMock()
store_read = AsyncMock(return_value=None)
with (
patch.object(global_mcp_server_manager, "is_config_declared_server", return_value=False),
patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=AsyncMock(return_value=None)),
patch("litellm.proxy._experimental.mcp_server.db.get_mcp_server", new=AsyncMock(return_value=None)),
patch("litellm.proxy._experimental.mcp_server.db.upsert_mcp_server_oauth_client_credentials", new=upsert),
patch(
"litellm.proxy._experimental.mcp_server.db.get_mcp_server_oauth_client_credentials",
new=store_read,
),
patch.object(global_mcp_server_manager, "update_server", new=AsyncMock()),
):
result = await _persist_dcr_client_registration(
temp, {"client_id": "temp-client", "client_secret": "s"}, "https://x/callback"
)
reused = await _reuse_persisted_dcr_client_if_available(
MCPServer(
server_id="temp-uuid",
name="temp",
server_name="temp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
client_id=None,
),
current_redirect_uri="https://x/callback",
)
assert result == "persisted"
assert temp.client_id == "temp-client"
upsert.assert_not_called()
store_read.assert_not_called()
assert reused is False
@pytest.mark.asyncio
async def test_hydrate_does_not_overwrite_explicit_config_client_id():
"""An explicit client_id set in config.yaml wins: hydration must not overwrite it with a stale
persisted store client, and must not even read the store when config already supplied a client."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
hydrate_config_server_dcr_client,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(
server_id="config_static",
name="config_static",
server_name="config_static",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
client_id="explicit-from-config",
)
store_read = AsyncMock(
return_value={"client_id": "stale-store-client", "client_secret": "x", "redirect_uris": []}
)
with (
patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
patch(
"litellm.proxy._experimental.mcp_server.db.get_mcp_server_oauth_client_credentials",
new=store_read,
),
):
applied = await hydrate_config_server_dcr_client(server)
assert applied is False
assert server.client_id == "explicit-from-config"
store_read.assert_not_called()
@pytest.mark.asyncio
async def test_reuse_does_not_inherit_store_client_when_a_row_exists():
"""Security: a server that HAS a LiteLLM_MCPServerTable row reads its DCR client only from that
row, never from the server-scoped store. server_id is caller-settable on create, so a submitted
server whose id collides with a config-declared server must not be able to load that config
server's client from the store and send it to its own token endpoint. A row that exists but has
no client_id yields no reusable client and must not fall back to the store."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_reuse_persisted_dcr_client_if_available,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
submitted = MCPServer(
server_id="collides_with_config",
name="submitted",
server_name="submitted",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
client_id=None,
)
row_without_client = MagicMock()
row_without_client.credentials = None
row_without_client.server_id = "collides_with_config"
store_lookup = AsyncMock(
return_value={"client_id": "config-secret-client", "client_secret": "leak", "redirect_uris": []}
)
with (
patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
patch(
"litellm.proxy._experimental.mcp_server.db.get_mcp_server",
new=AsyncMock(return_value=row_without_client),
),
patch(
"litellm.proxy._experimental.mcp_server.db.get_mcp_server_oauth_client_credentials",
new=store_lookup,
),
):
result = await _reuse_persisted_dcr_client_if_available(submitted, current_redirect_uri="https://x/callback")
assert result is False
assert submitted.client_id is None
store_lookup.assert_not_called()
@pytest.mark.asyncio
async def test_load_servers_from_config_hydrates_dcr_clients():
"""load_servers_from_config must invoke DCR-client hydration so config servers pick up their
persisted client on startup; deleting the call site leaves a restarted server with no client_id
and forces re-authentication on every token expiry."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
hydrate_spy = AsyncMock()
with patch.object(global_mcp_server_manager, "_hydrate_config_servers_dcr_clients", new=hydrate_spy):
await global_mcp_server_manager.load_servers_from_config({})
hydrate_spy.assert_awaited_once()
@pytest.mark.asyncio
async def test_reload_servers_from_database_hydrates_dcr_clients():
"""load_servers_from_config runs before the DB connects at startup, so its hydration no-ops;
reload_servers_from_database runs after the DB connects and must hydrate config servers' persisted
DCR clients too, or a fresh pod has no client_id for a config server and forces re-authentication
on the first token refresh."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
prisma = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
hydrate_spy = AsyncMock()
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=prisma,
),
patch.object(global_mcp_server_manager, "_hydrate_config_servers_dcr_clients", new=hydrate_spy),
):
await global_mcp_server_manager.reload_servers_from_database()
hydrate_spy.assert_awaited_once()

View file

@ -989,6 +989,7 @@ class TestRotateCredentials:
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[server])
mock_prisma.db.litellm_mcpservertable.update = AsyncMock()
mock_prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(return_value=[])
with (
patch(
@ -1036,6 +1037,7 @@ class TestRotateCredentials:
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[server])
mock_prisma.db.litellm_mcpservertable.update = AsyncMock()
mock_prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(return_value=[])
with (
patch(