mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
f25ce7437c
commit
ac86e054ab
3 changed files with 115 additions and 38 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue