fix(mcp): separate MCP and non-MCP tools in semantic filter hook

The semantic filter treated all tools equally — it didn't distinguish
between MCP tools (which it should filter) and non-MCP built-in tools
(which should pass through untouched). This caused two failures:

- On match: non-MCP tools were dropped, losing built-in functionality
- On zero match: all tools returned unfiltered, exceeding provider limits

Move MCP/non-MCP separation into the hook. The hook now passes only
MCP tools (identified by server-name prefix) to filter_tools() and
recombines with non-MCP tools after. filter_tools() returns empty
list on zero matches instead of all tools.
This commit is contained in:
Lance Hsu 2026-04-02 19:39:21 +08:00
parent f25ce7437c
commit ac86e054ab
No known key found for this signature in database
GPG key ID: CFF816BB8560B8B8
3 changed files with 115 additions and 38 deletions

View file

@ -6,7 +6,6 @@ Filters MCP tools semantically for /chat/completions and /responses endpoints.
from typing import TYPE_CHECKING, Any, Dict, List, Optional
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.utils import is_tool_name_prefixed
if TYPE_CHECKING:
from semantic_router.routers import SemanticRouter
@ -191,12 +190,9 @@ class SemanticMCPToolFilter:
matched_tool_names = self._extract_tool_names_from_matches(matches)
if not matched_tool_names:
# No semantic matches — drop MCP tools (prefixed) and keep only
# non-MCP tools to avoid exceeding provider tool limits.
return [
t for t in available_tools
if not is_tool_name_prefixed(self._extract_tool_info(t)[0])
]
# No semantic matches — return empty so only non-MCP tools
# (added back by the hook) reach the LLM.
return []
return self._get_tools_by_names(matched_tool_names, available_tools)

View file

@ -7,6 +7,7 @@ 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.constants import (
DEFAULT_MCP_SEMANTIC_FILTER_EMBEDDING_MODEL,
DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD,
@ -222,12 +223,25 @@ class SemanticToolFilterHook(CustomLogger):
f"with query: '{user_query[:50]}...'"
)
# Filter tools semantically
filtered_tools = await self.filter.filter_tools(
# Separate MCP tools (prefixed) from non-MCP tools — only filter
# MCP tools, always pass non-MCP tools through untouched.
def _tool_name(t):
return 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))]
if not mcp_tools:
return None
# Filter only MCP tools semantically
filtered_mcp_tools = await self.filter.filter_tools(
query=user_query,
available_tools=tools, # type: ignore
available_tools=mcp_tools, # type: ignore
)
filtered_tools = filtered_mcp_tools + non_mcp_tools
# Always update tools and emit header (even if count unchanged)
data["tools"] = filtered_tools

View file

@ -305,16 +305,16 @@ async def test_semantic_filter_hook_triggers_on_completion():
# Prepare data - completion request with tools
tools = [
MCPTool(name=f"tool_{i}", description=f"Tool {i}", inputSchema={"type": "object"})
MCPTool(name=f"server-tool_{i}", description=f"Tool {i}", inputSchema={"type": "object"})
for i in range(10)
]
# Build router with the tools before filtering
filter_instance._build_router(tools)
# Create hook
hook = SemanticToolFilterHook(filter_instance)
data = {
"model": "gpt-4",
"messages": [
@ -323,11 +323,11 @@ async def test_semantic_filter_hook_triggers_on_completion():
"tools": tools,
"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()
# Call hook
result = await hook.async_pre_call_hook(
user_api_key_dict=mock_user_api_key_dict,
@ -335,12 +335,12 @@ async def test_semantic_filter_hook_triggers_on_completion():
data=data,
call_type="completion",
)
# 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")
@ -394,13 +394,12 @@ async def test_semantic_filter_hook_skips_no_tools():
@pytest.mark.asyncio
async def test_semantic_filter_zero_matches_returns_only_non_mcp_tools():
async def test_semantic_filter_zero_matches_returns_empty():
"""
Test that when zero semantic matches are found, only non-MCP tools are returned.
Test that filter_tools returns an empty list when zero semantic matches are found.
MCP tools have server-name prefixes (e.g., "weather-get_forecast").
Non-MCP tools don't (e.g., "web_search"). On zero matches, the filter
should drop all MCP tools to avoid exceeding the 128 tool limit.
The hook is responsible for adding non-MCP tools back — filter_tools only
operates on MCP tools passed to it.
"""
from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
SemanticMCPToolFilter,
@ -418,27 +417,95 @@ async def test_semantic_filter_zero_matches_returns_only_non_mcp_tools():
# Mock the semantic router to return no matches
filter_instance.tool_router = Mock(return_value=[])
# Mix of MCP tools (prefixed) and non-MCP tools (no prefix)
available_tools = [
mcp_tools = [
MCPTool(name="weather-get_forecast", description="Get forecast", inputSchema={"type": "object"}),
MCPTool(name="email-send_email", description="Send email", inputSchema={"type": "object"}),
MCPTool(name="docs-read_document", description="Read doc", inputSchema={"type": "object"}),
MCPTool(name="web_search", description="Search the web", inputSchema={"type": "object"}),
MCPTool(name="code_interpreter", description="Run code", inputSchema={"type": "object"}),
]
filtered = await filter_instance.filter_tools(
query="hello",
available_tools=available_tools,
available_tools=mcp_tools,
)
# Should return only non-MCP tools
filtered_names = [t.name for t in filtered]
assert "web_search" in filtered_names
assert "code_interpreter" in filtered_names
assert len(filtered) == 2, f"Expected 2 non-MCP tools, got {len(filtered)}: {filtered_names}"
# MCP tools should be excluded
assert "weather-get_forecast" not in filtered_names
assert "email-send_email" not in filtered_names
assert "docs-read_document" not in filtered_names
assert len(filtered) == 0, f"Expected empty list on zero matches, got {len(filtered)}"
@pytest.mark.asyncio
async def test_hook_preserves_non_mcp_tools():
"""
Test that the hook passes non-MCP tools through untouched and only
filters MCP tools semantically.
"""
from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
SemanticMCPToolFilter,
)
from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
from litellm.types.utils import Embedding, EmbeddingResponse
mock_router = Mock()
def mock_embedding_sync(*args, **kwargs):
return EmbeddingResponse(
data=[Embedding(embedding=[0.1] * 1536, index=0, object="embedding")],
model="text-embedding-3-small",
object="list",
usage={"prompt_tokens": 10, "total_tokens": 10}
)
async def mock_embedding_async(*args, **kwargs):
return mock_embedding_sync()
mock_router.embedding = mock_embedding_sync
mock_router.aembedding = mock_embedding_async
filter_instance = SemanticMCPToolFilter(
embedding_model="text-embedding-3-small",
litellm_router_instance=mock_router,
top_k=3,
similarity_threshold=0.3,
enabled=True,
)
# MCP tools (prefixed) — these get filtered
mcp_tools = [
MCPTool(name=f"server-tool_{i}", description=f"MCP tool {i}", inputSchema={"type": "object"})
for i in range(10)
]
# Build router with MCP tools
filter_instance._build_router(mcp_tools)
hook = SemanticToolFilterHook(filter_instance)
# Request has both MCP and non-MCP tools
all_tools = mcp_tools + [
MCPTool(name="web_search", description="Search the web", inputSchema={"type": "object"}),
MCPTool(name="code_interpreter", description="Run code", inputSchema={"type": "object"}),
]
data = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Use MCP tool 1"}],
"tools": all_tools,
"metadata": {},
}
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}"