Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/sampling_handler.py: 11%
503 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"""
2MCP Sampling Handler
3Handles `sampling/createMessage` requests from upstream MCP servers by
4routing them through LiteLLM's internal completion infrastructure.
5This allows MCP servers to perform agentic reasoning (e.g., multi-step
6tool calling, chain-of-thought) without needing their own LLM API keys —
7LiteLLM acts as the LLM provider using its existing 100+ provider support,
8cost tracking, rate limiting, and model routing.
9MCP Spec Reference:
10 https://modelcontextprotocol.io/specification/2025-11-25/client/sampling
11"""
13import typing
14from collections.abc import Mapping, Sequence
15from typing import Any, Final, NamedTuple, Optional, Protocol, Union, runtime_checkable
17if typing.TYPE_CHECKING: 17 ↛ 18line 17 didn't jump to line 18 because the condition on line 17 was never true
18 from collections.abc import Awaitable, Callable
20 from fastapi import Request
21 from mcp.client.session import ClientRequestContext
22 from mcp.types import (
23 ContentBlock,
24 CreateMessageResult,
25 CreateMessageResultWithTools,
26 ErrorData,
27 SamplingMessageContentBlock,
28 TextContent,
29 ToolUseContent,
30 )
32 from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
33 from litellm.proxy._types import UserAPIKeyAuth
34 from litellm.types.utils import ModelResponse
36from fastapi import HTTPException
37from pydantic import TypeAdapter
39from litellm._logging import verbose_logger
41# Guard imports that require the mcp package
42try:
43 from mcp.types import (
44 CreateMessageRequestParams,
45 CreateMessageResult,
46 CreateMessageResultWithTools,
47 ErrorData,
48 ModelPreferences,
49 SamplingMessage,
50 TextContent,
51 Tool,
52 ToolChoice,
53 ToolUseContent,
54 )
56 MCP_SAMPLING_AVAILABLE = True
57except ImportError as _sampling_import_err:
58 MCP_SAMPLING_AVAILABLE = False
59 verbose_logger.warning(
60 "MCP sampling disabled: failed to import required types from mcp.types — %s. "
61 "This usually means the 'mcp' package is not installed or is an older version "
62 "that does not support sampling. Install/upgrade with: pip install 'mcp>=1.1'",
63 _sampling_import_err,
64 )
67def _resolve_model_from_preferences(
68 model_preferences: Optional["ModelPreferences"],
69 default_model: str | None = None,
70) -> str:
71 """
72 Resolve an LLM model name from MCP ModelPreferences.
73 Strategy:
74 1. Check hints for substring matches against known model names.
75 2. Fall back to priority-based selection (cost/speed/intelligence).
76 3. Fall back to the configured default model.
77 Args:
78 model_preferences: MCP ModelPreferences with hints and priorities.
79 default_model: Fallback model if no hint matches.
80 Returns:
81 A model string suitable for litellm.acompletion().
82 """
83 import litellm
85 # Build list of available model names from proxy Router or litellm.model_list
86 available_model_names: list[str] = []
87 try:
88 from litellm.proxy.proxy_server import llm_router
90 if llm_router is not None:
91 available_model_names = llm_router.get_model_names()
92 except Exception:
93 pass
94 if not available_model_names and litellm.model_list:
95 for entry in litellm.model_list:
96 if isinstance(entry, dict):
97 name = entry.get("model_name")
98 if name:
99 available_model_names.append(name)
100 elif isinstance(entry, str):
101 available_model_names.append(entry)
102 if model_preferences and model_preferences.hints:
103 for hint in model_preferences.hints:
104 hint_name: str | None = getattr(hint, "name", None)
105 if not hint_name:
106 continue
107 # Try direct match first
108 if hint_name in available_model_names:
109 verbose_logger.debug(
110 "MCP sampling model resolution: direct hint match '%s'",
111 hint_name,
112 )
113 return hint_name
114 # Try substring match against known models
115 for model_name in available_model_names:
116 if hint_name.lower() in model_name.lower():
117 verbose_logger.debug(
118 "MCP sampling model resolution: substring hint match '%s' -> '%s'",
119 hint_name,
120 model_name,
121 )
122 return model_name
123 verbose_logger.debug(
124 "MCP sampling model resolution: no hint matched from %s against %d available models",
125 [getattr(h, "name", None) for h in model_preferences.hints],
126 len(available_model_names),
127 )
129 # 2. Priority-based selection (cost/speed/intelligence)
130 if model_preferences and available_model_names and _has_priorities(model_preferences):
131 best: Final = _select_model_by_priority(available_model_names, model_preferences)
132 if best is not None:
133 verbose_logger.debug(
134 "MCP sampling model resolution: priority-based selection chose '%s'",
135 best,
136 )
137 return best
139 # 3. Use default model from caller
140 if default_model:
141 verbose_logger.debug(
142 "MCP sampling model resolution: using caller-provided default '%s'",
143 default_model,
144 )
145 return default_model
146 # Fall back to first available model
147 if available_model_names:
148 verbose_logger.debug(
149 "MCP sampling model resolution: no default configured, falling back to first available model '%s'",
150 available_model_names[0],
151 )
152 return available_model_names[0]
153 # Last resort - use LiteLLM default or raise error
154 default_sampling_model: Final[str | None] = getattr(litellm, "default_mcp_sampling_model", None)
155 if default_sampling_model:
156 verbose_logger.debug(
157 "MCP sampling model resolution: using litellm.default_mcp_sampling_model='%s'",
158 default_sampling_model,
159 )
160 return default_sampling_model
161 raise ValueError(
162 "No model could be resolved for MCP sampling. Please configure 'default_mcp_sampling_model' in your LiteLLM configuration."
163 )
166def _has_priorities(model_preferences: "ModelPreferences") -> bool:
167 """Return True if any priority weight is set (non-None and > 0)."""
168 return any(
169 (getattr(model_preferences, attr, None) or 0) > 0
170 for attr in ("costPriority", "speedPriority", "intelligencePriority")
171 )
174class _ScoredModel(NamedTuple):
175 name: str
176 cost: float
177 max_output: float
178 output_tps: float
181def _select_model_by_priority(
182 model_names: list[str],
183 model_preferences: "ModelPreferences",
184) -> str | None:
185 """Score available models by MCP priority weights and return the best.
187 Scoring strategy (per the MCP spec, priorities are 0-1 floats):
189 * **costPriority** — higher means "prefer cheaper models".
190 Metric: combined (input + output) cost per token from
191 ``model_prices_and_context_window.json``. Lower cost → higher score.
193 * **speedPriority** — higher means "prefer faster models".
194 Metric: ``output_tokens_per_second`` from model info when available;
195 otherwise a neutral score for every candidate, since no reliable
196 latency proxy exists (context-window size does not track speed).
198 * **intelligencePriority** — higher means "prefer smarter models".
199 Metric: ``max_output_tokens`` is used as a rough capability proxy
200 (frontier models expose larger context windows).
202 Each metric is min-max normalised across the candidate set so that
203 every model gets a 0-1 score per dimension. The final score is the
204 weighted sum of the three normalised dimensions.
206 Returns the highest-scoring model name, or None if scoring fails for
207 all candidates (e.g. no model_info available).
208 """
209 import litellm as _litellm
211 cost_weight: Final[float] = getattr(model_preferences, "costPriority", None) or 0.0
212 speed_weight: Final[float] = getattr(model_preferences, "speedPriority", None) or 0.0
213 intel_weight: Final[float] = getattr(model_preferences, "intelligencePriority", None) or 0.0
215 # Gather raw metrics for each model
216 scored: Final[list[_ScoredModel]] = []
217 for name in model_names:
218 try:
219 info = _litellm.get_model_info(name)
220 except Exception:
221 continue
222 input_cost = info.get("input_cost_per_token") or 0.0
223 output_cost = info.get("output_cost_per_token") or 0.0
224 total_cost = input_cost + output_cost
225 max_output = info.get("max_output_tokens") or info.get("max_tokens") or 0
226 output_tps = info.get("output_tokens_per_second") or 0.0
227 scored.append(
228 _ScoredModel(
229 name=name,
230 cost=total_cost,
231 max_output=max_output,
232 output_tps=output_tps,
233 )
234 )
236 if not scored:
237 return None
239 # Min-max normalisation helpers
240 def _normalise(values: list[float], invert: bool = False) -> list[float]:
241 """Normalise to [0, 1]. If *invert*, lower raw → higher score."""
242 lo, hi = min(values), max(values)
243 if hi == lo:
244 return [0.5] * len(values) # all equal → neutral score
245 normed = [(v - lo) / (hi - lo) for v in values]
246 if invert:
247 normed = [1.0 - n for n in normed]
248 return normed
250 costs: Final = [s.cost for s in scored]
251 max_outputs: Final = [float(s.max_output) for s in scored]
252 output_tps_values: Final = [s.output_tps for s in scored]
254 # costPriority: lower cost → higher score (invert)
255 cost_scores: Final = _normalise(costs, invert=True)
256 # speedPriority: use output_tokens_per_second if any model has it,
257 # otherwise a neutral score (no reliable latency proxy is available).
258 if any(v > 0 for v in output_tps_values):
259 speed_scores = _normalise(output_tps_values, invert=False)
260 else:
261 speed_scores = [0.5] * len(scored)
262 # intelligencePriority: higher max_output → smarter
263 intel_scores: Final = _normalise(max_outputs, invert=False)
265 best_name = None
266 best_score = -1.0
267 for i, entry in enumerate(scored):
268 score = cost_weight * cost_scores[i] + speed_weight * speed_scores[i] + intel_weight * intel_scores[i]
269 verbose_logger.debug(
270 "MCP priority scoring: model=%s cost_score=%.3f speed_score=%.3f intel_score=%.3f → weighted=%.3f",
271 entry.name,
272 cost_scores[i],
273 speed_scores[i],
274 intel_scores[i],
275 score,
276 )
277 if score > best_score:
278 best_score = score
279 best_name = entry.name
281 return best_name
284def _convert_mcp_content_to_openai(
285 content: "SamplingMessageContentBlock | Sequence[SamplingMessageContentBlock]",
286) -> "str | dict[str, object] | list[dict[str, object]]":
287 """
288 Convert MCP SamplingMessage content to OpenAI message content format.
289 Handles:
290 - TextContent → string or {"type": "text", "text": ...}
291 - ImageContent → {"type": "image_url", "image_url": {"url": "data:..."}}
292 - AudioContent → {"type": "input_audio", "input_audio": {...}}
293 - ToolUseContent → function call representation
294 - ToolResultContent → tool result representation
295 - List of mixed content → list of content parts
296 """
297 if isinstance(content, list):
298 parts: Final = []
299 for item in content:
300 converted = _convert_single_content(item)
301 if isinstance(converted, list):
302 parts.extend(converted)
303 else:
304 parts.append(converted)
305 return parts
306 return _convert_single_content(content)
309@runtime_checkable
310class _TextContentLike(Protocol):
311 @property
312 def text(self) -> object: ... 312 ↛ exitline 312 didn't return from function 'text' because
315def _convert_single_content(
316 content: object,
317) -> "dict[str, object] | list[dict[str, object]]":
318 """Convert a single MCP content item to OpenAI format.
320 For text/image/audio content, returns a single content-part dict.
321 For tool_use/tool_result, returns a dict with a ``_marker_type`` key
322 so the caller (``_convert_mcp_messages_to_openai``) can hoist it to
323 the correct message-level position (``tool_calls`` array or a
324 separate ``role: "tool"`` message).
325 """
326 import json
328 content_type: Final[str | None] = getattr(content, "type", None)
329 if content_type == "text":
330 if not isinstance(content, _TextContentLike):
331 raise AttributeError(f"{type(content).__name__!r} object has no attribute 'text'")
332 return {"type": "text", "text": content.text}
333 elif content_type == "image":
334 image_data: Final[str] = getattr(content, "data", "")
335 image_mime_type: Final[str] = getattr(content, "mime_type", "image/png")
336 return {
337 "type": "image_url",
338 "image_url": {"url": f"data:{image_mime_type};base64,{image_data}"},
339 }
340 elif content_type == "audio":
341 audio_data: Final[str] = getattr(content, "data", "")
342 audio_mime_type: Final[str] = getattr(content, "mime_type", "audio/wav")
343 # Map MIME type to OpenAI audio format
344 format_map: Final = {
345 "audio/wav": "wav",
346 "audio/mp3": "mp3",
347 "audio/mpeg": "mp3",
348 "audio/flac": "flac",
349 "audio/ogg": "ogg",
350 }
351 audio_format: Final = format_map.get(audio_mime_type, "wav")
352 return {
353 "type": "input_audio",
354 "input_audio": {"data": audio_data, "format": audio_format},
355 }
356 elif content_type == "tool_use":
357 # ToolUseContent → proper OpenAI function-call representation.
358 # The ``_marker_type`` key lets the message-level converter
359 # hoist this into the ``tool_calls`` array on the assistant
360 # message instead of embedding it inline as a content part.
361 tool_use_id: Final[str] = getattr(content, "id", f"call_{id(content)}")
362 tool_name: Final[str] = getattr(content, "name", "")
363 tool_input: Final[dict[str, object]] = getattr(content, "input", {})
364 return {
365 "_marker_type": "tool_use",
366 "id": tool_use_id,
367 "type": "function",
368 "function": {
369 "name": tool_name,
370 "arguments": json.dumps(tool_input, default=str),
371 },
372 }
373 elif content_type == "tool_result":
374 # ToolResultContent → proper OpenAI tool-role message.
375 # Marked so the message-level converter can emit it as a
376 # separate ``{"role": "tool", ...}`` message.
377 tool_result_use_id: Final = getattr(content, "tool_use_id", "")
378 nested_content: Final[Sequence[ContentBlock]] = getattr(content, "content", [])
379 if isinstance(nested_content, list):
380 text_parts = [getattr(c, "text", str(c)) for c in nested_content if getattr(c, "type", None) == "text"]
381 result_text = "\n".join(text_parts) if text_parts else ""
382 else:
383 result_text = str(nested_content)
384 return {
385 "_marker_type": "tool_result",
386 "role": "tool",
387 "tool_call_id": tool_result_use_id,
388 "content": result_text,
389 }
390 # Fallback: treat as text
391 return {"type": "text", "text": str(content)}
394def _convert_mcp_messages_to_openai(
395 messages: list["SamplingMessage"],
396 system_prompt: str | None = None,
397) -> "Sequence[Mapping[str, object]]":
398 """
399 Convert MCP SamplingMessage list to OpenAI messages format.
400 MCP messages use:
401 - role: "user" | "assistant"
402 - content: TextContent | ImageContent | AudioContent | ToolUseContent
403 | ToolResultContent | list[...]
404 OpenAI messages use:
405 - role: "system" | "user" | "assistant" | "tool"
406 - content: str | list[content_part]
407 """
408 openai_messages: Final[list[Mapping[str, object]]] = []
409 # Add system prompt if provided
410 if system_prompt:
411 openai_messages.append({"role": "system", "content": system_prompt})
412 for msg in messages:
413 role = msg.role
414 content = msg.content
415 # Handle tool use content from assistant
416 if role == "assistant" and _has_tool_use(content):
417 tool_calls = _extract_tool_calls(content)
418 if tool_calls:
419 openai_msg: dict[str, object] = {
420 "role": "assistant",
421 "tool_calls": tool_calls,
422 }
423 # Also include any text content alongside tool calls
424 text_parts = _extract_text_parts(content)
425 if text_parts:
426 openai_msg["content"] = text_parts
427 openai_messages.append(openai_msg)
428 continue
429 # Handle tool result content from user
430 if role == "user" and _has_tool_result(content):
431 tool_results = _extract_tool_results(content)
432 for tool_result in tool_results:
433 openai_messages.append(tool_result)
434 continue
435 # Standard text/image/audio message — also handles any stray
436 # tool_use / tool_result that slipped past the fast-path checks
437 # above (e.g. unexpected role, single non-list content).
438 converted = _convert_mcp_content_to_openai(content)
439 converted_parts: Sequence[Mapping[str, object]] = (
440 converted if isinstance(converted, list) else ([converted] if isinstance(converted, dict) else [])
441 )
443 # Separate marker items from regular content parts
444 tool_call_markers = []
445 tool_result_markers = []
446 regular_parts = []
447 for part in converted_parts:
448 marker = part.get("_marker_type") if isinstance(part, dict) else None
449 if marker == "tool_use":
450 # Strip the internal marker before emitting
451 tc = {k: v for k, v in part.items() if k != "_marker_type"}
452 tool_call_markers.append(tc)
453 elif marker == "tool_result":
454 tr = {k: v for k, v in part.items() if k != "_marker_type"}
455 tool_result_markers.append(tr)
456 else:
457 regular_parts.append(part)
459 # Emit assistant message with tool_calls if any were found
460 if tool_call_markers:
461 openai_msg_tc: dict[str, object] = {
462 "role": "assistant",
463 "tool_calls": tool_call_markers,
464 }
465 if regular_parts:
466 openai_msg_tc["content"] = regular_parts
467 openai_messages.append(openai_msg_tc)
468 elif regular_parts:
469 if isinstance(converted, str):
470 openai_messages.append({"role": role, "content": converted})
471 else:
472 openai_messages.append({"role": role, "content": regular_parts})
474 # Emit separate tool-result messages
475 for tr in tool_result_markers:
476 openai_messages.append(tr)
478 return openai_messages
481def _has_tool_use(content: "SamplingMessageContentBlock | Sequence[SamplingMessageContentBlock]") -> bool:
482 """Check if content contains ToolUseContent."""
483 if isinstance(content, list):
484 return any(getattr(c, "type", None) == "tool_use" for c in content)
485 content_type: Final[str | None] = getattr(content, "type", None)
486 return content_type == "tool_use"
489def _has_tool_result(content: "SamplingMessageContentBlock | Sequence[SamplingMessageContentBlock]") -> bool:
490 """Check if content contains ToolResultContent."""
491 if isinstance(content, list):
492 return any(getattr(c, "type", None) == "tool_result" for c in content)
493 content_type: Final[str | None] = getattr(content, "type", None)
494 return content_type == "tool_result"
497def _extract_tool_calls(
498 content: "SamplingMessageContentBlock | Sequence[SamplingMessageContentBlock]",
499) -> "Sequence[Mapping[str, object]]":
500 """Extract OpenAI-format tool_calls from MCP ToolUseContent."""
501 import json
503 items: Final = content if isinstance(content, list) else [content]
504 tool_calls: Final = []
505 for item in items:
506 if getattr(item, "type", None) == "tool_use":
507 tool_calls.append(
508 {
509 "id": getattr(item, "id", f"call_{id(item)}"),
510 "type": "function",
511 "function": {
512 "name": getattr(item, "name", ""),
513 "arguments": json.dumps(getattr(item, "input", {}), default=str),
514 },
515 }
516 )
517 return tool_calls
520def _extract_text_parts(
521 content: "SamplingMessageContentBlock | Sequence[SamplingMessageContentBlock]",
522) -> str | None:
523 """Extract text parts from mixed content."""
524 items: Final = content if isinstance(content, list) else [content]
525 texts: Final = []
526 for item in items:
527 if getattr(item, "type", None) == "text":
528 texts.append(getattr(item, "text", ""))
529 return "\n".join(texts) if texts else None
532def _extract_tool_results(
533 content: "SamplingMessageContentBlock | Sequence[SamplingMessageContentBlock]",
534) -> "Sequence[Mapping[str, object]]":
535 """Extract OpenAI-format tool messages from MCP ToolResultContent."""
536 items: Final = content if isinstance(content, list) else [content]
537 results: Final = []
538 for item in items:
539 if getattr(item, "type", None) == "tool_result":
540 tool_use_id = getattr(item, "tool_use_id", "")
541 # Extract text from nested content
542 nested_content: Sequence[ContentBlock] = getattr(item, "content", [])
543 if isinstance(nested_content, list):
544 text_parts = [getattr(c, "text", str(c)) for c in nested_content if getattr(c, "type", None) == "text"]
545 result_text = "\n".join(text_parts) if text_parts else ""
546 else:
547 result_text = str(nested_content)
548 results.append(
549 {
550 "role": "tool",
551 "tool_call_id": tool_use_id,
552 "content": result_text,
553 }
554 )
555 return results
558def _convert_mcp_tools_to_openai(
559 tools: list["Tool"] | None,
560) -> "Sequence[Mapping[str, object]] | None":
561 """
562 Convert MCP Tool definitions to OpenAI function calling format.
563 MCP Tool: {name, description, inputSchema}
564 OpenAI Tool: {type: "function", function: {name, description, parameters}}
565 """
566 if not tools:
567 return None
568 openai_tools: Final = []
569 for tool in tools:
570 openai_tool = {
571 "type": "function",
572 "function": {
573 "name": tool.name,
574 "description": tool.description or "",
575 "parameters": tool.input_schema
576 or {
577 "type": "object",
578 "properties": {},
579 },
580 },
581 }
582 openai_tools.append(openai_tool)
583 return openai_tools
586def _convert_mcp_tool_choice_to_openai(
587 tool_choice: Optional["ToolChoice"],
588) -> "str | None":
589 """
590 Convert MCP ToolChoice to OpenAI tool_choice format.
591 MCP: {mode: "auto"} | {mode: "required"} | {mode: "none"}
592 OpenAI: "auto" | "required" | "none"
593 """
594 if not tool_choice:
595 return None
596 mode: Final = getattr(tool_choice, "mode", "auto")
597 if mode == "auto":
598 return "auto"
599 elif mode == "required":
600 return "required"
601 elif mode == "none":
602 return "none"
603 return "auto"
606class _SamplingToolCallFunction(Protocol):
607 @property
608 def name(self) -> str | None: ... 608 ↛ exitline 608 didn't return from function 'name' because
610 @property
611 def arguments(self) -> object: ... 611 ↛ exitline 611 didn't return from function 'arguments' because
614class _SamplingToolCall(Protocol):
615 @property
616 def id(self) -> str | None: ... 616 ↛ exitline 616 didn't return from function 'id' because
618 @property
619 def function(self) -> _SamplingToolCallFunction: ... 619 ↛ exitline 619 didn't return from function 'function' because
622class _SamplingResponseMessage(Protocol):
623 @property
624 def content(self) -> str | None: ... 624 ↛ exitline 624 didn't return from function 'content' because
626 @property
627 def tool_calls(self) -> Sequence[_SamplingToolCall] | None: ... 627 ↛ exitline 627 didn't return from function 'tool_calls' because
630class _SamplingResponseChoice(Protocol):
631 @property
632 def message(self) -> _SamplingResponseMessage: ... 632 ↛ exitline 632 didn't return from function 'message' because
634 @property
635 def finish_reason(self) -> str | None: ... 635 ↛ exitline 635 didn't return from function 'finish_reason' because
638class _SamplingCompletionResponse(Protocol):
639 @property
640 def choices(self) -> Sequence[_SamplingResponseChoice]: ... 640 ↛ exitline 640 didn't return from function 'choices' because
642 @property
643 def model(self) -> str | None: ... 643 ↛ exitline 643 didn't return from function 'model' because
646_TOOL_ARGUMENTS_ADAPTER: Final = TypeAdapter(dict[str, object])
649def _parse_tool_arguments(arguments: object) -> "dict[str, object]":
650 """Decode OpenAI tool-call arguments into the MCP ``input`` mapping."""
651 import json
653 if not isinstance(arguments, str):
654 return _TOOL_ARGUMENTS_ADAPTER.validate_python(arguments)
655 try:
656 return _TOOL_ARGUMENTS_ADAPTER.validate_python(json.loads(arguments))
657 except (json.JSONDecodeError, TypeError):
658 return {"raw": arguments}
661def _convert_openai_response_to_mcp_result(
662 response: _SamplingCompletionResponse,
663 model_name: str,
664) -> Union["CreateMessageResult", "CreateMessageResultWithTools", "ErrorData"]:
665 """
666 Convert a litellm completion response to MCP CreateMessageResult.
667 Args:
668 response: The litellm ModelResponse.
669 model_name: The model that was used.
670 Returns:
671 MCP CreateMessageResult or CreateMessageResultWithTools.
672 """
673 if not response.choices:
674 verbose_logger.warning(
675 "MCP sampling: LLM returned empty choices list for model=%s (possible content filter or provider error)",
676 model_name,
677 )
678 return ErrorData(
679 code=-1,
680 message=(
681 f"LLM returned no choices for model '{model_name}'. "
682 "This may indicate content filtering or a provider-side error."
683 ),
684 )
685 choice: Final = response.choices[0]
686 message: Final = choice.message
687 # Determine stop reason
688 finish_reason: Final = getattr(choice, "finish_reason", "stop")
689 if finish_reason == "tool_calls":
690 stop_reason = "toolUse"
691 elif finish_reason == "length":
692 stop_reason = "maxTokens"
693 else:
694 stop_reason = "endTurn"
695 actual_model: Final[str] = getattr(response, "model", model_name) or model_name
696 # Check if response has tool calls
697 tool_calls: Final = message.tool_calls if hasattr(message, "tool_calls") else None
698 if tool_calls:
699 # Build ToolUseContent items
700 content_parts: Final[list[SamplingMessageContentBlock]] = []
701 # Include text content if present
702 if message.content:
703 content_parts.append(TextContent(type="text", text=message.content))
704 # Convert tool calls to MCP ToolUseContent
705 for tc in tool_calls:
706 content_parts.append(
707 ToolUseContent.model_validate(
708 {
709 "type": "tool_use",
710 "id": tc.id,
711 "name": tc.function.name,
712 "input": _parse_tool_arguments(tc.function.arguments),
713 }
714 )
715 )
716 return CreateMessageResultWithTools(
717 role="assistant",
718 content=content_parts,
719 model=actual_model,
720 stop_reason=stop_reason,
721 )
722 # Simple text response
723 text: Final = message.content or ""
724 return CreateMessageResult(
725 role="assistant",
726 content=TextContent(type="text", text=text),
727 model=actual_model,
728 stop_reason=stop_reason,
729 )
732async def _check_model_access(model: str, user_api_key_auth: "UserAPIKeyAuth | None") -> Optional["ErrorData"]:
733 """Enforce model-permission checks for MCP sampling requests.
735 Runs the same authorization checks as ``/chat/completions``:
736 key-level, team-level, per-member, user-level, and project-level
737 model restrictions. The model name comes from the upstream MCP
738 server (untrusted input).
740 Returns None if authorized, or an ErrorData describing the denial.
741 """
742 if user_api_key_auth is None:
743 return None
745 _api_key: Final = getattr(user_api_key_auth, "api_key", None)
746 _token: Final = getattr(user_api_key_auth, "token", None)
747 _user_role: Final = getattr(user_api_key_auth, "user_role", None)
749 _has_real_credential: Final = bool(_api_key) or bool(_token)
750 _is_admin: Final = _user_role in ("proxy_admin", "proxy_admin_viewer") if _user_role else False
752 if not _has_real_credential and not _is_admin:
753 verbose_logger.warning(
754 "MCP sampling: denying model access for model=%s — "
755 "auth context has no real LiteLLM credential (possible "
756 "OAuth passthrough placeholder). api_key=%s, token=%s, role=%s",
757 model,
758 bool(_api_key),
759 bool(_token),
760 _user_role,
761 )
762 return ErrorData(
763 code=-1,
764 message=(
765 "Model access denied: sampling requires a valid LiteLLM "
766 "API key or admin credential. OAuth-only sessions cannot "
767 "trigger proxy model calls without explicit authorization."
768 ),
769 )
771 try:
772 import litellm
773 from litellm.proxy._types import ModelAccessDeniedProxyException
774 from litellm.proxy.auth.auth_checks import (
775 _check_team_member_model_access,
776 can_key_call_model,
777 can_project_access_model,
778 can_team_access_model,
779 can_user_call_model,
780 get_project_object,
781 get_team_object,
782 get_user_object,
783 )
785 try:
786 from litellm.proxy.proxy_server import llm_router as _llm_router
787 except ImportError:
788 _llm_router = None
790 await can_key_call_model(
791 model=model,
792 llm_model_list=getattr(litellm, "model_list", None),
793 valid_token=user_api_key_auth,
794 llm_router=_llm_router,
795 )
797 _team_id: Final[str | None] = getattr(user_api_key_auth, "team_id", None)
798 _user_id: Final[str | None] = getattr(user_api_key_auth, "user_id", None)
799 _project_id: Final[str | None] = getattr(user_api_key_auth, "project_id", None)
801 try:
802 from litellm.proxy.proxy_server import (
803 prisma_client as _prisma_client,
804 )
805 from litellm.proxy.proxy_server import (
806 proxy_logging_obj as _proxy_logging_obj,
807 )
808 from litellm.proxy.proxy_server import (
809 user_api_key_cache as _user_api_key_cache,
810 )
811 except ImportError:
812 _prisma_client = None
813 _user_api_key_cache = None
814 _proxy_logging_obj = None
816 if _team_id and _prisma_client and _user_api_key_cache:
817 try:
818 team_obj = await get_team_object(
819 team_id=_team_id,
820 prisma_client=_prisma_client,
821 user_api_key_cache=_user_api_key_cache,
822 proxy_logging_obj=_proxy_logging_obj,
823 )
824 except Exception:
825 team_obj = None
827 if team_obj:
828 await can_team_access_model(
829 model=model,
830 team_object=team_obj,
831 llm_router=_llm_router,
832 team_model_aliases=getattr(user_api_key_auth, "team_model_aliases", None),
833 )
834 if _user_id and _proxy_logging_obj:
835 await _check_team_member_model_access(
836 model=model,
837 team_object=team_obj,
838 valid_token=user_api_key_auth,
839 llm_router=_llm_router,
840 prisma_client=_prisma_client,
841 user_api_key_cache=_user_api_key_cache,
842 proxy_logging_obj=_proxy_logging_obj,
843 )
844 elif not _team_id and _user_id and _prisma_client and _user_api_key_cache:
845 try:
846 user_obj = await get_user_object(
847 user_id=_user_id,
848 prisma_client=_prisma_client,
849 user_api_key_cache=_user_api_key_cache,
850 user_id_upsert=False,
851 proxy_logging_obj=_proxy_logging_obj,
852 )
853 except Exception:
854 user_obj = None
856 if user_obj:
857 await can_user_call_model(
858 model=model,
859 llm_router=_llm_router,
860 user_object=user_obj,
861 )
863 if _project_id and _prisma_client and _user_api_key_cache:
864 try:
865 project_obj = await get_project_object(
866 project_id=_project_id,
867 prisma_client=_prisma_client,
868 user_api_key_cache=_user_api_key_cache,
869 proxy_logging_obj=_proxy_logging_obj,
870 )
871 except Exception:
872 project_obj = None
874 if project_obj:
875 can_project_access_model(
876 model=model,
877 project_object=project_obj,
878 llm_router=_llm_router,
879 )
881 verbose_logger.debug(
882 "MCP sampling: model access check passed for model=%s",
883 model,
884 )
885 return None
886 except Exception as access_err:
887 if isinstance(access_err, ModelAccessDeniedProxyException):
888 verbose_logger.warning(
889 "MCP sampling: model access denied for model=%s: %s",
890 model,
891 access_err.sanitized_internal_message(),
892 )
893 return ErrorData(code=-1, message=access_err.message)
894 verbose_logger.warning("MCP sampling: model access denied for model=%s: %s", model, access_err)
895 return ErrorData(
896 code=-1,
897 message=(f"Model access denied: the API key is not authorized to use model '{model}'. {access_err}"),
898 )
901async def _run_budget_checks(
902 model: str,
903 user_api_key_auth: "UserAPIKeyAuth",
904 raw_headers: dict[str, str] | None = None,
905 client_ip: str | None = None,
906) -> Optional["ErrorData"]:
907 """Enforce key/team/user/org/global budget checks for sampling requests.
909 Runs the same ``common_checks`` path that ``/chat/completions`` uses,
910 so sampling cannot bypass budget limits.
912 Returns None if all checks pass, or an ErrorData describing the denial.
913 """
914 try:
915 import litellm
916 from litellm.proxy.auth.auth_checks import (
917 common_checks,
918 get_team_object,
919 get_user_object,
920 )
921 from litellm.proxy.proxy_server import (
922 general_settings,
923 )
924 from litellm.proxy.proxy_server import (
925 llm_router as _llm_router,
926 )
927 from litellm.proxy.proxy_server import (
928 prisma_client as _prisma_client,
929 )
930 from litellm.proxy.proxy_server import (
931 proxy_logging_obj as _proxy_logging_obj,
932 )
933 from litellm.proxy.proxy_server import (
934 user_api_key_cache as _user_api_key_cache,
935 )
936 except ImportError as import_err:
937 verbose_logger.warning("MCP sampling: budget check imports unavailable: %s", import_err)
938 return None # Can't enforce budgets without the modules
940 _team_id: Final[str | None] = getattr(user_api_key_auth, "team_id", None)
941 _user_id: Final[str | None] = getattr(user_api_key_auth, "user_id", None)
943 team_obj = None
944 if _team_id and _prisma_client and _user_api_key_cache:
945 try:
946 team_obj = await get_team_object(
947 team_id=_team_id,
948 prisma_client=_prisma_client,
949 user_api_key_cache=_user_api_key_cache,
950 proxy_logging_obj=_proxy_logging_obj,
951 )
952 except Exception:
953 pass
955 user_obj = None
956 if _user_id and _prisma_client and _user_api_key_cache:
957 try:
958 user_obj = await get_user_object(
959 user_id=_user_id,
960 prisma_client=_prisma_client,
961 user_api_key_cache=_user_api_key_cache,
962 user_id_upsert=False,
963 proxy_logging_obj=_proxy_logging_obj,
964 )
965 except Exception:
966 pass
968 dummy_request: Final = _build_sampling_request(
969 raw_headers=raw_headers,
970 client_ip=client_ip,
971 )
973 # Enforce virtual-key route restrictions: a key limited to MCP routes
974 # must not be able to trigger a /chat/completions call via sampling.
975 # This mirrors the RouteChecks.should_call_route gate that runs in
976 # user_api_key_auth before common_checks for regular requests.
977 try:
978 from litellm.proxy.auth.route_checks import RouteChecks
980 RouteChecks.should_call_route(
981 route="/chat/completions",
982 valid_token=user_api_key_auth,
983 request=dummy_request,
984 )
985 except HTTPException as route_err:
986 verbose_logger.warning(
987 "MCP sampling: route check denied /chat/completions for key: %s",
988 route_err.detail,
989 )
990 return ErrorData(
991 code=-1,
992 message=f"Sampling denied: virtual key is not allowed to call /chat/completions. {route_err.detail}",
993 )
995 global_proxy_spend: Final = getattr(litellm, "_global_proxy_spend", None)
997 # Build request body and merge x-litellm-tags from MCP headers BEFORE
998 # common_checks runs. _tag_max_budget_check inside common_checks only
999 # inspects request_body; without this pre-merge, header-supplied tags
1000 # bypass per-tag budget enforcement (mirroring the regular auth path).
1001 request_body: Final[dict[str, object]] = {"model": model}
1002 try:
1003 from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
1005 LiteLLMProxyRequestSetup.apply_client_tag_policy_pre_auth(
1006 request=dummy_request,
1007 request_data=request_body,
1008 user_api_key_dict=user_api_key_auth,
1009 )
1010 except Exception:
1011 # Non-fatal: tag merge is defense-in-depth; don't block sampling
1012 # if the merge utility is unavailable or fails.
1013 pass
1015 try:
1016 await common_checks(
1017 request_body=request_body,
1018 team_object=team_obj,
1019 user_object=user_obj,
1020 end_user_object=None,
1021 global_proxy_spend=global_proxy_spend,
1022 general_settings=general_settings or {},
1023 route="/chat/completions",
1024 llm_router=_llm_router,
1025 proxy_logging_obj=_proxy_logging_obj,
1026 valid_token=user_api_key_auth,
1027 request=dummy_request,
1028 )
1029 except Exception as budget_err:
1030 verbose_logger.warning(
1031 "MCP sampling: budget check failed for model=%s: %s",
1032 model,
1033 budget_err,
1034 )
1035 return ErrorData(
1036 code=-1,
1037 message=f"Sampling denied: {budget_err}",
1038 )
1040 verbose_logger.debug("MCP sampling: budget checks passed for model=%s", model)
1041 return None
1044def _build_sampling_request(
1045 raw_headers: dict[str, str] | None = None,
1046 client_ip: str | None = None,
1047) -> "Request":
1048 """The synthetic FastAPI Request for sampling sub-calls, carrying the original
1049 MCP connection's headers and client IP."""
1050 from litellm.proxy._experimental.mcp_server.utils import build_synthetic_mcp_request
1052 return build_synthetic_mcp_request(
1053 path="/mcp/sampling/createMessage",
1054 raw_headers=raw_headers,
1055 client_ip=client_ip,
1056 )
1059async def _build_completion_kwargs(
1060 params: "CreateMessageRequestParams",
1061 model: str,
1062 user_api_key_auth: "UserAPIKeyAuth",
1063 raw_headers: dict[str, str] | None,
1064 client_ip: str | None,
1065) -> dict[str, Any]:
1066 openai_messages: Final = _convert_mcp_messages_to_openai(
1067 messages=params.messages,
1068 system_prompt=params.system_prompt,
1069 )
1070 completion_kwargs: Final[dict[str, object]] = {
1071 "model": model,
1072 "messages": openai_messages,
1073 "max_tokens": params.max_tokens,
1074 }
1075 if params.temperature is not None:
1076 completion_kwargs["temperature"] = params.temperature
1077 if params.stop_sequences:
1078 completion_kwargs["stop"] = params.stop_sequences
1079 openai_tools: Final = _convert_mcp_tools_to_openai(params.tools)
1080 if openai_tools:
1081 completion_kwargs["tools"] = openai_tools
1082 openai_tool_choice: Final = _convert_mcp_tool_choice_to_openai(params.tool_choice)
1083 if openai_tool_choice is not None:
1084 completion_kwargs["tool_choice"] = openai_tool_choice
1085 completion_kwargs["metadata"] = {"mcp_metadata": params.metadata} if params.metadata else {}
1087 from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
1088 from litellm.proxy.proxy_server import proxy_config
1090 completion_kwargs["user"] = getattr(user_api_key_auth, "user_id", None)
1091 _dummy_request: Final = _build_sampling_request(raw_headers=raw_headers, client_ip=client_ip)
1092 return await add_litellm_data_to_request(
1093 data=completion_kwargs,
1094 request=_dummy_request,
1095 user_api_key_dict=user_api_key_auth,
1096 proxy_config=proxy_config,
1097 )
1100class _AcompletionCall(NamedTuple):
1101 fn: "Callable[..., Awaitable[ModelResponse | CustomStreamWrapper]]"
1104async def _run_guardrails_and_call_llm(
1105 completion_kwargs: dict[str, object],
1106 user_api_key_auth: "UserAPIKeyAuth",
1107) -> Any:
1108 try:
1109 from litellm.proxy.proxy_server import proxy_logging_obj as _plo
1111 if _plo is not None:
1112 completion_kwargs = await _plo.pre_call_hook(
1113 user_api_key_dict=user_api_key_auth,
1114 data=completion_kwargs,
1115 call_type="acompletion",
1116 )
1117 except ImportError:
1118 pass
1119 except Exception as guardrail_err:
1120 verbose_logger.warning(
1121 "MCP sampling: pre-call guardrail rejected request: %s",
1122 guardrail_err,
1123 )
1124 raise
1126 import litellm
1128 try:
1129 from litellm.proxy.proxy_server import llm_router
1131 if llm_router is not None:
1132 return await _AcompletionCall(fn=llm_router.acompletion).fn(**completion_kwargs)
1133 return await _AcompletionCall(fn=litellm.acompletion).fn(**completion_kwargs)
1134 except ImportError:
1135 return await _AcompletionCall(fn=litellm.acompletion).fn(**completion_kwargs)
1138async def handle_sampling_create_message(
1139 context: "ClientRequestContext",
1140 params: "CreateMessageRequestParams",
1141 default_model: str | None = None,
1142 user_api_key_auth: "UserAPIKeyAuth | None" = None,
1143 raw_headers: dict[str, str] | None = None,
1144 client_ip: str | None = None,
1145) -> Union["CreateMessageResult", "CreateMessageResultWithTools", "ErrorData"]:
1146 """
1147 Handle an MCP sampling/createMessage request by routing through LiteLLM.
1148 This is the main entry point called by the MCP client session when an
1149 upstream MCP server requests LLM inference.
1150 Args:
1151 context: MCP RequestContext (contains session info).
1152 params: The CreateMessageRequestParams from the MCP server.
1153 default_model: Default model to use if no preferences match.
1154 user_api_key_auth: Auth context for the requesting user.
1155 raw_headers: Original HTTP headers from the MCP connection.
1156 Forwarded into the internal acompletion call so that
1157 header-dependent guardrails, IP-routing, trace-id
1158 correlation, and forward_llm_provider_auth_headers
1159 work correctly for sampling sub-calls.
1160 client_ip: Original client IP address for IP-based guardrails.
1161 Returns:
1162 CreateMessageResult with the LLM's response, or ErrorData on failure.
1163 """
1164 if not MCP_SAMPLING_AVAILABLE:
1165 return ErrorData(
1166 code=-1,
1167 message="MCP sampling is not available (mcp package not installed)",
1168 )
1170 if user_api_key_auth is None:
1171 return ErrorData(
1172 code=-1,
1173 message=(
1174 "Sampling requires an authenticated user context. "
1175 "Internal or unauthenticated sessions cannot trigger "
1176 "upstream-initiated model calls."
1177 ),
1178 )
1180 try:
1181 model: Final = _resolve_model_from_preferences(
1182 model_preferences=params.model_preferences,
1183 default_model=default_model,
1184 )
1185 verbose_logger.info(
1186 "MCP sampling: resolved model=%s from preferences=%s",
1187 model,
1188 params.model_preferences,
1189 )
1191 access_denial: Final = await _check_model_access(model, user_api_key_auth)
1192 if access_denial is not None:
1193 return access_denial
1195 budget_denial: Final = await _run_budget_checks(
1196 model=model,
1197 user_api_key_auth=user_api_key_auth,
1198 raw_headers=raw_headers,
1199 client_ip=client_ip,
1200 )
1201 if budget_denial is not None:
1202 return budget_denial
1204 completion_kwargs: Final = await _build_completion_kwargs(
1205 params=params,
1206 model=model,
1207 user_api_key_auth=user_api_key_auth,
1208 raw_headers=raw_headers,
1209 client_ip=client_ip,
1210 )
1212 openai_messages: Final[Sequence[Mapping[str, object]]] = completion_kwargs["messages"]
1213 openai_tools: Final = completion_kwargs.get("tools")
1214 verbose_logger.debug(
1215 "MCP sampling: calling litellm.acompletion with model=%s, num_messages=%d, has_tools=%s",
1216 model,
1217 len(openai_messages),
1218 bool(openai_tools),
1219 )
1221 response: Final[_SamplingCompletionResponse] = await _run_guardrails_and_call_llm(
1222 completion_kwargs=completion_kwargs,
1223 user_api_key_auth=user_api_key_auth,
1224 )
1226 result: Final = _convert_openai_response_to_mcp_result(response=response, model_name=model)
1227 verbose_logger.info(
1228 "MCP sampling: completed successfully, model=%s, stopReason=%s",
1229 getattr(result, "model", "unknown"),
1230 getattr(result, "stop_reason", "unknown"),
1231 )
1232 return result
1233 except Exception as e:
1234 from litellm.exceptions import (
1235 AuthenticationError,
1236 BudgetExceededError,
1237 ContextWindowExceededError,
1238 PermissionDeniedError,
1239 RateLimitError,
1240 ServiceUnavailableError,
1241 )
1242 from litellm.proxy._types import ProxyException
1244 if isinstance(
1245 e,
1246 (
1247 HTTPException,
1248 BudgetExceededError,
1249 RateLimitError,
1250 AuthenticationError,
1251 PermissionDeniedError,
1252 ContextWindowExceededError,
1253 ServiceUnavailableError,
1254 ProxyException,
1255 ),
1256 ):
1257 raise
1259 verbose_logger.exception("MCP sampling handler failed: %s", e)
1260 return ErrorData(
1261 code=-1,
1262 message=f"Sampling failed: {e}",
1263 )