Coverage for agentos/mcp/sampling.py: 0%
201 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +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 dataclasses import dataclass, field
16from enum import Enum
17from typing import Any, Callable, Coroutine, Dict, List, Optional, Union
20# ── Sampling Data Models ────────────────────
23class SamplingRole(str, Enum):
24 """Sampling 消息角色。"""
25 USER = "user"
26 ASSISTANT = "assistant"
29@dataclass
30class SamplingContentBlock:
31 """Sampling 内容块。
33 支持 text 和 image 两种类型。
34 """
36 type: str = "text" # text | image
37 text: str = ""
38 data: str = "" # base64 image data
39 mime_type: str = ""
41 def to_dict(self) -> dict:
42 d: dict = {"type": self.type}
43 if self.type == "text":
44 d["text"] = self.text
45 elif self.type == "image":
46 d["data"] = self.data
47 d["mimeType"] = self.mime_type
48 return d
50 @classmethod
51 def from_dict(cls, d: dict) -> "SamplingContentBlock":
52 return cls(
53 type=d.get("type", "text"),
54 text=d.get("text", ""),
55 data=d.get("data", ""),
56 mime_type=d.get("mimeType", ""),
57 )
59 @classmethod
60 def text_block(cls, text: str) -> "SamplingContentBlock":
61 return cls(type="text", text=text)
63 @classmethod
64 def image_block(cls, base64_data: str, mime_type: str = "image/png") -> "SamplingContentBlock":
65 return cls(type="image", data=base64_data, mime_type=mime_type)
68@dataclass
69class SamplingMessage:
70 """Sampling 消息。"""
71 role: SamplingRole
72 content: Union[str, List[SamplingContentBlock]]
74 def to_dict(self) -> dict:
75 if isinstance(self.content, str):
76 return {
77 "role": self.role.value,
78 "content": {"type": "text", "text": self.content},
79 }
80 return {
81 "role": self.role.value,
82 "content": [c.to_dict() for c in self.content],
83 }
85 @classmethod
86 def from_dict(cls, d: dict) -> "SamplingMessage":
87 role = SamplingRole(d["role"])
88 content_raw = d["content"]
89 if isinstance(content_raw, str):
90 content = content_raw
91 elif isinstance(content_raw, dict):
92 content = [SamplingContentBlock.from_dict(content_raw)]
93 elif isinstance(content_raw, list):
94 content = [SamplingContentBlock.from_dict(c) for c in content_raw]
95 else:
96 content = str(content_raw)
97 return cls(role=role, content=content)
100@dataclass
101class SamplingRequest:
102 """Server 发起的 Sampling 请求。
104 符合 MCP sampling/createMessage 规范。
105 """
107 messages: List[SamplingMessage]
108 model_preferences: Optional[Dict[str, Any]] = None
109 system_prompt: str = ""
110 include_context: str = "none" # none | thisServer | allServers
111 temperature: float = 0.7
112 max_tokens: int = 4096
113 stop_sequences: List[str] = field(default_factory=list)
114 metadata: Dict[str, Any] = field(default_factory=dict)
116 def to_dict(self) -> dict:
117 d: dict = {
118 "messages": [m.to_dict() for m in self.messages],
119 "maxTokens": self.max_tokens,
120 }
121 if self.model_preferences:
122 d["modelPreferences"] = self.model_preferences
123 if self.system_prompt:
124 d["systemPrompt"] = self.system_prompt
125 if self.include_context != "none":
126 d["includeContext"] = self.include_context
127 if self.temperature != 0.7:
128 d["temperature"] = self.temperature
129 if self.stop_sequences:
130 d["stopSequences"] = self.stop_sequences
131 if self.metadata:
132 d["metadata"] = self.metadata
133 return d
135 @classmethod
136 def from_dict(cls, d: dict) -> "SamplingRequest":
137 return cls(
138 messages=[SamplingMessage.from_dict(m) for m in d.get("messages", [])],
139 model_preferences=d.get("modelPreferences"),
140 system_prompt=d.get("systemPrompt", ""),
141 include_context=d.get("includeContext", "none"),
142 temperature=d.get("temperature", 0.7),
143 max_tokens=d.get("maxTokens", 4096),
144 stop_sequences=d.get("stopSequences", []),
145 metadata=d.get("metadata", {}),
146 )
149@dataclass
150class SamplingResponse:
151 """Sampling 响应。
153 Client 调用 LLM 后返回给 Server 的结果。
154 """
156 model: str = ""
157 role: SamplingRole = SamplingRole.ASSISTANT
158 content: Union[str, List[SamplingContentBlock]] = ""
159 stop_reason: str = "" # endTurn | stopSequence | maxTokens
160 metadata: Dict[str, Any] = field(default_factory=dict)
162 def to_dict(self) -> dict:
163 content_block: dict
164 if isinstance(self.content, str):
165 content_block = {"type": "text", "text": self.content}
166 elif isinstance(self.content, list):
167 content_block = [c.to_dict() for c in self.content]
168 else:
169 content_block = {"type": "text", "text": str(self.content)}
171 return {
172 "model": self.model,
173 "role": self.role.value,
174 "content": content_block,
175 "stopReason": self.stop_reason or "endTurn",
176 }
178 @classmethod
179 def from_dict(cls, d: dict) -> "SamplingResponse":
180 content_raw = d.get("content", "")
181 if isinstance(content_raw, str):
182 content = content_raw
183 elif isinstance(content_raw, dict):
184 content = [SamplingContentBlock.from_dict(content_raw)]
185 elif isinstance(content_raw, list):
186 content = [SamplingContentBlock.from_dict(c) for c in content_raw]
187 else:
188 content = str(content_raw)
189 return cls(
190 model=d.get("model", ""),
191 role=SamplingRole(d.get("role", "assistant")),
192 content=content,
193 stop_reason=d.get("stopReason", ""),
194 )
197# ── Sampling Handler ────────────────────────
200class SamplingError(Exception):
201 """Sampling 错误。"""
202 def __init__(self, code: int, message: str):
203 self.code = code
204 self.message = message
205 super().__init__(f"Sampling Error [{code}]: {message}")
208# LLM 调用接口
209LLMCallFn = Callable[
210 [SamplingRequest],
211 Coroutine[Any, Any, SamplingResponse],
212]
215class MCPClientSampling:
216 """MCP Client 端 Sampling 支持。
218 在 MCPClient 上挂载此 handler 后,MCP Server 可通过
219 `sampling/createMessage` 向 Client 发起 LLM 请求。
221 Usage:
222 client = MCPClient()
223 sampling = MCPClientSampling(my_llm_call_fn)
224 client.set_sampling_handler(sampling)
225 """
227 def __init__(
228 self,
229 llm_call_fn: LLMCallFn,
230 default_model: str = "claude-sonnet-4-20250514",
231 max_tokens_limit: int = 8192,
232 allow_image_input: bool = True,
233 ):
234 self._llm_call = llm_call_fn
235 self.default_model = default_model
236 self.max_tokens_limit = max_tokens_limit
237 self.allow_image_input = allow_image_input
239 async def handle_create_message(
240 self,
241 params: dict[str, Any],
242 ) -> dict[str, Any]:
243 """处理 sampling/createMessage 请求。
245 Args:
246 params: MCP 请求参数
248 Returns:
249 MCP JSON-RPC 响应 result 部分
250 """
251 try:
252 request = SamplingRequest.from_dict(params)
253 except Exception as e:
254 raise SamplingError(-32602, f"Invalid sampling request: {e}")
256 # Validate max_tokens
257 if request.max_tokens > self.max_tokens_limit:
258 request.max_tokens = self.max_tokens_limit
260 # Validate image input
261 if not self.allow_image_input:
262 for msg in request.messages:
263 if isinstance(msg.content, list):
264 has_image = any(
265 c.type == "image" for c in msg.content
266 )
267 if has_image:
268 raise SamplingError(
269 -32000,
270 "Image input is not allowed by client policy",
271 )
273 try:
274 response = await self._llm_call(request)
275 except Exception as e:
276 raise SamplingError(-32001, f"LLM call failed: {e}")
278 return {
279 "model": response.model or self.default_model,
280 "role": response.role.value,
281 "content": response.to_dict()["content"],
282 "stopReason": response.stop_reason or "endTurn",
283 }
286# ── Mock LLM for Testing ───────────────────
289async def mock_llm_call(request: SamplingRequest) -> SamplingResponse:
290 """Mock LLM 调用函数(用于测试)。
292 简单回显最后一条用户消息的内容。
293 """
294 last_msg = request.messages[-1] if request.messages else None
295 if last_msg and isinstance(last_msg.content, str):
296 content = f"[Mock] Echo: {last_msg.content[:100]}"
297 elif last_msg and isinstance(last_msg.content, list):
298 text_parts = [c.text for c in last_msg.content if c.type == "text"]
299 content = f"[Mock] Echo: {' '.join(text_parts)[:100]}"
300 else:
301 content = "[Mock] No input messages."
303 return SamplingResponse(
304 model="mock-model",
305 content=content,
306 stop_reason="endTurn",
307 )
310# ── Resource Templates ──────────────────────
313@dataclass
314class MCPResourceTemplate:
315 """MCP Resource Template(URI 模板)。
317 Server 可暴露参数化的资源模板,Client 可按模板实例化资源。
318 """
320 uri_template: str
321 name: str = ""
322 description: str = ""
323 mime_type: str = ""
324 annotations: Dict[str, Any] = field(default_factory=dict)
326 def to_dict(self) -> dict:
327 d: dict = {
328 "uriTemplate": self.uri_template,
329 "name": self.name,
330 }
331 if self.description:
332 d["description"] = self.description
333 if self.mime_type:
334 d["mimeType"] = self.mime_type
335 if self.annotations:
336 d["annotations"] = self.annotations
337 return d
339 @classmethod
340 def from_dict(cls, d: dict) -> "MCPResourceTemplate":
341 return cls(
342 uri_template=d.get("uriTemplate", ""),
343 name=d.get("name", ""),
344 description=d.get("description", ""),
345 mime_type=d.get("mimeType", ""),
346 annotations=d.get("annotations", {}),
347 )
350# ── MCP Logging ────────────────────────────
353class MCPLogLevel(str, Enum):
354 """MCP 日志级别。"""
355 DEBUG = "debug"
356 INFO = "info"
357 NOTICE = "notice"
358 WARNING = "warning"
359 ERROR = "error"
360 CRITICAL = "critical"
361 ALERT = "alert"
362 EMERGENCY = "emergency"
365class MCPLoggingHandler:
366 """MCP Client 端日志接收处理。
368 Client 可通过 `logging/setLevel` 设置日志级别,
369 Server 通过 `notifications/message` 推送日志消息。
370 """
372 def __init__(self, max_level: MCPLogLevel = MCPLogLevel.INFO):
373 self._max_level = max_level
374 self._log_callback: Optional[Callable[[MCPLogLevel, str, str], None]] = None
376 def set_log_callback(
377 self,
378 callback: Callable[[MCPLogLevel, str, str], None],
379 ) -> None:
380 """设置日志回调函数。
382 Args:
383 callback: fn(level, logger_name, message)
384 """
385 self._log_callback = callback
387 @property
388 def max_level(self) -> MCPLogLevel:
389 return self._max_level
391 def set_level(self, level: MCPLogLevel) -> None:
392 """Client 设置日志级别。"""
393 self._max_level = level
395 def _level_rank(self, level: MCPLogLevel) -> int:
396 levels = list(MCPLogLevel)
397 return levels.index(level)
399 def should_log(self, level: MCPLogLevel) -> bool:
400 """判断给定级别的日志是否应被记录。"""
401 return self._level_rank(level) >= self._level_rank(self._max_level)
403 def handle_log_message(
404 self,
405 level: MCPLogLevel,
406 logger_name: str,
407 message: str,
408 ) -> None:
409 """处理来自 Server 的日志消息。"""
410 if self.should_log(level) and self._log_callback:
411 self._log_callback(level, logger_name, message)
414# ── MCP Roots ──────────────────────────────
417@dataclass
418class MCPRoot:
419 """MCP Root — Client 暴露给 Server 的可访问文件系统根。"""
420 uri: str
421 name: str = ""
423 def to_dict(self) -> dict:
424 d: dict = {"uri": self.uri}
425 if self.name:
426 d["name"] = self.name
427 return d
429 @classmethod
430 def from_dict(cls, d: dict) -> "MCPRoot":
431 return cls(uri=d.get("uri", ""), name=d.get("name", ""))