issues resolve

This commit is contained in:
Yug 2026-04-29 11:17:27 +05:30
parent 0cb5d9b1e4
commit 9b710c3502
8 changed files with 1060 additions and 386 deletions

File diff suppressed because it is too large Load diff

View file

@ -155,7 +155,11 @@ def _convert_single_content(content: Any) -> Dict[str, Any]:
# ToolResultContent → represents tool results
tool_content = getattr(content, "content", [])
if isinstance(tool_content, list) and tool_content:
texts = [getattr(c, "text", str(c)) for c in tool_content if getattr(c, "type", None) == "text"]
texts = [
getattr(c, "text", str(c))
for c in tool_content
if getattr(c, "type", None) == "text"
]
return {"type": "text", "text": "\n".join(texts) if texts else ""}
return {"type": "text", "text": str(tool_content)}
# Fallback: treat as text
@ -242,7 +246,9 @@ def _extract_tool_calls(content: Any) -> List[Dict[str, Any]]:
"type": "function",
"function": {
"name": getattr(item, "name", ""),
"arguments": json.dumps(getattr(item, "input", {}), default=str),
"arguments": json.dumps(
getattr(item, "input", {}), default=str
),
},
}
)
@ -269,7 +275,11 @@ def _extract_tool_results(content: Any) -> List[Dict[str, Any]]:
# Extract text from nested content
nested_content = getattr(item, "content", [])
if isinstance(nested_content, list):
text_parts = [getattr(c, "text", str(c)) for c in nested_content if getattr(c, "type", None) == "text"]
text_parts = [
getattr(c, "text", str(c))
for c in nested_content
if getattr(c, "type", None) == "text"
]
result_text = "\n".join(text_parts) if text_parts else ""
else:
result_text = str(nested_content)

File diff suppressed because it is too large Load diff

View file

