mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
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:
parent
24ae278990
commit
3cac5feab5
2 changed files with 202 additions and 44 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue