mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
refactor(mcp): satisfy catalog refresh lint and type gates
This commit is contained in:
parent
27d787d410
commit
515c124824
4 changed files with 63 additions and 47 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue