fix(mcp): preserve discovery published during catalog refresh

This commit is contained in:
Joshua Valluru 2026-09-22 13:37:09 -07:00
parent 74a470410d
commit ed7cdd151f
5 changed files with 134 additions and 16 deletions

View file

@ -133,21 +133,25 @@ class TargetCatalog:
self._retain_discovered_routing(snapshot)
def _retain_discovered_routing(self, snapshot: CatalogSnapshot) -> None:
self.manager.published_tool_routes = self.manager.published_tool_routes | self._unchanged_routing(
snapshot.servers, snapshot.routing
)
def _unchanged_routing(
self, servers: Mapping[str, MCPServer], routing: Mapping[str, str]
) -> MappingProxyType[str, str]:
from litellm.proxy._experimental.mcp_server.utils import normalize_server_name
current: Final = self.manager.config_mcp_servers | self.manager.registry
unchanged_owners: Final = frozenset(
owner
for key, server in snapshot.servers.items()
if current.get(key) == 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)
)
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
}
return MappingProxyType(
{name: owner for name, owner in routing.items() if normalize_server_name(owner) in unchanged_owners}
)
async def resolve(self, identifier: str, client_ip: str | None = None) -> MCPServer | None:
@ -202,9 +206,18 @@ class TargetCatalog:
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()
}
previous_config: Final = MappingProxyType(
{key: value.model_copy(deep=True) for key, value in self.manager.config_mcp_servers.items()}
)
previous_servers: Final = previous_config | MappingProxyType(
{key: value.model_copy(deep=True) for key, value in self.manager.registry.items()}
)
config_identities: Final = MappingProxyType(
{key: _configuration_identity(value) for key, value in previous_config.items()}
)
staged_config: Final = MappingProxyType(
{key: value.model_copy(deep=True) for key, value in previous_config.items()}
)
await self.manager.hydrate_config_servers_dcr_clients(tuple(staged_config.values()))
staged_routing: Final = dict(self.manager.published_tool_routes)
closed: Final = asyncio.Event()
@ -212,9 +225,24 @@ class TargetCatalog:
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
concurrent_routes: Final = self._unchanged_routing(
previous_servers, self.manager.published_tool_routes
)
retained_staged_routes: Final = self._unchanged_routing(
previous_config | self.manager.registry, staged_routing
)
self.manager.config_mcp_servers = {
key: value.model_copy(
update=staged_config[key].model_dump(
include=frozenset(("client_id", "client_secret", "token_endpoint_auth_method"))
)
)
if config_identities.get(key) == _configuration_identity(value)
else value
for key, value in self.manager.config_mcp_servers.items()
}
global_mcp_tool_registry.tools = staged_tools
self.manager.published_tool_routes = staged_routing
self.manager.published_tool_routes = concurrent_routes | retained_staged_routes
finally:
closed.set()
self._staged_routing.reset(routing_token)

View file

@ -2,7 +2,7 @@
import os
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
from contextlib import asynccontextmanager
from contextlib import asynccontextmanager, nullcontext
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
@ -909,6 +909,7 @@ async def test_get_tools_from_mcp_servers():
# Create a mock manager
mock_manager = AsyncMock()
mock_manager.catalog.operation = nullcontext
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["server1_id", "server2_id"]
)
@ -937,6 +938,7 @@ async def test_get_tools_from_mcp_servers():
# Test Case 2: Without specific MCP servers
# Create a different mock manager for the second test case
mock_manager_2 = AsyncMock()
mock_manager_2.catalog.operation = nullcontext
mock_manager_2.get_allowed_mcp_servers = AsyncMock(
return_value=["server1_id", "server2_id"]
)
@ -984,6 +986,7 @@ async def test_get_tools_from_mcp_servers():
# Test Case 3: With specific MCP servers and access groups
# Create a mock manager
mock_manager = AsyncMock()
mock_manager.catalog.operation = nullcontext
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["server1_id", "server2_id", "server3_id"]
)

View file

@ -66,8 +66,7 @@ async def _call_block(logging_obj, order: list, *, user_api_key_auth=mock.sentin
proxy_logging_obj = mock.MagicMock()
proxy_logging_obj.post_call_failure_hook.side_effect = _record_post_call_failure_hook
fake_proxy_server = types.ModuleType("litellm.proxy.proxy_server")
fake_proxy_server.proxy_logging_obj = proxy_logging_obj # pyright: ignore[reportAttributeAccessIssue]
fake_proxy_server = types.SimpleNamespace(proxy_logging_obj=proxy_logging_obj, prisma_client=None)
with mock.patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}):
with contextlib.suppress(HTTPException):

