diff --git a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py index e12c6cdbd56..37d3e0aedea 100644 --- a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py +++ b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py @@ -4,6 +4,7 @@ Semantic MCP Tool Filtering using semantic-router Filters MCP tools semantically for /chat/completions and /responses endpoints. """ +import asyncio from typing import TYPE_CHECKING, Any, Dict, List, Optional from litellm._logging import verbose_logger @@ -76,6 +77,7 @@ class SemanticMCPToolFilter: self.tool_router: Optional["SemanticRouter"] = None self.context_window_error: Optional[str] = None self._tool_map: Dict[str, Any] = {} # MCPTool objects or OpenAI function dicts + self._index_sync_lock = asyncio.Lock() async def build_router_from_mcp_registry(self) -> None: """Build semantic router from all MCP tools in the registry (no auth checks).""" @@ -182,6 +184,81 @@ class SemanticMCPToolFilter: return raise + def _has_tools_missing_from_index(self, tools: list[Any]) -> bool: + """Allocation-free check for any named tool not yet in the semantic index.""" + return any(name and name not in self._tool_map for name in (self._extract_tool_info(t)[0] for t in tools)) + + def _tools_missing_from_index(self, tools: list[Any]) -> dict[str, Any]: + """Map name -> tool for every named tool not yet in the semantic index.""" + return { + name: tool + for name, tool in ((self._extract_tool_info(t)[0], t) for t in tools) + if name and name not in self._tool_map + } + + async def _ensure_tools_indexed(self, available_tools: list[Any]) -> None: + """ + Index request-time tools the startup build never saw. + + The startup index lists every registered MCP server WITHOUT per-user + credentials, so servers requiring per-user auth (interactive OAuth + tokens, user-scoped env vars) contribute zero routes. Tools reaching + the filter came through an authenticated expansion; without indexing + them here they can never be selected, so requests either bypass + filtering entirely (N->N) or lose every tool to unrelated matches. + + Runs async-only (no synchronous embedding on the request path) and + never writes shared error state: an embedding failure here raises and + is scoped to the requesting call, so one request's oversized tool + description cannot poison the filter for other users on the worker. + """ + from semantic_router.routers import SemanticRouter + from semantic_router.routers.base import Route + + from litellm.router_strategy.auto_router.litellm_encoder import ( + LiteLLMRouterEncoder, + ) + + if not self._has_tools_missing_from_index(available_tools): + return + + async with self._index_sync_lock: + missing = self._tools_missing_from_index(available_tools) + if not missing: + return + + descriptions = {name: self._extract_tool_info(tool)[1] for name, tool in missing.items()} + routes = [ + Route( + name=name, + description=description, + utterances=[description], + score_threshold=self.similarity_threshold, + ) + for name, description in descriptions.items() + ] + + if self.tool_router is None: + router = SemanticRouter( + routes=[], + encoder=LiteLLMRouterEncoder( + litellm_router_instance=self.router_instance, + model_name=self.embedding_model, + score_threshold=self.similarity_threshold, + ), + auto_sync="local", + top_k=self.top_k, + ) + await router.aadd(routes) + self.tool_router = router + else: + await self.tool_router.aadd(routes) + + self._tool_map.update(missing) + verbose_logger.info( + f"Semantic tool filter indexed {len(routes)} request-time tools missing from the startup index" + ) + async def filter_tools( self, query: str, @@ -216,22 +293,34 @@ class SemanticMCPToolFilter: if not query or not query.strip(): return available_tools - # Router should be built on startup - if not, something went wrong - if self.tool_router is None: - verbose_logger.warning("Router not initialized - was build_router_from_mcp_registry() called on startup?") - return available_tools - # Run semantic filtering try: + await self._ensure_tools_indexed(available_tools) + + if self.tool_router is None: + verbose_logger.warning("Semantic router could not be built from the request's tools") + return available_tools + + available_names = [name for name in (self._extract_tool_info(t)[0] for t in available_tools) if name] + if not available_names: + return available_tools + limit = top_k or self.top_k - matches = self.tool_router(text=query, limit=limit) + if self.tool_router.top_k < limit: + self.tool_router.top_k = limit + matches = self.tool_router(text=query, limit=limit, route_filter=available_names) matched_tool_names = self._extract_tool_names_from_matches(matches) if not matched_tool_names: return available_tools - return self._get_tools_by_names(matched_tool_names, available_tools) + filtered_tools = self._get_tools_by_names(matched_tool_names, available_tools) + if not filtered_tools: + return available_tools + return filtered_tools + except SemanticToolFilterContextWindowError: + raise except Exception as e: if _is_context_window_error(e): verbose_logger.error( @@ -240,7 +329,7 @@ class SemanticMCPToolFilter: ) raise SemanticToolFilterContextWindowError( embedding_model=self.embedding_model, - stage="the user query", + stage="the user query or the MCP tool descriptions being indexed", original_error=str(e), ) from e verbose_logger.error(f"Semantic tool filter failed: {e}", exc_info=True) 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 9bc0a525326..0f392a54b4c 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 @@ -1664,3 +1664,252 @@ def test_is_context_window_error_detection_variants(): assert _is_context_window_error(ValueError("Invalid 'input[0]': maximum input length is 8192 tokens.")) assert not _is_context_window_error(ValueError("A generic API error occurred.")) assert not _is_context_window_error(None) + + +def _make_keyword_embedding_router(recorded_inputs): + """ + Mock litellm Router whose embeddings are deterministic keyword one-hots: + texts mentioning linear/issue/ticket embed to [1, 0], everything else to + [0, 1]. Lets tests assert real similarity ranking through the actual + semantic-router index. Every embedding input batch is appended to + recorded_inputs. + """ + from litellm.types.utils import Embedding, EmbeddingResponse + + def _vector(text): + lowered = text.lower() + if "kanban" in lowered: + return [0.6, 0.8] + if "linear" in lowered or "issue" in lowered or "ticket" in lowered: + return [1.0, 0.0] + return [0.0, 1.0] + + def mock_embedding_sync(*args, **kwargs): + texts = kwargs["input"] + recorded_inputs.append(list(texts)) + return EmbeddingResponse( + data=[Embedding(embedding=_vector(t), index=i, object="embedding") for i, t in enumerate(texts)], + 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(*args, **kwargs) + + mock_router = Mock() + mock_router.embedding = mock_embedding_sync + mock_router.aembedding = mock_embedding_async + return mock_router + + +def _make_keyword_filter(recorded_inputs, top_k: int = 3): + from litellm.proxy._experimental.mcp_server.semantic_tool_filter import ( + SemanticMCPToolFilter, + ) + + return SemanticMCPToolFilter( + embedding_model="text-embedding-3-small", + litellm_router_instance=_make_keyword_embedding_router(recorded_inputs), + top_k=top_k, + similarity_threshold=0.3, + enabled=True, + ) + + +def _linear_issue_tool(): + return MCPTool( + name="linear_stub-get_issue", + description="Get a Linear issue (ticket) by its identifier such as LIT-1234", + inputSchema={"type": "object"}, + ) + + +def _linear_list_tool(): + return MCPTool( + name="linear_stub-list_issues", + description="List Linear issues (tickets) in the workspace", + inputSchema={"type": "object"}, + ) + + +def _weather_tool(): + return MCPTool( + name="weather_stub-get_weather", + description="Get the current weather conditions for a city", + inputSchema={"type": "object"}, + ) + + +@pytest.mark.asyncio +async def test_filter_indexes_request_tools_when_startup_index_is_empty(): + """ + Regression test: the startup index is built by listing every MCP server + WITHOUT per-user credentials, so a gateway whose servers all require + per-user auth (e.g. interactive OAuth) starts with an empty index + (tool_router is None). filter_tools then failed open and returned all N + tools unfiltered (customer-visible as an N->N header and, past 128 tools, + an OpenAI 400 "tools array too long"). The authed request-time tools must + instead be indexed on first sight so filtering actually runs. + """ + filter_instance = _make_keyword_filter([]) + assert filter_instance.tool_router is None + + tools = [_linear_issue_tool(), _weather_tool()] + filtered = await filter_instance.filter_tools( + query="what is Linear ticket LIT-3794 about", + available_tools=tools, + ) + + assert [t.name for t in filtered] == ["linear_stub-get_issue"] + print("✅ Empty startup index is built from authed request-time tools") + + +@pytest.mark.asyncio +async def test_filter_indexes_tools_missing_from_partial_index(): + """ + Regression test: servers whose tools/list needs per-user auth contribute + zero routes to the startup index while anonymously listable servers are + indexed. Tools reaching the filter through the authed request-time + expansion must be added to the existing router (and only embedded once; + repeat requests embed just the query). + """ + recorded_inputs = [] + filter_instance = _make_keyword_filter(recorded_inputs) + filter_instance._build_router([_weather_tool()]) + assert filter_instance.tool_router is not None + + tools = [_linear_issue_tool(), _weather_tool()] + query = "what is Linear ticket LIT-3794 about" + + filtered = await filter_instance.filter_tools(query=query, available_tools=tools) + assert [t.name for t in filtered] == ["linear_stub-get_issue"] + + calls_after_first = len(recorded_inputs) + filtered_again = await filter_instance.filter_tools(query=query, available_tools=tools) + assert [t.name for t in filtered_again] == ["linear_stub-get_issue"] + assert len(recorded_inputs) == calls_after_first + 1 + + print("✅ Partial startup index is completed from request-time tools, embedding each tool once") + + +@pytest.mark.asyncio +async def test_filter_fails_open_when_matches_are_not_in_available_tools(): + """ + Regression test: when the semantic router's matches are all tools that are + NOT in the request's available_tools (an index/request mismatch), the + filter returned an empty list, stripping every tool from the request and + breaking it outright (observed live as a 3->0 header followed by a + provider 400). It must fail open with the full tool list instead, matching + the zero-match fallback. + """ + filter_instance = _make_keyword_filter([]) + filter_instance._build_router([_weather_tool()]) + + tools = [_linear_issue_tool(), _linear_list_tool()] + filtered = await filter_instance.filter_tools( + query="current weather in San Francisco", + available_tools=tools, + ) + + assert [t.name for t in filtered] == ["linear_stub-get_issue", "linear_stub-list_issues"] + print("✅ Matches outside available_tools fail open instead of dropping every tool") + + +@pytest.mark.asyncio +async def test_request_time_context_window_error_is_request_scoped(): + """ + Regression test: an oversized tool description hitting the embedding + context window while lazily indexing request-time tools must fail only + the requesting call. Previously the lazy path reused the startup build + and recorded the overflow in the shared context_window_error, after + which EVERY user's MCP requests on the worker were blocked with a 400 + until restart (index poisoning via a single request). + """ + from litellm.proxy._experimental.mcp_server.semantic_tool_filter import ( + SemanticToolFilterContextWindowError, + ) + + state = {"raise_context_error": True} + filter_instance = _make_context_window_filter(state) + tools = [ + MCPTool(name="tool_a", description="Tool A", inputSchema={"type": "object"}), + MCPTool(name="tool_b", description="Tool B", inputSchema={"type": "object"}), + ] + + with pytest.raises(SemanticToolFilterContextWindowError): + await filter_instance.filter_tools(query="send an email", available_tools=tools) + + assert filter_instance.context_window_error is None + assert filter_instance.tool_router is None + + state["raise_context_error"] = False + filtered = await filter_instance.filter_tools(query="send an email", available_tools=tools) + + assert len(filtered) > 0 + assert filter_instance.context_window_error is None + assert filter_instance.tool_router is not None + print("✅ Request-time context window overflow is scoped to the request, not the worker") + + +@pytest.mark.asyncio +async def test_foreign_index_routes_cannot_displace_available_tools(): + """ + Regression test: routes indexed from OTHER principals' tool listings must + not occupy the match candidate set for this request. Previously the + router matched over the whole shared index, so foreign routes that + embedded closer to the query displaced the caller's own tools from + top_k, degrading results to the fail-open list (or, before the + empty-result guard, stripping every tool). Matching is now scoped to the + request's own tool names via route_filter. + """ + filter_instance = _make_keyword_filter([], top_k=1) + foreign_tools = [ + MCPTool( + name=f"other_user-linear_tool_{i}", + description=f"Get a Linear issue variant {i}", + inputSchema={"type": "object"}, + ) + for i in range(6) + ] + filter_instance._build_router(foreign_tools) + + my_kanban = MCPTool( + name="mine-kanban_board", + description="Manage kanban board cards", + inputSchema={"type": "object"}, + ) + filtered = await filter_instance.filter_tools( + query="what is Linear ticket LIT-3794 about", + available_tools=[my_kanban, _weather_tool()], + ) + + assert [t.name for t in filtered] == ["mine-kanban_board"] + print("✅ Foreign index routes cannot displace the caller's own tools") + + +@pytest.mark.asyncio +async def test_top_k_above_router_default_is_respected(): + """ + Regression test: semantic-router's SemanticRouter defaults to top_k=5 at + the index-query layer, silently capping any configured filter top_k + above 5 regardless of the limit passed to __call__. The router must be + sized (and resized) to honor the configured top_k. + """ + filter_instance = _make_keyword_filter([], top_k=6) + tools = [ + MCPTool( + name=f"linear_stub-tool_{i}", + description=f"Work with Linear issues part {i}", + inputSchema={"type": "object"}, + ) + for i in range(6) + ] + + filtered = await filter_instance.filter_tools( + query="Linear ticket work", + available_tools=tools, + ) + + assert len(filtered) == 6 + print("✅ Configured top_k above the semantic-router default of 5 is honored")