From 9f46ec11b295fbc8b353145bba7fb92a5e439922 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Thu, 1 Oct 2026 02:43:31 -0700 Subject: [PATCH] fix(mcp): drop a listed catalog recorded across a server save _record_listed_tools ran after the awaited upstream fetch, so a PUT /v1/mcp/server that landed mid-fetch had its invalidation undone when the fetch completed: hooks then saw the pre-save description next to the post-save definition until the next listing, instead of name and arguments only. The manager now keeps a per-server listed-tools generation, bumped by _invalidate_server_definition_caches. _get_tools_from_server reads it before the fetch and _record_listed_tools skips the write when it moved; the next listing records normally. --- .../mcp_server/mcp_server_manager.py | 23 +++++++++-- .../mcp_server/test_mcp_server_manager.py | 38 +++++++++++++++++++ 2 files changed, 57 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index f7638846a81..eff17696cb0 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -2035,6 +2035,7 @@ class MCPServerManager: } """ self._listed_tools_by_server_id: dict[str, _ListedToolsByCaller] = {} # mutable-ok: refreshed per tools/list + self._listed_tools_generations: Mapping[str, int] = MappingProxyType({}) self._upstream_initialize_instructions_by_server_id: dict[str, str] = {} # Per-server monotonic timestamp of last upstream prefetch attempt (success, # empty result, or failure). Used to throttle re-probes for servers that do @@ -4504,6 +4505,7 @@ class MCPServerManager: raw_headers=raw_headers, oauth2_headers=oauth2_headers, ) + listed_generation: Final = self._listed_tools_generations.get(server.server_id, 0) try: # Tool *listing* must not be blocked by missing per-user env vars — @@ -4602,7 +4604,7 @@ class MCPServerManager: # through _create_prefixed_tools — that would add the prefix a second # time producing "test_petstore-test_petstore-getinventory". unprefixed_tools: Final = guarded_openapi - self._record_listed_tools(server, unprefixed_tools, listed_caller) + self._record_listed_tools(server, unprefixed_tools, listed_caller, listed_generation) if not add_prefix: return unprefixed_tools return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi] @@ -4618,7 +4620,7 @@ class MCPServerManager: raw_headers=raw_headers, ) prefixed_or_original_tools: Final = self._create_prefixed_tools( - guarded_tools, server, add_prefix=add_prefix, caller=listed_caller + guarded_tools, server, add_prefix=add_prefix, caller=listed_caller, generation=listed_generation ) return prefixed_or_original_tools @@ -4670,6 +4672,9 @@ class MCPServerManager: self._invalidate_discovery_lists(server_id) self._listed_tools_by_server_id.pop(server_id, None) + self._listed_tools_generations = MappingProxyType( + {**self._listed_tools_generations, server_id: self._listed_tools_generations.get(server_id, 0) + 1} + ) invalidate_oauth_metadata_cache(server_id) def _listed_tools_identity(self, server: MCPServer, caller: ListedToolsCaller | None) -> str | None: @@ -4715,8 +4720,17 @@ class MCPServerManager: ) def _record_listed_tools( - self, server: MCPServer, tools: Sequence[MCPTool], caller: ListedToolsCaller | None + self, + server: MCPServer, + tools: Sequence[MCPTool], + caller: ListedToolsCaller | None, + generation: int | None = None, ) -> None: + """Store the catalog served to ``caller``. ``generation`` is the server's listed-tools generation + read before the listing's upstream fetch; a server save that landed mid-fetch moved it, and the + pre-save catalog is then dropped rather than written over the invalidation.""" + if generation is not None and generation != self._listed_tools_generations.get(server.server_id, 0): + return identity: Final = self._listed_tools_identity(server, caller) listing: Final = MappingProxyType({tool.name: tool for tool in tools}) existing: Final = self._listed_tools_by_server_id.get(server.server_id, MappingProxyType({})) @@ -5638,6 +5652,7 @@ class MCPServerManager: server: MCPServer, add_prefix: bool = True, caller: ListedToolsCaller | None = None, + generation: int | None = None, ) -> list[MCPTool]: """ Create prefixed tools and update tool mapping. @@ -5663,7 +5678,7 @@ class MCPServerManager: for spelling in iter_known_tool_name_spellings(original_name, server): self.tool_name_to_mcp_server_name_mapping[spelling] = prefix - self._record_listed_tools(server, tools, caller) + self._record_listed_tools(server, tools, caller, generation) verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name) return prefixed_tools 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 8756ec421ae..e79c855b21b 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 @@ -7194,6 +7194,44 @@ class TestMCPServerManager: kept = manager.get_listed_tool(other, "ping") assert kept is not None and kept.description == "kept" + @pytest.mark.asyncio + async def test_server_save_during_an_in_flight_listing_is_not_undone_by_the_stale_record(self): + """A PUT /v1/mcp/server that lands while a listing awaits its upstream fetch drops the server's + catalog; the fetch completing afterwards must not write the pre-save catalog back, or hooks see + the old description next to the new definition until the next listing.""" + manager = MCPServerManager() + server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") + user = UserAPIKeyAuth(api_key="sk-litellm", user_id="saver") + fetch_started = asyncio.Event() + release_fetch = asyncio.Event() + + async def fetch(client, name): + fetch_started.set() + await release_fetch.wait() + return [MCPTool(name="turn", description="before save", inputSchema={})] + + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + manager._fetch_tools_with_timeout = fetch + caller = ListedToolsCaller(user_api_key_auth=user) + + async def list_tools() -> None: + await manager._get_tools_from_server(server=server, user_api_key_auth=user) + + listing = asyncio.create_task(list_tools()) + await fetch_started.wait() + manager._invalidate_server_definition_caches(server.server_id) + release_fetch.set() + await listing + + assert manager.get_listed_tool(server, "turn", caller) is None + + manager._fetch_tools_with_timeout = AsyncMock( + return_value=[MCPTool(name="turn", description="after save", inputSchema={})] + ) + await manager._get_tools_from_server(server=server, user_api_key_auth=user) + listed = manager.get_listed_tool(server, "turn", caller) + assert listed is not None and listed.description == "after save" + @pytest.mark.asyncio async def test_user_oauth_refresh_keeps_listed_tools(self): manager = MCPServerManager()