diff --git a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py index e8a96ee05e5..6d033f4eec9 100644 --- a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py +++ b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py @@ -5,14 +5,13 @@ This module provides semantic filtering for MCP tools to reduce context window s and improve tool selection accuracy. It leverages the existing semantic-router library and LiteLLMRouterEncoder to provide efficient tool filtering based on user queries. """ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union from litellm._logging import verbose_logger if TYPE_CHECKING: from mcp.types import Tool as MCPTool from semantic_router.routers import SemanticRouter - from semantic_router.schema import RouteChoice from litellm.router import Router @@ -20,11 +19,11 @@ if TYPE_CHECKING: class SemanticMCPToolFilter: """ Filters MCP tools using semantic-router library. - + Converts MCP tools to semantic-router Routes and uses SemanticRouter to find the most relevant tools for a given user query. """ - + def __init__( self, embedding_model: str, @@ -35,7 +34,7 @@ class SemanticMCPToolFilter: ): """ Initialize the semantic tool filter. - + Args: embedding_model: Model to use for generating embeddings (e.g., "text-embedding-3-small") litellm_router_instance: Router instance for embedding generation @@ -49,54 +48,83 @@ class SemanticMCPToolFilter: self.embedding_model = embedding_model self.router_instance = litellm_router_instance self.tool_router: Optional["SemanticRouter"] = None - self._tool_map: Dict[str, "MCPTool"] = {} # name -> tool - + self._tool_map: Dict[str, Union["MCPTool", Dict[str, Any]]] = {} # name -> tool + verbose_logger.debug( f"Initialized SemanticMCPToolFilter: enabled={enabled}, " f"top_k={top_k}, threshold={similarity_threshold}, " f"model={embedding_model}" ) - - def _mcp_tools_to_routes(self, tools: List["MCPTool"]) -> List: + + def _get_tool_info(self, tool: Union["MCPTool", Dict[str, Any]]) -> Tuple[str, str]: """ - Convert MCP tools to semantic-router Routes. - + Extract name and description from either MCP Tool or OpenAI function format. + Args: - tools: List of MCP tools - + tool: Either MCPTool object or OpenAI function dict + + Returns: + Tuple of (name, description) + """ + if isinstance(tool, dict): + # OpenAI function calling format: {"type": "function", "function": {"name": ..., "description": ...}} + if "function" in tool: + func = tool["function"] + name = func.get("name", "") + description = func.get("description", name) + else: + # Fallback for other dict formats + name = tool.get("name", "") + description = tool.get("description", name) + else: + # MCP Tool object + name = tool.name + description = tool.description or tool.name + + return name, description + + def _mcp_tools_to_routes( + self, tools: List[Union["MCPTool", Dict[str, Any]]] + ) -> List: + """ + Convert MCP tools or OpenAI function tools to semantic-router Routes. + + Args: + tools: List of MCP tools or OpenAI function dicts + Returns: List of Route objects """ from semantic_router.routers.base import Route - + routes = [] self._tool_map = {} - + for tool in tools: - self._tool_map[tool.name] = tool - + name, description = self._get_tool_info(tool) + self._tool_map[name] = tool + # Use tool description as both description and utterance - description = tool.description or tool.name utterances = [description] if description else [] - + routes.append( Route( - name=tool.name, + name=name, description=description, utterances=utterances, score_threshold=self.similarity_threshold, ) ) - + verbose_logger.debug(f"Converted {len(tools)} MCP tools to Routes") return routes - - def rebuild_router(self, tools: List["MCPTool"]) -> None: + + def rebuild_router(self, tools: List[Union["MCPTool", Dict[str, Any]]]) -> None: """ Rebuild semantic router with updated tools. - + This should be called whenever the tool list changes (server add/update/remove). - + Args: tools: Updated list of all available MCP tools """ @@ -105,15 +133,15 @@ class SemanticMCPToolFilter: from litellm.router_strategy.auto_router.litellm_encoder import ( LiteLLMRouterEncoder, ) - + if not tools: self.tool_router = None verbose_logger.debug("No tools provided, semantic router set to None") return - + try: routes = self._mcp_tools_to_routes(tools) - + self.tool_router = SemanticRouter( routes=routes, encoder=LiteLLMRouterEncoder( @@ -123,64 +151,66 @@ class SemanticMCPToolFilter: ), auto_sync="local", # Build index immediately ) - + verbose_logger.info( f"Rebuilt semantic router with {len(routes)} tool routes" ) - + except Exception as e: verbose_logger.error(f"Failed to rebuild semantic router: {e}") self.tool_router = None raise - + async def filter_tools( self, query: str, - available_tools: List["MCPTool"], + available_tools: List[Union["MCPTool", Dict[str, Any]]], top_k: Optional[int] = None, - ) -> List["MCPTool"]: + ) -> List[Union["MCPTool", Dict[str, Any]]]: """ Filter tools semantically based on query. - + Args: query: User query to match against tools available_tools: Full list of available tools top_k: Override default top_k (optional) - + Returns: Filtered and ordered list of tools (up to top_k) """ # Query semantic router with limit for top-k matches from semantic_router.schema import RouteChoice + if not self.enabled or not available_tools: return available_tools - + if not query or not query.strip(): verbose_logger.debug("Empty query, returning all tools") return available_tools - + top_k = top_k or self.top_k - + try: # Rebuild router if needed (first time or tools changed) if self.tool_router is None: verbose_logger.debug("Router not initialized, rebuilding...") self.rebuild_router(available_tools) - + if self.tool_router is None: verbose_logger.warning("Router rebuild failed, returning all tools") return available_tools - - - verbose_logger.debug(f"Querying semantic router with: '{query[:50]}...' (top_k={top_k})") + + verbose_logger.debug( + f"Querying semantic router with: '{query[:50]}...' (top_k={top_k})" + ) matches = self.tool_router(text=query, limit=top_k) - + if not matches: verbose_logger.warning( f"No tools matched query. Returning all {len(available_tools)} tools." ) return available_tools - + # Extract matched tool names matched_names: List[str] = [] if isinstance(matches, RouteChoice): @@ -188,41 +218,49 @@ class SemanticMCPToolFilter: matched_names = [matches.name] elif isinstance(matches, list): # semantic-router returns list of RouteChoice, take top_k - matched_names = [m.name for m in matches[:top_k] if hasattr(m, 'name') and m.name is not None] - + matched_names = [ + m.name + for m in matches[:top_k] + if hasattr(m, "name") and m.name is not None + ] + if not matched_names: - verbose_logger.warning("No matched tool names extracted, returning all tools") + verbose_logger.warning( + "No matched tool names extracted, returning all tools" + ) return available_tools - + # Filter available tools by matched names (preserve order from semantic router) matched_name_set = set(matched_names) - filtered = [ - tool for tool in available_tools - if tool.name in matched_name_set - ] - + filtered = [] + for tool in available_tools: + tool_name, _ = self._get_tool_info(tool) + if tool_name in matched_name_set: + filtered.append(tool) + # Reorder based on semantic router's ordering - name_to_tool = {tool.name: tool for tool in filtered} + name_to_tool = {} + for tool in filtered: + tool_name, _ = self._get_tool_info(tool) + name_to_tool[tool_name] = tool ordered_filtered = [ - name_to_tool[name] for name in matched_names - if name in name_to_tool + name_to_tool[name] for name in matched_names if name in name_to_tool ] return ordered_filtered if ordered_filtered else available_tools - + except Exception as e: verbose_logger.error( - f"Semantic tool filter failed: {e}. Returning all tools.", - exc_info=True + f"Semantic tool filter failed: {e}. Returning all tools.", exc_info=True ) return available_tools - + def extract_user_query(self, messages: List[Dict[str, Any]]) -> str: """ Extract user query from messages. - + Args: messages: List of message dictionaries - + Returns: Extracted query string """ @@ -230,11 +268,11 @@ class SemanticMCPToolFilter: for msg in reversed(messages): if msg.get("role") == "user": content = msg.get("content", "") - + # Handle string content if isinstance(content, str): return content - + # Handle content blocks (list) elif isinstance(content, list): texts = [] @@ -244,5 +282,5 @@ class SemanticMCPToolFilter: elif isinstance(block, str): texts.append(block) return " ".join(texts) - + return ""