diff --git a/litellm/proxy/hooks/mcp_semantic_filter/hook.py b/litellm/proxy/hooks/mcp_semantic_filter/hook.py index 96689f279a9..5b330606a08 100644 --- a/litellm/proxy/hooks/mcp_semantic_filter/hook.py +++ b/litellm/proxy/hooks/mcp_semantic_filter/hook.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py index 5612404e74a..0bf461f3ff5 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py @@ -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