Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/hooks/mcp_semantic_filter/hook.py: 22%
221 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2Semantic Tool Filter Hook
4Pre-call hook that filters MCP tools semantically before LLM inference.
5Reduces context window size and improves tool selection accuracy.
6"""
8from collections.abc import Awaitable, Callable, Collection, Iterable, Mapping, Sequence
9from typing import TYPE_CHECKING, Final, Optional
11from fastapi import HTTPException
12from typing_extensions import ReadOnly, TypedDict
14from litellm._logging import verbose_proxy_logger
15from litellm.constants import (
16 DEFAULT_MCP_SEMANTIC_FILTER_EMBEDDING_MODEL,
17 DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD,
18 DEFAULT_MCP_SEMANTIC_FILTER_TOP_K,
19)
20from litellm.integrations.custom_logger import CustomLogger
21from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
22 SemanticToolFilterContextWindowError,
23)
25if TYPE_CHECKING: 25 ↛ 26line 25 didn't jump to line 26 because the condition on line 25 was never true
26 from litellm.caching.caching import DualCache
27 from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
28 SemanticMCPToolFilter,
29 )
30 from litellm.proxy._types import UserAPIKeyAuth
31 from litellm.router import Router
34class SemanticToolFilterConfig(TypedDict, total=False):
35 enabled: ReadOnly[bool]
36 embedding_model: ReadOnly[str]
37 top_k: ReadOnly[int]
38 similarity_threshold: ReadOnly[float]
41def _truncate_csv_at_tool_name_boundary(tool_names_csv: str, max_length: int) -> str:
42 """Cap a CSV of tool names to max_length, dropping any name that does not fit whole."""
43 if len(tool_names_csv) <= max_length:
44 return tool_names_csv
46 head: Final = tool_names_csv[: max_length + 1]
47 if "," not in head:
48 return ""
50 return head.rsplit(",", 1)[0]
53class SemanticToolFilterHook(CustomLogger):
54 """
55 Pre-call hook that filters MCP tools semantically.
57 This hook:
58 1. Extracts the user query from messages
59 2. Filters tools based on semantic similarity to the query
60 3. Returns only the top-k most relevant tools to the LLM
61 """
63 def __init__(self, semantic_filter: "SemanticMCPToolFilter"):
64 """
65 Initialize the hook.
67 Args:
68 semantic_filter: SemanticMCPToolFilter instance
69 """
70 super().__init__()
71 self.filter = semantic_filter
73 verbose_proxy_logger.debug(
74 "Initialized SemanticToolFilterHook with filter: enabled=%s, top_k=%s",
75 semantic_filter.enabled,
76 semantic_filter.top_k,
77 )
79 def _should_expand_mcp_tools(self, tools: Iterable[Mapping[str, object]]) -> bool:
80 """
81 Check if tools contain MCP references with server_url="litellm_proxy".
83 Only expands MCP tools pointing to litellm proxy, not external MCP servers.
84 """
85 from litellm.responses.mcp.litellm_proxy_mcp_handler import (
86 LiteLLM_Proxy_MCP_Handler,
87 )
89 return LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(tools)
91 async def _expand_mcp_tools(
92 self,
93 tools: Iterable[Mapping[str, object]],
94 user_api_key_dict: "UserAPIKeyAuth",
95 ) -> list[dict[str, object]]:
96 """
97 Expand MCP references to actual tool definitions.
99 Reuses LiteLLM_Proxy_MCP_Handler._process_mcp_tools_to_openai_format
100 which internally does: parse -> fetch -> filter -> deduplicate -> transform
101 """
102 from litellm.responses.mcp.litellm_proxy_mcp_handler import (
103 LiteLLM_Proxy_MCP_Handler,
104 )
106 # Parse to separate MCP tools from other tools
107 mcp_tools, _ = await LiteLLM_Proxy_MCP_Handler._split_mcp_tools(tools)
109 if not mcp_tools:
110 return []
112 # Use single combined method instead of 3 separate calls
113 # This already handles: fetch -> filter by allowed_tools -> deduplicate -> transform
114 (
115 openai_tools,
116 _,
117 ) = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_to_openai_format(
118 user_api_key_auth=user_api_key_dict, mcp_tools_with_litellm_proxy=mcp_tools
119 )
121 # Convert Pydantic models to dicts for compatibility
122 openai_tools_as_dicts: Final[list[dict[str, object]]] = []
123 for tool in openai_tools:
124 if hasattr(tool, "model_dump"):
125 tool_dict = tool.model_dump(exclude_none=True)
126 verbose_proxy_logger.debug(
127 "Converted Pydantic tool to dict: %s -> dict with keys: %s",
128 type(tool).__name__,
129 list(tool_dict.keys()),
130 )
131 openai_tools_as_dicts.append(tool_dict)
132 elif hasattr(tool, "dict"):
133 tool_dict = tool.dict(exclude_none=True)
134 verbose_proxy_logger.debug("Converted Pydantic tool (v1) to dict: %s -> dict", type(tool).__name__)
135 openai_tools_as_dicts.append(tool_dict)
136 elif isinstance(tool, dict):
137 verbose_proxy_logger.debug("Tool is already a dict with keys: %s", list(tool.keys()))
138 openai_tools_as_dicts.append(tool)
139 else:
140 verbose_proxy_logger.warning("Tool is unknown type: %s, passing as-is", type(tool))
141 openai_tools_as_dicts.append(tool)
143 verbose_proxy_logger.debug(
144 "Expanded %s MCP reference(s) to %s tools (all as dicts)", len(mcp_tools), len(openai_tools_as_dicts)
145 )
147 return openai_tools_as_dicts
149 async def _filter_expanded_tools(
150 self,
151 data: dict,
152 expanded_tools: list[dict[str, object]],
153 ) -> list[dict[str, object]]:
154 """
155 Apply the semantic filter to expanded MCP tool definitions.
157 Expanded tools are flat OpenAI function dicts with a top-level
158 "name" (see transform_mcp_tool_to_openai_responses_api_tool), so
159 filter_tools can name-match them against the semantic router.
160 """
161 raw_messages: Final = data.get("messages") or data.get("input") or []
162 messages: Final = [{"role": "user", "content": raw_messages}] if isinstance(raw_messages, str) else raw_messages
163 user_query: Final = self.filter.extract_user_query(messages)
164 if not user_query:
165 verbose_proxy_logger.debug("No user query found, skipping semantic filter on expanded MCP tools")
166 return expanded_tools
168 return await self.filter.filter_tools(query=user_query, available_tools=expanded_tools)
170 def _selected_tool_names(self, filtered_tools: Sequence[object]) -> list[str]:
171 """Names of the semantically selected tools, as produced by the MCP expansion."""
172 names: Final = (self.filter._extract_tool_info(tool)[0] for tool in filtered_tools)
173 return [name for name in names if name]
175 @staticmethod
176 async def _narrow_mcp_references(
177 tools: Sequence[Mapping[str, object]],
178 selected_tool_names: list[str],
179 served_names: Callable[[Collection[str]], Awaitable[frozenset[str]]] | None = None,
180 ) -> list[object]:
181 """
182 Restrict each litellm_proxy MCP reference to the semantically selected tools.
184 The reference block is preserved rather than replaced with expanded tools, so the
185 MCP gateway still performs the expansion. That keeps the per-endpoint tool shape
186 and tool auto-execution intact. Expansion already applied any caller-supplied
187 allowed_tools, so this selection can only narrow a block further.
189 Whether an undecidable selection exposes every tool or none is owned by
190 SemanticMCPToolFilter.filter_tools, which returns the full set when nothing
191 matches; the same policy therefore governs references and plain tools. Passing an
192 empty selection through is safe rather than a hidden allow-all: the gateway reads
193 the union of every reference's allowed_tools and treats an empty union as unset.
194 """
195 from litellm.responses.mcp.litellm_proxy_mcp_handler import (
196 LiteLLM_Proxy_MCP_Handler,
197 )
199 via_gateway: Final = await (
200 LiteLLM_Proxy_MCP_Handler.routes_through_gateway(tools, served_names)
201 if served_names is not None
202 else LiteLLM_Proxy_MCP_Handler.routes_through_gateway(tools)
203 )
204 return [
205 {**tool, "allowed_tools": selected_tool_names} if isinstance(tool, dict) and routed else tool
206 for tool, routed in zip(tools, via_gateway, strict=True)
207 ]
209 def _is_mcp_tool(self, tool: object) -> bool:
210 """
211 Check whether *tool* is registered in the MCP semantic router.
213 Classification strategy (shape-first, lookup-second):
214 1. Chat Completions format dicts are always native.
215 2. Responses API function tools are always native.
216 3. Everything else is looked up by name in the MCP registry.
217 """
218 if isinstance(tool, dict) and tool.get("type") == "function" and isinstance(tool.get("function"), dict):
219 return False
220 if isinstance(tool, dict) and tool.get("type") == "function" and isinstance(tool.get("name"), str):
221 return False
222 name, _ = self.filter._extract_tool_info(tool)
223 return bool(name) and name in self.filter._tool_map
225 def _get_metadata_variable_name(self, data: dict) -> str:
226 if "litellm_metadata" in data:
227 return "litellm_metadata"
228 return "metadata"
230 def _emit_filter_metadata(
231 self,
232 data: dict,
233 mcp_tools: Sequence[object],
234 filtered_mcp_tools: Sequence[object],
235 native_tools: Sequence[object],
236 filtered_tools: Sequence[object],
237 ) -> None:
238 """
239 Emit response-header metadata when MCP tools were filtered.
241 Stats report MCP-only counts so downstream consumers see accurate
242 semantic filter metrics. Skips metadata entirely for purely-native
243 requests to avoid spurious headers.
244 """
245 if mcp_tools:
246 filter_stats: Final = f"{len(mcp_tools)}->{len(filtered_mcp_tools)}"
247 tool_names_csv: Final = self._get_tool_names_csv(filtered_mcp_tools)
249 _metadata_variable_name: Final = self._get_metadata_variable_name(data)
250 metadata: Final = data.setdefault(_metadata_variable_name, {})
251 metadata["litellm_semantic_filter_stats"] = filter_stats
252 metadata["litellm_semantic_filter_tools"] = tool_names_csv
254 verbose_proxy_logger.info(
255 "Semantic tool filter: %s MCP tools (%s native preserved, %s total)",
256 filter_stats,
257 len(native_tools),
258 len(filtered_tools),
259 )
260 else:
261 verbose_proxy_logger.info(
262 "Semantic tool filter: all %s tools are native, no MCP filtering applied", len(native_tools)
263 )
265 def _emit_filter_metadata_safe(
266 self,
267 data: dict,
268 mcp_tools: Sequence[object],
269 filtered_mcp_tools: Sequence[object],
270 native_tools: Sequence[object],
271 filtered_tools: Sequence[object],
272 ) -> None:
273 """
274 Emit filter metadata without letting an emission failure abort the
275 already-filtered request.
276 """
277 try:
278 self._emit_filter_metadata(
279 data=data,
280 mcp_tools=mcp_tools,
281 filtered_mcp_tools=filtered_mcp_tools,
282 native_tools=native_tools,
283 filtered_tools=filtered_tools,
284 )
285 except Exception as e:
286 verbose_proxy_logger.warning(
287 "Failed to emit semantic filter metadata: %s",
288 e,
289 exc_info=True,
290 )
292 async def async_pre_call_hook(
293 self,
294 user_api_key_dict: "UserAPIKeyAuth",
295 cache: "DualCache",
296 data: dict,
297 call_type: str,
298 ) -> Exception | str | dict | None:
299 """
300 Filter tools before LLM call based on user query.
302 This hook is called before the LLM request is made. It filters the
303 tools list to only include semantically relevant tools.
304 """
305 if call_type not in ("completion", "acompletion", "aresponses"): 305 ↛ 309line 305 didn't jump to line 309 because the condition on line 305 was always true
306 verbose_proxy_logger.debug("Skipping semantic filter for call_type=%s", call_type)
307 return None
309 tools: Final = data.get("tools")
310 if not tools:
311 verbose_proxy_logger.debug("No tools in request, skipping semantic filter")
312 return None
314 if self._should_expand_mcp_tools(tools):
315 verbose_proxy_logger.debug("Detected litellm_proxy MCP references, expanding before semantic filtering")
317 if not self.filter.enabled:
318 verbose_proxy_logger.debug("Semantic filter disabled, leaving MCP references untouched")
319 return None
321 try:
322 native_tools_before_expand = [t for t in tools if not (isinstance(t, dict) and t.get("type") == "mcp")]
324 expanded_tools: Final = await self._expand_mcp_tools(tools, user_api_key_dict)
326 if not expanded_tools:
327 verbose_proxy_logger.warning("No tools expanded from MCP references")
328 return None
330 filtered_expanded_tools = await self._filter_expanded_tools(data=data, expanded_tools=expanded_tools)
332 selected_tool_names: Final = self._selected_tool_names(filtered_expanded_tools)
333 narrowed_tools: Final = await self._narrow_mcp_references(tools, selected_tool_names)
334 data["tools"] = narrowed_tools
335 self._emit_filter_metadata_safe(
336 data=data,
337 mcp_tools=expanded_tools,
338 filtered_mcp_tools=filtered_expanded_tools,
339 native_tools=native_tools_before_expand,
340 filtered_tools=narrowed_tools,
341 )
342 verbose_proxy_logger.info(
343 "Expanded MCP references to %s tools (%s native preserved), semantic filter selected %s",
344 len(expanded_tools),
345 len(native_tools_before_expand),
346 len(filtered_expanded_tools),
347 )
348 return data
350 except SemanticToolFilterContextWindowError as e:
351 raise HTTPException(status_code=400, detail={"error": str(e)}) from e
352 except Exception as e:
353 verbose_proxy_logger.error("Failed to expand MCP references: %s", e, exc_info=True)
354 return None
356 messages = data.get("messages", [])
357 if not messages:
358 messages = data.get("input", [])
359 if not messages:
360 verbose_proxy_logger.debug("No messages in request, skipping semantic filter")
361 return None
363 if not self.filter.enabled:
364 verbose_proxy_logger.debug("Semantic filter disabled, skipping")
365 return None
367 try:
368 user_query: Final = self.filter.extract_user_query(messages)
369 if not user_query:
370 verbose_proxy_logger.debug("No user query found, skipping semantic filter")
371 return None
373 native_tools: Final[list[object]] = []
374 mcp_tools: Final[list[object]] = []
375 mcp_indices: Final[set[int]] = set()
376 for i, t in enumerate(tools):
377 if self._is_mcp_tool(t):
378 mcp_tools.append(t)
379 mcp_indices.add(i)
380 else:
381 native_tools.append(t)
383 verbose_proxy_logger.debug(
384 "Applying semantic filter: %s MCP tools, %s native tools, query: '%s...'",
385 len(mcp_tools),
386 len(native_tools),
387 user_query[:50],
388 )
390 if mcp_tools:
391 filtered_mcp_tools: list[object] = await self.filter.filter_tools(
392 query=user_query,
393 available_tools=mcp_tools,
394 )
395 else:
396 filtered_mcp_tools = []
398 filtered_mcp_names: Final[set[str]] = set()
399 for t in filtered_mcp_tools:
400 name, _ = self.filter._extract_tool_info(t)
401 if name:
402 filtered_mcp_names.add(name)
404 filtered_tools: Final[list[object]] = []
405 for i, t in enumerate(tools):
406 if i in mcp_indices:
407 name, _ = self.filter._extract_tool_info(t)
408 if name in filtered_mcp_names:
409 filtered_tools.append(t)
410 else:
411 filtered_tools.append(t)
413 data["tools"] = filtered_tools
415 self._emit_filter_metadata_safe(
416 data=data,
417 mcp_tools=mcp_tools,
418 filtered_mcp_tools=filtered_mcp_tools,
419 native_tools=native_tools,
420 filtered_tools=filtered_tools,
421 )
423 return data
425 except SemanticToolFilterContextWindowError as e:
426 raise HTTPException(status_code=400, detail={"error": str(e)}) from e
427 except Exception as e:
428 verbose_proxy_logger.warning("Semantic tool filter hook failed: %s. Proceeding with all tools.", e)
429 return None
431 async def async_post_call_response_headers_hook(
432 self,
433 data: dict,
434 user_api_key_dict: "UserAPIKeyAuth",
435 response: object,
436 request_headers: dict[str, str] | None = None,
437 litellm_call_info: dict[str, object] | None = None,
438 ) -> dict[str, str] | None:
439 """Add semantic filter stats and tool names to response headers."""
440 from litellm.constants import MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH
442 _metadata_variable_name: Final = self._get_metadata_variable_name(data)
443 metadata: Final = data.get(_metadata_variable_name, {})
445 filter_stats: Final = metadata.get("litellm_semantic_filter_stats")
446 if not filter_stats: 446 ↛ 449line 446 didn't jump to line 449 because the condition on line 446 was always true
447 return None
449 headers: Final = {"x-litellm-semantic-filter": filter_stats}
451 # Add CSV of filtered tool names (nginx-safe length)
452 tool_names_csv: Final = metadata.get("litellm_semantic_filter_tools", "")
453 header_safe_csv: Final = _truncate_csv_at_tool_name_boundary(
454 tool_names_csv=tool_names_csv,
455 max_length=MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH,
456 )
457 if header_safe_csv:
458 headers["x-litellm-semantic-filter-tools"] = header_safe_csv
460 return headers
462 def _get_tool_names_csv(self, tools: Sequence[object]) -> str:
463 """Extract tool names and return as CSV string."""
464 if not tools:
465 return ""
467 tool_names: Final = []
468 for tool in tools:
469 name = tool.get("name", "") if isinstance(tool, dict) else getattr(tool, "name", "")
470 if name:
471 tool_names.append(name)
473 return ",".join(tool_names)
475 @staticmethod
476 async def initialize_from_config(
477 config: SemanticToolFilterConfig | None,
478 llm_router: Optional["Router"],
479 ) -> Optional["SemanticToolFilterHook"]:
480 """
481 Initialize semantic tool filter from proxy config.
483 Args:
484 config: Proxy configuration dict (litellm_settings.mcp_semantic_tool_filter)
485 llm_router: LiteLLM router instance for embeddings
487 Returns:
488 SemanticToolFilterHook instance if enabled, None otherwise
489 """
490 from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
491 SemanticMCPToolFilter,
492 )
494 if not config or not config.get("enabled", False): 494 ↛ 495line 494 didn't jump to line 495 because the condition on line 494 was never true
495 verbose_proxy_logger.debug("Semantic tool filter not enabled in config")
496 return None
498 if llm_router is None: 498 ↛ 499line 498 didn't jump to line 499 because the condition on line 498 was never true
499 verbose_proxy_logger.warning("Cannot initialize semantic filter: llm_router is None")
500 return None
502 try:
503 embedding_model: Final = config.get("embedding_model", DEFAULT_MCP_SEMANTIC_FILTER_EMBEDDING_MODEL)
504 top_k: Final = config.get("top_k", DEFAULT_MCP_SEMANTIC_FILTER_TOP_K)
505 similarity_threshold = config.get("similarity_threshold", DEFAULT_MCP_SEMANTIC_FILTER_SIMILARITY_THRESHOLD)
507 semantic_filter: Final = SemanticMCPToolFilter(
508 embedding_model=embedding_model,
509 litellm_router_instance=llm_router,
510 top_k=top_k,
511 similarity_threshold=similarity_threshold,
512 enabled=True,
513 )
515 # Build router from MCP registry on startup
516 await semantic_filter.build_router_from_mcp_registry()
518 hook: Final = SemanticToolFilterHook(semantic_filter)
520 verbose_proxy_logger.info(
521 "✅ MCP Semantic Tool Filter enabled: embedding_model=%s, top_k=%s, similarity_threshold=%s",
522 embedding_model,
523 top_k,
524 similarity_threshold,
525 )
527 return hook
529 except ImportError as e:
530 verbose_proxy_logger.warning(
531 "semantic-router not installed. Install with: pip install 'litellm[semantic-router]'. Error: %s", e
532 )
533 return None
534 except Exception as e:
535 verbose_proxy_logger.exception("Failed to initialize MCP semantic tool filter: %s", e)
536 return None