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

This commit is contained in:
Joshua Valluru 2026-09-22 12:04:09 -07:00
parent 2ea43214c9
commit c6b5f6da7f
16 changed files with 478 additions and 69 deletions

View file

@ -13,6 +13,7 @@ from typing_extensions import assert_never
import litellm
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.catalog import with_mcp_catalog
from litellm.proxy._experimental.mcp_server.oauth_utils import (
get_passthrough_resource_metadata_url,
get_passthrough_www_authenticate,
@ -398,6 +399,7 @@ class MCPRequestHandler:
LITELLM_MCP_ACCESS_GROUPS_HEADER_NAME = SpecialHeaders.mcp_access_groups.value
@staticmethod
@with_mcp_catalog
async def process_mcp_request(
scope: Scope,
) -> tuple[

View file

@ -689,7 +689,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,94 @@
from __future__ import annotations
import asyncio
from collections.abc import AsyncIterator, Awaitable, Callable, Mapping
from contextlib import asynccontextmanager
from contextvars import ContextVar
from dataclasses import dataclass
from functools import wraps
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, ParamSpec, TypeVar
if TYPE_CHECKING:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.types.mcp_server.mcp_server_manager import MCPServer
_P: Final = ParamSpec("_P")
_R: Final = TypeVar("_R")
@dataclass(slots=True)
class _CatalogScope:
servers: Mapping[str, MCPServer]
owner_task_id: int
active: bool = True
class TargetCatalog:
def __init__(self, manager: MCPServerManager) -> None:
self._manager = manager
self._reload_lock = asyncio.Lock()
self._scope: ContextVar[_CatalogScope | None] = ContextVar("mcp_catalog_scope", default=None)
@property
def current(self) -> Mapping[str, MCPServer] | None:
scope: Final = self._scope.get()
return scope.servers if scope is not None and scope.active else None
async def refresh(self) -> None:
async with self._reload_lock:
await self._manager._reload_servers_from_database() # pyright: ignore[reportPrivateUsage] # existing staged loader
async def list(self) -> Mapping[str, MCPServer]:
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime proxy dependency
scope: Final = self._scope.get()
if scope is not None and scope.active and scope.owner_task_id == id(asyncio.current_task()):
return scope.servers
async with self._reload_lock:
if prisma_client is not None:
try:
await self._manager._reload_servers_from_database(reuse_unchanged=True) # pyright: ignore[reportPrivateUsage] # existing staged loader
except Exception as exc: # noqa: BLE001 # never serve an unverified database snapshot
from fastapi import HTTPException # noqa: PLC0415 # optional proxy dependency
raise HTTPException(
status_code=503, detail="MCP server configuration could not be refreshed"
) from exc
return MappingProxyType(self._manager.config_mcp_servers | self._manager.registry)
@asynccontextmanager
async def operation(self) -> AsyncIterator[Mapping[str, MCPServer]]:
existing: Final = self._scope.get()
if existing is not None and existing.active and existing.owner_task_id == id(asyncio.current_task()):
yield existing.servers
return
scope: Final = _CatalogScope(await self.list(), id(asyncio.current_task()))
token: Final = self._scope.set(scope)
try:
yield scope.servers
finally:
scope.active = False
self._scope.reset(token)
async def resolve(self, lookup: str, client_ip: str | None = None) -> MCPServer | None:
async with self.operation():
return self._manager.get_mcp_server_by_name(
lookup, client_ip=client_ip
) or self._manager.get_mcp_server_by_id(lookup, client_ip=client_ip)
def with_mcp_catalog(function: Callable[_P, Awaitable[_R]]) -> Callable[_P, Awaitable[_R]]:
@wraps(function)
async def wrapped(
*args: _P.args,
**kwargs: _P.kwargs, # kwargs-ok: ParamSpec preserves each wrapped endpoint keyword contract
) -> _R:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # manager imports catalog
global_mcp_server_manager,
)
async with global_mcp_server_manager.catalog.operation():
return await function(*args, **kwargs)
return wrapped

View file

@ -36,6 +36,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 with_mcp_catalog
from litellm.proxy._experimental.mcp_server.faults import (
CallerRejected,
CredentialSource,
@ -1875,6 +1876,7 @@ async def register_client_with_server(
@router.get("/authorize/mcp-session")
@with_mcp_catalog
async def authorize_mcp_session(
request: Request,
redirect_uri: str,
@ -1900,6 +1902,7 @@ async def authorize_mcp_session(
@router.get("/{mcp_server_name}/authorize")
@router.get("/authorize")
@with_mcp_catalog
async def authorize(
request: Request,
redirect_uri: str,
@ -1938,9 +1941,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:
@ -1974,6 +1979,7 @@ async def authorize(
@router.post("/{mcp_server_name}/token")
@router.post("/token")
@with_mcp_catalog
async def token_endpoint(
request: Request,
grant_type: str = Form(...),
@ -2067,6 +2073,7 @@ async def _vendor_credential_state(user_id: str, server_id: str) -> VendorCreden
@router.get("/authorize/flow")
@with_mcp_catalog
async def authorize_flow(request: Request, flow: str) -> Response:
return await describe_connect_flow(
request=request,
@ -2078,6 +2085,7 @@ async def authorize_flow(request: Request, flow: str) -> Response:
@router.post("/authorize/complete")
@with_mcp_catalog
async def authorize_complete(
request: Request,
flow: str = Form(...),
@ -2435,6 +2443,7 @@ def is_network_error(exc: Exception) -> bool:
return isinstance(exc, httpx.TransportError)
@with_mcp_catalog
async def _build_oauth_protected_resource_response(
request: Request,
mcp_server_name: str | None,
@ -2739,12 +2748,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
@ -2792,6 +2796,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}}")
@with_mcp_catalog
async def oauth_authorization_server_mcp_standard(request: Request, mcp_server_name: str):
"""
OAuth authorization server discovery endpoint using standard MCP URL pattern.
@ -2809,6 +2814,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")
@with_mcp_catalog
async def oauth_authorization_server_mcp(request: Request, mcp_server_name: str | None = None):
"""
OAuth authorization server discovery endpoint.
@ -2882,6 +2888,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")
@with_mcp_catalog
async def oauth_authorization_server_legacy(request: Request, mcp_server_name: str):
"""
OAuth authorization server discovery for legacy /{server_name}/mcp pattern.
@ -2903,46 +2910,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

@ -73,6 +73,7 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPServerAccess,
_is_mcp_admitted_user_subject,
)
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
MCP_ELICITATION_AVAILABLE,
@ -1860,6 +1861,7 @@ class MCPServerManager:
self._template_discovery_cache = _DiscoveryCache[ResourceTemplate](
discovery_ttl, discovery_clock, TypeAdapter(tuple[ResourceTemplate, ...])
)
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] = {}
@ -2287,11 +2289,12 @@ 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
snapshot: Final = self.catalog.current
return snapshot if snapshot is not None else self.config_mcp_servers | self.registry
def is_config_declared_server(self, server_id: str) -> bool:
"""True when server_id was declared in config.yaml (present in the in-memory config map).
@ -6507,14 +6510,15 @@ class MCPServerManager:
return None
async def reload_servers_from_database(self):
await self.catalog.refresh()
async def _reload_servers_from_database(self, *, reuse_unchanged: bool = False):
"""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")
@ -6593,7 +6597,8 @@ class MCPServerManager:
# 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)
if not reuse_unchanged or new_server is not previous_registry.get(server_id):
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
@ -6607,10 +6612,13 @@ class MCPServerManager:
dropped_registry_keys: Final = previous_registry.keys() - registered_registry.keys()
for registry_key in dropped_registry_keys:
self._cleanup_server_tool_routing_artifacts(previous_registry[registry_key])
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._upstream_initialize_instructions_by_server_id.pop(server_id, None)
self._upstream_initialize_instructions_probed_at.pop(server_id, None)
self._invalidate_discovery_lists(server_id)
self.registry = registered_registry
# A discovery task may have published into ``previous_registry`` while
@ -6834,7 +6842,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

@ -50,6 +50,7 @@ from litellm.proxy._experimental.mcp_server.byok_credential_cache import (
cache_byok_credential,
get_cached_byok_credential,
)
from litellm.proxy._experimental.mcp_server.catalog import with_mcp_catalog
from litellm.proxy._experimental.mcp_server.contracts import (
AuthorizedToolCall,
OperationContext,
@ -614,6 +615,7 @@ def apply_tool_overrides(
return tools
@with_mcp_catalog
async def _get_allowed_mcp_servers(
user_api_key_auth: UserAPIKeyAuth | None,
mcp_servers: Sequence[str] | None,
@ -1435,6 +1437,7 @@ async def filter_tools_by_key_team_permissions(
]
@with_mcp_catalog
async def _list_mcp_tools(
user_api_key_auth: UserAPIKeyAuth | None = None,
mcp_auth_header: str | None = None,
@ -1485,6 +1488,7 @@ async def _list_mcp_tools(
return AggregateToolListing(tools=[], outcomes={})
@with_mcp_catalog
async def _list_mcp_prompts(
user_api_key_auth: UserAPIKeyAuth | None = None,
mcp_auth_header: str | None = None,
@ -1526,6 +1530,7 @@ async def _list_mcp_prompts(
return managed_prompts
@with_mcp_catalog
async def _list_mcp_resources(
user_api_key_auth: UserAPIKeyAuth | None = None,
mcp_auth_header: str | None = None,
@ -1555,6 +1560,7 @@ async def _list_mcp_resources(
return managed_resources
@with_mcp_catalog
async def _list_mcp_resource_templates(
user_api_key_auth: UserAPIKeyAuth | None = None,
mcp_auth_header: str | None = None,
@ -3055,6 +3061,7 @@ class GatewayOperations:
@overload
async def execute(self, operation: ReadResourceRequest, context: OperationContext) -> ReadResourceResult: ...
@with_mcp_catalog
async def execute(self, operation: GatewayOperation, context: OperationContext) -> GatewayResult:
match operation:
case AuthorizedToolCall():

View file

@ -22,6 +22,7 @@ from litellm.exceptions import (
GuardrailRaisedException,
ModifyResponseException,
)
from litellm.proxy._experimental.mcp_server.catalog import with_mcp_catalog
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPServerListError,
MCPServerURLCredentialsError,
@ -861,6 +862,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)])
@with_mcp_catalog
async def list_tool_rest_api(
request: Request,
server_id: str | None = Query(None, description="The server id to list tools for"),
@ -1085,6 +1087,7 @@ if MCP_AVAILABLE:
}
@router.post("/tools/call", dependencies=[Depends(user_api_key_auth)])
@with_mcp_catalog
async def call_tool_rest_api(
request: Request,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),

View file

@ -45,6 +45,8 @@ from fastapi import (
from fastapi.responses import JSONResponse
from typing_extensions import ReadOnly, TypedDict
from litellm.proxy._experimental.mcp_server.catalog import with_mcp_catalog
try:
from prisma.errors import RecordNotFoundError, UniqueViolationError
except ImportError:
@ -1043,6 +1045,7 @@ if MCP_AVAILABLE:
tags=["mcp"],
description="MCP registry endpoint. Spec: https://github.com/modelcontextprotocol/registry",
)
@with_mcp_catalog
async def get_mcp_registry(request: Request):
if not _is_public_registry_enabled():
raise HTTPException(
@ -1166,6 +1169,7 @@ if MCP_AVAILABLE:
dependencies=[Depends(user_api_key_auth)],
response_model=list[LiteLLM_MCPServerTable],
)
@with_mcp_catalog
async def fetch_all_mcp_servers(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
team_id: str | None = Query(
@ -1284,6 +1288,7 @@ if MCP_AVAILABLE:
description="Health check for MCP servers",
dependencies=[Depends(user_api_key_auth)],
)
@with_mcp_catalog
async def health_check_servers(
server_ids: list[str] | None = Query(
None,
@ -1959,6 +1964,7 @@ if MCP_AVAILABLE:
return _redact_mcp_credentials(temp_record)
@with_mcp_catalog
async def _mcp_oauth_user_api_key_auth(request: Request) -> UserAPIKeyAuth:
"""
Auth dependency for MCP OAuth browser-navigation endpoints (/authorize, /token).
@ -2059,6 +2065,7 @@ if MCP_AVAILABLE:
request_data=request_data,
)
@with_mcp_catalog
async def _get_cached_temporary_mcp_server_or_404(
server_id: str,
user_api_key_dict: UserAPIKeyAuth,

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 with_mcp_catalog
from litellm.proxy._experimental.mcp_server.utils import (
iter_known_server_prefixes,
logging_safe_mcp_headers,
@ -178,6 +179,7 @@ class LiteLLM_Proxy_MCP_Handler:
)
@staticmethod
@with_mcp_catalog
async def routes_through_gateway(
tools: Iterable[Mapping[str, object]] | None,
served_names: Callable[[Collection[str]], Awaitable[frozenset[str]]] = _gateway_served_names,
@ -229,6 +231,7 @@ class LiteLLM_Proxy_MCP_Handler:
return user_api_key_auth
@staticmethod
@with_mcp_catalog
async def _get_mcp_tools_from_manager(
user_api_key_auth: "UserAPIKeyAuth | None",
mcp_tools_with_litellm_proxy: Iterable[Mapping[str, object]] | None,
@ -680,6 +683,7 @@ class LiteLLM_Proxy_MCP_Handler:
return result_text or "Tool executed successfully"
@staticmethod
@with_mcp_catalog
async def _execute_tool_calls(
tool_server_map: dict[str, str],
tool_calls: Sequence[object],

View file

@ -7482,9 +7482,11 @@ class TestGatewaySessionAdmission:
rpm_limit=rpm_limit,
)
)
prisma = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
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.prisma_client", prisma),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
):
yield get_user_object
@ -9571,10 +9573,12 @@ class TestScopedSessionAdmission:
rpm_limit=None,
)
)
prisma = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
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.prisma_client", prisma),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
):
auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope_dict)

View file

@ -632,6 +632,7 @@ async def test_execute_byok_tool_missing_credential_advertises_api_key_flow(monk
mcp_operations.byok_credential_cache.flush_cache()
server = MCPServer(server_id="byok-discovery", name="byok-discovery", transport=MCPTransport.http, is_byok=True)
prisma = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
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

@ -738,7 +738,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
@ -761,17 +762,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
@ -811,7 +813,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()
@ -12347,3 +12351,61 @@ async def test_identity_bound_authorize_unrelated_bearer_uses_browser_session(
proxy_server.prisma_client.db.litellm_mcpusercredentials.upsert.assert_not_called()
proxy_server.prisma_client.db.litellm_usertable.create.assert_not_called()
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)))
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()

View file

@ -14516,3 +14516,209 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie
assert captured["client_ip"] is None
finally:
auth_context_var.reset(token)
def _catalog_row(name="initial"):
return LiteLLM_MCPServerTable(
server_id="catalog-server",
server_name=name,
transport="http",
url="https://upstream.example/mcp",
created_at=datetime(2026, 1, 1),
updated_at=datetime(2026, 1, 1 if name == "initial" else 2),
)
def _catalog_database(monkeypatch, read_rows):
from types import SimpleNamespace
from litellm.proxy import proxy_server
client = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=SimpleNamespace(find_many=read_rows)))
monkeypatch.setattr(proxy_server, "prisma_client", client)
@pytest.mark.asyncio
async def test_catalog_pins_one_snapshot_and_next_operation_reads_current_rows(monkeypatch):
read_rows = AsyncMock(return_value=[_catalog_row()])
_catalog_database(monkeypatch, read_rows)
manager = MCPServerManager()
async with manager.catalog.operation():
initial = manager.get_mcp_server_by_id("catalog-server")
read_rows.return_value = [_catalog_row("updated")]
async with manager.catalog.operation():
assert await manager.catalog.resolve("initial") is initial
assert tuple((await manager.catalog.list()).values()) == (initial,)
read_rows.assert_awaited_once()
async with manager.catalog.operation():
assert manager.get_mcp_server_by_id("catalog-server").name == "updated"
assert await manager.catalog.resolve("initial") is None
assert manager.catalog.current is None
assert read_rows.await_count == 2
@pytest.mark.asyncio
async def test_catalog_database_failure_keeps_registry_but_rejects_operation(monkeypatch):
read_rows = AsyncMock(return_value=[_catalog_row()])
_catalog_database(monkeypatch, read_rows)
manager = MCPServerManager()
snapshot = await manager.catalog.list()
read_rows.side_effect = RuntimeError("private connection detail")
with pytest.raises(HTTPException) as exc:
async with manager.catalog.operation():
pytest.fail("An unverified database snapshot must not execute")
assert exc.value.status_code == 503
assert exc.value.detail == "MCP server configuration could not be refreshed"
assert manager.registry == dict(snapshot)
assert manager.catalog.current is None
@pytest.mark.asyncio
async def test_catalog_serializes_background_and_operation_refreshes(monkeypatch):
entered = asyncio.Event()
release = asyncio.Event()
async def first_read(**kwargs):
entered.set()
await release.wait()
return [_catalog_row()]
read_rows = AsyncMock(side_effect=first_read)
_catalog_database(monkeypatch, read_rows)
manager = MCPServerManager()
first = asyncio.create_task(manager.reload_servers_from_database())
await entered.wait()
read_rows.side_effect = None
read_rows.return_value = [_catalog_row("updated")]
second = asyncio.create_task(manager.catalog.list())
await asyncio.sleep(0)
assert read_rows.await_count == 1
release.set()
await first
latest = await second
assert latest["catalog-server"].name == "updated"
assert manager.registry["catalog-server"].name == "updated"
@pytest.mark.asyncio
async def test_catalog_cancellation_releases_refresh_lock(monkeypatch):
entered = asyncio.Event()
blocked = asyncio.Event()
async def blocked_read(**kwargs):
entered.set()
await blocked.wait()
return []
read_rows = AsyncMock(side_effect=blocked_read)
_catalog_database(monkeypatch, read_rows)
manager = MCPServerManager()
task = asyncio.create_task(manager.catalog.list())
await entered.wait()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
read_rows.side_effect = None
read_rows.return_value = [_catalog_row()]
snapshot = await asyncio.wait_for(manager.catalog.list(), timeout=1)
assert tuple(snapshot) == ("catalog-server",)
assert manager.catalog.current is None
@pytest.mark.asyncio
async def test_catalog_child_operation_does_not_inherit_stale_session_snapshot(monkeypatch):
read_rows = AsyncMock(return_value=[_catalog_row()])
_catalog_database(monkeypatch, read_rows)
manager = MCPServerManager()
release = asyncio.Event()
async def child_operation():
await release.wait()
async with manager.catalog.operation():
return manager.get_mcp_server_by_id("catalog-server")
async with manager.catalog.operation():
child = asyncio.create_task(child_operation())
read_rows.return_value = [_catalog_row("updated")]
release.set()
assert (await child).name == "updated"
assert read_rows.await_count == 2
@pytest.mark.asyncio
async def test_catalog_config_only_snapshot_cleans_up_after_error(monkeypatch):
from litellm.proxy import proxy_server
monkeypatch.setattr(proxy_server, "prisma_client", None)
manager = MCPServerManager()
configured = MCPServer(server_id="configured", name="configured", transport="stdio", command="echo")
manager.config_mcp_servers = {configured.server_id: configured}
async def fail_operation():
async with manager.catalog.operation():
assert await manager.catalog.resolve("configured") is configured
raise ValueError("stop operation")
with pytest.raises(ValueError, match="stop operation"):
await fail_operation()
assert manager.catalog.current is None
assert dict(await manager.catalog.list()) == {"configured": configured}
@pytest.mark.asyncio
async def test_catalog_delete_drops_derived_tool_mapping(monkeypatch):
read_rows = AsyncMock(return_value=[_catalog_row()])
_catalog_database(monkeypatch, read_rows)
manager = MCPServerManager()
await manager.catalog.list()
manager.tool_name_to_mcp_server_name_mapping = {"initial-echo": "initial"}
read_rows.return_value = []
assert not await manager.catalog.list()
assert manager.tool_name_to_mcp_server_name_mapping == {}
@pytest.mark.asyncio
async def test_catalog_unchanged_read_preserves_derived_initialize_instructions(monkeypatch):
read_rows = AsyncMock(return_value=[_catalog_row()])
_catalog_database(monkeypatch, read_rows)
manager = MCPServerManager()
await manager.catalog.list()
manager._upstream_initialize_instructions_by_server_id = {"catalog-server": "cached instructions"}
manager._upstream_initialize_instructions_probed_at = {"catalog-server": 123.0}
await manager.catalog.list()
assert manager._upstream_initialize_instructions_by_server_id == {"catalog-server": "cached instructions"}
assert manager._upstream_initialize_instructions_probed_at == {"catalog-server": 123.0}
@pytest.mark.asyncio
async def test_catalog_reuses_openapi_tools_until_configuration_or_background_refresh(monkeypatch, tmp_path):
from litellm.proxy._experimental.mcp_server import openapi_to_mcp_generator, tool_registry
registry = tool_registry.MCPToolRegistry()
monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry)
spec_path = tmp_path / "catalog.json"
spec_path.write_text(json.dumps({
"openapi": "3.0.0", "info": {"title": "Catalog", "version": "1"},
"paths": {"/echo": {"get": {"operationId": "echo"}}},
}))
row = _catalog_row().model_copy(update={"spec_path": str(spec_path)})
read_rows = AsyncMock(return_value=[row])
_catalog_database(monkeypatch, read_rows)
manager = MCPServerManager()
with patch.object(
openapi_to_mcp_generator, "load_openapi_spec_async",
wraps=openapi_to_mcp_generator.load_openapi_spec_async,
) as load_spec:
await manager.catalog.list()
first_tools = tuple(registry.list_tools())
assert len(first_tools) == 1
await manager.catalog.list()
assert load_spec.await_count == 1
assert [tool.name for tool in registry.list_tools()] == [tool.name for tool in first_tools]
await manager.reload_servers_from_database()
assert load_spec.await_count == 2
read_rows.return_value = [row.model_copy(update={"updated_at": datetime(2026, 1, 2)})]
await manager.catalog.list()
assert load_spec.await_count == 3
read_rows.return_value = []
await manager.catalog.list()
assert registry.list_tools() == []

View file

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

View file

@ -136,7 +136,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},
@ -2552,7 +2552,7 @@ class TestTemporaryMCPSessionEndpoints:
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)
with (
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
@ -2629,7 +2629,7 @@ 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)
fake_proxy_server = types.SimpleNamespace(master_key=None, prisma_client=None)
with (
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),
@ -2681,7 +2681,7 @@ 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)
fake_proxy_server = types.SimpleNamespace(master_key=None, prisma_client=None)
with (
patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}),

View file

@ -1,3 +1,4 @@
from contextlib import nullcontext
import subprocess
import sys
import textwrap
@ -28,10 +29,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
@ -49,7 +51,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
@ -378,6 +380,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"))
)
@ -506,6 +509,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=[]),
@ -559,6 +563,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=[]),