This commit is contained in:
Ishaan Jaffer 2026-02-06 17:44:56 -08:00
parent 04eef6869c
commit 6a0321f852

View file

@ -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)