Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/tool_search.py: 33%
246 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
1from __future__ import annotations
3import hashlib
4import json
5from collections.abc import Mapping, Sequence
6from dataclasses import dataclass
7from datetime import datetime
8from types import MappingProxyType
9from typing import TYPE_CHECKING, Any, Final, TypedDict
11from pydantic import ValidationError
12from typing_extensions import ReadOnly, Required, assert_never
14import litellm
15from litellm.llms.litellm_proxy.skills.skill_search import DEFAULT_SKILL_SEARCH_TOP_K
16from litellm.proxy.agent_endpoints.agent_search import DEFAULT_AGENT_SEARCH_TOP_K
17from litellm.proxy.common_utils.semantic_text_index import (
18 Embedder,
19 EmbeddingFailed,
20 SemanticTextIndex,
21 router_embedder,
22)
23from litellm.types.mcp import MCPToolSearchSettings
25if TYPE_CHECKING: 25 ↛ 26line 25 didn't jump to line 26 because the condition on line 25 was never true
26 from mcp.types import CallToolResult, Tool
28 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
29 from litellm.proxy._types import UserAPIKeyAuth
31MCP_TOOL_SEARCH_SETTINGS_KEY: Final[str] = "mcp_tool_search"
32MCP_TOOL_SEARCH_TOOL_NAME: Final[str] = "mcp_tool_search"
33MCP_TOOL_CALL_TOOL_NAME: Final[str] = "mcp_tool_call"
34MCP_PROXY_SEARCH_TOOL_NAME: Final[str] = "search_tools"
35MCP_PROXY_SCHEMA_TOOL_NAME: Final[str] = "get_tool_schema"
36MCP_PROXY_CALL_TOOL_NAME: Final[str] = "call_tool"
37MCP_PROXY_TOOL_NAMES: Final = frozenset(
38 (MCP_PROXY_SEARCH_TOOL_NAME, MCP_PROXY_SCHEMA_TOOL_NAME, MCP_PROXY_CALL_TOOL_NAME)
39)
40AGENT_SEARCH_TOOL_NAME: Final[str] = "agent_search"
41SKILL_SEARCH_TOOL_NAME: Final[str] = "skill_search"
42VIRTUAL_TOOL_NAMES: Final = frozenset(
43 (MCP_TOOL_SEARCH_TOOL_NAME, MCP_TOOL_CALL_TOOL_NAME, AGENT_SEARCH_TOOL_NAME, SKILL_SEARCH_TOOL_NAME)
44)
47def coerce_top_k(value: Any, default: int = 5) -> int:
48 try:
49 return int(value)
50 except (TypeError, ValueError):
51 return default
54class ToolSearchResult(TypedDict, total=False):
55 name: Required[ReadOnly[str]]
56 description: Required[ReadOnly[str]]
57 inputSchema: Required[ReadOnly[Mapping[str, object]]]
58 score: ReadOnly[float]
61class MCPProxySearchResult(TypedDict, total=False):
62 tool_id: Required[ReadOnly[str]]
63 name: Required[ReadOnly[str]]
64 description: Required[ReadOnly[str]]
65 score: ReadOnly[float]
68class MCPProxySchemaResult(MCPProxySearchResult, total=False):
69 inputSchema: Required[ReadOnly[Mapping[str, object]]]
70 outputSchema: ReadOnly[Mapping[str, object]]
73class MCPProxyToolIdentity(TypedDict):
74 server_id: ReadOnly[str]
75 tool_name: ReadOnly[str]
78@dataclass(frozen=True, slots=True)
79class MCPToolSearchHit:
80 tool: Tool
81 score: float | None = None
84@dataclass(frozen=True, slots=True)
85class SemanticToolRanker:
86 embed: Embedder
87 embedding_model: str
88 index: SemanticTextIndex
91global_mcp_tool_search_index: Final = SemanticTextIndex()
94def mcp_tool_search_settings() -> MCPToolSearchSettings | ValidationError:
95 try:
96 return MCPToolSearchSettings.model_validate(litellm.mcp_tool_search or {})
97 except ValidationError as exc:
98 return exc
101def _tool_result(tool: Tool) -> ToolSearchResult:
102 return {
103 "name": tool.name,
104 "description": tool.description or "",
105 "inputSchema": tool.input_schema,
106 }
109def _scored_result(tool: Tool, score: float) -> ToolSearchResult:
110 return {
111 "name": tool.name,
112 "description": tool.description or "",
113 "inputSchema": tool.input_schema,
114 "score": score,
115 }
118_MCP_PROXY_IDENTITY_META_KEY: Final[str] = "litellm.ai/proxy_tool_identity"
121def with_mcp_proxy_identity(tool: Tool, server_id: str) -> Tool:
122 identity: Final[MCPProxyToolIdentity] = {"server_id": server_id, "tool_name": tool.name}
123 return tool.model_copy(
124 update={ # mutable-ok: Pydantic update payload
125 "meta": {**(tool.meta or {}), _MCP_PROXY_IDENTITY_META_KEY: identity} # mutable-ok: metadata mapping
126 }
127 )
130def _mcp_proxy_identity(tool: Tool) -> MCPProxyToolIdentity:
131 identity: Final = None if tool.meta is None else tool.meta.get(_MCP_PROXY_IDENTITY_META_KEY)
132 if not isinstance(identity, Mapping):
133 raise TypeError("MCP proxy tool identity is missing")
134 server_id: Final = identity.get("server_id")
135 tool_name: Final = identity.get("tool_name")
136 if not isinstance(server_id, str) or not isinstance(tool_name, str):
137 raise TypeError("MCP proxy tool identity is invalid")
138 resolved: Final[MCPProxyToolIdentity] = {"server_id": server_id, "tool_name": tool_name}
139 return resolved
142def mcp_proxy_tool_id(tool: Tool) -> str:
143 identity: Final = _mcp_proxy_identity(tool)
144 return hashlib.sha256(f"{identity['server_id']}\0{identity['tool_name']}".encode()).hexdigest()[:32]
147def _proxy_search_result(hit: MCPToolSearchHit) -> MCPProxySearchResult:
148 base: Final[MCPProxySearchResult] = {
149 "tool_id": mcp_proxy_tool_id(hit.tool),
150 "name": hit.tool.name,
151 "description": hit.tool.description or "",
152 }
153 return {**base, "score": hit.score} if hit.score is not None else base # mutable-ok: wire result payload
156def _proxy_schema_result(tool: Tool) -> MCPProxySchemaResult:
157 base: Final[MCPProxySchemaResult] = {
158 "tool_id": mcp_proxy_tool_id(tool),
159 "name": tool.name,
160 "description": tool.description or "",
161 "inputSchema": tool.input_schema,
162 }
163 if tool.output_schema is None:
164 return base
165 return {**base, "outputSchema": tool.output_schema} # mutable-ok: wire schema payload
168def _tool_text(tool: Tool) -> str:
169 return "\n".join(part for part in (tool.name, tool.description or "") if part)
172def _keyword_score(query: str, tool: Tool) -> float:
173 haystack: Final = _tool_text(tool).lower()
174 return float(sum(1 for token in query.lower().split() if token in haystack))
177def _split_core_tools(tools: Sequence[Tool], core_tools: Sequence[str]) -> tuple[tuple[Tool, ...], tuple[Tool, ...]]:
178 by_name: Final = MappingProxyType({tool.name: tool for tool in tools})
179 core: Final = tuple(by_name[name] for name in dict.fromkeys(core_tools) if name in by_name)
180 rest: Final = tuple(tool for tool in tools if tool.name not in frozenset(core_tools))
181 return core, rest
184def _top_hits(
185 tools: Sequence[Tool], scores: Sequence[float], minimum: float, limit: int
186) -> tuple[tuple[float, Tool], ...]:
187 hits: Final = ((score, tool) for score, tool in zip(scores, tools, strict=True) if score >= minimum)
188 return tuple(sorted(hits, key=lambda hit: hit[0], reverse=True)[:limit])
191def search_tools(query: str, tools: Sequence[Tool], top_k: int = 5) -> tuple[ToolSearchResult, ...]:
192 """Keyword fallback used when no embedding model is configured: one point per query token found in the tool."""
193 if not query:
194 return ()
195 scores: Final = tuple(_keyword_score(query, tool) for tool in tools)
196 return tuple(_tool_result(tool) for _, tool in _top_hits(tools, scores, minimum=1.0, limit=top_k))
199async def rank_mcp_tools(
200 query: str,
201 tools: Sequence[Tool],
202 top_k: int,
203 settings: MCPToolSearchSettings,
204 ranker: SemanticToolRanker | None,
205) -> tuple[MCPToolSearchHit, ...] | EmbeddingFailed:
206 core, rest = _split_core_tools(tools, settings.core_tools)
207 core_hits: Final = tuple(MCPToolSearchHit(tool) for tool in core)
208 if not query:
209 return core_hits
210 limit: Final = min(top_k, settings.top_k)
211 if ranker is None:
212 scores: Final = tuple(_keyword_score(query, tool) for tool in rest)
213 return (
214 *core_hits,
215 *(MCPToolSearchHit(tool) for _, tool in _top_hits(rest, scores, minimum=1.0, limit=limit)),
216 )
217 semantic_scores: Final = await ranker.index.scores(
218 query, tuple(_tool_text(tool) for tool in rest), ranker.embed, ranker.embedding_model
219 )
220 if isinstance(semantic_scores, EmbeddingFailed):
221 return semantic_scores
222 return (
223 *core_hits,
224 *(
225 MCPToolSearchHit(tool, score)
226 for score, tool in _top_hits(rest, semantic_scores, settings.similarity_threshold, limit)
227 ),
228 )
231async def search_mcp_tools(
232 query: str,
233 tools: Sequence[Tool],
234 top_k: int,
235 settings: MCPToolSearchSettings,
236 ranker: SemanticToolRanker | None,
237) -> tuple[ToolSearchResult, ...] | EmbeddingFailed:
238 hits: Final = await rank_mcp_tools(query, tools, top_k, settings, ranker)
239 if isinstance(hits, EmbeddingFailed):
240 return hits
241 return tuple(
242 _scored_result(hit.tool, hit.score) if hit.score is not None else _tool_result(hit.tool) for hit in hits
243 )
246class _ToolParamSchema(TypedDict, total=False):
247 type: Required[ReadOnly[str]]
248 description: Required[ReadOnly[str]]
249 default: ReadOnly[int]
252class _ToolInputSchema(TypedDict):
253 type: ReadOnly[str]
254 properties: ReadOnly[Mapping[str, _ToolParamSchema]]
255 required: ReadOnly[Sequence[str]]
258class VirtualToolDefinition(TypedDict):
259 name: ReadOnly[str]
260 description: ReadOnly[str]
261 inputSchema: ReadOnly[_ToolInputSchema]
264def _json_array(*items: str) -> Sequence[str]:
265 return list(items) # mutable-ok: jsonschema's metaschema only accepts a JSON array for required
268_MCP_TOOL_SEARCH_DEFINITION: Final[VirtualToolDefinition] = {
269 "name": MCP_TOOL_SEARCH_TOOL_NAME,
270 "description": (
271 "Search for MCP tools by describing what you need. "
272 "Returns top matching tools with names, descriptions, and input schemas."
273 ),
274 "inputSchema": {
275 "type": "object",
276 "properties": {
277 "query": {
278 "type": "string",
279 "description": "What the tool should do, matched against names and descriptions.",
280 },
281 "top_k": {"type": "integer", "description": "Maximum number of results to return.", "default": 5},
282 },
283 "required": _json_array("query"),
284 },
285}
287_MCP_TOOL_CALL_DEFINITION: Final[VirtualToolDefinition] = {
288 "name": MCP_TOOL_CALL_TOOL_NAME,
289 "description": "Call an MCP tool by name with the given arguments.",
290 "inputSchema": {
291 "type": "object",
292 "properties": {
293 "tool_name": {"type": "string", "description": "The exact name of the MCP tool to call."},
294 "arguments": {"type": "object", "description": "Arguments to pass to the tool."},
295 },
296 "required": _json_array("tool_name"),
297 },
298}
300_AGENT_SEARCH_DEFINITION: Final[VirtualToolDefinition] = {
301 "name": AGENT_SEARCH_TOOL_NAME,
302 "description": "Find A2A agents by describing the task in natural language. Returns the best matching agents you can access, ranked by semantic similarity, each with its agent_id, name, description, skills, and score.",
303 "inputSchema": {
304 "type": "object",
305 "properties": {
306 "query": {"type": "string", "description": "The task the agent should be able to do, in natural language."},
307 "top_k": {
308 "type": "integer",
309 "description": "Maximum number of agents to return.",
310 "default": DEFAULT_AGENT_SEARCH_TOP_K,
311 },
312 },
313 "required": _json_array("query"),
314 },
315}
318_SKILL_SEARCH_DEFINITION: Final[VirtualToolDefinition] = {
319 "name": SKILL_SEARCH_TOOL_NAME,
320 "description": "Find registered skills by describing what you need in natural language. Returns the best "
321 "matching skills you can access, ranked by semantic similarity, each with its skill_id, display_title, "
322 "description, and score.",
323 "inputSchema": {
324 "type": "object",
325 "properties": {
326 "query": {"type": "string", "description": "What you need the skill to do, in natural language."},
327 "top_k": {
328 "type": "integer",
329 "description": "Maximum number of skills to return.",
330 "default": DEFAULT_SKILL_SEARCH_TOP_K,
331 },
332 },
333 "required": _json_array("query"),
334 },
335}
338_MCP_PROXY_SEARCH_DEFINITION: Final[VirtualToolDefinition] = {
339 "name": MCP_PROXY_SEARCH_TOOL_NAME,
340 "description": "Search accessible MCP tools by describing what you need. Returns opaque tool IDs.",
341 "inputSchema": {
342 "type": "object",
343 "properties": {"query": {"type": "string", "description": "What the tool should do."}},
344 "required": _json_array("query"),
345 },
346}
348_MCP_PROXY_SCHEMA_DEFINITION: Final[VirtualToolDefinition] = {
349 "name": MCP_PROXY_SCHEMA_TOOL_NAME,
350 "description": "Return the complete schema for an accessible MCP tool ID.",
351 "inputSchema": {
352 "type": "object",
353 "properties": {"tool_id": {"type": "string", "description": "Opaque ID from search_tools."}},
354 "required": _json_array("tool_id"),
355 },
356}
358_MCP_PROXY_CALL_DEFINITION: Final[VirtualToolDefinition] = {
359 "name": MCP_PROXY_CALL_TOOL_NAME,
360 "description": "Call an accessible MCP tool by opaque ID with schema-valid arguments.",
361 "inputSchema": {
362 "type": "object",
363 "properties": {
364 "tool_id": {"type": "string", "description": "Opaque ID from search_tools."},
365 "arguments": {"type": "object", "description": "Arguments validated against the selected tool schema."},
366 },
367 "required": _json_array("tool_id"),
368 },
369}
372def get_virtual_tool_definitions() -> tuple[VirtualToolDefinition, ...]:
373 return (_MCP_TOOL_SEARCH_DEFINITION, _MCP_TOOL_CALL_DEFINITION, _AGENT_SEARCH_DEFINITION, _SKILL_SEARCH_DEFINITION)
376def get_mcp_proxy_tool_definitions() -> tuple[VirtualToolDefinition, ...]:
377 return (_MCP_PROXY_SEARCH_DEFINITION, _MCP_PROXY_SCHEMA_DEFINITION, _MCP_PROXY_CALL_DEFINITION)
380def _text_tool_result(text: str, is_error: bool) -> CallToolResult:
381 from mcp.types import CallToolResult, TextContent
383 return CallToolResult(
384 content=[TextContent(type="text", text=text)], # mutable-ok: CallToolResult accepts only list content
385 is_error=is_error,
386 )
389async def handle_agent_search(query: str, top_k: int, user_api_key_dict: UserAPIKeyAuth) -> CallToolResult:
390 from litellm.proxy.agent_endpoints.agent_search import (
391 AgentSearchEmbeddingFailed,
392 AgentSearchHits,
393 AgentSearchNotConfigured,
394 agent_search_result,
395 global_agent_search_index,
396 search_agents,
397 )
398 from litellm.proxy.agent_endpoints.auth.agent_permission_handler import accessible_agents
399 from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user
400 from litellm.proxy.proxy_server import llm_router, proxy_logging_obj
402 await check_feature_access_for_user(user_api_key_dict, "agents")
403 outcome: Final = await search_agents(
404 query=query,
405 agents=await accessible_agents(user_api_key_dict),
406 top_k=max(top_k, 1),
407 router=llm_router,
408 embedding_model=litellm.agent_search_embedding_model,
409 index=global_agent_search_index,
410 user_api_key_dict=user_api_key_dict,
411 proxy_logging_obj=proxy_logging_obj,
412 )
413 match outcome:
414 case AgentSearchHits(hits):
415 results: Final = tuple(agent_search_result(hit).model_dump() for hit in hits)
416 return _text_tool_result(json.dumps(results), is_error=False)
417 case AgentSearchNotConfigured(reason) | AgentSearchEmbeddingFailed(reason):
418 return _text_tool_result(reason, is_error=True)
419 case _:
420 assert_never(outcome)
423async def handle_skill_search(query: str, top_k: int, user_api_key_dict: UserAPIKeyAuth) -> CallToolResult:
424 from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
425 from litellm.llms.litellm_proxy.skills.skill_search import (
426 MAX_SKILL_SEARCH_TOP_K,
427 SkillSearchEmbeddingFailed,
428 SkillSearchHits,
429 SkillSearchNotConfigured,
430 global_skill_search_index,
431 search_skills,
432 skill_search_result,
433 )
434 from litellm.proxy.proxy_server import llm_router, proxy_logging_obj
436 outcome: Final = await search_skills(
437 query=query,
438 skills=await LiteLLMSkillsHandler.list_skills_for_search(user_api_key_dict),
439 top_k=min(max(top_k, 1), MAX_SKILL_SEARCH_TOP_K),
440 router=llm_router,
441 embedding_model=litellm.skill_search_embedding_model,
442 index=global_skill_search_index,
443 user_api_key_dict=user_api_key_dict,
444 proxy_logging_obj=proxy_logging_obj,
445 )
446 match outcome:
447 case SkillSearchHits(hits):
448 results: Final = tuple(skill_search_result(hit).model_dump() for hit in hits)
449 return _text_tool_result(json.dumps(results), is_error=False)
450 case SkillSearchNotConfigured(reason) | SkillSearchEmbeddingFailed(reason):
451 return _text_tool_result(reason, is_error=True)
452 case _:
453 assert_never(outcome)
456async def handle_mcp_tool_search(
457 query: str,
458 top_k: int,
459 user_api_key_dict: UserAPIKeyAuth,
460 client_ip: str | None = None,
461 mcp_servers: list[str] | None = None,
462 mcp_auth_header: str | None = None,
463 mcp_server_auth_headers: dict[str, dict[str, str]] | None = None,
464 oauth2_headers: dict[str, str] | None = None,
465 raw_headers: dict[str, str] | None = None,
466) -> CallToolResult:
467 from litellm.proxy._experimental.mcp_server.operations import (
468 _list_mcp_tools,
469 )
470 from litellm.proxy.proxy_server import llm_router, proxy_logging_obj
472 settings: Final = mcp_tool_search_settings()
473 if isinstance(settings, ValidationError):
474 return _text_tool_result(
475 f"litellm_settings.{MCP_TOOL_SEARCH_SETTINGS_KEY} is invalid: {settings}", is_error=True
476 )
477 if settings.embedding_model is not None and llm_router is None:
478 return _text_tool_result(
479 f"litellm_settings.{MCP_TOOL_SEARCH_SETTINGS_KEY}.embedding_model needs a model_list so it can be called",
480 is_error=True,
481 )
482 ranker: Final = (
483 SemanticToolRanker(
484 embed=router_embedder(llm_router, settings.embedding_model, user_api_key_dict, proxy_logging_obj),
485 embedding_model=settings.embedding_model,
486 index=global_mcp_tool_search_index,
487 )
488 if settings.embedding_model is not None and llm_router is not None
489 else None
490 )
491 mcp_listing: Final = await _list_mcp_tools(
492 user_api_key_auth=user_api_key_dict,
493 mcp_servers=mcp_servers,
494 client_ip=client_ip,
495 mcp_auth_header=mcp_auth_header,
496 mcp_server_auth_headers=mcp_server_auth_headers,
497 oauth2_headers=oauth2_headers,
498 raw_headers=raw_headers,
499 )
500 results: Final = await search_mcp_tools(query, mcp_listing.tools, top_k, settings, ranker)
501 if isinstance(results, EmbeddingFailed):
502 return _text_tool_result(results.reason, is_error=True)
503 return _text_tool_result(json.dumps(results), is_error=False)
506async def handle_mcp_proxy_tool(
507 name: str,
508 arguments: dict[str, object], # mutable-ok: MCP dispatcher passes mutable call arguments
509 user_api_key_dict: UserAPIKeyAuth,
510 client_ip: str | None = None,
511 mcp_servers: list[str] | None = None, # mutable-ok: preserve MCP scope container for existing resolver
512 mcp_auth_header: str | None = None,
513 mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, # mutable-ok: preserve forwarded headers
514 oauth2_headers: dict[str, str] | None = None, # mutable-ok: preserve forwarded headers
515 raw_headers: dict[str, str] | None = None, # mutable-ok: preserve request headers
516 litellm_logging_obj: LiteLLMLoggingObj | None = None,
517) -> CallToolResult:
518 from fastapi import HTTPException
519 from jsonschema import ValidationError as JsonSchemaValidationError
520 from jsonschema import validate
522 from litellm.proxy import proxy_server
523 from litellm.proxy._experimental.mcp_server.operations import (
524 _list_mcp_tools,
525 )
527 listing: Final = await _list_mcp_tools(
528 user_api_key_auth=user_api_key_dict,
529 mcp_servers=mcp_servers,
530 client_ip=client_ip,
531 mcp_auth_header=mcp_auth_header,
532 mcp_server_auth_headers=mcp_server_auth_headers,
533 oauth2_headers=oauth2_headers,
534 raw_headers=raw_headers,
535 mcp_proxy_mode=True,
536 )
537 tools_by_id: Final = {mcp_proxy_tool_id(tool): tool for tool in listing.tools} # mutable-ok: lookup index
539 if name == MCP_PROXY_SEARCH_TOOL_NAME:
540 llm_router: Final = proxy_server.llm_router
541 proxy_logging_obj: Final = proxy_server.proxy_logging_obj
542 settings: Final = mcp_tool_search_settings()
543 if isinstance(settings, ValidationError):
544 return _text_tool_result(str(settings), is_error=True)
545 if settings.embedding_model is not None and llm_router is None:
546 return _text_tool_result(
547 f"litellm_settings.{MCP_TOOL_SEARCH_SETTINGS_KEY}.embedding_model needs a model_list so it can be called",
548 is_error=True,
549 )
550 ranker: Final = (
551 SemanticToolRanker(
552 embed=router_embedder(llm_router, settings.embedding_model, user_api_key_dict, proxy_logging_obj),
553 embedding_model=settings.embedding_model,
554 index=global_mcp_tool_search_index,
555 )
556 if settings.embedding_model is not None and llm_router is not None
557 else None
558 )
559 results: Final = await rank_mcp_tools(str(arguments.get("query", "")), listing.tools, 5, settings, ranker)
560 if isinstance(results, EmbeddingFailed):
561 return _text_tool_result(results.reason, is_error=True)
562 return _text_tool_result(json.dumps(tuple(_proxy_search_result(hit) for hit in results)), is_error=False)
564 tool_id: Final = arguments.get("tool_id")
565 tool: Final = tools_by_id.get(tool_id) if isinstance(tool_id, str) else None
566 if tool is None:
567 return _text_tool_result("Unknown or unauthorized tool_id", is_error=True)
569 if name == MCP_PROXY_SCHEMA_TOOL_NAME:
570 return _text_tool_result(json.dumps(_proxy_schema_result(tool)), is_error=False)
571 if name != MCP_PROXY_CALL_TOOL_NAME:
572 raise HTTPException(status_code=400, detail=f"Unknown MCP proxy tool: {name}")
574 tool_arguments: Final = arguments.get("arguments", {}) # mutable-ok: JSON Schema validator consumes mapping
575 if not isinstance(tool_arguments, dict):
576 return _text_tool_result("arguments must be an object", is_error=True)
577 try:
578 validate(instance=tool_arguments, schema=tool.input_schema)
579 except JsonSchemaValidationError as exc:
580 return _text_tool_result(f"Invalid arguments: {exc.message}", is_error=True)
582 return await handle_mcp_tool_call(
583 tool_name=_mcp_proxy_identity(tool)["tool_name"],
584 arguments=tool_arguments,
585 user_api_key_dict=user_api_key_dict,
586 requested_server_id=_mcp_proxy_identity(tool)["server_id"],
587 client_ip=client_ip,
588 mcp_servers=mcp_servers,
589 mcp_auth_header=mcp_auth_header,
590 mcp_server_auth_headers=mcp_server_auth_headers,
591 oauth2_headers=oauth2_headers,
592 raw_headers=raw_headers,
593 litellm_logging_obj=litellm_logging_obj,
594 )
597async def handle_mcp_tool_call(
598 tool_name: str,
599 arguments: dict[str, Any],
600 user_api_key_dict: UserAPIKeyAuth,
601 client_ip: str | None = None,
602 mcp_servers: list[str] | None = None,
603 mcp_auth_header: str | None = None,
604 mcp_server_auth_headers: dict[str, dict[str, str]] | None = None,
605 oauth2_headers: dict[str, str] | None = None,
606 raw_headers: dict[str, str] | None = None,
607 litellm_logging_obj: LiteLLMLoggingObj | None = None,
608 requested_server_id: str | None = None,
609 guardrail_context: Mapping[str, object] | None = None,
610) -> CallToolResult:
611 from litellm.proxy._experimental.mcp_server.operations import (
612 _get_allowed_mcp_servers,
613 execute_mcp_tool,
614 raise_denied_scoped_mcp_access,
615 )
617 allowed_mcp_servers: Final = await _get_allowed_mcp_servers(
618 user_api_key_auth=user_api_key_dict,
619 mcp_servers=mcp_servers,
620 client_ip=client_ip,
621 )
622 if mcp_servers and not allowed_mcp_servers:
623 await raise_denied_scoped_mcp_access(
624 requested_names=mcp_servers,
625 user_api_key_auth=user_api_key_dict,
626 client_ip=client_ip,
627 )
629 # Reject before dispatch when the key has no accessible servers; otherwise an
630 # unprefixed local tool name would fall through to the local registry in
631 # execute_mcp_tool, which has no server permission check.
632 if not allowed_mcp_servers:
633 from fastapi import HTTPException
635 raise HTTPException(status_code=403, detail="User not allowed to call this tool.")
637 return await execute_mcp_tool(
638 name=tool_name,
639 arguments=arguments,
640 allowed_mcp_servers=allowed_mcp_servers,
641 start_time=datetime.now(),
642 user_api_key_auth=user_api_key_dict,
643 mcp_auth_header=mcp_auth_header,
644 mcp_server_auth_headers=mcp_server_auth_headers,
645 oauth2_headers=oauth2_headers,
646 raw_headers=raw_headers,
647 client_ip=client_ip,
648 litellm_logging_obj=litellm_logging_obj,
649 requested_server_id=requested_server_id,
650 guardrail_context=guardrail_context,
651 )