fix(mcp): record only tools served by the bridge

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-10-02 08:58:48 +00:00
parent 8ffd7c107a
commit b0e6f66190
3 changed files with 104 additions and 18 deletions

View file

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

View file

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

View file

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