From 68eec0b3c5dde7a29d13c1b99e222a45837b0c73 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 2 Oct 2026 18:22:00 +0000 Subject: [PATCH] fix(mcp): keep bridge tool metadata request-local --- .../mcp_server/mcp_server_manager.py | 3 +- .../_experimental/mcp_server/operations.py | 62 ++-------- litellm/responses/main.py | 3 + .../responses/mcp/chat_completions_handler.py | 2 + .../mcp/litellm_proxy_mcp_handler.py | 17 ++- .../responses/mcp/mcp_streaming_iterator.py | 3 + .../mcp/test_litellm_proxy_mcp_handler.py | 117 ++++++++++++++++-- 7 files changed, 134 insertions(+), 73 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 1c6ac52f422..985e1a635e7 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -6664,6 +6664,7 @@ class MCPServerManager: wire_compat: WireCompat = WireCompat.LEGACY, *, catalog_auth_header: str | None | EllipsisType = ..., + listed_tool: MCPTool | None | EllipsisType = ..., ) -> CallToolResult | InputRequiredResult: """ Call a tool with the given name and arguments @@ -6717,7 +6718,7 @@ class MCPServerManager: raw_headers=raw_headers, litellm_logging_obj=litellm_logging_obj, guardrail_context=guardrail_context, - tool=self.get_listed_tool(mcp_server, name, listed_caller), + tool=self.get_listed_tool(mcp_server, name, listed_caller) if listed_tool is ... else listed_tool, ) if "arguments" in hook_result: arguments = hook_result["arguments"] diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 417d5280ef4..b60a6c99fd0 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 Callable, Mapping, Sequence +from collections.abc import Mapping, Sequence from datetime import datetime from typing import Any, Final, NoReturn, TypeAlias, overload @@ -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, @@ -958,7 +957,6 @@ 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. @@ -971,7 +969,6 @@ 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 @@ -1069,12 +1066,12 @@ async def _get_tools_from_mcp_servers( async def _fetch_and_filter_server_tools( server: MCPServer, - ) -> tuple[list[MCPTool], ServerOutcome, Callable[[frozenset[int]], None] | None]: + ) -> "tuple[list[MCPTool], ServerOutcome]": """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), None + return [], ServerListOk(tool_count=0) server_auth_header, extra_headers = _prepare_mcp_server_headers( server=server, @@ -1126,18 +1123,6 @@ 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, @@ -1149,7 +1134,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 and served_tool_selector is None, + record_listing=record_listing, ) filtered_tools = filter_tools_by_allowed_tools(tools, server) @@ -1159,14 +1144,6 @@ 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 @@ -1174,27 +1151,13 @@ 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)), - record_served_tools if defer_recording else None, - ) + return filtered_tools, ServerListOk(tool_count=len(filtered_tools)) 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 @@ -1203,25 +1166,20 @@ 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), None + return [], classify_list_exception(e) except Exception as e: verbose_logger.exception("Error getting tools from server %s: %s", server.name, e) - return [], classify_list_exception(e), None + return [], classify_list_exception(e) # 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, _, _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) + all_tools: Final[list[MCPTool]] = [tool for tools, _ in results for tool in tools] server_outcomes: Final[dict[str, ServerOutcome]] = { _aggregate_server_key(server): outcome - for server, (_, outcome, _record) in zip(allowed_mcp_servers, results) + for server, (_, outcome) in zip(allowed_mcp_servers, results) if server is not None } @@ -1229,7 +1187,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, _, _record) in zip(allowed_mcp_servers, results) + for server, (server_tools, _) in zip(allowed_mcp_servers, results) if server is not None } diff --git a/litellm/responses/main.py b/litellm/responses/main.py index d145f8cc6b8..68952d2e525 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -286,6 +286,7 @@ async def aresponses_api_with_mcp( call_params=call_params, previous_response_id=previous_response_id, tool_server_map=tool_server_map, + served_tools=original_mcp_tools, **kwargs, ) await mcp_streaming_response._create_initial_response_iterator() @@ -339,6 +340,7 @@ async def aresponses_api_with_mcp( tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( tool_server_map=tool_server_map, + served_tools=original_mcp_tools, tool_calls=tool_calls, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, @@ -395,6 +397,7 @@ async def aresponses_api_with_mcp( final_response = MCPEnhancedStreamingIterator( tool_server_map=tool_server_map, + served_tools=original_mcp_tools, base_iterator=final_response, mcp_events=tool_execution_events, user_api_key_auth=user_api_key_auth, diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index b17e0befba6..8aa4181cc8a 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -435,6 +435,7 @@ async def acompletion_with_mcp( # Execute tool calls self.tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( tool_server_map=self.tool_server_map, + served_tools=deduplicated_mcp_tools, tool_calls=self.tool_calls, user_api_key_auth=self.user_api_key_auth, mcp_auth_header=self.mcp_auth_header, @@ -609,6 +610,7 @@ async def acompletion_with_mcp( tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( tool_server_map=tool_server_map, tool_calls=tool_calls, + served_tools=deduplicated_mcp_tools, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 98de4ecd403..8e37466c6c9 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -318,13 +318,6 @@ 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, @@ -335,8 +328,6 @@ class LiteLLM_Proxy_MCP_Handler: litellm_trace_id=litellm_trace_id, request_tags=request_tags, raw_headers=raw_headers, - record_listing=True, - served_tool_selector=served_tools, ) tools: Final = listing.tools @@ -705,6 +696,7 @@ class LiteLLM_Proxy_MCP_Handler: litellm_trace_id: str | None = None, request_tags: list[str] | None = None, guardrail_context: Mapping[str, object] | None = None, + served_tools: Sequence[MCPTool] | None = None, ) -> list[MCPToolResult]: """Execute tool calls and return results.""" from fastapi import HTTPException @@ -869,6 +861,11 @@ class LiteLLM_Proxy_MCP_Handler: proxy_logging_obj=proxy_logging_obj, litellm_logging_obj=litellm_logging_obj, guardrail_context=guardrail_context, + listed_tool=( + next((tool for tool in served_tools if tool.name == tool_name), None) + if served_tools is not None + else ... + ), ) if proxy_logging_obj: @@ -1161,6 +1158,7 @@ class LiteLLM_Proxy_MCP_Handler: call_params: Mapping[str, object], previous_response_id: str | None, tool_server_map: dict[str, str], + served_tools: Sequence[MCPTool] | None = None, **kwargs, ) -> Any: """ @@ -1190,6 +1188,7 @@ class LiteLLM_Proxy_MCP_Handler: base_iterator=None, # Will be created internally mcp_events=mcp_discovery_events, # Pre-generated MCP discovery events tool_server_map=tool_server_map, + served_tools=served_tools, mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, user_api_key_auth=kwargs.get("user_api_key_auth") or kwargs.get("litellm_metadata", {}).get("user_api_key_auth"), diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index c0bb92cfb2a..b1f12233f33 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -281,6 +281,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]] | None = None, user_api_key_auth: "UserAPIKeyAuth | None" = None, original_request_params: dict[str, Any] | None = None, + served_tools: Sequence[MCPTool] | None = None, ): # MCP setup self.mcp_tools_with_litellm_proxy = mcp_tools_with_litellm_proxy or [] @@ -300,6 +301,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): self.mcp_discovery_generated = True # Events are already generated self.mcp_events = mcp_events # Store the initial MCP events for backward compatibility self.tool_server_map = tool_server_map + self.served_tools = tuple(served_tools) if served_tools is not None else None # Iterator references self.base_iterator: BaseResponsesAPIStreamingIterator | ResponsesAPIResponse | None = ( @@ -796,6 +798,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): # Execute the tools tool_results: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( tool_server_map=self.tool_server_map, + served_tools=self.served_tools, tool_calls=tool_calls, user_api_key_auth=self.user_api_key_auth, mcp_auth_header=self.mcp_auth_header, 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 ea540088b27..73bee304fc2 100644 --- a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -1,9 +1,10 @@ +import asyncio import importlib import subprocess import sys import textwrap import types -from typing import Any, Final, cast +from typing import Any, Final, Literal, cast from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -12,20 +13,26 @@ from mcp.types import CallToolResult, TextContent from mcp.types import Tool as MCPTool from openai.types.responses.tool_param import Mcp +import litellm +from litellm.caching.caching import DualCache +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.litellm_core_utils.litellm_logging import Logging from litellm.proxy._experimental.mcp_server import operations as mcp_operations from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing -from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller +from litellm.proxy._experimental.mcp_server.mcp_server_manager import ListedToolsCaller, MCPServerManager from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.utils import ProxyLogging from litellm.responses import main as responses_main from litellm.responses.mcp import litellm_proxy_mcp_handler as mcp_handler_module from litellm.responses.mcp.litellm_proxy_mcp_handler import ( LiteLLM_Proxy_MCP_Handler, ) +from litellm.types.guardrails import GuardrailEventHooks from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.mcp import MCPTransport from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.types.responses.main import OutputFunctionToolCall -from litellm.types.utils import ModelResponse +from litellm.types.utils import GenericGuardrailAPIInputs, ModelResponse class _DummyMCPResult: @@ -1316,6 +1323,7 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py @pytest.mark.asyncio +@pytest.mark.parametrize("real_listing", [False, True]) @pytest.mark.parametrize( ("allowed_tools", "expected_names"), [ @@ -1325,11 +1333,9 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py (["absent"], []), ], ) -async def test_get_mcp_tools_from_manager_records_the_served_catalog( - monkeypatch: pytest.MonkeyPatch, allowed_tools: list[str], expected_names: list[str] +async def test_bridge_listing_leaves_the_callers_catalog_unchanged( + monkeypatch: pytest.MonkeyPatch, allowed_tools: list[str], expected_names: list[str], real_listing: bool ) -> 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", alias="responses_slot", transport=MCPTransport.http @@ -1357,6 +1363,15 @@ async def test_get_mcp_tools_from_manager_records_the_served_catalog( patch.object(manager, "_fetch_tools_with_timeout", AsyncMock(return_value=upstream)), ): try: + if real_listing: + await manager._get_tools_from_server(server, user_api_key_auth=user, record_listing=True) + caller: Final = ListedToolsCaller(user_api_key_auth=user) + before: 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 bool(before) is real_listing tools, _server_names = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( user_api_key_auth=user, mcp_tools_with_litellm_proxy=[ @@ -1367,15 +1382,12 @@ async def test_get_mcp_tools_from_manager_records_the_served_catalog( } ], ) - 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 recorded == before assert ( manager.get_listed_tool( server, "echo", ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-other-caller")) @@ -1388,6 +1400,89 @@ async def test_get_mcp_tools_from_manager_records_the_served_catalog( assert [tool.name for tool in tools] == expected_names +class _BridgeMetadataGuardrail(CustomGuardrail): + def __init__(self) -> None: + super().__init__(guardrail_name="bridge-metadata", event_hook=GuardrailEventHooks.pre_mcp_call, default_on=True) + self.calls: tuple[tuple[object, object], ...] = () + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + logging_obj: Logging | None = None, + ) -> GenericGuardrailAPIInputs: + if request_data.get("mcp_arguments") == {"probe": "bridge"}: + self.calls += ((request_data.get("mcp_tool_description"), request_data.get("mcp_input_schema")),) + return inputs + + +@pytest.mark.asyncio +async def test_concurrent_bridge_calls_use_their_own_served_metadata(monkeypatch: pytest.MonkeyPatch) -> None: + manager: Final = MCPServerManager() + server: Final = MCPServer(server_id="bridge", name="bridge", transport=MCPTransport.http, url="http://upstream") + manager.registry = {server.server_id: server} + user: Final = UserAPIKeyAuth(api_key="sk-bridge", user_id="bridge-user") + upstream: Final = [ + MCPTool( + name="echo", + description="Echo text", + inputSchema={"type": "object", "properties": {"text": {"type": "string"}}}, + ), + MCPTool(name="status", description="Read status", inputSchema={"type": "object"}), + ] + client: Final = AsyncMock() + client.call_tool.return_value = CallToolResult(content=[TextContent(type="text", text="ok")]) + manager._create_mcp_client = AsyncMock(return_value=client) + manager._fetch_tools_with_timeout = AsyncMock(return_value=upstream) + guardrail: Final = _BridgeMetadataGuardrail() + logger: Final = ProxyLogging(user_api_key_cache=DualCache()) + monkeypatch.setattr(litellm, "callbacks", [guardrail]) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", logger) + monkeypatch.setattr(mcp_operations, "global_mcp_server_manager", manager) + monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])) + monkeypatch.setattr("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", manager) + first_listed: Final = asyncio.Event() + second_listed: Final = asyncio.Event() + + async def bridge(name: str, first: bool) -> None: + if not first: + await first_listed.wait() + tools, server_map = 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/bridge", "allowed_tools": [name]} + ], + ) + (first_listed if first else second_listed).set() + await second_listed.wait() + result: Final = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_server_map=server_map, + tool_calls=[ + {"type": "function_call", "name": f"bridge-{name}", "arguments": '{"probe":"bridge"}', "call_id": name} + ], + user_api_key_auth=user, + served_tools=tools, + ) + assert [entry["result"] for entry in result] == ["ok"] + + try: + await asyncio.gather(bridge("echo", True), bridge("status", False)) + assert sorted(guardrail.calls, key=str) == sorted( + ((tool.description, tool.input_schema) for tool in upstream), key=str + ) + await manager.call_tool("bridge", "echo", {"probe": "bridge"}, user_api_key_auth=user, proxy_logging_obj=logger) + assert guardrail.calls[-1] == (None, None) + await manager._get_tools_from_server( + server, user_api_key_auth=user, proxy_logging_obj=logger, record_listing=True + ) + await manager.call_tool("bridge", "echo", {"probe": "bridge"}, user_api_key_auth=user, proxy_logging_obj=logger) + assert guardrail.calls[-1] == (upstream[0].description, upstream[0].input_schema) + finally: + manager._drop_listed_tools(server.server_id) + ProxyLogging._callback_capabilities_cache.clear() + + def _toolset_gateway_manager(toolset_id: str, server_id: str) -> types.SimpleNamespace: return types.SimpleNamespace( get_registry=MagicMock(return_value={}),