From fcd71e0026bfd9885f96f6d560872d4c32063142 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Wed, 15 Apr 2026 18:19:17 -0700 Subject: [PATCH] style: black format test_mcp_server_manager.py --- .../mcp_server/test_mcp_server_manager.py | 169 ++++++++++++------ 1 file changed, 112 insertions(+), 57 deletions(-) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index aa95836a927..ac5349e7105 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -43,10 +43,10 @@ def _reload_mcp_manager_module(): # After reload, server.py still holds a stale reference to the old # global_mcp_server_manager. Update it so tests that exercise server.py # functions (e.g. _get_tools_from_mcp_servers) use the fresh instance. - server_module = sys.modules.get( - "litellm.proxy._experimental.mcp_server.server" - ) - if server_module is not None and hasattr(server_module, "global_mcp_server_manager"): + server_module = sys.modules.get("litellm.proxy._experimental.mcp_server.server") + if server_module is not None and hasattr( + server_module, "global_mcp_server_manager" + ): server_module.global_mcp_server_manager = reloaded.global_mcp_server_manager return reloaded @@ -223,9 +223,7 @@ class TestMCPServerManager: with caplog.at_level(logging.WARNING, logger="LiteLLM"): await manager.load_servers_from_config(config) - assert any( - "invalid alias 'bad/name'" in message for message in caplog.messages - ) + assert any("invalid alias 'bad/name'" in message for message in caplog.messages) @pytest.mark.asyncio async def test_load_servers_from_config_accepts_valid_alias(self, caplog): @@ -492,7 +490,12 @@ class TestMCPServerManager: mock_client = AsyncMock() mock_client.list_prompts = AsyncMock(return_value=[mock_prompt]) - with patch.object(manager, "_create_mcp_client", new_callable=AsyncMock, return_value=mock_client): + with patch.object( + manager, + "_create_mcp_client", + new_callable=AsyncMock, + return_value=mock_client, + ): prompts = await manager.get_prompts_from_server(server, add_prefix=True) mock_client.list_prompts.assert_awaited_once() @@ -520,7 +523,12 @@ class TestMCPServerManager: mock_client = AsyncMock() mock_client.get_prompt = AsyncMock(return_value=mock_result) - with patch.object(manager, "_create_mcp_client", new_callable=AsyncMock, return_value=mock_client): + with patch.object( + manager, + "_create_mcp_client", + new_callable=AsyncMock, + return_value=mock_client, + ): result = await manager.get_prompt_from_server( server=server, prompt_name="hello", @@ -551,13 +559,23 @@ class TestMCPServerManager: mock_client = AsyncMock() mock_resources = [Resource(name="file", uri="https://example.com/file")] mock_client.list_resources = AsyncMock(return_value=mock_resources) - prefixed_resources = [Resource(name="alias-server-file", uri="https://example.com/file")] + prefixed_resources = [ + Resource(name="alias-server-file", uri="https://example.com/file") + ] - with patch.object(manager, "_create_mcp_client", new_callable=AsyncMock, return_value=mock_client) as mock_create_client, patch.object( - manager, - "_create_prefixed_resources", - return_value=prefixed_resources, - ) as mock_prefix: + with ( + patch.object( + manager, + "_create_mcp_client", + new_callable=AsyncMock, + return_value=mock_client, + ) as mock_create_client, + patch.object( + manager, + "_create_prefixed_resources", + return_value=prefixed_resources, + ) as mock_prefix, + ): result = await manager.get_resources_from_server( server=server, mcp_auth_header="auth", @@ -602,11 +620,19 @@ class TestMCPServerManager: ) ] - with patch.object(manager, "_create_mcp_client", new_callable=AsyncMock, return_value=mock_client) as mock_create_client, patch.object( - manager, - "_create_prefixed_resource_templates", - return_value=prefixed_templates, - ) as mock_prefix: + with ( + patch.object( + manager, + "_create_mcp_client", + new_callable=AsyncMock, + return_value=mock_client, + ) as mock_create_client, + patch.object( + manager, + "_create_prefixed_resource_templates", + return_value=prefixed_templates, + ) as mock_prefix, + ): result = await manager.get_resource_templates_from_server( server=server, mcp_auth_header="auth", @@ -650,7 +676,12 @@ class TestMCPServerManager: ) mock_client.read_resource = AsyncMock(return_value=read_result) - with patch.object(manager, "_create_mcp_client", new_callable=AsyncMock, return_value=mock_client) as mock_create_client: + with patch.object( + manager, + "_create_mcp_client", + new_callable=AsyncMock, + return_value=mock_client, + ) as mock_create_client: result = await manager.read_resource_from_server( server=server, url="https://example.com/resource", @@ -661,7 +692,9 @@ class TestMCPServerManager: mock_create_client.assert_called_once() called_kwargs = mock_create_client.call_args.kwargs assert called_kwargs["extra_headers"] == {"X-Test": "1", "X-Static": "1"} - mock_client.read_resource.assert_awaited_once_with("https://example.com/resource") + mock_client.read_resource.assert_awaited_once_with( + "https://example.com/resource" + ) assert result is read_result @pytest.mark.asyncio @@ -724,22 +757,27 @@ class TestMCPServerManager: registration_url=None, ) - with patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", - return_value=mock_client, - ), patch.object( - manager, - "_fetch_oauth_metadata_from_resource", - AsyncMock(return_value=([], None)), - ), patch.object( - manager, - "_attempt_well_known_discovery", - AsyncMock(return_value=([], None)), - ), patch.object( - manager, - "_fetch_authorization_server_metadata", - AsyncMock(return_value=mock_metadata), - ) as mock_fetch_auth: + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client", + return_value=mock_client, + ), + patch.object( + manager, + "_fetch_oauth_metadata_from_resource", + AsyncMock(return_value=([], None)), + ), + patch.object( + manager, + "_attempt_well_known_discovery", + AsyncMock(return_value=([], None)), + ), + patch.object( + manager, + "_fetch_authorization_server_metadata", + AsyncMock(return_value=mock_metadata), + ) as mock_fetch_auth, + ): result = await manager._descovery_metadata(server_url) mock_fetch_auth.assert_awaited_once_with(["https://example.com"]) @@ -779,9 +817,8 @@ class TestMCPServerManager: assert server.scopes == ["config"] # config overrides discovery assert server.authorization_url == "https://config.example.com/auth" assert server.token_url == "https://discovered.example.com/token" - assert ( - server.registration_url == "https://discovered.example.com/register" - ) + assert server.registration_url == "https://discovered.example.com/register" + @pytest.mark.asyncio async def test_config_oauth_initialize_tool_name_to_mcp_server_name_mapping(self): manager = MCPServerManager() @@ -801,7 +838,7 @@ class TestMCPServerManager: # Initialize the tool mapping await manager._initialize_tool_name_to_mcp_server_name_mapping() assert manager.tool_name_to_mcp_server_name_mapping == {} - + @pytest.mark.asyncio async def test_list_tools_handles_missing_server_alias(self): """Test that list_tools handles servers without alias gracefully""" @@ -1017,7 +1054,9 @@ class TestMCPServerManager: # Capture the extra_headers passed to _create_mcp_client captured_extra_headers = None - async def capture_create_mcp_client(server, mcp_auth_header, extra_headers, stdio_env): + async def capture_create_mcp_client( + server, mcp_auth_header, extra_headers, stdio_env + ): nonlocal captured_extra_headers captured_extra_headers = extra_headers return mock_client @@ -1314,15 +1353,19 @@ class TestMCPServerManager: return tool_func - with patch( - "litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.create_tool_function", - side_effect=fake_create_tool_function, - ), patch( - "litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.build_input_schema", - return_value={"type": "object", "properties": {}, "required": []}, - ), patch( - "litellm.proxy._experimental.mcp_server.tool_registry.global_mcp_tool_registry.register_tool", - return_value=None, + with ( + patch( + "litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.create_tool_function", + side_effect=fake_create_tool_function, + ), + patch( + "litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.build_input_schema", + return_value={"type": "object", "properties": {}, "required": []}, + ), + patch( + "litellm.proxy._experimental.mcp_server.tool_registry.global_mcp_tool_registry.register_tool", + return_value=None, + ), ): await manager._register_openapi_tools( spec_path=str(spec_path), @@ -2161,7 +2204,9 @@ class TestMCPServerManager: # Register the server and map a tool to it manager.registry = {"test-server": server} manager.tool_name_to_mcp_server_name_mapping["test_tool"] = "test-server" - manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = "test-server" + manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = ( + "test-server" + ) # Create mock client that tracks call_tool usage mock_client = AsyncMock() @@ -2252,11 +2297,16 @@ class TestMCPServerManager: # Verify MCPRequestHandler.get_allowed_mcp_servers was called with user_api_key_auth mock_get_allowed.assert_called_once() call_args = mock_get_allowed.call_args - assert call_args[0][0] is user_api_key_auth # First positional arg should be user_api_key_auth + assert ( + call_args[0][0] is user_api_key_auth + ) # First positional arg should be user_api_key_auth assert call_args[0][0].user_id == "user-123" assert call_args[0][0].object_permission_id == "perm_123" assert call_args[0][0].object_permission is not None - assert call_args[0][0].object_permission.mcp_servers == ["test_server_1", "test_server_2"] + assert call_args[0][0].object_permission.mcp_servers == [ + "test_server_1", + "test_server_2", + ] # Verify result contains the expected servers assert "test_server_1" in result @@ -2494,7 +2544,10 @@ class TestMCPServerManagerUpstreamInstructionsCache: def test_get_returns_none_when_empty(self): """Empty cache returns None for any key.""" manager = MCPServerManager() - assert manager._upstream_initialize_instructions_by_server_id.get("nonexistent") is None + assert ( + manager._upstream_initialize_instructions_by_server_id.get("nonexistent") + is None + ) def test_remember_stores_stripped_value(self): """_remember_upstream_initialize_instructions stores a stripped string.""" @@ -2502,7 +2555,9 @@ class TestMCPServerManagerUpstreamInstructionsCache: fake_server = MagicMock(server_id="srv") fake_client = MagicMock(_last_initialize_instructions=" hello \n") manager._remember_upstream_initialize_instructions(fake_server, fake_client) - assert manager._upstream_initialize_instructions_by_server_id.get("srv") == "hello" + assert ( + manager._upstream_initialize_instructions_by_server_id.get("srv") == "hello" + ) def test_remember_ignores_empty_string(self): """Whitespace-only instructions are not stored."""