mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
clean up semantic tool filter
This commit is contained in:
parent
3967326922
commit
9f6257ca9c
1 changed files with 98 additions and 168 deletions
|
|
@ -1,11 +1,9 @@
|
||||||
"""
|
"""
|
||||||
Semantic MCP Tool Filtering using semantic-router
|
Semantic MCP Tool Filtering using semantic-router
|
||||||
|
|
||||||
This module provides semantic filtering for MCP tools to reduce context window size
|
Filters MCP tools semantically for /chat/completions and /responses endpoints.
|
||||||
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, Tuple, Union
|
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||||
|
|
||||||
from litellm._logging import verbose_logger
|
from litellm._logging import verbose_logger
|
||||||
|
|
||||||
|
|
@ -17,12 +15,7 @@ if TYPE_CHECKING:
|
||||||
|
|
||||||
|
|
||||||
class SemanticMCPToolFilter:
|
class SemanticMCPToolFilter:
|
||||||
"""
|
"""Filters MCP tools using semantic similarity to reduce context window size."""
|
||||||
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__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
|
|
@ -36,7 +29,7 @@ class SemanticMCPToolFilter:
|
||||||
Initialize the semantic tool filter.
|
Initialize the semantic tool filter.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
embedding_model: Model to use for generating embeddings (e.g., "text-embedding-3-small")
|
embedding_model: Model to use for embeddings (e.g., "text-embedding-3-small")
|
||||||
litellm_router_instance: Router instance for embedding generation
|
litellm_router_instance: Router instance for embedding generation
|
||||||
top_k: Maximum number of tools to return
|
top_k: Maximum number of tools to return
|
||||||
similarity_threshold: Minimum similarity score for filtering
|
similarity_threshold: Minimum similarity score for filtering
|
||||||
|
|
@ -48,87 +41,37 @@ class SemanticMCPToolFilter:
|
||||||
self.embedding_model = embedding_model
|
self.embedding_model = embedding_model
|
||||||
self.router_instance = litellm_router_instance
|
self.router_instance = litellm_router_instance
|
||||||
self.tool_router: Optional["SemanticRouter"] = None
|
self.tool_router: Optional["SemanticRouter"] = None
|
||||||
self._tool_map: Dict[str, Union["MCPTool", Dict[str, Any]]] = {} # name -> tool
|
self._tool_map: Dict[str, "MCPTool"] = {}
|
||||||
|
|
||||||
verbose_logger.debug(
|
async def build_router_from_mcp_registry(self) -> None:
|
||||||
f"Initialized SemanticMCPToolFilter: enabled={enabled}, "
|
"""Build semantic router from all MCP tools in the registry."""
|
||||||
f"top_k={top_k}, threshold={similarity_threshold}, "
|
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||||
f"model={embedding_model}"
|
global_mcp_server_manager,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _get_tool_info(self, tool: Union["MCPTool", Dict[str, Any]]) -> Tuple[str, str]:
|
try:
|
||||||
"""
|
# Fetch all MCP tools from the registry (no user auth = all servers)
|
||||||
Extract name and description from either MCP Tool or OpenAI function format.
|
tools = await global_mcp_server_manager.list_tools(
|
||||||
|
user_api_key_auth=None,
|
||||||
Args:
|
mcp_auth_header=None,
|
||||||
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:
|
|
||||||
name, description = self._get_tool_info(tool)
|
|
||||||
self._tool_map[name] = tool
|
|
||||||
|
|
||||||
# Use tool description as both description and utterance
|
|
||||||
utterances = [description] if description else []
|
|
||||||
|
|
||||||
routes.append(
|
|
||||||
Route(
|
|
||||||
name=name,
|
|
||||||
description=description,
|
|
||||||
utterances=utterances,
|
|
||||||
score_threshold=self.similarity_threshold,
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
|
|
||||||
verbose_logger.debug(f"Converted {len(tools)} MCP tools to Routes")
|
if not tools:
|
||||||
return routes
|
verbose_logger.warning("No MCP tools found in registry")
|
||||||
|
self.tool_router = None
|
||||||
|
return
|
||||||
|
|
||||||
def rebuild_router(self, tools: List[Union["MCPTool", Dict[str, Any]]]) -> None:
|
self._build_router(tools)
|
||||||
"""
|
|
||||||
Rebuild semantic router with updated tools.
|
|
||||||
|
|
||||||
This should be called whenever the tool list changes (server add/update/remove).
|
except Exception as e:
|
||||||
|
verbose_logger.error(f"Failed to build router from MCP registry: {e}")
|
||||||
|
self.tool_router = None
|
||||||
|
raise
|
||||||
|
|
||||||
Args:
|
def _build_router(self, tools: List["MCPTool"]) -> None:
|
||||||
tools: Updated list of all available MCP tools
|
"""Build semantic router with tools."""
|
||||||
"""
|
|
||||||
from semantic_router.routers import SemanticRouter
|
from semantic_router.routers import SemanticRouter
|
||||||
|
from semantic_router.routers.base import Route
|
||||||
|
|
||||||
from litellm.router_strategy.auto_router.litellm_encoder import (
|
from litellm.router_strategy.auto_router.litellm_encoder import (
|
||||||
LiteLLMRouterEncoder,
|
LiteLLMRouterEncoder,
|
||||||
|
|
@ -136,11 +79,26 @@ class SemanticMCPToolFilter:
|
||||||
|
|
||||||
if not tools:
|
if not tools:
|
||||||
self.tool_router = None
|
self.tool_router = None
|
||||||
verbose_logger.debug("No tools provided, semantic router set to None")
|
|
||||||
return
|
return
|
||||||
|
|
||||||
try:
|
try:
|
||||||
routes = self._mcp_tools_to_routes(tools)
|
# Convert tools to routes
|
||||||
|
routes = []
|
||||||
|
self._tool_map = {}
|
||||||
|
|
||||||
|
for tool in tools:
|
||||||
|
name = tool.name
|
||||||
|
description = tool.description or tool.name
|
||||||
|
self._tool_map[name] = tool
|
||||||
|
|
||||||
|
routes.append(
|
||||||
|
Route(
|
||||||
|
name=name,
|
||||||
|
description=description,
|
||||||
|
utterances=[description],
|
||||||
|
score_threshold=self.similarity_threshold,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
self.tool_router = SemanticRouter(
|
self.tool_router = SemanticRouter(
|
||||||
routes=routes,
|
routes=routes,
|
||||||
|
|
@ -149,138 +107,110 @@ class SemanticMCPToolFilter:
|
||||||
model_name=self.embedding_model,
|
model_name=self.embedding_model,
|
||||||
score_threshold=self.similarity_threshold,
|
score_threshold=self.similarity_threshold,
|
||||||
),
|
),
|
||||||
auto_sync="local", # Build index immediately
|
auto_sync="local",
|
||||||
)
|
)
|
||||||
|
|
||||||
verbose_logger.info(
|
verbose_logger.info(
|
||||||
f"Rebuilt semantic router with {len(routes)} tool routes"
|
f"Built semantic router with {len(routes)} MCP tools from registry"
|
||||||
)
|
)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
verbose_logger.error(f"Failed to rebuild semantic router: {e}")
|
verbose_logger.error(f"Failed to build semantic router: {e}")
|
||||||
self.tool_router = None
|
self.tool_router = None
|
||||||
raise
|
raise
|
||||||
|
|
||||||
async def filter_tools(
|
async def filter_tools(
|
||||||
self,
|
self,
|
||||||
query: str,
|
query: str,
|
||||||
available_tools: List[Union["MCPTool", Dict[str, Any]]],
|
available_tools: List["MCPTool"],
|
||||||
top_k: Optional[int] = None,
|
top_k: Optional[int] = None,
|
||||||
) -> List[Union["MCPTool", Dict[str, Any]]]:
|
) -> List["MCPTool"]:
|
||||||
"""
|
"""
|
||||||
Filter tools semantically based on query.
|
Filter tools semantically based on query.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
query: User query to match against tools
|
query: User query to match against tools
|
||||||
available_tools: Full list of available tools
|
available_tools: Full list of available MCP tools
|
||||||
top_k: Override default top_k (optional)
|
top_k: Override default top_k (optional)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Filtered and ordered list of tools (up to top_k)
|
Filtered and ordered list of tools (up to top_k)
|
||||||
"""
|
"""
|
||||||
# Query semantic router with limit for top-k matches
|
# Early returns for cases where we can't/shouldn't filter
|
||||||
from semantic_router.schema import RouteChoice
|
if not self.enabled:
|
||||||
|
|
||||||
if not self.enabled or not available_tools:
|
|
||||||
return available_tools
|
return available_tools
|
||||||
|
|
||||||
|
if not available_tools:
|
||||||
|
return available_tools
|
||||||
|
|
||||||
if not query or not query.strip():
|
if not query or not query.strip():
|
||||||
verbose_logger.debug("Empty query, returning all tools")
|
|
||||||
return available_tools
|
return available_tools
|
||||||
|
|
||||||
top_k = top_k or self.top_k
|
if self.tool_router is None:
|
||||||
|
verbose_logger.warning("Router not initialized, returning all tools")
|
||||||
|
return available_tools
|
||||||
|
|
||||||
|
# Run semantic filtering
|
||||||
try:
|
try:
|
||||||
# Rebuild router if needed (first time or tools changed)
|
limit = top_k or self.top_k
|
||||||
if self.tool_router is None:
|
matches = self.tool_router(text=query, limit=limit)
|
||||||
verbose_logger.debug("Router not initialized, rebuilding...")
|
matched_tool_names = self._extract_tool_names_from_matches(matches)
|
||||||
self.rebuild_router(available_tools)
|
|
||||||
|
if not matched_tool_names:
|
||||||
if self.tool_router is None:
|
|
||||||
verbose_logger.warning("Router rebuild failed, returning all tools")
|
|
||||||
return available_tools
|
return available_tools
|
||||||
|
|
||||||
verbose_logger.debug(
|
return self._get_tools_by_names(matched_tool_names)
|
||||||
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):
|
|
||||||
if matches.name:
|
|
||||||
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
|
|
||||||
]
|
|
||||||
|
|
||||||
if not matched_names:
|
|
||||||
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 = []
|
|
||||||
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 = {}
|
|
||||||
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
|
|
||||||
]
|
|
||||||
return ordered_filtered if ordered_filtered else available_tools
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
verbose_logger.error(
|
verbose_logger.error(f"Semantic tool filter failed: {e}", exc_info=True)
|
||||||
f"Semantic tool filter failed: {e}. Returning all tools.", exc_info=True
|
|
||||||
)
|
|
||||||
return available_tools
|
return available_tools
|
||||||
|
|
||||||
|
def _extract_tool_names_from_matches(self, matches) -> List[str]:
|
||||||
|
"""Extract tool names from semantic router match results."""
|
||||||
|
if not matches:
|
||||||
|
return []
|
||||||
|
|
||||||
|
# Handle single match
|
||||||
|
if hasattr(matches, "name") and matches.name:
|
||||||
|
return [matches.name]
|
||||||
|
|
||||||
|
# Handle list of matches
|
||||||
|
if isinstance(matches, list):
|
||||||
|
return [m.name for m in matches if hasattr(m, "name") and m.name]
|
||||||
|
|
||||||
|
return []
|
||||||
|
|
||||||
|
def _get_tools_by_names(self, tool_names: List[str]) -> List["MCPTool"]:
|
||||||
|
"""Get tools from tool map by their names, preserving order."""
|
||||||
|
return [
|
||||||
|
self._tool_map[name]
|
||||||
|
for name in tool_names
|
||||||
|
if name in self._tool_map
|
||||||
|
]
|
||||||
|
|
||||||
def extract_user_query(self, messages: List[Dict[str, Any]]) -> str:
|
def extract_user_query(self, messages: List[Dict[str, Any]]) -> str:
|
||||||
"""
|
"""
|
||||||
Extract user query from messages.
|
Extract user query from messages for /chat/completions or /responses.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
messages: List of message dictionaries
|
messages: List of message dictionaries (from 'messages' or 'input' field)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Extracted query string
|
Extracted query string
|
||||||
"""
|
"""
|
||||||
# Get the last user message
|
|
||||||
for msg in reversed(messages):
|
for msg in reversed(messages):
|
||||||
if msg.get("role") == "user":
|
if msg.get("role") == "user":
|
||||||
content = msg.get("content", "")
|
content = msg.get("content", "")
|
||||||
|
|
||||||
# Handle string content
|
|
||||||
if isinstance(content, str):
|
if isinstance(content, str):
|
||||||
return content
|
return content
|
||||||
|
|
||||||
# Handle content blocks (list)
|
if isinstance(content, list):
|
||||||
elif isinstance(content, list):
|
texts = [
|
||||||
texts = []
|
block.get("text", "") if isinstance(block, dict) else str(block)
|
||||||
for block in content:
|
for block in content
|
||||||
if isinstance(block, dict) and block.get("type") == "text":
|
if isinstance(block, (dict, str))
|
||||||
texts.append(block.get("text", ""))
|
]
|
||||||
elif isinstance(block, str):
|
|
||||||
texts.append(block)
|
|
||||||
return " ".join(texts)
|
return " ".join(texts)
|
||||||
|
|
||||||
return ""
|
return ""
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue