From 631b41481c8ce3b2108cc7ba928255a0e267d2f6 Mon Sep 17 00:00:00 2001 From: Genmin Date: Fri, 1 May 2026 07:34:21 -0700 Subject: [PATCH] fix: keep native tools out of MCP filtering --- litellm/proxy/_lazy_openapi_snapshot.py | 2 + .../proxy/hooks/mcp_semantic_filter/hook.py | 43 +++++-- .../mcp_server/test_semantic_tool_filter.py | 110 ++++++++++++++---- 3 files changed, 126 insertions(+), 29 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.py b/litellm/proxy/_lazy_openapi_snapshot.py index a9d5396cc29..c63ff8d0733 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.py +++ b/litellm/proxy/_lazy_openapi_snapshot.py @@ -86,6 +86,7 @@ def generate_snapshot() -> Dict[str, Dict]: import importlib from fastapi.openapi.utils import get_openapi + from litellm.proxy._lazy_features import LAZY_FEATURES from litellm.proxy.proxy_server import app, ensure_unique_openapi_operation_ids @@ -132,6 +133,7 @@ def generate_snapshot() -> Dict[str, Dict]: } return fragments + if __name__ == "__main__": fragments = generate_snapshot() SNAPSHOT_FILE.write_text(json.dumps(fragments, indent=2, sort_keys=True) + "\n") diff --git a/litellm/proxy/hooks/mcp_semantic_filter/hook.py b/litellm/proxy/hooks/mcp_semantic_filter/hook.py index 645e0eb84aa..598f89f795d 100644 --- a/litellm/proxy/hooks/mcp_semantic_filter/hook.py +++ b/litellm/proxy/hooks/mcp_semantic_filter/hook.py @@ -146,23 +146,50 @@ class SemanticToolFilterHook(CustomLogger): name = getattr(tool, "name", "") return name if isinstance(name, str) else "" - def _is_mcp_router_tool(self, tool: Any) -> bool: + @staticmethod + def _is_native_function_tool(tool: Any) -> bool: + if not isinstance(tool, dict): + return False + + if isinstance(tool.get("function"), dict): + return True + + return tool.get("type") == "function" and isinstance(tool.get("name"), str) + + @staticmethod + def _candidate_canonical_names(tool_name: str) -> List[str]: + return [ + tool_name[index + 1 :] + for index, char in enumerate(tool_name) + if char in {"_", "-"} and index + 1 < len(tool_name) + ] + + def _get_matching_mcp_canonical_name(self, tool: Any) -> Optional[str]: + if self._is_native_function_tool(tool): + return None + tool_name = self._get_tool_name(tool) if not tool_name: - return False + return None tool_map = getattr(self.filter, "_tool_map", {}) if tool_name in tool_map: - return True + return tool_name name_matches_canonical = getattr(self.filter, "_name_matches_canonical", None) if name_matches_canonical is None: - return False + return None - return any( - name_matches_canonical(tool_name, canonical_name) - for canonical_name in tool_map - ) + for canonical_name in self._candidate_canonical_names(tool_name): + if canonical_name in tool_map and name_matches_canonical( + tool_name, canonical_name + ): + return canonical_name + + return None + + def _is_mcp_router_tool(self, tool: Any) -> bool: + return self._get_matching_mcp_canonical_name(tool) is not None def _partition_mcp_router_tools( self, tools: List[Any] 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 49dda138977..b5aa6486757 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 @@ -200,6 +200,7 @@ async def test_semantic_filter_top_k_limiting(): # Should return at most 5 tools assert len(filtered) <= 5, f"Expected at most 5 tools, got {len(filtered)}" + @pytest.mark.asyncio async def test_semantic_filter_disabled(): """ @@ -558,6 +559,87 @@ async def test_semantic_filter_hook_preserves_native_openai_tools(): ) +@pytest.mark.asyncio +async def test_semantic_filter_hook_keeps_native_tool_name_collisions_native(): + """ + Responses API native tools can carry a top-level name. If that name happens + to equal an MCP canonical name, the native tool must not be filtered as MCP. + """ + from litellm.proxy._experimental.mcp_server.semantic_tool_filter import ( + SemanticMCPToolFilter, + ) + from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook + + filter_instance = SemanticMCPToolFilter( + embedding_model="text-embedding-3-small", + litellm_router_instance=Mock(), + top_k=3, + similarity_threshold=0.3, + enabled=True, + ) + mcp_tool = {"name": "github-search", "description": "Search GitHub"} + filter_instance._tool_map = {mcp_tool["name"]: mcp_tool} + filter_instance.filter_tools = AsyncMock(return_value=[]) # type: ignore[method-assign] + + native_tool = { + "type": "function", + "name": "github-search", + "description": "Caller-owned search tool", + "parameters": {"type": "object"}, + } + hook = SemanticToolFilterHook(filter_instance) + data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Search GitHub"}], + "tools": [native_tool], + "metadata": {}, + } + + result = await hook.async_pre_call_hook( + user_api_key_dict=Mock(), + cache=Mock(), + data=data, + call_type="completion", + ) + + assert result is None + filter_instance.filter_tools.assert_not_awaited() + assert data["tools"] == [native_tool] + + +def test_semantic_filter_hook_prefixed_match_does_not_scan_whole_tool_map(): + """ + Matching a client-prefixed MCP tool should check candidate suffixes from the + incoming name instead of scanning every registered canonical tool. + """ + from litellm.proxy._experimental.mcp_server.semantic_tool_filter import ( + SemanticMCPToolFilter, + ) + from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook + + filter_instance = SemanticMCPToolFilter( + embedding_model="text-embedding-3-small", + litellm_router_instance=Mock(), + top_k=3, + similarity_threshold=0.3, + enabled=True, + ) + filter_instance._tool_map = { + **{f"server-{index}": {"name": f"server-{index}"} for index in range(1000)}, + "github-search": {"name": "github-search"}, + } + filter_instance._name_matches_canonical = Mock( # type: ignore[method-assign] + wraps=filter_instance._name_matches_canonical + ) + + hook = SemanticToolFilterHook(filter_instance) + + assert hook._is_mcp_router_tool({"name": "client_github-search"}) + filter_instance._name_matches_canonical.assert_called_once_with( + "client_github-search", "github-search" + ) + + @pytest.mark.asyncio async def test_semantic_filter_hook_preserves_remaining_tool_order(): """ @@ -704,9 +786,7 @@ class TestGetToolsByNames: {"name": "send_email", "description": "send mail"}, ] - matched = filter_instance._get_tools_by_names( - ["send_email"], available_tools - ) + matched = filter_instance._get_tools_by_names(["send_email"], available_tools) assert len(matched) == 1 assert matched[0]["name"] == "send_email" @@ -718,9 +798,7 @@ class TestGetToolsByNames: client_name = "litellm_" + canonical available_tools = [{"name": client_name, "description": "scrape"}] - matched = filter_instance._get_tools_by_names( - [canonical], available_tools - ) + matched = filter_instance._get_tools_by_names([canonical], available_tools) assert len(matched) == 1 # Must return the incoming tool unchanged so the client-facing @@ -731,13 +809,9 @@ class TestGetToolsByNames: """Some clients use dash as alias separator; accept that too.""" filter_instance = self._make_filter() canonical = "weather_svc-get_weather" - available_tools = [ - {"name": "mcp-" + canonical, "description": "weather"} - ] + available_tools = [{"name": "mcp-" + canonical, "description": "weather"}] - matched = filter_instance._get_tools_by_names( - [canonical], available_tools - ) + matched = filter_instance._get_tools_by_names([canonical], available_tools) assert len(matched) == 1 assert matched[0]["name"] == "mcp-" + canonical @@ -767,9 +841,7 @@ class TestGetToolsByNames: {"name": "litellm_" + canonical, "description": "wrapped"}, ] - matched = filter_instance._get_tools_by_names( - [canonical], available_tools - ) + matched = filter_instance._get_tools_by_names([canonical], available_tools) assert len(matched) == 1 assert matched[0]["name"] == canonical @@ -782,9 +854,7 @@ class TestGetToolsByNames: separator-anchored suffixes of ``litellm_api-fs-read_file``. """ filter_instance = self._make_filter() - available_tools = [ - {"name": "litellm_api-fs-read_file", "description": "read"} - ] + available_tools = [{"name": "litellm_api-fs-read_file", "description": "read"}] matched = filter_instance._get_tools_by_names( ["fs-read_file", "api-fs-read_file"], available_tools @@ -805,9 +875,7 @@ class TestGetToolsByNames: {"name": "my_" + canonical, "description": "plain search"}, ] - matched = filter_instance._get_tools_by_names( - [canonical], available_tools - ) + matched = filter_instance._get_tools_by_names([canonical], available_tools) assert len(matched) == 1 assert matched[0]["name"] == "my_" + canonical