Coverage for agentos/subagent/parent_child.py: 47%

162 statements  

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

1""" 

2子Agent父子通信 — 状态共享、心跳、生命周期管理。 

3父Agent通过 ChildHandle 管控子Agent;子Agent通过 ChildContext 向父Agent报告。 

4""" 

5 

6from __future__ import annotations 

7 

8import asyncio 

9import time 

10from dataclasses import dataclass, field 

11from enum import Enum 

12from typing import Any, Callable, Awaitable 

13 

14 

15class ChildStatus(str, Enum): 

16 """子Agent运行状态。""" 

17 IDLE = "idle" 

18 RUNNING = "running" 

19 PAUSED = "paused" 

20 COMPLETED = "completed" 

21 FAILED = "failed" 

22 CANCELLED = "cancelled" 

23 TIMEOUT = "timeout" 

24 

25 

26@dataclass 

27class ChildHeartbeat: 

28 """子Agent心跳包。""" 

29 agent_id: str 

30 status: ChildStatus = ChildStatus.RUNNING 

31 progress: float = 0.0 # 0.0 ~ 1.0 

32 current_step: str = "" 

33 message: str = "" 

34 iteration: int = 0 

35 timestamp: float = field(default_factory=time.time) 

36 

37 

38@dataclass 

39class ChildInfo: 

40 """子Agent元信息(父Agent侧)。""" 

41 agent_id: str 

42 task: str 

43 mode: str 

44 status: ChildStatus = ChildStatus.IDLE 

45 spawned_at: float = field(default_factory=time.time) 

46 last_heartbeat: float = field(default_factory=time.time) 

47 heartbeat_interval: float = 2.0 # 期望心跳间隔(秒) 

48 timeout: float | None = None # 超时(秒),None=无超时 

49 progress: float = 0.0 

50 current_step: str = "" 

51 iterations: int = 0 

52 error: str | None = None 

53 output: str = "" 

54 

55 

56class SharedState: 

57 """父子共享状态(线程安全)。""" 

58 

59 def __init__(self): 

60 self._lock = asyncio.Lock() 

61 self._data: dict[str, Any] = {} 

62 

63 async def set(self, key: str, value: Any) -> None: 

64 async with self._lock: 

65 self._data[key] = value 

66 

67 async def get(self, key: str, default: Any = None) -> Any: 

68 async with self._lock: 

69 return self._data.get(key, default) 

70 

71 async def update(self, mapping: dict[str, Any]) -> None: 

72 async with self._lock: 

73 self._data.update(mapping) 

74 

75 async def snapshot(self) -> dict[str, Any]: 

76 async with self._lock: 

77 return dict(self._data) 

78 

79 def set_sync(self, key: str, value: Any) -> None: 

80 """同步写(非协程场景)。""" 

81 self._data[key] = value 

82 

83 def get_sync(self, key: str, default: Any = None) -> Any: 

84 """同步读(非协程场景)。""" 

85 return self._data.get(key, default) 

86 

87 

88class ChildContext: 

89 """子Agent视角 — 向父Agent报告状态、检查控制信号。""" 

90 

91 def __init__( 

92 self, 

93 agent_id: str, 

94 heartbeat_callback: Callable[[ChildHeartbeat], Awaitable[None]] | None = None, 

95 on_cancel: Callable[[], bool] | None = None, 

96 on_pause: Callable[[], Awaitable[None]] | None = None, 

97 shared_state: SharedState | None = None, 

98 ): 

99 self.agent_id = agent_id 

100 self._heartbeat_cb = heartbeat_callback 

101 self._cancel_check = on_cancel or (lambda: False) 

102 self._pause_cb = on_pause or (lambda: asyncio.sleep(0)) 

103 self.shared_state = shared_state or SharedState() 

104 self._cancelled = False 

105 self._paused = False 

106 self._progress = 0.0 

107 self._current_step = "" 

108 self._iteration = 0 

109 

110 @property 

111 def cancelled(self) -> bool: 

112 return self._cancelled 

113 

114 @property 

115 def paused(self) -> bool: 

116 return self._paused 

117 

118 @property 

119 def progress(self) -> float: 

120 return self._progress 

121 

122 async def report_progress( 

123 self, 

124 progress: float, 

125 step: str = "", 

126 message: str = "", 

127 ) -> None: 

128 """子Agent报告进度。""" 

129 self._progress = max(0.0, min(1.0, progress)) 

130 self._current_step = step 

131 if self._heartbeat_cb: 

132 await self._heartbeat_cb(ChildHeartbeat( 

133 agent_id=self.agent_id, 

134 status=ChildStatus.RUNNING, 

135 progress=self._progress, 

136 current_step=step, 

137 message=message, 

138 iteration=self._iteration, 

139 )) 

140 

141 async def step(self, iteration: int, step: str = "") -> None: 

142 """子Agent标记一个执行步。""" 

143 self._iteration = iteration 

144 self._current_step = step 

145 

146 async def check_control(self) -> ChildStatus: 

147 """检查父Agent控制信号,返回应执行的操作。""" 

148 if self._cancel_check(): 

149 self._cancelled = True 

150 return ChildStatus.CANCELLED 

151 if self._paused: 

152 await self._pause_cb() 

153 return ChildStatus.PAUSED 

154 return ChildStatus.RUNNING 

155 

156 async def send_heartbeat(self, message: str = "") -> None: 

157 """子Agent发送心跳。""" 

158 if self._heartbeat_cb: 

159 await self._heartbeat_cb(ChildHeartbeat( 

160 agent_id=self.agent_id, 

161 status=ChildStatus.RUNNING, 

162 progress=self._progress, 

163 current_step=self._current_step, 

164 message=message, 

165 iteration=self._iteration, 

166 )) 

167 

168 async def done(self, output: str = "") -> None: 

169 """子Agent标记完成。""" 

170 if self._heartbeat_cb: 

171 await self._heartbeat_cb(ChildHeartbeat( 

172 agent_id=self.agent_id, 

173 status=ChildStatus.COMPLETED, 

174 progress=1.0, 

175 current_step=self._current_step, 

176 message=output, 

177 iteration=self._iteration, 

178 )) 

179 

180 async def fail(self, error: str) -> None: 

181 """子Agent报告失败。""" 

182 if self._heartbeat_cb: 

183 await self._heartbeat_cb(ChildHeartbeat( 

184 agent_id=self.agent_id, 

185 status=ChildStatus.FAILED, 

186 progress=self._progress, 

187 current_step=self._current_step, 

188 message=error, 

189 iteration=self._iteration, 

190 )) 

191 

192 

193class ChildHandle: 

194 """父Agent视角 — 管控一个子Agent。""" 

195 

196 def __init__( 

197 self, 

198 agent_id: str, 

199 task: str, 

200 mode: str, 

201 timeout: float | None = None, 

202 heartbeat_interval: float = 2.0, 

203 ): 

204 self.info = ChildInfo( 

205 agent_id=agent_id, 

206 task=task, 

207 mode=mode, 

208 heartbeat_interval=heartbeat_interval, 

209 timeout=timeout, 

210 ) 

211 self._cancel_flag = False 

212 self._pause_flag = False 

213 self._resume_event = asyncio.Event() 

214 self._resume_event.set() # 默认未暂停 

215 self.shared_state = SharedState() 

216 self.context: ChildContext | None = None 

217 

218 @property 

219 def agent_id(self) -> str: 

220 return self.info.agent_id 

221 

222 @property 

223 def status(self) -> ChildStatus: 

224 return self.info.status 

225 

226 def create_context(self) -> ChildContext: 

227 """为子Agent创建 ChildContext。""" 

228 ctx = ChildContext( 

229 agent_id=self.agent_id, 

230 heartbeat_callback=self._receive_heartbeat, 

231 on_cancel=self._is_cancelled, 

232 on_pause=self._wait_if_paused, 

233 shared_state=self.shared_state, 

234 ) 

235 self.context = ctx 

236 return ctx 

237 

238 async def _receive_heartbeat(self, hb: ChildHeartbeat) -> None: 

239 """接收子Agent心跳。""" 

240 self.info.last_heartbeat = time.time() 

241 self.info.status = hb.status 

242 self.info.progress = hb.progress 

243 self.info.current_step = hb.current_step 

244 self.info.iterations = hb.iteration 

245 if hb.status == ChildStatus.FAILED: 

246 self.info.error = hb.message 

247 elif hb.status == ChildStatus.COMPLETED: 

248 self.info.output = hb.message 

249 

250 def _is_cancelled(self) -> bool: 

251 return self._cancel_flag 

252 

253 async def _wait_if_paused(self) -> None: 

254 await self._resume_event.wait() 

255 

256 async def cancel(self) -> None: 

257 """取消子Agent。""" 

258 self._cancel_flag = True 

259 self.info.status = ChildStatus.CANCELLED 

260 

261 async def pause(self) -> None: 

262 """暂停子Agent。""" 

263 self._pause_flag = True 

264 self._resume_event.clear() 

265 self.info.status = ChildStatus.PAUSED 

266 if self.context: 

267 self.context._paused = True 

268 

269 async def resume(self) -> None: 

270 """恢复子Agent。""" 

271 self._pause_flag = False 

272 self._resume_event.set() 

273 self.info.status = ChildStatus.RUNNING 

274 if self.context: 

275 self.context._paused = False 

276 

277 def check_timeout(self) -> bool: 

278 """检查是否超时,返回 True 表示已超时。""" 

279 if self.info.timeout is None: 

280 return False 

281 elapsed = time.time() - self.info.spawned_at 

282 return elapsed > self.info.timeout 

283 

284 def check_heartbeat_timeout(self) -> bool: 

285 """检查心跳是否超时(3倍心跳间隔无响应视为失联)。""" 

286 elapsed = time.time() - self.info.last_heartbeat 

287 return elapsed > self.info.heartbeat_interval * 3 

288 

289 def get_status(self) -> dict[str, Any]: 

290 """获取子Agent状态摘要。""" 

291 return { 

292 "agent_id": self.info.agent_id, 

293 "status": self.info.status.value, 

294 "progress": self.info.progress, 

295 "current_step": self.info.current_step, 

296 "iterations": self.info.iterations, 

297 "elapsed": time.time() - self.info.spawned_at, 

298 "error": self.info.error, 

299 }