Merge pull request #33768 from BerriAI/litellm_mcp_dcr_config_client_persist

fix(mcp): persist config.yaml DCR clients in a server-scoped store so refresh survives token expiry
This commit is contained in:
tin-berri 2026-07-18 11:04:41 -07:00 • committed by GitHub
commit 703327a544
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
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

@ -1100,6 +1100,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],
@ -1340,8 +1348,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.
@ -4935,6 +4967,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(