Coverage for agentos/api/websocket.py: 0%

263 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-09 09:19 +0800

1""" 

2WebSocket 双向流式通信 — Agent 实时交互层。 

3 

4基于 websockets 库,提供 Agent 与客户端之间的全双工实时通信。 

5支持流式进度报告、Agent 状态广播、父子 Agent 监控、暂停/恢复/取消。 

6 

7协议(JSON,双向): 

8 

9 Client → Server: 

10 {"type": "run", "task": "...", "session_id": "..."} 

11 {"type": "cancel", "session_id": "..."} 

12 {"type": "pause", "session_id": "..."} 

13 {"type": "resume", "session_id": "..."} 

14 {"type": "ping"} 

15 

16 Server → Client: 

17 {"type": "token", "text": "...", "seq": N} 

18 {"type": "progress", "value": 0.5, "step": "..."} 

19 {"type": "tool_call", "name": "...", "args": {...}} 

20 {"type": "tool_result", "name": "...", "result": ...} 

21 {"type": "status", "status": "running"|"paused"|"..."} 

22 {"type": "done", "output": "...", "iterations": N} 

23 {"type": "error", "message": "..."} 

24 {"type": "heartbeat"} 

25 {"type": "child_update", "agent_id": "...", "status": "..."} 

26 

27使用示例:: 

28 

29 from agentos.api.websocket import AgentWebSocket, serve_ws 

30 

31 mgr = SubAgentManager() 

32 

33 async def my_run(spec, ctx): 

34 await ctx.report_progress(0.5, "thinking") 

35 return "answer", 1 

36 

37 ws = AgentWebSocket(manager=mgr, run_func=my_run) 

38 await serve_ws(ws.handler, port=8765) 

39""" 

40 

41from __future__ import annotations 

42 

43import asyncio 

44import json 

45import time 

46import uuid 

47from collections.abc import Awaitable, Callable 

48from dataclasses import dataclass, field 

49from enum import StrEnum 

50from typing import Any 

51 

52import websockets 

53from websockets.server import WebSocketServerProtocol 

54 

55from agentos.subagent.manager import SubAgentManager, SubAgentSpec 

56from agentos.subagent.parent_child import ChildContext, ChildHandle, ChildStatus 

57 

58# ────────────────────────────────────────────── 

59# 消息协议 

60# ────────────────────────────────────────────── 

61 

62 

63class WSMsgType(StrEnum): 

64 """WebSocket 消息类型。""" 

65 

66 # Client → Server 

67 RUN = "run" 

68 CANCEL = "cancel" 

69 PAUSE = "pause" 

70 RESUME = "resume" 

71 PING = "ping" 

72 

73 # Server → Client 

74 TOKEN = "token" 

75 PROGRESS = "progress" 

76 TOOL_CALL = "tool_call" 

77 TOOL_RESULT = "tool_result" 

78 STATUS = "status" 

79 DONE = "done" 

80 ERROR = "error" 

81 HEARTBEAT = "heartbeat" 

82 CHILD_UPDATE = "child_update" 

83 

84 

85@dataclass 

86class WSMessage: 

87 """WebSocket 消息体。""" 

88 

89 type: str 

90 data: dict[str, Any] = field(default_factory=dict) 

91 

92 @classmethod 

93 def parse(cls, raw: str | bytes) -> WSMessage: 

94 payload = json.loads(raw if isinstance(raw, str) else raw.decode()) 

95 return cls( 

96 type=payload.get("type", ""), 

97 data={k: v for k, v in payload.items() if k != "type"}, 

98 ) 

99 

100 def serialize(self) -> str: 

101 return json.dumps({"type": self.type, **self.data}, ensure_ascii=False) 

102 

103 # ── 工厂方法 ────────────────────────── 

104 

105 @classmethod 

106 def token(cls, text: str, seq: int = 0) -> WSMessage: 

107 return cls(WSMsgType.TOKEN, {"text": text, "seq": seq}) 

108 

109 @classmethod 

110 def progress(cls, value: float, step: str = "", agent_id: str = "") -> WSMessage: 

