diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 60dc91a69cc..1f5eb9db66f 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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[ diff --git a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py index 2c63e0a96d8..4df115727c0 100644 --- a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py new file mode 100644 index 00000000000..673a1260c09 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 64bab0a7832..56c2896ea4f 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -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, + ) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 20c114a2f3e..9806bf1be68 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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. diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index fcee3483e15..05a630c38c0 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -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(): diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index c2f7bf7d531..1c476dc98fe 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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), diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 9ad78876043..f79e6bd69c4 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -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, diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 88a2b92c680..c9425828407 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -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], diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 35e7055bbc0..a18d17ef553 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -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) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py index a77b4c8d565..327a6370580 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py @@ -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: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index c1d5cedeba0..b3e10561c8d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -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() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 7725aca1948..30297322b0c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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() == [] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 233a8cc96ba..90178414fd6 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -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())) diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 53645e62034..acb248748f5 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -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}), diff --git a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py index 83537c236a3..6058f0bbd61 100644 --- a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -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=[]),