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

184 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-08 20:40 +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 StrEnum 

15from typing import Any 

16 

17 

18class TraceEvent(StrEnum): 

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( 

152 self, span_id: str, status: str = "done", data: dict | None = None 

153 ) -> TraceSpan | None: 

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

155 span = self._span_map.get(span_id) 

156 if not span: 

157 return None 

158 

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

160 span.duration_ms = span.end_ms - span.start_ms 

161 span.status = status 

162 if data: 

163 span.data.update(data) 

164 

165 # Pop from current stack 

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

167 self._current_spans.pop() 

168 

169 return span 

170 

171 def add_event( 

172 self, 

173 event: TraceEvent, 

174 name: str = "", 

175 data: dict | None = None, 

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

177 duration_ms: float = 0.0, 

178 ) -> TraceSpan: 

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

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

181 span.end_ms = span.start_ms + duration_ms 

182 span.duration_ms = duration_ms 

183 span.status = "done" 

184 

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

186 self._current_spans.pop() 

187 

188 return span 

189 

190 def to_dict(self) -> dict: 

191 return { 

192 "trace_id": self.id, 

193 "task_name": self.task_name, 

194 "total_spans": self.total_spans, 

195 "total_retries": self.total_retries, 

196 "total_errors": self.total_errors, 

197 "total_fallbacks": self.total_fallbacks, 

198 "created_at": self.created_at, 

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

200 } 

201 

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

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

204 

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

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

207 if span is None: 

208 span = self.root_span 

209 if span is None: 

210 return "(empty trace)" 

211 

212 lines = [] 

213 prefix = " " * indent 

214 status_icon = {"done": "O", "failed": "X", "started": ">", "aborted": "!"}.get( 

215 span.status, "?" 

216 ) 

217 lines.append( 

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

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

220 ) 

221 for child in span.children: 

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

223 return "\n".join(lines) 

224 

225 def summary(self) -> str: 

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

227 total_ms = 0.0 

228 if self.root_span and self.root_span.duration_ms: 

229 total_ms = self.root_span.duration_ms 

230 return ( 

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

232 f"{self.total_spans} spans, " 

233 f"{self.total_retries} retries, " 

234 f"{self.total_errors} errors, " 

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

236 ) 

237 

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

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

240 all_spans: list[TraceSpan] = [] 

241 

242 def collect(s: TraceSpan): 

243 all_spans.append(s) 

244 for c in s.children: 

245 collect(c) 

246 

247 if self.root_span: 

248 collect(self.root_span) 

249 

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

251 result = [] 

252 for s in sorted_spans[:top_n]: 

253 result.append( 

254 { 

255 "name": s.name, 

256 "event": s.event.value, 

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

258 "status": s.status, 

259 "tags": s.tags, 

260 } 

261 ) 

262 return result 

263 

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

265 """List all error spans.""" 

266 errors: list[dict] = [] 

267 

268 def collect(s: TraceSpan): 

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

270 errors.append( 

271 { 

272 "name": s.name, 

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

274 "tags": s.tags, 

275 } 

276 ) 

277 for c in s.children: 

278 collect(c) 

279 

280 if self.root_span: 

281 collect(self.root_span) 

282 

283 return errors 

284 

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

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

287 all_spans: list[TraceSpan] = [] 

288 

289 def collect(s: TraceSpan): 

290 all_spans.append(s) 

291 for c in s.children: 

292 collect(c) 

293 

294 if self.root_span: 

295 collect(self.root_span) 

296 

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

298 

299 timeline = [] 

300 for s in all_spans: 

301 timeline.append( 

302 { 

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

304 "event": s.event.value, 

305 "name": s.name, 

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

307 "status": s.status, 

308 } 

309 ) 

310 return timeline 

311 

312 

313class TraceCollector: 

314 """ 

315 Collects multiple execution traces and generates aggregate reports. 

316 

317 Usage: 

318 collector = TraceCollector() 

319 trace1 = await run_task("query_a") 

320 collector.add(trace1) 

321 trace2 = await run_task("query_b") 

322 collector.add(trace2) 

323 

324 print(collector.stats()) 

325 """ 

326 

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

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

329 self.max_traces = max_traces 

330 

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

332 self._traces[trace.id] = trace 

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

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

335 del self._traces[oldest] 

336 

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

338 return self._traces.get(trace_id) 

339 

340 def stats(self) -> dict: 

341 """Aggregate statistics across all traces.""" 

342 if not self._traces: 

343 return {"count": 0} 

344 

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

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

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

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

349 

350 durations = [] 

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

352 if t.root_span and t.root_span.duration_ms: 

353 durations.append(t.root_span.duration_ms) 

354 

355 return { 

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

357 "total_spans": total_spans, 

358 "total_retries": total_retries, 

359 "total_errors": total_errors, 

360 "total_fallbacks": total_fallbacks, 

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

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

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

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

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

366 } 

367 

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

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

370 failed = [] 

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

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

373 failed.append( 

374 { 

375 "trace_id": t.id, 

376 "task_name": t.task_name, 

377 "status": t.root_span.status, 

378 "errors": t.total_errors, 

379 } 

380 ) 

381 return failed