mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
style: black format test_mcp_server_manager.py
This commit is contained in:
parent
f768946549
commit
fcd71e0026
1 changed files with 112 additions and 57 deletions
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue