From 6a0321f852bf2f0e89189082c15010e6daebc973 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Fri, 6 Feb 2026 17:44:56 -0800 Subject: [PATCH] test mcp --- tests/mcp_tests/test_mcp_server.py | 191 +++++++++++++++++------------ 1 file changed, 114 insertions(+), 77 deletions(-) diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index 9e6c3f19e56..ea823df1fb2 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -627,7 +627,10 @@ def test_generate_stable_server_id(): @pytest.mark.asyncio async def test_list_tools_rest_api_server_not_found(): """Test the list_tools REST API when server is not found""" - from litellm.proxy._experimental.mcp_server.rest_endpoints import list_tool_rest_api + from litellm.proxy._experimental.mcp_server.rest_endpoints import ( + list_tool_rest_api, + global_mcp_server_manager, + ) from fastapi import Query from litellm.proxy._types import UserAPIKeyAuth @@ -641,21 +644,38 @@ async def test_list_tools_rest_api_server_not_found(): ), ) - # Mock request + # Mock request with proper client attribute (internal IP for no filtering) mock_request = MagicMock() mock_request.headers = {} + mock_request.client = MagicMock() + mock_request.client.host = "127.0.0.1" # Internal IP to bypass IP filtering - # Test with non-existent server ID - response = await list_tool_rest_api( - request=mock_request, - server_id="non_existent_server_id", - user_api_key_dict=mock_user_auth, - ) + # Mock the global_mcp_server_manager to allow the server ID but return None for the server + with patch( + "litellm.proxy._experimental.mcp_server.rest_endpoints.global_mcp_server_manager" + ) as mock_manager: + # Allow the server ID in permissions + mock_manager.get_allowed_mcp_servers = AsyncMock( + return_value=["non_existent_server_id"] + ) + # Mock filter_server_ids_by_ip to return input unchanged (no IP filtering in test) + mock_manager.filter_server_ids_by_ip = MagicMock( + side_effect=lambda server_ids, client_ip: server_ids + ) + # Return None when trying to get the server (server doesn't exist) + mock_manager.get_mcp_server_by_id = MagicMock(return_value=None) - assert isinstance(response, dict) - assert response["tools"] == [] - assert response["error"] == "server_not_found" - assert "Server with id non_existent_server_id not found" in response["message"] + # Test with non-existent server ID + response = await list_tool_rest_api( + request=mock_request, + server_id="non_existent_server_id", + user_api_key_dict=mock_user_auth, + ) + + assert isinstance(response, dict) + assert response["tools"] == [] + assert response["error"] == "server_not_found" + assert "Server with id non_existent_server_id not found" in response["message"] @pytest.mark.asyncio @@ -663,90 +683,75 @@ async def test_list_tools_rest_api_success(): """Test the list_tools REST API successful case""" from litellm.proxy._experimental.mcp_server.rest_endpoints import ( list_tool_rest_api, - global_mcp_server_manager, + ) + from litellm.proxy._experimental.mcp_server.server import ( + ListMCPToolsRestAPIResponseObject, ) from fastapi import Query from litellm.proxy._types import UserAPIKeyAuth - # Store original registry to restore after test - original_registry = global_mcp_server_manager.get_registry().copy() - original_tool_mapping = ( - global_mcp_server_manager.tool_name_to_mcp_server_name_mapping.copy() + # Mock successful tools + mock_tools = [ + ListMCPToolsRestAPIResponseObject( + name="test_tool", + description="A test tool", + inputSchema={"type": "object"}, + mcp_info={"server_name": "test_server"}, + ) + ] + + # Create a mock server + mock_server = MagicMock() + mock_server.server_id = "test-server-123" + mock_server.alias = "test_server" + mock_server.name = "test_server" + mock_server.mcp_info = {"server_name": "test_server"} + + # Mock UserAPIKeyAuth + mock_user_auth = UserAPIKeyAuth( + api_key="test", + user_id="test", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="dummy", + mcp_servers=["test-server-123"], + ), ) - try: - # Clear existing registry - global_mcp_server_manager.tool_name_to_mcp_server_name_mapping.clear() - global_mcp_server_manager.registry.clear() - global_mcp_server_manager.config_mcp_servers.clear() - # Mock successful tools - mock_tools = [ - MCPTool( - name="test_tool", - description="A test tool", - inputSchema={"type": "object"}, - ) - ] + # Mock request with proper client attribute (internal IP for no filtering) + mock_request = MagicMock() + mock_request.headers = {} + mock_request.client = MagicMock() + mock_request.client.host = "127.0.0.1" # Internal IP to bypass IP filtering - # Create mock client - mock_client = AsyncMock() - mock_client.list_tools = AsyncMock(return_value=mock_tools) - mock_client.__aenter__ = AsyncMock(return_value=mock_client) - mock_client.__aexit__ = AsyncMock(return_value=None) - - def mock_client_constructor(*args, **kwargs): - return mock_client + # Mock the global_mcp_server_manager + with patch( + "litellm.proxy._experimental.mcp_server.rest_endpoints.global_mcp_server_manager" + ) as mock_manager: + mock_manager.get_allowed_mcp_servers = AsyncMock( + return_value=["test-server-123"] + ) + mock_manager.get_mcp_server_by_id = MagicMock(return_value=mock_server) + # Mock filter_server_ids_by_ip to return input unchanged (no IP filtering in test) + mock_manager.filter_server_ids_by_ip = MagicMock( + side_effect=lambda server_ids, client_ip: server_ids + ) + # Mock the _get_tools_for_single_server function with patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient", - mock_client_constructor, - ): - # Load server config into global manager - await global_mcp_server_manager.load_servers_from_config( - { - "test_server": { - "url": "https://test-server.com/mcp", - "transport": MCPTransport.http, - } - } - ) - - # Mock UserAPIKeyAuth - mock_user_auth = UserAPIKeyAuth( - api_key="test", - user_id="test", - object_permission=LiteLLM_ObjectPermissionTable( - object_permission_id="dummy", - mcp_servers=list( - global_mcp_server_manager.get_all_mcp_server_ids() - ), - ), - ) - - # Get the server ID - server_id = list(global_mcp_server_manager.get_registry().keys())[0] - - # Mock request - mock_request = MagicMock() - mock_request.headers = {} + "litellm.proxy._experimental.mcp_server.rest_endpoints._get_tools_for_single_server" + ) as mock_get_tools: + mock_get_tools.return_value = mock_tools # Test successful case response = await list_tool_rest_api( request=mock_request, - server_id=server_id, + server_id="test-server-123", user_api_key_dict=mock_user_auth, ) assert isinstance(response, dict) assert len(response["tools"]) == 1 assert response["tools"][0].name == "test_tool" - finally: - # Restore original state - global_mcp_server_manager.registry = {} - global_mcp_server_manager.config_mcp_servers = original_registry - global_mcp_server_manager.tool_name_to_mcp_server_name_mapping = ( - original_tool_mapping - ) @pytest.mark.asyncio @@ -808,6 +813,10 @@ async def test_get_tools_from_mcp_servers(): ) mock_manager.get_mcp_server_by_id = lambda server_id: mock_server_1 if server_id == "server1_id" else mock_server_2 mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1]) + # Mock filter_server_ids_by_ip to return input unchanged (no IP filtering in test) + mock_manager.filter_server_ids_by_ip = MagicMock( + side_effect=lambda server_ids, client_ip: server_ids + ) with patch( "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", @@ -843,6 +852,10 @@ async def test_get_tools_from_mcp_servers(): mock_manager_2._get_tools_from_server = AsyncMock( side_effect=mock_get_tools_side_effect ) + # Mock filter_server_ids_by_ip to return input unchanged (no IP filtering in test) + mock_manager_2.filter_server_ids_by_ip = MagicMock( + side_effect=lambda server_ids, client_ip: server_ids + ) with patch( "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", @@ -867,6 +880,10 @@ async def test_get_tools_from_mcp_servers(): ) mock_manager.get_mcp_server_by_id = lambda server_id: mock_server_1 if server_id == "server1_id" else (mock_server_2 if server_id == "server2_id" else mock_server_3) mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1]) + # Mock filter_server_ids_by_ip to return input unchanged (no IP filtering in test) + mock_manager.filter_server_ids_by_ip = MagicMock( + side_effect=lambda server_ids, client_ip: server_ids + ) with patch( "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", @@ -1795,6 +1812,10 @@ async def test_list_tool_rest_api_with_server_specific_auth(): mock_server.mcp_info = {"server_name": "zapier"} mock_manager.get_mcp_server_by_id.return_value = mock_server + # Mock filter_server_ids_by_ip to return input unchanged (no IP filtering in test) + mock_manager.filter_server_ids_by_ip = MagicMock( + side_effect=lambda server_ids, client_ip: server_ids + ) mock_user_api_key_dict = UserAPIKeyAuth( api_key="test", @@ -1885,6 +1906,10 @@ async def test_list_tool_rest_api_with_default_auth(): mock_server.mcp_info = {"server_name": "unknown_server"} mock_manager.get_mcp_server_by_id.return_value = mock_server + # Mock filter_server_ids_by_ip to return input unchanged (no IP filtering in test) + mock_manager.filter_server_ids_by_ip = MagicMock( + side_effect=lambda server_ids, client_ip: server_ids + ) mock_user_api_key_dict = UserAPIKeyAuth( api_key="test", @@ -1991,6 +2016,10 @@ async def test_list_tool_rest_api_all_servers_with_auth(): server_id ) ) + # Mock filter_server_ids_by_ip to return input unchanged (no IP filtering in test) + mock_manager.filter_server_ids_by_ip = MagicMock( + side_effect=lambda server_ids, client_ip: server_ids + ) mock_user_api_key_dict = UserAPIKeyAuth( api_key="test", @@ -2234,6 +2263,10 @@ async def test_filter_tools_by_disallowed_tools_integration(): return_value=["test-server-456"] ) mock_manager.get_mcp_server_by_id = MagicMock(return_value=mock_server) + # Mock filter_server_ids_by_ip to return input unchanged (no IP filtering in test) + mock_manager.filter_server_ids_by_ip = MagicMock( + side_effect=lambda server_ids, client_ip: server_ids + ) # Mock the _get_tools_from_server method to return all tools mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools) @@ -2330,6 +2363,10 @@ async def test_filter_tools_no_restrictions_integration(): return_value=["test-server-000"] ) mock_manager.get_mcp_server_by_id = MagicMock(return_value=mock_server) + # Mock filter_server_ids_by_ip to return input unchanged (no IP filtering in test) + mock_manager.filter_server_ids_by_ip = MagicMock( + side_effect=lambda server_ids, client_ip: server_ids + ) # Mock the _get_tools_from_server method to return all tools mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools)