diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 7bcccd1841d..91a8c0d473b 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 23e853a32d3..22f8369895a 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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 diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index cae0a38c6e8..2e3e63bc451 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -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: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index f9de88610d7..ba16b70fff0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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): """ diff --git a/tests/test_litellm/responses/mcp/test_chat_completions_handler.py b/tests/test_litellm/responses/mcp/test_chat_completions_handler.py index ab4c5185057..3cef036214f 100644 --- a/tests/test_litellm/responses/mcp/test_chat_completions_handler.py +++ b/tests/test_litellm/responses/mcp/test_chat_completions_handler.py @@ -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 = [ diff --git a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py index a1347aa111c..44c1883343d 100644 --- a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -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, ) diff --git a/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py b/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py index 24edf12fffe..38e58dc1440 100644 --- a/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py +++ b/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py @@ -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",