diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 412a6de0059..a3b5e683e1d 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1209,6 +1209,19 @@ if MCP_AVAILABLE: # Remove prefix from tool name for logging and processing original_tool_name, server_name = split_server_prefix_from_name(name) + # Validate that extracted server_name exists in allowed servers + # If not, the tool name likely contains the separator itself + if server_name and not any( + server_name in [s for s in [server.name, server.alias] if s] + for server in allowed_mcp_servers + ): + verbose_logger.debug( + f"Server '{server_name}' from tool '{name}' not found in allowed servers. " + f"Treating as unprefixed tool." + ) + original_tool_name = name + server_name = "" + # If tool name is unprefixed, resolve its server so we can enforce permissions if not server_name: mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 4fc94000d61..6c8072c4c66 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -1544,3 +1544,72 @@ def test_filter_tools_by_allowed_tools(): assert len(filtered_tools) == 2 assert filtered_tools[0].name == "my_api_mcp-getpetbyid" assert filtered_tools[1].name == "my_api_mcp-findpetsbystatus" + +@pytest.mark.asyncio +async def test_call_mcp_tool_validates_extracted_server_name(): + """ + Test that call_mcp_tool validates extracted server names against allowed servers. + If a tool name contains the separator but the extracted server doesn't exist, + it should be treated as an unprefixed tool name. + """ + + try: + from litellm.proxy._experimental.mcp_server.server import ( + call_mcp_tool, + global_mcp_server_manager, + ) + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + except ImportError: + pytest.skip("MCP server not available") + + + # Create a real allowed server + allowed_server = MCPServer( + server_id="server-123", + name="backstage", + alias="backstage", + server_name="backstage", + url="https://backstage.com/mcp", + transport=MCPTransport.http, + mcp_info={"server_name": "backstage"}, + ) + + with patch.object( + global_mcp_server_manager, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + ) as mock_get_allowed, patch.object( + global_mcp_server_manager, + "get_mcp_server_by_id", + return_value=allowed_server, + ), patch.object( + global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + return_value=allowed_server, + ) as mock_get_server, patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_tool_registry" + ) as mock_tool_registry, patch( + "litellm.proxy._experimental.mcp_server.server._handle_managed_mcp_tool", + new_callable=AsyncMock, + ) as mock_handle_managed, patch( + "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed", + return_value=True, + ): + mock_get_allowed.return_value = [allowed_server.server_id] + mock_tool_registry.get_tool.return_value = None + + # Call with a tool name that contains separator but "get" is not a valid server + # This should be treated as unprefixed tool "get--catalog--entity" + result = await call_mcp_tool( + name="get-catalog-entity", + arguments={"id": "123"}, + mcp_servers=["backstage"], + ) + + mock_get_allowed.assert_awaited_once() + # Should resolve server from tool name since extracted "get" is not valid + assert mock_get_server.call_count >= 1 + # First call should use the full tool name (not split) since "get" is not a valid server + assert mock_get_server.call_args_list[0][0][0] == "get-catalog-entity" \ No newline at end of file