Coverage for agentos/swarm/execution_trace.py: 36%

184 statements  

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

1""" 

2v1.9.6: Execution Trace — full observability into Agent task execution. 

3 

4Captures every sub-task, gate evaluation, retry, and timing detail. 

5Supports timeline visualization, bottleneck detection, and debugging. 

6""" 

7 

8from __future__ import annotations 

9 

10import json 

11import time 

12import uuid 

13from dataclasses import dataclass, field 

14from enum import Enum 

15from typing import Any 

16 

17 

18class TraceEvent(str, Enum): 

19 """Event types in an execution trace.""" 

20 

21 TASK_START = "task_start" 

22 TASK_END = "task_end" 

23 SUBTASK_START = "subtask_start" 

24 SUBTASK_END = "subtask_end" 

25 DECOMPOSE = "decompose" 

26 FUSE = "fuse" 

27 RETRY = "retry" 

28 FALLBACK = "fallback" 

29 GATE_CHECK = "gate_check" 

30 HITL_BREAK = "hitl_break" 

31 HITL_RESUME = "hitl_resume" 

32 SANDBOX_RUN = "sandbox_run" 

33 ERROR = "error" 

34 ABORT = "abort" 

35 

36 

37@dataclass 

38class TraceSpan: 

39 """A single span in an execution trace.""" 

40 

41 id: str = field(default_factory=lambda: uuid.uuid4().hex[:8]) 

42 parent_id: str = "" 

43 event: TraceEvent = TraceEvent.TASK_START 

44 name: str = "" 

45 status: str = "started" # started, done, failed, aborted 

46 start_ms: float = field(default_factory=lambda: time.time() * 1000) 

47 end_ms: float = 0.0 

48 duration_ms: float = 0.0 

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

50 tags: list[str] = field(default_factory=list) 

51 children: list[TraceSpan] = field(default_factory=list) 

52 

53 @property 

54 def is_leaf(self) -> bool: 

55 return len(self.children) == 0 

56 

57 def to_dict(self) -> dict: 

58 return { 

59 "id": self.id, 

60 "parent_id": self.parent_id, 

61 "event": self.event.value, 

62 "name": self.name, 

63 "status": self.status, 

64 "start_ms": f"{self.start_ms:.1f}", 

65 "end_ms": f"{self.end_ms:.1f}" if self.end_ms else "-", 

66 "duration_ms": f"{self.duration_ms:.1f}", 

67 "data": {k: str(v)[:100] for k, v in self.data.items()}, 

68 "tags": self.tags, 

69 "children": [c.to_dict() for c in self.children] if self.children else [], 

70 } 

71 

72 def to_flat_list(self) -> list[dict]: 

73 """Flatten tree to list for tabular display.""" 

74 rows = [self.to_dict()] 

75 for child in self.children: 

76 rows.extend(child.to_flat_list()) 

77 return rows 

78 

79 

80@dataclass 

81class ExecutionTrace: 

82 """ 

83 Full execution trace for a single task execution. 

84 

85 Captures a tree of spans representing every step: decomposition, 

86 sub-task execution, fusion, retries, gate checks, etc. 

87 

88 Usage: 

89 trace = ExecutionTrace(task_name="research_query") 

90 

91 span = trace.start_span(TraceEvent.TASK_START, name="main") 

92 # ... do work ... 

93 trace.end_span(span.id, status="done") 

94 

95 print(trace.summary()) 

96 print(trace.to_json()) 

97 """ 

98 

99 id: str = field(default_factory=lambda: uuid.uuid4().hex[:12]) 

100 task_name: str = "" 

101 root_span: TraceSpan | None = None 

102 _span_map: dict[str, TraceSpan] = field(default_factory=dict) 

103 _current_spans: list[str] = field(default_factory=list) # stack 

104 total_spans: int = 0 

105 total_retries: int = 0 

106 total_errors: int = 0 

107 total_fallbacks: int = 0 

108 created_at: float = field(default_factory=time.time) 

109 

110 def start_span( 

111 self, 

112 event: TraceEvent, 

113 name: str = "", 

114 data: dict | None = None, 

115 tags: list[str] | None = None, 

116 ) -> TraceSpan: 

117 """Start a new span and add it to the trace tree.""" 

118 span = TraceSpan( 

119 event=event, 

120 name=name, 

121 data=data or {}, 

122 tags=tags or [], 

123 ) 

124 

125 # Determine parent 

126 if self._current_spans: 

127 span.parent_id = self._current_spans[-1] 

128 parent = self._span_map.get(span.parent_id) 

129 if parent: 

130 parent.children.append(span) 

131 elif self.root_span is None: 

132 self.root_span = span 

133 else: 

134 # Attach to root as sibling 

135 span.parent_id = self.root_span.id 

136 self.root_span.children.append(span) 

137 

138 self._span_map[span.id] = span 

139 self._current_spans.append(span.id) 

140 self.total_spans += 1 

141 

142 if event == TraceEvent.RETRY: 

143 self.total_retries += 1 

144 elif event == TraceEvent.ERROR: 

145 self.total_errors += 1 

146 elif event == TraceEvent.FALLBACK: 

147 self.total_fallbacks += 1 

148 

149 return span 

150 

151 def end_span(self, span_id: str, status: str = "done", data: dict | None = None) -> TraceSpan | None: 

152 """End a span and record its duration.""" 

153 span = self._span_map.get(span_id) 

154 if not span: 

155 return None 

156 

157 span.end_ms = time.time() * 1000 

158 span.duration_ms = span.end_ms - span.start_ms 

159 span.status = status 

160 if data: 

161 span.data.update(data) 

162 

163 # Pop from current stack 

164 if self._current_spans and self._current_spans[-1] == span_id: 

165 self._current_spans.pop() 

166 

167 return span 

168 

169 def add_event( 

170 self, 

171 event: TraceEvent, 

172 name: str = "", 

173 data: dict | None = None, 

174 tags: list[str] | None = None, 

175 duration_ms: float = 0.0, 

176 ) -> TraceSpan: 

177 """Quick-add a leaf event (start+end in one call).""" 

178 span = self.start_span(event, name, data, tags) 

179 span.end_ms = span.start_ms + duration_ms 

180 span.duration_ms = duration_ms 

181 span.status = "done" 

182 

183 if self._current_spans and self._current_spans[-1] == span.id: 

184 self._current_spans.pop() 

185 

186 return span 

187 

188 def to_dict(self) -> dict: 

189 return { 

190 "trace_id": self.id, 

191 "task_name": self.task_name, 

192 "total_spans": self.total_spans, 

193 "total_retries": self.total_retries, 

194 "total_errors": self.total_errors, 

195 "total_fallbacks": self.total_fallbacks, 

196 "created_at": self.created_at, 

197 "root": self.root_span.to_dict() if self.root_span else {}, 

198 } 

199 

200 def to_json(self, indent: int = 2) -> str: 

201 return json.dumps(self.to_dict(), indent=indent, ensure_ascii=False) 

202 

203 def to_tree_string(self, span: TraceSpan | None = None, indent: int = 0) -> str: 

204 """Render trace as indented ASCII tree.""" 

205 if span is None: 

206 span = self.root_span 

207 if span is None: 

208 return "(empty trace)" 

209 

210 lines = [] 

211 prefix = " " * indent 

212 status_icon = {"done": "O", "failed": "X", "started": ">", "aborted": "!"}.get(span.status, "?") 

213 lines.append( 

214 f"{prefix}{status_icon} [{span.event.value}] {span.name} " 

215 f"({span.duration_ms:.0f}ms) [{span.status}]" 

216 ) 

217 for child in span.children: 

218 lines.append(self.to_tree_string(child, indent + 1)) 

219 return "\n".join(lines) 

220 

221 def summary(self) -> str: 

222 """One-line summary of trace.""" 

223 total_ms = 0.0 

224 if self.root_span and self.root_span.duration_ms: 

225 total_ms = self.root_span.duration_ms 

226 return ( 

227 f"Trace[{self.id}] '{self.task_name}' " 

228 f"{self.total_spans} spans, " 

229 f"{self.total_retries} retries, " 

230 f"{self.total_errors} errors, " 

231 f"{total_ms:.0f}ms total" 

232 ) 

233 

234 def bottlenecks(self, top_n: int = 5) -> list[dict]: 

235 """Find slowest spans — helps identify bottlenecks.""" 

236 all_spans: list[TraceSpan] = [] 

237 

238 def collect(s: TraceSpan): 

239 all_spans.append(s) 

240 for c in s.children: 

241 collect(c) 

242 

243 if self.root_span: 

244 collect(self.root_span) 

245 

