mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
test mcp
This commit is contained in:
parent
04eef6869c
commit
6a0321f852
1 changed files with 114 additions and 77 deletions
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue