mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
test fixes
This commit is contained in:
parent
2579169ccb
commit
2ac13de66f
5 changed files with 27 additions and 26 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue