From 74a470410d49633fc5850fa7e7b911f022d68cd8 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Tue, 22 Sep 2026 12:40:45 -0700 Subject: [PATCH] fix(mcp): refresh shared catalog state for each operation --- .../mcp_server/auth/user_api_key_auth_mcp.py | 2 + .../mcp_server/byok_oauth_endpoints.py | 2 + .../proxy/_experimental/mcp_server/catalog.py | 394 +++++++++++++++ litellm/proxy/_experimental/mcp_server/db.py | 9 + .../mcp_server/discoverable_endpoints.py | 12 +- .../mcp_server/gateway_dcr_flow.py | 2 + .../mcp_server/mcp_server_manager.py | 275 +++-------- .../_experimental/mcp_server/operations.py | 28 +- .../mcp_server/rest_endpoints.py | 3 + .../mcp_server/semantic_tool_filter.py | 2 + .../proxy/_experimental/mcp_server/server.py | 3 + .../_experimental/mcp_server/tool_registry.py | 30 +- .../proxy/_experimental/mcp_server/utils.py | 6 +- .../mcp_security/mcp_security_guardrail.py | 11 +- .../mcp_management_endpoints.py | 268 ++++++----- litellm/proxy/proxy_server.py | 51 +- .../mcp/litellm_proxy_mcp_handler.py | 4 + .../types/mcp_server/mcp_server_manager.py | 2 +- .../auth/test_user_api_key_auth_mcp.py | 2 + .../mcp_server/test_byok_oauth_endpoints.py | 1 + .../mcp_server/test_discoverable_endpoints.py | 10 +- .../mcp_server/test_mcp_server.py | 8 +- .../mcp_server/test_mcp_server_manager.py | 449 +++++++++++++++--- .../mcp_server/test_rest_endpoints.py | 5 +- .../mcp_server/test_short_mcp_tool_prefix.py | 12 +- .../guardrail_hooks/test_mcp_security.py | 27 ++ .../test_mcp_management_endpoints.py | 154 ++++-- .../mcp/test_litellm_proxy_mcp_handler.py | 6 + 28 files changed, 1288 insertions(+), 490 deletions(-) create mode 100644 litellm/proxy/_experimental/mcp_server/catalog.py diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 60dc91a69cc..6b8006dbfc9 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -13,6 +13,7 @@ from typing_extensions import assert_never import litellm from litellm._logging import verbose_logger +from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, global_manager from litellm.proxy._experimental.mcp_server.oauth_utils import ( get_passthrough_resource_metadata_url, get_passthrough_www_authenticate, @@ -398,6 +399,7 @@ class MCPRequestHandler: LITELLM_MCP_ACCESS_GROUPS_HEADER_NAME = SpecialHeaders.mcp_access_groups.value @staticmethod + @catalog_operation(global_manager) async def process_mcp_request( scope: Scope, ) -> tuple[ diff --git a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py index 2c63e0a96d8..086dd946efe 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 catalog_operation, global_manager from litellm.proxy._experimental.mcp_server.db import store_user_credential from litellm.proxy._experimental.mcp_server.oauth_utils import ( BYOK_RESOURCE_METADATA_PATH, @@ -649,6 +650,7 @@ async def byok_protected_resource_metadata(request: Request) -> JSONResponse: @router.get("/v1/mcp/oauth/authorize", include_in_schema=False) +@catalog_operation(global_manager) async def byok_authorize_get( request: Request, client_id: str | None = None, diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py new file mode 100644 index 00000000000..0f5eca285e3 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -0,0 +1,394 @@ +"""Authoritative MCP catalog snapshots shared by legacy lookup adapters.""" + +from __future__ import annotations + +import asyncio +import hashlib +import json +from collections.abc import AsyncIterator, Awaitable, Callable, Mapping, Sequence +from contextlib import asynccontextmanager +from contextvars import ContextVar +from dataclasses import dataclass +from functools import wraps +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, ParamSpec, TypeVar + +from litellm._logging import verbose_logger + +if TYPE_CHECKING: + from mcp.types import Tool as SDKTool + from pydantic import BaseModel + + from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing, ServerOutcome + 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") + + +@dataclass(frozen=True, slots=True) +class CatalogSnapshot: + servers: Mapping[str, MCPServer] + identity: str + tools: Mapping[str, MCPTool] + routing: dict[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")) + | (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._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) -> dict[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) -> dict[str, MCPServer]: + snapshot: Final = self.current() + return ( + dict(snapshot.servers) if snapshot is not None else self.manager.config_mcp_servers | self.manager.registry + ) + + @asynccontextmanager + async def operation(self) -> AsyncIterator[CatalogSnapshot]: + current: Final = self.current() + if current is not None: + yield current + return + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is not None: + from litellm.proxy.proxy_server import should_load_db_object + + if should_load_db_object("mcp"): + await self.reload() + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + snapshot: Final = _snapshot(self.manager, self._database_identity) + 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) + + def _retain_discovered_routing(self, snapshot: CatalogSnapshot) -> None: + from litellm.proxy._experimental.mcp_server.utils import normalize_server_name + + current: Final = self.manager.config_mcp_servers | self.manager.registry + unchanged_owners: Final = frozenset( + owner + for key, server in snapshot.servers.items() + if current.get(key) == server + for owner in self.manager.owned_mapping_values(server) + ) + self.manager.published_tool_routes = self.manager.published_tool_routes | MappingProxyType( + { + name: owner + for name, owner in snapshot.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_id(identifier, client_ip) or self.manager.get_mcp_server_by_name( + identifier, 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 + ) + _check_oauth_revision(selected, registered) + resolved: Final = await resolve(selected) + _check_oauth_revision(selected, resolved) + return resolved + + @staticmethod + async def list( + servers: Sequence[MCPServer], + fetch: Callable[[MCPServer], Awaitable[tuple[list[SDKTool], ServerOutcome]]], + server_key: Callable[[MCPServer], str], + ) -> AggregateToolListing: + from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing + + tasks: Final = tuple(asyncio.ensure_future(fetch(server)) for server in servers) + try: + results: Final = await asyncio.gather(*tasks) + return AggregateToolListing( + tools=[tool for tools, _ in results for tool in tools], + outcomes={server_key(server): outcome for server, (_, outcome) in zip(servers, results)}, + ) + finally: + for task in tasks: + if not task.done(): + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + + async def reload(self) -> None: + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + async with self._refresh_lock: + staged_config: Final = { + key: value.model_copy(deep=True) for key, value in self.manager.config_mcp_servers.items() + } + await self.manager.hydrate_config_servers_dcr_clients(tuple(staged_config.values())) + staged_routing: Final = dict(self.manager.published_tool_routes) + closed: Final = asyncio.Event() + routing_token: Final = self._staged_routing.set((staged_routing, closed)) + try: + with global_mcp_tool_registry.catalog_scope(global_mcp_tool_registry.published_tools) as staged_tools: + await self._reload() + self.manager.config_mcp_servers = staged_config + global_mcp_tool_registry.tools = staged_tools + self.manager.published_tool_routes = staged_routing + finally: + closed.set() + self._staged_routing.reset(routing_token) + + async def _reload(self) -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + carry_forward_resolved_oauth_endpoints, + config_ids_capturing_db_identifiers, + oauth_endpoints_unresolved, + warn_on_server_name_fields, + ) + 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 LiteLLM_MCPServerTable, 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[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.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), + ) + 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, + ) + + # 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_discovery_lists(server_id) + self.manager.invalidate_oauth_discovery_state(server_id) + self._database_identity = database_identity + self.manager.registry = registered_registry + # 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 diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 30ee8b7a4fc..ee7298fe146 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -645,6 +645,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 _db_find_mcp_server_rows(prisma_client, 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 64bab0a7832..3378b1df521 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -36,6 +36,7 @@ from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( can_store_oauth_credential, oauth_authorization_uses_gateway_credential, ) +from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, global_manager from litellm.proxy._experimental.mcp_server.faults import ( CallerRejected, CredentialSource, @@ -1769,7 +1770,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: @@ -1900,6 +1901,7 @@ async def authorize_mcp_session( @router.get("/{mcp_server_name}/authorize") @router.get("/authorize") +@catalog_operation(global_manager) async def authorize( request: Request, redirect_uri: str, @@ -1974,6 +1976,7 @@ async def authorize( @router.post("/{mcp_server_name}/token") @router.post("/token") +@catalog_operation(global_manager) async def token_endpoint( request: Request, grant_type: str = Form(...), @@ -2078,6 +2081,7 @@ async def authorize_flow(request: Request, flow: str) -> Response: @router.post("/authorize/complete") +@catalog_operation(global_manager) async def authorize_complete( request: Request, flow: str = Form(...), @@ -2696,6 +2700,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}}") +@catalog_operation(global_manager) async def oauth_protected_resource_mcp_standard(request: Request, mcp_server_name: str): """ OAuth protected resource discovery endpoint using standard MCP URL pattern. @@ -2716,6 +2721,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") +@catalog_operation(global_manager) async def oauth_protected_resource_mcp(request: Request, mcp_server_name: str | None = None): """ OAuth protected resource discovery endpoint using LiteLLM legacy URL pattern. @@ -2792,6 +2798,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}}") +@catalog_operation(global_manager) async def oauth_authorization_server_mcp_standard(request: Request, mcp_server_name: str): """ OAuth authorization server discovery endpoint using standard MCP URL pattern. @@ -2809,6 +2816,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") +@catalog_operation(global_manager) async def oauth_authorization_server_mcp(request: Request, mcp_server_name: str | None = None): """ OAuth authorization server discovery endpoint. @@ -2882,6 +2890,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") +@catalog_operation(global_manager) async def oauth_authorization_server_legacy(request: Request, mcp_server_name: str): """ OAuth authorization server discovery for legacy /{server_name}/mcp pattern. @@ -2895,6 +2904,7 @@ async def oauth_authorization_server_legacy(request: Request, mcp_server_name: s @router.post("/{mcp_server_name}/register") @router.post("/register") +@catalog_operation(global_manager) 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) diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py index e66504af47a..624b1795a14 100644 --- a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -55,6 +55,7 @@ from typing_extensions import NotRequired, ReadOnly, TypedDict, assert_never 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, @@ -770,6 +771,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 20c114a2f3e..93899ea3a78 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -178,7 +178,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.llms.custom_http import httpxSpecialProvider from litellm.types.mcp import ( DEFAULT_SUBJECT_TOKEN_TYPE, @@ -287,7 +286,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] @@ -541,7 +540,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]: @@ -722,7 +721,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 @@ -781,7 +780,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 @@ -1427,7 +1426,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, @@ -1860,6 +1859,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] = {} @@ -1886,7 +1888,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", @@ -1900,13 +1902,11 @@ 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._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: @@ -1919,7 +1919,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() @@ -1932,8 +1932,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() @@ -2008,7 +2008,7 @@ class MCPServerManager: return resolved 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: @@ -2045,7 +2045,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) @@ -2056,7 +2056,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 ( @@ -2103,7 +2103,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 @@ -2119,7 +2119,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: @@ -2148,19 +2148,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 @@ -2209,7 +2214,7 @@ class MCPServerManager: 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") @@ -2291,7 +2296,15 @@ class MCPServerManager: """ 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) -> dict[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). @@ -2378,7 +2391,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, @@ -2584,10 +2597,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_discovery_lists(server_id) + self.invalidate_discovery_lists(server_id) self.config_mcp_servers[server_id] = new_server self._set_oauth_discovery_deferred( server_id, @@ -2608,13 +2621,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.""" @@ -2622,7 +2635,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( @@ -2787,17 +2800,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_discovery_lists(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_discovery_lists(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 @@ -2808,13 +2822,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) @@ -2845,7 +2859,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) @@ -2936,6 +2950,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)) @@ -3160,13 +3175,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) @@ -3194,10 +3210,10 @@ class MCPServerManager: # Re-decrypting plaintext would zero the values, so build with # env_vars_are_encrypted=False. 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_discovery_lists(mcp_server.server_id) + self.assign_unique_short_prefix(new_server) + self.invalidate_discovery_lists(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) self.prime_oauth_metadata_discovery(new_server) verbose_logger.debug("Added MCP Server: %s", new_server.name) @@ -3215,7 +3231,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: @@ -3227,14 +3243,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_discovery_lists(mcp_server.server_id) + self.assign_unique_short_prefix(new_server) + self.invalidate_discovery_lists(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) self.prime_oauth_metadata_discovery(new_server) verbose_logger.debug("Updated MCP Server: %s", new_server.name) @@ -4466,7 +4482,9 @@ 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 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) @@ -5256,7 +5274,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, @@ -6146,7 +6164,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 @@ -6445,7 +6463,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 @@ -6507,152 +6525,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), - ) - 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_discovery_lists(server_id) - self.registry = registered_registry - # 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 = [] diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index fcee3483e15..bbe2987b856 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -1,6 +1,5 @@ """Shared MCP operation policy and dispatch.""" -import asyncio import traceback import types import uuid @@ -50,6 +49,7 @@ from litellm.proxy._experimental.mcp_server.byok_credential_cache import ( cache_byok_credential, get_cached_byok_credential, ) +from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog, catalog_operation from litellm.proxy._experimental.mcp_server.contracts import ( AuthorizedToolCall, OperationContext, @@ -614,6 +614,7 @@ def apply_tool_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, @@ -930,6 +931,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, @@ -1151,24 +1153,17 @@ async def _get_tools_from_mcp_servers( verbose_logger.exception("Error getting tools from server %s: %s", server.name, e) return [], classify_list_exception(e) - # Fetch tools from all servers in parallel - tasks: Final = [_fetch_and_filter_server_tools(server) for server in allowed_mcp_servers] - results: Final = await asyncio.gather(*tasks) - - # Flatten results into single list - all_tools: Final[list[MCPTool]] = [tool for tools, _ in results for tool in tools] - server_outcomes: Final[dict[str, ServerOutcome]] = { - _aggregate_server_key(server): outcome - for server, (_, outcome) in zip(allowed_mcp_servers, results) - if server is not None - } + listing: Final = await TargetCatalog.list( + allowed_mcp_servers, _fetch_and_filter_server_tools, _aggregate_server_key + ) + all_tools: Final = listing.tools + server_outcomes: Final = listing.outcomes # If logging is enabled, enrich spend_logs_metadata with counts if litellm_logging_obj: per_server_tool_counts: Final[dict[str, int]] = { - _aggregate_server_key(server): len(server_tools) - for server, (server_tools, _) in zip(allowed_mcp_servers, results) - if server is not None + key: outcome.tool_count if isinstance(outcome, ServerListOk) else 0 + for key, outcome in server_outcomes.items() } metadata_dict: Final = litellm_logging_obj.model_call_details.get("metadata") @@ -1435,6 +1430,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, @@ -2271,6 +2267,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, @@ -3055,6 +3052,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 AuthorizedToolCall(): diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index c2f7bf7d531..74e530bb12f 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -22,6 +22,7 @@ from litellm.exceptions import ( GuardrailRaisedException, ModifyResponseException, ) +from litellm.proxy._experimental.mcp_server.catalog import catalog_operation, global_manager from litellm.proxy._experimental.mcp_server.exceptions import ( MCPServerListError, MCPServerURLCredentialsError, @@ -861,6 +862,7 @@ if MCP_AVAILABLE: return await _apply_toolset_scope(user_api_key_dict, toolset.toolset_id) @router.get("/tools/list", dependencies=[Depends(user_api_key_auth)]) + @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"), @@ -1085,6 +1087,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 433b693fcae..f160ce0f4a9 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1559,6 +1559,9 @@ if MCP_AVAILABLE: ) return user_api_key_auth.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": 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..2adf5db8f67 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, Iterator, 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,30 @@ class MCPToolRegistry: def __init__(self): # Registry to store all registered tools - self.tools: dict[str, MCPTool] = {} + self.published_tools: dict[str, MCPTool] = {} + self._catalog_tools: ContextVar[tuple[dict[str, MCPTool], asyncio.Event] | None] = ContextVar( + "mcp_catalog_tools", default=None + ) + + @property + def tools(self) -> dict[str, MCPTool]: + 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: + self.published_tools = tools + + @contextmanager + def catalog_scope(self, tools: Mapping[str, MCPTool]) -> Iterator[dict[str, MCPTool]]: + 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, diff --git a/litellm/proxy/_experimental/mcp_server/utils.py b/litellm/proxy/_experimental/mcp_server/utils.py index 6bd080f5216..b06098ef474 100644 --- a/litellm/proxy/_experimental/mcp_server/utils.py +++ b/litellm/proxy/_experimental/mcp_server/utils.py @@ -81,7 +81,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. @@ -114,7 +114,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. """ @@ -314,7 +314,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..aedc50e613a 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 @@ -46,7 +46,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,7 +90,7 @@ class MCPSecurityGuardrail(CustomGuardrail): return server_names @staticmethod - def _find_unregistered_mcp_servers(data: dict) -> set[str]: + async def _find_unregistered_mcp_servers(data: dict) -> set[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): @@ -104,7 +104,8 @@ class MCPSecurityGuardrail(CustomGuardrail): 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 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 9ad78876043..a1e72158794 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 AsyncIterator, Iterable, Mapping, Sequence +from contextlib import asynccontextmanager from dataclasses import dataclass from datetime import datetime, timedelta, timezone from typing import ( @@ -45,6 +46,8 @@ from fastapi import ( from fastapi.responses import JSONResponse from typing_extensions import ReadOnly, TypedDict +from litellm.proxy._experimental.mcp_server.catalog import catalog_operation + try: from prisma.errors import RecordNotFoundError, UniqueViolationError except ImportError: @@ -1061,7 +1064,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) @@ -1091,6 +1095,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]: @@ -1129,6 +1134,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]: @@ -1150,6 +1156,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 @@ -1284,6 +1291,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, @@ -1592,6 +1600,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, @@ -2016,9 +2025,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 @@ -2064,42 +2071,44 @@ if MCP_AVAILABLE: user_api_key_dict: UserAPIKeyAuth, request: Request | None = None, ) -> MCPServer: - server = await get_cached_temporary_mcp_server(server_id) - resolved_from_temp_cache: Final = server is not None - if server is None: - # Fall back to real DB/config server (e.g. for the user-side OAuth flow - # which calls these endpoints with a real server_id, not a temp session id). - from litellm.proxy.auth.ip_address_utils import IPAddressUtils + async with _oauth_server_operation(server_id, user_api_key_dict, request=request) as server: + return server - client_ip: Final = IPAddressUtils.get_mcp_client_ip(request) if request else None - server = global_mcp_server_manager.get_mcp_server_by_id( - server_id - ) or global_mcp_server_manager.get_mcp_server_by_name(server_id, client_ip=client_ip) - if server is None: - raise HTTPException( - status_code=status.HTTP_404_NOT_FOUND, - detail={"error": f"MCP server {server_id} not found"}, - ) + @asynccontextmanager + async def _oauth_server_operation( + server_id: str, + user_api_key_dict: UserAPIKeyAuth, + request: Request | None = None, + ) -> AsyncIterator[MCPServer]: + temporary: Final = await get_cached_temporary_mcp_server(server_id) + if temporary is not None: + if not _user_has_admin_view(user_api_key_dict): + raise HTTPException(status_code=403, detail={"error": f"Access denied to MCP server {server_id}"}) + yield temporary + return + async with global_mcp_server_manager.catalog.operation(): + yield await _resolve_saved_oauth_server(server_id, user_api_key_dict, request) - # Per-server access policy mirrors `fetch_mcp_server`: admin-view - # callers are unrestricted; non-admins must have the server in their - # allowed-servers set. Temporary cached servers come from the - # admin-only `/server/oauth/session` setup flow and are not exposed - # to non-admins. + @catalog_operation(lambda: global_mcp_server_manager) + async def _resolve_saved_oauth_server( + server_id: str, + user_api_key_dict: UserAPIKeyAuth, + request: Request | None, + ) -> MCPServer: + from litellm.proxy.auth.ip_address_utils import IPAddressUtils + + client_ip: Final = IPAddressUtils.get_mcp_client_ip(request) if request else None + server: Final = global_mcp_server_manager.get_mcp_server_by_id( + server_id + ) or global_mcp_server_manager.get_mcp_server_by_name(server_id, client_ip=client_ip) + if server is None: + raise HTTPException(status_code=404, detail={"error": f"MCP server {server_id} not found"}) if not _user_has_admin_view(user_api_key_dict): - if resolved_from_temp_cache: - raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail={"error": f"Access denied to MCP server {server_id}"}, - ) - allowed_server_ids: Final[set[str]] = set() - for auth_context in await build_effective_auth_contexts(user_api_key_dict): - allowed_server_ids.update(await global_mcp_server_manager.get_allowed_mcp_servers(auth_context)) - if server.server_id not in allowed_server_ids: - raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail={"error": f"Access denied to MCP server {server_id}"}, - ) + allowed_ids: Final[set[str]] = set() + for context in await build_effective_auth_contexts(user_api_key_dict): + allowed_ids.update(await global_mcp_server_manager.get_allowed_mcp_servers(context)) + if server.server_id not in allowed_ids: + raise HTTPException(status_code=403, detail={"error": f"Access denied to MCP server {server_id}"}) return server @router.get( @@ -2119,47 +2128,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", @@ -2179,47 +2188,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", @@ -2231,22 +2240,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[dict] = {**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}", @@ -2597,6 +2606,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 18bc02467d1..8e2bb3676d9 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -19677,30 +19677,33 @@ 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 + 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 + + 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 async def _is_mcp_access_group_cached(name: str) -> bool: @@ -19755,7 +19758,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/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 88a2b92c680..739b4087f5b 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, @@ -111,6 +112,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, @@ -229,6 +231,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, @@ -680,6 +683,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 cb32299b143..192391c6f44 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -204,7 +204,7 @@ class MCPServer(BaseModel): # 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/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 35e7055bbc0..8449fb94831 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -7485,6 +7485,7 @@ class TestGatewaySessionAdmission: 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 @@ -9575,6 +9576,7 @@ class TestScopedSessionAdmission: 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) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py index a77b4c8d565..90b92048b22 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py @@ -631,6 +631,7 @@ 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_mcpusercredentials.find_unique = AsyncMock(return_value=None) monkeypatch.setattr(proxy_server, "prisma_client", prisma) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index c1d5cedeba0..16698042361 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -9060,7 +9060,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() @@ -9085,7 +9085,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() @@ -9618,7 +9618,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 @@ -9635,7 +9635,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 @@ -11135,7 +11135,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 [] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index b715fe67e20..d80c5e84ee8 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -5535,7 +5535,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 @@ -5587,7 +5587,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() @@ -5659,7 +5659,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"), @@ -9534,7 +9534,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/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 7725aca1948..3615d9681f0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -44,7 +44,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, @@ -694,7 +694,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): @@ -720,7 +720,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): @@ -777,7 +777,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): @@ -809,7 +809,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): @@ -890,7 +890,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() @@ -907,9 +907,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={ @@ -918,12 +918,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): @@ -965,7 +965,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( @@ -998,7 +998,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", @@ -1017,7 +1017,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 @@ -1025,7 +1026,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( @@ -1034,16 +1035,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 @@ -1143,7 +1144,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): @@ -1370,7 +1371,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()) @@ -1386,7 +1387,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()) @@ -1408,7 +1409,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) @@ -1422,7 +1423,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()) @@ -7377,7 +7378,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()), @@ -7463,10 +7464,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 @@ -7486,11 +7487,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 @@ -7506,9 +7507,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): @@ -7553,7 +7554,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( @@ -7646,7 +7647,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: @@ -7662,19 +7663,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): @@ -7685,7 +7686,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( @@ -7707,7 +7708,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 @@ -7719,7 +7720,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: @@ -7742,7 +7743,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" @@ -7755,7 +7756,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" @@ -7767,7 +7768,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( @@ -7793,7 +7794,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 @@ -7807,7 +7808,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( @@ -7834,7 +7835,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" @@ -11975,7 +11976,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", @@ -12037,7 +12038,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: @@ -12395,7 +12396,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 @@ -13316,7 +13317,7 @@ def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: transport=MCPTransport.http, 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 @@ -13335,30 +13336,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 @@ -13719,12 +13720,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 @@ -14516,3 +14517,331 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie assert captured["client_ip"] is None finally: auth_context_var.reset(token) + + +@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.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=([row], [updated], [])) + 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_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.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=([row], [updated], [updated])) + 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.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=RuntimeError("database unavailable")) + 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.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=blocked_read) + 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.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row]) + + 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.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row]) + 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.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row]) + + 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_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.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=RuntimeError("unavailable")) + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma), + patch("litellm.proxy.proxy_server.general_settings", {}), + ): + with pytest.raises(RuntimeError, match="unavailable"): + await manager.catalog.resolve(server.server_id) + assert manager.registry == {server.server_id: server} + + +@pytest.mark.asyncio +async def test_catalog_list_failure_cancels_and_joins_other_upstream_fetches(): + from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog + + started: Final = asyncio.Event() + stopped: Final = asyncio.Event() + pending: Final = MCPServer(server_id="pending", name="pending", transport=MCPTransport.http) + broken: Final = MCPServer(server_id="broken", name="broken", transport=MCPTransport.http) + + async def fetch(server: MCPServer): + if server.server_id == broken.server_id: + await started.wait() + raise RuntimeError("unexpected fetch failure") + started.set() + try: + await asyncio.Event().wait() + finally: + stopped.set() + + with pytest.raises(RuntimeError, match="unexpected fetch failure"): + await TargetCatalog.list((pending, broken), fetch, lambda server: server.server_id) + assert stopped.is_set() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 233a8cc96ba..676cc31fe05 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -1559,9 +1559,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/test_litellm/proxy/_experimental/mcp_server/test_short_mcp_tool_prefix.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_short_mcp_tool_prefix.py index 941e5deee93..7b6007c405e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_short_mcp_tool_prefix.py +++ b/tests/test_litellm/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/test_litellm/proxy/guardrails/guardrail_hooks/test_mcp_security.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_mcp_security.py index d57a91d45bf..7ed0596fbe2 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_mcp_security.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_mcp_security.py @@ -205,3 +205,30 @@ 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.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=([row], [])) + 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/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 53645e62034..e21dce01119 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -1,3 +1,4 @@ +from contextlib import nullcontext import os import sys import types @@ -2629,6 +2630,7 @@ class TestTemporaryMCPSessionEndpoints: mock_manager = MagicMock() mock_manager.get_mcp_server_by_id.return_value = non_oauth_server mock_manager.get_mcp_server_by_name.return_value = None + mock_manager.catalog.resolve = AsyncMock(return_value=non_oauth_server) fake_proxy_server = types.SimpleNamespace(master_key=None) with ( @@ -2681,6 +2683,7 @@ class TestTemporaryMCPSessionEndpoints: mock_manager = MagicMock() mock_manager.get_mcp_server_by_id.return_value = internal_server mock_manager.get_mcp_server_by_name.return_value = None + mock_manager.catalog.resolve = AsyncMock(return_value=internal_server) fake_proxy_server = types.SimpleNamespace(master_key=None) with ( @@ -2738,6 +2741,63 @@ 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() + + 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) + 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 ( @@ -2754,8 +2814,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", @@ -2776,7 +2836,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, @@ -2804,8 +2864,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", @@ -2909,8 +2969,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", @@ -3055,8 +3115,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", @@ -3114,8 +3174,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", @@ -3168,8 +3228,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", @@ -3218,8 +3278,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", @@ -3266,8 +3326,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", @@ -3306,8 +3366,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", @@ -3351,8 +3411,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", @@ -3374,7 +3434,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, @@ -3405,8 +3465,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", @@ -3428,7 +3488,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, @@ -3465,8 +3525,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", @@ -3484,7 +3544,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, @@ -3528,8 +3588,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", @@ -3573,8 +3633,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", @@ -7940,3 +8000,33 @@ async def test_config_server_edit_preserves_api_contract_without_creating_rows(r prisma.tx.assert_not_called() assert server.model_dump() == original 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.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row]) + 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 diff --git a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py index 83537c236a3..9c6b2404d1e 100644 --- a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -1,3 +1,5 @@ +from contextlib import nullcontext + import subprocess import sys import textwrap @@ -32,6 +34,7 @@ def _setup_mcp_call_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock: 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 @@ -378,6 +381,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey post_call_failure_hook = _setup_proxy_logging(monkeypatch) fake_manager = types.SimpleNamespace( + catalog=types.SimpleNamespace(operation=nullcontext), get_registry=MagicMock(return_value={}), call_tool=AsyncMock(side_effect=HTTPException(status_code=500, detail="boom")) ) @@ -506,6 +510,7 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch # Patch manager methods used by _get_mcp_tools_from_manager to avoid needing full UserAPIKeyAuth fields. fake_manager = types.SimpleNamespace( + catalog=types.SimpleNamespace(operation=nullcontext), get_registry=MagicMock(return_value={}), get_allowed_mcp_servers=AsyncMock(return_value=[]), get_mcp_servers_from_ids=MagicMock(return_value=[]), @@ -559,6 +564,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=[]),