fix(mcp): retain discovered routes across cached catalog operations

This commit is contained in:
Joshua Valluru 2026-10-03 17:46:31 -07:00
parent 886aadefef
commit 57dfed6e39
3 changed files with 49 additions and 6 deletions

View file

@ -167,10 +167,11 @@ class TargetCatalog:
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
shared: Final = await self._fresh_snapshot()
initial_routing: Final = self._unchanged_routing(shared.servers, self.manager.published_tool_routes)
snapshot: Final = replace(
shared,
servers=MappingProxyType({key: value.model_copy(deep=True) for key, value in shared.servers.items()}),
routing=dict(shared.routing),
routing=dict(initial_routing),
)
closed: Final = asyncio.Event()
token: Final = self._operation.set((snapshot, closed))
@ -180,11 +181,12 @@ class TargetCatalog:
finally:
closed.set()
self._operation.reset(token)
self._retain_discovered_routing(snapshot)
self._retain_discovered_routing(snapshot, initial_routing)
def _retain_discovered_routing(self, snapshot: CatalogSnapshot) -> None:
def _retain_discovered_routing(self, snapshot: CatalogSnapshot, initial_routing: Mapping[str, str]) -> None:
self.manager.published_tool_routes = self.manager.published_tool_routes | self._unchanged_routing(
snapshot.servers, snapshot.routing
snapshot.servers,
{name: owner for name, owner in snapshot.routing.items() if initial_routing.get(name) != owner},
)
def _unchanged_routing(

View file

@ -16990,3 +16990,44 @@ async def test_upstream_preparation_honors_case_sensitive_extra_command(monkeypa
client: Final = await MCPServerManager()._create_mcp_client(server)
assert client.stdio_config is not None
assert client.stdio_config["command"] == "/opt/tools/CustomRunner"
@pytest.mark.asyncio
@pytest.mark.parametrize("overlap", [False, True])
async def test_catalog_cached_revision_retains_newly_discovered_tool_routes(monkeypatch, overlap):
from mcp.types import Tool
read_rows = AsyncMock(return_value=[])
_catalog_database(monkeypatch, read_rows, AsyncMock(return_value=_revision_row(7)))
manager = MCPServerManager()
first = MCPServer(server_id="first", name="first", transport=MCPTransport.http)
second = MCPServer(server_id="second", name="second", transport=MCPTransport.http)
manager.config_mcp_servers = {server.server_id: server for server in (first, second)}
manager.published_tool_routes = {"search": "first"}
ready = asyncio.Event()
release = asyncio.Event()
async def read_catalog():
async with manager.catalog.operation():
assert manager._get_mcp_server_from_tool_name("search").server_id == "first"
ready.set()
await release.wait()
reader = asyncio.create_task(read_catalog()) if overlap else None
try:
if reader is not None:
await asyncio.wait_for(ready.wait(), 2)
async with manager.catalog.operation():
manager._create_prefixed_tools([Tool(name="search", input_schema={})], second)
assert manager.published_tool_routes["search"] == "second"
release.set()
if reader is not None:
await asyncio.wait_for(reader, 2)
async with manager.catalog.operation():
assert manager._get_mcp_server_from_tool_name("search").server_id == "second"
assert manager.published_tool_routes["search"] == "second"
read_rows.assert_awaited_once()
finally:
if reader is not None:
reader.cancel()
await asyncio.gather(reader, return_exceptions=True)

View file

@ -7735,8 +7735,8 @@ async def test_execute_mcp_tool_rest_server_id_injects_requested_server_credenti
),
patch.object(
mcp_operations.global_mcp_server_manager,
"get_registry",
return_value={
"registry",
{
requested_server.server_id: requested_server,
collision_server.server_id: collision_server,
},