From 2b5eb794fca2727aff280c18799aec7813f71121 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Tue, 14 Apr 2026 12:40:55 -0700 Subject: [PATCH] fix(mcp): set instructions=None in test_add_update_server_with_alias mock --- tests/mcp_tests/test_mcp_server.py | 192 +++++++++++++++++------------ 1 file changed, 115 insertions(+), 77 deletions(-) diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index 7c5dcc66a8..d0a4e9d0c9 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -417,17 +417,22 @@ async def test_streamable_http_mcp_handler_mock(): # Mock extract_mcp_auth_context to bypass auth checks in the handler mock_auth_context = (None, None, None, {}, {}, {}) - with patch( - "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", - True, - ), patch( - "litellm.proxy._experimental.mcp_server.server.session_manager", - mock_session_manager, - ), patch( - "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", - AsyncMock(return_value=mock_auth_context), - ), patch( - "litellm.proxy._experimental.mcp_server.server.set_auth_context", + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.session_manager", + mock_session_manager, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + AsyncMock(return_value=mock_auth_context), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + ), ): from litellm.proxy._experimental.mcp_server.server import ( handle_streamable_http_mcp, @@ -471,17 +476,22 @@ async def test_sse_mcp_handler_mock(): [], ) - with patch( - "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", - True, - ), patch( - "litellm.proxy._experimental.mcp_server.server.sse_session_manager", - mock_sse_session_manager, - ), patch( - "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", - new=AsyncMock(return_value=mock_auth_result), - ), patch( - "litellm.proxy._experimental.mcp_server.server.set_auth_context", + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.sse_session_manager", + mock_sse_session_manager, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new=AsyncMock(return_value=mock_auth_result), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + ), ): from litellm.proxy._experimental.mcp_server.server import handle_sse_mcp @@ -833,7 +843,9 @@ async def test_get_tools_from_mcp_servers(): mock_manager.get_allowed_mcp_servers = AsyncMock( return_value=["server1_id", "server2_id"] ) - 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_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_with_info to return input unchanged (no IP filtering in test) mock_manager.filter_server_ids_by_ip_with_info = MagicMock( @@ -859,7 +871,10 @@ async def test_get_tools_from_mcp_servers(): mock_manager_2.get_allowed_mcp_servers = AsyncMock( return_value=["server1_id", "server2_id"] ) - mock_manager_2.get_mcp_server_by_id = lambda server_id: mock_server_1 if server_id == "server1_id" else mock_server_2 + mock_manager_2.get_mcp_server_by_id = lambda server_id: ( + mock_server_1 if server_id == "server1_id" else mock_server_2 + ) + async def mock_get_tools_side_effect( server, mcp_auth_header=None, @@ -900,7 +915,11 @@ async def test_get_tools_from_mcp_servers(): mock_manager.get_allowed_mcp_servers = AsyncMock( return_value=["server1_id", "server2_id", "server3_id"] ) - 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_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_with_info to return input unchanged (no IP filtering in test) mock_manager.filter_server_ids_by_ip_with_info = MagicMock( @@ -1050,15 +1069,15 @@ async def test_mcp_server_manager_access_groups_from_config(): # Should find config_server for group-a, both for group-b, other_server for group-c import asyncio - server_ids_a = await MCPRequestHandler._get_mcp_servers_from_access_groups([ - "group-a" - ]) - server_ids_b = await MCPRequestHandler._get_mcp_servers_from_access_groups([ - "group-b" - ]) - server_ids_c = await MCPRequestHandler._get_mcp_servers_from_access_groups([ - "group-c" - ]) + server_ids_a = await MCPRequestHandler._get_mcp_servers_from_access_groups( + ["group-a"] + ) + server_ids_b = await MCPRequestHandler._get_mcp_servers_from_access_groups( + ["group-b"] + ) + server_ids_c = await MCPRequestHandler._get_mcp_servers_from_access_groups( + ["group-c"] + ) assert any(config_server.server_id == sid for sid in server_ids_a) assert set(server_ids_b) == set( [ @@ -1474,6 +1493,7 @@ async def test_add_update_server_with_alias(): mock_mcp_server.byok_api_key_help_url = None mock_mcp_server.created_at = None mock_mcp_server.updated_at = None + mock_mcp_server.instructions = None # Add server to manager await test_manager.add_server(mock_mcp_server) @@ -2151,8 +2171,12 @@ async def test_list_tool_rest_api_all_servers_with_auth(): for call_args in mock_get_tools.call_args_list } - assert server_auth_map.get(mock_zapier_server) == "Bearer zapier_token" - assert server_auth_map.get(mock_slack_server) == "Bearer slack_token" + assert ( + server_auth_map.get(mock_zapier_server) == "Bearer zapier_token" + ) + assert ( + server_auth_map.get(mock_slack_server) == "Bearer slack_token" + ) @pytest.mark.asyncio @@ -2690,26 +2714,33 @@ async def test_call_mcp_tool_uses_manager_permission_lookup(): expected_response = [TextContent(type="text", text="ok")] - with patch.object( - global_mcp_server_manager, - "get_allowed_mcp_servers", - new_callable=AsyncMock, - ) as mock_get_allowed, patch.object( - global_mcp_server_manager, - "get_mcp_server_by_id", - return_value=mock_server, - ), patch.object( - global_mcp_server_manager, - "_get_mcp_server_from_tool_name", - return_value=mock_server, - ) as mock_get_server, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_tool_registry" - ) as mock_tool_registry, patch( - "litellm.proxy._experimental.mcp_server.server._handle_managed_mcp_tool", - new_callable=AsyncMock, - ) as mock_handle_managed, patch( - "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed", - return_value=True, + with ( + patch.object( + global_mcp_server_manager, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + ) as mock_get_allowed, + patch.object( + global_mcp_server_manager, + "get_mcp_server_by_id", + return_value=mock_server, + ), + patch.object( + global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + return_value=mock_server, + ) as mock_get_server, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_tool_registry" + ) as mock_tool_registry, + patch( + "litellm.proxy._experimental.mcp_server.server._handle_managed_mcp_tool", + new_callable=AsyncMock, + ) as mock_handle_managed, + patch( + "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed", + return_value=True, + ), ): mock_get_allowed.return_value = [mock_server.server_id] mock_tool_registry.get_tool.return_value = None @@ -2759,27 +2790,34 @@ async def test_call_mcp_tool_resolves_unprefixed_tool_name_and_checks_permission expected_response = [TextContent(type="text", text="ok")] - with patch.object( - global_mcp_server_manager, - "get_allowed_mcp_servers", - new_callable=AsyncMock, - ) as mock_get_allowed, patch.object( - global_mcp_server_manager, - "get_mcp_server_by_id", - return_value=mock_server, - ), patch.object( - global_mcp_server_manager, - "_get_mcp_server_from_tool_name", - return_value=mock_server, - ) as mock_get_server, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_tool_registry" - ) as mock_tool_registry, patch( - "litellm.proxy._experimental.mcp_server.server._handle_managed_mcp_tool", - new_callable=AsyncMock, - ) as mock_handle_managed, patch( - "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed", - return_value=True, - ) as mock_is_allowed: + with ( + patch.object( + global_mcp_server_manager, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + ) as mock_get_allowed, + patch.object( + global_mcp_server_manager, + "get_mcp_server_by_id", + return_value=mock_server, + ), + patch.object( + global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + return_value=mock_server, + ) as mock_get_server, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_tool_registry" + ) as mock_tool_registry, + patch( + "litellm.proxy._experimental.mcp_server.server._handle_managed_mcp_tool", + new_callable=AsyncMock, + ) as mock_handle_managed, + patch( + "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed", + return_value=True, + ) as mock_is_allowed, + ): mock_get_allowed.return_value = [mock_server.server_id] mock_tool_registry.get_tool.return_value = None mock_handle_managed.return_value = expected_response