mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
style: black format test_mcp_server.py
This commit is contained in:
parent
fcd71e0026
commit
537e72c742
1 changed files with 305 additions and 182 deletions
|
|
@ -5,7 +5,12 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from mcp import ReadResourceResult, Resource
|
||||
from mcp.types import BlobResourceContents, Prompt, ResourceTemplate, TextResourceContents
|
||||
from mcp.types import (
|
||||
BlobResourceContents,
|
||||
Prompt,
|
||||
ResourceTemplate,
|
||||
TextResourceContents,
|
||||
)
|
||||
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_MCPServerTable,
|
||||
|
|
@ -157,15 +162,19 @@ async def test_get_prompts_from_mcp_servers_success():
|
|||
server_b.auth_type = None
|
||||
server_b.extra_headers = None
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[server_a, server_b]),
|
||||
) as mock_allowed, patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
|
||||
return_value=(None, None),
|
||||
) as mock_headers, patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
) as mock_manager:
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[server_a, server_b]),
|
||||
) as mock_allowed,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
|
||||
return_value=(None, None),
|
||||
) as mock_headers,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
) as mock_manager,
|
||||
):
|
||||
mock_manager.get_prompts_from_server = AsyncMock(
|
||||
side_effect=[
|
||||
[Prompt(name="hello", description="hi")],
|
||||
|
|
@ -213,15 +222,19 @@ async def test_get_resources_from_mcp_servers_success():
|
|||
server_b.auth_type = None
|
||||
server_b.extra_headers = None
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[server_a, server_b]),
|
||||
) as mock_allowed, patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
|
||||
return_value=(None, None),
|
||||
) as mock_headers, patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
) as mock_manager:
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[server_a, server_b]),
|
||||
) as mock_allowed,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
|
||||
return_value=(None, None),
|
||||
) as mock_headers,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
) as mock_manager,
|
||||
):
|
||||
mock_manager.get_resources_from_server = AsyncMock(
|
||||
side_effect=[
|
||||
[
|
||||
|
|
@ -274,15 +287,19 @@ async def test_get_resource_templates_from_mcp_servers_success():
|
|||
server.auth_type = None
|
||||
server.extra_headers = None
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[server]),
|
||||
) as mock_allowed, patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
|
||||
return_value=(None, None),
|
||||
) as mock_headers, patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
) as mock_manager:
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[server]),
|
||||
) as mock_allowed,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
|
||||
return_value=(None, None),
|
||||
) as mock_headers,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
) as mock_manager,
|
||||
):
|
||||
mock_manager.get_resource_templates_from_server = AsyncMock(
|
||||
return_value=[
|
||||
ResourceTemplate(
|
||||
|
|
@ -320,15 +337,19 @@ async def test_mcp_get_prompt_success():
|
|||
|
||||
prompt_result = MagicMock(name="prompt_result")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[server]),
|
||||
) as mock_allowed, patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
|
||||
return_value=({"Authorization": "token"}, {"X-Test": "1"}),
|
||||
) as mock_headers, patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
) as mock_manager:
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[server]),
|
||||
) as mock_allowed,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
|
||||
return_value=({"Authorization": "token"}, {"X-Test": "1"}),
|
||||
) as mock_headers,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
) as mock_manager,
|
||||
):
|
||||
mock_manager.get_prompt_from_server = AsyncMock(return_value=prompt_result)
|
||||
|
||||
result = await mcp_get_prompt(
|
||||
|
|
@ -378,15 +399,19 @@ async def test_mcp_read_resource_success():
|
|||
]
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[server]),
|
||||
) as mock_allowed, patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
|
||||
return_value=({"Authorization": "token"}, {"X-Test": "1"}),
|
||||
) as mock_headers, patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
) as mock_manager:
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[server]),
|
||||
) as mock_allowed,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
|
||||
return_value=({"Authorization": "token"}, {"X-Test": "1"}),
|
||||
) as mock_headers,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
) as mock_manager,
|
||||
):
|
||||
mock_manager.read_resource_from_server = AsyncMock(return_value=read_result)
|
||||
|
||||
result = await mcp_read_resource(
|
||||
|
|
@ -591,7 +616,10 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails():
|
|||
working_server if server_id == "working_server" else failing_server
|
||||
)
|
||||
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
|
||||
server_ids,
|
||||
0,
|
||||
)
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server,
|
||||
|
|
@ -693,7 +721,10 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing():
|
|||
failing_server1 if server_id == "failing_server1" else failing_server2
|
||||
)
|
||||
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
|
||||
server_ids,
|
||||
0,
|
||||
)
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server,
|
||||
|
|
@ -830,12 +861,14 @@ async def test_concurrent_initialize_session_managers():
|
|||
mcp_server._sse_session_manager_cm = None
|
||||
|
||||
# Mock the session managers to avoid actual MCP initialization
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.session_manager"
|
||||
) as mock_session_manager, patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.sse_session_manager"
|
||||
) as mock_sse_session_manager, patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.verbose_logger"
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.session_manager"
|
||||
) as mock_session_manager,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.sse_session_manager"
|
||||
) as mock_sse_session_manager,
|
||||
patch("litellm.proxy._experimental.mcp_server.server.verbose_logger"),
|
||||
):
|
||||
# Mock the run() method to return a mock context manager
|
||||
mock_cm = AsyncMock()
|
||||
|
|
@ -961,15 +994,19 @@ async def test_mcp_routing_with_conflicting_alias_and_group_name():
|
|||
return_value=[specific_server.server_id, other_server.server_id]
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_allowed_mcp_servers",
|
||||
mock_get_allowed,
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler._get_mcp_servers_from_access_groups",
|
||||
mock_db_lookup,
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server",
|
||||
mock_get_tools_spy,
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_allowed_mcp_servers",
|
||||
mock_get_allowed,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler._get_mcp_servers_from_access_groups",
|
||||
mock_db_lookup,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server",
|
||||
mock_get_tools_spy,
|
||||
),
|
||||
):
|
||||
mcp_servers_from_path = _get_mcp_servers_in_path(test_path)
|
||||
|
||||
|
|
@ -1062,17 +1099,21 @@ async def test_oauth2_headers_passed_to_mcp_client():
|
|||
async def mock_fetch_tools_with_timeout(client, server_name):
|
||||
return [] # Return empty list of tools
|
||||
|
||||
with patch.object(
|
||||
global_mcp_server_manager,
|
||||
"_create_mcp_client",
|
||||
side_effect=mock_create_mcp_client,
|
||||
) as mock_create_client, patch.object(
|
||||
global_mcp_server_manager,
|
||||
"_fetch_tools_with_timeout",
|
||||
side_effect=mock_fetch_tools_with_timeout,
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[oauth2_server]),
|
||||
with (
|
||||
patch.object(
|
||||
global_mcp_server_manager,
|
||||
"_create_mcp_client",
|
||||
side_effect=mock_create_mcp_client,
|
||||
) as mock_create_client,
|
||||
patch.object(
|
||||
global_mcp_server_manager,
|
||||
"_fetch_tools_with_timeout",
|
||||
side_effect=mock_fetch_tools_with_timeout,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[oauth2_server]),
|
||||
),
|
||||
):
|
||||
# Call _get_tools_from_mcp_servers which should eventually call _create_mcp_client
|
||||
await _get_tools_from_mcp_servers(
|
||||
|
|
@ -1138,7 +1179,10 @@ async def test_list_tools_single_server_unprefixed_names():
|
|||
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
|
||||
mock_manager.get_mcp_server_by_id = MagicMock(return_value=server)
|
||||
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
|
||||
server_ids,
|
||||
0,
|
||||
)
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server,
|
||||
|
|
@ -1216,7 +1260,10 @@ async def test_list_tools_multiple_servers_prefixed_names():
|
|||
server1 if server_id == "server1" else server2
|
||||
)
|
||||
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
|
||||
server_ids,
|
||||
0,
|
||||
)
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server,
|
||||
|
|
@ -1270,12 +1317,15 @@ async def test_mcp_manager_allows_public_servers_without_permissions():
|
|||
)
|
||||
manager.registry = {public_server.server_id: public_server}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.common_utils._user_has_admin_view",
|
||||
return_value=False,
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[]),
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.common_utils._user_has_admin_view",
|
||||
return_value=False,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[]),
|
||||
),
|
||||
):
|
||||
allowed = await manager.get_allowed_mcp_servers(UserAPIKeyAuth())
|
||||
|
||||
|
|
@ -1302,12 +1352,15 @@ async def test_mcp_manager_returns_public_when_permission_lookup_fails():
|
|||
)
|
||||
manager.registry = {public_server.server_id: public_server}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.common_utils._user_has_admin_view",
|
||||
return_value=False,
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers",
|
||||
AsyncMock(side_effect=Exception("boom")),
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.common_utils._user_has_admin_view",
|
||||
return_value=False,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers",
|
||||
AsyncMock(side_effect=Exception("boom")),
|
||||
),
|
||||
):
|
||||
allowed = await manager.get_allowed_mcp_servers(UserAPIKeyAuth())
|
||||
|
||||
|
|
@ -1342,12 +1395,15 @@ async def test_mcp_manager_merges_public_and_restricted_servers():
|
|||
scoped_server.server_id: scoped_server,
|
||||
}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.common_utils._user_has_admin_view",
|
||||
return_value=False,
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=["restricted"]),
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.common_utils._user_has_admin_view",
|
||||
return_value=False,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPRequestHandler.get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=["restricted"]),
|
||||
),
|
||||
):
|
||||
allowed = await manager.get_allowed_mcp_servers(UserAPIKeyAuth())
|
||||
|
||||
|
|
@ -1399,12 +1455,15 @@ async def test_call_mcp_tool_user_unauthorized_access():
|
|||
return another_server_obj
|
||||
return None
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=["allowed_server", "another_server"]),
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_id",
|
||||
side_effect=mock_get_server_by_id,
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=["allowed_server", "another_server"]),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_id",
|
||||
side_effect=mock_get_server_by_id,
|
||||
),
|
||||
):
|
||||
# Try to call a tool from "restricted_server" - should raise HTTPException with 403 status
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
|
|
@ -1467,7 +1526,10 @@ async def test_list_tools_filters_by_key_team_permissions():
|
|||
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
|
||||
mock_manager.get_mcp_server_by_id = lambda server_id: server
|
||||
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
|
||||
server_ids,
|
||||
0,
|
||||
)
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server,
|
||||
|
|
@ -1573,7 +1635,10 @@ async def test_list_tools_with_team_tool_permissions_inheritance():
|
|||
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
|
||||
mock_manager.get_mcp_server_by_id = lambda server_id: server
|
||||
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
|
||||
server_ids,
|
||||
0,
|
||||
)
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server,
|
||||
|
|
@ -1665,7 +1730,10 @@ async def test_list_tools_with_no_tool_permissions_shows_all():
|
|||
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"])
|
||||
mock_manager.get_mcp_server_by_id = lambda server_id: server
|
||||
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
|
||||
server_ids,
|
||||
0,
|
||||
)
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server,
|
||||
|
|
@ -1760,7 +1828,10 @@ async def test_list_tools_strips_prefix_when_matching_permissions():
|
|||
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["gitmcp_server"])
|
||||
mock_manager.get_mcp_server_by_id = MagicMock(return_value=server)
|
||||
# Mock filter_server_ids_by_ip to return server_ids unchanged (no IP filtering)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (server_ids, 0)
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
|
||||
server_ids,
|
||||
0,
|
||||
)
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server,
|
||||
|
|
@ -2002,12 +2073,15 @@ class TestMCPServerManagerReload:
|
|||
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(
|
||||
return_value=[db_row]
|
||||
)
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=mock_prisma,
|
||||
), patch.object(
|
||||
manager, "build_mcp_server_from_table", AsyncMock()
|
||||
) as mock_build:
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=mock_prisma,
|
||||
),
|
||||
patch.object(
|
||||
manager, "build_mcp_server_from_table", AsyncMock()
|
||||
) as mock_build,
|
||||
):
|
||||
await manager.reload_servers_from_database()
|
||||
|
||||
mock_build.assert_not_awaited()
|
||||
|
|
@ -2045,14 +2119,17 @@ class TestMCPServerManagerReload:
|
|||
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(
|
||||
return_value=[db_row]
|
||||
)
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=mock_prisma,
|
||||
), patch.object(
|
||||
manager,
|
||||
"build_mcp_server_from_table",
|
||||
AsyncMock(return_value=rebuilt_server),
|
||||
) as mock_build:
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
|
||||
return_value=mock_prisma,
|
||||
),
|
||||
patch.object(
|
||||
manager,
|
||||
"build_mcp_server_from_table",
|
||||
AsyncMock(return_value=rebuilt_server),
|
||||
) as mock_build,
|
||||
):
|
||||
await manager.reload_servers_from_database()
|
||||
|
||||
mock_build.assert_awaited_once_with(db_row)
|
||||
|
|
@ -2090,26 +2167,32 @@ async def test_call_mcp_tool_logs_failure_via_post_call_failure_hook():
|
|||
|
||||
user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
|
||||
|
||||
with patch.object(
|
||||
global_mcp_server_manager,
|
||||
"get_allowed_mcp_servers",
|
||||
new_callable=AsyncMock,
|
||||
return_value=[mock_server.server_id],
|
||||
), patch.object(
|
||||
global_mcp_server_manager,
|
||||
"get_mcp_server_by_id",
|
||||
return_value=mock_server,
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers_from_mcp_server_names",
|
||||
new_callable=AsyncMock,
|
||||
return_value=[mock_server],
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.execute_mcp_tool",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=Exception("boom"),
|
||||
), patch(
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj",
|
||||
proxy_logging_mock,
|
||||
with (
|
||||
patch.object(
|
||||
global_mcp_server_manager,
|
||||
"get_allowed_mcp_servers",
|
||||
new_callable=AsyncMock,
|
||||
return_value=[mock_server.server_id],
|
||||
),
|
||||
patch.object(
|
||||
global_mcp_server_manager,
|
||||
"get_mcp_server_by_id",
|
||||
return_value=mock_server,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers_from_mcp_server_names",
|
||||
new_callable=AsyncMock,
|
||||
return_value=[mock_server],
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.execute_mcp_tool",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=Exception("boom"),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj",
|
||||
proxy_logging_mock,
|
||||
),
|
||||
):
|
||||
with pytest.raises(Exception):
|
||||
await call_mcp_tool(
|
||||
|
|
@ -2157,23 +2240,30 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab
|
|||
dummy_logging_obj.model_call_details = {"metadata": {"spend_logs_metadata": {}}}
|
||||
dummy_logging_obj.async_success_handler = AsyncMock()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
new=AsyncMock(return_value=[server_a]),
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
|
||||
return_value=(None, None),
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
) as mock_manager, patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools",
|
||||
side_effect=lambda tools, _server: tools,
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions",
|
||||
new=AsyncMock(side_effect=lambda tools, **_: tools),
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.function_setup",
|
||||
return_value=(dummy_logging_obj, None),
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
new=AsyncMock(return_value=[server_a]),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
|
||||
return_value=(None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
) as mock_manager,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools",
|
||||
side_effect=lambda tools, _server: tools,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions",
|
||||
new=AsyncMock(side_effect=lambda tools, **_: tools),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.function_setup",
|
||||
return_value=(dummy_logging_obj, None),
|
||||
),
|
||||
):
|
||||
mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1])
|
||||
|
||||
|
|
@ -2188,7 +2278,9 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab
|
|||
|
||||
assert tools == [tool_1]
|
||||
dummy_logging_obj.async_success_handler.assert_awaited_once()
|
||||
assert dummy_logging_obj.async_success_handler.await_args.kwargs["result"] == [tool_1]
|
||||
assert dummy_logging_obj.async_success_handler.await_args.kwargs["result"] == [
|
||||
tool_1
|
||||
]
|
||||
|
||||
spend_meta = dummy_logging_obj.model_call_details["metadata"]["spend_logs_metadata"]
|
||||
assert spend_meta["tool_count_total"] == 1
|
||||
|
|
@ -2381,26 +2473,34 @@ async def test_get_tools_from_mcp_servers_injects_stored_oauth2_token():
|
|||
oauth2_server.extra_headers = None
|
||||
|
||||
# Simulate the DB returning a valid credential for this user+server
|
||||
prefetched_creds = {SERVER_ID: {"access_token": STORED_TOKEN, "server_id": SERVER_ID}}
|
||||
prefetched_creds = {
|
||||
SERVER_ID: {"access_token": STORED_TOKEN, "server_id": SERVER_ID}
|
||||
}
|
||||
|
||||
tool_1 = MagicMock()
|
||||
tool_1.name = "atlassian_test-search"
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
new=AsyncMock(return_value=[oauth2_server]),
|
||||
), patch(
|
||||
# Patch the bulk prefetch so no real DB connection is needed
|
||||
"litellm.proxy._experimental.mcp_server.server._prefetch_oauth_creds_for_user",
|
||||
new=AsyncMock(return_value=prefetched_creds),
|
||||
) as mock_prefetch, patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
) as mock_manager, patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools",
|
||||
side_effect=lambda tools, _server: tools,
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions",
|
||||
new=AsyncMock(side_effect=lambda tools, **_: tools),
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
new=AsyncMock(return_value=[oauth2_server]),
|
||||
),
|
||||
patch(
|
||||
# Patch the bulk prefetch so no real DB connection is needed
|
||||
"litellm.proxy._experimental.mcp_server.server._prefetch_oauth_creds_for_user",
|
||||
new=AsyncMock(return_value=prefetched_creds),
|
||||
) as mock_prefetch,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
) as mock_manager,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools",
|
||||
side_effect=lambda tools, _server: tools,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions",
|
||||
new=AsyncMock(side_effect=lambda tools, **_: tools),
|
||||
),
|
||||
):
|
||||
mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1])
|
||||
|
||||
|
|
@ -2481,24 +2581,34 @@ class TestMergeGatewayInitializeInstructions:
|
|||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
global_mcp_server_manager._upstream_initialize_instructions_by_server_id["s1"] = "upstream"
|
||||
|
||||
global_mcp_server_manager._upstream_initialize_instructions_by_server_id[
|
||||
"s1"
|
||||
] = "upstream"
|
||||
try:
|
||||
s = _make_instruction_server(instructions="yaml wins")
|
||||
assert self._merge([s]) == "yaml wins"
|
||||
finally:
|
||||
global_mcp_server_manager._upstream_initialize_instructions_by_server_id.pop("s1", None)
|
||||
global_mcp_server_manager._upstream_initialize_instructions_by_server_id.pop(
|
||||
"s1", None
|
||||
)
|
||||
|
||||
def test_upstream_cache_used_when_no_yaml(self):
|
||||
"""Upstream cached instructions are used when no YAML override is set."""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
global_mcp_server_manager._upstream_initialize_instructions_by_server_id["s1"] = "from upstream"
|
||||
|
||||
global_mcp_server_manager._upstream_initialize_instructions_by_server_id[
|
||||
"s1"
|
||||
] = "from upstream"
|
||||
try:
|
||||
s = _make_instruction_server(instructions=None)
|
||||
assert self._merge([s]) == "from upstream"
|
||||
finally:
|
||||
global_mcp_server_manager._upstream_initialize_instructions_by_server_id.pop("s1", None)
|
||||
global_mcp_server_manager._upstream_initialize_instructions_by_server_id.pop(
|
||||
"s1", None
|
||||
)
|
||||
|
||||
def test_spec_path_servers_skipped(self):
|
||||
"""OpenAPI (spec_path) servers do not contribute instructions."""
|
||||
|
|
@ -2512,8 +2622,12 @@ class TestMergeGatewayInitializeInstructions:
|
|||
|
||||
def test_multiple_servers_merged_with_labels(self):
|
||||
"""Multiple servers get label-prefixed and separator-joined."""
|
||||
s1 = _make_instruction_server(server_id="a", name="a", alias="Alpha", instructions="instr A")
|
||||
s2 = _make_instruction_server(server_id="b", name="b", alias="Beta", instructions="instr B")
|
||||
s1 = _make_instruction_server(
|
||||
server_id="a", name="a", alias="Alpha", instructions="instr A"
|
||||
)
|
||||
s2 = _make_instruction_server(
|
||||
server_id="b", name="b", alias="Beta", instructions="instr B"
|
||||
)
|
||||
result = self._merge([s1, s2])
|
||||
assert result is not None
|
||||
assert "[Alpha]" in result and "[Beta]" in result
|
||||
|
|
@ -2532,17 +2646,26 @@ class TestMergeGatewayInitializeInstructions:
|
|||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
global_mcp_server_manager._upstream_initialize_instructions_by_server_id["c"] = "cached C"
|
||||
|
||||
global_mcp_server_manager._upstream_initialize_instructions_by_server_id[
|
||||
"c"
|
||||
] = "cached C"
|
||||
try:
|
||||
s_yaml = _make_instruction_server(server_id="a", name="a", alias="A", instructions="yaml A")
|
||||
s_spec = _make_instruction_server(server_id="b", name="b", alias="B", spec_path="/spec.json", url=None)
|
||||
s_yaml = _make_instruction_server(
|
||||
server_id="a", name="a", alias="A", instructions="yaml A"
|
||||
)
|
||||
s_spec = _make_instruction_server(
|
||||
server_id="b", name="b", alias="B", spec_path="/spec.json", url=None
|
||||
)
|
||||
s_cached = _make_instruction_server(server_id="c", name="c", alias="C")
|
||||
result = self._merge([s_yaml, s_spec, s_cached])
|
||||
assert "yaml A" in result
|
||||
assert "cached C" in result
|
||||
assert "[B]" not in result
|
||||
finally:
|
||||
global_mcp_server_manager._upstream_initialize_instructions_by_server_id.pop("c", None)
|
||||
global_mcp_server_manager._upstream_initialize_instructions_by_server_id.pop(
|
||||
"c", None
|
||||
)
|
||||
|
||||
|
||||
class TestGatewayCreateInitializationOptions:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue