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

1""" 

2Semantic Tool Filter Hook 

3 

4Pre-call hook that filters MCP tools semantically before LLM inference. 

5Reduces context window size and improves tool selection accuracy. 

6""" 

7 

8from collections.abc import Awaitable, Callable, Collection, Iterable, Mapping, Sequence 

9from typing import TYPE_CHECKING, Final, Optional 

10 

11from fastapi import HTTPException 

12from typing_extensions import ReadOnly, TypedDict 

13 

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) 

24 

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 

32 

33 

34class SemanticToolFilterConfig(TypedDict, total=False): 

35 enabled: ReadOnly[bool] 

36 embedding_model: ReadOnly[str] 

37 top_k: ReadOnly[int] 

38 similarity_threshold: ReadOnly[float] 

39 

40 

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 

45 

46 head: Final = tool_names_csv[: max_length + 1] 

47 if "," not in head: 

48 return "" 

49 

50 return head.rsplit(",", 1)[0] 

51 

52 

53class SemanticToolFilterHook(CustomLogger): 

54 """ 

55 Pre-call hook that filters MCP tools semantically. 

56 

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 """ 

62 

63 def __init__(self, semantic_filter: "SemanticMCPToolFilter"): 

64 """ 

65 Initialize the hook. 

66 

67 Args: 

68 semantic_filter: SemanticMCPToolFilter instance 

69 """ 

70 super().__init__() 

71 self.filter = semantic_filter 

72 

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 ) 

78 

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". 

82 

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 ) 

88 

89 return LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(tools) 

90 

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. 

98 

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 ) 

105 

106 # Parse to separate MCP tools from other tools 

107 mcp_tools, _ = await LiteLLM_Proxy_MCP_Handler._split_mcp_tools(tools) 

108 

109 if not mcp_tools: 

110 return [] 

111 

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 ) 

120 

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) 

142 

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 ) 

146 

147 return openai_tools_as_dicts 

148 

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. 

156 

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 

167 

168 return await self.filter.filter_tools(query=user_query, available_tools=expanded_tools) 

169 

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] 

174 

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. 

183 

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. 

188 

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 ) 

198 

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 ] 

208 

209 def _is_mcp_tool(self, tool: object) -> bool: 

210 """ 

211 Check whether *tool* is registered in the MCP semantic router. 

212 

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 

224 

225 def _get_metadata_variable_name(self, data: dict) -> str: 

226 if "litellm_metadata" in data: 

227 return "litellm_metadata" 

228 return "metadata" 

229 

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. 

240 

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) 

248 

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 

253 

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 ) 

264 

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 ) 

291 

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. 

301 

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 

308 

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 

313 

314 if self._should_expand_mcp_tools(tools): 

315 verbose_proxy_logger.debug("Detected litellm_proxy MCP references, expanding before semantic filtering") 

316 

317 if not self.filter.enabled: 

318 verbose_proxy_logger.debug("Semantic filter disabled, leaving MCP references untouched") 

319 return None 

320 

321 try: 

322 native_tools_before_expand = [t for t in tools if not (isinstance(t, dict) and t.get("type") == "mcp")] 

323 

324 expanded_tools: Final = await self._expand_mcp_tools(tools, user_api_key_dict) 

325 

326 if not expanded_tools: 

327 verbose_proxy_logger.warning("No tools expanded from MCP references") 

328 return None 

329 

330 filtered_expanded_tools = await self._filter_expanded_tools(data=data, expanded_tools=expanded_tools) 

331 

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 

349 

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 

355 

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 

362 

363 if not self.filter.enabled: 

364 verbose_proxy_logger.debug("Semantic filter disabled, skipping") 

365 return None 

366 

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 

372 

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) 

382 

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 ) 

389 

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 = [] 

397 

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) 

403 

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) 

412 

413 data["tools"] = filtered_tools 

414 

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 ) 

422 

423 return data 

424 

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 

430 

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 

441 

442 _metadata_variable_name: Final = self._get_metadata_variable_name(data) 

443 metadata: Final = data.get(_metadata_variable_name, {}) 

444 

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 

448 

449 headers: Final = {"x-litellm-semantic-filter": filter_stats} 

450 

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 

459 

460 return headers 

461 

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 "" 

466 

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) 

472 

473 return ",".join(tool_names) 

474 

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. 

482 

483 Args: 

484 config: Proxy configuration dict (litellm_settings.mcp_semantic_tool_filter) 

485 llm_router: LiteLLM router instance for embeddings 

486 

487 Returns: 

488 SemanticToolFilterHook instance if enabled, None otherwise 

489 """ 

490 from litellm.proxy._experimental.mcp_server.semantic_tool_filter import ( 

491 SemanticMCPToolFilter, 

492 ) 

493 

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 

497 

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 

501 

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) 

506 

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 ) 

514 

515 # Build router from MCP registry on startup 

516 await semantic_filter.build_router_from_mcp_registry() 

517 

518 hook: Final = SemanticToolFilterHook(semantic_filter) 

519 

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 ) 

526 

527 return hook 

528 

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