111 return cls(WSMsgType.PROGRESS, {"value": value, "step": step, "agent_id": agent_id}) 

112 

113 @classmethod 

114 def tool_call(cls, name: str, args: dict) -> WSMessage: 

115 return cls(WSMsgType.TOOL_CALL, {"name": name, "args": args}) 

116 

117 @classmethod 

118 def tool_result(cls, name: str, result: Any) -> WSMessage: 

119 return cls(WSMsgType.TOOL_RESULT, {"name": name, "result": result}) 

120 

121 @classmethod 

122 def status(cls, status: str, agent_id: str = "") -> WSMessage: 

123 return cls(WSMsgType.STATUS, {"status": status, "agent_id": agent_id}) 

124 

125 @classmethod 

126 def done(cls, output: str, iterations: int = 0, agent_id: str = "") -> WSMessage: 

127 return cls( 

128 WSMsgType.DONE, {"output": output, "iterations": iterations, "agent_id": agent_id} 

129 ) 

130 

131 @classmethod 

132 def error(cls, message: str, code: str = "UNKNOWN") -> WSMessage: 

133 return cls(WSMsgType.ERROR, {"message": message, "code": code}) 

134 

135 @classmethod 

136 def heartbeat(cls) -> WSMessage: 

137 return cls(WSMsgType.HEARTBEAT, {"ts": time.time()}) 

138 

139 @classmethod 

140 def child_update( 

141 cls, agent_id: str, status: str, progress: float = 0, step: str = "" 

142 ) -> WSMessage: 

143 return cls( 

144 WSMsgType.CHILD_UPDATE, 

145 { 

146 "agent_id": agent_id, 

147 "status": status, 

148 "progress": progress, 

149 "step": step, 

150 }, 

151 ) 

152 

153 

154# ────────────────────────────────────────────── 

155# 会话管理 

156# ────────────────────────────────────────────── 

157 

158 

159@dataclass 

160class WSSession: 

161 """单个 WebSocket 连接的会话。""" 

162 

163 session_id: str = field(default_factory=lambda: uuid.uuid4().hex[:12]) 

164 connected_at: float = field(default_factory=time.time) 

165 last_active: float = field(default_factory=time.time) 

166 running_task: asyncio.Task | None = None 

167 running_handle: ChildHandle | None = None 

168 poll_task: asyncio.Task | None = None 

169 metadata: dict[str, Any] = field(default_factory=dict) 

170 

171 @property 

172 def is_busy(self) -> bool: 

173 return self.running_task is not None and not self.running_task.done() 

174 

175 def touch(self): 

176 self.last_active = time.time() 

177 

178 

179# ────────────────────────────────────────────── 

180# WebSocket Agent 核心 

181# ────────────────────────────────────────────── 

182 

183 

184class AgentWebSocket: 

185 """Agent WebSocket 服务。 

186 

187 Args: 

188 manager: SubAgentManager 实例 

189 run_func: 自定义执行函数 (spec, ctx) -> (output, iterations) 

190 heartbeat_interval: WebSocket 心跳间隔(秒) 

191 poll_interval: 子 Agent 状态轮询间隔(秒) 

192 max_message_size: 最大消息大小(字节) 

193 """ 

194 

195 def __init__( 

196 self, 

197 manager: SubAgentManager | None = None, 

198 run_func: Callable[[SubAgentSpec, ChildContext], Awaitable[tuple[str, int]]] | None = None, 

199 heartbeat_interval: float = 15.0, 

200 poll_interval: float = 0.5, 

201 max_message_size: int = 2**20, 

202 ): 

203 self._mgr = manager or SubAgentManager() 

204 self._run = run_func 

205 self._heartbeat_interval = heartbeat_interval 

206 self._poll_interval = poll_interval 

207 self._max_message_size = max_message_size 

208 self._sessions: dict[str, WSSession] = {} 

209 self._conn_session: dict[WebSocketServerProtocol, str] = {} 

210 

211 # ── 主 handler ──────────────────────── 

212 

213 async def handler(self, websocket: WebSocketServerProtocol) -> None: 

214 """单连接 handler。""" 

215 session = WSSession() 

216 self._sessions[session.session_id] = session 

217 self._conn_session[websocket] = session.session_id 

218 

219 heartbeat_task = asyncio.create_task(self._heartbeat_loop(websocket)) 

220 

221 try: 

222 await self._send(websocket, WSMessage.status("connected", session.session_id)) 

223 

224 async for raw in websocket: 

225 session.touch() 

226 try: 

227 msg = WSMessage.parse(raw) 

228 await self._dispatch(websocket, session, msg) 

229 except json.JSONDecodeError: 

230 await self._send(websocket, WSMessage.error("Invalid JSON", "PARSE_ERROR")) 

231 except Exception as e: 

232 await self._send(websocket, WSMessage.error(str(e))) 

233 

234 except websockets.exceptions.ConnectionClosed: 

235 pass 

236 finally: 

237 heartbeat_task.cancel() 

238 try: 

239 await heartbeat_task 

240 except asyncio.CancelledError: 

241 pass 

242 await self._cleanup_session(websocket, session) 

243 

244 # ── 消息分发 ────────────────────────── 

245 

246 async def _dispatch( 

247 self, 

248 ws: WebSocketServerProtocol, 

249 session: WSSession, 

250 msg: WSMessage, 

251 ) -> None: 

252 handlers: dict[str, Callable] = { 

253 WSMsgType.RUN: self._handle_run, 

254 WSMsgType.CANCEL: self._handle_cancel, 

255 WSMsgType.PAUSE: self._handle_pause, 

256 WSMsgType.RESUME: self._handle_resume, 

257 WSMsgType.PING: self._handle_ping, 

258 } 

259 

260 handler = handlers.get(msg.type) 

261 if handler: 

262 await handler(ws, session, msg) 

263 else: 

264 await self._send(ws, WSMessage.error(f"Unknown type: {msg.type}", "UNKNOWN_TYPE")) 

265 

266 # ── run ─────────────────────────────── 

267 

268 async def _handle_run( 

269 self, 

270 ws: WebSocketServerProtocol, 

271 session: WSSession, 

272 msg: WSMessage, 

273 ) -> None: 

274 if session.is_busy: 

275 await self._send(ws, WSMessage.error("Session busy", "BUSY")) 

276 return 

277 

278 task = msg.data.get("task", "") 

279 if not task: 

280 await self._send(ws, WSMessage.error("Missing 'task'", "INVALID")) 

281 return 

282 

283 await self._send(ws, WSMessage.status("running", session.session_id)) 

284 

285 session.running_task = asyncio.create_task(self._run_agent(session, task)) 

286 session.poll_task = asyncio.create_task(self._poll_agent(ws, session)) 

287 

288 try: 

289 await session.running_task 

290 except asyncio.CancelledError: 

291 await self._send(ws, WSMessage.status("cancelled", session.session_id)) 

292 return 

293 

294 async def _run_agent(self, session: WSSession, task: str) -> None: 

295 """启动 Agent 并在完成后推送结果。""" 

296 

297 async def capturing_run(spec: SubAgentSpec, ctx: ChildContext) -> tuple[str, int]: 

298 if self._run: 

299 return await self._run(spec, ctx) 

300 # 默认 fallback 

301 await ctx.report_progress(0.5, "processing") 

302 await ctx.report_progress(1.0, "done") 

303 return f"Agent received: {task}", 1 

304 

305 result = await self._mgr.spawn_fork(task=task, run_func=capturing_run) 

306 session.running_handle = self._mgr.get_handle(result.agent_id) 

307 

308 # 停止轮询 

309 if session.poll_task and not session.poll_task.done(): 

310 session.poll_task.cancel() 

311 

312 # 确定最终状态并发送 

313 if result.error: 

314 await self.broadcast_to_session(session, WSMessage.error(result.error)) 

315 await self.broadcast_to_session(session, WSMessage.status("failed", result.agent_id)) 

316 else: 

317 await self.broadcast_to_session( 

318 session, 

319 WSMessage.done( 

320 output=result.output, 

321 iterations=result.iterations, 

322 agent_id=result.agent_id, 

323 ), 

324 ) 

325 await self.broadcast_to_session(session, WSMessage.status("completed", result.agent_id)) 

326 

327 async def _poll_agent( 

328 self, 

329 ws: WebSocketServerProtocol, 

330 session: WSSession, 

331 ) -> None: 

332 """轮询子 Agent 状态并流式推送进度。""" 

333 try: 

334 last_progress = -1.0 

335 last_step = "" 

336 while session.running_task and not session.running_task.done(): 

337 handle = session.running_handle 

338 if handle is None: 

339 # spawn_fork 尚未返回,检查 manager 中是否有新 agent 

340 children = self._mgr.list_children() 

341 if children: 

342 latest = children[-1] 

343 sid = latest.get("agent_id", "") 

344 handle = self._mgr.get_handle(sid) 

345 if handle and handle.status not in (ChildStatus.IDLE,): 

346 session.running_handle = handle 

347 

348 if handle: 

349 cur_progress = handle.info.progress 

350 cur_step = handle.info.current_step 

351 if cur_progress != last_progress or cur_step != last_step: 

352 await self._send( 

353 ws, 

354 WSMessage.progress( 

355 value=cur_progress, 

356 step=cur_step, 

357 agent_id=handle.agent_id, 

358 ), 

359 ) 

360 last_progress = cur_progress 

361 last_step = cur_step 

362 

363 # 推送状态变化 

364 if handle.status == ChildStatus.FAILED: 

365 await self._send( 

366 ws, 

367 WSMessage.error( 

368 handle.info.error or "Agent failed", 

369 ), 

370 ) 

371 break 

372 elif handle.status == ChildStatus.CANCELLED: 

373 break 

374 

375 await asyncio.sleep(self._poll_interval) 

376 except asyncio.CancelledError: 

377 pass 

378 

379 # ── cancel / pause / resume ─────────── 

380 

381 async def _handle_cancel( 

382 self, 

383 ws: WebSocketServerProtocol, 

384 session: WSSession, 

385 msg: WSMessage, 

386 ) -> None: 

387 if session.running_handle: 

388 await session.running_handle.cancel() 

389 if session.running_task and not session.running_task.done(): 

390 session.running_task.cancel() 

391 if session.poll_task and not session.poll_task.done(): 

392 session.poll_task.cancel() 

393 await self._send(ws, WSMessage.status("cancelled", session.session_id)) 

394 

395 async def _handle_pause( 

396 self, 

397 ws: WebSocketServerProtocol, 

398 session: WSSession, 

399 msg: WSMessage, 

400 ) -> None: 

401 if session.running_handle: 

402 await session.running_handle.pause() 

403 await self._send(ws, WSMessage.status("paused", session.session_id)) 

404 else: 

405 await self._send(ws, WSMessage.error("No agent to pause", "IDLE")) 

406 

407 async def _handle_resume( 

408 self, 

409 ws: WebSocketServerProtocol, 

410 session: WSSession, 

411 msg: WSMessage, 

412 ) -> None: 

413 if session.running_handle: 

414 await session.running_handle.resume() 

415 await self._send(ws, WSMessage.status("running", session.session_id)) 

416 else: 

417 await self._send(ws, WSMessage.error("No agent to resume", "IDLE")) 

418 

419 async def _handle_ping( 

420 self, 

421 ws: WebSocketServerProtocol, 

422 session: WSSession, 

423 msg: WSMessage, 

424 ) -> None: 

425 await self._send(ws, WSMessage.heartbeat()) 

426 

427 # ── 心跳与广播 ──────────────────────── 

428 

429 async def _heartbeat_loop(self, ws: WebSocketServerProtocol) -> None: 

430 try: 

431 while True: 

432 await asyncio.sleep(self._heartbeat_interval) 

433 await self._send(ws, WSMessage.heartbeat()) 

434 except (websockets.exceptions.ConnectionClosed, asyncio.CancelledError): 

435 pass 

436 

437 async def broadcast(self, msg: WSMessage, exclude_session: str = "") -> None: 

438 """向所有连接的客户端广播。""" 

439 dead: list[WebSocketServerProtocol] = [] 

440 for ws, sid in list(self._conn_session.items()): 

441 if sid == exclude_session: 

442 continue 

443 try: 

444 await ws.send(msg.serialize()) 

445 except websockets.exceptions.ConnectionClosed: 

446 dead.append(ws) 

447 for ws in dead: 

448 await self._cleanup_ws(ws) 

449 

450 async def broadcast_to_session( 

451 self, 

452 session: WSSession, 

453 msg: WSMessage, 

454 ) -> None: 

455 """向指定会话对应的 WebSocket 发送消息。""" 

456 for ws, sid in self._conn_session.items(): 

457 if sid == session.session_id: 

458 try: 

459 await ws.send(msg.serialize()) 

460 except websockets.exceptions.ConnectionClosed: 

461 pass 

462 return 

463 

464 async def broadcast_child_status(self) -> None: 

465 """广播所有子 Agent 状态。""" 

466 children = self._mgr.list_children() 

467 for child in children: 

468 await self.broadcast( 

469 WSMessage.child_update( 

470 agent_id=child.get("agent_id", ""), 

471 status=child.get("status", "unknown"), 

472 progress=child.get("progress", 0), 

473 step=child.get("current_step", ""), 

474 ) 

475 ) 

476 

477 # ── 辅助 ────────────────────────────── 

478 

479 async def _send(self, ws: WebSocketServerProtocol, msg: WSMessage) -> None: 

480 try: 

481 await ws.send(msg.serialize()) 

482 except websockets.exceptions.ConnectionClosed: 

483 pass 

484 

485 async def _cleanup_session(self, ws: WebSocketServerProtocol, session: WSSession) -> None: 

486 if session.running_task and not session.running_task.done(): 

487 session.running_task.cancel() 

488 if session.poll_task and not session.poll_task.done(): 

489 session.poll_task.cancel() 

490 self._conn_session.pop(ws, None) 

491 self._sessions.pop(session.session_id, None) 

492 

493 async def _cleanup_ws(self, ws: WebSocketServerProtocol) -> None: 

494 sid = self._conn_session.pop(ws, None) 

495 if sid: 

496 session = self._sessions.pop(sid, None) 

497 if session: 

498 if session.running_task and not session.running_task.done(): 

499 session.running_task.cancel() 

500 if session.poll_task and not session.poll_task.done(): 

501 session.poll_task.cancel() 

502 

503 # ── 属性 ────────────────────────────── 

504 

505 @property 

506 def manager(self) -> SubAgentManager: 

507 return self._mgr 

508 

509 @property 

510 def active_connections(self) -> int: 

511 return len(self._conn_session) 

512 

513 @property 

514 def active_sessions(self) -> int: 

515 return len(self._sessions) 

516 

517 

518# ────────────────────────────────────────────── 

519# 便捷启动 

520# ────────────────────────────────────────────── 

521 

522 

523async def serve_ws( 

524 ws_handler, 

525 host: str = "0.0.0.0", 

526 port: int = 8765, 

527 **kwargs, 

528): 

529 """启动 WebSocket 服务。 

530 

531 Args: 

532 ws_handler: AgentWebSocket.handler 或兼容的 coroutine handler 

533 host: 监听地址 

534 port: 监听端口 

535 

536 Example:: 

537 

538 mgr = SubAgentManager() 

539 ws = AgentWebSocket(manager=mgr) 

540 await serve_ws(ws.handler, port=8765) 

541 """ 

542 async with websockets.serve( 

543 ws_handler, 

544 host=host, 

545 port=port, 

546 max_size=kwargs.pop("max_size", 2**20), 

547 ping_interval=kwargs.pop("ping_interval", 20), 

548 **kwargs, 

549 ): 

550 print(f"WebSocket server listening on ws://{host}:{port}") 

551 await asyncio.Future() # run forever