@ -28,19 +28,25 @@ async def main():
return ElicitResult(action="accept", content=user_response)
async with sse_client("http://localhost:4000/mcp/sse", headers={"Authorization": "Bearer sk-1234"}) as (
async with sse_client(
"http://localhost:4000/mcp/sse", headers={"Authorization": "Bearer sk-1234"}
) as (
read_stream,
write_stream,
):
logger.info("SSE connection established.")
async with ClientSession(read_stream, write_stream, elicitation_callback=my_elicitation_callback) as session:
async with ClientSession(
read_stream, write_stream, elicitation_callback=my_elicitation_callback
) as session:
await session.initialize()
logger.info("Initialized!")
logger.info("\n--- Testing Complex Pipeline (Elicitation + Sampling) ---")
logger.info("Calling 'test_server-test_complex_pipeline'...")
try:
result = await session.call_tool("test_server-test_complex_pipeline", arguments={})
result = await session.call_tool(
"test_server-test_complex_pipeline", arguments={}
)
logger.info("\nFINAL TOOL RESULT:")
logger.info("==================")
logger.info(result.content[0].text)

View file

@ -31,7 +31,11 @@ async def test_auth_context_persistence():
@pytest.mark.asyncio
async def test_get_or_extract_auth_context_fallback():
"""Test get_or_extract_auth_context fallback to server object."""
from litellm.proxy._experimental.mcp_server.server import server, MCPAuthenticatedUser, auth_context_var
from litellm.proxy._experimental.mcp_server.server import (
server,
MCPAuthenticatedUser,
auth_context_var,
)
auth_data = UserAPIKeyAuth(api_key="fallback-key")
auth_user = MCPAuthenticatedUser(user_api_key_auth=auth_data)
@ -61,7 +65,9 @@ async def test_extract_mcp_auth_context_with_key():
mock_user_auth = UserAPIKeyAuth(api_key="sk-123")
with patch("litellm.proxy.auth.auth_checks.common_checks", new_callable=AsyncMock) as mock_auth:
with patch(
"litellm.proxy.auth.auth_checks.common_checks", new_callable=AsyncMock
) as mock_auth:
mock_auth.return_value = mock_user_auth
result = await extract_mcp_auth_context(mock_scope, "/mcp/sse")

View file

@ -22,7 +22,10 @@ def test_convert_image_content():
mock_image.mimeType = "image/jpeg"
result = _convert_single_content(mock_image)
assert result == {"type": "image_url", "image_url": {"url": "data:image/jpeg;base64,base64data"}}
assert result == {
"type": "image_url",
"image_url": {"url": "data:image/jpeg;base64,base64data"},
}
def test_convert_audio_content():
@ -32,7 +35,10 @@ def test_convert_audio_content():
mock_audio.mimeType = "audio/mp3"
result = _convert_single_content(mock_audio)
assert result == {"type": "input_audio", "input_audio": {"data": "audiobase64", "format": "mp3"}}
assert result == {
"type": "input_audio",
"input_audio": {"data": "audiobase64", "format": "mp3"},
}
def test_convert_list_content():
@ -64,7 +70,10 @@ def test_resolve_model_from_hints():
original_router = proxy_server.llm_router
try:
proxy_server.llm_router = MagicMock()
proxy_server.llm_router.get_model_names.return_value = ["gpt-4", "claude-3-5-sonnet"]
proxy_server.llm_router.get_model_names.return_value = [
"gpt-4",
"claude-3-5-sonnet",
]
result = _resolve_model_from_preferences(mock_prefs)
assert result == "claude-3-5-sonnet"
finally:

View file

@ -78,7 +78,9 @@ def _make_params(
return params
def _make_completion_response(content="Hello!", model="gpt-4o-mini", finish_reason="stop", tool_calls=None):
def _make_completion_response(
content="Hello!", model="gpt-4o-mini", finish_reason="stop", tool_calls=None
):
"""Create a mock litellm completion response."""
response = MagicMock()
choice = MagicMock()
@ -155,7 +157,9 @@ class TestConvertMCPMessagesToOpenAI:
)
user_msg = _make_sampling_message("user", _make_text_content("Hi"))
assistant_msg = _make_sampling_message("assistant", _make_text_content("Hello!"))
assistant_msg = _make_sampling_message(
"assistant", _make_text_content("Hello!")
)
result = _convert_mcp_messages_to_openai([user_msg, assistant_msg])
assert len(result) == 2
assert result[0]["role"] == "user"
@ -173,7 +177,9 @@ class TestResolveModel:
_resolve_model_from_preferences,
)
result = _resolve_model_from_preferences(None, default_model="claude-3.5-sonnet")
result = _resolve_model_from_preferences(
None, default_model="claude-3.5-sonnet"
)
assert result == "claude-3.5-sonnet"
@patch("litellm.proxy.proxy_server.llm_router", None)
@ -309,7 +315,9 @@ class TestConvertOpenAIResponseToMCPResult:
_convert_openai_response_to_mcp_result,
)
response = _make_completion_response(content="Partial...", finish_reason="length")
response = _make_completion_response(
content="Partial...", finish_reason="length"
)
result = _convert_openai_response_to_mcp_result(response, "gpt-4o")
assert result.stopReason == "maxTokens"
@ -322,7 +330,9 @@ class TestConvertOpenAIResponseToMCPResult:
tc.id = "call_123"
tc.function.name = "get_weather"
tc.function.arguments = '{"city": "NYC"}'
response = _make_completion_response(content=None, finish_reason="tool_calls", tool_calls=[tc])
response = _make_completion_response(
content=None, finish_reason="tool_calls", tool_calls=[tc]
)
result = _convert_openai_response_to_mcp_result(response, "gpt-4o")
assert result.stopReason == "toolUse"
assert isinstance(result.content, list)

View file

@ -35,14 +35,18 @@ from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPSer
def _reload_mcp_manager_module():
utils_module = sys.modules["litellm.proxy._experimental.mcp_server.utils"]
manager_module = sys.modules["litellm.proxy._experimental.mcp_server.mcp_server_manager"]
manager_module = sys.modules[
"litellm.proxy._experimental.mcp_server.mcp_server_manager"
]
importlib.reload(utils_module)
reloaded = importlib.reload(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"):
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
@ -303,7 +307,9 @@ class TestMCPServerManager:
# Mock get_allowed_mcp_servers to return our test servers
manager.get_allowed_mcp_servers = AsyncMock(return_value=["github", "zapier"])
manager.get_mcp_server_by_id = MagicMock(side_effect=lambda x: server1 if x == "github" else server2)
manager.get_mcp_server_by_id = MagicMock(
side_effect=lambda x: server1 if x == "github" else server2
)
# Mock _get_tools_from_server to return different results
async def mock_get_tools_from_server(
@ -331,7 +337,9 @@ class TestMCPServerManager:
"zapier": "zapier-api-key",
}
result = await manager.list_tools(mcp_server_auth_headers=mcp_server_auth_headers)
result = await manager.list_tools(
mcp_server_auth_headers=mcp_server_auth_headers
)
# Verify that both servers were called with their specific auth headers
assert len(result) == 3 # 2 from github + 1 from zapier
@ -402,7 +410,9 @@ class TestMCPServerManager:
mcp_protocol_version=None,
raw_headers=None,
):
assert mcp_auth_header == "server-specific-token" # Should use server-specific header
assert (
mcp_auth_header == "server-specific-token"
) # Should use server-specific header
tool = MagicMock()
tool.name = "github_tool_1"
return [tool]
@ -433,7 +443,9 @@ class TestMCPServerManager:
)
mock_client = AsyncMock()
mock_client.call_tool = AsyncMock(return_value=CallToolResult(content=[], isError=False))
mock_client.call_tool = AsyncMock(
return_value=CallToolResult(content=[], isError=False)
)
captured_extra_headers = None
async def capture_create_mcp_client(
@ -547,7 +559,9 @@ 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(
@ -678,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
@ -725,7 +741,9 @@ class TestMCPServerManager:
)
def raise_http_error():
raise httpx.HTTPStatusError("unauthorized", request=request, response=response_obj)
raise httpx.HTTPStatusError(
"unauthorized", request=request, response=response_obj
)
response_obj.raise_for_status = MagicMock(side_effect=raise_http_error)
@ -843,7 +861,9 @@ class TestMCPServerManager:
mcp_protocol_version=None,
raw_headers=None,
):
assert mcp_auth_header == "server-specific-token" # Should use server-specific header via server_name
assert (
mcp_auth_header == "server-specific-token"
) # Should use server-specific header via server_name
tool = MagicMock()
tool.name = "github_tool_1"
return [tool]
@ -910,7 +930,9 @@ class TestMCPServerManager:
# Mock failed client.run_with_session
mock_client = AsyncMock()
mock_client.run_with_session = AsyncMock(side_effect=Exception("Connection timeout"))
mock_client.run_with_session = AsyncMock(
side_effect=Exception("Connection timeout")
)
manager._create_mcp_client = AsyncMock(return_value=mock_client)
# Perform health check
@ -1032,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
@ -1374,7 +1398,9 @@ class TestMCPServerManager:
proxy_logging_obj = MagicMock()
# Mock the async methods that pre_call_tool_check calls
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(
return_value={}
)
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
@ -1418,8 +1444,13 @@ class TestMCPServerManager:
)
assert exc_info.value.status_code == 403
assert "Tool blocked_tool is not allowed for server test-server" in exc_info.value.detail["error"]
assert "Contact proxy admin to allow this tool" in exc_info.value.detail["error"]
assert (
"Tool blocked_tool is not allowed for server test-server"
in exc_info.value.detail["error"]
)
assert (
"Contact proxy admin to allow this tool" in exc_info.value.detail["error"]
)
@pytest.mark.asyncio
async def test_pre_call_tool_check_disallowed_tools_list_allows_tool(self):
@ -1443,7 +1474,9 @@ class TestMCPServerManager:
proxy_logging_obj = MagicMock()
# Mock the async methods that pre_call_tool_check calls
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(
return_value={}
)
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
@ -1487,8 +1520,13 @@ class TestMCPServerManager:
)
assert exc_info.value.status_code == 403
assert "Tool banned_tool is not allowed for server test-server" in exc_info.value.detail["error"]
assert "Contact proxy admin to allow this tool" in exc_info.value.detail["error"]
assert (
"Tool banned_tool is not allowed for server test-server"
in exc_info.value.detail["error"]
)
assert (
"Contact proxy admin to allow this tool" in exc_info.value.detail["error"]
)
@pytest.mark.asyncio
async def test_pre_call_tool_check_no_restrictions_allows_any_tool(self):
@ -1512,7 +1550,9 @@ class TestMCPServerManager:
proxy_logging_obj = MagicMock()
# Mock the async methods that pre_call_tool_check calls
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(
return_value={}
)
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
@ -1549,7 +1589,9 @@ class TestMCPServerManager:
proxy_logging_obj = MagicMock()
# Mock the async methods that pre_call_tool_check calls
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(
return_value={}
)
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
@ -1575,7 +1617,10 @@ class TestMCPServerManager:
)
assert exc_info.value.status_code == 403
assert "Tool tool3 is not allowed for server test-server" in exc_info.value.detail["error"]
assert (
"Tool tool3 is not allowed for server test-server"
in exc_info.value.detail["error"]
)
async def test_get_tools_from_server_add_prefix(self):
"""Verify _get_tools_from_server respects add_prefix True/False."""
@ -1606,7 +1651,9 @@ class TestMCPServerManager:
assert tools_prefixed[0].name == "zapier-send_email"
# Case 2: add_prefix=False (single-server) -> expect unprefixed
tools_unprefixed = await manager._get_tools_from_server(server, add_prefix=False)
tools_unprefixed = await manager._get_tools_from_server(
server, add_prefix=False
)
assert len(tools_unprefixed) == 1
assert tools_unprefixed[0].name == "send_email"
@ -1641,9 +1688,13 @@ class TestMCPServerManager:
# Mapping should include both original and prefixed names -> resolves calls either way
assert manager.tool_name_to_mcp_server_name_mapping["create_issue"] == "jira"
assert manager.tool_name_to_mcp_server_name_mapping["jira-create_issue"] == "jira"
assert (
manager.tool_name_to_mcp_server_name_mapping["jira-create_issue"] == "jira"
)
assert manager.tool_name_to_mcp_server_name_mapping["close_issue"] == "jira"
assert manager.tool_name_to_mcp_server_name_mapping["jira-close_issue"] == "jira"
assert (
manager.tool_name_to_mcp_server_name_mapping["jira-close_issue"] == "jira"
)
def test_get_mcp_server_from_tool_name_with_prefixed_and_unprefixed(self):
"""After mapping is populated, manager resolves both prefixed and unprefixed tool names to the same server."""
@ -1673,7 +1724,9 @@ class TestMCPServerManager:
assert resolved_server_unpref.server_id == server.server_id
# Prefixed resolution
resolved_server_pref = manager._get_mcp_server_from_tool_name("zapier-create_zap")
resolved_server_pref = manager._get_mcp_server_from_tool_name(
"zapier-create_zap"
)
assert resolved_server_pref is not None
assert resolved_server_pref.server_id == server.server_id
@ -1718,7 +1771,9 @@ class TestMCPServerManager:
new=AsyncMock(return_value=[tool1, tool2, tool3]),
):
# Call the REST endpoint helper
filtered_response = await _get_tools_for_single_server(server, server_auth_header=None)
filtered_response = await _get_tools_for_single_server(
server, server_auth_header=None
)
# Verify only allowed tools are in the response
assert len(filtered_response) == 2
@ -1768,7 +1823,9 @@ class TestMCPServerManager:
new=AsyncMock(return_value=[tool1, tool2, tool3]),
):
# Call the REST endpoint helper
all_tools_response = await _get_tools_for_single_server(server, server_auth_header=None)
all_tools_response = await _get_tools_for_single_server(
server, server_auth_header=None
)
# Verify all tools are returned (no filtering)
assert len(all_tools_response) == 3
@ -1813,7 +1870,9 @@ class TestMCPServerManager:
new=AsyncMock(return_value=[tool1, tool2]),
):
# Call the REST endpoint helper
all_tools_response = await _get_tools_for_single_server(server, server_auth_header=None)
all_tools_response = await _get_tools_for_single_server(
server, server_auth_header=None
)
# Verify all tools are returned (no filtering)
assert len(all_tools_response) == 2
@ -1883,7 +1942,9 @@ class TestMCPServerManager:
)
proxy_logging = MagicMock()
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(
return_value={}
)
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging.pre_call_hook = AsyncMock(return_value=None)
@ -1926,7 +1987,9 @@ class TestMCPServerManager:
)
proxy_logging = MagicMock()
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(
return_value={}
)
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging.pre_call_hook = AsyncMock(return_value=None)
@ -2071,7 +2134,9 @@ class TestMCPServerManager:
proxy_logging_obj = MagicMock()
# Mock the async methods that pre_call_tool_check calls
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(
return_value={}
)
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
@ -2107,8 +2172,13 @@ class TestMCPServerManager:
)
assert exc_info.value.status_code == 403
assert "Tool deletepet is not allowed for server my_api_mcp" in exc_info.value.detail["error"]
assert "Contact proxy admin to allow this tool" in exc_info.value.detail["error"]
assert (
"Tool deletepet is not allowed for server my_api_mcp"
in exc_info.value.detail["error"]
)
assert (
"Contact proxy admin to allow this tool" in exc_info.value.detail["error"]
)
@pytest.mark.asyncio
async def test_call_tool_without_broken_pipe_error(self):
@ -2133,7 +2203,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()
@ -2157,7 +2229,9 @@ class TestMCPServerManager:
# Mock proxy logging
proxy_logging_obj = MagicMock()
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(
return_value={}
)
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
proxy_logging_obj.during_call_hook = AsyncMock(return_value=None)
@ -2222,7 +2296,9 @@ 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
@ -2467,7 +2543,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."""
@ -2475,7 +2554,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."""
@ -2553,7 +2634,9 @@ class TestMCPServerManagerExpandPermissionList:
def test_expands_server_name(self):
manager = MCPServerManager()
manager.config_mcp_servers["id-usw1"] = self._make_server("id-usw1", server_name="a")
manager.config_mcp_servers["id-usw1"] = self._make_server(
"id-usw1", server_name="a"
)
assert manager.expand_permission_list(["a"]) == ["id-usw1"]
@ -2577,7 +2660,9 @@ class TestMCPServerManagerExpandPermissionList:
def test_name_collision_expands_to_all_matches(self):
"""Two servers sharing a server_name both resolve — the documented behavior."""
manager = MCPServerManager()
manager.config_mcp_servers["id-config"] = self._make_server("id-config", server_name="shared")
manager.config_mcp_servers["id-config"] = self._make_server(
"id-config", server_name="shared"
)
manager.registry["id-db"] = self._make_server("id-db", server_name="shared")
assert sorted(manager.expand_permission_list(["shared"])) == [
@ -2587,7 +2672,9 @@ class TestMCPServerManagerExpandPermissionList:
def test_searches_config_and_registry_union(self):
manager = MCPServerManager()
manager.config_mcp_servers["cfg-id"] = self._make_server("cfg-id", server_name="a")
manager.config_mcp_servers["cfg-id"] = self._make_server(
"cfg-id", server_name="a"
)
manager.registry["reg-id"] = self._make_server("reg-id", server_name="b")
assert manager.expand_permission_list(["a"]) == ["cfg-id"]
@ -2599,15 +2686,23 @@ class TestMCPServerManagerExpandPermissionList:
servers whose server_name happens to equal that id.
"""
manager = MCPServerManager()
manager.config_mcp_servers["id-1"] = self._make_server("id-1", server_name="other_name")
manager.config_mcp_servers["id-2"] = self._make_server("id-2", server_name="id-1")
manager.config_mcp_servers["id-1"] = self._make_server(
"id-1", server_name="other_name"
)
manager.config_mcp_servers["id-2"] = self._make_server(
"id-2", server_name="id-1"
)
assert manager.expand_permission_list(["id-1"]) == ["id-1"]
def test_mixed_ids_and_names_in_same_list(self):
manager = MCPServerManager()
manager.config_mcp_servers["uuid-1"] = self._make_server("uuid-1", server_name="a")
manager.config_mcp_servers["uuid-2"] = self._make_server("uuid-2", server_name="b")
manager.config_mcp_servers["uuid-1"] = self._make_server(
"uuid-1", server_name="a"
)
manager.config_mcp_servers["uuid-2"] = self._make_server(
"uuid-2", server_name="b"
)
# ["uuid-1", "b"] -> uuid-1 passes through, "b" resolves to uuid-2
assert sorted(manager.expand_permission_list(["uuid-1", "b"])) == [
@ -2618,7 +2713,9 @@ class TestMCPServerManagerExpandPermissionList:
def test_deduplicates_overlapping_id_and_name_entries(self):
"""If a list references the same server by both id and name, return it once."""
manager = MCPServerManager()
manager.config_mcp_servers["uuid-1"] = self._make_server("uuid-1", server_name="a")
manager.config_mcp_servers["uuid-1"] = self._make_server(
"uuid-1", server_name="a"
)
assert manager.expand_permission_list(["uuid-1", "a"]) == ["uuid-1"]
@ -2628,10 +2725,14 @@ class TestMCPServerManagerExpandPermissionList:
the cross-region portability the customer is asking for.
"""
usw1 = MCPServerManager()
usw1.config_mcp_servers["hash-usw1"] = self._make_server("hash-usw1", server_name="a")
usw1.config_mcp_servers["hash-usw1"] = self._make_server(
"hash-usw1", server_name="a"
)
usc1 = MCPServerManager()
usc1.config_mcp_servers["hash-usc1"] = self._make_server("hash-usc1", server_name="a")
usc1.config_mcp_servers["hash-usc1"] = self._make_server(
"hash-usc1", server_name="a"
)
assert usw1.expand_permission_list(["a"]) == ["hash-usw1"]
assert usc1.expand_permission_list(["a"]) == ["hash-usc1"]
@ -2660,14 +2761,18 @@ class TestMCPServerManagerExpandToolPermissions:
concrete server_id, otherwise `.get(server_id)` misses and the tool
restriction is silently dropped (caller treats None as allow-all)."""
manager = MCPServerManager()
manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="my-alias")
manager.config_mcp_servers["uuid-a"] = self._make_server(
"uuid-a", server_name="my-alias"
)
result = manager.expand_tool_permissions({"my-alias": ["read_file"]})
assert result == {"uuid-a": ["read_file"]}
def test_passes_through_existing_server_id_key(self):
manager = MCPServerManager()
manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alpha")
manager.config_mcp_servers["uuid-a"] = self._make_server(
"uuid-a", server_name="alpha"
)
result = manager.expand_tool_permissions({"uuid-a": ["read_file"]})
assert result == {"uuid-a": ["read_file"]}
@ -2686,7 +2791,9 @@ class TestMCPServerManagerExpandToolPermissions:
"""Two servers sharing a server_name both match; their tool lists get
the restriction (matches the list-expansion collision semantics)."""
manager = MCPServerManager()
manager.config_mcp_servers["uuid-1"] = self._make_server("uuid-1", server_name="shared")
manager.config_mcp_servers["uuid-1"] = self._make_server(
"uuid-1", server_name="shared"
)
manager.registry["uuid-2"] = self._make_server("uuid-2", server_name="shared")
result = manager.expand_tool_permissions({"shared": ["read_file"]})
@ -2699,9 +2806,13 @@ class TestMCPServerManagerExpandToolPermissions:
both refer to the same server, the tool lists are unioned rather
than one overwriting the other."""
manager = MCPServerManager()
manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alias-a")
manager.config_mcp_servers["uuid-a"] = self._make_server(
"uuid-a", server_name="alias-a"
)
result = manager.expand_tool_permissions({"uuid-a": ["read_file"], "alias-a": ["write_file"]})
result = manager.expand_tool_permissions(
{"uuid-a": ["read_file"], "alias-a": ["write_file"]}
)
assert sorted(result["uuid-a"]) == ["read_file", "write_file"]
@pytest.mark.asyncio
@ -2740,7 +2851,9 @@ class TestMCPServerManagerExpandToolPermissions:
"litellm.proxy._experimental.mcp_server.elicitation_handler.handle_elicitation_request",
new_callable=AsyncMock,
) as mock_handle,
patch("litellm.proxy._experimental.mcp_server.server.get_active_mcp_session") as mock_get_session,
patch(
"litellm.proxy._experimental.mcp_server.server.get_active_mcp_session"
) as mock_get_session,
):
mock_session = MagicMock()
mock_session.capabilities = "test_capabilities"