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 acba06afee4..9df6408b0d7 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 @@ -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: