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.
This commit is contained in:
Yucheng He 2026-10-01 02:43:31 -07:00
parent 4b0b326dc2
commit 9f46ec11b2
2 changed files with 57 additions and 4 deletions

View file

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

View file

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