246 sorted_spans = sorted(all_spans, key=lambda s: s.duration_ms, reverse=True) 

247 result = [] 

248 for s in sorted_spans[:top_n]: 

249 result.append({ 

250 "name": s.name, 

251 "event": s.event.value, 

252 "duration_ms": round(s.duration_ms, 1), 

253 "status": s.status, 

254 "tags": s.tags, 

255 }) 

256 return result 

257 

258 def errors_list(self) -> list[dict]: 

259 """List all error spans.""" 

260 errors: list[dict] = [] 

261 

262 def collect(s: TraceSpan): 

263 if s.event == TraceEvent.ERROR or s.status == "failed": 

264 errors.append({ 

265 "name": s.name, 

266 "data": {k: str(v)[:100] for k, v in s.data.items()}, 

267 "tags": s.tags, 

268 }) 

269 for c in s.children: 

270 collect(c) 

271 

272 if self.root_span: 

273 collect(self.root_span) 

274 

275 return errors 

276 

277 def timeline(self) -> list[dict]: 

278 """Generate a flat timeline of all spans sorted by start time.""" 

279 all_spans: list[TraceSpan] = [] 

280 

281 def collect(s: TraceSpan): 

282 all_spans.append(s) 

283 for c in s.children: 

284 collect(c) 

285 

286 if self.root_span: 

287 collect(self.root_span) 

288 

289 all_spans.sort(key=lambda s: s.start_ms) 

290 

291 timeline = [] 

292 for s in all_spans: 

293 timeline.append({ 

294 "time_ms": f"{s.start_ms:.1f}", 

295 "event": s.event.value, 

296 "name": s.name, 

297 "duration_ms": f"{s.duration_ms:.1f}", 

298 "status": s.status, 

299 }) 

300 return timeline 

301 

302 

303class TraceCollector: 

304 """ 

305 Collects multiple execution traces and generates aggregate reports. 

306 

307 Usage: 

308 collector = TraceCollector() 

309 trace1 = await run_task("query_a") 

310 collector.add(trace1) 

311 trace2 = await run_task("query_b") 

312 collector.add(trace2) 

313 

314 print(collector.stats()) 

315 """ 

316 

317 def __init__(self, max_traces: int = 100): 

318 self._traces: dict[str, ExecutionTrace] = {} 

319 self.max_traces = max_traces 

320 

321 def add(self, trace: ExecutionTrace) -> None: 

322 self._traces[trace.id] = trace 

323 if len(self._traces) > self.max_traces: 

324 oldest = next(iter(self._traces)) 

325 del self._traces[oldest] 

326 

327 def get(self, trace_id: str) -> ExecutionTrace | None: 

328 return self._traces.get(trace_id) 

329 

330 def stats(self) -> dict: 

331 """Aggregate statistics across all traces.""" 

332 if not self._traces: 

333 return {"count": 0} 

334 

335 total_spans = sum(t.total_spans for t in self._traces.values()) 

336 total_retries = sum(t.total_retries for t in self._traces.values()) 

337 total_errors = sum(t.total_errors for t in self._traces.values()) 

338 total_fallbacks = sum(t.total_fallbacks for t in self._traces.values()) 

339 

340 durations = [] 

341 for t in self._traces.values(): 

342 if t.root_span and t.root_span.duration_ms: 

343 durations.append(t.root_span.duration_ms) 

344 

345 return { 

346 "count": len(self._traces), 

347 "total_spans": total_spans, 

348 "total_retries": total_retries, 

349 "total_errors": total_errors, 

350 "total_fallbacks": total_fallbacks, 

351 "retry_rate": round(total_retries / max(total_spans, 1), 3), 

352 "error_rate": round(total_errors / max(total_spans, 1), 3), 

353 "avg_duration_ms": round(sum(durations) / len(durations), 1) if durations else 0, 

354 "max_duration_ms": round(max(durations), 1) if durations else 0, 

355 "min_duration_ms": round(min(durations), 1) if durations else 0, 

356 } 

357 

358 def failed_tasks(self) -> list[dict]: 

359 """List tasks that ended with errors.""" 

360 failed = [] 

361 for t in self._traces.values(): 

362 if t.root_span and t.root_span.status in ("failed", "aborted"): 

363 failed.append({ 

364 "trace_id": t.id, 

365 "task_name": t.task_name, 

366 "status": t.root_span.status, 

367 "errors": t.total_errors, 

368 }) 

369 return failed