Coverage for agentos/mcp/sampling.py: 0%
202 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 20:40 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 20:40 +0800
1"""
2AgentOS v1.14.0 — MCP Sampling 支持。
4MCP Sampling 允许 MCP Server 向 Client 发起 LLM 请求。
5这是 MCP 协议中 server→client 方向的核心能力,对标 Claude Desktop 的采样功能。
7协议流程:
81. Server 发送 `sampling/createMessage` 请求到 Client
92. Client 调用 LLM 生成回复
103. Client 返回结果给 Server
11"""
13from __future__ import annotations
15from collections.abc import Callable, Coroutine
16from dataclasses import dataclass, field
17from enum import StrEnum
18from typing import Any
20# ── Sampling Data Models ────────────────────
23class SamplingRole(StrEnum):
24 """Sampling 消息角色。"""
26 USER = "user"
27 ASSISTANT = "assistant"
30@dataclass
31class SamplingContentBlock:
32 """Sampling 内容块。
34 支持 text 和 image 两种类型。
35 """
37 type: str = "text" # text | image
38 text: str = ""
39 data: str = "" # base64 image data
40 mime_type: str = ""
42 def to_dict(self) -> dict:
43 d: dict = {"type": self.type}
44 if self.type == "text":
45 d["text"] = self.text
46 elif self.type == "image":
47 d["data"] = self.data
48 d["mimeType"] = self.mime_type
49 return d
51 @classmethod
52 def from_dict(cls, d: dict) -> SamplingContentBlock:
53 return cls(
54 type=d.get("type", "text"),
55 text=d.get("text", ""),
56 data=d.get("data", ""),
57 mime_type=d.get("mimeType", ""),
58 )
60 @classmethod
61 def text_block(cls, text: str) -> SamplingContentBlock:
62 return cls(type="text", text=text)
64 @classmethod
65 def image_block(cls, base64_data: str, mime_type: str = "image/png") -> SamplingContentBlock:
66 return cls(type="image", data=base64_data, mime_type=mime_type)
69@dataclass
70class SamplingMessage:
71 """Sampling 消息。"""
73 role: SamplingRole
74 content: str | list[SamplingContentBlock]
76 def to_dict(self) -> dict:
77 if isinstance(self.content, str):
78 return {
79 "role": self.role.value,
80 "content": {"type": "text", "text": self.content},
81 }
82 return {
83 "role": self.role.value,
84 "content": [c.to_dict() for c in self.content],
85 }
87 @classmethod
88 def from_dict(cls, d: dict) -> SamplingMessage:
89 role = SamplingRole(d["role"])
90 content_raw = d["content"]
91 if isinstance(content_raw, str):
92 content = content_raw
93 elif isinstance(content_raw, dict):
94 content = [SamplingContentBlock.from_dict(content_raw)]
95 elif isinstance(content_raw, list):
96 content = [SamplingContentBlock.from_dict(c) for c in content_raw]
97 else:
98 content = str(content_raw)
99 return cls(role=role, content=content)
102@dataclass
103class SamplingRequest:
104 """Server 发起的 Sampling 请求。
106 符合 MCP sampling/createMessage 规范。
107 """
109 messages: list[SamplingMessage]
110 model_preferences: dict[str, Any] | None = None
111 system_prompt: str = ""
112 include_context: str = "none" # none | thisServer | allServers
113 temperature: float = 0.7
114 max_tokens: int = 4096
115 stop_sequences: list[str] = field(default_factory=list)
116 metadata: dict[str, Any] = field(default_factory=dict)
118 def to_dict(self) -> dict:
119 d: dict = {
120 "messages": [m.to_dict() for m in self.messages],
121 "maxTokens": self.max_tokens,
122 }
123 if self.model_preferences:
124 d["modelPreferences"] = self.model_preferences
125 if self.system_prompt:
126 d["systemPrompt"] = self.system_prompt
127 if self.include_context != "none":
128 d["includeContext"] = self.include_context
129 if self.temperature != 0.7:
130 d["temperature"] = self.temperature
131 if self.stop_sequences:
132 d["stopSequences"] = self.stop_sequences
133 if self.metadata:
134 d["metadata"] = self.metadata
135 return d
137 @classmethod
138 def from_dict(cls, d: dict) -> SamplingRequest:
139 return cls(
140 messages=[SamplingMessage.from_dict(m) for m in d.get("messages", [])],
141 model_preferences=d.get("modelPreferences"),
142 system_prompt=d.get("systemPrompt", ""),
143 include_context=d.get("includeContext", "none"),
144 temperature=d.get("temperature", 0.7),
145 max_tokens=d.get("maxTokens", 4096),
146 stop_sequences=d.get("stopSequences", []),
147 metadata=d.get("metadata", {}),
148 )
151@dataclass
152class SamplingResponse:
153 """Sampling 响应。
155 Client 调用 LLM 后返回给 Server 的结果。
156 """
158 model: str = ""
159 role: SamplingRole = SamplingRole.ASSISTANT
160 content: str | list[SamplingContentBlock] = ""
161 stop_reason: str = "" # endTurn | stopSequence | maxTokens
162 metadata: dict[str, Any] = field(default_factory=dict)
164 def to_dict(self) -> dict:
165 content_block: dict
166 if isinstance(self.content, str):
167 content_block = {"type": "text", "text": self.content}
168 elif isinstance(self.content, list):
169 content_block = [c.to_dict() for c in self.content]
170 else:
171 content_block = {"type": "text", "text": str(self.content)}
173 return {
174 "model": self.model,
175 "role": self.role.value,
176 "content": content_block,
177 "stopReason": self.stop_reason or "endTurn",
178 }
180 @classmethod
181 def from_dict(cls, d: dict) -> SamplingResponse:
182 content_raw = d.get("content", "")
183 if isinstance(content_raw, str):
184 content = content_raw
185 elif isinstance(content_raw, dict):
186 content = [SamplingContentBlock.from_dict(content_raw)]
187 elif isinstance(content_raw, list):
188 content = [SamplingContentBlock.from_dict(c) for c in content_raw]
189 else:
190 content = str(content_raw)
191 return cls(
192 model=d.get("model", ""),
193 role=SamplingRole(d.get("role", "assistant")),
194 content=content,
195 stop_reason=d.get("stopReason", ""),
196 )
199# ── Sampling Handler ────────────────────────
202class SamplingError(Exception):
203 """Sampling 错误。"""
205 def __init__(self, code: int, message: str):
206 self.code = code
207 self.message = message
208 super().__init__(f"Sampling Error [{code}]: {message}")
211# LLM 调用接口
212LLMCallFn = Callable[
213 [SamplingRequest],
214 Coroutine[Any, Any, SamplingResponse],
215]
218class MCPClientSampling:
219 """MCP Client 端 Sampling 支持。
221 在 MCPClient 上挂载此 handler 后,MCP Server 可通过
222 `sampling/createMessage` 向 Client 发起 LLM 请求。
224 Usage:
225 client = MCPClient()
226 sampling = MCPClientSampling(my_llm_call_fn)
227 client.set_sampling_handler(sampling)
228 """
230 def __init__(
231 self,
232 llm_call_fn: LLMCallFn,
233 default_model: str = "claude-sonnet-4-20250514",
234 max_tokens_limit: int = 8192,
235 allow_image_input: bool = True,
236 ):
237 self._llm_call = llm_call_fn
238 self.default_model = default_model
239 self.max_tokens_limit = max_tokens_limit
240 self.allow_image_input = allow_image_input
242 async def handle_create_message(
243 self,
244 params: dict[str, Any],
245 ) -> dict[str, Any]:
246 """处理 sampling/createMessage 请求。
248 Args:
249 params: MCP 请求参数
251 Returns:
252 MCP JSON-RPC 响应 result 部分
253 """
254 try:
255 request = SamplingRequest.from_dict(params)
256 except Exception as e:
257 raise SamplingError(-32602, f"Invalid sampling request: {e}")
259 # Validate max_tokens
260 if request.max_tokens > self.max_tokens_limit:
261 request.max_tokens = self.max_tokens_limit
263 # Validate image input
264 if not self.allow_image_input:
265 for msg in request.messages:
266 if isinstance(msg.content, list):
267 has_image = any(c.type == "image" for c in msg.content)
268 if has_image:
269 raise SamplingError(
270 -32000,
271 "Image input is not allowed by client policy",
272 )
274 try:
275 response = await self._llm_call(request)
276 except Exception as e:
277 raise SamplingError(-32001, f"LLM call failed: {e}")
279 return {
280 "model": response.model or self.default_model,
281 "role": response.role.value,
282 "content": response.to_dict()["content"],
283 "stopReason": response.stop_reason or "endTurn",
284 }
287# ── Mock LLM for Testing ───────────────────
290async def mock_llm_call(request: SamplingRequest) -> SamplingResponse:
291 """Mock LLM 调用函数(用于测试)。
293 简单回显最后一条用户消息的内容。
294 """
295 last_msg = request.messages[-1] if request.messages else None
296 if last_msg and isinstance(last_msg.content, str):
297 content = f"[Mock] Echo: {last_msg.content[:100]}"
298 elif last_msg and isinstance(last_msg.content, list):
299 text_parts = [c.text for c in last_msg.content if c.type == "text"]
300 content = f"[Mock] Echo: {' '.join(text_parts)[:100]}"
301 else:
302 content = "[Mock] No input messages."
304 return SamplingResponse(
305 model="mock-model",
306 content=content,
307 stop_reason="endTurn",
308 )
311# ── Resource Templates ──────────────────────
314@dataclass
315class MCPResourceTemplate:
316 """MCP Resource Template(URI 模板)。
318 Server 可暴露参数化的资源模板,Client 可按模板实例化资源。
319 """
321 uri_template: str
322 name: str = ""
323 description: str = ""
324 mime_type: str = ""
325 annotations: dict[str, Any] = field(default_factory=dict)
327 def to_dict(self) -> dict:
328 d: dict = {
329 "uriTemplate": self.uri_template,
330 "name": self.name,
331 }
332 if self.description:
333 d["description"] = self.description
334 if self.mime_type:
335 d["mimeType"] = self.mime_type
336 if self.annotations:
337 d["annotations"] = self.annotations
338 return d
340 @classmethod
341 def from_dict(cls, d: dict) -> MCPResourceTemplate:
342 return cls(
343 uri_template=d.get("uriTemplate", ""),
344 name=d.get("name", ""),
345 description=d.get("description", ""),
346 mime_type=d.get("mimeType", ""),
347 annotations=d.get("annotations", {}),
348 )
351# ── MCP Logging ────────────────────────────
354class MCPLogLevel(StrEnum):
355 """MCP 日志级别。"""
357 DEBUG = "debug"
358 INFO = "info"
359 NOTICE = "notice"
360 WARNING = "warning"
361 ERROR = "error"
362 CRITICAL = "critical"
363 ALERT = "alert"
364 EMERGENCY = "emergency"
367class MCPLoggingHandler:
368 """MCP Client 端日志接收处理。
370 Client 可通过 `logging/setLevel` 设置日志级别,
371 Server 通过 `notifications/message` 推送日志消息。
372 """
374 def __init__(self, max_level: MCPLogLevel = MCPLogLevel.INFO):
375 self._max_level = max_level
376 self._log_callback: Callable[[MCPLogLevel, str, str], None] | None = None
378 def set_log_callback(
379 self,
380 callback: Callable[[MCPLogLevel, str, str], None],
381 ) -> None:
382 """设置日志回调函数。
384 Args:
385 callback: fn(level, logger_name, message)
386 """
387 self._log_callback = callback
389 @property
390 def max_level(self) -> MCPLogLevel:
391 return self._max_level
393 def set_level(self, level: MCPLogLevel) -> None:
394 """Client 设置日志级别。"""
395 self._max_level = level
397 def _level_rank(self, level: MCPLogLevel) -> int:
398 levels = list(MCPLogLevel)
399 return levels.index(level)
401 def should_log(self, level: MCPLogLevel) -> bool:
402 """判断给定级别的日志是否应被记录。"""
403 return self._level_rank(level) >= self._level_rank(self._max_level)
405 def handle_log_message(
406 self,
407 level: MCPLogLevel,
408 logger_name: str,
409 message: str,
410 ) -> None:
411 """处理来自 Server 的日志消息。"""
412 if self.should_log(level) and self._log_callback:
413 self._log_callback(level, logger_name, message)
416# ── MCP Roots ──────────────────────────────
419@dataclass
420class MCPRoot:
421 """MCP Root — Client 暴露给 Server 的可访问文件系统根。"""
423 uri: str
424 name: str = ""
426 def to_dict(self) -> dict:
427 d: dict = {"uri": self.uri}
428 if self.name:
429 d["name"] = self.name
430 return d
432 @classmethod
433 def from_dict(cls, d: dict) -> MCPRoot:
434 return cls(uri=d.get("uri", ""), name=d.get("name", ""))