Check if MCP server name exists after splitting {mcpServerName-toolName}

This commit is contained in:
Aaron Chen 2025-11-26 13:42:03 +11:00
parent 658445b93f
commit fe6258a4c9
No known key found for this signature in database
GPG key ID: E695E2915D55C2A9
2 changed files with 82 additions and 0 deletions

View file

@ -1209,6 +1209,19 @@ if MCP_AVAILABLE:
# Remove prefix from tool name for logging and processing
original_tool_name, server_name = split_server_prefix_from_name(name)
# Validate that extracted server_name exists in allowed servers
# If not, the tool name likely contains the separator itself
if server_name and not any(
server_name in [s for s in [server.name, server.alias] if s]
for server in allowed_mcp_servers
):
verbose_logger.debug(
f"Server '{server_name}' from tool '{name}' not found in allowed servers. "
f"Treating as unprefixed tool."
)
original_tool_name = name
server_name = ""
# If tool name is unprefixed, resolve its server so we can enforce permissions
if not server_name:
mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name)

View file

@ -1544,3 +1544,72 @@ def test_filter_tools_by_allowed_tools():
assert len(filtered_tools) == 2
assert filtered_tools[0].name == "my_api_mcp-getpetbyid"
assert filtered_tools[1].name == "my_api_mcp-findpetsbystatus"
@pytest.mark.asyncio
async def test_call_mcp_tool_validates_extracted_server_name():
"""
Test that call_mcp_tool validates extracted server names against allowed servers.
If a tool name contains the separator but the extracted server doesn't exist,
it should be treated as an unprefixed tool name.
"""
try:
from litellm.proxy._experimental.mcp_server.server import (
call_mcp_tool,
global_mcp_server_manager,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
except ImportError:
pytest.skip("MCP server not available")
# Create a real allowed server
allowed_server = MCPServer(
server_id="server-123",
name="backstage",
alias="backstage",
server_name="backstage",
url="https://backstage.com/mcp",
transport=MCPTransport.http,
mcp_info={"server_name": "backstage"},
)
with patch.object(
global_mcp_server_manager,
"get_allowed_mcp_servers",
new_callable=AsyncMock,
) as mock_get_allowed, patch.object(
global_mcp_server_manager,
"get_mcp_server_by_id",
return_value=allowed_server,
), patch.object(
global_mcp_server_manager,
"_get_mcp_server_from_tool_name",
return_value=allowed_server,
) as mock_get_server, patch(
"litellm.proxy._experimental.mcp_server.server.global_mcp_tool_registry"
) as mock_tool_registry, patch(
"litellm.proxy._experimental.mcp_server.server._handle_managed_mcp_tool",
new_callable=AsyncMock,
) as mock_handle_managed, patch(
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed",
return_value=True,
):
mock_get_allowed.return_value = [allowed_server.server_id]
mock_tool_registry.get_tool.return_value = None
# Call with a tool name that contains separator but "get" is not a valid server
# This should be treated as unprefixed tool "get--catalog--entity"
result = await call_mcp_tool(
name="get-catalog-entity",
arguments={"id": "123"},
mcp_servers=["backstage"],
)
mock_get_allowed.assert_awaited_once()
# Should resolve server from tool name since extracted "get" is not valid
assert mock_get_server.call_count >= 1
# First call should use the full tool name (not split) since "get" is not a valid server
assert mock_get_server.call_args_list[0][0][0] == "get-catalog-entity"