fix(mcp): keep bridge tool metadata request-local

This commit is contained in:
Devin AI 2026-10-02 18:22:00 +00:00
parent b0e6f66190
commit 68eec0b3c5
7 changed files with 134 additions and 73 deletions

View file

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

View file

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

View file

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

View file

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

View file

@ -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"),

View file

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

View file

@ -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={}),