fix(mcp): validate tool prefix against registry instead of heuristic

is_tool_name_prefixed() just checks if "-" exists in the name, which
misclassifies non-MCP tools with hyphens (e.g., text-to-speech). Replace
with _is_mcp_tool() that splits on the first "-" and checks if the
prefix is a known registered MCP server name from the registry.
This commit is contained in:
Lance Hsu 2026-04-04 00:49:52 +08:00
parent 24ae278990
commit 3cac5feab5
No known key found for this signature in database
GPG key ID: CFF816BB8560B8B8
2 changed files with 202 additions and 44 deletions

View file

@ -8,7 +8,11 @@ Reduces context window size and improves tool selection accuracy.
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from litellm._logging import verbose_proxy_logger
from litellm.proxy._experimental.mcp_server.utils import is_tool_name_prefixed
from litellm.proxy._experimental.mcp_server.utils import (
MCP_TOOL_PREFIX_SEPARATOR,
get_server_prefix,
normalize_server_name,
)
from litellm.constants import (
DEFAULT_MCP_SEMANTIC_FILTER_EMBEDDING_MODEL,
DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD,
@ -44,12 +48,35 @@ class SemanticToolFilterHook(CustomLogger):
"""
super().__init__()
self.filter = semantic_filter
self._registered_server_prefixes: Optional[set] = None
verbose_proxy_logger.debug(
f"Initialized SemanticToolFilterHook with filter: "
f"enabled={semantic_filter.enabled}, top_k={semantic_filter.top_k}"
)
def _get_registered_server_prefixes(self) -> set:
"""Get the set of known MCP server prefixes from the registry."""
if self._registered_server_prefixes is None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
registry = global_mcp_server_manager.get_registry()
self._registered_server_prefixes = {
normalize_server_name(get_server_prefix(server))
for server in registry.values()
if get_server_prefix(server)
}
return self._registered_server_prefixes
def _is_mcp_tool(self, tool_name: str) -> bool:
"""Check if a tool is an MCP tool by validating its prefix against the registry."""
if MCP_TOOL_PREFIX_SEPARATOR not in tool_name:
return False
prefix = tool_name.split(MCP_TOOL_PREFIX_SEPARATOR, 1)[0]
return normalize_server_name(prefix) in self._get_registered_server_prefixes()
def _should_expand_mcp_tools(self, tools: List[Any]) -> bool:
"""
Check if tools contain MCP references with server_url="litellm_proxy".
@ -231,10 +258,8 @@ class SemanticToolFilterHook(CustomLogger):
t.get("name", "") if isinstance(t, dict) else getattr(t, "name", "")
)
mcp_tools = [t for t in tools if is_tool_name_prefixed(_tool_name(t))]
non_mcp_tools = [
t for t in tools if not is_tool_name_prefixed(_tool_name(t))
]
mcp_tools = [t for t in tools if self._is_mcp_tool(_tool_name(t))]
non_mcp_tools = [t for t in tools if not self._is_mcp_tool(_tool_name(t))]
if not mcp_tools:
return None

View file

@ -324,24 +324,34 @@ async def test_semantic_filter_hook_triggers_on_completion():
"metadata": {}, # Hook needs metadata field to store filter stats
}
# Mock user API key dict and cache
mock_user_api_key_dict = Mock()
mock_cache = Mock()
# Mock registry so "server" prefix is recognized as MCP
mock_server = Mock()
mock_server.alias = "server"
mock_server.server_name = "server"
mock_server.server_id = "server-id"
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_manager:
mock_manager.get_registry.return_value = {"server-id": mock_server}
# Call hook
result = await hook.async_pre_call_hook(
user_api_key_dict=mock_user_api_key_dict,
cache=mock_cache,
data=data,
call_type="completion",
)
# Mock user API key dict and cache
mock_user_api_key_dict = Mock()
mock_cache = Mock()
# Assertions
assert result is not None, "Hook should return modified data"
assert "tools" in result, "Result should contain tools"
assert len(result["tools"]) < len(tools), f"Hook should filter tools, got {len(result['tools'])}/{len(tools)}"
# Call hook
result = await hook.async_pre_call_hook(
user_api_key_dict=mock_user_api_key_dict,
cache=mock_cache,
data=data,
call_type="completion",
)
print(f"✅ Hook triggered correctly: {len(tools)} -> {len(result['tools'])} tools")
# Assertions
assert result is not None, "Hook should return modified data"
assert "tools" in result, "Result should contain tools"
assert len(result["tools"]) < len(tools), f"Hook should filter tools, got {len(result['tools'])}/{len(tools)}"
print(f"✅ Hook triggered correctly: {len(tools)} -> {len(result['tools'])} tools")
@ -472,16 +482,25 @@ async def test_hook_falls_back_to_top_k_when_only_mcp_and_zero_matches():
"metadata": {},
}
result = await hook.async_pre_call_hook(
user_api_key_dict=Mock(),
cache=Mock(),
data=data,
call_type="completion",
)
mock_server = Mock()
mock_server.alias = "server"
mock_server.server_name = "server"
mock_server.server_id = "server-id"
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_manager:
mock_manager.get_registry.return_value = {"server-id": mock_server}
assert result is not None
# Should fall back to first top_k (3) MCP tools, not empty
assert len(result["tools"]) == 3, f"Expected 3 tools (top_k), got {len(result['tools'])}"
result = await hook.async_pre_call_hook(
user_api_key_dict=Mock(),
cache=Mock(),
data=data,
call_type="completion",
)
assert result is not None
# Should fall back to first top_k (3) MCP tools, not empty
assert len(result["tools"]) == 3, f"Expected 3 tools (top_k), got {len(result['tools'])}"
@pytest.mark.asyncio
@ -581,22 +600,136 @@ async def test_hook_preserves_non_mcp_tools():
"metadata": {},
}
mock_user_api_key_dict = Mock()
mock_cache = Mock()
mock_server = Mock()
mock_server.alias = "server"
mock_server.server_name = "server"
mock_server.server_id = "server-id"
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_manager:
mock_manager.get_registry.return_value = {"server-id": mock_server}
result = await hook.async_pre_call_hook(
user_api_key_dict=mock_user_api_key_dict,
cache=mock_cache,
data=data,
call_type="completion",
mock_user_api_key_dict = Mock()
mock_cache = Mock()
result = await hook.async_pre_call_hook(
user_api_key_dict=mock_user_api_key_dict,
cache=mock_cache,
data=data,
call_type="completion",
)
assert result is not None
filtered_names = [t.name if hasattr(t, 'name') else t.get('name', '') for t in result["tools"]]
# Non-MCP tools should always be present
assert "web_search" in filtered_names, f"web_search missing from {filtered_names}"
assert "code_interpreter" in filtered_names, f"code_interpreter missing from {filtered_names}"
# MCP tools should be filtered (fewer than 10)
mcp_count = sum(1 for n in filtered_names if "-" in n)
assert mcp_count <= 3, f"Expected at most 3 MCP tools (top_k=3), got {mcp_count}"
@pytest.mark.asyncio
async def test_hook_does_not_filter_hyphenated_non_mcp_tools():
"""
Test that non-MCP tools with hyphens (e.g., text-to-speech) are NOT
misclassified as MCP tools when their prefix doesn't match a registered server.
"""
from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
SemanticMCPToolFilter,
)
from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
mock_router = Mock()
filter_instance = SemanticMCPToolFilter(
embedding_model="text-embedding-3-small",
litellm_router_instance=mock_router,
top_k=3,
similarity_threshold=0.3,
enabled=True,
)
assert result is not None
filtered_names = [t.name if hasattr(t, 'name') else t.get('name', '') for t in result["tools"]]
# Non-MCP tools should always be present
assert "web_search" in filtered_names, f"web_search missing from {filtered_names}"
assert "code_interpreter" in filtered_names, f"code_interpreter missing from {filtered_names}"
# MCP tools should be filtered (fewer than 10)
mcp_count = sum(1 for n in filtered_names if "-" in n)
assert mcp_count <= 3, f"Expected at most 3 MCP tools (top_k=3), got {mcp_count}"
# Mock the semantic router to return no matches
filter_instance.tool_router = Mock(return_value=[])
hook = SemanticToolFilterHook(filter_instance)
# Mock registry with only "weather" as a registered server
mock_server = Mock()
mock_server.alias = "weather"
mock_server.server_name = "weather"
mock_server.server_id = "weather-id"
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_manager:
mock_manager.get_registry.return_value = {"weather-id": mock_server}
tools = [
MCPTool(name="weather-get_forecast", description="Get forecast", inputSchema={"type": "object"}),
MCPTool(name="text-to-speech", description="Convert text to speech", inputSchema={"type": "object"}),
MCPTool(name="code-review", description="Review code", inputSchema={"type": "object"}),
MCPTool(name="web_search", description="Search the web", inputSchema={"type": "object"}),
]
data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "hello"}],
"tools": tools,
"metadata": {},
}
result = await hook.async_pre_call_hook(
user_api_key_dict=Mock(),
cache=Mock(),
data=data,
call_type="completion",
)
assert result is not None
filtered_names = [
t.name if hasattr(t, "name") else t.get("name", "")
for t in result["tools"]
]
# Hyphenated non-MCP tools should pass through (not filtered)
assert "text-to-speech" in filtered_names, f"text-to-speech should not be filtered: {filtered_names}"
assert "code-review" in filtered_names, f"code-review should not be filtered: {filtered_names}"
# Non-hyphenated non-MCP tool also passes through
assert "web_search" in filtered_names, f"web_search should not be filtered: {filtered_names}"
def test_is_mcp_tool_with_registered_prefix():
"""Test _is_mcp_tool correctly identifies MCP tools by checking against registry."""
from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
SemanticMCPToolFilter,
)
from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
mock_router = Mock()
filter_instance = SemanticMCPToolFilter(
embedding_model="text-embedding-3-small",
litellm_router_instance=mock_router,
top_k=3,
similarity_threshold=0.3,
enabled=True,
)
hook = SemanticToolFilterHook(filter_instance)
# Mock registry
mock_server = Mock()
mock_server.alias = "weather"
mock_server.server_name = "weather"
mock_server.server_id = "weather-id"
with patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_manager:
mock_manager.get_registry.return_value = {"weather-id": mock_server}
# MCP tool with registered prefix
assert hook._is_mcp_tool("weather-get_forecast") is True
# Non-MCP tool with hyphen but unregistered prefix
assert hook._is_mcp_tool("text-to-speech") is False
assert hook._is_mcp_tool("code-review") is False
# Non-MCP tool without hyphen
assert hook._is_mcp_tool("web_search") is False