Coverage for agentos/swarm/execution_trace.py: 36%
184 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 21:26 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 21:26 +0800
1"""
2v1.9.6: Execution Trace — full observability into Agent task execution.
4Captures every sub-task, gate evaluation, retry, and timing detail.
5Supports timeline visualization, bottleneck detection, and debugging.
6"""
8from __future__ import annotations
10import json
11import time
12import uuid
13from dataclasses import dataclass, field
14from enum import StrEnum
15from typing import Any
18class TraceEvent(StrEnum):
19 """Event types in an execution trace."""
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"
37@dataclass
38class TraceSpan:
39 """A single span in an execution trace."""
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)
53 @property
54 def is_leaf(self) -> bool:
55 return len(self.children) == 0
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 }
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
80@dataclass
81class ExecutionTrace:
82 """
83 Full execution trace for a single task execution.
85 Captures a tree of spans representing every step: decomposition,
86 sub-task execution, fusion, retries, gate checks, etc.
88 Usage:
89 trace = ExecutionTrace(task_name="research_query")
91 span = trace.start_span(TraceEvent.TASK_START, name="main")
92 # ... do work ...
93 trace.end_span(span.id, status="done")
95 print(trace.summary())
96 print(trace.to_json())
97 """
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)
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 )
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)
138 self._span_map[span.id] = span
139 self._current_spans.append(span.id)
140 self.total_spans += 1
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
149 return span
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
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)
165 # Pop from current stack
166 if self._current_spans and self._current_spans[-1] == span_id:
167 self._current_spans.pop()
169 return span
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"
185 if self._current_spans and self._current_spans[-1] == span.id:
186 self._current_spans.pop()
188 return span
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 }
202 def to_json(self, indent: int = 2) -> str:
203 return json.dumps(self.to_dict(), indent=indent, ensure_ascii=False)
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)"
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)
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 )
238 def bottlenecks(self, top_n: int = 5) -> list[dict]:
239 """Find slowest spans — helps identify bottlenecks."""
240 all_spans: list[TraceSpan] = []
242 def collect(s: TraceSpan):
243 all_spans.append(s)
244 for c in s.children:
245 collect(c)
247 if self.root_span:
248 collect(self.root_span)
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
264 def errors_list(self) -> list[dict]:
265 """List all error spans."""
266 errors: list[dict] = []
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)
280 if self.root_span:
281 collect(self.root_span)
283 return errors
285 def timeline(self) -> list[dict]:
286 """Generate a flat timeline of all spans sorted by start time."""
287 all_spans: list[TraceSpan] = []
289 def collect(s: TraceSpan):
290 all_spans.append(s)
291 for c in s.children:
292 collect(c)
294 if self.root_span:
295 collect(self.root_span)
297 all_spans.sort(key=lambda s: s.start_ms)
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
313class TraceCollector:
314 """
315 Collects multiple execution traces and generates aggregate reports.
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)
324 print(collector.stats())
325 """
327 def __init__(self, max_traces: int = 100):
328 self._traces: dict[str, ExecutionTrace] = {}
329 self.max_traces = max_traces
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]
337 def get(self, trace_id: str) -> ExecutionTrace | None:
338 return self._traces.get(trace_id)
340 def stats(self) -> dict:
341 """Aggregate statistics across all traces."""
342 if not self._traces:
343 return {"count": 0}
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())
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)
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 }
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