From b0e6f66190803425c634154c00aeed36883847c2 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 08:58:48 +0000 Subject: [PATCH] fix(mcp): record only tools served by the bridge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../_experimental/mcp_server/operations.py | 62 ++++++++++++++++--- .../mcp/litellm_proxy_mcp_handler.py | 8 +++ .../mcp/test_litellm_proxy_mcp_handler.py | 52 +++++++++++++--- 3 files changed, 104 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index b60a6c99fd0..417d5280ef4 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -4,7 +4,7 @@ import asyncio import traceback import types import uuid -from collections.abc import Mapping, Sequence +from collections.abc import Callable, Mapping, Sequence from datetime import datetime from typing import Any, Final, NoReturn, TypeAlias, overload @@ -81,6 +81,7 @@ 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, @@ -957,6 +958,7 @@ async def _get_tools_from_mcp_servers( mcp_proxy_mode: bool = False, *, record_listing: bool = False, + served_tool_selector: Callable[[list[MCPTool]], list[MCPTool]] | None = None, ) -> AggregateToolListing: """ Helper method to fetch tools from MCP servers based on server filtering criteria. @@ -969,6 +971,7 @@ async def _get_tools_from_mcp_servers( oauth2_headers: Optional dict of oauth2 headers record_listing: Record each served catalog into the caller's listed-tools slot; only a listing actually served to the caller sets it + served_tool_selector: Select existing tool objects from the aggregate for deferred recording Returns: AggregateToolListing: Combined tools from filtered servers plus each server's @@ -1066,12 +1069,12 @@ async def _get_tools_from_mcp_servers( async def _fetch_and_filter_server_tools( server: MCPServer, - ) -> "tuple[list[MCPTool], ServerOutcome]": + ) -> tuple[list[MCPTool], ServerOutcome, Callable[[frozenset[int]], None] | None]: """Fetch and filter tools from a single server, classifying any failure into that server's outcome so the aggregate can keep serving the healthy subset without a broken server masquerading as an empty one.""" if server is None: - return [], ServerListOk(tool_count=0) + return [], ServerListOk(tool_count=0), None server_auth_header, extra_headers = _prepare_mcp_server_headers( server=server, @@ -1123,6 +1126,18 @@ async def _get_tools_from_mcp_servers( try: from litellm.proxy.proxy_server import proxy_logging_obj + defer_recording: Final = record_listing and served_tool_selector is not None + listed_generation: Final = ( + global_mcp_server_manager._listed_tools_generations.get(server.server_id, 0) + if defer_recording + else None + ) + listed_caller: Final = ListedToolsCaller( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=catalog_auth_header, + raw_headers=raw_headers, + oauth2_headers=oauth2_headers, + ) tools: Final = await global_mcp_server_manager._get_tools_from_server( server=server, mcp_auth_header=server_auth_header, @@ -1134,7 +1149,7 @@ async def _get_tools_from_mcp_servers( oauth2_headers=oauth2_headers, proxy_logging_obj=proxy_logging_obj, catalog_auth_header=catalog_auth_header, - record_listing=record_listing, + record_listing=record_listing and served_tool_selector is None, ) filtered_tools = filter_tools_by_allowed_tools(tools, server) @@ -1144,6 +1159,14 @@ async def _get_tools_from_mcp_servers( user_api_key_auth=user_api_key_auth, ) + unprefixed_tools: Final = ( + tuple( + tool.model_copy(update={"name": strip_known_server_prefix(tool.name, server) or tool.name}) + for tool in filtered_tools + ) + if defer_recording + else () + ) if mcp_proxy_mode: from litellm.proxy._experimental.mcp_server.tool_search import with_mcp_proxy_identity @@ -1151,13 +1174,27 @@ async def _get_tools_from_mcp_servers( else: filtered_tools = apply_display_name_overrides(filtered_tools, server) + catalog_entries: Final = tuple(zip(filtered_tools, unprefixed_tools)) + + def record_served_tools(served_ids: frozenset[int]) -> None: + global_mcp_server_manager._record_listed_tools( + server, + [original for exposed, original in catalog_entries if id(exposed) in served_ids], + listed_caller, + listed_generation, + ) + verbose_logger.debug( "Successfully fetched %s tools from server %s, %s after filtering", len(tools), server.name, len(filtered_tools), ) - return filtered_tools, ServerListOk(tool_count=len(filtered_tools)) + return ( + filtered_tools, + ServerListOk(tool_count=len(filtered_tools)), + record_served_tools if defer_recording else None, + ) except MCPUpstreamAuthError as e: # Absorb so one unauthenticated server does not empty every other server's # tools. Surfacing the upstream 401 to the client as a re-auth challenge is @@ -1166,20 +1203,25 @@ async def _get_tools_from_mcp_servers( # error). Single-server routes surface it via the request-scope preemptive # check in _raise_preemptive_401_for_unauthenticated_servers instead. verbose_logger.debug("MCP list_tools: omitting %s; it needs upstream auth", server.name) - return [], classify_list_exception(e) + return [], classify_list_exception(e), None except Exception as e: verbose_logger.exception("Error getting tools from server %s: %s", server.name, e) - return [], classify_list_exception(e) + return [], classify_list_exception(e), None # Fetch tools from all servers in parallel tasks: Final = [_fetch_and_filter_server_tools(server) for server in allowed_mcp_servers] results: Final = await asyncio.gather(*tasks) # Flatten results into single list - all_tools: Final[list[MCPTool]] = [tool for tools, _ in results for tool in tools] + all_tools: Final[list[MCPTool]] = [tool for tools, _, _record in results for tool in tools] + if record_listing and served_tool_selector is not None: + served_ids: Final = frozenset(id(tool) for tool in served_tool_selector(all_tools)) + for _tools, _outcome, record in results: + if record is not None: + record(served_ids) server_outcomes: Final[dict[str, ServerOutcome]] = { _aggregate_server_key(server): outcome - for server, (_, outcome) in zip(allowed_mcp_servers, results) + for server, (_, outcome, _record) in zip(allowed_mcp_servers, results) if server is not None } @@ -1187,7 +1229,7 @@ async def _get_tools_from_mcp_servers( if litellm_logging_obj: per_server_tool_counts: Final[dict[str, int]] = { _aggregate_server_key(server): len(server_tools) - for server, (server_tools, _) in zip(allowed_mcp_servers, results) + for server, (server_tools, _, _record) in zip(allowed_mcp_servers, results) if server is not None } diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 2ca1c32f6dc..98de4ecd403 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -318,6 +318,13 @@ class LiteLLM_Proxy_MCP_Handler: # names), so use None and let the auth object's mcp_servers do the filtering. effective_server_filter: Final = None if resolved_toolset_ids else (resolved_mcp_servers or None) + def served_tools(tools: list[MCPTool]) -> list[MCPTool]: + filtered: Final = LiteLLM_Proxy_MCP_Handler._filter_mcp_tools_by_allowed_tools( + tools, mcp_tools_with_litellm_proxy + ) + deduplicated, _server_map = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools(filtered, []) + return deduplicated + listing: Final = await _get_tools_from_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, @@ -329,6 +336,7 @@ class LiteLLM_Proxy_MCP_Handler: request_tags=request_tags, raw_headers=raw_headers, record_listing=True, + served_tool_selector=served_tools, ) tools: Final = listing.tools diff --git a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py index 44b7034ca34..ea540088b27 100644 --- a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -1316,13 +1316,30 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py @pytest.mark.asyncio -async def test_get_mcp_tools_from_manager_records_the_served_catalog(monkeypatch: pytest.MonkeyPatch) -> None: +@pytest.mark.parametrize( + ("allowed_tools", "expected_names"), + [ + ([], ["responses_slot-echo", "responses_slot-status"]), + (["echo"], ["responses_slot-echo"]), + (["responses_slot-echo"], ["responses_slot-echo"]), + (["absent"], []), + ], +) +async def test_get_mcp_tools_from_manager_records_the_served_catalog( + monkeypatch: pytest.MonkeyPatch, allowed_tools: list[str], expected_names: list[str] +) -> None: """The Responses bridge serves the listing to the model and its own tools/call reads the slot, so this listing records the caller's catalog.""" manager: Final = mcp_operations.global_mcp_server_manager - server: Final = MCPServer(server_id="responses-slot", name="responses-slot", transport=MCPTransport.http) + server: Final = MCPServer( + server_id="responses-slot", name="responses_slot", alias="responses_slot", transport=MCPTransport.http + ) user: Final = UserAPIKeyAuth(api_key="sk-responses-slot", user_id="responder") - upstream: Final = [MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"})] + upstream: Final = [ + MCPTool(name="echo", description="Echo text back", inputSchema={"type": "object"}), + MCPTool(name="status", description="Report status", inputSchema={"type": "object"}), + MCPTool(name="echo", description="Duplicate echo", inputSchema={"type": "object", "properties": {}}), + ] fake_manager: Final = types.SimpleNamespace( get_registry=MagicMock(return_value={}), get_allowed_mcp_servers=AsyncMock(return_value=[]), @@ -1340,16 +1357,35 @@ async def test_get_mcp_tools_from_manager_records_the_served_catalog(monkeypatch patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)), ): try: - tools, _server_names = await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( + tools, _server_names = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( user_api_key_auth=user, - mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/responses-slot"}], + mcp_tools_with_litellm_proxy=[ + { + "type": "mcp", + "server_url": "litellm_proxy/mcp/responses-slot", + "allowed_tools": allowed_tools, + } + ], + ) + caller: Final = ListedToolsCaller(user_api_key_auth=user) + recorded: Final = { + tool.name: (listed.description, listed.input_schema) + for tool in upstream + if (listed := manager.get_listed_tool(server, tool.name, caller)) is not None + } + assert recorded == { + tool.name.removeprefix("responses_slot-"): (tool.description, tool.input_schema) for tool in tools + } + assert ( + manager.get_listed_tool( + server, "echo", ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-other-caller")) + ) + is None ) - listed: Final = manager.get_listed_tool(server, "echo", ListedToolsCaller(user_api_key_auth=user)) finally: manager._drop_listed_tools(server.server_id) - assert [tool.name for tool in tools] == ["responses-slot-echo"] - assert listed is not None and listed.description == "Echo text back" + assert [tool.name for tool in tools] == expected_names def _toolset_gateway_manager(toolset_id: str, server_id: str) -> types.SimpleNamespace: