diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index eff17696cb0..20d6128d3fb 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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(), ) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index be2ca076de8..151e3a87939 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -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, 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 e79c855b21b..3ad7652fb8c 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 @@ -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()