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.
This commit is contained in:
Yucheng He 2026-10-01 12:52:57 -07:00
parent 71be87b893
commit 46639297ad
2 changed files with 64 additions and 3 deletions

View file

@ -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.

View file

@ -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()