From 46639297ad7b52c4a2f2c5fb3704ac78864b1d07 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Thu, 1 Oct 2026 12:52:57 -0700 Subject: [PATCH] fix(mcp): keep discovery and OAuth metadata caches across an OpenAPI spec re-read add_server and update_server ran the full server-definition invalidation a second time after the awaited OpenAPI spec fetch, which also dropped the prompts/resources/templates discovery entries and the OAuth protected-resource metadata filled under the already-published definition, so the next request went upstream again. Only the listed-tool catalog recorded during the fetch holds pre-save entries, so the post-fetch pass now drops just that catalog and bumps its generation via the new _drop_listed_tools helper, which the full invalidation also calls. --- .../mcp_server/mcp_server_manager.py | 9 ++- .../mcp_server/test_mcp_server_manager.py | 58 +++++++++++++++++++ 2 files changed, 64 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 20d6128d3fb..ce9a5f62bf2 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3352,7 +3352,7 @@ class MCPServerManager: self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) if new_server.spec_path: - self._invalidate_server_definition_caches(mcp_server.server_id) + self._drop_listed_tools(mcp_server.server_id) self.prime_oauth_metadata_discovery(new_server) verbose_logger.debug("Added MCP Server: %s", new_server.name) @@ -3391,7 +3391,7 @@ class MCPServerManager: self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) if new_server.spec_path: - self._invalidate_server_definition_caches(mcp_server.server_id) + self._drop_listed_tools(mcp_server.server_id) self.prime_oauth_metadata_discovery(new_server) verbose_logger.debug("Updated MCP Server: %s", new_server.name) @@ -4675,9 +4675,12 @@ class MCPServerManager: ) self._invalidate_discovery_lists(server_id) + self._drop_listed_tools(server_id) + invalidate_oauth_metadata_cache(server_id) + + def _drop_listed_tools(self, server_id: str) -> None: self._listed_tools_by_server_id.pop(server_id, None) 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: """Key the listed-tool cache by every request input that can change the served catalog. 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 3ad7652fb8c..632e3415245 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 @@ -5,6 +5,7 @@ import json import logging import os import sys +import time from collections.abc import AsyncIterator from datetime import datetime from pathlib import Path @@ -40,6 +41,7 @@ from mcp.types import Tool as MCPTool from pydantic import AnyUrl, TypeAdapter from litellm.constants import MCP_METADATA_TIMEOUT +from litellm.proxy._experimental.mcp_server import discoverable_endpoints from litellm.proxy._experimental.mcp_server.tool_outcome import TextResult from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( ListedToolsCaller, @@ -7267,6 +7269,62 @@ class TestMCPServerManager: assert manager.registry["srv"] is new assert manager.get_listed_tool(new, "search", caller) is None + @pytest.mark.asyncio + @pytest.mark.parametrize("already_registered", [False, True], ids=["add_server", "update_server"]) + async def test_openapi_spec_re_read_keeps_discovery_and_oauth_metadata_filled_during_the_fetch( + self, already_registered: bool + ): + """The listed-tool catalog recorded during the spec fetch holds pre-save entries, but a prompts + discovery or OAuth protected-resource fetch answered in that window already saw the published + definition; dropping those too sends the next request upstream again.""" + manager = MCPServerManager() + if already_registered: + manager.registry["srv"] = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://old", spec_path="/old.json" + ) + new = MCPServer( + server_id="srv", name="srv", transport=MCPTransport.http, url="http://new", spec_path="/new.json" + ) + caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="lister")) + metadata_key: Final = (new.server_id, new.url) + prompt_fetches = 0 + + async def fetch_prompts() -> list[Prompt]: + nonlocal prompt_fetches + prompt_fetches += 1 + return [Prompt(name="greet")] + + async def register_while_discovery_fills(server: MCPServer, *, initialize_mapping: bool = True) -> None: + manager._record_listed_tools( + server, + [MCPTool(name="search", description="pre-save", inputSchema={})], + caller, + manager._listed_tools_generations.get(server.server_id, 0), + ) + await manager._prompt_discovery_cache.get((server.server_id, None), fetch_prompts) + discoverable_endpoints._OAUTH_METADATA_CACHE[metadata_key] = (time.time() + 300, {"resource": new.url}) + + manager.build_mcp_server_from_table = AsyncMock(return_value=new) + manager._maybe_register_openapi_tools = register_while_discovery_fills + manager.prime_oauth_metadata_discovery = MagicMock() + record = LiteLLM_MCPServerTable( + server_id="srv", server_name="srv", url="http://new", transport=MCPTransport.http + ) + save = manager.update_server if already_registered else manager.add_server + + try: + await save(record) + + assert manager.registry["srv"] is new + assert manager.get_listed_tool(new, "search", caller) is None + prompts = await manager._prompt_discovery_cache.get((new.server_id, None), fetch_prompts) + assert [prompt.name for prompt in prompts] == ["greet"] + assert prompt_fetches == 1, "the prompts list filled after the save was published went upstream again" + cached_metadata = discoverable_endpoints._OAUTH_METADATA_CACHE.get(metadata_key) + assert cached_metadata is not None and cached_metadata[1] == {"resource": new.url} + finally: + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(metadata_key, None) + @pytest.mark.asyncio async def test_user_oauth_refresh_keeps_listed_tools(self): manager = MCPServerManager()