diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260923000000_add_mcp_catalog_revision_trigger/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260923000000_add_mcp_catalog_revision_trigger/migration.sql new file mode 100644 index 00000000000..bc43353063e --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260923000000_add_mcp_catalog_revision_trigger/migration.sql @@ -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(); 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 f575265a5e7..19e7f716f81 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 @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py index 5db9a92e51d..aece755e4c4 100644 --- a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py new file mode 100644 index 00000000000..1f3cd77cdfa --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index a2aeef3f67f..3b442795c55 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 26b2d296713..f4bcc57366c 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -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, + ) diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py index c4777179256..2ea1bd062ce 100644 --- a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -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]: diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 4e7054cab98..e301aa53e50 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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. diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 42cca179cd8..fdea39bf5bc 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -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(): diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 305b8a27c92..e2807d06fbb 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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), diff --git a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py index dcf1b01bc25..0644238071e 100644 --- a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py +++ b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py @@ -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 ( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 2090a0c7421..42c2e55a1ee 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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, diff --git a/litellm/proxy/_experimental/mcp_server/tool_registry.py b/litellm/proxy/_experimental/mcp_server/tool_registry.py index e9e28c8a782..dd25f09cc57 100644 --- a/litellm/proxy/_experimental/mcp_server/tool_registry.py +++ b/litellm/proxy/_experimental/mcp_server/tool_registry.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/utils.py b/litellm/proxy/_experimental/mcp_server/utils.py index 23a6e193d5d..63dfd090f58 100644 --- a/litellm/proxy/_experimental/mcp_server/utils.py +++ b/litellm/proxy/_experimental/mcp_server/utils.py @@ -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: diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_security/mcp_security_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_security/mcp_security_guardrail.py index 35480f5e397..8482567b6da 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/mcp_security/mcp_security_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_security/mcp_security_guardrail.py @@ -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) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index e481f331e66..6d8635960a1 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -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, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index a620d9320fa..2f607342e22 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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 diff --git a/litellm/repositories/config_repository.py b/litellm/repositories/config_repository.py index c0be8ac55a9..8768d79b0dd 100644 --- a/litellm/repositories/config_repository.py +++ b/litellm/repositories/config_repository.py @@ -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: ... diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index ba8e236b405..3293851fd9f 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 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], diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index c54eda505df..5b1519bbccb 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -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. diff --git a/litellm/types/mcp_server/tool_registry.py b/litellm/types/mcp_server/tool_registry.py index df2dda27a38..738d37c9503 100644 --- a/litellm/types/mcp_server/tool_registry.py +++ b/litellm/types/mcp_server/tool_registry.py @@ -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): diff --git a/tests/integration/mcp/test_mcp_management.py b/tests/integration/mcp/test_mcp_management.py index bef882eae31..b9d55f86371 100644 --- a/tests/integration/mcp/test_mcp_management.py +++ b/tests/integration/mcp/test_mcp_management.py @@ -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()) == () diff --git a/tests/mcp_tests/test_proxy_mcp_e2e.py b/tests/mcp_tests/test_proxy_mcp_e2e.py index d22d54f0bce..21d6633f9f2 100644 --- a/tests/mcp_tests/test_proxy_mcp_e2e.py +++ b/tests/mcp_tests/test_proxy_mcp_e2e.py @@ -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": diff --git a/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 35004c8b787..1a4a01a0141 100644 --- a/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -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" diff --git a/tests/unit/proxy/_experimental/mcp_server/conftest.py b/tests/unit/proxy/_experimental/mcp_server/conftest.py index 51cab559797..90bb783ad4e 100644 --- a/tests/unit/proxy/_experimental/mcp_server/conftest.py +++ b/tests/unit/proxy/_experimental/mcp_server/conftest.py @@ -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) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py index 55456028bb0..b2364647e53 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py @@ -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: diff --git a/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py index 2e255ddf853..f57b5121fa5 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py @@ -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)] ) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 6d71256ff18..195c1ae2c74 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_block_recording.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_block_recording.py index b8aadef430f..746130fb892 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_block_recording.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_block_recording.py @@ -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): diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py index 316988ef175..941d67cb587 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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"] ) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 8daec99e7ad..4359762ba33 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -50,7 +50,7 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( _deserialize_json_dict, _flow_endpoints_missing, _mcp_oauth_discovery_on_startup_enabled, - _oauth_endpoints_unresolved, + oauth_endpoints_unresolved, _deserialize_json_list, _normalize_mcp_server_cost_info, _obo_retry_applies, @@ -877,7 +877,7 @@ class TestMCPServerManager: assert resolved[0].scopes == ["mcp.read"] assert manager.config_mcp_servers[server.server_id] is resolved[0] assert server.authorization_url is None - assert manager._oauth_discovery_slot(server.server_id) is None + assert manager.oauth_discovery_slot(server.server_id) is None @pytest.mark.asyncio async def test_table_oauth_discovery_can_be_deferred_until_first_use(self): @@ -903,7 +903,7 @@ class TestMCPServerManager: server = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) discovery.assert_not_awaited() - assert manager._oauth_discovery_slot(server.server_id) is not None + assert manager.oauth_discovery_slot(server.server_id) is not None manager.registry[server.server_id] = server with patch.object(manager, "_descovery_metadata", new=discovery): @@ -960,7 +960,7 @@ class TestMCPServerManager: assert len({id(resolution) for resolution in resolutions}) == 1 assert resolutions[0].authorization_url == "https://idp.example.com/authorize" assert resolutions[0].token_url == "https://idp.example.com/token" - assert manager._oauth_discovery_slot(server.server_id) is None + assert manager.oauth_discovery_slot(server.server_id) is None @pytest.mark.asyncio async def test_lazy_oauth_discovery_timeout_is_bounded(self): @@ -992,7 +992,7 @@ class TestMCPServerManager: assert exc.value.status_code == 503 assert "timed out" in str(exc.value.detail) discovery.assert_awaited_once_with(server) - assert manager._oauth_discovery_slot(server.server_id) is not None + assert manager.oauth_discovery_slot(server.server_id) is not None @pytest.mark.asyncio async def test_cancelling_one_waiter_does_not_cancel_shared_discovery(self): @@ -1073,7 +1073,7 @@ class TestMCPServerManager: assert replacement.token_url is None assert resolved.authorization_url == "https://idp.example.com/authorize" assert resolved.token_url == "https://idp.example.com/token" - assert manager._oauth_discovery_slot(replacement.server_id) is None + assert manager.oauth_discovery_slot(replacement.server_id) is None def test_registry_swap_reconcile_keeps_slot_for_issuer_anchored_server_without_url(self): manager = MCPServerManager() @@ -1090,9 +1090,9 @@ class TestMCPServerManager: manager.registry[server.server_id] = server manager._set_oauth_discovery_deferred(server.server_id, True) - manager._reconcile_oauth_discovery_slots_for_servers([server]) + manager.reconcile_oauth_discovery_slots_for_servers([server]) - assert manager._oauth_discovery_slot(server.server_id) is not None + assert manager.oauth_discovery_slot(server.server_id) is not None resolved = server.model_copy( update={ @@ -1101,12 +1101,12 @@ class TestMCPServerManager: } ) manager.registry[resolved.server_id] = resolved - manager._reconcile_oauth_discovery_slots_for_servers([resolved]) + manager.reconcile_oauth_discovery_slots_for_servers([resolved]) - assert manager._oauth_discovery_slot(server.server_id) is None + assert manager.oauth_discovery_slot(server.server_id) is None def _assert_oauth_discovery_state_removed(self, manager, server_id): - assert manager._oauth_discovery_slot(server_id) is None + assert manager.oauth_discovery_slot(server_id) is None @pytest.mark.asyncio async def test_deactivated_server_clears_lazy_oauth_discovery_state(self): @@ -1148,7 +1148,7 @@ class TestMCPServerManager: with ( patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", + "litellm.proxy._experimental.mcp_server.db.MCPServerRepository", return_value=repository, ), patch( @@ -1181,7 +1181,7 @@ class TestMCPServerManager: manager.registry[server.server_id] = server previous_registry = manager.registry manager._set_oauth_discovery_deferred(server.server_id, True) - old_generation = manager._oauth_discovery_slot(server.server_id).generation + old_generation = manager.oauth_discovery_slot(server.server_id).generation resolved = server.model_copy( update={ "authorization_url": "https://idp.example.com/authorize", @@ -1200,7 +1200,8 @@ class TestMCPServerManager: raw_row = MagicMock() raw_row.model_dump.return_value = row.model_dump() repository = MagicMock() - repository.table.find_many = AsyncMock(return_value=[raw_row]) + other_row: Final = row.model_copy(update={"server_id": "changed-other", "server_name": "changed_other", "auth_type": MCPAuth.none}) + repository.table.find_many = AsyncMock(return_value=[raw_row, other_row]) async def publish_while_staged(*_args, **_kwargs): assert manager.registry is previous_registry @@ -1208,7 +1209,7 @@ class TestMCPServerManager: with ( patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", + "litellm.proxy._experimental.mcp_server.db.MCPServerRepository", return_value=repository, ), patch( @@ -1217,16 +1218,16 @@ class TestMCPServerManager: ), patch.object( manager, - "_maybe_register_openapi_tools", + "maybe_register_openapi_tools", new=AsyncMock(side_effect=publish_while_staged), ), - patch.object(manager, "_prime_oauth_metadata_discovery_for_servers"), + patch.object(manager, "prime_oauth_metadata_discovery_for_servers"), ): await manager.reload_servers_from_database() assert previous_registry[server.server_id] is resolved assert manager.registry[server.server_id] is server - retry_slot = manager._oauth_discovery_slot(server.server_id) + retry_slot = manager.oauth_discovery_slot(server.server_id) assert retry_slot is not None assert retry_slot.generation > old_generation @@ -1277,7 +1278,8 @@ class TestMCPServerManager: table = SimpleNamespace( find_many=AsyncMock(return_value=[_row(cached.server_id, corrupted), _row("healthy-sibling", stored)]) ) - prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table)) + prisma = SimpleNamespace(db=SimpleNamespace(litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None)), litellm_mcpservertable=table)) + prisma.writer_db = prisma.db monkeypatch.setattr(proxy_server, "prisma_client", prisma) with caplog.at_level(logging.DEBUG, logger="LiteLLM"): @@ -1326,7 +1328,7 @@ class TestMCPServerManager: assert manager.config_mcp_servers[server.server_id].authorization_url == "https://idp.example.com/authorize" assert manager.config_mcp_servers[server.server_id].token_url is None assert manager.config_mcp_servers[server.server_id].scopes is None - assert manager._oauth_discovery_slot(server.server_id) is not None + assert manager.oauth_discovery_slot(server.server_id) is not None @pytest.mark.asyncio async def test_create_mcp_client_triggers_deferred_oauth_discovery(self): @@ -1553,7 +1555,7 @@ class TestMCPServerManager: manager = MCPServerManager() with ( patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)), - patch.object(manager, "_hydrate_config_servers_dcr_clients", new=AsyncMock()), + patch.object(manager, "hydrate_config_servers_dcr_clients", new=AsyncMock()), caplog.at_level(logging.WARNING, logger="LiteLLM"), ): await manager.load_servers_from_config(self._id_jag_config()) @@ -1569,7 +1571,7 @@ class TestMCPServerManager: manager = MCPServerManager() with ( patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)), - patch.object(manager, "_hydrate_config_servers_dcr_clients", new=AsyncMock()), + patch.object(manager, "hydrate_config_servers_dcr_clients", new=AsyncMock()), caplog.at_level(logging.WARNING, logger="LiteLLM"), ): await manager.load_servers_from_config(self._id_jag_config()) @@ -1593,7 +1595,7 @@ class TestMCPServerManager: } with ( patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)), - patch.object(manager, "_hydrate_config_servers_dcr_clients", new=AsyncMock()), + patch.object(manager, "hydrate_config_servers_dcr_clients", new=AsyncMock()), caplog.at_level(logging.WARNING, logger="LiteLLM"), ): await manager.load_servers_from_config(config) @@ -1607,7 +1609,7 @@ class TestMCPServerManager: manager = MCPServerManager() with ( patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)), - patch.object(manager, "_hydrate_config_servers_dcr_clients", new=AsyncMock()), + patch.object(manager, "hydrate_config_servers_dcr_clients", new=AsyncMock()), caplog.at_level(logging.WARNING, logger="LiteLLM"), ): await manager.load_servers_from_config(self._id_jag_config()) @@ -7358,7 +7360,7 @@ class TestMCPServerManager: manager.record_listed_tools(server, [MCPTool(name="echo", description="old", inputSchema={})], None) manager.record_listed_tools(other, [MCPTool(name="ping", description="kept", inputSchema={})], None) - manager._invalidate_server_definition_caches(server.server_id) + manager.invalidate_server_definition_caches(server.server_id) assert manager.get_listed_tool(server, "echo") is None kept = manager.get_listed_tool(other, "ping") @@ -7389,7 +7391,7 @@ class TestMCPServerManager: listing = asyncio.create_task(list_tools()) await fetch_started.wait() - manager._invalidate_server_definition_caches(server.server_id) + manager.invalidate_server_definition_caches(server.server_id) release_fetch.set() await listing @@ -7426,7 +7428,7 @@ class TestMCPServerManager: ) manager.build_mcp_server_from_table = AsyncMock(return_value=new) - manager._maybe_register_openapi_tools = register_while_a_listing_records + manager.maybe_register_openapi_tools = register_while_a_listing_records manager.prime_oauth_metadata_discovery = MagicMock() record = LiteLLM_MCPServerTable( server_id="srv", server_name="srv", url="http://new", transport=MCPTransport.http @@ -7473,7 +7475,7 @@ class TestMCPServerManager: discoverable_endpoints._OAUTH_METADATA_CACHE[metadata_key] = (time.time() + 300, {"resource": new.url}) manager.build_mcp_server_from_table = AsyncMock(return_value=new) - manager._maybe_register_openapi_tools = register_while_discovery_fills + manager.maybe_register_openapi_tools = register_while_discovery_fills manager.prime_oauth_metadata_discovery = MagicMock() record = LiteLLM_MCPServerTable( server_id="srv", server_name="srv", url="http://new", transport=MCPTransport.http @@ -8831,7 +8833,7 @@ class TestMCPServerTimestamps: with ( patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=update_mcp_server_mock), patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", + "litellm.proxy._experimental.mcp_server.db.MCPServerRepository", return_value=repo_instance, ), patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), @@ -8917,10 +8919,10 @@ class TestMCPServerTimestamps: client_id="cid", client_secret="csec", ) - assert _oauth_endpoints_unresolved(m2m_shaped) is False + assert oauth_endpoints_unresolved(m2m_shaped) is False interactive_unresolved = m2m_shaped.model_copy(update={"client_id": None, "client_secret": None}) - assert _oauth_endpoints_unresolved(interactive_unresolved) is True + assert oauth_endpoints_unresolved(interactive_unresolved) is True def test_dcr_bridge_relay_arm_needs_its_registration_endpoint(self): """A dcr_bridge server with no admin-configured client can only register callers through the @@ -8940,11 +8942,11 @@ class TestMCPServerTimestamps: token_url="https://idp.example.com/token", registration_url=None, ) - assert _oauth_endpoints_unresolved(relay_arm) is True + assert oauth_endpoints_unresolved(relay_arm) is True assert ( - _oauth_endpoints_unresolved(relay_arm.model_copy(update={"registration_url": "https://idp/reg"})) is False + oauth_endpoints_unresolved(relay_arm.model_copy(update={"registration_url": "https://idp/reg"})) is False ) - assert _oauth_endpoints_unresolved(relay_arm.model_copy(update={"client_id": "admin-client"})) is False + assert oauth_endpoints_unresolved(relay_arm.model_copy(update={"client_id": "admin-client"})) is False def test_entra_obo_without_scopes_is_unresolved(self): """entra_obo token exchange fails closed without a scope, and scopes can come from resource @@ -8960,9 +8962,9 @@ class TestMCPServerTimestamps: token_url="https://idp.example.com/token", scopes=None, ) - assert _oauth_endpoints_unresolved(entra) is True - assert _oauth_endpoints_unresolved(entra.model_copy(update={"scopes": ["api://app/.default"]})) is False - assert _oauth_endpoints_unresolved(entra.model_copy(update={"token_exchange_profile": "rfc8693"})) is False + assert oauth_endpoints_unresolved(entra) is True + assert oauth_endpoints_unresolved(entra.model_copy(update={"scopes": ["api://app/.default"]})) is False + assert oauth_endpoints_unresolved(entra.model_copy(update={"token_exchange_profile": "rfc8693"})) is False @pytest.mark.asyncio async def test_reload_fast_path_retries_unresolved_oauth_servers(self): @@ -9007,7 +9009,7 @@ class TestMCPServerTimestamps: build_mock = AsyncMock(return_value=previous_entry) with ( patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", + "litellm.proxy._experimental.mcp_server.db.MCPServerRepository", return_value=repo_instance, ), patch( @@ -9100,7 +9102,7 @@ class TestMCPServerTimestamps: def test_carry_forward_skips_when_url_or_auth_type_changed(self): """Stale endpoints from a different upstream or auth mode must not carry forward.""" from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - _carry_forward_resolved_oauth_endpoints, + carry_forward_resolved_oauth_endpoints, ) def make_server(url: str, auth_type: MCPAuth, authorization_url: Optional[str]) -> MCPServer: @@ -9116,19 +9118,19 @@ class TestMCPServerTimestamps: previous = make_server("https://old.example.com/mcp", MCPAuth.oauth2, "https://idp.example.com/authorize") url_changed = make_server("https://new.example.com/mcp", MCPAuth.oauth2, None) - _carry_forward_resolved_oauth_endpoints(new_server=url_changed, previous_server=previous) + carry_forward_resolved_oauth_endpoints(new_server=url_changed, previous_server=previous) assert url_changed.authorization_url is None auth_changed = make_server("https://old.example.com/mcp", MCPAuth.true_passthrough, None) - _carry_forward_resolved_oauth_endpoints(new_server=auth_changed, previous_server=previous) + carry_forward_resolved_oauth_endpoints(new_server=auth_changed, previous_server=previous) assert auth_changed.authorization_url is None same = make_server("https://old.example.com/mcp", MCPAuth.oauth2, None) - _carry_forward_resolved_oauth_endpoints(new_server=same, previous_server=previous) + carry_forward_resolved_oauth_endpoints(new_server=same, previous_server=previous) assert same.authorization_url == "https://idp.example.com/authorize" explicit = make_server("https://old.example.com/mcp", MCPAuth.oauth2, "https://configured.example.com/auth") - _carry_forward_resolved_oauth_endpoints(new_server=explicit, previous_server=previous) + carry_forward_resolved_oauth_endpoints(new_server=explicit, previous_server=previous) assert explicit.authorization_url == "https://configured.example.com/auth" def test_carry_forward_does_not_revive_token_url_across_authorization_url_change(self): @@ -9139,7 +9141,7 @@ class TestMCPServerTimestamps: endpoint recreates the RFC 9700 mix-up, durably, and the discovery gate alone cannot catch it because the stale endpoint comes from the registry, not from discovery.""" from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - _carry_forward_resolved_oauth_endpoints, + carry_forward_resolved_oauth_endpoints, ) previous = MCPServer( @@ -9161,7 +9163,7 @@ class TestMCPServerTimestamps: authorization_url="https://idp-b.example.com/authorize", ) - _carry_forward_resolved_oauth_endpoints(new_server=repointed, previous_server=previous) + carry_forward_resolved_oauth_endpoints(new_server=repointed, previous_server=previous) assert repointed.authorization_url == "https://idp-b.example.com/authorize" assert repointed.token_url is None @@ -9173,7 +9175,7 @@ class TestMCPServerTimestamps: consistent group, and a rebuild that re-pins the same authorize endpoint (formatting aside) keeps carrying the corroborated token endpoint.""" from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - _carry_forward_resolved_oauth_endpoints, + carry_forward_resolved_oauth_endpoints, ) def previous() -> MCPServer: @@ -9198,7 +9200,7 @@ class TestMCPServerTimestamps: auth_type=MCPAuth.oauth2, authorization_url=None, ) - _carry_forward_resolved_oauth_endpoints(new_server=blipped, previous_server=previous()) + carry_forward_resolved_oauth_endpoints(new_server=blipped, previous_server=previous()) assert blipped.authorization_url == "https://idp.example.com/authorize" assert blipped.token_url == "https://idp.example.com/token" assert blipped.registration_url == "https://idp.example.com/register" @@ -9221,7 +9223,7 @@ class TestMCPServerTimestamps: auth_type=MCPAuth.oauth2, authorization_url="https://IDP.example.com:443/authorize/", ) - _carry_forward_resolved_oauth_endpoints(new_server=same_authorize, previous_server=previous()) + carry_forward_resolved_oauth_endpoints(new_server=same_authorize, previous_server=previous()) assert same_authorize.token_url == "https://idp.example.com/token" assert same_authorize.registration_url == "https://idp.example.com/register" @@ -9233,7 +9235,7 @@ class TestMCPServerTimestamps: scopes still carry as last-known-good. Anchoring is keyed on the explicit issuer_is_anchored flag, not on issuer truthiness, so a discovered issuer does not trip this fail-closed branch.""" from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - _carry_forward_resolved_oauth_endpoints, + carry_forward_resolved_oauth_endpoints, ) previous = MCPServer( @@ -9259,7 +9261,7 @@ class TestMCPServerTimestamps: issuer_is_anchored=True, ) - _carry_forward_resolved_oauth_endpoints(new_server=failed_rebuild, previous_server=previous) + carry_forward_resolved_oauth_endpoints(new_server=failed_rebuild, previous_server=previous) assert failed_rebuild.authorization_url is None assert failed_rebuild.token_url is None @@ -9273,7 +9275,7 @@ class TestMCPServerTimestamps: regression the explicit issuer_is_anchored flag prevents: keying fail-closed on issuer truthiness alone would drop the working endpoints the moment the server learned its issuer.""" from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - _carry_forward_resolved_oauth_endpoints, + carry_forward_resolved_oauth_endpoints, ) previous = MCPServer( @@ -9300,7 +9302,7 @@ class TestMCPServerTimestamps: authorization_url=None, ) - _carry_forward_resolved_oauth_endpoints(new_server=blipped_rebuild, previous_server=previous) + carry_forward_resolved_oauth_endpoints(new_server=blipped_rebuild, previous_server=previous) assert blipped_rebuild.authorization_url == "https://idp.example.com/authorize" assert blipped_rebuild.token_url == "https://idp.example.com/token" @@ -13589,7 +13591,7 @@ class TestClientForwardedDiscoveryFailureIsNotFatal: assert resolved is manager.config_mcp_servers[server.server_id] assert resolved.authorization_url is None assert resolved.token_url is None - assert manager._oauth_discovery_slot(server.server_id) is not None + assert manager.oauth_discovery_slot(server.server_id) is not None @pytest.mark.parametrize( "auth_type, serves_the_listing", @@ -13651,7 +13653,7 @@ class TestClientForwardedDiscoveryFailureIsNotFatal: assert resolved.token_url == "https://idp.example.com/token" assert resolved.registration_url == "https://idp.example.com/register" assert manager.config_mcp_servers[server.server_id].authorization_url == "https://idp.example.com/authorize" - assert manager._oauth_discovery_slot(server.server_id) is None + assert manager.oauth_discovery_slot(server.server_id) is None class TestResolveOpenapiToolAuth: @@ -14041,7 +14043,7 @@ class TestConfigServerIdPinning: ) with ( patch( # test-quality-ok: the db reload path has no seam but its own repository - "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", + "litellm.proxy._experimental.mcp_server.db.MCPServerRepository", return_value=repository, ), patch( # test-quality-ok: same, the prisma client is fetched inside the reload @@ -14994,7 +14996,7 @@ def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: auth_type=MCPAuth.oauth2, ) manager._set_oauth_discovery_deferred(original.server_id, True) - original_slot: Final = manager._oauth_discovery_slot(original.server_id) + original_slot: Final = manager.oauth_discovery_slot(original.server_id) assert original_slot is not None replacement: Final = original.model_copy(update={"url": "https://new.example.com/mcp"}) manager.registry[original.server_id] = replacement @@ -15017,30 +15019,30 @@ async def test_temporary_oauth_discovery_expires_without_more_requests() -> None ) manager._set_oauth_discovery_deferred(server.server_id, True) resolved: Final = await manager.ensure_oauth_metadata_discovered(server) - assert manager._oauth_discovery_slot(server.server_id) is not None + assert manager.oauth_discovery_slot(server.server_id) is not None loop: Final = asyncio.get_running_loop() expired: Final = loop.create_future() with patch.object(loop, "time", return_value=loop.time() + 301): loop.call_later(0, expired.set_result, None) await expired assert resolved.authorization_url == server.authorization_url - assert manager._oauth_discovery_slot(server.server_id) is None + assert manager.oauth_discovery_slot(server.server_id) is None def test_old_temporary_discovery_expiry_preserves_replacement() -> None: manager: Final = MCPServerManager() manager._set_oauth_discovery_deferred("reused-session", True) - old_slot: Final = manager._oauth_discovery_slot("reused-session") + old_slot: Final = manager.oauth_discovery_slot("reused-session") assert old_slot is not None manager._set_oauth_discovery_deferred("reused-session", True) - replacement: Final = manager._oauth_discovery_slot("reused-session") + replacement: Final = manager.oauth_discovery_slot("reused-session") manager._expire_temporary_oauth_discovery("reused-session", old_slot.generation) - assert manager._oauth_discovery_slot("reused-session") is replacement + assert manager.oauth_discovery_slot("reused-session") is replacement assert replacement is not None manager._expire_temporary_oauth_discovery("reused-session", replacement.generation) - assert manager._oauth_discovery_slot("reused-session") is None + assert manager.oauth_discovery_slot("reused-session") is None manager._expire_temporary_oauth_discovery("reused-session", replacement.generation) - assert manager._oauth_discovery_slot("reused-session") is None + assert manager.oauth_discovery_slot("reused-session") is None @pytest.mark.asyncio @@ -15408,12 +15410,12 @@ async def test_discovery_cache_invalidation_during_fetch_does_not_repopulate_old with _mcp_upstream(upstream.respond): task: Final = asyncio.create_task(manager.get_prompts_from_server(_discovery_server(), None)) await asyncio.wait_for(upstream.entered.wait(), timeout=5) - manager._invalidate_discovery_lists("discovery") + manager.invalidate_discovery_lists("discovery") upstream.release.set() assert (await task)[0].name == "discovery-example" assert len(await manager.get_prompts_from_server(_discovery_server(), None)) == 1 assert upstream.initializes == 2 - manager._invalidate_discovery_lists("discovery") + manager.invalidate_discovery_lists("discovery") assert len(await manager.get_prompts_from_server(_discovery_server(), None)) == 1 assert upstream.initializes == 3 @@ -15769,6 +15771,7 @@ class TestProtectedCredentialPreparation: credential: str | None, dispatch: str, ) -> None: + from litellm.proxy._experimental.mcp_server import operations from litellm.proxy._experimental.mcp_server.server import _handle_local_mcp_tool from litellm.proxy._experimental.mcp_server.utils import add_server_prefix_to_name, get_server_prefix @@ -15791,6 +15794,8 @@ class TestProtectedCredentialPreparation: authentication_token=credential, ) manager: Final = MCPServerManager() + manager.config_mcp_servers = {server.server_id: server} + monkeypatch.setattr(operations, "global_mcp_server_manager", manager) await manager._register_openapi_tools(str(spec_path), server, server.url) monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="unexpected success") @@ -16475,6 +16480,1170 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie 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, read_revision=None): + from types import SimpleNamespace + + from litellm.proxy import proxy_server + + if read_revision is None: + read_revision = AsyncMock(return_value=None) + client = SimpleNamespace( + db=SimpleNamespace( + litellm_mcpservertable=SimpleNamespace(find_many=read_rows), + litellm_config=SimpleNamespace(find_unique=read_revision), + ) + ) + client.writer_db = client.db + 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 +@pytest.mark.parametrize("failed", [False, True]) +async def test_catalog_coalesces_waiters_only_behind_a_read_started_after_their_arrival(monkeypatch, failed): + started = asyncio.Event() + release = asyncio.Event() + old_row = _catalog_row() + new_row = _catalog_row("updated") + read_rows = AsyncMock(return_value=[old_row]) + _catalog_database(monkeypatch, read_rows) + manager = MCPServerManager() + previous = await manager.catalog.list() + read_rows.reset_mock() + + async def read_current_rows(**kwargs): + first = not started.is_set() + if first: + started.set() + await release.wait() + if failed: + raise RuntimeError("unavailable database") + return [old_row if first else new_row] + + read_rows.side_effect = read_current_rows + first = asyncio.create_task(manager.catalog.list()) + await started.wait() + followers = [asyncio.create_task(manager.catalog.list()) for _ in range(8)] + await asyncio.sleep(0) + release.set() + results = await asyncio.gather(first, *followers, return_exceptions=True) + if failed: + assert all(isinstance(result, HTTPException) and result.status_code == 503 for result in results) + assert manager.registry == dict(previous) + else: + assert results[0]["catalog-server"].name == "initial" + assert all(result["catalog-server"].name == "updated" for result in results[1:]) + assert read_rows.await_count == 2 + + read_rows.side_effect = None + read_rows.return_value = [new_row] + assert (await manager.catalog.list())["catalog-server"].name == "updated" + assert read_rows.await_count == 3 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("interruption", ["background", "cancel_reader", "cancel_waiter"]) +async def test_catalog_queued_refresh_preserves_freshness_through_cancellation(monkeypatch, interruption): + started = asyncio.Event() + release = asyncio.Event() + + async def read_current_rows(**kwargs): + first = not started.is_set() + if first: + started.set() + await release.wait() + return [_catalog_row("initial" if first else "updated")] + + read_rows = AsyncMock(side_effect=read_current_rows) + _catalog_database(monkeypatch, read_rows) + manager = MCPServerManager() + first = asyncio.create_task(manager.catalog.list()) + await started.wait() + background = asyncio.create_task(manager.reload_servers_from_database()) if interruption == "background" else None + waiter = asyncio.create_task(manager.catalog.list()) + survivors = [asyncio.create_task(manager.catalog.list()) for _ in range(4)] + await asyncio.sleep(0) + if interruption == "cancel_reader": + first.cancel() + elif interruption == "cancel_waiter": + waiter.cancel() + release.set() + first_result, waiter_result, *results = await asyncio.gather(first, waiter, *survivors, return_exceptions=True) + if background is not None: + await background + if interruption == "cancel_reader": + assert isinstance(first_result, asyncio.CancelledError) + else: + assert first_result["catalog-server"].name == "initial" + if interruption == "cancel_waiter": + assert isinstance(waiter_result, asyncio.CancelledError) + else: + assert waiter_result["catalog-server"].name == "updated" + assert all(result["catalog-server"].name == "updated" for result in results) + assert read_rows.await_count == 2 + 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") == 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() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("rediscover", [False, True]) +async def test_catalog_publishes_rediscovered_routes_despite_concurrent_owner_change(monkeypatch, rediscover): + from mcp.types import Tool + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "prisma_client", None) + manager = MCPServerManager() + selected = MCPServer(server_id="selected", name="selected", transport=MCPTransport.http) + competing = MCPServer(server_id="competing", name="competing", transport=MCPTransport.http) + manager.config_mcp_servers = {server.server_id: server for server in (selected, competing)} + manager.published_tool_routes = {"shared_tool": selected.name} + async with manager.catalog.operation(): + manager.published_tool_routes["shared_tool"] = competing.name + if rediscover: + manager._create_prefixed_tools([Tool(name="shared_tool", input_schema={})], selected) + assert manager._get_mcp_server_from_tool_name("shared_tool").server_id == selected.server_id + async with manager.catalog.operation(): + expected = selected if rediscover else competing + assert manager._get_mcp_server_from_tool_name("shared_tool").server_id == expected.server_id + + +@pytest.mark.asyncio +@pytest.mark.parametrize("anchored", [False, True]) +async def test_catalog_reload_preserves_concurrent_config_discovery_and_routes(anchored): + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="config-race", name="config_race", transport=MCPTransport.http, + url="https://upstream.example/mcp", auth_type=MCPAuth.oauth2, + issuer="https://issuer.example" if anchored else None, issuer_is_anchored=anchored) + manager.config_mcp_servers = {server.server_id: server} + manager._set_oauth_discovery_deferred(server.server_id, True) + generation: Final = manager.oauth_discovery_slot(server.server_id).generation + resolved: Final = server.model_copy(update={"authorization_url": "https://issuer.example/authorize", + "token_url": "https://issuer.example/token", "scopes": ["read"], "issuer": "https://issuer.example"}) + + async def hydrate(target: MCPServer) -> bool: + target.client_id = "persisted-client" + return True + + async def read_rows(**_kwargs): + assert manager._publish_resolved_oauth_server(resolved, generation) is resolved + manager.published_tool_routes["config_race-search"] = "config_race" + return [] + + prisma: Final = MagicMock() + prisma.writer_db = prisma.db + prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=read_rows) + prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), + patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.hydrate_config_server_dcr_client", side_effect=hydrate), + ): + await manager.reload_servers_from_database() + current: Final = manager.get_mcp_server_by_id(server.server_id) + assert current.authorization_url == resolved.authorization_url + assert current.token_url == resolved.token_url + assert current.scopes == ["read"] + assert current.issuer == "https://issuer.example" + assert current.client_id == "persisted-client" + assert manager.oauth_discovery_slot(server.server_id) is None + assert manager.published_tool_routes["config_race-search"] == "config_race" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", ["credentials", "delete"]) +async def test_catalog_reload_does_not_restore_replaced_config_credentials_or_routes(change): + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="config-replaced", name="config_replaced", transport=MCPTransport.http, + client_id="previous-client") + manager.config_mcp_servers = {server.server_id: server} + manager.published_tool_routes["config_replaced-search"] = "config_replaced" + + async def read_rows(**_kwargs): + manager.config_mcp_servers = {} if change == "delete" else { + server.server_id: server.model_copy(update={"client_id": "new-client"})} + return [] + + prisma: Final = MagicMock() + prisma.writer_db = prisma.db + prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=read_rows) + prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + await manager.reload_servers_from_database() + current: Final = manager.get_mcp_server_by_id(server.server_id) + if change == "delete": + assert current is None + else: + assert current.client_id == "new-client" + assert "config_replaced-search" not in manager.published_tool_routes + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", ["discovery", "update", "delete"]) +async def test_catalog_operation_retains_routes_only_for_same_configured_target(change): + from mcp.types import Tool + + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="route-race", name="route_race", transport=MCPTransport.http, + url="https://upstream.example/mcp", auth_type=MCPAuth.oauth2) + manager.registry = {server.server_id: server} + with patch("litellm.proxy.proxy_server.prisma_client", None): + async with manager.catalog.operation(): + tools: Final = manager._create_prefixed_tools([Tool(name="search", inputSchema={})], server) + if change == "delete": + manager.registry = {} + else: + manager.registry[server.server_id] = server.model_copy(update=( + {"authorization_url": "https://issuer.example/authorize", "token_url": "https://issuer.example/token"} + if change == "discovery" else {"url": "https://changed.example/mcp"})) + assert (tools[0].name in manager.published_tool_routes) is (change == "discovery") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", ["concurrent", "staged", "deleted", "staged_after_delete", "staged_with_concurrent"]) +@pytest.mark.parametrize("mapped", [False, True]) +async def test_catalog_reload_keeps_new_route_owner_over_earlier_route(change, mapped, monkeypatch): + from litellm.proxy._experimental.mcp_server import tool_registry + + manager: Final = MCPServerManager() + registry: Final = tool_registry.MCPToolRegistry() + monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry) + servers: Final = {name: MCPServer(server_id=name, name=name, transport=MCPTransport.http, + client_id="configured-client") for name in ("old_owner", "new_owner", "concurrent_owner")} + manager.config_mcp_servers = servers + manager.published_tool_routes = {"shared-search": "old_owner"} if mapped else {} + + async def original_handler(): + return "original response" + + async def new_handler(): + return "new response" + + async def concurrent_handler(): + return "concurrent response" + + registry.register_tool("shared-search", "Search", {}, original_handler) + original_tool: Final = registry.get_tool("shared-search") + + async def read_rows(**_kwargs): + if change == "concurrent": + manager.published_tool_routes["shared-search"] = "new_owner" + registry.published_tools["shared-search"] = original_tool.model_copy(update={"handler": new_handler}) + elif change in ("staged", "staged_after_delete", "staged_with_concurrent"): + if change == "staged_after_delete": + manager.published_tool_routes.clear() + registry.published_tools.clear() + elif change == "staged_with_concurrent": + manager.published_tool_routes["shared-search"] = "concurrent_owner" + registry.published_tools["shared-search"] = original_tool.model_copy(update={"handler": concurrent_handler}) + manager.tool_name_to_mcp_server_name_mapping["shared-search"] = "new_owner" + registry.register_tool("shared-search", "Search", {}, new_handler) + else: + manager.published_tool_routes.clear() + registry.published_tools.clear() + if not mapped: + manager.published_tool_routes.clear() + manager.tool_name_to_mcp_server_name_mapping.clear() + return [] + + prisma: Final = MagicMock() + prisma.writer_db = prisma.db + prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=read_rows) + prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + await manager.reload_servers_from_database() + if change == "deleted" or not mapped: + assert "shared-search" not in manager.published_tool_routes + else: + assert manager.published_tool_routes["shared-search"] == "new_owner" + if change == "deleted": + assert registry.get_tool("shared-search") is None + else: + assert await registry.get_tool("shared-search").handler() == "new response" + + +@pytest.mark.asyncio +async def test_catalog_observes_committed_update_and_delete_without_background_reload(): + from datetime import timedelta + + timestamp: Final = datetime.now() + row: Final = LiteLLM_MCPServerTable( + server_id="catalog-server", server_name="catalog_server", alias="catalog_server", + transport=MCPTransport.http, url="https://first.example.com/mcp", updated_at=timestamp, + ) + updated: Final = row.model_copy(update={"url": "https://second.example.com/mcp", "updated_at": timestamp + timedelta(seconds=1)}) + prisma: Final = MagicMock() + prisma.writer_db = prisma.db + prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=([row], [updated], [])) + prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) + manager: Final = MCPServerManager() + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), + patch("litellm.proxy.proxy_server.general_settings", {}), + ): + first = await manager.catalog.resolve(row.server_id) + manager.registry[row.server_id].short_prefix = "a12" + second = await manager.catalog.resolve(row.server_id) + deleted = await manager.catalog.resolve(row.server_id) + + assert first is not None and first.url == row.url + assert second is not None and second.url == updated.url + assert second.short_prefix == "a12" + assert deleted is None + assert prisma.db.litellm_mcpservertable.find_many.await_count == 3 + + +@pytest.mark.asyncio +async def test_catalog_rebuilt_unchanged_server_keeps_discovered_tool_routes(): + from mcp.types import Tool + + row: Final = LiteLLM_MCPServerTable(server_id="unchanged-routes", alias="unchanged_routes", + url="https://upstream.example/mcp", transport=MCPTransport.http) + manager: Final = MCPServerManager() + 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) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + await manager.reload_servers_from_database() + server: Final = manager.get_mcp_server_by_id(row.server_id) + tools: Final = manager._create_prefixed_tools([Tool(name="search", inputSchema={})], server) + await manager.reload_servers_from_database() + assert manager.server_exposes_tool(manager.get_mcp_server_by_id(row.server_id), tools[0].name) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", ["same", "changed", "deleted"]) +@pytest.mark.parametrize("openapi", [False, True]) +async def test_catalog_reload_retains_routes_discovered_for_a_late_server(change, openapi, monkeypatch): + from datetime import timedelta + + from mcp.types import Tool + from litellm.proxy._experimental.mcp_server import tool_registry + + row: Final = LiteLLM_MCPServerTable(server_id="late-routes", alias="late_routes", + url="https://upstream.example/mcp", transport=MCPTransport.http, updated_at=datetime.now(), + spec_path="late-openapi.json" if openapi else None) + manager: Final = MCPServerManager() + registry: Final = tool_registry.MCPToolRegistry() + monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry) + server: Final = await manager.build_mcp_server_from_table(row) + start: Final = asyncio.Event() + published: Final = asyncio.Event() + + async def handler(): + return "late server response" + + async def publish(): + await start.wait() + manager.registry[server.server_id] = server + manager._create_prefixed_tools([Tool(name="search", inputSchema={})], server) + if openapi: + registry.register_tool("late_routes-search", "Search", {}, handler) + published.set() + + async def read_rows(**_kwargs): + start.set() + await published.wait() + if change == "deleted": + return [] + return [row if change == "same" else row.model_copy(update={ + "url": "https://updated.example/mcp", "spec_path": None, + "updated_at": row.updated_at + timedelta(seconds=1)})] + + prisma: Final = MagicMock() + prisma.writer_db = prisma.db + prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=read_rows) + prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) + task: Final = asyncio.create_task(publish()) + try: + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + await manager.catalog.list() + finally: + task.cancel() + await asyncio.gather(task, return_exceptions=True) + if change == "deleted": + assert manager.get_mcp_server_by_id(row.server_id) is None + assert "late_routes-search" not in manager.published_tool_routes + else: + assert manager.server_exposes_tool(manager.get_mcp_server_by_id(row.server_id), "late_routes-search") is (change == "same") + tool: Final = registry.get_tool("late_routes-search") + if openapi and change == "same": + assert tool is not None + assert await tool.handler() == "late server response" + else: + assert tool is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("retain_operation", [False, True]) +async def test_catalog_openapi_refresh_does_not_restore_removed_operations(tmp_path, monkeypatch, respx_mock, retain_operation): + from litellm.proxy._experimental.mcp_server import tool_registry + + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + manager: Final = MCPServerManager() + registry: Final = tool_registry.MCPToolRegistry() + monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry) + spec_path: Final = tmp_path / "openapi.json" + paths: Final = {"/removed": {"get": {"operationId": "removed"}}, "/retained": {"get": {"operationId": "retained"}}} + spec: Final = {"openapi": "3.0.0", "info": {"title": "Refresh", "version": "1"}, "paths": paths} + spec_path.write_text(json.dumps(spec)) + row: Final = LiteLLM_MCPServerTable(server_id="spec-refresh", alias="spec_refresh", + url="https://upstream.example", transport=MCPTransport.http, spec_path=str(spec_path)) + 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) + upstream: Final = respx_mock.get("https://upstream.example/retained").respond(200, json={"value": "retained"}) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + await manager.reload_servers_from_database() + assert manager.server_exposes_tool(manager.registry[row.server_id], "spec_refresh-removed") + assert registry.get_tool("spec_refresh-removed") is not None + paths.pop("/removed") + if not retain_operation: + paths.clear() + spec_path.write_text(json.dumps(spec)) + for _ in range(2): + await manager.reload_servers_from_database() + for name in ("removed", "spec_refresh-removed"): + assert name not in manager.published_tool_routes + assert not manager.server_exposes_tool(manager.registry[row.server_id], name) + assert registry.get_tool("spec_refresh-removed") is None + retained = registry.get_tool("spec_refresh-retained") + if retain_operation: + assert manager.server_exposes_tool(manager.registry[row.server_id], "spec_refresh-retained") + assert retained is not None + from litellm.proxy._experimental.mcp_server.tool_outcome import JsonResult + + retained_result = await retained.handler() + assert isinstance(retained_result, JsonResult) + assert retained_result.value == {"value": "retained"} + else: + assert retained is None + assert manager.published_tool_routes == {} + assert upstream.call_count == (2 if retain_operation else 0) + + +@pytest.mark.asyncio +async def test_catalog_lookup_uses_one_snapshot_until_operation_finishes(): + from datetime import timedelta + + row: Final = LiteLLM_MCPServerTable( + server_id="snapshot-server", alias="snapshot_server", transport=MCPTransport.http, + url="https://first.example.com/mcp", updated_at=datetime.now(), + ) + updated: Final = row.model_copy(update={"url": "https://second.example.com/mcp", "updated_at": row.updated_at + timedelta(seconds=1)}) + prisma: Final = MagicMock() + prisma.writer_db = prisma.db + prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=([row], [updated], [updated])) + prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) + manager: Final = MCPServerManager() + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), + patch("litellm.proxy.proxy_server.general_settings", {}), + ): + async with manager.catalog.operation() as before: + await manager.reload_servers_from_database() + during = await manager.catalog.resolve(row.server_id) + assert during is not None and during.url == row.url + assert manager.registry[row.server_id].url == updated.url + async with manager.catalog.operation() as after: + current = manager.get_mcp_server_by_id(row.server_id) + assert current is not None and current.url == updated.url + assert before.identity != after.identity + assert prisma.db.litellm_mcpservertable.find_many.await_count == 3 + + +@pytest.mark.asyncio +async def test_catalog_failed_reload_preserves_published_discovery_state(): + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="healthy", name="healthy", transport=MCPTransport.http) + manager.registry = {server.server_id: server} + manager._upstream_initialize_instructions_by_server_id[server.server_id] = "healthy instructions" + prisma: Final = MagicMock() + prisma.writer_db = prisma.db + prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=RuntimeError("database unavailable")) + prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + with pytest.raises(RuntimeError, match="database unavailable"): + await manager.reload_servers_from_database() + assert manager.get_mcp_server_by_id(server.server_id) is server + assert manager._upstream_initialize_instructions_by_server_id == {server.server_id: "healthy instructions"} + + +@pytest.mark.asyncio +async def test_catalog_cancellation_retains_state_and_releases_refresh_lock(): + entered: Final = asyncio.Event() + release: Final = asyncio.Event() + + async def blocked_read(**kwargs: object) -> list[LiteLLM_MCPServerTable]: + entered.set() + await release.wait() + return [] + + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="healthy", name="healthy", transport=MCPTransport.http) + manager.registry = {server.server_id: server} + manager._upstream_initialize_instructions_by_server_id[server.server_id] = "healthy instructions" + prisma: Final = MagicMock() + prisma.writer_db = prisma.db + prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=blocked_read) + prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + pending = asyncio.create_task(manager.reload_servers_from_database()) + await asyncio.wait_for(entered.wait(), timeout=2) + pending.cancel() + with pytest.raises(asyncio.CancelledError): + await pending + assert manager.registry == {server.server_id: server} + assert manager._upstream_initialize_instructions_by_server_id == {server.server_id: "healthy instructions"} + release.set() + await asyncio.wait_for(manager.reload_servers_from_database(), timeout=2) + assert manager.registry == {} + + +@pytest.mark.asyncio +async def test_catalog_background_lookup_after_operation_exit_observes_deletion(): + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="expired-snapshot", name="expired_snapshot", transport=MCPTransport.http) + manager.registry = {server.server_id: server} + release: Final = asyncio.Event() + + async def lookup_after_exit() -> MCPServer | None: + await release.wait() + return await manager.catalog.resolve(server.server_id) + + with patch("litellm.proxy.proxy_server.prisma_client", None): + async with manager.catalog.operation(): + pending: Final = asyncio.create_task(lookup_after_exit()) + manager.registry = {} + release.set() + assert await pending is None + + +@pytest.mark.asyncio +async def test_catalog_cancelled_openapi_refresh_retains_tools_and_discovery(monkeypatch): + from datetime import timedelta + from litellm.proxy._experimental.mcp_server import tool_registry + + manager: Final = MCPServerManager() + stamp: Final = datetime.now() + server: Final = MCPServer(server_id="staged", name="staged", transport=MCPTransport.http, + url="https://before.example/mcp", spec_path="before.json", updated_at=stamp, auth_type=MCPAuth.oauth2) + manager.registry = {server.server_id: server} + manager._set_oauth_discovery_deferred(server.server_id, True) + original_slot: Final = manager.oauth_discovery_slot(server.server_id) + manager._upstream_initialize_instructions_by_server_id[server.server_id] = "keep instructions" + manager.tool_name_to_mcp_server_name_mapping = {"staged-existing": "staged"} + registry: Final = tool_registry.MCPToolRegistry() + registry.register_tool("staged-existing", "existing", {}, lambda: "existing") + monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry) + row: Final = LiteLLM_MCPServerTable(server_id=server.server_id, alias="staged", transport=MCPTransport.http, + url="https://after.example/mcp", spec_path="after.json", updated_at=stamp + timedelta(seconds=1), auth_type=MCPAuth.oauth2) + 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) + + async def cancelled_registration(*args: object, **kwargs: object) -> None: + registry.register_tool("staged-new", "new", {}, lambda: "new") + raise asyncio.CancelledError() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), + patch.object(manager, "maybe_register_openapi_tools", side_effect=cancelled_registration), + patch.object(manager, "invalidate_discovery_lists") as invalidate, + ): + with pytest.raises(asyncio.CancelledError): + await manager.reload_servers_from_database() + invalidate.assert_not_called() + assert manager.registry == {server.server_id: server} + assert [tool.name for tool in registry.list_tools()] == ["staged-existing"] + assert manager.tool_name_to_mcp_server_name_mapping == {"staged-existing": "staged"} + assert manager._upstream_initialize_instructions_by_server_id == {server.server_id: "keep instructions"} + + assert manager.oauth_discovery_slot(server.server_id) is original_slot + + +@pytest.mark.asyncio +async def test_catalog_snapshot_identity_is_independent_of_worker_oauth_discovery(): + row: Final = LiteLLM_MCPServerTable(server_id="identity-server", alias="identity_server", transport=MCPTransport.http, + url="https://upstream.example/mcp", auth_type=MCPAuth.oauth2, updated_at=datetime.now()) + 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) + manager: Final = MCPServerManager() + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), + patch("litellm.proxy.proxy_server.general_settings", {}), + patch.object(manager, "prime_oauth_metadata_discovery_for_servers"), + ): + async with manager.catalog.operation() as unresolved: + assert manager.get_mcp_server_by_id(row.server_id).authorization_url is None + resolved: Final = manager.registry[row.server_id].model_copy(update={ + "authorization_url": "https://idp.example/authorize", "token_url": "https://idp.example/token"}) + manager.registry[row.server_id] = resolved + async with manager.catalog.operation() as discovered: + assert manager.get_mcp_server_by_id(row.server_id).authorization_url == resolved.authorization_url + assert unresolved.identity == discovered.identity + + +@pytest.mark.asyncio +async def test_catalog_failed_openapi_row_does_not_publish_partial_handlers(monkeypatch): + from litellm.proxy._experimental.mcp_server import tool_registry + + manager: Final = MCPServerManager() + registry: Final = tool_registry.MCPToolRegistry() + monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry) + row: Final = LiteLLM_MCPServerTable(server_id="broken", alias="broken", transport=MCPTransport.http, + url="https://upstream.example/mcp", spec_path="broken.json") + 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) + + async def broken_registration(server: MCPServer, **kwargs: object) -> None: + registry.register_tool("broken-partial", "partial", {}, lambda: "must not run") + manager.tool_name_to_mcp_server_name_mapping["broken-partial"] = "broken" + raise ValueError("invalid remaining operation") + + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), + patch.object(manager, "maybe_register_openapi_tools", side_effect=broken_registration), + ): + await manager.reload_servers_from_database() + assert manager.registry == {} + assert registry.list_tools() == [] + 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() == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", ["edit", "delete"]) +async def test_catalog_rejects_configuration_switch_before_client_creation(monkeypatch, change): + read_rows = AsyncMock(return_value=[_catalog_row()]) + _catalog_database(monkeypatch, read_rows) + manager = MCPServerManager() + async with manager.catalog.operation(): + admitted = manager.get_mcp_server_by_id("catalog-server") + read_rows.return_value = [_catalog_row("updated")] if change == "edit" else [] + await asyncio.create_task(manager.reload_servers_from_database()) + with patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as factory: + with pytest.raises(HTTPException) as exc: + await manager._create_mcp_client(admitted) + assert exc.value.status_code == 503 + assert exc.value.detail == "MCP server configuration changed; retry the operation" + factory.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("dispatch", ["managed", "local"]) +async def test_catalog_changed_openapi_handler_never_dispatches(monkeypatch, tmp_path, respx_mock, dispatch): + from litellm.proxy._experimental.mcp_server import operations, tool_registry + from litellm.proxy._experimental.mcp_server.utils import add_server_prefix_to_name, get_server_prefix + + registry = tool_registry.MCPToolRegistry() + monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry) + monkeypatch.setattr(operations, "global_mcp_tool_registry", registry) + spec_path = tmp_path / "catalog-race.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() + monkeypatch.setattr(operations, "global_mcp_server_manager", manager) + destination = respx_mock.get("https://changed.example/echo").respond(200, text="must not execute") + async with manager.catalog.operation(): + admitted = manager.get_mcp_server_by_id("catalog-server") + read_rows.return_value = [row.model_copy(update={ + "url": "https://changed.example", "updated_at": datetime(2026, 1, 2), + })] + await asyncio.create_task(manager.reload_servers_from_database()) + with pytest.raises(HTTPException) as exc: + if dispatch == "managed": + await manager._call_openapi_tool_handler(admitted, "echo", {}) + else: + await operations._handle_local_mcp_tool( + add_server_prefix_to_name("echo", get_server_prefix(admitted)), {} + ) + assert exc.value.status_code == 503 + assert destination.call_count == 0 + + +@pytest.mark.asyncio +async def test_catalog_rejects_old_admission_in_a_new_operation(monkeypatch): + read_rows = AsyncMock(return_value=[_catalog_row()]) + _catalog_database(monkeypatch, read_rows) + manager = MCPServerManager() + original = (await manager.catalog.list())["catalog-server"] + read_rows.return_value = [_catalog_row("updated")] + async with manager.catalog.operation(): + with pytest.raises(HTTPException) as exc: + await manager._create_mcp_client(original) + assert exc.value.detail == "MCP server configuration changed; retry the operation" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", [False, True]) +async def test_catalog_deferred_discovery_preserves_admitted_configuration(monkeypatch, change): + row = _catalog_row().model_copy(update={"auth_type": MCPAuth.true_passthrough}) + read_rows = AsyncMock(return_value=[row]) + _catalog_database(monkeypatch, read_rows) + manager = MCPServerManager() + entered = asyncio.Event() + release = asyncio.Event() + metadata = MCPOAuthMetadata( + authorization_url="https://issuer.example/authorize", token_url="https://issuer.example/token", + ) + + async def discover(server): + entered.set() + await release.wait() + return metadata + + with patch.object(manager, "_discover_oauth_metadata_for_server", side_effect=discover): + async with manager.catalog.operation(): + admitted = manager.get_mcp_server_by_id("catalog-server") + resolving = asyncio.create_task(manager.ensure_oauth_metadata_discovered(admitted)) + await asyncio.wait_for(entered.wait(), timeout=1) + if change: + read_rows.return_value = [row.model_copy(update={"updated_at": datetime(2026, 1, 2)})] + await manager.reload_servers_from_database() + release.set() + if change: + with pytest.raises(HTTPException) as exc: + await asyncio.wait_for(resolving, timeout=1) + assert exc.value.detail == "MCP server configuration changed; retry the operation" + else: + resolved = await asyncio.wait_for(resolving, timeout=1) + assert resolved.authorization_url == metadata.authorization_url + assert resolved.token_url == metadata.token_url + assert resolved.updated_at == admitted.updated_at + + +async def test_catalog_cancelled_config_hydration_preserves_published_credentials(): + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="config-hydration", name="config_hydration", transport=MCPTransport.http, + client_id="previous-client") + manager.config_mcp_servers = {server.server_id: server} + + async def cancelled_hydration(target: MCPServer) -> bool: + target.client_id = "unpublished-client" + raise asyncio.CancelledError() + + with patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.hydrate_config_server_dcr_client", side_effect=cancelled_hydration): + with pytest.raises(asyncio.CancelledError): + await manager.reload_servers_from_database() + assert manager.get_mcp_server_by_id(server.server_id).client_id == "previous-client" + + +@pytest.mark.asyncio +async def test_catalog_oauth_resolution_cannot_replace_an_operation_target_after_update(): + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="pinned-oauth", name="pinned_oauth", transport=MCPTransport.http, + url="https://before.example/mcp", auth_type=MCPAuth.oauth2, + authorization_url="https://before.example/authorize", token_url="https://before.example/token") + manager.registry = {server.server_id: server} + with patch("litellm.proxy.proxy_server.prisma_client", None): + async with manager.catalog.operation(): + selected: Final = manager.get_mcp_server_by_id(server.server_id) + manager.registry[server.server_id] = server.model_copy(update={ + "url": "https://after.example/mcp", "authorization_url": "https://after.example/authorize"}) + resolved: Final = await manager.ensure_oauth_metadata_discovered(selected) + assert resolved.url == "https://before.example/mcp" + assert resolved.authorization_url == "https://before.example/authorize" + assert manager.registry[server.server_id].url == "https://after.example/mcp" + + +@pytest.mark.asyncio +async def test_catalog_unresolved_oauth_snapshot_fails_closed_if_target_was_replaced(): + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="replaced-oauth", name="replaced_oauth", transport=MCPTransport.http, + url="https://before.example/mcp", auth_type=MCPAuth.oauth2) + manager.registry = {server.server_id: server} + with ( + patch("litellm.proxy.proxy_server.prisma_client", None), + patch.object(manager, "_discover_oauth_metadata_for_server", new_callable=AsyncMock) as discovery, + ): + async with manager.catalog.operation(): + selected: Final = manager.get_mcp_server_by_id(server.server_id) + manager.registry[server.server_id] = server.model_copy(update={ + "url": "https://after.example/mcp", "authorization_url": "https://after.example/authorize", + "token_url": "https://after.example/token"}) + with pytest.raises(HTTPException) as exc: + await manager.ensure_oauth_metadata_discovered(selected) + assert exc.value.status_code == 503 + discovery.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_catalog_unresolved_oauth_snapshot_accepts_discovery_for_the_same_target(): + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="same-oauth", name="same_oauth", transport=MCPTransport.http, + url="https://upstream.example/mcp", auth_type=MCPAuth.oauth2) + manager.registry = {server.server_id: server} + discovered: Final = server.model_copy(update={"authorization_url": "https://issuer.example/authorize", + "token_url": "https://issuer.example/token", "scopes": ["read"]}) + with ( + patch("litellm.proxy.proxy_server.prisma_client", None), + patch.object(manager, "_ensure_oauth_metadata_discovered", return_value=discovered) as discovery, + ): + async with manager.catalog.operation(): + resolved: Final = await manager.ensure_oauth_metadata_discovered(server) + assert resolved.authorization_url == "https://issuer.example/authorize" + assert resolved.token_url == "https://issuer.example/token" + assert resolved.scopes == ["read"] + discovery.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_catalog_fresh_lookup_does_not_fall_back_to_stale_grants_when_database_fails(): + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="stale-grant", name="stale_grant", transport=MCPTransport.http, + allow_all_keys=True) + manager.registry = {server.server_id: server} + prisma: Final = MagicMock() + prisma.writer_db = prisma.db + prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=RuntimeError("unavailable")) + prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), + patch("litellm.proxy.proxy_server.general_settings", {}), + ): + with pytest.raises(HTTPException) as error: + await manager.catalog.resolve(server.server_id) + assert error.value.status_code == 503 + assert "unavailable" not in error.value.detail + assert manager.registry == {server.server_id: server} + + +@pytest.mark.asyncio +async def test_catalog_cancelled_registration_does_not_publish_partial_handlers(monkeypatch): + from litellm.proxy._experimental.mcp_server import tool_registry + + manager = MCPServerManager() + registry = tool_registry.MCPToolRegistry() + registry.register_tool("existing", "existing", {}, lambda: "existing") + monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry) + _catalog_database(monkeypatch, AsyncMock(return_value=[_catalog_row()])) + + async def cancel_after_registration(*args, **kwargs): + registry.register_tool("partial", "partial", {}, lambda: "partial") + raise asyncio.CancelledError() + + monkeypatch.setattr(manager, "maybe_register_openapi_tools", cancel_after_registration) + with pytest.raises(asyncio.CancelledError): + await manager.reload_servers_from_database() + assert [tool.name for tool in registry.list_tools()] == ["existing"] + assert manager.registry == {} + + + +def _revision_row(revision): + from types import SimpleNamespace + + return SimpleNamespace(reload_revision=revision) + + +@pytest.mark.asyncio +async def test_catalog_unchanged_revision_skips_table_reload(monkeypatch): + read_rows = AsyncMock(side_effect=([_catalog_row()], [_catalog_row("updated")])) + read_revision = AsyncMock(return_value=_revision_row(7)) + _catalog_database(monkeypatch, read_rows, read_revision) + manager = MCPServerManager() + assert (await manager.catalog.list())["catalog-server"].name == "initial" + assert (await manager.catalog.list())["catalog-server"].name == "initial" + read_rows.assert_awaited_once() + assert read_revision.await_count == 2 + + +@pytest.mark.asyncio +async def test_catalog_changed_revision_reloads_table(monkeypatch): + read_rows = AsyncMock(side_effect=([_catalog_row()], [_catalog_row("updated")])) + read_revision = AsyncMock(side_effect=[_revision_row(7), _revision_row(8)]) + _catalog_database(monkeypatch, read_rows, read_revision) + manager = MCPServerManager() + assert (await manager.catalog.list())["catalog-server"].name == "initial" + assert (await manager.catalog.list())["catalog-server"].name == "updated" + assert read_rows.await_count == 2 + assert read_revision.await_count == 2 + + +@pytest.mark.asyncio +async def test_catalog_absent_revision_row_reloads_table_every_operation(monkeypatch): + read_rows = AsyncMock(side_effect=([_catalog_row()], [_catalog_row("updated")])) + read_revision = AsyncMock(return_value=None) + _catalog_database(monkeypatch, read_rows, read_revision) + manager = MCPServerManager() + assert (await manager.catalog.list())["catalog-server"].name == "initial" + assert (await manager.catalog.list())["catalog-server"].name == "updated" + assert read_rows.await_count == 2 + + +@pytest.mark.asyncio +async def test_catalog_failed_refresh_does_not_apply_the_read_revision(monkeypatch): + read_rows = AsyncMock(side_effect=([_catalog_row()], RuntimeError("unavailable"), [_catalog_row("updated")])) + read_revision = AsyncMock(side_effect=[_revision_row(7), _revision_row(8), _revision_row(8)]) + _catalog_database(monkeypatch, read_rows, read_revision) + manager = MCPServerManager() + assert (await manager.catalog.list())["catalog-server"].name == "initial" + with pytest.raises(HTTPException) as error: + await manager.catalog.list() + assert error.value.status_code == 503 + assert (await manager.catalog.list())["catalog-server"].name == "updated" + assert read_rows.await_count == 3 + + +@pytest.mark.asyncio +async def test_catalog_reload_applies_revision_so_next_operation_skips_read(monkeypatch): + read_rows = AsyncMock(return_value=[_catalog_row()]) + read_revision = AsyncMock(return_value=_revision_row(7)) + _catalog_database(monkeypatch, read_rows, read_revision) + manager = MCPServerManager() + await manager.reload_servers_from_database() + assert (await manager.catalog.list())["catalog-server"].name == "initial" + read_rows.assert_awaited_once() + assert read_revision.await_count == 2 + + +@pytest.mark.asyncio +async def test_catalog_waiter_that_observed_newer_revision_refreshes(monkeypatch): + entered = asyncio.Event() + release = asyncio.Event() + waiter_saw_new_revision = asyncio.Event() + current_revision = [5] + reads = [0] + + async def read_rows(**_kwargs): + reads[0] += 1 + if not entered.is_set(): + entered.set() + await release.wait() + return [_catalog_row()] + if reads[0] == 2: + return [_catalog_row()] + return [ + _catalog_row(), + _catalog_row("added").model_copy(update={"server_id": "added-server"}), + ] + + async def read_current_revision(**_kwargs): + if current_revision[0] == 6: + waiter_saw_new_revision.set() + return _revision_row(current_revision[0]) + + read_rows_mock = AsyncMock(side_effect=read_rows) + read_revision = AsyncMock(side_effect=read_current_revision) + _catalog_database(monkeypatch, read_rows_mock, read_revision) + manager = MCPServerManager() + first = asyncio.create_task(manager.catalog.list()) + await entered.wait() + middle = asyncio.create_task(manager.catalog.list()) + await asyncio.sleep(0) + current_revision[0] = 6 + last = asyncio.create_task(manager.catalog.list()) + await waiter_saw_new_revision.wait() + release.set() + _, _, servers = await asyncio.gather(first, middle, last) + assert "added-server" in servers + assert reads[0] == 3 + + class TestSharedIdentifierPrefixWarning: """Two stored rows sharing lowercased alias-or-server_name publish one tool prefix; reload must surface them once so the ambiguity is visible.""" @@ -16508,8 +17677,11 @@ class TestSharedIdentifierPrefixWarning: updated_at=datetime.now(), ), ] - raw_rows = [MagicMock(model_dump=lambda row=row: row.model_dump()) for row in rows] + raw_rows = [MagicMock(model_dump=lambda row=row, **kwargs: row.model_dump(**kwargs)) for row in rows] repository = MagicMock() + prisma = MagicMock() + prisma.writer_db = prisma.db + prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) repository.table.find_many = AsyncMock(return_value=raw_rows) async def build_from_table(table, **_kwargs): @@ -16523,17 +17695,18 @@ class TestSharedIdentifierPrefixWarning: ) with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", + "litellm.proxy._experimental.mcp_server.db.MCPServerRepository", return_value=repository, ), patch( "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", - return_value=MagicMock(), + return_value=prisma, ), patch.object(manager, "build_mcp_server_from_table", new=build_from_table), - patch.object(manager, "_maybe_register_openapi_tools", new=AsyncMock()), - patch.object(manager, "_prime_oauth_metadata_discovery_for_servers"), + patch.object(manager, "maybe_register_openapi_tools", new=AsyncMock()), + patch.object(manager, "prime_oauth_metadata_discovery_for_servers"), caplog.at_level(logging.WARNING, logger="LiteLLM"), ): await manager.reload_servers_from_database() @@ -16564,22 +17737,26 @@ async def test_reload_warns_once_about_a_blocked_stdio_row_that_is_rebuilt_every monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", flag) manager = MCPServerManager() repository = MagicMock() + prisma = MagicMock() + prisma.writer_db = prisma.db + prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) async def build_from_table(table, **_kwargs): return MCPServer(server_id=table.server_id, name=table.server_name, transport=table.transport) with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", + "litellm.proxy._experimental.mcp_server.db.MCPServerRepository", return_value=repository, ), patch( "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", - return_value=MagicMock(), + return_value=prisma, ), patch.object(manager, "build_mcp_server_from_table", new=build_from_table), - patch.object(manager, "_maybe_register_openapi_tools", new=AsyncMock()), - patch.object(manager, "_prime_oauth_metadata_discovery_for_servers"), + patch.object(manager, "maybe_register_openapi_tools", new=AsyncMock()), + patch.object(manager, "prime_oauth_metadata_discovery_for_servers"), caplog.at_level(logging.WARNING, logger="LiteLLM"), ): for transport in transports: @@ -17250,3 +18427,380 @@ async def test_upstream_preparation_honors_case_sensitive_extra_command(monkeypa client: Final = await MCPServerManager()._create_mcp_client(server) assert client.stdio_config is not None assert client.stdio_config["command"] == "/opt/tools/CustomRunner" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("overlap", [False, True]) +async def test_catalog_cached_revision_retains_newly_discovered_tool_routes(monkeypatch, overlap): + from mcp.types import Tool + + read_rows = AsyncMock(return_value=[]) + _catalog_database(monkeypatch, read_rows, AsyncMock(return_value=_revision_row(7))) + manager = MCPServerManager() + first = MCPServer(server_id="first", name="first", transport=MCPTransport.http) + second = MCPServer(server_id="second", name="second", transport=MCPTransport.http) + manager.config_mcp_servers = {server.server_id: server for server in (first, second)} + manager.published_tool_routes = {"search": "first"} + ready = asyncio.Event() + release = asyncio.Event() + + async def read_catalog(): + async with manager.catalog.operation(): + assert manager._get_mcp_server_from_tool_name("search").server_id == "first" + ready.set() + await release.wait() + + reader = asyncio.create_task(read_catalog()) if overlap else None + try: + if reader is not None: + await asyncio.wait_for(ready.wait(), 2) + async with manager.catalog.operation(): + manager._create_prefixed_tools([Tool(name="search", input_schema={})], second) + assert manager.published_tool_routes["search"] == "second" + release.set() + if reader is not None: + await asyncio.wait_for(reader, 2) + async with manager.catalog.operation(): + assert manager._get_mcp_server_from_tool_name("search").server_id == "second" + assert manager.published_tool_routes["search"] == "second" + read_rows.assert_awaited_once() + finally: + if reader is not None: + reader.cancel() + await asyncio.gather(reader, return_exceptions=True) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("poisoned_route", [False, True]) +async def test_cached_discovery_cannot_authorize_another_servers_local_handler(monkeypatch, poisoned_route): + from datetime import datetime + + from fastapi import HTTPException + from mcp.types import Tool + + from litellm.proxy._experimental.mcp_server import operations + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + _catalog_database(monkeypatch, AsyncMock(return_value=[]), AsyncMock(return_value=_revision_row(7))) + manager = MCPServerManager() + allowed = MCPServer(server_id="allowed", name="allowed", transport=MCPTransport.http) + private = MCPServer(server_id="private", name="private", transport=MCPTransport.http) + manager.config_mcp_servers = {server.server_id: server for server in (allowed, private)} + manager.published_tool_routes = {"getsecret": "private", "private-getsecret": "private"} + handler = AsyncMock(return_value="private result") + monkeypatch.setattr(global_mcp_tool_registry, "published_tools", {}) + global_mcp_tool_registry.register_tool("private-getsecret", "Private tool", {}, handler) + monkeypatch.setattr(operations, "global_mcp_server_manager", manager) + check = AsyncMock(return_value={}) + monkeypatch.setattr(manager, "pre_call_tool_check", check) + async with manager.catalog.operation(): + manager._create_prefixed_tools([Tool(name="private-getsecret", input_schema={})], allowed) + if poisoned_route: + manager.published_tool_routes["private-getsecret"] = "allowed" + async with manager.catalog.operation(): + with pytest.raises(HTTPException) as denied: + await operations._execute_mcp_tool( + name="private-getsecret", arguments={}, allowed_mcp_servers=[allowed], start_time=datetime.now() + ) + assert denied.value.status_code == 403 + handler.assert_not_awaited() + check.assert_not_awaited() + result = await operations._execute_mcp_tool( + name="private-getsecret", arguments={}, allowed_mcp_servers=[private], start_time=datetime.now() + ) + assert result.is_error is False + handler.assert_awaited_once_with() + assert check.await_args.kwargs["server"].server_id == private.server_id + assert manager._get_mcp_server_from_tool_name("allowed-private-getsecret").server_id == allowed.server_id + + +@pytest.mark.parametrize("local_handler", [False, True]) +def test_discovery_preserves_registered_tool_namespaces(monkeypatch, local_handler): + from mcp.types import Tool + + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + manager = MCPServerManager() + allowed = MCPServer(server_id="allowed", name="allowed", transport=MCPTransport.http) + private = MCPServer(server_id="private", name="private", alias="private-alias", transport=MCPTransport.http) + manager.registry = {server.server_id: server for server in (allowed, private)} + manager.published_tool_routes = {"private-getsecret": "private", "private-alias-getsecret": "private"} + monkeypatch.setattr(global_mcp_tool_registry, "published_tools", {}) + if local_handler: + global_mcp_tool_registry.register_tool("orphan", "Orphan local tool", {}, AsyncMock()) + tools = [Tool(name=name, input_schema={}) for name in ("private-getsecret", "private-alias-getsecret", "orphan")] + listed = manager._create_prefixed_tools(tools, allowed) + assert [tool.name for tool in listed] == ["allowed-" + tool.name for tool in tools] + assert manager.published_tool_routes["private-getsecret"] == "private" + assert manager.published_tool_routes["private-alias-getsecret"] == "private" + assert manager._get_mcp_server_from_tool_name("allowed-private-getsecret").server_id == "allowed" + assert ("orphan" in manager.published_tool_routes) is not local_handler + + +@pytest.mark.asyncio +async def test_cached_catalog_excludes_routes_for_servers_outside_its_snapshot(monkeypatch): + read_rows = AsyncMock(return_value=[]) + _catalog_database(monkeypatch, read_rows, AsyncMock(return_value=_revision_row(7))) + manager = MCPServerManager() + pinned = MCPServer(server_id="pinned", name="pinned", transport=MCPTransport.http) + manager.config_mcp_servers = {pinned.server_id: pinned} + async with manager.catalog.operation(): + assert manager.get_mcp_server_by_id("pinned") is not None + published = MCPServer(server_id="published", name="published", transport=MCPTransport.http) + manager.registry[published.server_id] = published + manager.published_tool_routes = {"search": "published", "known": "pinned"} + async with manager.catalog.operation(): + assert manager.get_mcp_server_by_id("published") is None + assert "search" not in manager.tool_name_to_mcp_server_name_mapping + assert manager._get_mcp_server_from_tool_name("known").server_id == "pinned" + read_rows.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_overlapping_server_prefix_cannot_authorize_registered_openapi_handler(monkeypatch, tmp_path): + from datetime import datetime + + from fastapi import HTTPException + + from litellm.proxy._experimental.mcp_server import operations + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + manager = MCPServerManager() + private = MCPServer(server_id="private", name="billing", alias="billing", transport=MCPTransport.http) + allowed = MCPServer(server_id="allowed", name="billing_admin", alias="billing-admin", transport=MCPTransport.http) + manager.config_mcp_servers = {server.server_id: server for server in (private, allowed)} + monkeypatch.setattr(global_mcp_tool_registry, "published_tools", {}) + monkeypatch.setattr(operations, "global_mcp_server_manager", manager) + spec = tmp_path / "spec.json" + spec.write_text(json.dumps({"openapi": "3.0.0", "paths": {"/export": {"get": {"operationId": "admin-export"}}}})) + await manager._register_openapi_tools(str(spec), private, "https://example.com") + tool = global_mcp_tool_registry.get_tool("billing-admin-export") + assert tool is not None + handler = AsyncMock(return_value="private result") + tool.handler = handler + check = AsyncMock(return_value={}) + monkeypatch.setattr(manager, "pre_call_tool_check", check) + + with pytest.raises(HTTPException) as denied: + await operations._execute_mcp_tool( + name="billing-admin-export", arguments={}, allowed_mcp_servers=[allowed], start_time=datetime.now() + ) + assert denied.value.status_code == 403 + handler.assert_not_awaited() + check.assert_not_awaited() + result = await operations._execute_mcp_tool( + name="billing-admin-export", arguments={}, allowed_mcp_servers=[private], start_time=datetime.now() + ) + assert result.is_error is False + handler.assert_awaited_once_with() + assert check.await_args.kwargs["server"].server_id == private.server_id + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", ["update", "delete"]) +async def test_catalog_observes_writer_changes_while_read_replica_lags(monkeypatch, change): + from types import SimpleNamespace + + from litellm.proxy import proxy_server + from litellm.proxy.db.routing_prisma_wrapper import _RoutedActions + + original: Final = _catalog_row() + changed: Final = _catalog_row("updated") + reader: Final = SimpleNamespace( + litellm_mcpservertable=SimpleNamespace(find_many=AsyncMock(return_value=[original])), + litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=_revision_row(7))), + ) + writer: Final = SimpleNamespace( + litellm_mcpservertable=SimpleNamespace(find_many=AsyncMock(return_value=[original])), + litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=_revision_row(7))), + ) + routed: Final = SimpleNamespace(**{ + name: _RoutedActions(getattr(writer, name), getattr(reader, name), lambda: True) + for name in ("litellm_mcpservertable", "litellm_config") + }) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=routed, writer_db=writer)) + manager: Final = MCPServerManager() + assert (await manager.catalog.resolve(original.server_id)).name == "initial" + writer.litellm_mcpservertable.find_many.return_value = [changed] if change == "update" else [] + writer.litellm_config.find_unique.return_value = _revision_row(8) + selected: Final = await manager.catalog.resolve(original.server_id) + assert (selected.name if selected is not None else None) == ("updated" if change == "update" else None) + writer.litellm_config.find_unique.side_effect = RuntimeError("writer unavailable") + with pytest.raises(HTTPException) as unavailable: + await manager.catalog.resolve(original.server_id) + assert unavailable.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_catalog_revision_failure_preserves_state_and_recovers(monkeypatch): + reader: Final = AsyncMock(side_effect=([_catalog_row()], [_catalog_row("updated")])) + revisions: Final = AsyncMock(side_effect=[_revision_row(7), RuntimeError("private database detail"), _revision_row(8)]) + _catalog_database(monkeypatch, reader, revisions) + manager: Final = MCPServerManager() + assert (await manager.catalog.resolve("catalog-server")).name == "initial" + with pytest.raises(HTTPException) as unavailable: + await manager.catalog.resolve("catalog-server") + assert unavailable.value.status_code == 503 + assert unavailable.value.detail == "MCP server configuration could not be refreshed" + assert manager.registry["catalog-server"].name == "initial" + assert (await manager.catalog.resolve("catalog-server")).name == "updated" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("groups", [[], ["catalog-group"]]) +async def test_access_group_resolution_uses_the_operation_catalog(monkeypatch, groups): + from litellm.proxy._experimental.mcp_server import mcp_server_manager + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler + + row = _catalog_row().model_copy(update={"mcp_access_groups": groups}) + _catalog_database(monkeypatch, AsyncMock(return_value=[row])) + manager = MCPServerManager() + monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager) + stale = AsyncMock(return_value=set() if groups else {row.server_id}) + monkeypatch.setattr(MCPRequestHandler, "_get_db_server_ids_for_access_groups", stale) + async with manager.catalog.operation(): + allowed = await MCPRequestHandler._get_mcp_servers_from_access_groups(["catalog-group"]) + assert set(allowed) == ({row.server_id} if groups else set()) + stale.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "level,inheritance", + [ + ("key", True), + ("team", True), + ("user", True), + ("org", True), + ("end_user", True), + ("agent", True), + ("key", False), + ("team", False), + ], +) +@pytest.mark.parametrize("additive", [False, True]) +@pytest.mark.parametrize("declared,direct", [(False, False), (True, False), (True, True)]) +async def test_empty_access_group_scope_cannot_inherit_unrelated_server( + monkeypatch, level, inheritance, declared, direct, additive +): + from types import SimpleNamespace + + from litellm.proxy._experimental.mcp_server import mcp_server_manager + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth + + row = _catalog_row() + _catalog_database(monkeypatch, AsyncMock(return_value=[row])) + manager = MCPServerManager() + monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager) + permission = LiteLLM_ObjectPermissionTable( + object_permission_id="group-scope", + mcp_access_groups=["missing-group"] if declared else [], + mcp_servers=[row.server_id] if direct else [], + ) + auth = UserAPIKeyAuth(api_key="test", team_id="team", org_id="org", end_user_id="end", agent_id="agent") + methods = { + "key": "_get_allowed_mcp_servers_for_key", + "team": "_get_allowed_mcp_servers_for_team", + "user": "_get_allowed_mcp_servers_for_user", + "org": "_get_allowed_mcp_servers_for_org", + "end_user": "_get_allowed_mcp_servers_for_end_user", + "agent": "get_allowed_mcp_servers_for_agent", + } + selected = getattr(MCPRequestHandler, methods[level]) + for name, method in methods.items(): + monkeypatch.setattr( + MCPRequestHandler, + method, + AsyncMock(return_value=[row.server_id] if inheritance and name in ("key", "team") else []), + ) + monkeypatch.setattr( + MCPRequestHandler, + "_get_key_access_group_mcp_server_extras", + AsyncMock(return_value=[row.server_id] if additive else []), + ) + monkeypatch.setattr(MCPRequestHandler, "_get_agent_access_group_server_ceiling", AsyncMock(return_value=None)) + monkeypatch.setattr(MCPRequestHandler, "_key_or_team_declares_toolsets", AsyncMock(return_value=False)) + monkeypatch.setattr( + MCPRequestHandler, "_get_key_object_permission", lambda auth: permission if level == "key" else None + ) + if level == "team": + + async def team_servers(auth): + return list( + await MCPRequestHandler._team_granted_servers(SimpleNamespace(object_permission=permission), []) + ) + + monkeypatch.setattr(MCPRequestHandler, methods[level], team_servers) + else: + monkeypatch.setattr(MCPRequestHandler, methods[level], selected) + if level != "key": + monkeypatch.setattr( + MCPRequestHandler, f"_get_{level}_object_permission", AsyncMock(return_value=permission) + ) + async with manager.catalog.operation(): + access = await MCPRequestHandler.get_mcp_server_access(auth) + assert set(access.server_ids) == ( + {row.server_id} + if direct or (not declared and inheritance) or (additive and level in ("key", "team")) + else set() + ) + if declared: + assert access.scope == "scoped" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("issuer_parameter_supported", [False, True]) +@pytest.mark.parametrize("anchored", [False, True]) +async def test_catalog_accepts_discovered_issuer_response_support( + monkeypatch: pytest.MonkeyPatch, issuer_parameter_supported: bool, anchored: bool +) -> None: + from litellm.proxy import proxy_server + from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata + + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="discovered-issuer", name="discovered_issuer", transport=MCPTransport.http, + url="https://resource.example/mcp", auth_type=MCPAuth.oauth2, + issuer="https://issuer.example" if anchored else None, issuer_is_anchored=anchored, + ) + manager.registry = {server.server_id: server} + manager._set_oauth_discovery_deferred(server.server_id, True) + metadata: Final = MCPOAuthMetadata( + discovered_issuer="https://issuer.example", + authorization_url="https://issuer.example/authorize", token_url="https://issuer.example/token", + authorization_response_iss_parameter_supported=issuer_parameter_supported, + ) + discovery: Final = AsyncMock(return_value=metadata) + monkeypatch.setattr(proxy_server, "prisma_client", None) + monkeypatch.setattr(manager, "_discover_oauth_metadata_for_server", discovery) + async with manager.catalog.operation(): + resolved: Final = await manager.ensure_oauth_metadata_discovered(server) + repeated: Final = await manager.ensure_oauth_metadata_discovered(server) + assert resolved.issuer == "https://issuer.example" + assert resolved.authorization_response_iss_parameter_supported is issuer_parameter_supported + assert repeated == resolved + assert manager.registry[server.server_id] == resolved + discovery.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_catalog_rejects_a_changed_anchored_issuer_during_discovery(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="anchored-issuer", name="anchored_issuer", transport=MCPTransport.http, + url="https://resource.example/mcp", auth_type=MCPAuth.oauth2, + issuer="https://original.example", issuer_is_anchored=True, + ) + manager.registry = {server.server_id: server} + discovery: Final = AsyncMock() + monkeypatch.setattr(proxy_server, "prisma_client", None) + monkeypatch.setattr(manager, "_discover_oauth_metadata_for_server", discovery) + async with manager.catalog.operation(): + manager.registry[server.server_id] = server.model_copy(update={"issuer": "https://replacement.example"}) + with pytest.raises(HTTPException) as rejected: + await manager.ensure_oauth_metadata_discovered(server) + assert rejected.value.status_code == 503 + discovery.assert_not_awaited() diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index 7b9fddebb36..189424a9f90 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py index c469a82e889..6430d5f9259 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py @@ -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=[]) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_oauth2_flow_backfill.py b/tests/unit/proxy/_experimental/mcp_server/test_oauth2_flow_backfill.py index c1239c228aa..ce079878328 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_oauth2_flow_backfill.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_oauth2_flow_backfill.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py b/tests/unit/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py index 0dc52b13950..a576bcbb3e2 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_operations.py b/tests/unit/proxy/_experimental/mcp_server/test_operations.py index f710bc7f1d7..e3dad9c513e 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_operations.py @@ -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() diff --git a/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py index bd3c34f5abf..0e38d03ca3a 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -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") diff --git a/tests/unit/proxy/_experimental/mcp_server/test_short_mcp_tool_prefix.py b/tests/unit/proxy/_experimental/mcp_server/test_short_mcp_tool_prefix.py index 941e5deee93..7b6007c405e 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_short_mcp_tool_prefix.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_short_mcp_tool_prefix.py @@ -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" diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_mcp_security.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_mcp_security.py index d57a91d45bf..e4c7bcba951 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_mcp_security.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_mcp_security.py @@ -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"] diff --git a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py index 081a2d8ce73..44cc1e80b09 100644 --- a/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -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 diff --git a/tests/unit/proxy/test_dynamic_mcp_route.py b/tests/unit/proxy/test_dynamic_mcp_route.py index da7b8e01f46..83963fbd962 100644 --- a/tests/unit/proxy/test_dynamic_mcp_route.py +++ b/tests/unit/proxy/test_dynamic_mcp_route.py @@ -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 diff --git a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py index 7445df1b066..d215ce292aa 100644 --- a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -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=[]), diff --git a/tests/unit/responses/mcp/test_mcp_streaming_iterator.py b/tests/unit/responses/mcp/test_mcp_streaming_iterator.py index c2c204e7024..3982081706f 100644 --- a/tests/unit/responses/mcp/test_mcp_streaming_iterator.py +++ b/tests/unit/responses/mcp/test_mcp_streaming_iterator.py @@ -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 diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 9a7bafc7b3c..cbaea37d409 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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) */