fix: _expand_mcp_tools

This commit is contained in:
Ishaan Jaffer 2026-02-02 15:52:54 -08:00
parent 8a02888249
commit 8913390fe6

View file

@ -4,7 +4,7 @@ Semantic Tool Filter Hook
Pre-call hook that filters MCP tools semantically before LLM inference.
Reduces context window size and improves tool selection accuracy.
"""
from typing import TYPE_CHECKING, Any, Dict, Optional, Union
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from litellm._logging import verbose_proxy_logger
from litellm.constants import (
@ -48,6 +48,52 @@ class SemanticToolFilterHook(CustomLogger):
f"enabled={semantic_filter.enabled}, top_k={semantic_filter.top_k}"
)
def _should_expand_mcp_tools(self, tools: List[Any]) -> bool:
"""
Check if tools contain MCP references with server_url="litellm_proxy".
Only expands MCP tools pointing to litellm proxy, not external MCP servers.
"""
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
LiteLLM_Proxy_MCP_Handler,
)
return LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(tools)
async def _expand_mcp_tools(
self,
tools: List[Any],
user_api_key_dict: "UserAPIKeyAuth",
) -> List[Dict[str, Any]]:
"""
Expand MCP references to actual tool definitions.
Reuses LiteLLM_Proxy_MCP_Handler._process_mcp_tools_to_openai_format
which internally does: parse -> fetch -> filter -> deduplicate -> transform
"""
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
LiteLLM_Proxy_MCP_Handler,
)
# Parse to separate MCP tools from other tools
mcp_tools, _ = LiteLLM_Proxy_MCP_Handler._parse_mcp_tools(tools)
if not mcp_tools:
return []
# Use single combined method instead of 3 separate calls
# This already handles: fetch -> filter by allowed_tools -> deduplicate -> transform
openai_tools, _ = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_to_openai_format(
user_api_key_auth=user_api_key_dict,
mcp_tools_with_litellm_proxy=mcp_tools
)
verbose_proxy_logger.debug(
f"Expanded {len(mcp_tools)} MCP reference(s) to {len(openai_tools)} tools"
)
return openai_tools
async def async_pre_call_hook(
self,
user_api_key_dict: "UserAPIKeyAuth",
@ -70,8 +116,8 @@ class SemanticToolFilterHook(CustomLogger):
Returns:
Modified data dict with filtered tools, or None if no changes
"""
# Only filter chat completions with tools
if call_type not in ("completion", "acompletion"):
# Only filter endpoints that support tools
if call_type not in ("completion", "acompletion", "aresponses"):
verbose_proxy_logger.debug(
f"Skipping semantic filter for call_type={call_type}"
)
@ -83,18 +129,49 @@ class SemanticToolFilterHook(CustomLogger):
verbose_proxy_logger.debug("No tools in request, skipping semantic filter")
return None
# Skip filtering if tools are MCP references (type: "mcp")
# These will be expanded by the responses API handler
if isinstance(tools, list) and len(tools) > 0:
first_tool = tools[0]
if isinstance(first_tool, dict) and first_tool.get("type") == "mcp":
verbose_proxy_logger.debug(
"Skipping semantic filter for MCP tool references (type: mcp)"
original_tool_count = len(tools)
# Check for MCP references (server_url="litellm_proxy") and expand them
if self._should_expand_mcp_tools(tools):
verbose_proxy_logger.debug(
"Detected litellm_proxy MCP references, expanding before semantic filtering"
)
try:
expanded_tools = await self._expand_mcp_tools(
tools, user_api_key_dict
)
if not expanded_tools:
verbose_proxy_logger.warning(
"No tools expanded from MCP references"
)
return None
verbose_proxy_logger.info(
f"Expanded {len(tools)} MCP reference(s) to {len(expanded_tools)} tools"
)
# Debug: log first expanded tool format
if expanded_tools:
verbose_proxy_logger.debug(
f"First expanded tool format: {type(expanded_tools[0])}, keys: {list(expanded_tools[0].keys()) if isinstance(expanded_tools[0], dict) else 'not a dict'}"
)
# Update tools for filtering
tools = expanded_tools
original_tool_count = len(tools)
except Exception as e:
verbose_proxy_logger.error(
f"Failed to expand MCP references: {e}", exc_info=True
)
return None
# Check if messages are present
# Check if messages are present (try both "messages" and "input" for responses API)
messages = data.get("messages", [])
if not messages:
messages = data.get("input", [])
if not messages:
verbose_proxy_logger.debug("No messages in request, skipping semantic filter")
return None
@ -119,24 +196,25 @@ class SemanticToolFilterHook(CustomLogger):
# Filter tools semantically
filtered_tools = await self.filter.filter_tools(
query=user_query,
available_tools=tools,
available_tools=tools, # type: ignore
)
# Only modify data if filtering actually reduced the tool count
if len(filtered_tools) < len(tools):
data["tools"] = filtered_tools
verbose_proxy_logger.info(
f"Semantic tool filter: {len(tools)} -> {len(filtered_tools)} tools"
)
return data
else:
verbose_proxy_logger.debug(
f"Semantic filter did not reduce tool count ({len(tools)}), "
"returning original"
)
return None
# Always update tools and emit header (even if count unchanged)
data["tools"] = filtered_tools
# Store filter stats for response header
filter_stats = f"{original_tool_count}->{len(filtered_tools)}"
# Store in proxy_server_request (internal field not sent to providers)
if "proxy_server_request" not in data:
data["proxy_server_request"] = {}
data["proxy_server_request"]["semantic_filter_stats"] = filter_stats
verbose_proxy_logger.info(
f"Semantic tool filter: {filter_stats} tools"
)
return data
except Exception as e:
verbose_proxy_logger.warning(
@ -162,16 +240,20 @@ class SemanticToolFilterHook(CustomLogger):
from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
SemanticMCPToolFilter,
)
verbose_proxy_logger.info(f"🔍 initialize_from_config called: config={config}, llm_router={llm_router is not None}")
if not config or not config.get("enabled", False):
verbose_proxy_logger.debug("Semantic tool filter not enabled in config")
verbose_proxy_logger.warning(f"❌ Semantic tool filter not enabled: config={config}")
return None
if llm_router is None:
verbose_proxy_logger.warning(
"Cannot initialize semantic filter: llm_router is None"
"❌ Cannot initialize semantic filter: llm_router is None"
)
return None
verbose_proxy_logger.info("✅ Config and router available, creating filter...")
try:
embedding_model = config.get(