fix(proxy): Correctly parse multi-part MCP server aliases from URL paths

This commit is contained in:
iabhi4 2025-09-14 15:04:39 -07:00
parent 03f2be1e20
commit 4ba3a21042
2 changed files with 80 additions and 1 deletions

View file

@ -578,7 +578,7 @@ if MCP_AVAILABLE:
"""
import re
mcp_servers_from_path: Optional[List[str]] = None
mcp_path_match = re.match(r"^/mcp/([^/]+)(/.*)?$", path)
mcp_path_match = re.match(r"^/mcp/([^/]+/[^/]+|[^/]+)(/.*)?$", path)
if mcp_path_match:
mcp_servers_str = mcp_path_match.group(1)
if mcp_servers_str:

View file

@ -342,3 +342,82 @@ async def test_concurrent_initialize_session_managers():
mcp_server._SESSION_MANAGERS_INITIALIZED = original_initialized
mcp_server._session_manager_cm = original_session_cm
mcp_server._sse_session_manager_cm = original_sse_session_cm
@pytest.mark.asyncio
async def test_mcp_routing_with_conflicting_alias_and_group_name():
"""
Tests (GH #14536) where an MCP server alias (e.g., "group/id")
conflicts with an access group name (e.g., "group").
"""
try:
from litellm.proxy._experimental.mcp_server.server import (
_get_mcp_servers_in_path,
_get_tools_from_mcp_servers,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.types.mcp_server.mcp_server_manager import MCPServer
from litellm.proxy._types import MCPTransport, MCPSpecVersion
except ImportError:
pytest.skip("MCP server not available")
global_mcp_server_manager.registry.clear()
# Create two in-memory servers
specific_server = MCPServer(
server_id="specific_server_id",
name="custom_solutions/user_123",
alias="custom_solutions/user_123",
transport=MCPTransport.http,
spec_version=MCPSpecVersion.jun_2025,
)
other_server = MCPServer(
server_id="other_server_in_group_id",
name="custom_solutions/another_user_456",
alias="custom_solutions/another_user_456",
transport=MCPTransport.http,
spec_version=MCPSpecVersion.jun_2025,
)
global_mcp_server_manager.registry[specific_server.server_id] = specific_server
global_mcp_server_manager.registry[other_server.server_id] = other_server
user_key = UserAPIKeyAuth(api_key="sk-test", team_id="team_custom_solutions")
# Define the request path that triggers the bug
test_path = "/mcp/custom_solutions/user_123/chat/completions"
# This mock will be our "spy" to see which servers are ultimately contacted
mock_get_tools_spy = AsyncMock(return_value=[])
# Mock the function that checks DB for an access group named "custom_solutions"
mock_db_lookup = AsyncMock(return_value=[specific_server.server_id, other_server.server_id])
mock_get_allowed = AsyncMock(return_value=[specific_server.server_id, other_server.server_id])
with patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_allowed_mcp_servers",
mock_get_allowed,
), patch(
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler._get_mcp_servers_from_access_groups",
mock_db_lookup,
), patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server",
mock_get_tools_spy,
):
mcp_servers_from_path = _get_mcp_servers_in_path(test_path)
await _get_tools_from_mcp_servers(
user_api_key_auth=user_key,
mcp_servers=mcp_servers_from_path,
mcp_auth_header=None,
)
# Get the list of actual server objects that the orchestrator tried to contact
called_servers = [call.kwargs["server"] for call in mock_get_tools_spy.call_args_list]
assert len(called_servers) == 1, "Should have resolved to exactly one server."
assert (
called_servers[0].server_id == specific_server.server_id
), "Should have contacted the specific server alias, not the group."