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

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 dataclasses import dataclass, field 

16from enum import Enum 

17from typing import Any, Callable, Coroutine, Dict, List, Optional, Union 

18 

19 

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

21 

22 

23class SamplingRole(str, Enum): 

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

25 USER = "user" 

26 ASSISTANT = "assistant" 

27 

28 

29@dataclass 

30class SamplingContentBlock: 

31 """Sampling 内容块。 

32 

33 支持 text 和 image 两种类型。 

34 """ 

35 

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

37 text: str = "" 

38 data: str = "" # base64 image data 

39 mime_type: str = "" 

40 

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 

49 

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 ) 

58 

59 @classmethod 

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

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

62 

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) 

66 

67 

68@dataclass 

69class SamplingMessage: 

70 """Sampling 消息。""" 

71 role: SamplingRole 

72 content: Union[str, List[SamplingContentBlock]] 

73 

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 } 

84 

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) 

98 

99 

100@dataclass 

101class SamplingRequest: 

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

103 

104 符合 MCP sampling/createMessage 规范。 

105 """ 

106 

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) 

115 

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 

134 

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 ) 

147 

148 

149@dataclass 

150class SamplingResponse: 

151 """Sampling 响应。 

152 

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

154 """ 

155 

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) 

161 

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

170 

171 return { 

172 "model": self.model, 

173 "role": self.role.value, 

174 "content": content_block, 

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

176 } 

177 

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 ) 

195 

196 

197# ── Sampling Handler ──────────────────────── 

198 

199 

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

206 

207 

208# LLM 调用接口 

209LLMCallFn = Callable[ 

210 [SamplingRequest], 

211 Coroutine[Any, Any, SamplingResponse], 

212] 

213 

214 

215class MCPClientSampling: 

216 """MCP Client 端 Sampling 支持。 

217 

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

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

220 

221 Usage: 

222 client = MCPClient() 

223 sampling = MCPClientSampling(my_llm_call_fn) 

224 client.set_sampling_handler(sampling) 

225 """ 

226 

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 

238 

239 async def handle_create_message( 

240 self, 

241 params: dict[str, Any], 

242 ) -> dict[str, Any]: 

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

244 

245 Args: 

246 params: MCP 请求参数 

247 

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

255 

256 # Validate max_tokens 

257 if request.max_tokens > self.max_tokens_limit: 

258 request.max_tokens = self.max_tokens_limit 

259 

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 ) 

272 

273 try: 

274 response = await self._llm_call(request) 

275 except Exception as e: 

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

277 

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 } 

284 

285 

286# ── Mock LLM for Testing ─────────────────── 

287 

288 

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

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

291 

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

302 

303 return SamplingResponse( 

304 model="mock-model", 

305 content=content, 

306 stop_reason="endTurn", 

307 ) 

308 

309 

310# ── Resource Templates ────────────────────── 

311 

312 

313@dataclass 

314class MCPResourceTemplate: 

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

316 

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

318 """ 

319 

320 uri_template: str 

321 name: str = "" 

322 description: str = "" 

323 mime_type: str = "" 

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

325 

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 

338 

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 ) 

348 

349 

350# ── MCP Logging ──────────────────────────── 

351 

352 

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" 

363 

364 

365class MCPLoggingHandler: 

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

367 

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

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

370 """ 

371 

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 

375 

376 def set_log_callback( 

377 self, 

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

379 ) -> None: 

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

381 

382 Args: 

383 callback: fn(level, logger_name, message) 

384 """ 

385 self._log_callback = callback 

386 

387 @property 

388 def max_level(self) -> MCPLogLevel: 

389 return self._max_level 

390 

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

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

393 self._max_level = level 

394 

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

396 levels = list(MCPLogLevel) 

397 return levels.index(level) 

398 

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

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

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

402 

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) 

412 

413 

414# ── MCP Roots ────────────────────────────── 

415 

416 

417@dataclass 

418class MCPRoot: 

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

420 uri: str 

421 name: str = "" 

422 

423 def to_dict(self) -> dict: 

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

425 if self.name: 

426 d["name"] = self.name 

427 return d 

428 

429 @classmethod 

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

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