mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
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:
parent
0353b24509
commit
e9fd8fccdf
7 changed files with 254 additions and 54 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue