refactor(mcp): satisfy catalog refresh lint and type gates

This commit is contained in:
Joshua Valluru 2026-10-03 14:25:57 -07:00
parent 27d787d410
commit 515c124824
4 changed files with 63 additions and 47 deletions

View file

@ -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

View file

@ -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)

View file

@ -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))

View file

@ -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