diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 2fe40bb197e..65d3a2caa16 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -541,7 +541,7 @@ if MCP_AVAILABLE: static_headers=request.static_headers, ) - client = global_mcp_server_manager._create_mcp_client( + client = await global_mcp_server_manager._create_mcp_client( server=server_model, mcp_auth_header=mcp_auth_header, extra_headers=merged_headers, diff --git a/tests/mcp_tests/test_mcp_auth_priority.py b/tests/mcp_tests/test_mcp_auth_priority.py index bc217a7903b..ad6e9438edd 100644 --- a/tests/mcp_tests/test_mcp_auth_priority.py +++ b/tests/mcp_tests/test_mcp_auth_priority.py @@ -12,7 +12,8 @@ from litellm.types.mcp import MCPAuth, MCPTransport, MCPSpecVersion from litellm.types.mcp_server.mcp_server_manager import MCPServer -def test_mcp_server_works_without_config_auth_value(): +@pytest.mark.asyncio +async def test_mcp_server_works_without_config_auth_value(): """ Test that MCP servers work without auth_value in config when headers are provided. This validates that auth_value is truly optional in config.yaml. @@ -32,7 +33,7 @@ def test_mcp_server_works_without_config_auth_value(): manager = MCPServerManager() # Test that it works with only header auth - client = manager._create_mcp_client( + client = await manager._create_mcp_client( server=server_without_config_auth, mcp_auth_header="Bearer token_from_header_only", ) @@ -58,7 +59,7 @@ async def test_mcp_server_config_auth_value_header_used(token_key): await manager.load_servers_from_config(config) server = next(iter(manager.config_mcp_servers.values())) - client = manager._create_mcp_client(server) + client = await manager._create_mcp_client(server) headers = client._get_auth_headers() assert headers["Authorization"] == "Bearer example_token" 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 9aa7de89eff..07c3dfcc763 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 @@ -886,7 +886,7 @@ async def test_oauth2_headers_passed_to_mcp_client(): # This will capture the arguments passed to _create_mcp_client captured_client_args = {} - def mock_create_mcp_client( + async def mock_create_mcp_client( server, mcp_auth_header=None, extra_headers=None, 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 f241b2aa0d5..abb8dd49159 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 @@ -1,5 +1,5 @@ -import json import importlib +import json import logging import os import sys @@ -21,8 +21,8 @@ from mcp.types import ( Prompt, ResourceTemplate, TextResourceContents, - Tool as MCPTool, ) +from mcp.types import Tool as MCPTool from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( MCPServerManager, @@ -92,7 +92,7 @@ class TestMCPServerManager: assert added_server.args == ["-m", "server"] assert added_server.env == {"DEBUG": "1", "TEST": "1"} - def test_create_mcp_client_stdio(self): + async def test_create_mcp_client_stdio(self): """Test creating MCP client for stdio transport""" manager = MCPServerManager() @@ -106,7 +106,7 @@ class TestMCPServerManager: env={"NODE_ENV": "test"}, ) - client = manager._create_mcp_client(stdio_server) + client = await manager._create_mcp_client(stdio_server) assert client.transport_type == MCPTransport.stdio assert client.stdio_config is not None @@ -404,14 +404,14 @@ class TestMCPServerManager: ) captured_extra_headers = None - def capture_create_mcp_client( + async def capture_create_mcp_client( server, mcp_auth_header, extra_headers, stdio_env ): # pragma: no cover - helper nonlocal captured_extra_headers captured_extra_headers = extra_headers return mock_client - manager._create_mcp_client = MagicMock(side_effect=capture_create_mcp_client) + manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) result = await manager._call_regular_mcp_tool( mcp_server=server, @@ -446,7 +446,7 @@ class TestMCPServerManager: mock_client = AsyncMock() mock_client.list_prompts = AsyncMock(return_value=[mock_prompt]) - with patch.object(manager, "_create_mcp_client", 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() @@ -474,7 +474,7 @@ class TestMCPServerManager: mock_client = AsyncMock() mock_client.get_prompt = AsyncMock(return_value=mock_result) - with patch.object(manager, "_create_mcp_client", 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", @@ -507,7 +507,7 @@ class TestMCPServerManager: mock_client.list_resources = AsyncMock(return_value=mock_resources) prefixed_resources = [Resource(name="alias-server-file", uri="https://example.com/file")] - with patch.object(manager, "_create_mcp_client", return_value=mock_client) as mock_create_client, patch.object( + 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, @@ -556,7 +556,7 @@ class TestMCPServerManager: ) ] - with patch.object(manager, "_create_mcp_client", return_value=mock_client) as mock_create_client, patch.object( + 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, @@ -604,7 +604,7 @@ class TestMCPServerManager: ) mock_client.read_resource = AsyncMock(return_value=read_result) - with patch.object(manager, "_create_mcp_client", 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", @@ -816,7 +816,7 @@ class TestMCPServerManager: # Mock successful client.run_with_session mock_client = AsyncMock() mock_client.run_with_session = AsyncMock(return_value="ok") - manager._create_mcp_client = MagicMock(return_value=mock_client) + manager._create_mcp_client = AsyncMock(return_value=mock_client) # Perform health check result = await manager.health_check_server("test-server") @@ -850,7 +850,7 @@ class TestMCPServerManager: mock_client.run_with_session = AsyncMock( side_effect=Exception("Connection timeout") ) - manager._create_mcp_client = MagicMock(return_value=mock_client) + manager._create_mcp_client = AsyncMock(return_value=mock_client) # Perform health check result = await manager.health_check_server("test-server") @@ -898,7 +898,7 @@ class TestMCPServerManager: manager.get_mcp_server_by_id = MagicMock(return_value=server) # _create_mcp_client should not be called for OAuth2 servers - manager._create_mcp_client = MagicMock() + manager._create_mcp_client = AsyncMock() # Perform health check result = await manager.health_check_server("oauth2-server") @@ -931,7 +931,7 @@ class TestMCPServerManager: manager.get_mcp_server_by_id = MagicMock(return_value=server) # _create_mcp_client should not be called - manager._create_mcp_client = MagicMock() + manager._create_mcp_client = AsyncMock() # Perform health check result = await manager.health_check_server("no-token-server") @@ -971,12 +971,12 @@ class TestMCPServerManager: # Capture the extra_headers passed to _create_mcp_client captured_extra_headers = None - 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 - manager._create_mcp_client = MagicMock(side_effect=capture_create_mcp_client) + manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) # Perform health check result = await manager.health_check_server("test-server") @@ -1310,7 +1310,7 @@ class TestMCPServerManager: ) # Mock client creation and fetching tools - manager._create_mcp_client = MagicMock(return_value=object()) + manager._create_mcp_client = AsyncMock(return_value=object()) # Tools returned upstream (unprefixed from provider) upstream_tool = MCPTool( @@ -1895,7 +1895,7 @@ class TestMCPServerManager: mock_client.call_tool.side_effect = mock_call_tool # Mock _create_mcp_client to return our mock client - manager._create_mcp_client = MagicMock(return_value=mock_client) + manager._create_mcp_client = AsyncMock(return_value=mock_client) # Mock user auth with no restrictions user_api_key_auth = MagicMock() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 3091f03a53d..f6c8115f7bc 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -76,7 +76,7 @@ def _route_has_dependency(route, dependency) -> bool: class TestExecuteWithMcpClient: @pytest.mark.asyncio async def test_redacts_stack_trace(self, monkeypatch): - def fake_create_client(*args, **kwargs): + async def fake_create_client(*args, **kwargs): return object() monkeypatch.setattr( @@ -114,7 +114,7 @@ class TestExecuteWithMcpClient: def fake_build_stdio_env(server, raw_headers): return None - def fake_create_client(*args, **kwargs): + async def fake_create_client(*args, **kwargs): captured["extra_headers"] = kwargs.get("extra_headers") return object()