fix(mcp): retain rediscovered routes across concurrent updates

This commit is contained in:
Joshua Valluru 2026-10-05 16:43:42 -07:00
parent c72b7b8824
commit 47e47f62ae
3 changed files with 43 additions and 8 deletions

View file

@ -5,7 +5,8 @@ from __future__ import annotations
import asyncio
import hashlib
import json
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence
from collections import UserDict
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, MutableMapping, Sequence
from contextlib import asynccontextmanager
from contextvars import ContextVar
from dataclasses import dataclass, replace
@ -27,12 +28,23 @@ _P = ParamSpec("_P")
_R = TypeVar("_R")
class _OperationRoutes(UserDict[str, str]):
def __init__(self, initial: Mapping[str, str]) -> None:
super().__init__()
self.data = dict(initial)
self.written_names: set[str] = set() # mutable-ok: record route writes without copying the journal per tool
def __setitem__(self, name: str, owner: str) -> None:
self.data[name] = owner
self.written_names.add(name)
@dataclass(frozen=True, slots=True)
class CatalogSnapshot:
servers: Mapping[str, MCPServer]
identity: str
tools: Mapping[str, MCPTool]
routing: dict[str, str] # mutable-ok: tool routes are remapped in place through this snapshot's live dict
routing: MutableMapping[str, str]
def _configuration_identity(server: MCPServer) -> str:
@ -96,7 +108,7 @@ class TargetCatalog:
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]: # mutable-ok: returns the live routing dict that callers remap in place
def routing(self) -> MutableMapping[str, str]:
staged: Final = self._staged_routing.get()
if staged is not None and not staged[1].is_set():
return staged[0]
@ -166,10 +178,11 @@ class TargetCatalog:
shared: Final = await self._fresh_snapshot()
initial_routing: Final = self._unchanged_routing(shared.servers, self.manager.published_tool_routes)
routing: Final = _OperationRoutes(initial_routing)
snapshot: Final = replace(
shared,
servers=MappingProxyType({key: value.model_copy(deep=True) for key, value in shared.servers.items()}),
routing=dict(initial_routing),
routing=routing,
)
closed: Final = asyncio.Event()
token: Final = self._operation.set((snapshot, closed))
@ -179,12 +192,12 @@ class TargetCatalog:
finally:
closed.set()
self._operation.reset(token)
self._retain_discovered_routing(snapshot, initial_routing)
self._retain_discovered_routing(snapshot, frozenset(routing.written_names))
def _retain_discovered_routing(self, snapshot: CatalogSnapshot, initial_routing: Mapping[str, str]) -> None:
def _retain_discovered_routing(self, snapshot: CatalogSnapshot, written_names: frozenset[str]) -> None:
self.manager.published_tool_routes = self.manager.published_tool_routes | self._unchanged_routing(
snapshot.servers,
{name: owner for name, owner in snapshot.routing.items() if initial_routing.get(name) != owner},
{name: owner for name, owner in snapshot.routing.items() if name in written_names},
)
def _unchanged_routing(

View file

@ -2417,7 +2417,7 @@ class MCPServerManager:
@property
def tool_name_to_mcp_server_name_mapping(
self,
) -> dict[str, str]: # mutable-ok: callers remap tool routes through this dict
) -> MutableMapping[str, str]:
return self.catalog.routing()
@tool_name_to_mcp_server_name_mapping.setter

View file

@ -16716,6 +16716,28 @@ async def test_catalog_delete_drops_derived_tool_mapping(monkeypatch):
assert not await manager.catalog.list()
@pytest.mark.asyncio
@pytest.mark.parametrize("rediscover", [False, True])
async def test_catalog_publishes_rediscovered_routes_despite_concurrent_owner_change(monkeypatch, rediscover):
from mcp.types import Tool
from litellm.proxy import proxy_server
monkeypatch.setattr(proxy_server, "prisma_client", None)
manager = MCPServerManager()
selected = MCPServer(server_id="selected", name="selected", transport=MCPTransport.http)
competing = MCPServer(server_id="competing", name="competing", transport=MCPTransport.http)
manager.config_mcp_servers = {server.server_id: server for server in (selected, competing)}
manager.published_tool_routes = {"shared_tool": selected.name}
async with manager.catalog.operation():
manager.published_tool_routes["shared_tool"] = competing.name
if rediscover:
manager._create_prefixed_tools([Tool(name="shared_tool", input_schema={})], selected)
assert manager._get_mcp_server_from_tool_name("shared_tool").server_id == selected.server_id
async with manager.catalog.operation():
expected = selected if rediscover else competing
assert manager._get_mcp_server_from_tool_name("shared_tool").server_id == expected.server_id
@pytest.mark.asyncio
@pytest.mark.parametrize("anchored", [False, True])
async def test_catalog_reload_preserves_concurrent_config_discovery_and_routes(anchored):