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

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

12 

13import typing 

14from collections.abc import Mapping, Sequence 

15from typing import Any, Final, NamedTuple, Optional, Protocol, Union, runtime_checkable 

16 

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 

19 

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 ) 

31 

32 from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper 

33 from litellm.proxy._types import UserAPIKeyAuth 

34 from litellm.types.utils import ModelResponse 

35 

36from fastapi import HTTPException 

37from pydantic import TypeAdapter 

38 

39from litellm._logging import verbose_logger 

40 

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 ) 

55 

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 ) 

65 

66 

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 

84 

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 

89 

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 ) 

128 

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 

138 

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 ) 

164 

165 

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 ) 

172 

173 

174class _ScoredModel(NamedTuple): 

175 name: str 

176 cost: float 

177 max_output: float 

178 output_tps: float 

179 

180 

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. 

186 

187 Scoring strategy (per the MCP spec, priorities are 0-1 floats): 

188 

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. 

192 

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

197 

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

201 

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. 

205 

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 

210 

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 

214 

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 ) 

235 

236 if not scored: 

237 return None 

238 

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 

249 

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] 

253 

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) 

264 

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 

280 

281 return best_name 

282 

283 

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) 

307 

308 

309@runtime_checkable 

310class _TextContentLike(Protocol): 

311 @property 

312 def text(self) -> object: ... 312 ↛ exitline 312 didn't return from function 'text' because

313 

314 

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. 

319 

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 

327 

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)} 

392 

393 

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 ) 

442 

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) 

458 

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}) 

473 

474 # Emit separate tool-result messages 

475 for tr in tool_result_markers: 

476 openai_messages.append(tr) 

477 

478 return openai_messages 

479 

480 

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" 

487 

488 

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" 

495 

496 

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 

502 

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 

518 

519 

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 

530 

531 

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 

556 

557 

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 

584 

585 

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" 

604 

605 

606class _SamplingToolCallFunction(Protocol): 

607 @property 

608 def name(self) -> str | None: ... 608 ↛ exitline 608 didn't return from function 'name' because

609 

610 @property 

611 def arguments(self) -> object: ... 611 ↛ exitline 611 didn't return from function 'arguments' because

612 

613 

614class _SamplingToolCall(Protocol): 

615 @property 

616 def id(self) -> str | None: ... 616 ↛ exitline 616 didn't return from function 'id' because

617 

618 @property 

619 def function(self) -> _SamplingToolCallFunction: ... 619 ↛ exitline 619 didn't return from function 'function' because

620 

621 

622class _SamplingResponseMessage(Protocol): 

623 @property 

624 def content(self) -> str | None: ... 624 ↛ exitline 624 didn't return from function 'content' because

625 

626 @property 

627 def tool_calls(self) -> Sequence[_SamplingToolCall] | None: ... 627 ↛ exitline 627 didn't return from function 'tool_calls' because

628 

629 

630class _SamplingResponseChoice(Protocol): 

631 @property 

632 def message(self) -> _SamplingResponseMessage: ... 632 ↛ exitline 632 didn't return from function 'message' because

633 

634 @property 

635 def finish_reason(self) -> str | None: ... 635 ↛ exitline 635 didn't return from function 'finish_reason' because

636 

637 

638class _SamplingCompletionResponse(Protocol): 

639 @property 

640 def choices(self) -> Sequence[_SamplingResponseChoice]: ... 640 ↛ exitline 640 didn't return from function 'choices' because

641 

642 @property 

643 def model(self) -> str | None: ... 643 ↛ exitline 643 didn't return from function 'model' because

644 

645 

646_TOOL_ARGUMENTS_ADAPTER: Final = TypeAdapter(dict[str, object]) 

647 

648 

649def _parse_tool_arguments(arguments: object) -> "dict[str, object]": 

650 """Decode OpenAI tool-call arguments into the MCP ``input`` mapping.""" 

651 import json 

652 

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} 

659 

660 

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 ) 

730 

731 

732async def _check_model_access(model: str, user_api_key_auth: "UserAPIKeyAuth | None") -> Optional["ErrorData"]: 

733 """Enforce model-permission checks for MCP sampling requests. 

734 

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

739 

740 Returns None if authorized, or an ErrorData describing the denial. 

741 """ 

742 if user_api_key_auth is None: 

743 return None 

744 

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) 

748 

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 

751 

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 ) 

770 

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 ) 

784 

785 try: 

786 from litellm.proxy.proxy_server import llm_router as _llm_router 

787 except ImportError: 

788 _llm_router = None 

789 

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 ) 

796 

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) 

800 

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 

815 

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 

826 

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 

855 

856 if user_obj: 

857 await can_user_call_model( 

858 model=model, 

859 llm_router=_llm_router, 

860 user_object=user_obj, 

861 ) 

862 

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 

873 

874 if project_obj: 

875 can_project_access_model( 

876 model=model, 

877 project_object=project_obj, 

878 llm_router=_llm_router, 

879 ) 

880 

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 ) 

899 

900 

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. 

908 

909 Runs the same ``common_checks`` path that ``/chat/completions`` uses, 

910 so sampling cannot bypass budget limits. 

911 

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 

939 

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) 

942 

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 

954 

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 

967 

968 dummy_request: Final = _build_sampling_request( 

969 raw_headers=raw_headers, 

970 client_ip=client_ip, 

971 ) 

972 

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 

979 

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 ) 

994 

995 global_proxy_spend: Final = getattr(litellm, "_global_proxy_spend", None) 

996 

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 

1004 

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 

1014 

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 ) 

1039 

1040 verbose_logger.debug("MCP sampling: budget checks passed for model=%s", model) 

1041 return None 

1042 

1043 

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 

1051 

1052 return build_synthetic_mcp_request( 

1053 path="/mcp/sampling/createMessage", 

1054 raw_headers=raw_headers, 

1055 client_ip=client_ip, 

1056 ) 

1057 

1058 

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 {} 

1086 

1087 from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request 

1088 from litellm.proxy.proxy_server import proxy_config 

1089 

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 ) 

1098 

1099 

1100class _AcompletionCall(NamedTuple): 

1101 fn: "Callable[..., Awaitable[ModelResponse | CustomStreamWrapper]]" 

1102 

1103 

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 

1110 

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 

1125 

1126 import litellm 

1127 

1128 try: 

1129 from litellm.proxy.proxy_server import llm_router 

1130 

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) 

1136 

1137 

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 ) 

1169 

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 ) 

1179 

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 ) 

1190 

1191 access_denial: Final = await _check_model_access(model, user_api_key_auth) 

1192 if access_denial is not None: 

1193 return access_denial 

1194 

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 

1203 

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 ) 

1211 

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 ) 

1220 

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 ) 

1225 

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 

1243 

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 

1258 

1259 verbose_logger.exception("MCP sampling handler failed: %s", e) 

1260 return ErrorData( 

1261 code=-1, 

1262 message=f"Sampling failed: {e}", 

1263 )