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

1""" 

2AgentOS v1.14.0 — MCP Sampling 支持。 

3 

4MCP Sampling 允许 MCP Server 向 Client 发起 LLM 请求。 

5这是 MCP 协议中 server→client 方向的核心能力,对标 Claude Desktop 的采样功能。 

6 

7协议流程: 

81. Server 发送 `sampling/createMessage` 请求到 Client 

92. Client 调用 LLM 生成回复 

103. Client 返回结果给 Server 

11""" 

12 

13from __future__ import annotations 

14 

15from collections.abc import Callable, Coroutine 

16from dataclasses import dataclass, field 

17from enum import StrEnum 

18from typing import Any 

19 

20# ── Sampling Data Models ──────────────────── 

21 

22 

23class SamplingRole(StrEnum): 

24 """Sampling 消息角色。""" 

25 

26 USER = "user" 

27 ASSISTANT = "assistant" 

28 

29 

30@dataclass 

31class SamplingContentBlock: 

32 """Sampling 内容块。 

33 

34 支持 text 和 image 两种类型。 

35 """ 

36 

37 type: str = "text" # text | image 

38 text: str = "" 

39 data: str = "" # base64 image data 

40 mime_type: str = "" 

41 

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 

50 

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 ) 

59 

60 @classmethod 

61 def text_block(cls, text: str) -> SamplingContentBlock: 

62 return cls(type="text", text=text) 

63 

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) 

67 

68 

69@dataclass 

70class SamplingMessage: 

71 """Sampling 消息。""" 

72 

73 role: SamplingRole 

74 content: str | list[SamplingContentBlock] 

75 

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 } 

86 

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) 

100 

101 

102@dataclass 

103class SamplingRequest: 

104 """Server 发起的 Sampling 请求。 

105 

106 符合 MCP sampling/createMessage 规范。 

107 """ 

108 

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) 

117 

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 

136 

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 ) 

149 

150 

151@dataclass 

152class SamplingResponse: 

153 """Sampling 响应。 

154 

155 Client 调用 LLM 后返回给 Server 的结果。 

156 """ 

157 

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) 

163 

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

172 

173 return { 

174 "model": self.model, 

175 "role": self.role.value, 

176 "content": content_block, 

177 "stopReason": self.stop_reason or "endTurn", 

178 } 

179 

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 ) 

197 

198 

199# ── Sampling Handler ──────────────────────── 

200 

201 

202class SamplingError(Exception): 

203 """Sampling 错误。""" 

204 

205 def __init__(self, code: int, message: str): 

206 self.code = code 

207 self.message = message 

208 super().__init__(f"Sampling Error [{code}]: {message}") 

209 

210 

211# LLM 调用接口 

212LLMCallFn = Callable[ 

213 [SamplingRequest], 

214 Coroutine[Any, Any, SamplingResponse], 

215] 

216 

217 

218class MCPClientSampling: 

219 """MCP Client 端 Sampling 支持。 

220 

221 在 MCPClient 上挂载此 handler 后,MCP Server 可通过 

222 `sampling/createMessage` 向 Client 发起 LLM 请求。 

223 

224 Usage: 

225 client = MCPClient() 

226 sampling = MCPClientSampling(my_llm_call_fn) 

227 client.set_sampling_handler(sampling) 

228 """ 

229 

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 

241 

242 async def handle_create_message( 

243 self, 

244 params: dict[str, Any], 

245 ) -> dict[str, Any]: 

246 """处理 sampling/createMessage 请求。 

247 

248 Args: 

249 params: MCP 请求参数 

250 

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

258 

259 # Validate max_tokens 

260 if request.max_tokens > self.max_tokens_limit: 

261 request.max_tokens = self.max_tokens_limit 

262 

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 ) 

273 

274 try: 

275 response = await self._llm_call(request) 

276 except Exception as e: 

277 raise SamplingError(-32001, f"LLM call failed: {e}") 

278 

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 } 

285 

286 

287# ── Mock LLM for Testing ─────────────────── 

288 

289 

290async def mock_llm_call(request: SamplingRequest) -> SamplingResponse: 

291 """Mock LLM 调用函数(用于测试)。 

292 

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

303 

304 return SamplingResponse( 

305 model="mock-model", 

306 content=content, 

307 stop_reason="endTurn", 

308 ) 

309 

310 

311# ── Resource Templates ────────────────────── 

312 

313 

314@dataclass 

315class MCPResourceTemplate: 

316 """MCP Resource Template(URI 模板)。 

317 

318 Server 可暴露参数化的资源模板,Client 可按模板实例化资源。 

319 """ 

320 

321 uri_template: str 

322 name: str = "" 

323 description: str = "" 

324 mime_type: str = "" 

325 annotations: dict[str, Any] = field(default_factory=dict) 

326 

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 

339 

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 ) 

349 

350 

351# ── MCP Logging ──────────────────────────── 

352 

353 

354class MCPLogLevel(StrEnum): 

355 """MCP 日志级别。""" 

356 

357 DEBUG = "debug" 

358 INFO = "info" 

359 NOTICE = "notice" 

360 WARNING = "warning" 

361 ERROR = "error" 

362 CRITICAL = "critical" 

363 ALERT = "alert" 

364 EMERGENCY = "emergency" 

365 

366 

367class MCPLoggingHandler: 

368 """MCP Client 端日志接收处理。 

369 

370 Client 可通过 `logging/setLevel` 设置日志级别, 

371 Server 通过 `notifications/message` 推送日志消息。 

372 """ 

373 

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 

377 

378 def set_log_callback( 

379 self, 

380 callback: Callable[[MCPLogLevel, str, str], None], 

381 ) -> None: 

382 """设置日志回调函数。 

383 

384 Args: 

385 callback: fn(level, logger_name, message) 

386 """ 

387 self._log_callback = callback 

388 

389 @property 

390 def max_level(self) -> MCPLogLevel: 

391 return self._max_level 

392 

393 def set_level(self, level: MCPLogLevel) -> None: 

394 """Client 设置日志级别。""" 

395 self._max_level = level 

396 

397 def _level_rank(self, level: MCPLogLevel) -> int: 

398 levels = list(MCPLogLevel) 

399 return levels.index(level) 

400 

401 def should_log(self, level: MCPLogLevel) -> bool: 

402 """判断给定级别的日志是否应被记录。""" 

403 return self._level_rank(level) >= self._level_rank(self._max_level) 

404 

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) 

414 

415 

416# ── MCP Roots ────────────────────────────── 

417 

418 

419@dataclass 

420class MCPRoot: 

421 """MCP Root — Client 暴露给 Server 的可访问文件系统根。""" 

422 

423 uri: str 

424 name: str = "" 

425 

426 def to_dict(self) -> dict: 

427 d: dict = {"uri": self.uri} 

428 if self.name: 

429 d["name"] = self.name 

430 return d 

431 

432 @classmethod 

433 def from_dict(cls, d: dict) -> MCPRoot: 

434 return cls(uri=d.get("uri", ""), name=d.get("name", ""))