fix(mcp): drop the catalog again once a saved OpenAPI server's registry is rebuilt

add_server and update_server publish the saved definition before the OpenAPI registry entries are
rebuilt from the spec, so a listing recorded during that fetch held the pre-save entries under the
new generation. The generation is bumped a second time after the registry refresh.

The during-hook task no longer accepts a listed entry, the one-line wrapper over get_listed_tool is
inlined at its two call sites, and the per-server generation map is a plain dict.
This commit is contained in:
Yucheng He 2026-10-01 03:57:06 -07:00
parent 9f46ec11b2
commit 71be87b893
3 changed files with 46 additions and 20 deletions

View file

@ -2035,7 +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._listed_tools_generations: dict[str, int] = {} # mutable-ok: bumped per server save
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
@ -3351,6 +3351,8 @@ class MCPServerManager:
self._invalidate_server_definition_caches(mcp_server.server_id)
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.prime_oauth_metadata_discovery(new_server)
verbose_logger.debug("Added MCP Server: %s", new_server.name)
@ -3388,6 +3390,8 @@ class MCPServerManager:
self._invalidate_server_definition_caches(mcp_server.server_id)
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.prime_oauth_metadata_discovery(new_server)
verbose_logger.debug("Updated MCP Server: %s", new_server.name)
@ -4672,9 +4676,7 @@ 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}
)
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:
@ -4727,8 +4729,7 @@ class MCPServerManager:
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."""
read before the listing's upstream fetch; the record is skipped when it no longer matches."""
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)
@ -6059,7 +6060,6 @@ class MCPServerManager:
start_time: datetime.datetime,
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
guardrail_context: Mapping[str, object] | None = None,
tool: MCPTool | None = None,
):
"""Create and return a during hook task for MCP tool calls.
@ -6074,8 +6074,6 @@ class MCPServerManager:
tool_name=name,
arguments=arguments,
server_name=server_name_from_prefix,
tool_description=tool.description if tool is not None else None,
tool_input_schema=tool.input_schema if tool is not None else None,
start_time=start_time.timestamp() if start_time else None,
hidden_params=HiddenParams(),
)

View file

@ -81,7 +81,6 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
outcome_wire_value,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
ListedToolsCaller,
MCPServerManager,
_caller_authorization_fans_out,
_client_forwarded_authorization_headers,
@ -1610,12 +1609,6 @@ async def _list_mcp_resource_templates(
return managed_resource_templates
def _registered_tool_metadata(name: str, server: MCPServer, caller: ListedToolsCaller) -> MCPTool | None:
"""The tool as ``tools/list`` served it to this caller (pinned, overridden, guardrail-masked), or None when
no listing was recorded so the call hands the hooks name and arguments only."""
return global_mcp_server_manager.get_listed_tool(server, name, caller)
def _resolve_display_name_to_original(
name: str,
allowed_mcp_servers: list[MCPServer],
@ -2086,9 +2079,9 @@ async def _execute_mcp_tool(
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
guardrail_context=guardrail_context,
tool=_registered_tool_metadata(
original_tool_name,
tool=global_mcp_server_manager.get_listed_tool(
mcp_server,
original_tool_name,
listed_tools_caller_for(
mcp_server,
user_api_key_auth,
@ -2210,9 +2203,9 @@ async def _execute_mcp_tool(
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
guardrail_context=guardrail_context,
tool=_registered_tool_metadata(
original_tool_name,
tool=global_mcp_server_manager.get_listed_tool(
prefix_server,
original_tool_name,
listed_tools_caller_for(
prefix_server,
user_api_key_auth,

View file

@ -7232,6 +7232,41 @@ class TestMCPServerManager:
listed = manager.get_listed_tool(server, "turn", caller)
assert listed is not None and listed.description == "after save"
@pytest.mark.asyncio
async def test_update_server_refreshing_openapi_tools_drops_a_listing_recorded_during_the_spec_fetch(self):
"""An OpenAPI server's registry entries are rebuilt after the save is published, so a listing that
records while the spec is fetched holds the pre-save entries; the catalog is dropped again once the
registry is current."""
manager = MCPServerManager()
old = MCPServer(
server_id="srv", name="srv", transport=MCPTransport.http, url="http://old", spec_path="/old.json"
)
manager.registry[old.server_id] = old
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"))
async def register_while_a_listing_records(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),
)
manager.build_mcp_server_from_table = AsyncMock(return_value=new)
manager._maybe_register_openapi_tools = register_while_a_listing_records
manager.prime_oauth_metadata_discovery = MagicMock()
record = LiteLLM_MCPServerTable(
server_id="srv", server_name="srv", url="http://new", transport=MCPTransport.http
)
await manager.update_server(record)
assert manager.registry["srv"] is new
assert manager.get_listed_tool(new, "search", caller) is None
@pytest.mark.asyncio
async def test_user_oauth_refresh_keeps_listed_tools(self):
manager = MCPServerManager()