diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py index 4b43d97d2e2..1395c4b6b08 100644 --- a/litellm/proxy/_experimental/mcp_server/catalog.py +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -5,11 +5,12 @@ from __future__ import annotations import asyncio import hashlib import json -from collections.abc import AsyncIterator, Awaitable, Callable, Mapping, Sequence +from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence from contextlib import asynccontextmanager from contextvars import ContextVar from dataclasses import dataclass, replace from functools import wraps +from itertools import chain from types import MappingProxyType from typing import TYPE_CHECKING, Final, ParamSpec, TypeVar @@ -158,7 +159,7 @@ class TargetCatalog: raise HTTPException(status_code=503, detail="MCP server configuration changed; retry the operation") @asynccontextmanager - async def operation(self) -> AsyncIterator[CatalogSnapshot]: + async def operation(self) -> AsyncGenerator[CatalogSnapshot]: current: Final = self.current() scoped: Final = self._operation.get() if current is not None and scoped is not None and scoped[2] == id(asyncio.current_task()): @@ -193,13 +194,13 @@ class TargetCatalog: 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 + unchanged: Final = ( + server for key, server in servers.items() if (candidate := current.get(key)) is not None and _configuration_identity(candidate) == _configuration_identity(server) - for owner in self.manager.owned_mapping_values(server) ) + unchanged_owners: Final = frozenset(chain.from_iterable(map(self.manager.owned_mapping_values, unchanged))) return MappingProxyType( {name: owner for name, owner in routing.items() if normalize_server_name(owner) in unchanged_owners} ) @@ -282,11 +283,13 @@ class TargetCatalog: try: with global_mcp_tool_registry.catalog_scope(initial_tools) as staged_tools: await self._reload(reuse_unchanged=reuse_unchanged) - refreshed_openapi_owners: Final = frozenset( - owner + refreshed_openapi: Final = ( + server for server in self.manager.registry.values() if server.spec_path and server is not live_registry.get(server.server_id) - for owner in self.manager.owned_mapping_values(server) + ) + refreshed_openapi_owners: Final = frozenset( + chain.from_iterable(map(self.manager.owned_mapping_values, refreshed_openapi)) ) live_routes: Final = self._unchanged_routing( previous_servers, @@ -357,32 +360,13 @@ class TargetCatalog: closed.set() self._staged_routing.reset(routing_token) - async def _reload(self, *, reuse_unchanged: bool) -> None: + async def _stage_servers(self, rows: Sequence[BaseModel], *, reuse_unchanged: bool) -> dict[str, MCPServer]: + from litellm.proxy._experimental.mcp_server.db import LiteLLM_MCPServerTable from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - _warn_on_shared_identifier_prefixes, 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]] = {} @@ -390,7 +374,7 @@ class TargetCatalog: # 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: + for row in rows: try: server = LiteLLM_MCPServerTable.model_validate(row.model_dump()) existing_server = previous_registry.get(server.server_id) @@ -416,7 +400,7 @@ class TargetCatalog: alias=getattr(server, "alias", None), server_name=getattr(server, "server_name", None), ) - self.manager._warn_if_newly_blocked_stdio(server, existing_server) + self.manager.warn_if_newly_blocked_stdio(server, existing_server) verbose_logger.debug("Building server from DB: %s (%s)", server.server_id, server.server_name) # raw_rows come straight from the DB, so their global env var # values (like credentials) are still encrypted here, unlike the @@ -439,6 +423,35 @@ class TargetCatalog: e, ) + return new_registry + + async def _reload(self, *, reuse_unchanged: bool) -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + config_ids_capturing_db_identifiers, + warn_on_shared_identifier_prefixes, + ) + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + get_prisma_client_or_throw, + ) + + verbose_logger.debug("Loading MCP servers from database into registry...") + + prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") + # Load only "active", legacy "approved", and NULL (no approval workflow) rows. + # Pending/rejected servers are excluded at the DB level so we never load them. + from litellm.proxy._experimental.mcp_server.db import get_runtime_mcp_server_rows + + raw_rows: Final[Sequence[BaseModel]] = await get_runtime_mcp_server_rows(prisma_client) + database_identity: Final = hashlib.sha256( + json.dumps( + tuple(sorted(json.dumps(row.model_dump(mode="json"), sort_keys=True, default=str) for row in raw_rows)) + ).encode() + ).hexdigest() + verbose_logger.info("Found %s MCP servers in database", len(raw_rows)) + + previous_registry: Final = self.manager.registry + new_registry: Final = await self._stage_servers(raw_rows, reuse_unchanged=reuse_unchanged) + # Assign short prefixes against the full candidate set without # publishing the staged registry to concurrent callers. registered_registry: Final[dict[str, MCPServer]] = {} @@ -470,14 +483,13 @@ class TargetCatalog: for server_id in previous_registry.keys() | registered_registry.keys(): if previous_registry.get(server_id) != registered_registry.get(server_id): - self.manager._invalidate_server_definition_caches(server_id) + self.manager.invalidate_server_definition_caches(server_id) self.manager.invalidate_oauth_discovery_state(server_id) self._database_identity = database_identity self.manager.registry = registered_registry if not reuse_unchanged: - self.manager._upstream_initialize_instructions_by_server_id.clear() - self.manager._upstream_initialize_instructions_probed_at.clear() - _warn_on_shared_identifier_prefixes(registered_registry.values()) + self.manager.clear_initialize_instructions() + warn_on_shared_identifier_prefixes(registered_registry.values()) # A discovery task may have published into ``previous_registry`` while # this replacement was being staged. Reconcile every published entry # synchronously after the swap so a lost publication cannot also leave diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index d75cb269b39..e043708dd7a 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1458,7 +1458,7 @@ def warn_on_server_name_fields( _warn("server_name", server_name) -def _warn_on_shared_identifier_prefixes(servers: Iterable[MCPServer]) -> None: +def warn_on_shared_identifier_prefixes(servers: Iterable[MCPServer]) -> None: """Warn once per identifier that several servers share. ``get_server_prefix`` resolves alias first, so two servers sharing a @@ -2644,7 +2644,7 @@ class MCPServerManager: self.assign_unique_short_prefix(new_server) _warn_legacy_delegate_auth_if_applicable(new_server, source="config") _warn_config_id_jag_server_outruns_sso(new_server) - self._invalidate_server_definition_caches(server_id) + self.invalidate_server_definition_caches(server_id) self.config_mcp_servers[server_id] = new_server self._set_oauth_discovery_deferred( server_id, @@ -2844,7 +2844,7 @@ class MCPServerManager: mappings make ``_get_mcp_server_from_tool_name`` resolve to a prefix that no longer exists in the live registry. """ - self._invalidate_server_definition_caches(server.server_id) + self.invalidate_server_definition_caches(server.server_id) self.remove_server_tool_routing(server) def remove_server_tool_routing(self, server: MCPServer) -> None: @@ -3257,10 +3257,10 @@ class MCPServerManager: # `credentials` field is the only one still encrypted here). # Re-decrypting plaintext would zero the values, so build with # env_vars_are_encrypted=False. - self._warn_if_newly_blocked_stdio(mcp_server, None) + self.warn_if_newly_blocked_stdio(mcp_server, None) new_server: Final = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False) self.assign_unique_short_prefix(new_server) - self._invalidate_server_definition_caches(mcp_server.server_id) + self.invalidate_server_definition_caches(mcp_server.server_id) self.registry[mcp_server.server_id] = new_server await self.maybe_register_openapi_tools(new_server) self.prime_oauth_metadata_discovery(new_server) @@ -3297,7 +3297,7 @@ class MCPServerManager: previous_server=self.registry[mcp_server.server_id], ) self.assign_unique_short_prefix(new_server) - self._invalidate_server_definition_caches(mcp_server.server_id) + self.invalidate_server_definition_caches(mcp_server.server_id) self.registry[mcp_server.server_id] = new_server await self.maybe_register_openapi_tools(new_server) self.prime_oauth_metadata_discovery(new_server) @@ -4343,6 +4343,10 @@ class MCPServerManager: ) raise_classified_list_failure(e, server.name, suppress_challenge=server.is_dcr_bridge) + def clear_initialize_instructions(self) -> None: + self._upstream_initialize_instructions_by_server_id.clear() + self._upstream_initialize_instructions_probed_at.clear() + def invalidate_discovery_lists(self, server_id: str) -> None: self._upstream_initialize_instructions_by_server_id.pop(server_id, None) self._upstream_initialize_instructions_probed_at.pop(server_id, None) @@ -4350,7 +4354,7 @@ class MCPServerManager: self._resource_discovery_cache.invalidate(server_id) self._template_discovery_cache.invalidate(server_id) - def _invalidate_server_definition_caches(self, server_id: str) -> None: + def invalidate_server_definition_caches(self, server_id: str) -> None: from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( # noqa: PLC0415 # lazy: discoverable_endpoints lazily imports this module's manager singleton invalidate_oauth_metadata_cache, ) @@ -4389,7 +4393,7 @@ class MCPServerManager: return server.server_id, hashlib.sha256(material.encode()).hexdigest() @staticmethod - def _warn_if_newly_blocked_stdio(row: LiteLLM_MCPServerTable, previous: MCPServer | None) -> None: + def warn_if_newly_blocked_stdio(row: LiteLLM_MCPServerTable, previous: MCPServer | None) -> None: if previous is None or previous.transport != row.transport: warn_if_mcp_stdio_blocked(row.alias or row.server_name, row.transport) diff --git a/litellm/proxy/_experimental/mcp_server/tool_registry.py b/litellm/proxy/_experimental/mcp_server/tool_registry.py index 2adf5db8f67..830ea7293f5 100644 --- a/litellm/proxy/_experimental/mcp_server/tool_registry.py +++ b/litellm/proxy/_experimental/mcp_server/tool_registry.py @@ -1,6 +1,6 @@ import asyncio import json -from collections.abc import Callable, Iterator, Mapping +from collections.abc import Callable, Generator, Mapping from contextlib import contextmanager from contextvars import ContextVar from typing import TYPE_CHECKING, Any, Final @@ -40,7 +40,7 @@ class MCPToolRegistry: self.published_tools = tools @contextmanager - def catalog_scope(self, tools: Mapping[str, MCPTool]) -> Iterator[dict[str, MCPTool]]: + def catalog_scope(self, tools: Mapping[str, MCPTool]) -> Generator[dict[str, MCPTool]]: detached: Final = dict(tools) closed: Final = asyncio.Event() token: Final = self._catalog_tools.set((detached, closed)) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 4ab82b375fc..1b48ad67ff8 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -19,7 +19,7 @@ import functools import importlib import json import os -from collections.abc import AsyncIterator, Iterable, Mapping, Sequence +from collections.abc import AsyncGenerator, Iterable, Mapping, Sequence from contextlib import asynccontextmanager from dataclasses import dataclass from datetime import datetime, timedelta, timezone @@ -2227,7 +2227,7 @@ if MCP_AVAILABLE: server_id: str, user_api_key_dict: UserAPIKeyAuth, request: Request | None = None, - ) -> AsyncIterator[MCPServer]: + ) -> AsyncGenerator[MCPServer]: if await get_cached_temporary_mcp_server(server_id) is not None: yield await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request) return