diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py index 606ceb22444..f319d0530da 100644 --- a/litellm/proxy/_experimental/mcp_server/catalog.py +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -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( diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 8e7f694853b..f3299a03686 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index 6c6c1717777..036a1c80042 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -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, },