Coverage for agentos/api/websocket.py: 0%
263 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 23:17 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 23:17 +0800
1"""
2WebSocket 双向流式通信 — Agent 实时交互层。
4基于 websockets 库,提供 Agent 与客户端之间的全双工实时通信。
5支持流式进度报告、Agent 状态广播、父子 Agent 监控、暂停/恢复/取消。
7协议(JSON,双向):
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"}
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": "..."}
27使用示例::
29 from agentos.api.websocket import AgentWebSocket, serve_ws
31 mgr = SubAgentManager()
33 async def my_run(spec, ctx):
34 await ctx.report_progress(0.5, "thinking")
35 return "answer", 1
37 ws = AgentWebSocket(manager=mgr, run_func=my_run)
38 await serve_ws(ws.handler, port=8765)
39"""
41from __future__ import annotations
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
52import websockets
53from websockets.server import WebSocketServerProtocol
55from agentos.subagent.manager import SubAgentManager, SubAgentSpec
56from agentos.subagent.parent_child import ChildContext, ChildHandle, ChildStatus
58# ──────────────────────────────────────────────
59# 消息协议
60# ──────────────────────────────────────────────
63class WSMsgType(StrEnum):
64 """WebSocket 消息类型。"""
66 # Client → Server
67 RUN = "run"
68 CANCEL = "cancel"
69 PAUSE = "pause"
70 RESUME = "resume"
71 PING = "ping"
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"
85@dataclass
86class WSMessage:
87 """WebSocket 消息体。"""
89 type: str
90 data: dict[str, Any] = field(default_factory=dict)
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 )
100 def serialize(self) -> str:
101 return json.dumps({"type": self.type, **self.data}, ensure_ascii=False)
103 # ── 工厂方法 ──────────────────────────
105 @classmethod
106 def token(cls, text: str, seq: int = 0) -> WSMessage:
107 return cls(WSMsgType.TOKEN, {"text": text, "seq": seq})
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})
113 @classmethod
114 def tool_call(cls, name: str, args: dict) -> WSMessage:
115 return cls(WSMsgType.TOOL_CALL, {"name": name, "args": args})
117 @classmethod
118 def tool_result(cls, name: str, result: Any) -> WSMessage:
119 return cls(WSMsgType.TOOL_RESULT, {"name": name, "result": result})
121 @classmethod
122 def status(cls, status: str, agent_id: str = "") -> WSMessage:
123 return cls(WSMsgType.STATUS, {"status": status, "agent_id": agent_id})
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 )
131 @classmethod
132 def error(cls, message: str, code: str = "UNKNOWN") -> WSMessage:
133 return cls(WSMsgType.ERROR, {"message": message, "code": code})
135 @classmethod
136 def heartbeat(cls) -> WSMessage:
137 return cls(WSMsgType.HEARTBEAT, {"ts": time.time()})
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 )
154# ──────────────────────────────────────────────
155# 会话管理
156# ──────────────────────────────────────────────
159@dataclass
160class WSSession:
161 """单个 WebSocket 连接的会话。"""
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)
171 @property
172 def is_busy(self) -> bool:
173 return self.running_task is not None and not self.running_task.done()
175 def touch(self):
176 self.last_active = time.time()
179# ──────────────────────────────────────────────
180# WebSocket Agent 核心
181# ──────────────────────────────────────────────
184class AgentWebSocket:
185 """Agent WebSocket 服务。
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 """
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] = {}
211 # ── 主 handler ────────────────────────
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
219 heartbeat_task = asyncio.create_task(self._heartbeat_loop(websocket))
221 try:
222 await self._send(websocket, WSMessage.status("connected", session.session_id))
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)))
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)
244 # ── 消息分发 ──────────────────────────
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 }
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"))
266 # ── run ───────────────────────────────
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
278 task = msg.data.get("task", "")
279 if not task:
280 await self._send(ws, WSMessage.error("Missing 'task'", "INVALID"))
281 return
283 await self._send(ws, WSMessage.status("running", session.session_id))
285 session.running_task = asyncio.create_task(self._run_agent(session, task))
286 session.poll_task = asyncio.create_task(self._poll_agent(ws, session))
288 try:
289 await session.running_task
290 except asyncio.CancelledError:
291 await self._send(ws, WSMessage.status("cancelled", session.session_id))
292 return
294 async def _run_agent(self, session: WSSession, task: str) -> None:
295 """启动 Agent 并在完成后推送结果。"""
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
305 result = await self._mgr.spawn_fork(task=task, run_func=capturing_run)
306 session.running_handle = self._mgr.get_handle(result.agent_id)
308 # 停止轮询
309 if session.poll_task and not session.poll_task.done():
310 session.poll_task.cancel()
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))
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
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
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
375 await asyncio.sleep(self._poll_interval)
376 except asyncio.CancelledError:
377 pass
379 # ── cancel / pause / resume ───────────
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))
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"))
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"))
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())
427 # ── 心跳与广播 ────────────────────────
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
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)
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
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 )
477 # ── 辅助 ──────────────────────────────
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
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)
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()
503 # ── 属性 ──────────────────────────────
505 @property
506 def manager(self) -> SubAgentManager:
507 return self._mgr
509 @property
510 def active_connections(self) -> int:
511 return len(self._conn_session)
513 @property
514 def active_sessions(self) -> int:
515 return len(self._sessions)
518# ──────────────────────────────────────────────
519# 便捷启动
520# ──────────────────────────────────────────────
523async def serve_ws(
524 ws_handler,
525 host: str = "0.0.0.0",
526 port: int = 8765,
527 **kwargs,
528):
529 """启动 WebSocket 服务。
531 Args:
532 ws_handler: AgentWebSocket.handler 或兼容的 coroutine handler
533 host: 监听地址
534 port: 监听端口
536 Example::
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