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
« 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.
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 Enum
15from typing import Any
18class TraceEvent(str, Enum):
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(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
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)
163 # Pop from current stack
164 if self._current_spans and self._current_spans[-1] == span_id:
165 self._current_spans.pop()
167 return span
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"
183 if self._current_spans and self._current_spans[-1] == span.id:
184 self._current_spans.pop()
186 return span
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 }
200 def to_json(self, indent: int = 2) -> str:
201 return json.dumps(self.to_dict(), indent=indent, ensure_ascii=False)
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)"
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)
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 )
234 def bottlenecks(self, top_n: int = 5) -> list[dict]:
235 """Find slowest spans — helps identify bottlenecks."""
236 all_spans: list[TraceSpan] = []
238 def collect(s: TraceSpan):
239 all_spans.append(s)
240 for c in s.children:
241 collect(c)
243 if self.root_span:
244 collect(self.root_span)
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
258 def errors_list(self) -> list[dict]:
259 """List all error spans."""
260 errors: list[dict] = []
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)
272 if self.root_span:
273 collect(self.root_span)
275 return errors
277 def timeline(self) -> list[dict]:
278 """Generate a flat timeline of all spans sorted by start time."""
279 all_spans: list[TraceSpan] = []
281 def collect(s: TraceSpan):
282 all_spans.append(s)
283 for c in s.children:
284 collect(c)
286 if self.root_span:
287 collect(self.root_span)
289 all_spans.sort(key=lambda s: s.start_ms)
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
303class TraceCollector:
304 """
305 Collects multiple execution traces and generates aggregate reports.
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)
314 print(collector.stats())
315 """
317 def __init__(self, max_traces: int = 100):
318 self._traces: dict[str, ExecutionTrace] = {}
319 self.max_traces = max_traces
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]
327 def get(self, trace_id: str) -> ExecutionTrace | None:
328 return self._traces.get(trace_id)
330 def stats(self) -> dict:
331 """Aggregate statistics across all traces."""
332 if not self._traces:
333 return {"count": 0}
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())
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)
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 }
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