fix(mcp): resolve responses-API tool dispatch by server_id and gate it on scope

The responses API surface still routed tool calls by server name. It built
the caller's reachable set as MCPServer objects, narrowed by both the key's
grants and the requested server filter, then discarded that identity by
flattening to display names. `tool_server_map` carried a name, and dispatch
re-resolved it with `get_mcp_server_by_name`, which walks the whole registry
and returns the first match. Server names are not unique, so a tool listed
from a reachable server could dispatch to a same-named server the caller
cannot reach, sending that server's upstream credential.

That is the bug this PR already fixed for MCP JSON-RPC, on the one surface
that had not been converted. The fix is the same: stop discarding identity.
`tool_server_map` now carries the server_id resolved within the caller's
reachable set through `resolve_tool_route`, the same scoped resolver the
JSON-RPC path uses, so a name two reachable servers share stays ambiguous
here too rather than silently picking one. Dispatch looks the server up by
id. A tool with no reachable owner fails closed and reports a result for its
tool call, matching how every other failure in that loop is surfaced, rather
than being dropped.

`resolved_server` is a parameter this PR introduced, and it let a caller
hand `call_tool` any server at all. That is the same class of defect one
layer down, so the check belongs at the chokepoint rather than at each
caller: `call_tool` now takes the reachable set the server was resolved
against and refuses to dispatch outside it, covering the caller's server and
one it resolves by name itself, which also walks the whole registry. Supplying
`resolved_server` without that set is rejected, so caller-supplied identity
always arrives with its provenance. Both callers already computed the set, so
nothing recomputes it.
This commit is contained in:
Tin Chi Lo 2026-07-25 17:17:59 -07:00
parent 0353b24509
commit e9fd8fccdf
7 changed files with 254 additions and 54 deletions

View file

@ -4686,6 +4686,7 @@ class MCPServerManager:
raw_headers: Optional[dict[str, str]] = None,
host_progress_callback: Optional[Callable] = None,
resolved_server: MCPServer | None = None,
allowed_server_ids: frozenset[str] | None = None,
) -> CallToolResult:
"""
Call a tool with the given name and arguments
@ -4702,17 +4703,38 @@ class MCPServerManager:
dispatched to verbatim, so identity never round-trips through the
non-unique ``server_name``. Callers without one fall back to
resolving by name.
allowed_server_ids: The reachable set ``resolved_server`` was resolved
against. Every dispatch is checked against it, both the caller's
server and one this resolves by name, because ``server_name`` is not
unique and resolving it walks the whole registry. Without it a caller
that resolved scope-blind, or a name collision between a reachable
server and an unreachable one, dispatches the wrong upstream
credential. Required whenever ``resolved_server`` is given, so
caller-supplied identity always arrives with its provenance.
Returns:
CallToolResult from the MCP server
"""
start_time = datetime.datetime.now()
if resolved_server is not None and allowed_server_ids is None:
raise ValueError("call_tool requires allowed_server_ids alongside resolved_server")
mcp_server = (
resolved_server
if resolved_server is not None
else self._resolve_mcp_server_for_tool_call(server_name, name)
)
if allowed_server_ids is not None and mcp_server.server_id not in allowed_server_ids:
raise HTTPException(
status_code=403,
detail={
"error": "server_out_of_scope",
"message": (
f"Tool '{name}' resolved to MCP server '{mcp_server.name}', which is not among "
f"the servers this caller may reach."
),
},
)
# Resolved before any hook runs so a missing BYOK credential (401) never
# leaves during-hook side effects (audit logging, rate-limit bookkeeping)

View file

@ -2841,6 +2841,7 @@ if MCP_AVAILABLE:
litellm_logging_obj=litellm_logging_obj,
host_progress_callback=host_progress_callback,
resolved_server=mcp_server,
allowed_server_ids=frozenset(server.server_id for server in allowed_mcp_servers),
)
# Fall back to local tool registry with original name (legacy support)
@ -3154,6 +3155,7 @@ if MCP_AVAILABLE:
litellm_logging_obj: Optional[Any] = None,
host_progress_callback: Optional[Callable] = None,
resolved_server: MCPServer | None = None,
allowed_server_ids: frozenset[str] | None = None,
) -> CallToolResult:
"""Handle tool execution for managed server tools"""
# Import here to avoid circular import
@ -3171,6 +3173,7 @@ if MCP_AVAILABLE:
proxy_logging_obj=proxy_logging_obj,
host_progress_callback=host_progress_callback,
resolved_server=resolved_server,
allowed_server_ids=allowed_server_ids,
)
verbose_logger.debug("CALL TOOL RESULT: %s", call_tool_result)
return call_tool_result

View file

@ -192,7 +192,7 @@ class LiteLLM_Proxy_MCP_Handler:
Returns:
List of MCP tools
List names of allowed MCP servers
List of server_ids for the MCP servers this caller may reach
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
@ -282,21 +282,18 @@ class LiteLLM_Proxy_MCP_Handler:
allowed_mcp_servers=allowed_mcp_servers,
)
server_names: List[str] = []
for server in allowed_mcp_servers:
if server is None:
continue
server_name = (
getattr(server, "server_name", None) or getattr(server, "alias", None) or getattr(server, "name", None)
)
if isinstance(server_name, str):
server_names.append(server_name)
# Carry server_id, not a display name. Names are not unique, so a name handed
# to the dispatch site gets re-resolved against the whole registry and can
# land on a server this caller cannot reach.
server_ids: List[str] = [
server.server_id for server in allowed_mcp_servers if server is not None and server.server_id
]
return tools, server_names
return tools, server_ids
@staticmethod
def _deduplicate_mcp_tools(
mcp_tools: List[MCPTool], allowed_mcp_servers: List[str]
mcp_tools: List[MCPTool], allowed_server_ids_list: List[str]
) -> tuple[List[MCPTool], dict[str, str]]:
"""
Deduplicate MCP tools by name, keeping the first occurrence of each tool.
@ -306,11 +303,16 @@ class LiteLLM_Proxy_MCP_Handler:
Returns:
List of deduplicated MCP tools
The returned dictionary maps each tool_name to the server_name
The returned dictionary maps each tool_name to the owning server_id
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
seen_names = set()
deduplicated_tools = []
tool_server_map: dict[str, str] = {}
allowed_server_ids = frozenset(allowed_server_ids_list)
for tool in mcp_tools:
if isinstance(tool, dict):
@ -321,10 +323,16 @@ class LiteLLM_Proxy_MCP_Handler:
if tool_name and tool_name not in seen_names:
seen_names.add(tool_name)
deduplicated_tools.append(tool)
if len(allowed_mcp_servers) == 1:
tool_server_map[tool_name] = allowed_mcp_servers[0]
else:
_, tool_server_map[tool_name] = split_server_prefix_from_name(tool_name)
if len(allowed_server_ids_list) == 1:
tool_server_map[tool_name] = allowed_server_ids_list[0]
continue
# Same scoped resolver the MCP JSON-RPC surface uses, so a name two
# reachable servers share stays ambiguous here too instead of
# silently picking one. Unresolvable names are left out of the map
# and fail closed at dispatch.
route = global_mcp_server_manager.resolve_tool_route(tool_name, allowed_server_ids=allowed_server_ids)
if route.kind == "resolved":
tool_server_map[tool_name] = route.server.server_id
return deduplicated_tools, tool_server_map
@ -658,11 +666,25 @@ class LiteLLM_Proxy_MCP_Handler:
# Import here to avoid circular import
from litellm.proxy.proxy_server import proxy_logging_obj
server_name = tool_server_map[tool_name]
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(
server_name
) or global_mcp_server_manager._get_mcp_server_from_tool_name(tool_name)
# tool_server_map carries the server_id resolved within this caller's
# reachable set at listing time, so look it up by identity. A name
# lookup here would walk the whole registry and could land on a
# same-named server the caller cannot reach.
server_id = tool_server_map.get(tool_name)
mcp_server = global_mcp_server_manager.get_mcp_server_by_id(server_id) if server_id else None
if mcp_server is None:
# Fail closed, and report it the way every other failure in this
# loop does so the model still sees a result for its tool call.
verbose_logger.warning(f"No reachable MCP server owns tool {tool_name}")
tool_results.append(
{
"tool_call_id": tool_call_id,
"result": f"Tool call failed: no MCP server reachable by this caller serves '{tool_name}'.",
"name": tool_name,
}
)
continue
server_name = mcp_server.name
resolved_tool_name = (
_resolve_display_name_to_original(tool_name, [mcp_server]) if mcp_server else tool_name
)
@ -781,6 +803,7 @@ class LiteLLM_Proxy_MCP_Handler:
raw_headers=raw_headers,
proxy_logging_obj=proxy_logging_obj,
resolved_server=mcp_server,
allowed_server_ids=frozenset(tool_server_map.values()),
)
if litellm_logging_obj:

View file

@ -5121,6 +5121,63 @@ class TestMCPServerManager:
assert "Tool deletepet is not allowed for server my_api_mcp" in exc_info.value.detail["error"]
assert "Contact proxy admin to allow this tool" in exc_info.value.detail["error"]
@pytest.mark.asyncio
async def test_call_tool_refuses_a_resolved_server_outside_the_callers_scope(self):
"""Dispatch is gated at the chokepoint, so a caller cannot hand over any server.
``resolved_server`` exists so identity does not round-trip through the
non-unique ``server_name``, but a caller that resolved it scope-blind would
otherwise send an unreachable server's upstream credential.
"""
manager = MCPServerManager()
zulu = MCPServer(server_id="id-zulu", name="echo_zulu", transport=MCPTransport.http)
manager.registry = {"id-zulu": zulu}
with pytest.raises(HTTPException) as excinfo:
await manager.call_tool(
server_name="echo_zulu",
name="echo",
arguments={},
resolved_server=zulu,
allowed_server_ids=frozenset({"id-alpha"}),
)
assert excinfo.value.status_code == 403
assert excinfo.value.detail["error"] == "server_out_of_scope"
@pytest.mark.asyncio
async def test_call_tool_refuses_a_name_resolved_server_outside_the_callers_scope(self):
"""The gate covers name resolution too, which walks the whole registry."""
manager = MCPServerManager()
zulu = MCPServer(server_id="id-zulu", name="echo_zulu", server_name="echo_zulu", transport=MCPTransport.http)
manager.registry = {"id-zulu": zulu}
manager._replace_server_tool_routes("id-zulu", {"echo", "echo_zulu-echo"})
with pytest.raises(HTTPException) as excinfo:
await manager.call_tool(
server_name="echo_zulu",
name="echo",
arguments={},
allowed_server_ids=frozenset({"id-alpha"}),
)
assert excinfo.value.status_code == 403
@pytest.mark.asyncio
async def test_call_tool_requires_the_scope_a_resolved_server_was_resolved_against(self):
"""Caller-supplied identity must arrive with its provenance, or it is unverifiable."""
manager = MCPServerManager()
alpha = MCPServer(server_id="id-alpha", name="echo_alpha", transport=MCPTransport.http)
manager.registry = {"id-alpha": alpha}
with pytest.raises(ValueError, match="allowed_server_ids"):
await manager.call_tool(
server_name="echo_alpha",
name="echo",
arguments={},
resolved_server=alpha,
)
@pytest.mark.asyncio
async def test_call_tool_without_broken_pipe_error(self):
"""

View file

@ -1069,6 +1069,14 @@ async def test_execute_tool_calls_sets_proxy_server_request_arguments(monkeypatc
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.call_tool",
mock_call_tool,
)
# tool_server_map carries server_ids, so dispatch resolves the server by identity.
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServer
from litellm.types.mcp import MCPTransport
monkeypatch.setattr(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.get_mcp_server_by_id",
lambda sid: MCPServer(server_id=sid, name=sid, server_name=sid, transport=MCPTransport.http),
)
# Create test data
tool_calls = [

View file

@ -27,11 +27,28 @@ def _setup_mcp_call_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock:
proxy_module = types.SimpleNamespace(proxy_logging_obj=object())
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module)
# tool_server_map carries server_ids, so dispatch resolves by identity. A real
# MCPServer is used rather than a stub so the prefix-stripping and display-name
# paths downstream see the attributes they actually read.
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServer
from litellm.types.mcp import MCPTransport
def _server_by_id(server_id):
if server_id is None:
return None
return MCPServer(
server_id=server_id,
name=server_id,
server_name=server_id,
transport=MCPTransport.http,
)
fake_manager = types.SimpleNamespace(
call_tool=AsyncMock(return_value=_DummyMCPResult()),
# Newer logging path calls this to enrich spend logs metadata
_get_mcp_server_from_tool_name=MagicMock(return_value=None),
get_mcp_server_by_name=MagicMock(return_value=None),
get_mcp_server_by_id=MagicMock(side_effect=_server_by_id),
)
monkeypatch.setattr(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
@ -40,6 +57,21 @@ def _setup_mcp_call_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock:
return fake_manager.call_tool
def _mcp_server(server_id: str, server_name: str, alias=None, tool_name_to_display_name=None):
"""A real MCPServer, so tests exercise the attributes production actually reads."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServer
from litellm.types.mcp import MCPTransport
return MCPServer(
server_id=server_id,
name=server_name,
server_name=server_name,
alias=alias,
transport=MCPTransport.http,
tool_name_to_display_name=tool_name_to_display_name,
)
def _setup_proxy_logging(monkeypatch: pytest.MonkeyPatch) -> AsyncMock:
"""Patch proxy_logging_obj so failure hook can be asserted."""
proxy_logging_obj = MagicMock()
@ -61,22 +93,75 @@ def test_deduplicate_mcp_tools_single_allowed_server():
assert server_map == {"search": "everything"}
@pytest.mark.parametrize(
"tool_name,expected_server",
[
("alpha-tool", "alpha"),
("beta-another_tool", "beta"),
],
)
def test_deduplicate_mcp_tools_prefixed_names(tool_name, expected_server):
tools = [{"name": tool_name}]
_, server_map = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools(
tools,
["alpha", "beta"],
def _manager_with(monkeypatch, servers, routes):
"""A real manager holding `servers`, with `routes` as {server_id: {tool names}}."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
assert server_map[tool_name] == expected_server
manager = MCPServerManager()
manager.registry = {server.server_id: server for server in servers}
for server_id, names in routes.items():
manager._replace_server_tool_routes(server_id, names)
monkeypatch.setattr(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
manager,
)
return manager
def test_deduplicate_mcp_tools_records_the_owning_server_id(monkeypatch):
"""A multi-server listing records server_ids, not names, since names are not unique."""
_manager_with(
monkeypatch,
[_mcp_server(server_id="id-alpha", server_name="alpha"), _mcp_server(server_id="id-beta", server_name="beta")],
{"id-alpha": {"alpha-tool"}, "id-beta": {"beta-another_tool"}},
)
_, server_map = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools(
[{"name": "alpha-tool"}, {"name": "beta-another_tool"}],
["id-alpha", "id-beta"],
)
assert server_map == {"alpha-tool": "id-alpha", "beta-another_tool": "id-beta"}
def test_deduplicate_mcp_tools_ignores_a_same_named_server_out_of_scope(monkeypatch):
"""The reviewed case: a reachable server sharing its name with an unreachable one.
Recording the name would let dispatch re-resolve it against the whole registry and
land on the unreachable server, sending that server's upstream credential.
"""
reachable = _mcp_server(server_id="id-reachable", server_name="shared")
unreachable = _mcp_server(server_id="id-unreachable", server_name="shared")
_manager_with(
monkeypatch,
[reachable, unreachable],
{"id-reachable": {"echo"}, "id-unreachable": {"echo"}},
)
_, server_map = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools(
[{"name": "echo"}],
["id-reachable"],
)
assert server_map == {"echo": "id-reachable"}
def test_deduplicate_mcp_tools_omits_a_name_two_reachable_servers_share(monkeypatch):
"""Ambiguity within the caller's own scope must not silently pick one server."""
_manager_with(
monkeypatch,
[_mcp_server(server_id="id-alpha", server_name="alpha"), _mcp_server(server_id="id-beta", server_name="beta")],
{"id-alpha": {"echo"}, "id-beta": {"echo"}},
)
_, server_map = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools(
[{"name": "echo"}],
["id-alpha", "id-beta"],
)
assert "echo" not in server_map
def test_extract_tool_calls_from_chat_response_handles_tool_calls():
@ -288,19 +373,14 @@ async def test_execute_tool_calls_strips_prefix_when_alias_differs_from_server_n
monkeypatch,
):
call_tool_mock = _setup_mcp_call_environment(monkeypatch)
fake_server = types.SimpleNamespace(
fake_server = _mcp_server(
alias="my_deepwiki",
server_name="deepwiki_test",
server_id="test-server-id",
short_prefix=None,
mcp_info=None,
tool_name_to_display_name=None,
)
from litellm.proxy._experimental.mcp_server import mcp_server_manager as _msm
_msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock(
return_value=fake_server
)
_msm.global_mcp_server_manager.get_mcp_server_by_id = MagicMock(return_value=fake_server)
tool_name = "my_deepwiki-read_wiki_structure"
tool_calls = [
@ -311,7 +391,7 @@ async def test_execute_tool_calls_strips_prefix_when_alias_differs_from_server_n
]
await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
tool_server_map={tool_name: "deepwiki_test"},
tool_server_map={tool_name: "test-server-id"},
tool_calls=tool_calls,
user_api_key_auth=None,
)
@ -324,26 +404,24 @@ async def test_execute_tool_calls_strips_prefix_when_alias_differs_from_server_n
@pytest.mark.asyncio
async def test_execute_tool_calls_reverse_maps_display_name(monkeypatch):
call_tool_mock = _setup_mcp_call_environment(monkeypatch)
colliding_server = types.SimpleNamespace(
alias=None,
colliding_server = _mcp_server(
server_name="other_mcp",
server_id="other-server-id",
short_prefix=None,
mcp_info=None,
tool_name_to_display_name={"search": "search_docs"},
)
fake_server = types.SimpleNamespace(
alias=None,
fake_server = _mcp_server(
server_name="deepwiki_mcp",
server_id="test-server-id",
short_prefix=None,
mcp_info=None,
tool_name_to_display_name={"read_wiki_structure": "browse_repo_docs"},
)
from litellm.proxy._experimental.mcp_server import mcp_server_manager as _msm
# Dispatch resolves the server_id recorded at listing time. A name lookup would
# be free to return the collider instead, which is the bug this guards.
by_id = {"test-server-id": fake_server, "other-server-id": colliding_server}
_msm.global_mcp_server_manager.get_mcp_server_by_id = MagicMock(side_effect=by_id.get)
_msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock(return_value=colliding_server)
_msm.global_mcp_server_manager.get_mcp_server_by_name = MagicMock(return_value=fake_server)
_msm.global_mcp_server_manager.get_mcp_server_by_name = MagicMock(return_value=colliding_server)
tool_name = "browse_repo_docs"
tool_calls = [
@ -354,7 +432,7 @@ async def test_execute_tool_calls_reverse_maps_display_name(monkeypatch):
]
await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
tool_server_map={tool_name: "deepwiki_mcp"},
tool_server_map={tool_name: "test-server-id"},
tool_calls=tool_calls,
user_api_key_auth=None,
)

View file

@ -68,10 +68,19 @@ def _text_only_stream(text: str, response_id: str = "resp-1") -> _FakeAsyncStrea
def _mock_mcp_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock:
"""Patch the MCP tool-call plumbing so _execute_tool_calls can run in tests."""
call_tool = AsyncMock(return_value=CallToolResult(content=[TextContent(type="text", text="ok")], isError=False))
# tool_server_map carries server_ids now, so dispatch resolves by identity.
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServer
from litellm.types.mcp import MCPTransport
fake_manager = types.SimpleNamespace(
call_tool=call_tool,
_get_mcp_server_from_tool_name=MagicMock(return_value=None),
get_mcp_server_by_name=MagicMock(return_value=None),
get_mcp_server_by_id=MagicMock(
side_effect=lambda sid: (
MCPServer(server_id=sid, name=sid, server_name=sid, transport=MCPTransport.http) if sid else None
)
),
)
monkeypatch.setattr(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",