View file

@ -14519,6 +14519,92 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie
auth_context_var.reset(token)
@pytest.mark.asyncio
@pytest.mark.parametrize("anchored", [False, True])
async def test_catalog_reload_preserves_concurrent_config_discovery_and_routes(anchored):
manager: Final = MCPServerManager()
server: Final = MCPServer(server_id="config-race", name="config_race", transport=MCPTransport.http,
url="https://upstream.example/mcp", auth_type=MCPAuth.oauth2,
issuer="https://issuer.example" if anchored else None, issuer_is_anchored=anchored)
manager.config_mcp_servers = {server.server_id: server}
manager._set_oauth_discovery_deferred(server.server_id, True)
generation: Final = manager.oauth_discovery_slot(server.server_id).generation
resolved: Final = server.model_copy(update={"authorization_url": "https://issuer.example/authorize",
"token_url": "https://issuer.example/token", "scopes": ["read"], "issuer": "https://issuer.example"})
async def hydrate(target: MCPServer) -> bool:
target.client_id = "persisted-client"
return True
async def read_rows(**_kwargs):
assert manager._publish_resolved_oauth_server(resolved, generation) is resolved
manager.published_tool_routes["config_race-search"] = "config_race"
return []
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=read_rows)
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.hydrate_config_server_dcr_client", side_effect=hydrate),
):
await manager.reload_servers_from_database()
current: Final = manager.get_mcp_server_by_id(server.server_id)
assert current.authorization_url == resolved.authorization_url
assert current.token_url == resolved.token_url
assert current.scopes == ["read"]
assert current.issuer == "https://issuer.example"
assert current.client_id == "persisted-client"
assert manager.oauth_discovery_slot(server.server_id) is None
assert manager.published_tool_routes["config_race-search"] == "config_race"
@pytest.mark.asyncio
@pytest.mark.parametrize("change", ["credentials", "delete"])
async def test_catalog_reload_does_not_restore_replaced_config_credentials_or_routes(change):
manager: Final = MCPServerManager()
server: Final = MCPServer(server_id="config-replaced", name="config_replaced", transport=MCPTransport.http,
client_id="previous-client")
manager.config_mcp_servers = {server.server_id: server}
manager.published_tool_routes["config_replaced-search"] = "config_replaced"
async def read_rows(**_kwargs):
manager.config_mcp_servers = {} if change == "delete" else {
server.server_id: server.model_copy(update={"client_id": "new-client"})}
return []
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=read_rows)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
await manager.reload_servers_from_database()
current: Final = manager.get_mcp_server_by_id(server.server_id)
if change == "delete":
assert current is None
else:
assert current.client_id == "new-client"
assert "config_replaced-search" not in manager.published_tool_routes
@pytest.mark.asyncio
@pytest.mark.parametrize("change", ["discovery", "update", "delete"])
async def test_catalog_operation_retains_routes_only_for_same_configured_target(change):
from mcp.types import Tool
manager: Final = MCPServerManager()
server: Final = MCPServer(server_id="route-race", name="route_race", transport=MCPTransport.http,
url="https://upstream.example/mcp", auth_type=MCPAuth.oauth2)
manager.registry = {server.server_id: server}
with patch("litellm.proxy.proxy_server.prisma_client", None):
async with manager.catalog.operation():
tools: Final = manager._create_prefixed_tools([Tool(name="search", inputSchema={})], server)
if change == "delete":
manager.registry = {}
else:
manager.registry[server.server_id] = server.model_copy(update=(
{"authorization_url": "https://issuer.example/authorize", "token_url": "https://issuer.example/token"}
if change == "discovery" else {"url": "https://changed.example/mcp"}))
assert (tools[0].name in manager.published_tool_routes) is (change == "discovery")
@pytest.mark.asyncio
async def test_catalog_observes_committed_update_and_delete_without_background_reload():
from datetime import timedelta

View file

@ -1,5 +1,6 @@
import sys
import types
from contextlib import nullcontext
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
@ -77,6 +78,7 @@ def _mock_mcp_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock:
"""Patch the MCP tool-call plumbing so _execute_tool_calls can run in tests."""
call_tool = AsyncMock(return_value=CallToolResult(content=[TextContent(type="text", text="ok")], isError=False))
fake_manager = types.SimpleNamespace(
catalog=types.SimpleNamespace(operation=nullcontext),
get_registry=MagicMock(return_value={}),
call_tool=call_tool,
_get_mcp_server_from_tool_name=MagicMock(return_value=None),