fix(mcp): keep worker MCP configurations consistent via a catalog revision (#42568)

* fix(mcp): refresh server catalog for each gateway operation

* fix(mcp): reject configuration changes during scoped dispatch

* fix(mcp): refresh shared catalog state for each operation

* fix(mcp): refresh catalog before native alias routing

* docs(mcp): clarify native route resolution order

* test(mcp): align streaming fixture and generated API documentation

* fix(mcp): preserve discovery published during catalog refresh

* fix(mcp): coalesce queued catalog reads without a stale window

* fix(mcp): reconcile concurrent route changes when publishing catalog

* test(mcp): provide catalog scope in post-call hook fixtures

* fix(mcp): retain valid live routes and handlers during refresh

* fix(mcp): keep refreshed OpenAPI operation membership authoritative

* test(mcp): preserve logging fixtures after catalog integration

* fix(mcp): isolate transport test admission state and exhaust OAuth outcomes

* refactor(mcp): narrow catalog consistency change to ticket scope

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* feat(mcp): gate catalog refresh on a database revision marker

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): refresh waiters that observed a newer catalog revision

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(mcp): awaitable catalog revision doubles in prisma mocks

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(mcp): satisfy catalog refresh lint and type gates

* test(mcp): provide catalog scope in toolset fixtures

* test(mcp): exercise temporary OAuth through catalog operations

* fix(mcp): share active catalog snapshots across discovery tasks

* fix(mcp): retain discovered routes across cached catalog operations

* fix(mcp): preserve local tool ownership across discovery

* refactor(mcp): satisfy tightened immutability lint ceiling

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): authorize local handlers by registered server ownership

* fix(mcp): check registered handler freshness and isolate fixtures

* fix(mcp): read catalog revisions and snapshots from the writer

* fix(mcp): resolve access groups from the operation catalog

* fix(mcp): preserve empty access group restrictions

* fix(mcp): retain rediscovered routes across concurrent updates

* fix(mcp): enforce registered ownership for local tool dispatch

* fix(mcp): preserve issuer discovery across catalog snapshots

---------

Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
joshua-berri 2026-10-06 18:21:29 -07:00 • committed by GitHub
parent 48c88bae04
commit 950da7c2ec
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
44 changed files with 3529 additions and 793 deletions

View file

@ -0,0 +1,13 @@
CREATE OR REPLACE FUNCTION litellm_bump_mcp_catalog_revision() RETURNS TRIGGER AS $$
BEGIN
INSERT INTO "LiteLLM_Config" ("param_name", "reload_revision")
VALUES ('mcp_catalog', 1)
ON CONFLICT ("param_name") DO UPDATE SET "reload_revision" = "LiteLLM_Config"."reload_revision" + 1;
RETURN NULL;
END;
$$ LANGUAGE plpgsql;
DROP TRIGGER IF EXISTS "litellm_mcp_catalog_revision" ON "LiteLLM_MCPServerTable";
CREATE TRIGGER "litellm_mcp_catalog_revision"
AFTER INSERT OR UPDATE OR DELETE ON "LiteLLM_MCPServerTable"
FOR EACH STATEMENT EXECUTE PROCEDURE litellm_bump_mcp_catalog_revision();

View file

@ -1,5 +1,6 @@
import re
from collections.abc import Mapping, Sequence
from collections.abc import Set as AbstractSet
from dataclasses import dataclass
from datetime import datetime, timezone
from types import MappingProxyType
@ -15,6 +16,7 @@ import litellm
from litellm._internal_context import with_service_target
from litellm._logging import verbose_logger
from litellm.constants import MCP_ALL_TOOLS_WILDCARD
from litellm.proxy._experimental.mcp_server.catalog import global_manager
from litellm.proxy._experimental.mcp_server.oauth_utils import (
get_passthrough_resource_metadata_url,
get_passthrough_www_authenticate,
@ -448,188 +450,191 @@ class MCPRequestHandler:
Raises:
HTTPException: If headers are invalid or missing required headers
"""
headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope)
async with global_manager().catalog.operation():
headers: Final = MCPRequestHandler._safe_get_headers_from_scope(scope)
# Check if there is an explicit LiteLLM API key (primary header)
has_explicit_litellm_key: Final = headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY) is not None
# Check if there is an explicit LiteLLM API key (primary header)
has_explicit_litellm_key: Final = (
headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY) is not None
)
litellm_api_key: Final = MCPRequestHandler.get_litellm_api_key_from_headers(headers) or ""
litellm_api_key: Final = MCPRequestHandler.get_litellm_api_key_from_headers(headers) or ""
# Get the old mcp_auth_header for backward compatibility
mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers)
# Get the old mcp_auth_header for backward compatibility
mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(headers)
# Get the new server-specific auth headers
mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers)
# Get the new server-specific auth headers
mcp_server_auth_headers = MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers)
# Get the oauth2 headers
oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers)
# Get the oauth2 headers
oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers)
# Parse MCP servers from header
mcp_servers_header: Final = headers.get(MCPRequestHandler.LITELLM_MCP_SERVERS_HEADER_NAME)
verbose_logger.debug("Raw MCP servers header: %s", mcp_servers_header)
mcp_servers = None
if mcp_servers_header is not None:
try:
mcp_servers = [s.strip() for s in mcp_servers_header.split(",") if s.strip()]
verbose_logger.debug("Parsed MCP servers: %s", mcp_servers)
except Exception as e:
verbose_logger.debug("Error parsing mcp_servers header: %s", e)
mcp_servers = None
if mcp_servers_header == "" or (mcp_servers is not None and len(mcp_servers) == 0):
mcp_servers = []
# Create a proper Request object with mock body method to avoid ASGI receive channel issues
request: Final = Request(scope=scope)
# Parse MCP servers from header
mcp_servers_header: Final = headers.get(MCPRequestHandler.LITELLM_MCP_SERVERS_HEADER_NAME)
verbose_logger.debug("Raw MCP servers header: %s", mcp_servers_header)
mcp_servers = None
if mcp_servers_header is not None:
try:
mcp_servers = [s.strip() for s in mcp_servers_header.split(",") if s.strip()]
verbose_logger.debug("Parsed MCP servers: %s", mcp_servers)
except Exception as e:
verbose_logger.debug("Error parsing mcp_servers header: %s", e)
mcp_servers = None
if mcp_servers_header == "" or (mcp_servers is not None and len(mcp_servers) == 0):
mcp_servers = []
# Create a proper Request object with mock body method to avoid ASGI receive channel issues
request: Final = Request(scope=scope)
async def mock_body():
return b"{}"
async def mock_body():
return b"{}"
request.body = mock_body
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415
get_request_route,
)
request.body = mock_body
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415
get_request_route,
)
request_route: Final = get_request_route(request)
# Only OAuth metadata routes registered under /.well-known/ are public.
if request_route.startswith("/.well-known/"):
validated_user_api_key_auth = UserAPIKeyAuth()
elif (
has_explicit_litellm_key
and oauth2_headers
and is_bridge_envelope_shaped(oauth2_headers["Authorization"])
and (
dual_bridge_target := MCPRequestHandler._single_dcr_bridge_delegate_target(
request_route: Final = get_request_route(request)
# Only OAuth metadata routes registered under /.well-known/ are public.
if request_route.startswith("/.well-known/"):
validated_user_api_key_auth = UserAPIKeyAuth()
elif (
has_explicit_litellm_key
and oauth2_headers
and is_bridge_envelope_shaped(oauth2_headers["Authorization"])
and (
dual_bridge_target := MCPRequestHandler._single_dcr_bridge_delegate_target(
path=request_route,
mcp_servers=mcp_servers,
client_ip=IPAddressUtils.get_mcp_client_ip(request),
)
)
is not None
):
(
validated_user_api_key_auth,
mcp_server_auth_headers,
) = await MCPRequestHandler._admit_dcr_bridge_dual_credential(
server=dual_bridge_target.server,
requested_name=dual_bridge_target.requested_name,
authorization_value=oauth2_headers["Authorization"],
litellm_api_key=litellm_api_key,
mcp_server_auth_headers=mcp_server_auth_headers,
request=request,
route=request_route,
)
elif has_explicit_litellm_key:
# An explicit x-litellm-api-key is always a LiteLLM credential, even
# for a delegated server, so validate it: identity / spend / rate
# limits resolve and any stored upstream token can be forwarded.
validated_user_api_key_auth = await user_api_key_auth(
api_key=f"Bearer {_get_bearer_token_or_received_api_key(litellm_api_key)}",
request=request,
)
elif MCPRequestHandler._target_servers_are_true_passthrough(
path=request_route,
mcp_servers=mcp_servers,
client_ip=IPAddressUtils.get_mcp_client_ip(request),
) or (
MCPRequestHandler._single_dcr_bridge_delegate_target(
path=request_route,
mcp_servers=mcp_servers,
client_ip=IPAddressUtils.get_mcp_client_ip(request),
)
)
is not None
):
(
validated_user_api_key_auth,
mcp_server_auth_headers,
) = await MCPRequestHandler._admit_dcr_bridge_dual_credential(
server=dual_bridge_target.server,
requested_name=dual_bridge_target.requested_name,
authorization_value=oauth2_headers["Authorization"],
litellm_api_key=litellm_api_key,
mcp_server_auth_headers=mcp_server_auth_headers,
request=request,
route=request_route,
)
elif has_explicit_litellm_key:
# An explicit x-litellm-api-key is always a LiteLLM credential, even
# for a delegated server, so validate it: identity / spend / rate
# limits resolve and any stored upstream token can be forwarded.
validated_user_api_key_auth = await user_api_key_auth(
api_key=f"Bearer {_get_bearer_token_or_received_api_key(litellm_api_key)}",
request=request,
)
elif MCPRequestHandler._target_servers_are_true_passthrough(
path=request_route,
mcp_servers=mcp_servers,
client_ip=IPAddressUtils.get_mcp_client_ip(request),
) or (
MCPRequestHandler._single_dcr_bridge_delegate_target(
path=request_route,
mcp_servers=mcp_servers,
client_ip=IPAddressUtils.get_mcp_client_ip(request),
)
is not None
and not oauth2_headers
and not mcp_server_auth_headers
and not mcp_auth_header
):
validated_user_api_key_auth = UserAPIKeyAuth()
elif (
bridge_delegate_target := MCPRequestHandler._single_dcr_bridge_delegate_target(
path=request_route,
mcp_servers=mcp_servers,
client_ip=IPAddressUtils.get_mcp_client_ip(request),
)
) is not None and oauth2_headers:
(
validated_user_api_key_auth,
mcp_server_auth_headers,
) = await MCPRequestHandler._admit_dcr_bridge_authorization(
server=bridge_delegate_target.server,
requested_name=bridge_delegate_target.requested_name,
authorization_value=oauth2_headers["Authorization"],
litellm_api_key=litellm_api_key,
mcp_server_auth_headers=mcp_server_auth_headers,
request=request,
route=request_route,
)
elif oauth2_headers and is_session_bearer_shaped(oauth2_headers["Authorization"]):
# A gateway DCR session bearer at any MCP scope: open the identity-only session
# token and admit under the live litellm user; downstream grant resolution
# intersects the admitted subject's servers with any path or header target, so a
# per-server scope narrows and never broadens. One that does not open fails
# closed with the scope's invalid_token challenge; a non-session bearer falls
# through to the oauth2 arm.
validated_user_api_key_auth = await MCPRequestHandler._admit_gateway_session(
authorization_value=oauth2_headers["Authorization"],
request=request,
route=request_route,
mcp_servers=mcp_servers,
)
elif oauth2_headers:
# Authorization on a non-delegated server: the bearer must be a real
# LiteLLM credential, so a failed validation is a genuine 401/403 and
# propagates unless a fallback in _admission_failure_fallback applies.
try:
validated_user_api_key_auth = await user_api_key_auth(api_key=litellm_api_key, request=request)
except (HTTPException, ProxyException) as e:
validated_user_api_key_auth = _admission_failure_fallback(
request=request,
request_route=request_route,
is not None
and not oauth2_headers
and not mcp_server_auth_headers
and not mcp_auth_header
):
validated_user_api_key_auth = UserAPIKeyAuth()
elif (
bridge_delegate_target := MCPRequestHandler._single_dcr_bridge_delegate_target(
path=request_route,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
exc=e,
bearer_presented=True,
client_ip=IPAddressUtils.get_mcp_client_ip(request),
)
else:
try:
validated_user_api_key_auth = await user_api_key_auth(api_key=litellm_api_key, request=request)
except (HTTPException, ProxyException) as exc:
validated_user_api_key_auth = _admission_failure_fallback(
) is not None and oauth2_headers:
(
validated_user_api_key_auth,
mcp_server_auth_headers,
) = await MCPRequestHandler._admit_dcr_bridge_authorization(
server=bridge_delegate_target.server,
requested_name=bridge_delegate_target.requested_name,
authorization_value=oauth2_headers["Authorization"],
litellm_api_key=litellm_api_key,
mcp_server_auth_headers=mcp_server_auth_headers,
request=request,
request_route=request_route,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
exc=exc,
bearer_presented=False,
route=request_route,
)
elif oauth2_headers and is_session_bearer_shaped(oauth2_headers["Authorization"]):
# A gateway DCR session bearer at any MCP scope: open the identity-only session
# token and admit under the live litellm user; downstream grant resolution
# intersects the admitted subject's servers with any path or header target, so a
# per-server scope narrows and never broadens. One that does not open fails
# closed with the scope's invalid_token challenge; a non-session bearer falls
# through to the oauth2 arm.
validated_user_api_key_auth = await MCPRequestHandler._admit_gateway_session(
authorization_value=oauth2_headers["Authorization"],
request=request,
route=request_route,
mcp_servers=mcp_servers,
)
elif oauth2_headers:
# Authorization on a non-delegated server: the bearer must be a real
# LiteLLM credential, so a failed validation is a genuine 401/403 and
# propagates unless a fallback in _admission_failure_fallback applies.
try:
validated_user_api_key_auth = await user_api_key_auth(api_key=litellm_api_key, request=request)
except (HTTPException, ProxyException) as e:
validated_user_api_key_auth = _admission_failure_fallback(
request=request,
request_route=request_route,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
exc=e,
bearer_presented=True,
)
else:
try:
validated_user_api_key_auth = await user_api_key_auth(api_key=litellm_api_key, request=request)
except (HTTPException, ProxyException) as exc:
validated_user_api_key_auth = _admission_failure_fallback(
request=request,
request_route=request_route,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
exc=exc,
bearer_presented=False,
)
# Leak-defense (single chokepoint): a gateway admission credential (session bearer or bridge
# envelope) is NEVER a valid upstream token. Scrub it from EVERY egress context so no
# client-forwarded, OBO, or passthrough path can send it upstream for replay. Anchored to the
# credential SHAPE, so a legitimate upstream/passthrough token is forwarded unchanged.
raw_headers = dict(headers)
(
oauth2_headers,
raw_headers,
mcp_auth_header,
mcp_server_auth_headers,
) = MCPRequestHandler._scrub_gateway_admission_credentials(
admitted=_is_mcp_admitted_user_subject(validated_user_api_key_auth),
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
)
# Leak-defense (single chokepoint): a gateway admission credential (session bearer or bridge
# envelope) is NEVER a valid upstream token. Scrub it from EVERY egress context so no
# client-forwarded, OBO, or passthrough path can send it upstream for replay. Anchored to the
# credential SHAPE, so a legitimate upstream/passthrough token is forwarded unchanged.
raw_headers = dict(headers)
(
oauth2_headers,
raw_headers,
mcp_auth_header,
mcp_server_auth_headers,
) = MCPRequestHandler._scrub_gateway_admission_credentials(
admitted=_is_mcp_admitted_user_subject(validated_user_api_key_auth),
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
)
return (
validated_user_api_key_auth,
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
oauth2_headers,
raw_headers,
)
return (
validated_user_api_key_auth,
mcp_auth_header,
mcp_servers,
mcp_server_auth_headers,
oauth2_headers,
raw_headers,
)
@staticmethod
def _is_gateway_admission_credential(value: str | None) -> bool:
@ -1613,9 +1618,13 @@ class MCPRequestHandler:
# Get allowed servers from key and team
allowed_mcp_servers_for_key = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth)
# The key explicitly opted out of every MCP server. This overrides
# team inheritance and additive grants (mirrors no-default-models).
if SpecialMCPServerNames.no_mcp_servers.value in allowed_mcp_servers_for_key:
# Only an explicit opt-out overrides additive grants. An empty group
# scope uses the same marker to restrict inheritance below.
if SpecialMCPServerNames.no_mcp_servers.value in allowed_mcp_servers_for_key and (
user_api_key_auth is None
or (permission := await MCPRequestHandler._key_object_permission_hydrated(user_api_key_auth)) is None
or SpecialMCPServerNames.no_mcp_servers.value in (permission.mcp_servers or [])
):
return MCPServerAccess(server_ids=(), scope="scoped")
allowed_mcp_servers_for_team = await MCPRequestHandler._get_allowed_mcp_servers_for_team(user_api_key_auth)
@ -1742,7 +1751,7 @@ class MCPRequestHandler:
declares_key_mcp_scope: Final = getattr(key_object_permission, "mcp_servers", None) is not None
return MCPServerAccess(
server_ids=tuple(set(allowed_mcp_servers)),
server_ids=tuple(set(allowed_mcp_servers) - {SpecialMCPServerNames.no_mcp_servers.value}),
scope=(
"scoped"
if has_lower_level_mcp_restrictions or org_restricts or declares_key_mcp_scope
@ -2543,6 +2552,20 @@ class MCPRequestHandler:
verbose_logger.warning("Failed to get key access group MCP server grants: %s", e)
return []
@staticmethod
def _server_grants_with_group_scope(
servers: AbstractSet[str],
permission: LiteLLM_ObjectPermissionTable,
) -> AbstractSet[str]:
"""Keep a configured empty group scope distinct from an unrestricted level.
Combine same-level grants first; the existing zero-server marker prevents
inheritance and ceiling substitution without discarding independent grants.
"""
if servers or not permission.mcp_access_groups:
return servers
return {SpecialMCPServerNames.no_mcp_servers.value}
@staticmethod
async def _get_allowed_mcp_servers_for_key(
user_api_key_auth: UserAPIKeyAuth | None = None,
@ -2625,7 +2648,7 @@ class MCPRequestHandler:
# Combine all lists
all_servers: Final = direct_mcp_servers + access_group_servers + tool_perm_servers + toolset_servers
return list(set(all_servers))
return list(MCPRequestHandler._server_grants_with_group_scope(set(all_servers), key_object_permission))
except Exception as e:
verbose_logger.warning("Failed to get allowed MCP servers for key: %s", e)
return []
@ -2698,7 +2721,7 @@ class MCPRequestHandler:
team_access_group_servers: list[str],
*,
requires_fresh_policy: bool = False,
) -> set[str]:
) -> AbstractSet[str]:
"""The raw MCP-server set a team grants (before any org ceiling): its object_permission (direct
``mcp_servers``, the ``all_proxy_servers`` sentinel → the full registry, legacy access groups,
tool-perm-referenced servers, toolset-referenced servers) unioned with its unified
@ -2719,13 +2742,14 @@ class MCPRequestHandler:
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
object_permissions, requires_fresh_policy=requires_fresh_policy
)
return (
servers: Final = (
set(global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or []))
| set(legacy_access_group_servers)
| set(global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys())
| toolset_grants.keys()
| set(team_access_group_servers)
)
return MCPRequestHandler._server_grants_with_group_scope(servers, object_permissions)
@staticmethod
async def _allowed_mcp_servers_for_single_team(
@ -2935,7 +2959,7 @@ class MCPRequestHandler:
all_servers: Final = tuple(
{*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *toolset_grants}
)
return list(set(all_servers))
return list(MCPRequestHandler._server_grants_with_group_scope(set(all_servers), object_permissions))
except Exception as e:
# None = ceiling UNRESOLVED, distinct from [] = org places no restriction. Collapsing them
# let a DB fault silently drop a ceiling; the caller picks fail-open/closed from this signal.
@ -3033,7 +3057,7 @@ class MCPRequestHandler:
# Combine all lists
all_servers: Final = direct_mcp_servers + access_group_servers + tool_perm_servers
return list(set(all_servers))
return list(MCPRequestHandler._server_grants_with_group_scope(set(all_servers), object_permission))
except Exception as e:
verbose_logger.warning("Failed to get allowed MCP servers for end_user: %s", e)
return []
@ -3168,7 +3192,12 @@ class MCPRequestHandler:
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
object_permissions, requires_fresh_policy=fresh
)
return tuple({*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *toolset_grants})
return tuple(
MCPRequestHandler._server_grants_with_group_scope(
{*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *toolset_grants},
object_permissions,
)
)
except Exception as e: # noqa: BLE001 # any resolution fault is an unresolved ceiling, never "no ceiling"
verbose_logger.warning("Failed to get allowed MCP servers for user: %s", e)
return None
@ -3515,7 +3544,12 @@ class MCPRequestHandler:
obj_perm, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
)
inline_tools: Final = global_mcp_server_manager.expand_tool_permissions(obj_perm.mcp_tool_permissions)
return list({*expanded_direct_servers, *access_group_servers, *toolset_grants, *inline_tools})
return list(
MCPRequestHandler._server_grants_with_group_scope(
{*expanded_direct_servers, *access_group_servers, *toolset_grants, *inline_tools},
obj_perm,
)
)
except Exception as e:
if managed_agent_policy(user_api_key_auth) is not None or isinstance(e, UnloadableEntitlementError):
raise
@ -3638,7 +3672,7 @@ class MCPRequestHandler:
requires_fresh_policy: bool = False,
) -> list[str]:
"""
Resolve MCP access groups to server IDs by querying BOTH the MCP server table (DB) AND config-loaded servers.
Resolve MCP access groups against the operation snapshot, or database and config outside an operation.
``requires_fresh_policy`` reads the writer and propagates a read fault instead of resolving to no servers.
"""
from litellm.proxy.proxy_server import prisma_client
@ -3649,6 +3683,10 @@ class MCPRequestHandler:
global_mcp_server_manager,
)
snapshot: Final = global_mcp_server_manager.catalog.current()
if snapshot is not None:
return list(MCPRequestHandler._get_config_server_ids_for_access_groups(snapshot.servers, access_groups))
# Use the new helper for config-loaded servers
server_ids: Final = MCPRequestHandler._get_config_server_ids_for_access_groups(
global_mcp_server_manager.config_mcp_servers, access_groups

View file

@ -25,6 +25,7 @@ from fastapi import APIRouter, Depends, Form, HTTPException, Request
from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse
from litellm._logging import verbose_proxy_logger
from litellm.proxy._experimental.mcp_server.catalog import public_catalog_operation
from litellm.proxy._experimental.mcp_server.db import store_user_credential
from litellm.proxy._experimental.mcp_server.oauth_utils import (
BYOK_RESOURCE_METADATA_PATH,
@ -738,6 +739,7 @@ async def byok_protected_resource_metadata(request: Request) -> JSONResponse:
@router.get("/v1/mcp/oauth/authorize", include_in_schema=False)
@public_catalog_operation
async def byok_authorize_get(
request: Request,
client_id: str | None = None,
@ -778,7 +780,7 @@ async def byok_authorize_get(
global_mcp_server_manager,
)
registry: Final = global_mcp_server_manager.get_registry()
registry: Final = await global_mcp_server_manager.catalog.list()
if server_id in registry:
srv: Final = registry[server_id]
server_name = srv.server_name or srv.name

View file

@ -0,0 +1,578 @@
"""Authoritative MCP catalog snapshots shared by legacy lookup adapters."""
from __future__ import annotations
import asyncio
import hashlib
import json
from collections import UserDict
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, MutableMapping, Sequence
from contextlib import asynccontextmanager
from contextvars import ContextVar
from dataclasses import dataclass, replace
from functools import wraps
from itertools import chain
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, ParamSpec, TypeVar
from litellm._logging import verbose_logger
if TYPE_CHECKING:
from pydantic import BaseModel
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.types.mcp_server.mcp_server_manager import MCPServer
from litellm.types.mcp_server.tool_registry import MCPTool
_P = ParamSpec("_P")
_R = TypeVar("_R")
class _OperationRoutes(UserDict[str, str]):
def __init__(self, initial: Mapping[str, str]) -> None:
super().__init__()
self.data = dict(initial)
self.written_names: set[str] = set() # mutable-ok: record route writes without copying the journal per tool
def __setitem__(self, name: str, owner: str) -> None:
self.data[name] = owner
self.written_names.add(name)
@dataclass(frozen=True, slots=True)
class CatalogSnapshot:
servers: Mapping[str, MCPServer]
identity: str
tools: Mapping[str, MCPTool]
routing: MutableMapping[str, str]
def _configuration_identity(server: MCPServer) -> str:
return json.dumps(
server.model_dump(
mode="json",
exclude=frozenset(
(
"short_prefix",
"scopes",
"authorization_url",
"token_url",
"registration_url",
"authorization_response_iss_parameter_supported",
)
)
| (frozenset() if server.issuer_is_anchored else frozenset(("issuer",))),
),
sort_keys=True,
)
def _check_oauth_revision(selected: MCPServer, candidate: MCPServer | None) -> None:
from fastapi import HTTPException
if candidate is None or _configuration_identity(selected) != _configuration_identity(candidate):
raise HTTPException(status_code=503, detail="OAuth metadata discovery changed repeatedly; retry shortly")
def _snapshot(manager: MCPServerManager, database_identity: str) -> CatalogSnapshot:
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
servers: Final = manager.config_mcp_servers | manager.registry
detached: Final = MappingProxyType({key: value.model_copy(deep=True) for key, value in servers.items()})
serialized: Final = json.dumps(
(
database_identity,
tuple(sorted(manager.registry)),
tuple((key, _configuration_identity(value)) for key, value in sorted(manager.config_mcp_servers.items())),
),
sort_keys=True,
)
return CatalogSnapshot(
detached,
hashlib.sha256(serialized.encode()).hexdigest(),
MappingProxyType(dict(global_mcp_tool_registry.published_tools)),
dict(manager.published_tool_routes),
)
class TargetCatalog:
def __init__(self, manager: MCPServerManager) -> None:
self.manager = manager
self._refresh_lock = asyncio.Lock()
self._database_identity = ""
self._arrival_ticket = 0
self._completed_ticket = 0
self._shared_snapshot: CatalogSnapshot | None = None
self._applied_revision: int | None = None
self._warned_shadowed_config_server_ids: frozenset[str] = frozenset()
self._warned_capturing_config_server_ids: frozenset[str] = frozenset()
self._operation: ContextVar[tuple[CatalogSnapshot, asyncio.Event] | None] = ContextVar(
"mcp_catalog_snapshot", default=None
)
self._staged_routing: ContextVar[tuple[dict[str, str], asyncio.Event] | None] = ContextVar(
"mcp_catalog_routing", default=None
)
def current(self) -> CatalogSnapshot | None:
scoped: Final = self._operation.get()
return scoped[0] if scoped is not None and not scoped[1].is_set() else None
def routing(self) -> MutableMapping[str, str]:
staged: Final = self._staged_routing.get()
if staged is not None and not staged[1].is_set():
return staged[0]
snapshot: Final = self.current()
return snapshot.routing if snapshot is not None else self.manager.published_tool_routes
def registry(self) -> Mapping[str, MCPServer]:
snapshot: Final = self.current()
return snapshot.servers if snapshot is not None else self.manager.config_mcp_servers | self.manager.registry
async def _fresh_snapshot(self) -> CatalogSnapshot:
from fastapi import HTTPException
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return _snapshot(self.manager, self._database_identity)
from litellm.proxy.proxy_server import should_load_db_object
if not should_load_db_object("mcp"):
return _snapshot(self.manager, self._database_identity)
from litellm.proxy._experimental.mcp_server.db import get_mcp_catalog_revision
try:
revision: Final = await get_mcp_catalog_revision(prisma_client)
if revision is not None and revision == self._applied_revision and self._shared_snapshot is not None:
return self._shared_snapshot
self._arrival_ticket += 1
arrival: Final = self._arrival_ticket
async with self._refresh_lock:
if arrival > self._completed_ticket or revision != self._applied_revision:
await self._publish_refresh(revision, reuse_unchanged=True)
if self._shared_snapshot is None:
raise HTTPException(status_code=503, detail="MCP server configuration could not be refreshed")
return self._shared_snapshot
except Exception as exc:
raise HTTPException(status_code=503, detail="MCP server configuration could not be refreshed") from exc
async def list(self) -> Mapping[str, MCPServer]:
async with self.operation() as snapshot:
return snapshot.servers
def assert_current(self, server: MCPServer) -> None:
from fastapi import HTTPException
snapshot: Final = self.current()
if snapshot is None or server.server_id not in snapshot.servers:
return
expected: Final = snapshot.servers[server.server_id]
registered: Final = self.manager.registry.get(server.server_id) or self.manager.config_mcp_servers.get(
server.server_id
)
if (
registered is None
or registered.updated_at != expected.updated_at
or server.updated_at != expected.updated_at
):
raise HTTPException(status_code=503, detail="MCP server configuration changed; retry the operation")
@asynccontextmanager
async def operation(self) -> AsyncGenerator[CatalogSnapshot]:
current: Final = self.current()
if current is not None:
yield current
return
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
shared: Final = await self._fresh_snapshot()
initial_routing: Final = self._unchanged_routing(shared.servers, self.manager.published_tool_routes)
routing: Final = _OperationRoutes(initial_routing)
snapshot: Final = replace(
shared,
servers=MappingProxyType({key: value.model_copy(deep=True) for key, value in shared.servers.items()}),
routing=routing,
)
closed: Final = asyncio.Event()
token: Final = self._operation.set((snapshot, closed))
try:
with global_mcp_tool_registry.catalog_scope(snapshot.tools):
yield snapshot
finally:
closed.set()
self._operation.reset(token)
self._retain_discovered_routing(snapshot, frozenset(routing.written_names))
def _retain_discovered_routing(self, snapshot: CatalogSnapshot, written_names: frozenset[str]) -> None:
self.manager.published_tool_routes = self.manager.published_tool_routes | self._unchanged_routing(
snapshot.servers,
{name: owner for name, owner in snapshot.routing.items() if name in written_names},
)
def _unchanged_routing(
self, servers: Mapping[str, MCPServer], routing: Mapping[str, str]
) -> MappingProxyType[str, str]:
from litellm.proxy._experimental.mcp_server.utils import normalize_server_name
current: Final = self.manager.config_mcp_servers | self.manager.registry
unchanged: Final = (
server
for key, server in servers.items()
if (candidate := current.get(key)) is not None
and _configuration_identity(candidate) == _configuration_identity(server)
)
unchanged_owners: Final = frozenset(chain.from_iterable(map(self.manager.owned_mapping_values, unchanged)))
return MappingProxyType(
{name: owner for name, owner in routing.items() if normalize_server_name(owner) in unchanged_owners}
)
async def resolve(self, identifier: str, client_ip: str | None = None) -> MCPServer | None:
async with self.operation():
return self.manager.get_mcp_server_by_name(
identifier, client_ip=client_ip
) or self.manager.get_mcp_server_by_id(identifier, client_ip=client_ip)
async def resolve_oauth_metadata(
self,
server: MCPServer,
resolve: Callable[[MCPServer], Awaitable[MCPServer]],
) -> MCPServer:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import oauth_endpoints_unresolved
snapshot: Final = self.current()
selected: Final = snapshot.servers.get(server.server_id) if snapshot is not None else None
if selected is None:
return await resolve(server)
if not oauth_endpoints_unresolved(selected):
return selected
registered: Final = self.manager.registry.get(server.server_id) or self.manager.config_mcp_servers.get(
server.server_id
)
self.assert_current(server)
_check_oauth_revision(selected, registered)
resolved: Final = await resolve(selected)
_check_oauth_revision(selected, resolved)
return resolved
async def reload(self) -> None:
from litellm.proxy._experimental.mcp_server.db import get_mcp_catalog_revision
from litellm.proxy.proxy_server import prisma_client
revision: Final = await get_mcp_catalog_revision(prisma_client) if prisma_client is not None else None
async with self._refresh_lock:
await self._publish_refresh(revision)
async def _publish_refresh(self, revision: int | None, *, reuse_unchanged: bool = False) -> None:
covered: Final = self._arrival_ticket
token: Final = self._operation.set(None)
try:
await self._reload_and_publish(reuse_unchanged=reuse_unchanged)
except Exception:
self._shared_snapshot = None
self._completed_ticket = covered
raise
else:
self._shared_snapshot = _snapshot(self.manager, self._database_identity)
self._applied_revision = revision
self._completed_ticket = covered
finally:
self._operation.reset(token)
async def _reload_and_publish(self, *, reuse_unchanged: bool) -> None:
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
from litellm.proxy._experimental.mcp_server.utils import normalize_server_name
previous_config: Final = MappingProxyType(
{key: value.model_copy(deep=True) for key, value in self.manager.config_mcp_servers.items()}
)
live_registry: Final = self.manager.registry
previous_servers: Final = previous_config | MappingProxyType(
{key: value.model_copy(deep=True) for key, value in live_registry.items()}
)
config_identities: Final = MappingProxyType(
{key: _configuration_identity(value) for key, value in previous_config.items()}
)
staged_config: Final = MappingProxyType(
{key: value.model_copy(deep=True) for key, value in previous_config.items()}
)
await self.manager.hydrate_config_servers_dcr_clients(tuple(staged_config.values()))
initial_routing: Final = MappingProxyType(dict(self.manager.published_tool_routes))
staged_routing: Final = dict(initial_routing)
closed: Final = asyncio.Event()
routing_token: Final = self._staged_routing.set((staged_routing, closed))
initial_tools: Final = MappingProxyType(dict(global_mcp_tool_registry.published_tools))
try:
with global_mcp_tool_registry.catalog_scope(initial_tools) as staged_tools:
await self._reload(reuse_unchanged=reuse_unchanged)
refreshed_openapi: Final = (
server
for server in self.manager.registry.values()
if server.spec_path and server is not live_registry.get(server.server_id)
)
refreshed_openapi_owners: Final = frozenset(
chain.from_iterable(map(self.manager.owned_mapping_values, refreshed_openapi))
)
live_routes: Final = self._unchanged_routing(
previous_servers,
MappingProxyType(
{
name: owner
for name, owner in self.manager.published_tool_routes.items()
if normalize_server_name(owner) not in refreshed_openapi_owners
}
),
)
concurrent_routes: Final = self._unchanged_routing(
self.manager.config_mcp_servers | live_registry,
MappingProxyType(
{
name: owner
for name, owner in self.manager.published_tool_routes.items()
if initial_routing.get(name) != owner
and (name not in staged_routing or staged_routing.get(name) == initial_routing.get(name))
}
),
)
removed_routes: Final = initial_routing.keys() - self.manager.published_tool_routes.keys()
retained_staged_routes: Final = self._unchanged_routing(
previous_config | self.manager.registry,
MappingProxyType(
{
name: owner
for name, owner in staged_routing.items()
if name not in removed_routes or owner != initial_routing[name]
}
),
)
concurrent_tools: Final = MappingProxyType(
{
name: tool
for name, tool in global_mcp_tool_registry.published_tools.items()
if (
name in live_routes
or name in concurrent_routes
or (name not in initial_routing and name not in self.manager.published_tool_routes)
)
and initial_tools.get(name) is staged_tools.get(name)
}
)
self.manager.config_mcp_servers = {
key: value.model_copy(
update=staged_config[key].model_dump(
include=frozenset(("client_id", "client_secret", "token_endpoint_auth_method"))
)
)
if config_identities.get(key) == _configuration_identity(value)
else value
for key, value in self.manager.config_mcp_servers.items()
}
global_mcp_tool_registry.tools = (
MappingProxyType(
{
name: tool
for name, tool in staged_tools.items()
if name in global_mcp_tool_registry.published_tools or tool is not initial_tools.get(name)
}
)
| concurrent_tools
)
self.manager.published_tool_routes = live_routes | retained_staged_routes | concurrent_routes
finally:
closed.set()
self._staged_routing.reset(routing_token)
async def _stage_servers(
self, rows: Sequence[BaseModel], *, reuse_unchanged: bool
) -> dict[str, MCPServer]: # mutable-ok: assign_unique_short_prefix requires a dict registry
from litellm.proxy._experimental.mcp_server.db import LiteLLM_MCPServerTable
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
carry_forward_resolved_oauth_endpoints,
oauth_endpoints_unresolved,
warn_on_server_name_fields,
)
previous_registry: Final = self.manager.registry
new_registry: Final[dict[str, MCPServer]] = {}
# Stage one: build every server. Stage two assigns short prefixes
# against the *full* set so dedup is deterministic regardless of
# iteration order.
for row in rows:
try:
server = LiteLLM_MCPServerTable.model_validate(row.model_dump())
existing_server = previous_registry.get(server.server_id)
if (
existing_server is not None
and (reuse_unchanged or not existing_server.spec_path)
and existing_server.updated_at is not None
and server.updated_at is not None
and existing_server.updated_at == server.updated_at
and (
self.manager.oauth_discovery_slot(server.server_id) is not None
or not oauth_endpoints_unresolved(existing_server)
)
):
# Re-use existing server instance to avoid re-running build_mcp_server_from_table()
# which can perform network discovery for OAuth2 servers.
new_registry[server.server_id] = existing_server
continue
warn_on_server_name_fields(
server_id=server.server_id,
alias=getattr(server, "alias", None),
server_name=getattr(server, "server_name", None),
)
self.manager.warn_if_newly_blocked_stdio(server, existing_server)
verbose_logger.debug("Building server from DB: %s (%s)", server.server_id, server.server_name)
# raw_rows come straight from the DB, so their global env var
# values (like credentials) are still encrypted here, unlike the
# already-decrypted records add_server/update_server are handed.
# Decrypt them while building the registry entry.
new_server = await self.manager.build_mcp_server_from_table(
server, env_vars_are_encrypted=True, register_oauth_discovery=False
)
# Carry the cached short_prefix from the previous registry entry
# (if any) so the prefix is stable across reloads.
if existing_server is not None and existing_server.short_prefix:
new_server.short_prefix = existing_server.short_prefix
carry_forward_resolved_oauth_endpoints(new_server=new_server, previous_server=existing_server)
new_registry[server.server_id] = new_server
except Exception as e:
verbose_logger.exception(
"Skipping MCP server %s (%s) during DB reload: %s",
getattr(row, "server_id", None),
getattr(row, "alias", None),
e,
)
return new_registry
async def _reload(self, *, reuse_unchanged: bool) -> None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
config_ids_capturing_db_identifiers,
warn_on_shared_identifier_prefixes,
)
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
get_prisma_client_or_throw,
)
verbose_logger.debug("Loading MCP servers from database into registry...")
prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy")
# Load only "active", legacy "approved", and NULL (no approval workflow) rows.
# Pending/rejected servers are excluded at the DB level so we never load them.
from litellm.proxy._experimental.mcp_server.db import get_runtime_mcp_server_rows
raw_rows: Final[Sequence[BaseModel]] = await get_runtime_mcp_server_rows(prisma_client)
database_identity: Final = hashlib.sha256(
json.dumps(
tuple(sorted(json.dumps(row.model_dump(mode="json"), sort_keys=True, default=str) for row in raw_rows))
).encode()
).hexdigest()
verbose_logger.info("Found %s MCP servers in database", len(raw_rows))
previous_registry: Final = self.manager.registry
new_registry: Final = await self._stage_servers(raw_rows, reuse_unchanged=reuse_unchanged)
# Assign short prefixes against the full candidate set without
# publishing the staged registry to concurrent callers.
registered_registry: Final[dict[str, MCPServer]] = {}
for server_id, new_server in new_registry.items():
try:
if new_server is not previous_registry.get(server_id):
self.manager.assign_unique_short_prefix(new_server, registry=new_registry)
# Register OpenAPI tools *after* the final short prefix is assigned
# so the tools are stored in the global registry under the same
# prefix that lookups will use.
if new_server is not previous_registry.get(server_id):
if previous_server := previous_registry.get(server_id):
self.manager.remove_server_tool_routing(previous_server)
await self.manager.maybe_register_openapi_tools(new_server, initialize_mapping=False)
registered_registry[server_id] = new_server
except Exception as e:
self.manager.remove_server_tool_routing(new_server)
verbose_logger.exception(
"Skipping MCP server %s (%s) during DB reload: %s",
new_server.server_id,
getattr(new_server, "alias", None),
e,
)
dropped_registry_keys: Final = previous_registry.keys() - registered_registry.keys()
for registry_key in dropped_registry_keys:
self.manager.remove_server_tool_routing(previous_registry[registry_key])
self.manager.invalidate_oauth_discovery_state(previous_registry[registry_key].server_id)
for server_id in previous_registry.keys() | registered_registry.keys():
if previous_registry.get(server_id) != registered_registry.get(server_id):
self.manager.invalidate_server_definition_caches(server_id)
self.manager.invalidate_oauth_discovery_state(server_id)
self._database_identity = database_identity
self.manager.registry = registered_registry
if not reuse_unchanged:
self.manager.clear_initialize_instructions()
warn_on_shared_identifier_prefixes(registered_registry.values())
# A discovery task may have published into ``previous_registry`` while
# this replacement was being staged. Reconcile every published entry
# synchronously after the swap so a lost publication cannot also leave
# the replacement unresolved with no retry slot.
registered_servers: Final = tuple(registered_registry.values())
self.manager.reconcile_oauth_discovery_slots_for_servers(registered_servers)
self.manager.prime_oauth_metadata_discovery_for_servers(registered_servers)
verbose_logger.debug("MCP registry refreshed (%s servers in registry)", len(registered_registry))
# get_registry() is ``config_mcp_servers | registry``, so a database row sharing an id with a
# config.yaml server hides that server everywhere. Only reachable once an operator pins
# ``server_id`` in config.yaml; say so rather than letting the server disappear silently.
shadowed_config_server_ids: Final = frozenset(
self.manager.config_mcp_servers.keys() & registered_registry.keys()
)
if shadowed_config_server_ids and shadowed_config_server_ids != self._warned_shadowed_config_server_ids:
verbose_logger.warning(
"config.yaml MCP server_id(s) %s are also database-backed MCP servers. The database "
"entry takes precedence, so the config.yaml server is unreachable. Give the config "
"entry a different server_id.",
", ".join(sorted(shadowed_config_server_ids)),
)
self._warned_shadowed_config_server_ids = shadowed_config_server_ids
# The mirror image of the block above: a config server_id that is a database server's name
# answers that server's grants instead, because ids are matched before names.
capturing_config_server_ids: Final = config_ids_capturing_db_identifiers(
self.manager.config_mcp_servers.keys(), registered_registry.values()
)
if capturing_config_server_ids and capturing_config_server_ids != self._warned_capturing_config_server_ids:
verbose_logger.warning(
"config.yaml MCP server_id(s) %s are the name or alias of a database-backed MCP "
"server. Permission entries naming them resolve to the config.yaml server, not the "
"database one. Give the config entry a different server_id.",
", ".join(sorted(capturing_config_server_ids)),
)
self._warned_capturing_config_server_ids = capturing_config_server_ids
def catalog_operation(
manager: Callable[[], MCPServerManager],
) -> Callable[[Callable[_P, Awaitable[_R]]], Callable[_P, Awaitable[_R]]]:
def decorate(function: Callable[_P, Awaitable[_R]]) -> Callable[_P, Awaitable[_R]]:
@wraps(function)
async def wrapped(*args: _P.args, **kwargs: _P.kwargs) -> _R: # kwargs-ok: preserves ParamSpec
async with manager().catalog.operation():
return await function(*args, **kwargs)
return wrapped
return decorate
def global_manager() -> MCPServerManager:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
return global_mcp_server_manager
def public_catalog_operation(function: Callable[_P, Awaitable[_R]]) -> Callable[_P, Awaitable[_R]]:
return catalog_operation(global_manager)(function)

View file

@ -42,6 +42,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
encrypt_value_helper,
)
from litellm.proxy.utils import PrismaClient
from litellm.repositories.config_repository import ConfigRepository
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import (
@ -723,6 +724,20 @@ async def _db_find_mcp_server_row(
return await _mcp_server_table_actions(prisma_client).find_unique(where={"server_id": server_id})
MCP_CATALOG_REVISION_PARAM_NAME: Final = "mcp_catalog"
async def get_mcp_catalog_revision(prisma_client: PrismaClient) -> int | None:
"""The ``LiteLLM_Config`` revision the MCP server table trigger last published, or None when
no write has ever bumped it (or the trigger is not installed), meaning always reload."""
row: Final = await ConfigRepository(prisma_client, use_writer=True).table.find_unique(
where={"param_name": MCP_CATALOG_REVISION_PARAM_NAME}
)
if row is None:
return None
return int(row.reload_revision or 0)
async def _db_update_mcp_server_row(
prisma_client: PrismaClient,
server_id: str,
@ -853,6 +868,15 @@ async def get_all_mcp_servers(
return list(_readable_mcp_servers(mcp_servers))
async def get_runtime_mcp_server_rows(
prisma_client: PrismaClient,
) -> Sequence["prisma_db_models.LiteLLM_MCPServerTable"]:
where: Final[prisma_db_types.LiteLLM_MCPServerTableWhereInput] = {
"OR": [{"approval_status": None}, {"approval_status": {"in": ["active", "approved"]}}]
}
return await MCPServerRepository(prisma_client, use_writer=True).table.find_many(where=where)
async def get_mcp_server(prisma_client: PrismaClient, server_id: str) -> LiteLLM_MCPServerTable | None:
"""
Returns the matching mcp server from the db iff exists

View file

@ -37,6 +37,7 @@ from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
can_store_oauth_credential,
oauth_authorization_uses_gateway_credential,
)
from litellm.proxy._experimental.mcp_server.catalog import public_catalog_operation
from litellm.proxy._experimental.mcp_server.faults import (
CallerRejected,
CredentialSource,
@ -1894,7 +1895,7 @@ async def resolve_ephemeral_dcr_client(
def _register_flow_needed_endpoint(mcp_server: MCPServer) -> str | None:
"""The register flow's deferred-discovery join gate. A DCR bridge with no admin-configured
client can only register callers through the upstream's registration endpoint
(``_oauth_endpoints_unresolved`` keeps its discovery slot armed for exactly this shape), so
(``oauth_endpoints_unresolved`` keeps its discovery slot armed for exactly this shape), so
the flow must keep joining discovery while registration is still missing instead of silently
degrading to the dummy short-circuit. Every other shape only needs the authorization url."""
if mcp_server.is_dcr_bridge and not mcp_server.client_id and mcp_server.effective_registration_url is None:
@ -2020,6 +2021,7 @@ async def register_client_with_server(
@router.get("/authorize/mcp-session")
@public_catalog_operation
async def authorize_mcp_session(
request: Request,
redirect_uri: str,
@ -2045,6 +2047,7 @@ async def authorize_mcp_session(
@router.get("/{mcp_server_name}/authorize")
@router.get("/authorize")
@public_catalog_operation
async def authorize(
request: Request,
redirect_uri: str,
@ -2083,9 +2086,11 @@ async def authorize(
resource=resource,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
lookup_name: Final[str | None] = mcp_server_name or client_id
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
mcp_server = _resolve_mcp_server_by_name_or_id(lookup_name, client_ip) if lookup_name else None
mcp_server = await global_mcp_server_manager.catalog.resolve(lookup_name, client_ip) if lookup_name else None
if mcp_server is None and mcp_server_name is None:
mcp_server = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if mcp_server is None:
@ -2119,6 +2124,7 @@ async def authorize(
@router.post("/{mcp_server_name}/token")
@router.post("/token")
@public_catalog_operation
async def token_endpoint(
request: Request,
grant_type: str = Form(...),
@ -2212,6 +2218,7 @@ async def _vendor_credential_state(user_id: str, server_id: str) -> VendorCreden
@router.get("/authorize/flow")
@public_catalog_operation
async def authorize_flow(request: Request, flow: str) -> Response:
return await describe_connect_flow(
request=request,
@ -2223,6 +2230,7 @@ async def authorize_flow(request: Request, flow: str) -> Response:
@router.post("/authorize/complete")
@public_catalog_operation
async def authorize_complete(
request: Request,
flow: str = Form(...),
@ -2610,6 +2618,7 @@ def is_network_error(exc: Exception) -> bool:
return isinstance(exc, httpx.TransportError)
@public_catalog_operation
async def _build_oauth_protected_resource_response(
request: Request,
mcp_server_name: str | None,
@ -2873,6 +2882,7 @@ async def oauth_authorization_server_aggregate(request: Request):
# Standard MCP pattern: /.well-known/oauth-protected-resource/mcp/{server_name}
# This is the pattern expected by standard MCP clients (mcp-inspector, VSCode Copilot)
@router.get(f"/.well-known/oauth-protected-resource{well_known_root_suffix()}/mcp/{{mcp_server_name}}")
@public_catalog_operation
async def oauth_protected_resource_mcp_standard(request: Request, mcp_server_name: str):
"""
OAuth protected resource discovery endpoint using standard MCP URL pattern.
@ -2893,6 +2903,7 @@ async def oauth_protected_resource_mcp_standard(request: Request, mcp_server_nam
# LiteLLM legacy pattern: /.well-known/oauth-protected-resource/{server_name}/mcp
# Kept for backward compatibility with existing deployments
@router.get(f"/.well-known/oauth-protected-resource{well_known_root_suffix()}/{{mcp_server_name}}/mcp")
@public_catalog_operation
async def oauth_protected_resource_mcp(request: Request, mcp_server_name: str | None = None):
"""
OAuth protected resource discovery endpoint using LiteLLM legacy URL pattern.
@ -2916,12 +2927,7 @@ def _build_oauth_authorization_server_response(
*,
issuer_path: str | None = None,
) -> dict:
"""Build OAuth authorization server metadata response (gateway-as-AS shape).
Synchronous because the body only does dict construction and synchronous
registry lookups; unlike :func:`_build_oauth_protected_resource_response`
it does not need to await any upstream IO.
"""
"""Build OAuth authorization server metadata response (gateway-as-AS shape)."""
request_base_url: Final = get_request_base_url(request)
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
explicitly_named: Final = mcp_server_name is not None
@ -2969,6 +2975,7 @@ def _build_oauth_authorization_server_response(
# Standard MCP pattern: /.well-known/oauth-authorization-server/mcp/{server_name}
@router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/mcp/{{mcp_server_name}}")
@public_catalog_operation
async def oauth_authorization_server_mcp_standard(request: Request, mcp_server_name: str):
"""
OAuth authorization server discovery endpoint using standard MCP URL pattern.
@ -2986,6 +2993,7 @@ async def oauth_authorization_server_mcp_standard(request: Request, mcp_server_n
# LiteLLM legacy pattern and root endpoint
@router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/{{mcp_server_name}}")
@router.get("/.well-known/oauth-authorization-server")
@public_catalog_operation
async def oauth_authorization_server_mcp(request: Request, mcp_server_name: str | None = None):
"""
OAuth authorization server discovery endpoint.
@ -3059,6 +3067,7 @@ async def jwks_json(request: Request):
# Additional legacy pattern support
@router.get(f"/.well-known/oauth-authorization-server{well_known_root_suffix()}/{{mcp_server_name}}/mcp")
@public_catalog_operation
async def oauth_authorization_server_legacy(request: Request, mcp_server_name: str):
"""
OAuth authorization server discovery for legacy /{server_name}/mcp pattern.
@ -3072,6 +3081,7 @@ async def oauth_authorization_server_legacy(request: Request, mcp_server_name: s
@router.post("/{mcp_server_name}/register")
@router.post("/register")
@public_catalog_operation
async def register_client(request: Request, mcp_server_name: str | None = None):
# Get the correct base URL considering X-Forwarded-* headers
request_base_url: Final = get_request_base_url(request)
@ -3080,46 +3090,44 @@ async def register_client(request: Request, mcp_server_name: str | None = None):
data: Final[dict] = {**request_data}
client_redirect_uris: Final = client_supplied_redirect_uris(data.get("redirect_uris"))
dummy_return: Final = {
"client_id": mcp_server_name or "dummy_client",
"client_secret": "dummy",
"redirect_uris": client_redirect_uris or [f"{request_base_url}/callback"],
}
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
if not mcp_server_name:
# A real DCR request carries redirect_uris (RFC 7591): route it to the aggregate DCR
# endpoint the aggregate authorization-server metadata advertises. A single-server
# deployment registers at /{server}/register instead (its bare-origin discovery
# advertises that), so this does not affect it. A request without redirect_uris is not
# a DCR request, so the legacy single-server-or-dummy fallback is kept for it.
if data.get("redirect_uris"):
return await register_aggregate_client(
request=request, request_body=data, token_exchange_available=token_exchange_available()
)
resolved: Final = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if resolved:
return await register_client_with_server(
request=request,
mcp_server=resolved,
client_name=data.get("client_name", ""),
grant_types=data.get("grant_types", []),
response_types=data.get("response_types", []),
token_endpoint_auth_method=data.get("token_endpoint_auth_method", ""),
fallback_client_id=resolved.server_name or resolved.name,
client_redirect_uris=client_redirect_uris,
)
return dummy_return
if not mcp_server_name and data.get("redirect_uris"):
return await register_aggregate_client(
request=request, request_body=data, token_exchange_available=token_exchange_available()
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
mcp_server: Final = _resolve_mcp_server_by_name_or_id(mcp_server_name, client_ip)
if mcp_server is None:
return dummy_return
return await register_client_with_server(
request=request,
mcp_server=mcp_server,
client_name=data.get("client_name", ""),
grant_types=data.get("grant_types", []),
response_types=data.get("response_types", []),
token_endpoint_auth_method=data.get("token_endpoint_auth_method", ""),
fallback_client_id=mcp_server_name,
client_redirect_uris=client_redirect_uris,
)
async with global_mcp_server_manager.catalog.operation():
dummy_return: Final = {
"client_id": mcp_server_name or "dummy_client",
"client_secret": "dummy",
"redirect_uris": client_redirect_uris or [f"{request_base_url}/callback"],
}
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
if not mcp_server_name:
resolved: Final = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if resolved:
return await register_client_with_server(
request=request,
mcp_server=resolved,
client_name=data.get("client_name", ""),
grant_types=data.get("grant_types", []),
response_types=data.get("response_types", []),
token_endpoint_auth_method=data.get("token_endpoint_auth_method", ""),
fallback_client_id=resolved.server_name or resolved.name,
client_redirect_uris=client_redirect_uris,
)
return dummy_return
mcp_server: Final = _resolve_mcp_server_by_name_or_id(mcp_server_name, client_ip)
if mcp_server is None:
return dummy_return
return await register_client_with_server(
request=request,
mcp_server=mcp_server,
client_name=data.get("client_name", ""),
grant_types=data.get("grant_types", []),
response_types=data.get("response_types", []),
token_endpoint_auth_method=data.get("token_endpoint_auth_method", ""),
fallback_client_id=mcp_server_name,
client_redirect_uris=client_redirect_uris,
)

View file

@ -56,6 +56,7 @@ from typing_extensions import NotRequired, ReadOnly, TypedDict, assert_never
from litellm._internal_context import with_service_target
from litellm._logging import verbose_logger
from litellm.caching.caching import DualCache
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, global_manager
from litellm.proxy._experimental.mcp_server.oauth_utils import (
TOKEN_NO_CACHE_HEADERS,
canonical_resource_uri,
@ -774,6 +775,7 @@ def _open_flow_for(
return flow
@catalog_operation(global_manager)
async def _flow_target(
flow: _ConnectFlow, lookup_server_reachability: LookupServerReachability
) -> tuple[Literal["unscoped", "interactive", "m2m", "stale"], MCPServer | None]:

View file

@ -49,7 +49,7 @@ from mcp.types import (
)
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl, BaseModel, Field, TypeAdapter
from typing_extensions import ReadOnly
from typing_extensions import ReadOnly, assert_never
import litellm
from litellm._internal_context import with_service_target
@ -195,7 +195,6 @@ from litellm.proxy.middleware.per_request_root_path_middleware import (
get_request_root_path,
)
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.table_repositories import MCPServerRepository
from litellm.types.integrations.slack_alerting import AlertType
from litellm.types.llms.base import LiteLLMBaseModel
from litellm.types.llms.custom_http import httpxSpecialProvider
@ -318,7 +317,7 @@ def _requires_oauth_discovery(
use_issuer_anchor: bool,
server: MCPServer,
) -> bool:
return _has_oauth_discovery_source(server_url, use_issuer_anchor) and _oauth_endpoints_unresolved(server)
return _has_oauth_discovery_source(server_url, use_issuer_anchor) and oauth_endpoints_unresolved(server)
_StringList: TypeAlias = list[str]
@ -573,7 +572,7 @@ def _config_identifier_owners(
)
def _config_ids_capturing_db_identifiers(
def config_ids_capturing_db_identifiers(
config_server_ids: Container[str],
db_servers: Iterable[MCPServer],
) -> frozenset[str]:
@ -753,7 +752,7 @@ def _flow_endpoints_missing(
return authorization_url is None or token_url is None
def _oauth_endpoints_unresolved(server: MCPServer) -> bool:
def oauth_endpoints_unresolved(server: MCPServer) -> bool:
"""``_flow_endpoints_missing`` over a built registry entry, for the reload fast-path check.
The flow comes from ``effective_oauth2_flow``, the one column-first, shape-fallback judge every
@ -812,7 +811,7 @@ def _endpoints_corroborate_authorization_url(
) == _normalized_authorize_endpoint(trusted_authorization_url)
def _carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_server: MCPServer | None) -> None:
def carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_server: MCPServer | None) -> None:
"""Keep the last known good OAuth endpoints when a rebuild's re-discovery comes back empty.
A rebuild wholesale-replaces the registry entry, so without this a transient upstream outage
@ -1511,7 +1510,7 @@ def _obo_retry_applies(server: MCPServer, subject_token: str | None) -> bool:
return server.auth_type == MCPAuth.oauth2_token_exchange and bool(subject_token)
def _warn_on_server_name_fields(
def warn_on_server_name_fields(
*,
server_id: str,
alias: str | None,
@ -1537,7 +1536,7 @@ def _warn_on_server_name_fields(
_warn("server_name", server_name)
def _warn_on_shared_identifier_prefixes(servers: Iterable[MCPServer]) -> None:
def warn_on_shared_identifier_prefixes(servers: Iterable[MCPServer]) -> None:
"""Warn once per identifier that several servers share.
``get_server_prefix`` resolves alias first, so two servers sharing a
@ -1971,6 +1970,9 @@ class MCPServerManager:
self._template_discovery_cache = _DiscoveryCache[ResourceTemplate](
discovery_ttl, discovery_clock, TypeAdapter(tuple[ResourceTemplate, ...])
)
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
self.catalog = TargetCatalog(self)
self.registry: dict[str, MCPServer] = {}
self._openapi_health_probes: Callable[[str], _OpenAPIHealthProbe] = lru_cache(maxsize=128)(_OpenAPIHealthProbe)
self.config_mcp_servers: dict[str, MCPServer] = {}
@ -1997,7 +1999,7 @@ class MCPServerManager:
# semaphore so an edited limit rebuilds it instead of keeping the old cap
# until restart.
self._server_call_semaphores: dict[str, tuple[int, asyncio.Semaphore]] = {}
self.tool_name_to_mcp_server_name_mapping: dict[str, str] = {}
self.published_tool_routes: dict[str, str] = {}
"""
{
"gmail_send_email": "zapier_mcp_server",
@ -2013,14 +2015,12 @@ class MCPServerManager:
# Last set of config server ids found shadowed by database rows. reload_servers_from_database
# runs on the config-reload timer, so this keeps a standing misconfiguration from re-logging
# the same warning every interval; a change in the set logs again.
self._warned_shadowed_config_server_ids: frozenset[str] = frozenset()
self._warned_capturing_config_server_ids: frozenset[str] = frozenset()
self._catalog_alert_signatures: Mapping[tuple[str, AlertType], str] = MappingProxyType({})
self._oauth_discovery_on_startup = _mcp_oauth_discovery_on_startup_enabled()
self._oauth_discovery_generation_counter = 0
self._oauth_discovery_slots: tuple[_OAuthDiscoverySlot, ...] = ()
def _oauth_discovery_slot(self, server_id: str) -> _OAuthDiscoverySlot | None:
def oauth_discovery_slot(self, server_id: str) -> _OAuthDiscoverySlot | None:
return next((slot for slot in self._oauth_discovery_slots if slot.server_id == server_id), None)
def _remove_oauth_discovery_slot(self, server_id: str) -> None:
@ -2033,7 +2033,7 @@ class MCPServerManager:
)
def _set_oauth_discovery_deferred(self, server_id: str, discovery_deferred: bool) -> None:
previous: Final = self._oauth_discovery_slot(server_id)
previous: Final = self.oauth_discovery_slot(server_id)
self._remove_oauth_discovery_slot(server_id)
if previous is not None and previous.task is not None and not previous.task.done():
previous.task.cancel()
@ -2046,8 +2046,8 @@ class MCPServerManager:
)
)
def _invalidate_oauth_discovery_state(self, server_id: str) -> None:
previous: Final = self._oauth_discovery_slot(server_id)
def invalidate_oauth_discovery_state(self, server_id: str) -> None:
previous: Final = self.oauth_discovery_slot(server_id)
self._remove_oauth_discovery_slot(server_id)
if previous is not None and previous.task is not None and not previous.task.done():
previous.task.cancel()
@ -2129,7 +2129,7 @@ class MCPServerManager:
)
def _oauth_discovery_slot_is_current(self, server_id: str, generation: int) -> bool:
slot: Final = self._oauth_discovery_slot(server_id)
slot: Final = self.oauth_discovery_slot(server_id)
return slot is not None and slot.generation == generation
def _expire_temporary_oauth_discovery(self, server_id: str, generation: int) -> None:
@ -2166,7 +2166,7 @@ class MCPServerManager:
if not self._oauth_discovery_slot_is_current(server.server_id, generation):
return _OAuthDiscoveryStale(server_id=server.server_id)
current: Final = self._registered_server(server)
if not _oauth_endpoints_unresolved(current):
if not oauth_endpoints_unresolved(current):
published: Final = self._publish_resolved_oauth_server(current, generation)
return (
_OAuthDiscoveryResolved(server=published)
@ -2177,7 +2177,7 @@ class MCPServerManager:
if not self._oauth_discovery_slot_is_current(server.server_id, generation):
return _OAuthDiscoveryStale(server_id=server.server_id)
candidate: Final = self._merge_discovered_oauth_metadata(self._registered_server(server), metadata)
if _oauth_endpoints_unresolved(candidate):
if oauth_endpoints_unresolved(candidate):
return None
published_candidate: Final = self._publish_resolved_oauth_server(candidate, generation)
return (
@ -2224,7 +2224,7 @@ class MCPServerManager:
return outcome
def _record_oauth_discovery_failure(self, server_id: str, generation: int) -> None:
slot: Final = self._oauth_discovery_slot(server_id)
slot: Final = self.oauth_discovery_slot(server_id)
if slot is None or slot.generation != generation:
return
consecutive_failures: Final = slot.consecutive_failures + 1
@ -2240,7 +2240,7 @@ class MCPServerManager:
self,
server: MCPServer,
) -> tuple[asyncio.Task[_OAuthDiscoveryOutcome], int] | None:
slot: Final = self._oauth_discovery_slot(server.server_id)
slot: Final = self.oauth_discovery_slot(server.server_id)
if slot is None:
return None
if slot.task is not None:
@ -2269,19 +2269,24 @@ class MCPServerManager:
"""
self._get_or_start_oauth_discovery_task(server)
def _prime_oauth_metadata_discovery_for_servers(self, servers: Sequence[MCPServer]) -> None:
def prime_oauth_metadata_discovery_for_servers(self, servers: Sequence[MCPServer]) -> None:
for server in servers:
self.prime_oauth_metadata_discovery(server)
def _reconcile_oauth_discovery_slots_for_servers(self, servers: Sequence[MCPServer]) -> None:
def reconcile_oauth_discovery_slots_for_servers(self, servers: Sequence[MCPServer]) -> None:
"""Align retry slots after an atomic registry replacement."""
for server in servers:
should_defer = _requires_oauth_discovery(server.url, server.issuer_is_anchored, server)
has_slot = self._oauth_discovery_slot(server.server_id) is not None
has_slot = self.oauth_discovery_slot(server.server_id) is not None
if should_defer != has_slot:
self._set_oauth_discovery_deferred(server.server_id, should_defer)
async def ensure_oauth_metadata_discovered(self, server: MCPServer, *, _retry_stale: bool = True) -> MCPServer:
return await self.catalog.resolve_oauth_metadata(
server, lambda selected: self._ensure_oauth_metadata_discovered(selected, _retry_stale=_retry_stale)
)
async def _ensure_oauth_metadata_discovered(self, server: MCPServer, *, _retry_stale: bool = True) -> MCPServer:
"""Join the bounded discovery task and return the resolved server.
Concurrent callers share one task per server. A failed attempt remains
@ -2300,6 +2305,7 @@ class MCPServerManager:
incomplete metadata for a server whose OAuth flow the gateway
runs itself.
"""
self.catalog.assert_current(server)
acquisition: Final = self._get_or_start_oauth_discovery_task(server)
if acquisition is None:
return self._registered_server(server)
@ -2312,11 +2318,13 @@ class MCPServerManager:
raise
match outcome:
case _OAuthDiscoveryResolved(resolved_server):
self.catalog.assert_current(resolved_server)
return resolved_server
case _OAuthDiscoveryStale():
return await self._rejoin_oauth_metadata_discovery(server, retry_stale=_retry_stale)
case _OAuthDiscoveryFailed(timed_out=timed_out):
current: Final = self._registered_server(server)
self.catalog.assert_current(current)
if current.is_client_forwarded_token:
return current
server_ref: Final = current.alias or current.server_name or current.name or current.server_id
@ -2326,11 +2334,13 @@ class MCPServerManager:
detail=f"OAuth metadata discovery {reason} for MCP server {server_ref!r}",
)
return assert_never(outcome)
async def _rejoin_oauth_metadata_discovery(self, server: MCPServer, *, retry_stale: bool) -> MCPServer:
if retry_stale:
return await self.ensure_oauth_metadata_discovered(server, _retry_stale=False)
current: Final = self._registered_server(server)
if not _oauth_endpoints_unresolved(current) or current.is_client_forwarded_token:
if not oauth_endpoints_unresolved(current) or current.is_client_forwarded_token:
return current
raise HTTPException(status_code=503, detail="OAuth metadata discovery changed repeatedly; retry shortly")
@ -2408,11 +2418,21 @@ class MCPServerManager:
e,
)
def get_registry(self) -> dict[str, MCPServer]:
def get_registry(self) -> Mapping[str, MCPServer]:
"""
Get the registered MCP Servers from the registry and union with the config MCP Servers
"""
return self.config_mcp_servers | self.registry
return self.catalog.registry()
@property
def tool_name_to_mcp_server_name_mapping(
self,
) -> MutableMapping[str, str]:
return self.catalog.routing()
@tool_name_to_mcp_server_name_mapping.setter
def tool_name_to_mcp_server_name_mapping(self, mapping: dict[str, str]) -> None:
self.published_tool_routes = mapping
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).
@ -2499,7 +2519,7 @@ class MCPServerManager:
)
assigned_server_ids[server_id] = server_name
_warn_on_server_name_fields(
warn_on_server_name_fields(
server_id=server_id,
alias=alias,
server_name=server_name,
@ -2715,10 +2735,10 @@ class MCPServerManager:
token_validation=server_config.get("token_validation", None),
oauth_identity_binding=server_config.get("oauth_identity_binding", None),
)
self._assign_unique_short_prefix(new_server)
self.assign_unique_short_prefix(new_server)
_warn_legacy_delegate_auth_if_applicable(new_server, source="config")
_warn_config_id_jag_server_outruns_sso(new_server)
self._invalidate_server_definition_caches(server_id)
self.invalidate_server_definition_caches(server_id)
self.config_mcp_servers[server_id] = new_server
self._set_oauth_discovery_deferred(
server_id,
@ -2739,13 +2759,13 @@ class MCPServerManager:
"Loaded MCP Servers: %s", json.dumps(_redacted_registry_dump(self.config_mcp_servers), indent=4)
)
await self._hydrate_config_servers_dcr_clients()
await self.hydrate_config_servers_dcr_clients()
self._prime_oauth_metadata_discovery_for_servers(tuple(self.config_mcp_servers.values()))
self.prime_oauth_metadata_discovery_for_servers(tuple(self.config_mcp_servers.values()))
self.initialize_tool_name_to_mcp_server_name_mapping()
async def _hydrate_config_servers_dcr_clients(self) -> None:
async def hydrate_config_servers_dcr_clients(self, servers: Sequence[MCPServer] | None = None) -> 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."""
@ -2753,7 +2773,7 @@ class MCPServerManager:
hydrate_config_server_dcr_client,
)
for server in self.config_mcp_servers.values():
for server in servers if servers is not None else self.config_mcp_servers.values():
try:
if await hydrate_config_server_dcr_client(server):
verbose_logger.debug(
@ -2894,6 +2914,7 @@ class MCPServerManager:
description=description,
input_schema=input_schema,
handler=tool_func,
server_id=server.server_id,
)
# Update tool name to server name mapping (for both prefixed and base names)
@ -2918,17 +2939,18 @@ class MCPServerManager:
mappings make ``_get_mcp_server_from_tool_name`` resolve to a prefix that
no longer exists in the live registry.
"""
from litellm.proxy._experimental.mcp_server.tool_registry import (
global_mcp_tool_registry,
)
self.invalidate_server_definition_caches(server.server_id)
self.remove_server_tool_routing(server)
def remove_server_tool_routing(self, server: MCPServer) -> None:
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
self._invalidate_server_definition_caches(server.server_id)
prefix_root: Final = normalize_server_name(get_server_prefix(server))
if server.spec_path and prefix_root:
openapi_key_prefix: Final = prefix_root + MCP_TOOL_PREFIX_SEPARATOR
global_mcp_tool_registry.unregister_tools_with_prefix(openapi_key_prefix)
owned_normalized: Final = self._owned_mapping_values(server)
owned_normalized: Final = self.owned_mapping_values(server)
stale_mapping_keys: Final = tuple(
tool_name
@ -2939,13 +2961,13 @@ class MCPServerManager:
for key in stale_mapping_keys:
del self.tool_name_to_mcp_server_name_mapping[key]
def _owned_mapping_values(self, server: MCPServer) -> frozenset[str]:
def owned_mapping_values(self, server: MCPServer) -> frozenset[str]:
return frozenset(
normalize_server_name(value) for value in (*iter_known_server_prefixes(server), server.name) if value
)
def server_exposes_tool(self, server: MCPServer, tool_name: str) -> bool:
owned: Final = self._owned_mapping_values(server)
owned: Final = self.owned_mapping_values(server)
mapped_owners: Final = (
self.tool_name_to_mcp_server_name_mapping.get(spelling)
for spelling in iter_known_tool_name_spellings(tool_name, server)
@ -2976,7 +2998,7 @@ class MCPServerManager:
if evicted is not None:
verbose_logger.debug("Removed MCP Server: %s", mcp_server.server_id or mcp_server.server_name)
self._cleanup_server_tool_routing_artifacts(evicted)
self._invalidate_oauth_discovery_state(evicted.server_id)
self.invalidate_oauth_discovery_state(evicted.server_id)
else:
verbose_logger.warning("Server ID %s not found in registry", mcp_server.server_id)
@ -3067,6 +3089,7 @@ class MCPServerManager:
*,
credentials_are_encrypted: bool = True,
env_vars_are_encrypted: bool | None = None,
register_oauth_discovery: bool = True,
) -> MCPServer:
_mcp_info: Final[MCPInfo] = mcp_server.mcp_info or {}
env_dict: Final = _deserialize_json_dict(getattr(mcp_server, "env", None))
@ -3300,13 +3323,14 @@ class MCPServerManager:
max_concurrent_requests=getattr(mcp_server, "max_concurrent_requests", None),
)
_warn_legacy_delegate_auth_if_applicable(new_server, source="database")
self._set_oauth_discovery_deferred(
new_server.server_id,
_requires_oauth_discovery(server_url, use_issuer_anchor, new_server),
)
if register_oauth_discovery:
self._set_oauth_discovery_deferred(
new_server.server_id,
_requires_oauth_discovery(server_url, use_issuer_anchor, new_server),
)
return new_server
async def _maybe_register_openapi_tools(self, server: MCPServer, *, initialize_mapping: bool = True):
async def maybe_register_openapi_tools(self, server: MCPServer, *, initialize_mapping: bool = True):
"""Register OpenAPI tools if the server has a spec_path configured."""
if server.spec_path:
verbose_logger.info("Loading OpenAPI spec from %s for server %s", server.spec_path, server.name)
@ -3333,12 +3357,12 @@ class MCPServerManager:
# `credentials` field is the only one still encrypted here).
# Re-decrypting plaintext would zero the values, so build with
# env_vars_are_encrypted=False.
self._warn_if_newly_blocked_stdio(mcp_server, None)
self.warn_if_newly_blocked_stdio(mcp_server, None)
new_server: Final = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False)
self._assign_unique_short_prefix(new_server)
self._invalidate_server_definition_caches(mcp_server.server_id)
self.assign_unique_short_prefix(new_server)
self.invalidate_server_definition_caches(mcp_server.server_id)
self.registry[mcp_server.server_id] = new_server
await self._maybe_register_openapi_tools(new_server)
await self.maybe_register_openapi_tools(new_server)
if new_server.spec_path:
self._drop_listed_tools(mcp_server.server_id)
self.prime_oauth_metadata_discovery(new_server)
@ -3358,7 +3382,7 @@ class MCPServerManager:
evicted = self.registry.pop(mcp_server.server_name, None)
if evicted is not None:
self._cleanup_server_tool_routing_artifacts(evicted)
self._invalidate_oauth_discovery_state(evicted.server_id)
self.invalidate_oauth_discovery_state(evicted.server_id)
return
try:
if mcp_server.server_id in self.registry:
@ -3370,14 +3394,14 @@ class MCPServerManager:
existing_prefix: Final = self.registry[mcp_server.server_id].short_prefix
if existing_prefix and not new_server.short_prefix:
new_server.short_prefix = existing_prefix
_carry_forward_resolved_oauth_endpoints(
carry_forward_resolved_oauth_endpoints(
new_server=new_server,
previous_server=self.registry[mcp_server.server_id],
)
self._assign_unique_short_prefix(new_server)
self._invalidate_server_definition_caches(mcp_server.server_id)
self.assign_unique_short_prefix(new_server)
self.invalidate_server_definition_caches(mcp_server.server_id)
self.registry[mcp_server.server_id] = new_server
await self._maybe_register_openapi_tools(new_server)
await self.maybe_register_openapi_tools(new_server)
if new_server.spec_path:
self._drop_listed_tools(mcp_server.server_id)
self.prime_oauth_metadata_discovery(new_server)
@ -4200,6 +4224,7 @@ class MCPServerManager:
client_ip: str | None = None,
protocol_version_override: MCPUpstreamProtocol | None = None,
) -> MCPClient:
self.catalog.assert_current(server)
record_auth_resolution(server.server_id, AuthResolution.unresolved)
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
return await prepare_upstream_client(
@ -4433,17 +4458,23 @@ class MCPServerManager:
)
raise_classified_list_failure(e, server.name, suppress_challenge=server.is_dcr_bridge)
def _invalidate_discovery_lists(self, server_id: str) -> None:
def clear_initialize_instructions(self) -> None:
self._upstream_initialize_instructions_by_server_id.clear()
self._upstream_initialize_instructions_probed_at.clear()
def invalidate_discovery_lists(self, server_id: str) -> None:
self._upstream_initialize_instructions_by_server_id.pop(server_id, None)
self._upstream_initialize_instructions_probed_at.pop(server_id, None)
self._prompt_discovery_cache.invalidate(server_id)
self._resource_discovery_cache.invalidate(server_id)
self._template_discovery_cache.invalidate(server_id)
def _invalidate_server_definition_caches(self, server_id: str) -> None:
def invalidate_server_definition_caches(self, server_id: str) -> None:
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( # noqa: PLC0415 # lazy: discoverable_endpoints lazily imports this module's manager singleton
invalidate_oauth_metadata_cache,
)
self._invalidate_discovery_lists(server_id)
self.invalidate_discovery_lists(server_id)
self._drop_listed_tools(server_id)
invalidate_oauth_metadata_cache(server_id)
@ -4559,7 +4590,7 @@ class MCPServerManager:
return server.server_id, hashlib.sha256(material.encode()).hexdigest()
@staticmethod
def _warn_if_newly_blocked_stdio(row: LiteLLM_MCPServerTable, previous: MCPServer | None) -> None:
def warn_if_newly_blocked_stdio(row: LiteLLM_MCPServerTable, previous: MCPServer | None) -> None:
if previous is None or previous.transport != row.transport:
warn_if_mcp_stdio_blocked(row.alias or row.server_name, row.transport)
@ -5336,7 +5367,7 @@ class MCPServerManager:
_SHORT_PREFIX_MAX_REHASH_ATTEMPTS = 1024
def _assign_unique_short_prefix(
def assign_unique_short_prefix(
self,
server: MCPServer,
registry: dict[str, MCPServer] | None = None,
@ -5467,6 +5498,8 @@ class MCPServerManager:
Returns:
List of tools with prefixed names
"""
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
prefixed_tools: Final = []
prefix: Final = get_server_prefix(server)
@ -5479,6 +5512,11 @@ class MCPServerManager:
# short ID) so call_tool can resolve regardless of which form a
# caller / cached client is using.
for spelling in iter_known_tool_name_spellings(original_name, server):
namespace_owner = self.server_owning_tool_name_prefix(spelling)
if namespace_owner is not None and namespace_owner.server_id != server.server_id:
continue
if namespace_owner is None and global_mcp_tool_registry.get_tool(spelling) is not None:
continue
self.tool_name_to_mcp_server_name_mapping[spelling] = prefix
verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name)
@ -5682,6 +5720,7 @@ class MCPServerManager:
# Registration used add_server_prefix_to_name(base, get_server_prefix(server)),
# and tool_name is the bare base name by the time call_tool reaches here, so
# rebuilding the key the same way reproduces it exactly
self.catalog.assert_current(server)
registry_key: Final = add_server_prefix_to_name(tool_name, get_server_prefix(server))
tool: Final = global_mcp_tool_registry.get_tool(registry_key)
if tool is None:
@ -6300,7 +6339,7 @@ class MCPServerManager:
failure is logged, never raised, because the DB write already succeeded and the TTL remains
the backstop.
"""
self._invalidate_discovery_lists(server_id)
self.invalidate_discovery_lists(server_id)
try:
await self._per_user_oauth_token_store.invalidate(user_id, server_id)
except Exception as exc: # noqa: BLE001 - cache drop is best-effort; TTL is the backstop
@ -6612,7 +6651,7 @@ class MCPServerManager:
Note: This now handles prefixed tool names
"""
for server in self.get_registry().values():
if self._oauth_discovery_slot(server.server_id) is not None:
if self.oauth_discovery_slot(server.server_id) is not None:
continue
if server.needs_user_oauth_token:
# Skip OAuth2 servers that rely on user-provided tokens
@ -6649,6 +6688,14 @@ class MCPServerManager:
Returns:
MCPServer if found, None otherwise
"""
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
registered_tool: Final = global_mcp_tool_registry.get_tool(tool_name)
if registered_tool is not None:
if registered_tool.server_id is not None:
return self.get_mcp_server_by_id(registered_tool.server_id)
return self.server_owning_tool_name_prefix(tool_name)
registry_servers: Final = list(self.get_registry().values())
prefix_to_server: Final = self._known_prefix_to_server()
@ -6677,154 +6724,7 @@ class MCPServerManager:
return None
async def reload_servers_from_database(self):
"""Re-synchronize the in-memory MCP server registry with the database."""
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
get_prisma_client_or_throw,
)
verbose_logger.debug("Loading MCP servers from database into registry...")
self._upstream_initialize_instructions_by_server_id.clear()
self._upstream_initialize_instructions_probed_at.clear()
# perform authz check to filter the mcp servers user has access to
prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy")
# Load only "active", legacy "approved", and NULL (no approval workflow) rows.
# Pending/rejected servers are excluded at the DB level so we never load them.
from litellm.proxy._experimental.mcp_server.db import LiteLLM_MCPServerTable
raw_rows: Final[Sequence[BaseModel]] = await MCPServerRepository(prisma_client).table.find_many(
where={
"OR": [
{"approval_status": None},
{"approval_status": {"in": ["active", "approved"]}},
]
}
)
verbose_logger.info("Found %s MCP servers in database", len(raw_rows))
previous_registry: Final = self.registry
new_registry: Final[dict[str, MCPServer]] = {}
# Stage one: build every server. Stage two assigns short prefixes
# against the *full* set so dedup is deterministic regardless of
# iteration order.
for row in raw_rows:
try:
server = LiteLLM_MCPServerTable.model_validate(row.model_dump())
existing_server = previous_registry.get(server.server_id)
if (
existing_server is not None
and existing_server.updated_at is not None
and server.updated_at is not None
and existing_server.updated_at == server.updated_at
and (
self._oauth_discovery_slot(server.server_id) is not None
or not _oauth_endpoints_unresolved(existing_server)
)
):
# Re-use existing server instance to avoid re-running build_mcp_server_from_table()
# which can perform network discovery for OAuth2 servers.
new_registry[server.server_id] = existing_server
continue
_warn_on_server_name_fields(
server_id=server.server_id,
alias=getattr(server, "alias", None),
server_name=getattr(server, "server_name", None),
)
self._warn_if_newly_blocked_stdio(server, existing_server)
verbose_logger.debug("Building server from DB: %s (%s)", server.server_id, server.server_name)
# raw_rows come straight from the DB, so their global env var
# values (like credentials) are still encrypted here, unlike the
# already-decrypted records add_server/update_server are handed.
# Decrypt them while building the registry entry.
new_server = await self.build_mcp_server_from_table(server, env_vars_are_encrypted=True)
# Carry the cached short_prefix from the previous registry entry
# (if any) so the prefix is stable across reloads.
if existing_server is not None and existing_server.short_prefix:
new_server.short_prefix = existing_server.short_prefix
_carry_forward_resolved_oauth_endpoints(new_server=new_server, previous_server=existing_server)
new_registry[server.server_id] = new_server
except Exception as e:
verbose_logger.exception(
"Skipping MCP server %s (%s) during DB reload: %s",
getattr(row, "server_id", None),
getattr(row, "alias", None),
e,
)
# Assign short prefixes against the full candidate set without
# publishing the staged registry to concurrent callers.
registered_registry: Final[dict[str, MCPServer]] = {}
registered_openapi_tools = False
for server_id, new_server in new_registry.items():
try:
self._assign_unique_short_prefix(new_server, registry=new_registry)
# Register OpenAPI tools *after* the final short prefix is assigned
# so the tools are stored in the global registry under the same
# prefix that lookups will use.
await self._maybe_register_openapi_tools(new_server, initialize_mapping=False)
registered_registry[server_id] = new_server
if new_server.spec_path:
registered_openapi_tools = True
except Exception as e:
verbose_logger.exception(
"Skipping MCP server %s (%s) during DB reload: %s",
new_server.server_id,
getattr(new_server, "alias", None),
e,
)
dropped_registry_keys: Final = previous_registry.keys() - registered_registry.keys()
for registry_key in dropped_registry_keys:
self._invalidate_oauth_discovery_state(previous_registry[registry_key].server_id)
for server_id in previous_registry.keys() | registered_registry.keys():
if previous_registry.get(server_id) != registered_registry.get(server_id):
self._invalidate_server_definition_caches(server_id)
self.registry = registered_registry
_warn_on_shared_identifier_prefixes(registered_registry.values())
# A discovery task may have published into ``previous_registry`` while
# this replacement was being staged. Reconcile every published entry
# synchronously after the swap so a lost publication cannot also leave
# the replacement unresolved with no retry slot.
registered_servers: Final = tuple(registered_registry.values())
self._reconcile_oauth_discovery_slots_for_servers(registered_servers)
self._prime_oauth_metadata_discovery_for_servers(registered_servers)
if registered_openapi_tools:
self.initialize_tool_name_to_mcp_server_name_mapping()
verbose_logger.debug("MCP registry refreshed (%s servers in registry)", len(registered_registry))
# get_registry() is ``config_mcp_servers | registry``, so a database row sharing an id with a
# config.yaml server hides that server everywhere. Only reachable once an operator pins
# ``server_id`` in config.yaml; say so rather than letting the server disappear silently.
shadowed_config_server_ids: Final = frozenset(self.config_mcp_servers.keys() & registered_registry.keys())
if shadowed_config_server_ids and shadowed_config_server_ids != self._warned_shadowed_config_server_ids:
verbose_logger.warning(
"config.yaml MCP server_id(s) %s are also database-backed MCP servers. The database "
"entry takes precedence, so the config.yaml server is unreachable. Give the config "
"entry a different server_id.",
", ".join(sorted(shadowed_config_server_ids)),
)
self._warned_shadowed_config_server_ids = shadowed_config_server_ids
# The mirror image of the block above: a config server_id that is a database server's name
# answers that server's grants instead, because ids are matched before names.
capturing_config_server_ids: Final = _config_ids_capturing_db_identifiers(
self.config_mcp_servers.keys(), registered_registry.values()
)
if capturing_config_server_ids and capturing_config_server_ids != self._warned_capturing_config_server_ids:
verbose_logger.warning(
"config.yaml MCP server_id(s) %s are the name or alias of a database-backed MCP "
"server. Permission entries naming them resolve to the config.yaml server, not the "
"database one. Give the config entry a different server_id.",
", ".join(sorted(capturing_config_server_ids)),
)
self._warned_capturing_config_server_ids = capturing_config_server_ids
await self._hydrate_config_servers_dcr_clients()
await self.catalog.reload()
def get_mcp_servers_from_ids(self, server_ids: list[str]) -> list[MCPServer]:
servers: Final = []
@ -7015,7 +6915,7 @@ class MCPServerManager:
return server
return None
def get_filtered_registry(self, client_ip: str | None = None) -> dict[str, MCPServer]:
def get_filtered_registry(self, client_ip: str | None = None) -> Mapping[str, MCPServer]:
"""
Get registry filtered by client IP access control.

View file

@ -62,6 +62,7 @@ from litellm.proxy._experimental.mcp_server.capabilities import (
build_discovery,
configured_versions,
)
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, global_manager
from litellm.proxy._experimental.mcp_server.contracts import (
AuthorizedToolCall,
OperationContext,
@ -627,6 +628,7 @@ def apply_display_name_overrides(
return tools
@catalog_operation(lambda: global_mcp_server_manager)
async def _get_allowed_mcp_servers(
user_api_key_auth: UserAPIKeyAuth | None,
mcp_servers: Sequence[str] | None,
@ -943,6 +945,7 @@ def _aggregate_server_key(server: MCPServer) -> str:
return get_server_prefix(server) or "unknown"
@catalog_operation(lambda: global_mcp_server_manager)
async def _get_tools_from_mcp_servers(
user_api_key_auth: UserAPIKeyAuth | None,
mcp_auth_header: str | None,
@ -1474,6 +1477,7 @@ async def filter_tools_by_key_team_permissions(
]
@catalog_operation(lambda: global_mcp_server_manager)
async def _list_mcp_tools(
user_api_key_auth: UserAPIKeyAuth | None = None,
mcp_auth_header: str | None = None,
@ -1529,6 +1533,7 @@ async def _list_mcp_tools(
return AggregateToolListing(tools=[], outcomes={})
@catalog_operation(global_manager)
async def _list_mcp_prompts(
user_api_key_auth: UserAPIKeyAuth | None = None,
mcp_auth_header: str | None = None,
@ -1570,6 +1575,7 @@ async def _list_mcp_prompts(
return managed_prompts
@catalog_operation(global_manager)
async def _list_mcp_resources(
user_api_key_auth: UserAPIKeyAuth | None = None,
mcp_auth_header: str | None = None,
@ -1599,6 +1605,7 @@ async def _list_mcp_resources(
return managed_resources
@catalog_operation(global_manager)
async def _list_mcp_resource_templates(
user_api_key_auth: UserAPIKeyAuth | None = None,
mcp_auth_header: str | None = None,
@ -2087,6 +2094,9 @@ async def _execute_mcp_tool(
),
)
if local_tool.server_id is not None and local_tool.server_id != mcp_server.server_id:
raise HTTPException(status_code=403, detail="User not allowed to call this tool.")
# `pre_call_tool_check` calls into `proxy_logging_obj` for the
# pre-call guardrail hooks, so source it from the canonical
# `proxy_server` module the same way `_handle_managed_mcp_tool`
@ -2217,6 +2227,12 @@ async def _execute_mcp_tool(
),
)
if (
registered_local_tool.server_id is not None
and registered_local_tool.server_id != prefix_server.server_id
):
raise HTTPException(status_code=403, detail="User not allowed to call this tool.")
from litellm.proxy.proxy_server import proxy_logging_obj
hook_result = await global_mcp_server_manager.pre_call_tool_check(
@ -2398,6 +2414,7 @@ async def fire_mcp_tool_call_failure_logging(
@client
@catalog_operation(lambda: global_mcp_server_manager)
async def call_mcp_tool(
name: str,
arguments: dict[str, object] | None = None,
@ -2694,6 +2711,15 @@ async def _handle_local_mcp_tool(
tool: Final = global_mcp_tool_registry.get_tool(name)
if not tool:
raise HTTPException(status_code=404, detail=f"Tool '{name}' not found")
server: Final = (
global_mcp_server_manager.get_mcp_server_by_id(tool.server_id)
if tool.server_id is not None
else global_mcp_server_manager.server_owning_tool_name_prefix(name)
)
if tool.server_id is not None and server is None:
raise HTTPException(status_code=503, detail="MCP server configuration changed; retry the operation")
if server is not None:
global_mcp_server_manager.catalog.assert_current(server)
try:
if inspect.iscoroutinefunction(tool.handler):
@ -3217,6 +3243,7 @@ class GatewayOperations:
@overload
async def execute(self, operation: ReadResourceRequest, context: OperationContext) -> ReadResourceResult: ...
@catalog_operation(lambda: global_mcp_server_manager)
async def execute(self, operation: GatewayOperation, context: OperationContext) -> GatewayResult:
match operation:
case DiscoverRequest():

View file

@ -23,6 +23,7 @@ from litellm.exceptions import (
GuardrailRaisedException,
ModifyResponseException,
)
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, global_manager
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPServerListError,
MCPServerURLCredentialsError,
@ -950,6 +951,7 @@ if MCP_AVAILABLE:
return await _apply_toolset_scope(user_api_key_dict, toolset.toolset_id)
@router.get("/tools/list", dependencies=[Depends(user_api_key_auth)])
@catalog_operation(global_manager)
async def list_tool_rest_api(
request: Request,
server_id: str | None = Query(None, description="The server id to list tools for"),
@ -1174,6 +1176,7 @@ if MCP_AVAILABLE:
}
@router.post("/tools/call", dependencies=[Depends(user_api_key_auth)])
@catalog_operation(global_manager)
async def call_tool_rest_api(
request: Request,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),

View file

@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Any, Final
from litellm._logging import verbose_logger
from litellm.exceptions import ContextWindowExceededError
from litellm.litellm_core_utils.exception_mapping_utils import ExceptionCheckers
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, global_manager
from litellm.proxy._experimental.mcp_server.faults import iter_exception_tree
from litellm.proxy._experimental.mcp_server.utils import MCP_TOOL_PREFIX_SEPARATOR
@ -78,6 +79,7 @@ class SemanticMCPToolFilter:
self._tool_map: dict[str, object] = {} # MCPTool objects or OpenAI function dicts
self._index_sync_lock = asyncio.Lock()
@catalog_operation(global_manager)
async def build_router_from_mcp_registry(self) -> None:
"""Build semantic router from all MCP tools in the registry (no auth checks)."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (

View file

@ -1599,6 +1599,9 @@ if MCP_AVAILABLE:
await operations.global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=[toolset_id])
)
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation
@catalog_operation(lambda: operations.global_mcp_server_manager)
async def _raise_preemptive_401_for_unauthenticated_servers(
scope: Scope,
mcp_servers: list[str] | None,

View file

@ -1,5 +1,8 @@
import asyncio
import json
from collections.abc import Callable
from collections.abc import Callable, Generator, Mapping
from contextlib import contextmanager
from contextvars import ContextVar
from typing import TYPE_CHECKING, Any, Final
from litellm._logging import verbose_logger
@ -22,7 +25,32 @@ class MCPToolRegistry:
def __init__(self):
# Registry to store all registered tools
self.tools: dict[str, MCPTool] = {}
self.published_tools: dict[str, MCPTool] = {} # mutable-ok: register_tool writes entries through .tools
self._catalog_tools: ContextVar[tuple[dict[str, MCPTool], asyncio.Event] | None] = ContextVar(
"mcp_catalog_tools", default=None
)
@property
def tools(self) -> dict[str, MCPTool]: # mutable-ok: register_tool mutates the returned mapping
scoped: Final = self._catalog_tools.get()
return scoped[0] if scoped is not None and not scoped[1].is_set() else self.published_tools
@tools.setter
def tools(self, tools: dict[str, MCPTool]) -> None: # mutable-ok: stored dict is mutated by register_tool
self.published_tools = tools
@contextmanager
def catalog_scope(
self, tools: Mapping[str, MCPTool]
) -> Generator[dict[str, MCPTool]]: # mutable-ok: yields the mutable staged copy
detached: Final = dict(tools)
closed: Final = asyncio.Event()
token: Final = self._catalog_tools.set((detached, closed))
try:
yield detached
finally:
closed.set()
self._catalog_tools.reset(token)
def register_tool(
self,
@ -30,6 +58,8 @@ class MCPToolRegistry:
description: str,
input_schema: dict[str, Any],
handler: Callable,
*,
server_id: str | None = None,
) -> None:
"""
Register a new tool in the registry
@ -39,6 +69,7 @@ class MCPToolRegistry:
description=description,
input_schema=input_schema,
handler=handler,
server_id=server_id,
)
verbose_logger.debug("Registered tool: %s", name)

View file

@ -76,7 +76,7 @@ MCP_TOOL_PREFIX_FORMAT: Final = "{server_name}{separator}{tool_name}"
# principle hash to the same three chars; that natural-hash collision
# IS a routing-correctness issue (the second registrant would otherwise
# have its tools misrouted to the first), so registration goes through
# ``MCPServerManager._assign_unique_short_prefix`` which rehashes with
# ``MCPServerManager.assign_unique_short_prefix`` which rehashes with
# a deterministic attempt counter until it finds an unused prefix and
# caches the result on ``MCPServer.short_prefix``. A collision is
# logged at INFO when it happens.
@ -109,7 +109,7 @@ def compute_short_server_prefix(server_id: str, attempt: int = 0) -> str:
and whose remaining characters are drawn from the full base62
alphabet. Pass ``attempt > 0`` to rehash to a different prefix when
the natural hash collides with a prefix already assigned to another
server (see ``MCPServerManager._assign_unique_short_prefix``). An
server (see ``MCPServerManager.assign_unique_short_prefix``). An
empty ``server_id`` raises ``ValueError`` — short prefixes require a
stable identifier to be deterministic.
"""
@ -309,7 +309,7 @@ def get_server_prefix(server: object) -> str:
When the short-prefix mode is enabled (``LITELLM_USE_SHORT_MCP_TOOL_PREFIX``)
a three-character base62 ID is returned. We prefer the cached
``server.short_prefix`` value when set — that field is populated at
registration time by ``MCPServerManager._assign_unique_short_prefix``
registration time by ``MCPServerManager.assign_unique_short_prefix``
and resolves natural-hash collisions deterministically — and only fall
back to the natural hash for ad-hoc / temp-server objects without a
cached value. In default mode the historical behaviour is preserved:

View file

@ -5,6 +5,7 @@ Validates that MCP servers referenced in request tools are registered
on the LiteLLM gateway. Blocks or alerts when unregistered servers are found.
"""
from collections.abc import Mapping
from typing import Any, Final, Literal
from fastapi import HTTPException
@ -46,7 +47,7 @@ class MCPSecurityGuardrail(CustomGuardrail):
if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_call) is not True:
return data
unregistered: Final = self._find_unregistered_mcp_servers(data)
unregistered: Final = await self._find_unregistered_mcp_servers(data)
if not unregistered:
return data
@ -90,21 +91,22 @@ class MCPSecurityGuardrail(CustomGuardrail):
return server_names
@staticmethod
def _find_unregistered_mcp_servers(data: dict) -> set[str]:
async def _find_unregistered_mcp_servers(data: Mapping[str, object]) -> frozenset[str]:
"""Check tools in data against the MCP server registry. Returns set of unregistered server names."""
tools: Final = data.get("tools")
if not tools or not isinstance(tools, list):
return set()
return frozenset()
requested_servers: Final = MCPSecurityGuardrail._extract_mcp_server_names_from_tools(tools)
if not requested_servers:
return set()
return frozenset()
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
registry: Final = global_mcp_server_manager.get_registry()
registered_names: Final = set(registry.keys())
async with global_mcp_server_manager.catalog.operation():
registry: Final = global_mcp_server_manager.get_registry()
registered_names: Final = set(registry.keys())
return requested_servers - registered_names
return frozenset(requested_servers - registered_names)

View file

@ -19,7 +19,8 @@ import functools
import importlib
import json
import os
from collections.abc import Iterable, Mapping, Sequence
from collections.abc import AsyncGenerator, Iterable, Mapping, Sequence
from contextlib import asynccontextmanager
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import (
@ -47,6 +48,8 @@ from fastapi.responses import JSONResponse
from pydantic import TypeAdapter
from typing_extensions import ReadOnly, TypedDict
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, public_catalog_operation
try:
from prisma.errors import RecordNotFoundError, UniqueViolationError
except ImportError:
@ -1139,6 +1142,7 @@ if MCP_AVAILABLE:
tags=["mcp"],
description="MCP registry endpoint. Spec: https://github.com/modelcontextprotocol/registry",
)
@public_catalog_operation
async def get_mcp_registry(request: Request):
if not _is_public_registry_enabled():
raise HTTPException(
@ -1157,7 +1161,8 @@ if MCP_AVAILABLE:
registry_servers.append({"server": _build_builtin_registry_entry(base_url)})
# Centralized IP-based filtering: external callers only see public servers
registered_servers: Final = list(global_mcp_server_manager.get_filtered_registry(client_ip).values())
async with global_mcp_server_manager.catalog.operation():
registered_servers: Final = list(global_mcp_server_manager.get_filtered_registry(client_ip).values())
registered_servers.sort(key=_build_mcp_registry_server_name)
@ -1187,6 +1192,7 @@ if MCP_AVAILABLE:
return "view_all"
return "restricted"
@catalog_operation(lambda: global_mcp_server_manager)
async def _get_team_scoped_mcp_server_list(
team_id: str,
) -> list[LiteLLM_MCPServerTable]:
@ -1225,6 +1231,7 @@ if MCP_AVAILABLE:
return _redact_mcp_credentials_list(servers)
@catalog_operation(lambda: global_mcp_server_manager)
async def _resolve_accessible_mcp_servers(
user_api_key_dict: UserAPIKeyAuth,
) -> list[LiteLLM_MCPServerTable]:
@ -1246,6 +1253,7 @@ if MCP_AVAILABLE:
aggregated.setdefault(server.server_id, server)
return list(aggregated.values())
@catalog_operation(lambda: global_mcp_server_manager)
async def _connected_app_reachable_server_ids(user_api_key_dict: UserAPIKeyAuth) -> frozenset[str]:
"""Server ids a connected app authorized by this dashboard user is served on the aggregate
MCP endpoint, resolved through the one owner of the admitted subject so the page and the
@ -1262,6 +1270,7 @@ if MCP_AVAILABLE:
dependencies=[Depends(user_api_key_auth)],
response_model=list[LiteLLM_MCPServerTable],
)
@catalog_operation(lambda: global_mcp_server_manager)
async def fetch_all_mcp_servers(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
team_id: str | None = Query(
@ -1378,6 +1387,7 @@ if MCP_AVAILABLE:
description="Health check for MCP servers",
dependencies=[Depends(user_api_key_auth)],
)
@catalog_operation(lambda: global_mcp_server_manager)
async def health_check_servers(
server_ids: list[str] | None = Query(
None,
@ -1784,6 +1794,7 @@ if MCP_AVAILABLE:
dependencies=[Depends(user_api_key_auth)],
response_model=LiteLLM_MCPServerTable,
)
@catalog_operation(lambda: global_mcp_server_manager)
async def fetch_mcp_server(
request: Request,
server_id: str,
@ -2145,6 +2156,7 @@ if MCP_AVAILABLE:
return _redact_mcp_credentials(temp_record)
@public_catalog_operation
async def _mcp_oauth_user_api_key_auth(request: Request) -> UserAPIKeyAuth:
"""
Auth dependency for MCP OAuth browser-navigation endpoints (/authorize, /token).
@ -2202,9 +2214,7 @@ if MCP_AVAILABLE:
server_id: Final[str] = request.path_params.get("server_id", "")
if server_id:
_s = global_mcp_server_manager.get_mcp_server_by_id(server_id)
if not _s:
_s = global_mcp_server_manager.get_mcp_server_by_name(server_id)
_s = await global_mcp_server_manager.catalog.resolve(server_id)
if (
_s
and getattr(_s, "auth_type", None) == MCPAuth.oauth2
@ -2271,6 +2281,18 @@ if MCP_AVAILABLE:
assert authorized.runtime is not None
return authorized.runtime
@asynccontextmanager
async def _oauth_server_operation(
server_id: str,
user_api_key_dict: UserAPIKeyAuth,
request: Request | None = None,
) -> AsyncGenerator[MCPServer]:
if await get_cached_temporary_mcp_server(server_id) is not None:
yield await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request)
return
async with global_mcp_server_manager.catalog.operation():
yield await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request)
@router.get(
"/server/oauth/{server_id}/authorize",
include_in_schema=False,
@ -2288,47 +2310,47 @@ if MCP_AVAILABLE:
response_type: str | None = None,
scope: str | None = None,
):
mcp_server: Final = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request)
_raise_if_not_oauth2(mcp_server)
# Use the server's stored client_id when the caller doesn't supply one
stored_or_supplied_client_id: Final = mcp_server.client_id or client_id or ""
ephemeral_dcr_client: Final = (
await resolve_ephemeral_dcr_client(
async with _oauth_server_operation(server_id, user_api_key_dict, request=request) as mcp_server:
_raise_if_not_oauth2(mcp_server)
# Use the server's stored client_id when the caller doesn't supply one
stored_or_supplied_client_id: Final = mcp_server.client_id or client_id or ""
ephemeral_dcr_client: Final = (
await resolve_ephemeral_dcr_client(
request=request,
mcp_server=mcp_server,
code_challenge=code_challenge,
code_challenge_method=code_challenge_method,
redirect_uri=redirect_uri,
)
if not stored_or_supplied_client_id
else None
)
resolved_client_id: Final = stored_or_supplied_client_id or (
ephemeral_dcr_client.client_id if ephemeral_dcr_client else ""
)
if not resolved_client_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": "missing_client_id",
"message": (
"No client_id available for this MCP server. "
"Either configure the server with a client_id or supply one in the request."
),
},
)
return await authorize_with_server(
request=request,
mcp_server=mcp_server,
client_id=resolved_client_id,
redirect_uri=redirect_uri,
state=state,
code_challenge=code_challenge,
code_challenge_method=code_challenge_method,
redirect_uri=redirect_uri,
response_type=response_type,
scope=scope,
ephemeral_dcr_client=ephemeral_dcr_client,
)
if not stored_or_supplied_client_id
else None
)
resolved_client_id: Final = stored_or_supplied_client_id or (
ephemeral_dcr_client.client_id if ephemeral_dcr_client else ""
)
if not resolved_client_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": "missing_client_id",
"message": (
"No client_id available for this MCP server. "
"Either configure the server with a client_id or supply one in the request."
),
},
)
return await authorize_with_server(
request=request,
mcp_server=mcp_server,
client_id=resolved_client_id,
redirect_uri=redirect_uri,
state=state,
code_challenge=code_challenge,
code_challenge_method=code_challenge_method,
response_type=response_type,
scope=scope,
ephemeral_dcr_client=ephemeral_dcr_client,
)
@router.post(
"/server/oauth/{server_id}/token",
@ -2348,47 +2370,47 @@ if MCP_AVAILABLE:
refresh_token: str | None = Form(None),
scope: str | None = Form(None),
):
mcp_server: Final = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request)
_raise_if_not_oauth2(mcp_server)
# Sealed passthrough codes exist only for the authorization_code grant. A refresh_token
# grant must never open one: the minted client is unrecoverable after the single flow by
# contract, so an expired browser-held token re-runs authorize instead.
sealed_code: Final = (
redeem_passthrough_authorization_code(code=code, mcp_server=mcp_server, code_verifier=code_verifier)
if grant_type == "authorization_code"
else None
)
resolved_code: Final = sealed_code.upstream_code if sealed_code else code
# A sealed flow ran the gateway /callback as its upstream redirect (bridge short-circuit
# or plain flow alike), so the exchange must present that binding, not the browser page.
resolved_redirect_uri: Final = f"{get_request_base_url(request)}/callback" if sealed_code else redirect_uri
caller_client_id: Final = sealed_code.client_id if sealed_code else client_id
caller_client_secret: Final = sealed_code.client_secret if sealed_code else client_secret
resolved_client_id: Final = mcp_server.client_id or caller_client_id or ""
if not resolved_client_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": "missing_client_id",
"message": (
"No client_id available for this MCP server. "
"Either configure the server with a client_id or supply one in the request."
),
},
async with _oauth_server_operation(server_id, user_api_key_dict, request=request) as mcp_server:
_raise_if_not_oauth2(mcp_server)
# Sealed passthrough codes exist only for the authorization_code grant. A refresh_token
# grant must never open one: the minted client is unrecoverable after the single flow by
# contract, so an expired browser-held token re-runs authorize instead.
sealed_code: Final = (
redeem_passthrough_authorization_code(code=code, mcp_server=mcp_server, code_verifier=code_verifier)
if grant_type == "authorization_code"
else None
)
resolved_code: Final = sealed_code.upstream_code if sealed_code else code
# A sealed flow ran the gateway /callback as its upstream redirect (bridge short-circuit
# or plain flow alike), so the exchange must present that binding, not the browser page.
resolved_redirect_uri: Final = f"{get_request_base_url(request)}/callback" if sealed_code else redirect_uri
caller_client_id: Final = sealed_code.client_id if sealed_code else client_id
caller_client_secret: Final = sealed_code.client_secret if sealed_code else client_secret
resolved_client_id: Final = mcp_server.client_id or caller_client_id or ""
if not resolved_client_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={
"error": "missing_client_id",
"message": (
"No client_id available for this MCP server. "
"Either configure the server with a client_id or supply one in the request."
),
},
)
return await exchange_token_with_server(
request=request,
mcp_server=mcp_server,
grant_type=grant_type,
code=resolved_code,
redirect_uri=resolved_redirect_uri,
client_id=resolved_client_id,
client_secret=caller_client_secret,
code_verifier=code_verifier,
refresh_token=refresh_token,
scope=scope,
client_token_endpoint_auth_method=sealed_code.token_endpoint_auth_method if sealed_code else None,
)
return await exchange_token_with_server(
request=request,
mcp_server=mcp_server,
grant_type=grant_type,
code=resolved_code,
redirect_uri=resolved_redirect_uri,
client_id=resolved_client_id,
client_secret=caller_client_secret,
code_verifier=code_verifier,
refresh_token=refresh_token,
scope=scope,
client_token_endpoint_auth_method=sealed_code.token_endpoint_auth_method if sealed_code else None,
)
@router.post(
"/server/oauth/{server_id}/register",
@ -2400,22 +2422,22 @@ if MCP_AVAILABLE:
server_id: str,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
mcp_server: Final = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request)
request_data: Final = await _read_request_body(request=request)
data: Final[dict] = {**request_data}
client_redirect_uris: Final = client_supplied_redirect_uris(data.get("redirect_uris"))
async with _oauth_server_operation(server_id, user_api_key_dict, request=request) as mcp_server:
request_data: Final = await _read_request_body(request=request)
data: Final[Mapping[str, object]] = {**request_data}
client_redirect_uris: Final = client_supplied_redirect_uris(data.get("redirect_uris"))
return await register_client_with_server(
request=request,
mcp_server=mcp_server,
client_name=data.get("client_name", ""),
grant_types=data.get("grant_types", []),
response_types=data.get("response_types", []),
token_endpoint_auth_method=data.get("token_endpoint_auth_method", ""),
fallback_client_id=server_id,
persist_credentials=_user_is_full_admin(user_api_key_dict),
client_redirect_uris=client_redirect_uris,
)
return await register_client_with_server(
request=request,
mcp_server=mcp_server,
client_name=data.get("client_name", ""),
grant_types=data.get("grant_types", []),
response_types=data.get("response_types", []),
token_endpoint_auth_method=data.get("token_endpoint_auth_method", ""),
fallback_client_id=server_id,
persist_credentials=_user_is_full_admin(user_api_key_dict),
client_redirect_uris=client_redirect_uris,
)
@router.delete(
"/server/{server_id}",
@ -2791,6 +2813,7 @@ if MCP_AVAILABLE:
# ── Per-user MCP env var endpoints ────────────────────────────────────────
@catalog_operation(lambda: global_mcp_server_manager)
async def _authorize_and_fetch_mcp_server(
prisma_client,
user_api_key_dict: UserAPIKeyAuth,

View file

@ -336,6 +336,7 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.openai_like.model_info import MODEL_INFO_REFRESH_SECONDS
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._experimental.mcp_server.byok_credential_cache import byok_credential_cache
from litellm.proxy._experimental.mcp_server.catalog import public_catalog_operation
from litellm.proxy._experimental.mcp_server.stdio_gate import MCP_STDIO_ENABLED_ENV_VAR, is_mcp_stdio_flag_key
from litellm.proxy._lazy_features import attach_lazy_features, reserve_lazy_slot
from litellm.proxy._types import *
@ -20368,30 +20369,26 @@ async def _resolve_mcp_csv_tokens(csv_segment: str, client_ip: str | None) -> li
all-unmatched server filter falls back to the full ``allowed_mcp_servers``
list and silently broadens the request scope).
"""
from litellm.constants import DEFAULT_MCP_NAMESPACE_CSV_MAX_TOKENS
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._experimental.mcp_server.catalog import global_manager
seen: Final[set] = set()
deduped: Final[list[str]] = []
for raw in csv_segment.split(","):
token = raw.strip()
if not token or token in seen:
continue
seen.add(token)
deduped.append(token)
if len(deduped) >= DEFAULT_MCP_NAMESPACE_CSV_MAX_TOKENS:
break
async with global_manager().catalog.operation():
from litellm.constants import DEFAULT_MCP_NAMESPACE_CSV_MAX_TOKENS
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
resolved: Final[list[str]] = []
for token in deduped:
if global_mcp_server_manager.get_mcp_server_by_name(token, client_ip=client_ip):
resolved.append(token)
continue
if await _is_mcp_access_group_cached(token):
resolved.append(token)
return resolved
deduped: Final = tuple(
token for token in dict.fromkeys(raw.strip() for raw in csv_segment.split(",")) if token
)[:DEFAULT_MCP_NAMESPACE_CSV_MAX_TOKENS]
resolved: Final[list[str]] = [] # mutable-ok: sequential await per token cannot live in a comprehension
for token in deduped:
if global_mcp_server_manager.get_mcp_server_by_name(token, client_ip=client_ip):
resolved.append(token)
continue
if await _is_mcp_access_group_cached(token):
resolved.append(token)
return resolved
async def _is_mcp_access_group_cached(name: str) -> bool:
@ -20428,12 +20425,13 @@ async def _is_mcp_access_group_cached(name: str) -> bool:
"/{mcp_server_name}/mcp",
methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"],
)
@public_catalog_operation
async def dynamic_mcp_route(mcp_server_name: str, request: Request):
"""Handle /{name}/mcp for MCP server aliases, toolsets, MCP access group tags, and comma-separated lists.
Resolution order:
1. Registered MCP server alias / name
2. Comma-separated list (short-circuits before any DB call)
2. Comma-separated list
3. Toolset name (DB lookup, cached)
4. MCP access group tag (DB lookup, cached)
"""
@ -20446,7 +20444,9 @@ async def dynamic_mcp_route(mcp_server_name: str, request: Request):
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
# 1. Registered MCP server alias
if global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name, client_ip=client_ip):
async with global_mcp_server_manager.catalog.operation():
server: Final = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name, client_ip=client_ip)
if server is not None:
return await _mcp_forward_as_path(mcp_server_name, request)
# 2. Comma-separated list — validate every token resolves to a known

View file

@ -22,6 +22,9 @@ class _ConfigRow(Protocol):
@property
def param_value(self) -> object: ...
@property
def reload_revision(self) -> int | None: ...
class _ConfigTable(Protocol):
async def find_unique(self, *, where: Mapping[str, str]) -> _ConfigRow | None: ...

View file

@ -10,6 +10,7 @@ from openai.types.responses.function_tool_param import FunctionToolParam
from litellm._logging import verbose_logger
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, global_manager
from litellm.proxy._experimental.mcp_server.utils import (
iter_known_server_prefixes,
logging_safe_mcp_headers,
@ -112,6 +113,7 @@ async def _toolset_exists(name: str) -> bool:
return False
@catalog_operation(global_manager)
async def _gateway_served_names(
names: Collection[str],
servers: Callable[[], Collection[MCPServer]] = _registered_mcp_servers,
@ -185,6 +187,7 @@ class LiteLLM_Proxy_MCP_Handler:
_split_mcp_tools = split_mcp_tools
@staticmethod
@catalog_operation(global_manager)
async def routes_through_gateway(
tools: Iterable[Mapping[str, object]] | None,
served_names: Callable[[Collection[str]], Awaitable[frozenset[str]]] = _gateway_served_names,
@ -238,6 +241,7 @@ class LiteLLM_Proxy_MCP_Handler:
return user_api_key_auth
@staticmethod
@catalog_operation(global_manager)
async def _get_mcp_tools_from_manager(
user_api_key_auth: "UserAPIKeyAuth | None",
mcp_tools_with_litellm_proxy: Iterable[Mapping[str, object]] | None,
@ -704,6 +708,7 @@ class LiteLLM_Proxy_MCP_Handler:
return result_text or "Tool executed successfully"
@staticmethod
@catalog_operation(global_manager)
async def execute_tool_calls(
tool_server_map: dict[str, str],
tool_calls: Sequence[object],

View file

@ -238,7 +238,7 @@ class MCPServer(LiteLLMBaseModel):
# None or a value <= 0 means unlimited.
max_concurrent_requests: int | None = None
# Resolved short-ID tool prefix when LITELLM_USE_SHORT_MCP_TOOL_PREFIX is
# enabled. Set by ``MCPServerManager._assign_unique_short_prefix`` at
# enabled. Set by ``MCPServerManager.assign_unique_short_prefix`` at
# registration time so that natural-hash collisions between two
# different ``server_id`` values are bumped deterministically. Left
# ``None`` in default-prefix mode.

View file

@ -1,7 +1,7 @@
from collections.abc import Callable
from typing import Any, ClassVar
from pydantic import ConfigDict
from pydantic import ConfigDict, Field
from litellm.types.llms.base import LiteLLMBaseModel
@ -12,6 +12,7 @@ class MCPTool(LiteLLMBaseModel):
description: str
input_schema: dict[str, Any]
handler: Callable
server_id: str | None = Field(default=None, frozen=True)
class ToolSchema(LiteLLMBaseModel):

View file

@ -283,38 +283,40 @@ def test_access_group_membership_follows_edits(gateway: Gateway) -> None:
assert tool_calls(peer.drain()) == ()
def test_peer_worker_observes_create_edit_and_delete_without_restart(gateway: Gateway, peer: Gateway) -> None:
with mcp_peer() as first, mcp_peer() as second, gateway.scenario() as scenario:
def test_peer_worker_observes_create_edit_and_delete_without_restart(gateway: Gateway, tmp_path: Path) -> None:
config: Final = tmp_path / "peer.yaml"
config.write_text(yaml.safe_dump({
"model_list": [],
"general_settings": {"master_key": gateway.key, "store_model_in_db": True, "proxy_config_reload_interval_seconds": 3600},
}))
with (
owned_proxy(gateway, tmp_path, {}, config=config, database_setup=()) as peer,
mcp_peer() as first,
mcp_peer() as second,
gateway.scenario() as scenario,
):
alias: Final = "mgmt" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, first, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
eventually(
lambda: peer.client.get(
"/mcp-rest/tools/list", headers={"x-litellm-api-key": key}, params={"server_id": identity}
),
lambda value: value.status_code == 200 and value.json() != [],
seconds=40,
listing: Final = peer.client.get(
"/mcp-rest/tools/list", headers={"x-litellm-api-key": key}, params={"server_id": identity}
)
assert listing.status_code == 200 and listing.json() != [], listing.text
names: Final = tool_names(peer, key, identity)
assert call_tool(peer, key, identity, names["add"], ADD).status_code == 200
assert len(tool_calls(first.drain())) == 1
moved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "url": second.url})
assert moved.status_code == 202, moved.text
eventually(
lambda: call_tool(peer, key, identity, names["add"], ADD),
lambda value: value.status_code == 200 and len(tool_calls(second.drain())) == 1,
seconds=40,
)
called: Final = call_tool(peer, key, identity, names["add"], ADD)
assert called.status_code == 200, called.text
assert tool_calls(first.drain()) == ()
assert len(tool_calls(second.drain())) == 1
delete_mcp(gateway, identity)
eventually(
lambda: peer.client.get(
"/mcp-rest/tools/list", headers={"x-litellm-api-key": key}, params={"server_id": identity}
),
lambda value: value.status_code >= 400 or value.json() == [],
seconds=40,
deleted: Final = peer.client.get(
"/mcp-rest/tools/list", headers={"x-litellm-api-key": key}, params={"server_id": identity}
)
second.drain()
assert call_tool(peer, key, identity, names["add"], ADD).status_code >= 400
assert deleted.status_code in (403, 404) or (deleted.status_code == 200 and deleted.json() == []), deleted.text
assert call_tool(peer, key, identity, names["add"], ADD).status_code in (403, 404)
assert tool_calls(second.drain()) == ()

View file

@ -235,9 +235,11 @@ def _clear_proxy_database_env() -> typing.Iterator[None]:
async def _initialize_proxy(config_path: str) -> None:
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
cleanup_router_config_variables()
global_mcp_server_manager.catalog = TargetCatalog(global_mcp_server_manager)
await initialize(config=config_path, debug=True)
for server_id, upstream in tuple(global_mcp_server_manager.registry.items()):
if upstream.server_name != "math_restricted":

View file

@ -7990,9 +7990,13 @@ class TestGatewaySessionAdmission:
rpm_limit=rpm_limit,
)
)
prisma = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
patch("litellm.proxy.auth.auth_checks.get_user_object", get_user_object),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.should_load_db_object", return_value=False),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
):
yield get_user_object
@ -10143,10 +10147,14 @@ class TestScopedSessionAdmission:
rpm_limit=None,
)
)
prisma = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY),
patch("litellm.proxy.auth.auth_checks.get_user_object", get_user_object),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.should_load_db_object", return_value=False),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
):
auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope_dict)
@ -10200,3 +10208,21 @@ async def test_managed_agent_permission_resolution_outage_is_not_an_unrestricted
)
with pytest.raises(RuntimeError, match="policy unavailable"):
await resolution
@pytest.mark.asyncio
async def test_unreadable_empty_key_scope_cannot_gain_additive_grants(monkeypatch):
auth = UserAPIKeyAuth(api_key="test-key", object_permission_id="key-scope")
monkeypatch.setattr(
MCPRequestHandler,
"_get_allowed_mcp_servers_for_key",
AsyncMock(return_value=[SpecialMCPServerNames.no_mcp_servers.value]),
)
monkeypatch.setattr(MCPRequestHandler, "_key_object_permission_hydrated", AsyncMock(return_value=None))
monkeypatch.setattr(MCPRequestHandler, "_get_allowed_mcp_servers_for_team", AsyncMock(return_value=[]))
monkeypatch.setattr(
MCPRequestHandler, "_get_key_access_group_mcp_server_extras", AsyncMock(return_value=["unrelated-server"])
)
access = await MCPRequestHandler.get_mcp_server_access(auth)
assert access.server_ids == ()
assert access.scope == "scoped"

View file

@ -81,10 +81,18 @@ def config_only_mcp_manager_factory():
@pytest.fixture(autouse=True)
def _hermetic_mcp_server_registry():
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
saved_tools = global_mcp_tool_registry.published_tools
global_mcp_tool_registry.published_tools = {}
saved_catalog = global_mcp_server_manager.catalog
global_mcp_server_manager.catalog = TargetCatalog(global_mcp_server_manager)
saved_registry = dict(global_mcp_server_manager.registry)
saved_config_servers = dict(global_mcp_server_manager.config_mcp_servers)
saved_tool_mapping = dict(global_mcp_server_manager.tool_name_to_mcp_server_name_mapping)
@ -103,6 +111,8 @@ def _hermetic_mcp_server_registry():
global_mcp_server_manager.tool_name_to_mcp_server_name_mapping.clear()
global_mcp_server_manager.tool_name_to_mcp_server_name_mapping.update(saved_tool_mapping)
global_mcp_server_manager._oauth_discovery_slots = saved_oauth_slots
global_mcp_server_manager.catalog = saved_catalog
global_mcp_tool_registry.published_tools = saved_tools
@pytest.fixture(autouse=True)

View file

@ -655,7 +655,10 @@ async def test_execute_byok_tool_missing_credential_advertises_api_key_flow(monk
monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com/proxy")
mcp_operations.byok_credential_cache.flush_cache()
server = MCPServer(server_id="byok-discovery", name="byok-discovery", transport=MCPTransport.http, is_byok=True)
monkeypatch.setattr(proxy_server, "should_load_db_object", lambda _kind: False)
prisma = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=None)
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
with pytest.raises(HTTPException) as exc_info:

View file

@ -507,7 +507,7 @@ class _MapTable:
@pytest.mark.parametrize("quoted", [False, True])
async def test_secret_maps_create_update_round_trip(map_algorithm: str, field: str, quoted: bool) -> None:
table: Final = _MapTable(quoted=quoted)
prisma: Final = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table))
prisma: Final = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table, litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None))))
original: Final = {"TOKEN": " sensitive-secret\n", "PREFIX": "v2:gcm:literal", "TEMPLATE": "Bearer ${TOKEN}"}
create: Final = NewMCPServerRequest.model_validate({
"server_id": "srv-map", "transport": "http", "url": "https://up.example.com/mcp", field: original,
@ -580,7 +580,7 @@ async def test_secret_map_rotation_migrates_rekeys_and_preserves_corrupt(
{"server_id": "encrypted", field: old, other: None},
)
prisma: Final = SimpleNamespace(db=SimpleNamespace(
litellm_mcpservertable=table, litellm_mcpserveroauthclient=SimpleNamespace(find_many=AsyncMock(return_value=[]))
litellm_mcpservertable=table, litellm_mcpserveroauthclient=SimpleNamespace(find_many=AsyncMock(return_value=[])), litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None))
))
await rotate_mcp_server_credentials_master_key(prisma, touched_by="test", new_master_key="rotated-map-key")
assert table.rows["broken"][field] == corrupt
@ -608,7 +608,7 @@ async def test_bulk_reads_isolate_corrupt_secret_maps(reader, field, map_algorit
]
snapshot = [row.model_dump() for row in rows]
table = SimpleNamespace(find_many=AsyncMock(return_value=rows))
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table))
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table, litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None))))
result = await reader(prisma, ["broken", "healthy"]) if reader is get_mcp_servers else await reader(prisma)
items = result.items if reader is get_mcp_submissions else result
assert [row.server_id for row in items] == ["healthy"]
@ -627,7 +627,7 @@ async def test_bulk_reads_do_not_swallow_unrelated_validation_errors(reader):
row = _prisma_map_row({"server_id": "invalid", "transport": "unsupported"})
table = SimpleNamespace(find_many=AsyncMock(return_value=[row]))
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table))
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table, litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None))))
request = reader(prisma, ["invalid"]) if reader is get_mcp_servers else reader(prisma)
with pytest.raises(ValidationError, match="transport"):
await request
@ -1545,6 +1545,7 @@ async def test_master_key_rotation_reencrypts_oauth_client_store(monkeypatch):
prisma = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(
return_value=[SimpleNamespace(server_id="config_faros", credentials=blob_old)]
)

View file

@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from fastapi import HTTPException, Request
from litellm.proxy._types import LiteLLM_MCPServerTable
from litellm.types.mcp import MCPAuth
@ -92,7 +92,7 @@ def mock_mcp_client_ip():
@pytest.fixture(autouse=True)
def isolate_global_mcp_registry():
def isolate_global_mcp_registry(monkeypatch):
"""Restore the module-global MCP server registry after each test.
Tests here register servers on ``global_mcp_server_manager`` directly; without a
@ -101,6 +101,8 @@ def isolate_global_mcp_registry():
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
monkeypatch.setattr(global_mcp_server_manager, "catalog", TargetCatalog(global_mcp_server_manager))
snapshot = dict(global_mcp_server_manager.registry)
yield
global_mcp_server_manager.registry.clear()
@ -115,7 +117,8 @@ def _mock_callback_request(base_url: str = "http://localhost:3000/"):
and trusted ``X-Forwarded-*`` headers). A simple MagicMock with the
right attributes is sufficient.
"""
req = MagicMock()
req = MagicMock(spec=Request)
req.client = None
req.base_url = base_url
req.headers = {}
req.cookies = {}
@ -739,7 +742,8 @@ async def test_token_endpoint_forwards_code_verifier():
@pytest.mark.asyncio
async def test_register_client_without_mcp_server_name_returns_dummy():
@pytest.mark.parametrize("server_name", [None, "missing"])
async def test_register_client_without_mcp_server_name_returns_dummy(server_name):
try:
from fastapi import Request
@ -762,17 +766,18 @@ async def test_register_client_without_mcp_server_name_returns_dummy():
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body",
new=AsyncMock(return_value={}),
):
result = await register_client(request=mock_request)
result = await register_client(request=mock_request, mcp_server_name=server_name)
assert result == {
"client_id": "dummy_client",
"client_id": server_name or "dummy_client",
"client_secret": "dummy",
"redirect_uris": ["https://proxy.litellm.example/callback"],
}
@pytest.mark.asyncio
async def test_register_client_returns_existing_server_credentials():
@pytest.mark.parametrize("use_root", [False, True])
async def test_register_client_returns_existing_server_credentials(use_root):
try:
from fastapi import Request
@ -812,7 +817,9 @@ async def test_register_client_returns_existing_server_credentials():
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body",
new=AsyncMock(return_value={}),
):
result = await register_client(request=mock_request, mcp_server_name=oauth2_server.server_name)
result = await register_client(
request=mock_request, mcp_server_name=None if use_root else oauth2_server.server_name
)
finally:
global_mcp_server_manager.registry.clear()
@ -4885,7 +4892,7 @@ async def test_token_endpoint_authorization_code_missing_code():
)
global_mcp_server_manager.registry[server.server_id] = server
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://proxy.example/"
mock_request.headers = {}
@ -8245,7 +8252,7 @@ async def test_authorize_endpoint_rejects_non_oauth2_server():
server = _access_group_none_server()
global_mcp_server_manager.registry[server.server_id] = server
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -8284,7 +8291,7 @@ async def test_token_endpoint_rejects_non_oauth2_server():
server = _access_group_none_server()
global_mcp_server_manager.registry[server.server_id] = server
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -8326,7 +8333,7 @@ async def test_register_client_rejects_non_oauth2_server():
server = _access_group_none_server()
global_mcp_server_manager.registry[server.server_id] = server
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -8399,7 +8406,7 @@ async def test_oauth_authorization_server_404_for_non_oauth2_server():
server = _access_group_none_server()
global_mcp_server_manager.registry[server.server_id] = server
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -8447,7 +8454,7 @@ async def test_oauth_protected_resource_passthrough_none_auth_not_404():
)
global_mcp_server_manager.registry[passthrough_server.server_id] = passthrough_server
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -8481,7 +8488,7 @@ async def test_oauth_protected_resource_404_for_unknown_server_name():
pytest.skip("MCP discoverable endpoints not available")
global_mcp_server_manager.registry.clear()
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -8509,7 +8516,7 @@ async def test_oauth_authorization_server_404_for_unknown_server_name():
pytest.skip("MCP discoverable endpoints not available")
global_mcp_server_manager.registry.clear()
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -9518,7 +9525,7 @@ async def test_load_servers_from_config_hydrates_dcr_clients():
)
hydrate_spy = AsyncMock()
with patch.object(global_mcp_server_manager, "_hydrate_config_servers_dcr_clients", new=hydrate_spy):
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()
@ -9535,7 +9542,9 @@ async def test_reload_servers_from_database_hydrates_dcr_clients():
)
prisma = MagicMock()
prisma.writer_db = prisma.db
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
hydrate_spy = AsyncMock()
with (
@ -9543,7 +9552,7 @@ async def test_reload_servers_from_database_hydrates_dcr_clients():
"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),
patch.object(global_mcp_server_manager, "hydrate_config_servers_dcr_clients", new=hydrate_spy),
):
await global_mcp_server_manager.reload_servers_from_database()
@ -9814,7 +9823,7 @@ async def test_authorize_wall_names_the_fix_for_urlless_servers():
auth_type=MCPAuth.oauth2,
spec_path="https://example.com/openapi.yaml",
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -9852,7 +9861,7 @@ async def test_token_wall_names_the_fix_for_urlless_servers():
spec_path="https://example.com/openapi.yaml",
authorization_url="https://accounts.google.com/o/oauth2/v2/auth",
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -9892,7 +9901,7 @@ async def test_register_wall_names_the_fix_for_urlless_servers():
auth_type=MCPAuth.oauth2,
spec_path="https://example.com/openapi.yaml",
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -9931,7 +9940,7 @@ async def test_authorize_wall_points_at_discovery_failure_for_url_servers():
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -9967,7 +9976,7 @@ async def test_token_wall_points_at_discovery_failure_for_url_servers():
auth_type=MCPAuth.oauth2,
authorization_url="https://idp.example.com/authorize",
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -10007,7 +10016,7 @@ async def test_authorize_wall_names_the_issuer_for_anchored_servers():
issuer="https://idp.example.com",
issuer_is_anchored=True,
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -10054,7 +10063,7 @@ async def test_authorize_uses_admin_entered_github_oauth_urls_after_issuer_yield
configured_authorization_url="https://github.com/login/oauth/authorize",
configured_token_url="https://github.com/login/oauth/access_token",
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -10076,7 +10085,7 @@ def test_oauth_endpoints_count_admin_entered_urls_as_resolved():
"""A leftover issuer empties the resolved authorize/token fields but must not keep the
server on the deferred-discovery retry path when the admin already stored those URLs."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
_oauth_endpoints_unresolved,
oauth_endpoints_unresolved,
)
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@ -10093,7 +10102,7 @@ def test_oauth_endpoints_count_admin_entered_urls_as_resolved():
configured_authorization_url="https://github.com/login/oauth/authorize",
configured_token_url="https://github.com/login/oauth/access_token",
)
assert _oauth_endpoints_unresolved(server) is False
assert oauth_endpoints_unresolved(server) is False
@pytest.mark.asyncio
@ -10146,7 +10155,7 @@ async def test_token_exchange_with_configured_token_url_never_joins_discovery(mo
"get_async_httpx_client",
lambda llm_provider: fake_http_client,
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -10271,7 +10280,7 @@ async def test_bridge_authorize_relays_with_registration_url_resolved_by_deferre
"ensure_oauth_metadata_discovered",
resolve_discovery,
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
@ -11593,7 +11602,7 @@ def test_discovery_advertises_the_exchange_grant_only_where_the_gateway_can_serv
litellm_jwtauth=LiteLLM_JWTAuth(virtual_key_claim_field=virtual_key_claim_field),
)
monkeypatch.setattr("litellm.proxy.proxy_server.jwt_handler", handler)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {"enable_jwt_auth": jwt_auth_enabled})
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {"enable_jwt_auth": jwt_auth_enabled, "supported_db_objects": []})
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", object())
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
exchange_grant = ["urn:ietf:params:oauth:grant-type:token-exchange"] if exchange_servable else []
@ -12807,6 +12816,65 @@ async def test_identity_bound_authorize_unrelated_bearer_uses_browser_session(
proxy_server.prisma_client.db.litellm_teamtable.create.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize("change", ["create", "update", "delete"])
async def test_authorize_observes_committed_peer_server_changes(monkeypatch, change):
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from fastapi import HTTPException, Request
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
from litellm.proxy._types import LiteLLM_MCPServerTable
monkeypatch.setenv("LITELLM_SALT_KEY", "catalog-consistency-test-key")
stamp = datetime.now(timezone.utc)
old_server = _create_id_lookup_oauth2_server()
old_server.updated_at = stamp
row = LiteLLM_MCPServerTable(
server_id=old_server.server_id,
server_name=old_server.server_name,
alias=old_server.alias,
url="https://upstream.example/mcp",
transport="http",
auth_type="oauth2",
authorization_url="https://new-provider.example/authorize",
token_url="https://new-provider.example/token",
scopes=["read"],
credentials={"client_id": "current-client", "client_secret": "current-secret"},
created_at=stamp,
updated_at=stamp + timedelta(seconds=1),
)
read_rows = AsyncMock(return_value=[] if change == "delete" else [row])
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=SimpleNamespace(find_many=read_rows), litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None))))
prisma.writer_db = prisma.db
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
monkeypatch.setattr(
global_mcp_server_manager, "registry", {} if change == "create" else {old_server.server_id: old_server}
)
monkeypatch.setattr(global_mcp_server_manager, "config_mcp_servers", {})
request = Request(
{"type": "http", "scheme": "https", "server": ("gateway.example", 443), "path": "/authorize", "headers": []}
)
if change == "delete":
with pytest.raises(HTTPException) as exc:
await discoverable_endpoints.authorize(
request, "http://localhost/callback", mcp_server_name=old_server.server_id
)
assert exc.value.status_code == 404
else:
response = await discoverable_endpoints.authorize(
request, "http://localhost/callback", mcp_server_name=old_server.server_id
)
assert response.status_code == 307
assert response.headers["location"].startswith("https://new-provider.example/authorize?")
assert "client_id=current-client" in response.headers["location"]
read_rows.assert_awaited_once()
@pytest.mark.asyncio
async def test_update_server_drops_cached_upstream_oauth_metadata():
from litellm.proxy._experimental.mcp_server import discoverable_endpoints

View file

@ -66,8 +66,7 @@ async def _call_block(logging_obj, order: list, *, user_api_key_auth=mock.sentin
proxy_logging_obj = mock.MagicMock()
proxy_logging_obj.post_call_failure_hook.side_effect = _record_post_call_failure_hook
fake_proxy_server = types.ModuleType("litellm.proxy.proxy_server")
fake_proxy_server.proxy_logging_obj = proxy_logging_obj # pyright: ignore[reportAttributeAccessIssue]
fake_proxy_server = types.SimpleNamespace(proxy_logging_obj=proxy_logging_obj, prisma_client=None)
with mock.patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}):
with contextlib.suppress(HTTPException):

View file

@ -2,7 +2,7 @@
import os
import pytest
from unittest.mock import ANY, AsyncMock, MagicMock, patch
from contextlib import asynccontextmanager
from contextlib import asynccontextmanager, nullcontext
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
@ -918,6 +918,7 @@ async def test_get_tools_from_mcp_servers():
# Create a mock manager
mock_manager = AsyncMock()
mock_manager.catalog.operation = nullcontext
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["server1_id", "server2_id"]
)
@ -946,6 +947,7 @@ async def test_get_tools_from_mcp_servers():
# Test Case 2: Without specific MCP servers
# Create a different mock manager for the second test case
mock_manager_2 = AsyncMock()
mock_manager_2.catalog.operation = nullcontext
mock_manager_2.get_allowed_mcp_servers = AsyncMock(
return_value=["server1_id", "server2_id"]
)
@ -996,6 +998,7 @@ async def test_get_tools_from_mcp_servers():
# Test Case 3: With specific MCP servers and access groups
# Create a mock manager
mock_manager = AsyncMock()
mock_manager.catalog.operation = nullcontext
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["server1_id", "server2_id", "server3_id"]
)

View file

@ -5405,7 +5405,9 @@ class TestMCPServerManagerReload:
db_row = _make_db_mcp_server("server-1", timestamp)
mock_prisma = MagicMock()
mock_prisma.writer_db = mock_prisma.db
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[db_row])
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
@ -5447,7 +5449,9 @@ class TestMCPServerManagerReload:
)
mock_prisma = MagicMock()
mock_prisma.writer_db = mock_prisma.db
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[db_row])
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
@ -5461,7 +5465,7 @@ class TestMCPServerManagerReload:
):
await manager.reload_servers_from_database()
mock_build.assert_awaited_once_with(db_row, env_vars_are_encrypted=True)
mock_build.assert_awaited_once_with(db_row, env_vars_are_encrypted=True, register_oauth_discovery=False)
assert manager.registry["server-1"] is rebuilt_server
@pytest.mark.asyncio
@ -5500,9 +5504,11 @@ class TestMCPServerManagerReload:
return another_healthy_server
mock_prisma = MagicMock()
mock_prisma.writer_db = mock_prisma.db
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(
return_value=[healthy_row, bad_row, another_healthy_row]
)
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
@ -5513,7 +5519,7 @@ class TestMCPServerManagerReload:
"build_mcp_server_from_table",
AsyncMock(side_effect=build_server),
),
patch.object(manager, "_maybe_register_openapi_tools", AsyncMock()),
patch.object(manager, "maybe_register_openapi_tools", AsyncMock()),
caplog.at_level("ERROR", logger="LiteLLM"),
):
await manager.reload_servers_from_database()
@ -5572,7 +5578,9 @@ class TestMCPServerManagerReload:
raise RuntimeError("blocked address")
mock_prisma = MagicMock()
mock_prisma.writer_db = mock_prisma.db
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[healthy_row, bad_openapi_row])
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
@ -5585,7 +5593,7 @@ class TestMCPServerManagerReload:
),
patch.object(
manager,
"_maybe_register_openapi_tools",
"maybe_register_openapi_tools",
AsyncMock(side_effect=register_openapi_tools),
),
caplog.at_level("ERROR", logger="LiteLLM"),
@ -7754,8 +7762,8 @@ async def test_execute_mcp_tool_rest_server_id_injects_requested_server_credenti
),
patch.object(
mcp_operations.global_mcp_server_manager,
"get_registry",
return_value={
"registry",
{
requested_server.server_id: requested_server,
collision_server.server_id: collision_server,
},
@ -8019,6 +8027,7 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details():
fake_tool.name = "list_pets"
fake_tool.description = "test tool"
fake_tool.input_schema = {"type": "object"}
fake_tool.server_id = fake_server.server_id
start_time = datetime.now(timezone.utc)
litellm_logging_obj, _ = function_setup(
@ -9000,6 +9009,7 @@ async def test_stateful_mcp_tool_call_uses_current_requests_otel_destinations(_m
server = MCPServer(
server_id="otel-context-test",
name="otelcontext",
server_name="otelcontext",
transport=MCPTransport.http,
allow_all_keys=True,
)
@ -9091,6 +9101,7 @@ async def test_get_active_submitted_mcp_server_ids_for_user_queries_active_rows(
row.server_id = "submitted-1"
prisma_client = MagicMock()
prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
prisma_client.db.litellm_config.find_unique = AsyncMock(return_value=None)
result = await get_active_submitted_mcp_server_ids_for_user(prisma_client, "submitter-user")
@ -9111,6 +9122,7 @@ async def test_get_active_submitted_mcp_server_ids_for_user_empty_user_id_skips_
prisma_client = MagicMock()
prisma_client.db.litellm_mcpservertable.find_many = AsyncMock()
prisma_client.db.litellm_config.find_unique = AsyncMock(return_value=None)
assert await get_active_submitted_mcp_server_ids_for_user(prisma_client, "") == []
prisma_client.db.litellm_mcpservertable.find_many.assert_not_awaited()
@ -10026,7 +10038,7 @@ class TestPreemptive401ModeAware:
assert resolved.authorization_url == "https://idp.example.com/authorize"
assert resolved.token_url == "https://idp.example.com/token"
assert resolved.registration_url == "https://idp.example.com/register"
assert manager._oauth_discovery_slot(server.server_id) is None
assert manager.oauth_discovery_slot(server.server_id) is None
assert exc.value.status_code == 401
@pytest.mark.asyncio

View file

@ -995,6 +995,7 @@ class TestRotateCredentials:
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[server])
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
mock_prisma.db.litellm_mcpservertable.update = AsyncMock()
mock_prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(return_value=[])
@ -1043,6 +1044,7 @@ class TestRotateCredentials:
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[server])
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
mock_prisma.db.litellm_mcpservertable.update = AsyncMock()
mock_prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(return_value=[])

View file

@ -134,6 +134,7 @@ def _byok_key_row(server_id):
def _mock_prisma(null_rows, token_rows):
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=null_rows)
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
mock_prisma.db.litellm_mcpservertable.update_many = AsyncMock(return_value=MagicMock())
mock_prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=token_rows)
return mock_prisma

View file

@ -34,6 +34,7 @@ def _row(**overrides):
def _prisma(rows):
prisma_client = MagicMock()
prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=rows)
prisma_client.db.litellm_config.find_unique = AsyncMock(return_value=None)
prisma_client.db.litellm_mcpservertable.update = AsyncMock()
return prisma_client

View file

@ -622,6 +622,7 @@ async def test_local_tool_json_array_is_converted_once_for_the_caller_revision(c
body = '["a","b"]'
tool = MagicMock()
tool.server_id = None
tool.handler = AsyncMock(return_value=parse_http_body(body))
with patch.object(global_mcp_tool_registry, "get_tool", return_value=tool):
result = await operations._handle_local_mcp_tool("reports-list_tags", {}, WireCompat(compat))
@ -778,3 +779,172 @@ async def test_list_mcp_tools_records_the_catalog_only_when_asked(
manager._drop_listed_tools(server.server_id)
assert [tool.name for tool in listing.tools] == ["listing-slot-echo"]
assert (listed is not None) is recorded
@pytest.mark.asyncio
async def test_discovery_keeps_one_catalog_revision_across_concurrent_listings(monkeypatch):
from mcp.types import (
DiscoverRequest, ListToolsResult, ListPromptsResult, ListResourcesResult,
ListResourceTemplatesResult,
)
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import operations
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
manager = MCPServerManager()
original = MCPServer(server_id="catalog-server", name="before", transport=MCPTransport.http)
updated = original.model_copy(update={"name": "after"})
manager.registry = {original.server_id: original}
monkeypatch.setattr(proxy_server, "prisma_client", None)
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
observed = []
responses = (ListToolsResult(tools=[]), ListPromptsResult(prompts=[]),
ListResourcesResult(resources=[]), ListResourceTemplatesResult(resource_templates=[]))
def listing(index):
async def run(*args, **kwargs):
async with manager.catalog.operation():
observed.append(manager.get_mcp_server_by_id(original.server_id).name)
manager.registry = {updated.server_id: updated}
return responses[index]
return run
with (
patch.object(operations, "_execute_handle_list_tools", side_effect=listing(0)),
patch.object(operations, "_execute_list_prompts", side_effect=listing(1)),
patch.object(operations, "_execute_list_resources", side_effect=listing(2)),
patch.object(operations, "_execute_list_resource_templates", side_effect=listing(3)),
):
result = await GatewayOperations().execute(DiscoverRequest(), prepare_context(UserAPIKeyAuth(user_id="scoped")))
assert result.capabilities.model_dump(exclude_none=True) == {}
assert observed == ["before"] * 4
async with manager.catalog.operation():
assert manager.get_mcp_server_by_id(original.server_id).name == "after"
@pytest.mark.asyncio
@pytest.mark.parametrize("changed_server_id", ["private", "allowed"])
async def test_local_handler_freshness_tracks_registered_owner_with_overlapping_alias(monkeypatch, changed_server_id):
from datetime import datetime, timedelta, timezone
from fastapi import HTTPException
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
monkeypatch.setattr(proxy_server, "prisma_client", None)
manager = MCPServerManager()
now = datetime.now(timezone.utc)
private = MCPServer(server_id="private", name="billing", alias="billing", transport=MCPTransport.http, updated_at=now)
allowed = MCPServer(server_id="allowed", name="billing_admin", alias="billing-admin", transport=MCPTransport.http, updated_at=now)
manager.config_mcp_servers = {server.server_id: server for server in (private, allowed)}
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
monkeypatch.setattr(global_mcp_tool_registry, "published_tools", {})
handler = AsyncMock(return_value="private result")
global_mcp_tool_registry.register_tool("billing-admin-export", "export", {}, handler, server_id=private.server_id)
async with manager.catalog.operation():
manager.config_mcp_servers[changed_server_id] = manager.config_mcp_servers[changed_server_id].model_copy(update={"updated_at": now + timedelta(seconds=1)})
if changed_server_id == "private":
with pytest.raises(HTTPException) as denied:
await operations._handle_local_mcp_tool("billing-admin-export", {})
assert denied.value.status_code == 503
handler.assert_not_awaited()
else:
result = await operations._handle_local_mcp_tool("billing-admin-export", {})
assert result.is_error is False
handler.assert_awaited_once_with()
@pytest.mark.asyncio
@pytest.mark.parametrize("owner_present", [False, True])
@pytest.mark.parametrize("requested_server", [False, True])
async def test_local_call_cannot_borrow_another_servers_authority(
monkeypatch: pytest.MonkeyPatch, owner_present: bool, requested_server: bool
) -> None:
from datetime import datetime
from fastapi import HTTPException
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
manager: Final = MCPServerManager()
allowed: Final = MCPServer(server_id="allowed", name="allowed", transport=MCPTransport.http)
owner: Final = MCPServer(server_id="private", name="private", transport=MCPTransport.http)
manager.config_mcp_servers = {
server.server_id: server for server in ((allowed, owner) if owner_present else (allowed,))
}
handler: Final = AsyncMock(return_value="private result")
monkeypatch.setattr(proxy_server, "prisma_client", None)
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
monkeypatch.setattr(global_mcp_tool_registry, "published_tools", {})
tool_name: Final = "export" if requested_server else "private-export"
global_mcp_tool_registry.register_tool(tool_name, "export", {}, handler, server_id=owner.server_id)
async with manager.catalog.operation():
with pytest.raises(HTTPException) as denied:
await operations._execute_mcp_tool(
name=tool_name if requested_server else "allowed-private-export",
arguments={},
allowed_mcp_servers=[allowed],
start_time=datetime.now(),
requested_server_id=allowed.server_id if requested_server else None,
)
assert denied.value.status_code == 403
handler.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("server_owned", [False, True])
@pytest.mark.parametrize("requested_server", [False, True])
async def test_local_call_preserves_matching_and_legacy_handlers(
monkeypatch: pytest.MonkeyPatch, server_owned: bool, requested_server: bool
) -> None:
from datetime import datetime
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
manager: Final = MCPServerManager()
server: Final = MCPServer(server_id="allowed", name="allowed", transport=MCPTransport.http)
manager.config_mcp_servers = {server.server_id: server}
handler: Final = AsyncMock(return_value="allowed result")
monkeypatch.setattr(proxy_server, "prisma_client", None)
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
monkeypatch.setattr(global_mcp_tool_registry, "published_tools", {})
global_mcp_tool_registry.register_tool(
"export", "export", {}, handler, server_id=server.server_id if server_owned else None
)
async with manager.catalog.operation():
result: Final = await operations._execute_mcp_tool(
name="export" if requested_server else "allowed-export",
arguments={},
allowed_mcp_servers=[server],
start_time=datetime.now(),
requested_server_id=server.server_id if requested_server else None,
)
assert result.is_error is False
handler.assert_awaited_once_with()
@pytest.mark.asyncio
async def test_local_handler_rejects_an_owner_absent_from_the_catalog(monkeypatch: pytest.MonkeyPatch) -> None:
from fastapi import HTTPException
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
manager: Final = MCPServerManager()
handler: Final = AsyncMock(return_value="private result")
monkeypatch.setattr(proxy_server, "prisma_client", None)
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
monkeypatch.setattr(global_mcp_tool_registry, "published_tools", {})
global_mcp_tool_registry.register_tool("private-export", "export", {}, handler, server_id="private")
async with manager.catalog.operation():
with pytest.raises(HTTPException) as denied:
await operations._handle_local_mcp_tool("private-export", {})
assert denied.value.status_code == 503
handler.assert_not_awaited()

View file

@ -1509,6 +1509,7 @@ class TestListToolsRestAPI:
mcp_info={"server_name": "stub"},
)
stub_server.available_on_public_internet = True
monkeypatch.setattr(rest_endpoints.global_mcp_server_manager, "registry", {"server-1": stub_server})
mock_transport_ctx = AsyncMock()
mock_transport_ctx.__aenter__ = AsyncMock(return_value=(MagicMock(), MagicMock()))
@ -1577,9 +1578,8 @@ class TestListToolsRestAPI:
)
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_mcp_server_by_id",
lambda server_id: stub_server if server_id == "server-1" else None,
raising=False,
"registry",
{stub_server.server_id: stub_server},
)
request = _build_request(path="/mcp-rest/tools/list", method="GET")

View file

@ -345,7 +345,7 @@ class TestManagerShortPrefix:
class TestShortPrefixCollisionResolution:
"""``_assign_unique_short_prefix`` must rehash on collision.
"""``assign_unique_short_prefix`` must rehash on collision.
The dedup path is exercised by forcing two distinct ``server_id``
values to both hash to the same natural prefix via a monkeypatched
@ -355,7 +355,7 @@ class TestShortPrefixCollisionResolution:
def test_no_op_when_flag_off(self):
manager = MCPServerManager()
server = _make_server(server_id="abc")
manager._assign_unique_short_prefix(server)
manager.assign_unique_short_prefix(server)
assert server.short_prefix is None
def test_assigns_natural_hash_when_no_collision(self, monkeypatch):
@ -364,7 +364,7 @@ class TestShortPrefixCollisionResolution:
monkeypatch.setenv("LITELLM_USE_SHORT_MCP_TOOL_PREFIX", "true")
manager = MCPServerManager()
server = _make_server(server_id="abc")
manager._assign_unique_short_prefix(server)
manager.assign_unique_short_prefix(server)
assert server.short_prefix == mcp_utils.compute_short_server_prefix("abc")
@ -394,9 +394,9 @@ class TestShortPrefixCollisionResolution:
# Pretend both are already in the registry so dedup sees both.
manager.registry[first.server_id] = first
manager._assign_unique_short_prefix(first)
manager.assign_unique_short_prefix(first)
manager.registry[second.server_id] = second
manager._assign_unique_short_prefix(second)
manager.assign_unique_short_prefix(second)
assert first.short_prefix == "AAA"
assert second.short_prefix == "AAB"
@ -408,7 +408,7 @@ class TestShortPrefixCollisionResolution:
server = _make_server(server_id="abc")
server.short_prefix = "ZZZ" # pretend a previous registration set this
manager._assign_unique_short_prefix(server)
manager.assign_unique_short_prefix(server)
assert server.short_prefix == "ZZZ"

View file

@ -205,3 +205,32 @@ class TestInitializeGuardrail:
assert isinstance(result, MCPSecurityGuardrail)
assert result.on_violation == expected
assert result in litellm.callbacks
@pytest.mark.asyncio
@pytest.mark.parametrize("call_type", ["acompletion", "aresponses"])
async def test_guardrail_observes_saved_server_creation_and_deletion_on_another_worker(guardrail, call_type):
from unittest.mock import AsyncMock
from litellm.proxy._types import LiteLLM_MCPServerTable
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
manager = MCPServerManager()
row = LiteLLM_MCPServerTable(server_id="peer-server", alias="peer_server", transport="http",
url="https://upstream.example/mcp")
prisma = MagicMock()
prisma.writer_db = prisma.db
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=([row], []))
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
data = {"tools": [{"type": "mcp", "server_url": "litellm_proxy/mcp/peer-server"}],
"guardrails": ["test-mcp-security"]}
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", manager),
):
result = await guardrail.async_pre_call_hook(UserAPIKeyAuth(), MagicMock(), data, call_type)
assert result == data
with pytest.raises(HTTPException) as exc:
await guardrail.async_pre_call_hook(UserAPIKeyAuth(), MagicMock(), data, call_type)
assert exc.value.status_code == 400
assert exc.value.detail["unregistered_servers"] == ["peer-server"]

View file

@ -1,3 +1,6 @@
from contextlib import nullcontext
import asyncio
import os
import sys
@ -16,7 +19,7 @@ import httpx
import pytest
from pydantic import BaseModel, TypeAdapter, ValidationError
from respx import MockRouter
from fastapi import FastAPI, HTTPException
from fastapi import FastAPI, HTTPException, Request
from fastapi.testclient import TestClient
import litellm
@ -152,7 +155,7 @@ def create_mcp_router_test_client() -> TestClient:
def patch_proxy_general_settings(settings: dict):
fake_proxy_server_module = types.SimpleNamespace(general_settings=settings)
fake_proxy_server_module = types.SimpleNamespace(general_settings=settings, prisma_client=None)
return patch.dict(
sys.modules,
{"litellm.proxy.proxy_server": fake_proxy_server_module},
@ -2732,14 +2735,15 @@ class TestTemporaryMCPSessionEndpoints:
algorithm="HS256",
)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.client = None
mock_request.headers = {}
mock_request.cookies = {"token": token_cookie}
expected_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key=api_key_in_cookie
)
fake_proxy_server = types.SimpleNamespace(master_key=master_key)
fake_proxy_server = types.SimpleNamespace(master_key=master_key, prisma_client=None, general_settings={})
with (
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
@ -2772,7 +2776,8 @@ class TestTemporaryMCPSessionEndpoints:
)
expected_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.client = None
mock_request.headers = {"Authorization": "Bearer sk-header-key"}
mock_request.cookies = {}
@ -2806,7 +2811,8 @@ class TestTemporaryMCPSessionEndpoints:
)
expected_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.client = None
mock_request.headers = {}
mock_request.cookies = {}
mock_request.path_params = {"server_id": "server-1"}
@ -2816,7 +2822,8 @@ class TestTemporaryMCPSessionEndpoints:
mock_manager = MagicMock()
mock_manager.get_mcp_server_by_id.return_value = non_oauth_server
mock_manager.get_mcp_server_by_name.return_value = None
fake_proxy_server = types.SimpleNamespace(master_key=None)
mock_manager.catalog.resolve = AsyncMock(return_value=non_oauth_server)
fake_proxy_server = types.SimpleNamespace(master_key=None, prisma_client=None, general_settings={})
with (
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
@ -2854,7 +2861,8 @@ class TestTemporaryMCPSessionEndpoints:
)
expected_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
mock_request = MagicMock()
mock_request = MagicMock(spec=Request)
mock_request.client = None
mock_request.headers = {}
mock_request.cookies = {}
mock_request.path_params = {"server_id": "server-1"}
@ -2868,7 +2876,8 @@ class TestTemporaryMCPSessionEndpoints:
mock_manager = MagicMock()
mock_manager.get_mcp_server_by_id.return_value = internal_server
mock_manager.get_mcp_server_by_name.return_value = None
fake_proxy_server = types.SimpleNamespace(master_key=None)
mock_manager.catalog.resolve = AsyncMock(return_value=internal_server)
fake_proxy_server = types.SimpleNamespace(master_key=None, prisma_client=None, general_settings={})
with (
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
@ -2925,6 +2934,65 @@ class TestTemporaryMCPSessionEndpoints:
}
assert dependency_names == {None, "user_api_key_dict"}
@pytest.mark.asyncio
async def test_authorize_saved_server_on_cold_worker_without_oauth_session(self):
from collections.abc import Mapping
from starlette.requests import Request
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.proxy.management_endpoints.mcp_management_endpoints import mcp_authorize
row: Final = LiteLLM_MCPServerTable(
server_id="saved-server",
server_name="saved_server",
alias="saved_server",
transport=MCPTransport.http,
url="https://upstream.example.com/mcp",
auth_type=MCPAuth.oauth2,
oauth2_flow="authorization_code",
authorization_url="https://upstream.example.com/authorize",
token_url="https://upstream.example.com/token",
approval_status="active",
)
manager: Final = MCPServerManager()
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
async def persisted_rows(*, where: Mapping[str, object]) -> list[LiteLLM_MCPServerTable]:
return [] if where.get("approval_status") == "draft" else [row]
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=persisted_rows)
prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=None)
prisma.db.litellm_mcpservertable.find_first = AsyncMock(return_value=None)
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
request: Final = Request(
{"type": "http", "method": "GET", "scheme": "http", "server": ("localhost", 4000),
"path": "/v1/mcp/server/oauth/saved-server/authorize", "headers": [],
"query_string": b"", "client": ("127.0.0.1", 1234)}
)
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.master_key", "sk-unit-test-catalog"),
patch.object(mgmt_endpoints, "global_mcp_server_manager", manager),
):
response = await mcp_authorize(
request=request,
server_id=row.server_id,
user_api_key_dict=generate_mock_user_api_key_auth(),
client_id="client-id",
redirect_uri="http://localhost:9876/callback",
state="saved-server-test",
code_challenge=None,
code_challenge_method=None,
response_type="code",
scope=None,
)
assert response.status_code == 307
assert response.headers["location"].startswith("https://upstream.example.com/authorize?")
@pytest.mark.asyncio
async def test_mcp_authorize_proxies_to_discoverable_endpoint(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
@ -2941,8 +3009,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
) as get_server,
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.authorize_with_server",
@ -2963,7 +3031,7 @@ class TestTemporaryMCPSessionEndpoints:
)
assert result is authorize_response
get_server.assert_awaited_once_with("server-1", admin_auth, request=request)
get_server.assert_called_once_with("server-1", admin_auth, request=request)
authorize_mock.assert_awaited_once_with(
request=request,
mcp_server=server,
@ -2991,8 +3059,8 @@ class TestTemporaryMCPSessionEndpoints:
admin_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
patches = [
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.authorize_with_server",
@ -3096,8 +3164,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.mint_ephemeral_dcr_client",
@ -3242,8 +3310,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch("litellm.proxy.proxy_server.master_key", "sk-lit4581-test-master-key"),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
@ -3301,8 +3369,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch("litellm.proxy.proxy_server.master_key", "sk-lit4581-test-master-key"),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
@ -3355,8 +3423,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch("litellm.proxy.proxy_server.master_key", "sk-lit4581-test-master-key"),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
@ -3405,8 +3473,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch("litellm.proxy.proxy_server.master_key", "sk-lit4581-test-master-key"),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
@ -3453,8 +3521,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.authorize_with_server",
@ -3493,8 +3561,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
@ -3538,8 +3606,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
) as get_server,
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
@ -3561,7 +3629,7 @@ class TestTemporaryMCPSessionEndpoints:
)
assert result is exchange_response
get_server.assert_awaited_once_with("server-1", admin_auth, request=request)
get_server.assert_called_once_with("server-1", admin_auth, request=request)
exchange_mock.assert_awaited_once_with(
request=request,
mcp_server=server,
@ -3592,8 +3660,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
) as get_server,
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
@ -3615,7 +3683,7 @@ class TestTemporaryMCPSessionEndpoints:
)
assert result is exchange_response
get_server.assert_awaited_once_with("server-1", admin_auth, request=request)
get_server.assert_called_once_with("server-1", admin_auth, request=request)
exchange_mock.assert_awaited_once_with(
request=request,
mcp_server=server,
@ -3652,8 +3720,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
) as get_server,
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body",
@ -3671,7 +3739,7 @@ class TestTemporaryMCPSessionEndpoints:
)
assert result is register_response
get_server.assert_awaited_once_with("server-1", admin_auth, request=request)
get_server.assert_called_once_with("server-1", admin_auth, request=request)
read_body.assert_awaited_once_with(request=request)
register_mock.assert_awaited_once_with(
request=request,
@ -3715,8 +3783,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body",
@ -3760,8 +3828,8 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
"litellm.proxy.management_endpoints.mcp_management_endpoints._oauth_server_operation",
return_value=nullcontext(server),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body",
@ -8307,6 +8375,38 @@ async def test_config_server_edit_preserves_api_contract_without_creating_rows(r
assert manager.registry == {}
@pytest.mark.asyncio
async def test_saved_server_authorize_denial_does_not_dispatch_upstream():
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
manager: Final = MCPServerManager()
row: Final = LiteLLM_MCPServerTable(server_id="denied-peer", alias="denied_peer", transport=MCPTransport.http,
auth_type=MCPAuth.oauth2, url="https://upstream.example/mcp",
authorization_url="https://upstream.example/authorize", token_url="https://upstream.example/token")
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
user: Final = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.INTERNAL_USER)
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch.object(mgmt_endpoints, "global_mcp_server_manager", manager),
patch.object(mgmt_endpoints, "get_cached_temporary_mcp_server", AsyncMock(return_value=None)),
patch.object(mgmt_endpoints, "build_effective_auth_contexts", AsyncMock(return_value=(user,))),
patch.object(manager, "get_allowed_mcp_servers", AsyncMock(return_value=[])),
patch.object(mgmt_endpoints, "authorize_with_server", new_callable=AsyncMock) as authorize,
patch.object(mgmt_endpoints, "resolve_ephemeral_dcr_client", new_callable=AsyncMock) as register,
):
with pytest.raises(HTTPException) as exc:
await mgmt_endpoints.mcp_authorize(request=None, server_id=row.server_id, user_api_key_dict=user,
client_id="client", redirect_uri="http://localhost/callback")
assert exc.value.status_code == 403
authorize.assert_not_awaited()
register.assert_not_awaited()
assert prisma.db.litellm_mcpservertable.find_many.await_count == 1
class TestDuplicateIdentifierRejection:
"""server_name/alias must be unique across live servers, case-insensitive.
@ -8801,11 +8901,15 @@ def _mock_mcp_resolution_prisma_client(
object_permission: LiteLLM_ObjectPermissionTable | None = None,
) -> MagicMock:
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
prisma.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=SimpleNamespace(object_permission=key_permission)
)
def matches_server_filter(name: str, condition: object) -> bool:
if name == "OR" and condition == [{"approval_status": None}, {"approval_status": {"in": ["active", "approved"]}}]:
return server.approval_status in (None, "active", "approved")
if name == "submitted_by":
return server.submitted_by == condition
if name == "server_id":
@ -9213,7 +9317,7 @@ class TestMCPServerResolutionRegressions:
allowed_id: Final = "lit3974-allowed-config"
denied_id: Final = "lit3974-denied-config"
prisma: Final = _mock_mcp_resolution_prisma_client(
generate_mock_mcp_server_db_record(server_id=denied_id),
generate_mock_mcp_server_db_record(server_id=denied_id, alias="denied_alias"),
LiteLLM_ObjectPermissionTable(
object_permission_id="lit3974-alias-permission",
mcp_servers=[allowed_id],
@ -9376,6 +9480,15 @@ class TestMCPServerResolutionCharacterization:
}
)
prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=hidden)
eligible: Final = approval_status in (None, "active", "approved")
existing_find_many: Final = prisma.db.litellm_mcpservertable.find_many
async def current_rows(**kwargs: object) -> list[LiteLLM_MCPServerTable]:
if kwargs.get("where") == {"OR": [{"approval_status": None}, {"approval_status": {"in": ["active", "approved"]}}]}:
return [hidden] if eligible else []
return await existing_find_many(**kwargs)
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=current_rows)
if not registered:
manager.config_mcp_servers = {}
health: Final = AsyncMock(return_value=hidden)
@ -9390,7 +9503,13 @@ class TestMCPServerResolutionCharacterization:
patch("litellm.proxy.proxy_server.general_settings", {"user_mcp_management_mode": "view_all"}),
):
listed: Final = await mgmt_endpoints.fetch_all_mcp_servers(auth, team_id=None)
assert (server_id in {item.server_id for item in listed}) is registered
assert (server_id in {item.server_id for item in listed}) is (registered or eligible)
if caller != "admin" and eligible:
public_detail: Final = await mgmt_endpoints.fetch_mcp_server(_make_mock_request(), server_id, auth)
assert public_detail.credentials is None
assert public_detail.url is None
assert public_detail.static_headers is None
return
if caller != "admin":
with pytest.raises(HTTPException) as error:
await mgmt_endpoints.fetch_mcp_server(_make_mock_request(), server_id, auth)
@ -10491,7 +10610,7 @@ class TestMCPServerResolutionCharacterization:
effects: Final = _ResolutionEffects()
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_cache_temporary_mcp_server_in_redis,
_get_cached_temporary_mcp_server_or_404,
_oauth_server_operation,
_TemporaryMCPServerEntry,
)
@ -10576,11 +10695,8 @@ class TestMCPServerResolutionCharacterization:
):
if expected_status in (403, 404):
with pytest.raises(HTTPException) as exc_info:
await _get_cached_temporary_mcp_server_or_404(
server_id,
auth,
request=_make_mock_request(),
)
async with _oauth_server_operation(server_id, auth, request=_make_mock_request()):
pytest.fail("Denied OAuth resolution must not enter the operation")
expected_detail: Final = (
{"error": f"MCP server {server_id} not found"}
@ -10594,20 +10710,16 @@ class TestMCPServerResolutionCharacterization:
effects.assert_no_writes()
assert httpx_mock.calls.call_count == 0, f"{source}/{caller}: no upstream HTTP"
else:
resolved: Final = await _get_cached_temporary_mcp_server_or_404(
server_id,
auth,
request=_make_mock_request(),
)
expected_alias: Final = (
db_server.alias
if source == "temp_draft"
else "LIT3974 OAuth"
if source == "config"
else temp_server.alias
)
assert resolved.server_id == server_id, f"{source}/{caller}: resolved OAuth server"
assert resolved.alias == expected_alias, f"{source}/{caller}: resolved OAuth display name"
async with _oauth_server_operation(server_id, auth, request=_make_mock_request()) as resolved:
expected_alias: Final = (
db_server.alias
if source == "temp_draft"
else "LIT3974 OAuth"
if source == "config"
else temp_server.alias
)
assert resolved.server_id == server_id, f"{source}/{caller}: resolved OAuth server"
assert resolved.alias == expected_alias, f"{source}/{caller}: resolved OAuth display name"
finally:
mgmt_endpoints.litellm.cache = original_cache

View file

@ -3,8 +3,7 @@ Tests for the dynamic_mcp_route handler in proxy_server.py.
Covers the resolution order:
1. Registered MCP server alias → forwards to /mcp/{name}
2. Comma-separated list → short-circuits before any DB call;
forwarded to /mcp/{segment}
2. Comma-separated list → forwards to /mcp/{segment}
3. Toolset name (cached) → sets toolset scope, forwards to /mcp
4. MCP access group tag (cached) → forwards to /mcp/{name} when the group
resolves to at least one server
@ -611,3 +610,66 @@ def test_aggregate_mcp_route_returns_404_when_mcp_unavailable():
assert response.status_code == 404
assert handler_calls == []
@pytest.mark.asyncio
@pytest.mark.parametrize("change", ["create", "rename", "delete"])
@pytest.mark.parametrize("csv", [False, True])
async def test_dynamic_route_observes_committed_peer_catalog_changes(monkeypatch, change, csv):
from datetime import datetime
from types import SimpleNamespace
from starlette.responses import Response
from litellm.caching.caching import DualCache
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import mcp_server_manager
from litellm.proxy._types import LiteLLM_MCPServerTable
old_row = LiteLLM_MCPServerTable(
server_id="catalog-route", server_name="previous", alias="previous", transport="http",
url="https://previous.example/mcp", updated_at=datetime(2026, 1, 1),
)
new_row = old_row.model_copy(update={
"server_name": "current", "alias": "current", "url": "https://current.example/mcp",
"updated_at": datetime(2026, 1, 2),
})
async def find_rows(*, where):
if "mcp_access_groups" in where or change == "delete":
return []
return [new_row]
read_rows = AsyncMock(side_effect=find_rows)
prisma = SimpleNamespace(db=SimpleNamespace(
litellm_mcpservertable=SimpleNamespace(find_many=read_rows),
litellm_mcptoolsettable=SimpleNamespace(find_first=AsyncMock(return_value=None)),
litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None)),
))
prisma.writer_db = prisma.db
manager = mcp_server_manager.MCPServerManager()
if change != "create":
manager.registry = {old_row.server_id: await manager.build_mcp_server_from_table(old_row)}
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache())
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager)
name = "previous" if change == "delete" else "current"
segment = f"{name},missing" if csv else name
request = _make_request(f"/{segment}/mcp")
async def forwarded(path_segment, request):
if change != "delete":
assert manager.get_mcp_server_by_name(path_segment).url == new_row.url
return Response(content=b"forwarded", status_code=200)
with patch(_FORWARD, new=AsyncMock(side_effect=forwarded)) as forward:
if change == "delete":
with pytest.raises(HTTPException) as exc:
await proxy_server.dynamic_mcp_route(segment, request)
assert exc.value.status_code == 404
forward.assert_not_awaited()
else:
response = await proxy_server.dynamic_mcp_route(segment, request)
assert response.status_code == 200
forward.assert_awaited_once_with(name, request)
assert sum("OR" in call.kwargs["where"] for call in read_rows.await_args_list) == 1

View file

@ -4,6 +4,7 @@ import subprocess
import sys
import textwrap
import types
from contextlib import nullcontext
from typing import Any, Final, Literal, cast
from unittest.mock import AsyncMock, MagicMock, patch
@ -42,10 +43,11 @@ class _DummyMCPResult:
def _setup_mcp_call_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock:
"""Patch MCP globals so _execute_tool_calls can run in tests."""
proxy_module = types.SimpleNamespace(proxy_logging_obj=object())
proxy_module = types.SimpleNamespace(proxy_logging_obj=object(), prisma_client=None)
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module)
fake_manager = types.SimpleNamespace(
catalog=types.SimpleNamespace(operation=nullcontext),
get_registry=MagicMock(return_value={}),
call_tool=AsyncMock(return_value=_DummyMCPResult()),
# Newer logging path calls this to enrich spend logs metadata
@ -63,7 +65,7 @@ def _setup_proxy_logging(monkeypatch: pytest.MonkeyPatch) -> AsyncMock:
"""Patch proxy_logging_obj so failure hook can be asserted."""
proxy_logging_obj = MagicMock()
proxy_logging_obj.post_call_failure_hook = AsyncMock()
proxy_module = types.SimpleNamespace(proxy_logging_obj=proxy_logging_obj)
proxy_module = types.SimpleNamespace(proxy_logging_obj=proxy_logging_obj, prisma_client=None)
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module)
return proxy_logging_obj.post_call_failure_hook
@ -386,6 +388,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey
post_call_failure_hook = _setup_proxy_logging(monkeypatch)
fake_manager = types.SimpleNamespace(
catalog=types.SimpleNamespace(operation=nullcontext),
get_registry=MagicMock(return_value={}),
call_tool=AsyncMock(side_effect=HTTPException(status_code=500, detail="boom")),
)
@ -504,6 +507,7 @@ async def test_execute_tool_calls_applies_post_call_hook_content(monkeypatch):
isError=False,
)
fake_manager = types.SimpleNamespace(
catalog=types.SimpleNamespace(operation=nullcontext),
get_registry=MagicMock(return_value={}),
call_tool=AsyncMock(return_value=result),
_get_mcp_server_from_tool_name=MagicMock(return_value=None),
@ -545,6 +549,7 @@ async def test_execute_tool_calls_returns_proxy_result_without_logging(monkeypat
)
fake_manager = types.SimpleNamespace(
catalog=types.SimpleNamespace(operation=nullcontext),
get_registry=MagicMock(return_value={}),
call_tool=AsyncMock(return_value=result),
_get_mcp_server_from_tool_name=MagicMock(return_value=None),
@ -578,6 +583,7 @@ async def test_execute_tool_calls_passes_logging_details_to_proxy_hook(monkeypat
)
fake_manager = types.SimpleNamespace(
catalog=types.SimpleNamespace(operation=nullcontext),
get_registry=MagicMock(return_value={}),
call_tool=AsyncMock(return_value=result),
_get_mcp_server_from_tool_name=MagicMock(return_value=None),
@ -613,6 +619,7 @@ async def test_execute_tool_calls_continues_when_post_call_logging_fails(monkeyp
result = CallToolResult(content=[TextContent(type="text", text="ok")], isError=False)
fake_manager = types.SimpleNamespace(
catalog=types.SimpleNamespace(operation=nullcontext),
get_registry=MagicMock(return_value={}),
call_tool=AsyncMock(return_value=result),
_get_mcp_server_from_tool_name=MagicMock(return_value=None),
@ -668,6 +675,7 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch
# Patch manager methods used by _get_mcp_tools_from_manager to avoid needing full UserAPIKeyAuth fields.
fake_manager = types.SimpleNamespace(
catalog=types.SimpleNamespace(operation=nullcontext),
get_registry=MagicMock(return_value={}),
get_allowed_mcp_servers=AsyncMock(return_value=[]),
get_mcp_servers_from_ids=MagicMock(return_value=[]),
@ -725,6 +733,7 @@ async def test_get_mcp_tools_from_manager_forwards_request_tags(monkeypatch):
mock_get_tools,
)
fake_manager = types.SimpleNamespace(
catalog=types.SimpleNamespace(operation=nullcontext),
get_registry=MagicMock(return_value={}),
get_allowed_mcp_servers=AsyncMock(return_value=[]),
get_mcp_servers_from_ids=MagicMock(return_value=[]),
@ -1287,6 +1296,7 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py
"authorization": "proxy-sentinel",
}
manager: Final = types.SimpleNamespace(
catalog=types.SimpleNamespace(operation=nullcontext),
get_registry=MagicMock(return_value={}),
get_allowed_mcp_servers=AsyncMock(return_value=[]),
get_mcp_servers_from_ids=MagicMock(return_value=[]),
@ -1347,6 +1357,7 @@ async def test_bridge_listing_leaves_the_callers_catalog_unchanged(
MCPTool(name="echo", description="Duplicate echo", inputSchema={"type": "object", "properties": {}}),
]
fake_manager: Final = types.SimpleNamespace(
catalog=types.SimpleNamespace(operation=nullcontext),
get_registry=MagicMock(return_value={}),
get_allowed_mcp_servers=AsyncMock(return_value=[]),
get_mcp_servers_from_ids=MagicMock(return_value=[]),
@ -1485,6 +1496,7 @@ async def test_concurrent_bridge_calls_use_their_own_served_metadata(monkeypatch
def _toolset_gateway_manager(toolset_id: str, server_id: str) -> types.SimpleNamespace:
return types.SimpleNamespace(
catalog=types.SimpleNamespace(operation=nullcontext),
get_registry=MagicMock(return_value={}),
get_allowed_mcp_servers=AsyncMock(return_value=[]),
get_mcp_servers_from_ids=MagicMock(return_value=[]),

View file

@ -1,5 +1,6 @@
import sys
import types
from contextlib import nullcontext
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
@ -77,6 +78,7 @@ def _mock_mcp_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock:
"""Patch the MCP tool-call plumbing so _execute_tool_calls can run in tests."""
call_tool = AsyncMock(return_value=CallToolResult(content=[TextContent(type="text", text="ok")], isError=False))
fake_manager = types.SimpleNamespace(
catalog=types.SimpleNamespace(operation=nullcontext),
get_registry=MagicMock(return_value={}),
call_tool=call_tool,
_get_mcp_server_from_tool_name=MagicMock(return_value=None),
@ -89,7 +91,7 @@ def _mock_mcp_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock:
monkeypatch.setitem(
sys.modules,
"litellm.proxy.proxy_server",
types.SimpleNamespace(proxy_logging_obj=MagicMock()),
types.SimpleNamespace(proxy_logging_obj=MagicMock(), prisma_client=None),
)
return call_tool

View file

@ -24697,7 +24697,7 @@ export interface paths {
*
* Resolution order:
* 1. Registered MCP server alias / name
* 2. Comma-separated list (short-circuits before any DB call)
* 2. Comma-separated list
* 3. Toolset name (DB lookup, cached)
* 4. MCP access group tag (DB lookup, cached)
*/
@ -24708,7 +24708,7 @@ export interface paths {
*
* Resolution order:
* 1. Registered MCP server alias / name
* 2. Comma-separated list (short-circuits before any DB call)
* 2. Comma-separated list
* 3. Toolset name (DB lookup, cached)
* 4. MCP access group tag (DB lookup, cached)
*/
@ -24719,7 +24719,7 @@ export interface paths {
*
* Resolution order:
* 1. Registered MCP server alias / name
* 2. Comma-separated list (short-circuits before any DB call)
* 2. Comma-separated list
* 3. Toolset name (DB lookup, cached)
* 4. MCP access group tag (DB lookup, cached)
*/
@ -24730,7 +24730,7 @@ export interface paths {
*
* Resolution order:
* 1. Registered MCP server alias / name
* 2. Comma-separated list (short-circuits before any DB call)
* 2. Comma-separated list
* 3. Toolset name (DB lookup, cached)
* 4. MCP access group tag (DB lookup, cached)
*/
@ -24741,7 +24741,7 @@ export interface paths {
*
* Resolution order:
* 1. Registered MCP server alias / name
* 2. Comma-separated list (short-circuits before any DB call)
* 2. Comma-separated list
* 3. Toolset name (DB lookup, cached)
* 4. MCP access group tag (DB lookup, cached)
*/
@ -24752,7 +24752,7 @@ export interface paths {
*
* Resolution order:
* 1. Registered MCP server alias / name
* 2. Comma-separated list (short-circuits before any DB call)
* 2. Comma-separated list
* 3. Toolset name (DB lookup, cached)
* 4. MCP access group tag (DB lookup, cached)
*/
@ -24763,7 +24763,7 @@ export interface paths {
*
* Resolution order:
* 1. Registered MCP server alias / name
* 2. Comma-separated list (short-circuits before any DB call)
* 2. Comma-separated list
* 3. Toolset name (DB lookup, cached)
* 4. MCP access group tag (DB lookup, cached)
*/