From d3edb5839dea5c1cdbb42461766f8d1f1f29a2e4 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Tue, 22 Sep 2026 14:25:58 -0700 Subject: [PATCH] fix(mcp): keep refreshed OpenAPI operation membership authoritative --- .../proxy/_experimental/mcp_server/catalog.py | 18 +++++++- .../mcp_server/test_mcp_server_manager.py | 43 +++++++++++++++++++ 2 files changed, 60 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py index 35d6d266a5d..014d84875b5 100644 --- a/litellm/proxy/_experimental/mcp_server/catalog.py +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -204,6 +204,7 @@ class TargetCatalog: async def reload(self) -> None: from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + from litellm.proxy._experimental.mcp_server.utils import normalize_server_name async with self._refresh_lock: previous_config: Final = MappingProxyType( @@ -228,7 +229,22 @@ class TargetCatalog: try: with global_mcp_tool_registry.catalog_scope(initial_tools) as staged_tools: await self._reload() - live_routes: Final = self._unchanged_routing(previous_servers, self.manager.published_tool_routes) + refreshed_openapi_owners: Final = frozenset( + owner + 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) + ) + live_routes: Final = self._unchanged_routing( + previous_servers, + MappingProxyType( + { + name: owner + for name, owner in self.manager.published_tool_routes.items() + if normalize_server_name(owner) not in refreshed_openapi_owners + } + ), + ) concurrent_routes: Final = self._unchanged_routing( self.manager.config_mcp_servers | live_registry, MappingProxyType( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 6f903aebbf5..b3d80beec6b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -14773,6 +14773,49 @@ async def test_catalog_reload_retains_routes_discovered_for_a_late_server(change assert tool is None +@pytest.mark.asyncio +@pytest.mark.parametrize("retain_operation", [False, True]) +async def test_catalog_openapi_refresh_does_not_restore_removed_operations(tmp_path, monkeypatch, respx_mock, retain_operation): + from litellm.proxy._experimental.mcp_server import tool_registry + + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + manager: Final = MCPServerManager() + registry: Final = tool_registry.MCPToolRegistry() + monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry) + spec_path: Final = tmp_path / "openapi.json" + paths: Final = {"/removed": {"get": {"operationId": "removed"}}, "/retained": {"get": {"operationId": "retained"}}} + spec: Final = {"openapi": "3.0.0", "info": {"title": "Refresh", "version": "1"}, "paths": paths} + spec_path.write_text(json.dumps(spec)) + row: Final = LiteLLM_MCPServerTable(server_id="spec-refresh", alias="spec_refresh", + url="https://upstream.example", transport=MCPTransport.http, spec_path=str(spec_path)) + prisma: Final = MagicMock() + prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row]) + upstream: Final = respx_mock.get("https://upstream.example/retained").respond(200, json={"value": "retained"}) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + await manager.reload_servers_from_database() + assert manager.server_exposes_tool(manager.registry[row.server_id], "spec_refresh-removed") + assert registry.get_tool("spec_refresh-removed") is not None + paths.pop("/removed") + if not retain_operation: + paths.clear() + spec_path.write_text(json.dumps(spec)) + for _ in range(2): + await manager.reload_servers_from_database() + for name in ("removed", "spec_refresh-removed"): + assert name not in manager.published_tool_routes + assert not manager.server_exposes_tool(manager.registry[row.server_id], name) + assert registry.get_tool("spec_refresh-removed") is None + retained = registry.get_tool("spec_refresh-retained") + if retain_operation: + assert manager.server_exposes_tool(manager.registry[row.server_id], "spec_refresh-retained") + assert retained is not None + assert json.loads(await retained.handler()) == {"value": "retained"} + else: + assert retained is None + assert manager.published_tool_routes == {} + assert upstream.call_count == (2 if retain_operation else 0) + + @pytest.mark.asyncio async def test_catalog_lookup_uses_one_snapshot_until_operation_finishes(): from datetime import timedelta