style: black format test_mcp_server_manager.py

This commit is contained in:
Ishaan Jaffer 2026-04-15 18:19:17 -07:00
parent f768946549
commit fcd71e0026
No known key found for this signature in database

View file

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