mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix: _expand_mcp_tools
This commit is contained in:
parent
8a02888249
commit
8913390fe6
1 changed files with 111 additions and 29 deletions
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue