test fixes

This commit is contained in:
Ishaan Jaffer 2026-02-09 14:56:04 -08:00
parent 2579169ccb
commit 2ac13de66f
5 changed files with 27 additions and 26 deletions

View file

@ -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,

View file

@ -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"

View file

@ -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,

View file

@ -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()

View file

@ -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()