deer-agent-framework 0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (52) hide show
  1. deer/__init__.py +36 -0
  2. deer/builtins/__init__.py +7 -0
  3. deer/builtins/python_manager/agent.py +29 -0
  4. deer/builtins/python_manager/tools.py +54 -0
  5. deer/core/__init__.py +1 -0
  6. deer/core/agent.py +463 -0
  7. deer/core/ui.py +31 -0
  8. deer/drivers/__init__.py +66 -0
  9. deer/drivers/base_driver.py +56 -0
  10. deer/drivers/gemini_driver.py +62 -0
  11. deer/drivers/ollama_driver.py +69 -0
  12. deer/executor/__init__.py +1 -0
  13. deer/executor/executor.py +168 -0
  14. deer/executor/logic.py +75 -0
  15. deer/executor/logic_secure.py +102 -0
  16. deer/main.py +71 -0
  17. deer/planner/__init__.py +1 -0
  18. deer/planner/planner.py +71 -0
  19. deer/prompts/__init__.py +6 -0
  20. deer/prompts/error_explain.py +30 -0
  21. deer/prompts/goal_improvement.py +65 -0
  22. deer/prompts/goal_validation.py +31 -0
  23. deer/prompts/humanizer.py +20 -0
  24. deer/prompts/planner.py +92 -0
  25. deer/prompts/response_improvement.py +23 -0
  26. deer/schema/__init__.py +2 -0
  27. deer/schema/io.py +40 -0
  28. deer/schema/plan.py +75 -0
  29. deer/tools/__init__.py +4 -0
  30. deer/tools/base.py +141 -0
  31. deer/tools/builtin/__init__.py +3 -0
  32. deer/tools/builtin/file_manager.py +136 -0
  33. deer/tools/builtin/git_manager.py +67 -0
  34. deer/tools/builtin/search_manager.py +123 -0
  35. deer/tools/decorators.py +127 -0
  36. deer/tools/registry.py +114 -0
  37. deer/tracing/__init__.py +2 -0
  38. deer/tracing/logging_config.py +20 -0
  39. deer/tracing/store.py +21 -0
  40. deer/utils/__init__.py +0 -0
  41. deer/utils/console.py +11 -0
  42. deer/utils/plots/__init__.py +7 -0
  43. deer/utils/plots/plot_traces.py +1034 -0
  44. deer/validator/__init__.py +1 -0
  45. deer/validator/plan_validator.py +23 -0
  46. deer/validator/rules.py +98 -0
  47. deer_agent_framework-0.0.dist-info/METADATA +163 -0
  48. deer_agent_framework-0.0.dist-info/RECORD +52 -0
  49. deer_agent_framework-0.0.dist-info/WHEEL +5 -0
  50. deer_agent_framework-0.0.dist-info/entry_points.txt +2 -0
  51. deer_agent_framework-0.0.dist-info/licenses/LICENSE +24 -0
  52. deer_agent_framework-0.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,1034 @@
1
+ import os
2
+ import pickle
3
+ from collections import Counter
4
+ from collections.abc import Iterable
5
+ from typing import Any
6
+
7
+ import matplotlib.colors as mcolors
8
+ import matplotlib.patches as patches
9
+ import matplotlib.pyplot as plt
10
+ from matplotlib.path import Path
11
+
12
+ import pandas as pd
13
+
14
+ NODE_WIDTH: int = 11
15
+ MIN_CONTAINER_HEIGHT: int = 3
16
+ MAX_CONTAINER_HEIGHT: int = 10
17
+ ROW_GAP: int = 0
18
+
19
+
20
+ def scale_range(
21
+ value: float,
22
+ input_min: float,
23
+ input_max: float,
24
+ output_min: float,
25
+ output_max: float,
26
+ ) -> float:
27
+ """Scale a value from one range into another.
28
+
29
+ Parameters
30
+ ----------
31
+ value : float
32
+ Value to scale.
33
+ input_min : float
34
+ Minimum input range value.
35
+ input_max : float
36
+ Maximum input range value.
37
+ output_min : float
38
+ Minimum output range value.
39
+ output_max : float
40
+ Maximum output range value.
41
+
42
+ Returns
43
+ -------
44
+ float
45
+ Scaled value.
46
+
47
+ Examples
48
+ --------
49
+ >>> scale_range(5, 0, 10, 0, 100)
50
+ 50.0
51
+ """
52
+ if input_min == input_max:
53
+ # PSS: Prevent division by zero.
54
+ return output_min
55
+
56
+ return (
57
+ (value - input_min) * (output_max - output_min) / (input_max - input_min)
58
+ ) + output_min
59
+
60
+
61
+ def hide_axis_spines(axis: plt.Axes) -> None:
62
+ """Hide the top and right spines of a Matplotlib axis.
63
+
64
+ Parameters
65
+ ----------
66
+ axis : matplotlib.axes.Axes
67
+ Axis where the spines will be hidden.
68
+
69
+ Returns
70
+ -------
71
+ None
72
+ """
73
+ axis.spines["top"].set_visible(False)
74
+ axis.spines["right"].set_visible(False)
75
+
76
+
77
+ def draw_bezier_flow(
78
+ axis: plt.Axes,
79
+ source_x: float,
80
+ target_x: float,
81
+ source_y: float,
82
+ source_height: float,
83
+ target_y: float,
84
+ target_height: float,
85
+ color: str,
86
+ alpha: float = 0.35,
87
+ ) -> None:
88
+ """Draw a Bézier flow between two Sankey nodes.
89
+
90
+ Parameters
91
+ ----------
92
+ axis : matplotlib.axes.Axes
93
+ Matplotlib axis.
94
+ source_x : float
95
+ Source x coordinate.
96
+ target_x : float
97
+ Target x coordinate.
98
+ source_y : float
99
+ Source y coordinate.
100
+ source_height : float
101
+ Height of the source flow.
102
+ target_y : float
103
+ Target y coordinate.
104
+ target_height : float
105
+ Height of the target flow.
106
+ color : str
107
+ Flow color.
108
+ alpha : float, optional
109
+ Flow transparency.
110
+ """
111
+ vertices = [
112
+ (source_x, source_y),
113
+ (
114
+ source_x + (target_x - source_x) / 2,
115
+ source_y,
116
+ ),
117
+ (
118
+ source_x + (target_x - source_x) / 2,
119
+ target_y,
120
+ ),
121
+ (target_x, target_y),
122
+ (target_x, target_y - target_height),
123
+ (
124
+ source_x + (target_x - source_x) / 2,
125
+ target_y - target_height,
126
+ ),
127
+ (
128
+ source_x + (target_x - source_x) / 2,
129
+ source_y - source_height,
130
+ ),
131
+ (source_x, source_y - source_height),
132
+ (source_x, source_y),
133
+ ]
134
+
135
+ path_codes = [
136
+ Path.MOVETO,
137
+ Path.CURVE4,
138
+ Path.CURVE4,
139
+ Path.CURVE4,
140
+ Path.LINETO,
141
+ Path.CURVE4,
142
+ Path.CURVE4,
143
+ Path.CURVE4,
144
+ Path.CLOSEPOLY,
145
+ ]
146
+
147
+ bezier_path = Path(vertices, path_codes)
148
+
149
+ axis.add_patch(
150
+ patches.PathPatch(
151
+ bezier_path,
152
+ facecolor=color,
153
+ edgecolor="none",
154
+ alpha=alpha,
155
+ )
156
+ )
157
+
158
+
159
+ def precompute_global_heights(
160
+ trace_histories: list[list[str]],
161
+ tool_names: list[str],
162
+ step_columns: list[str],
163
+ ) -> tuple[
164
+ dict[str, float],
165
+ dict[str, float],
166
+ float,
167
+ ]:
168
+ """Precompute consistent row heights across all columns.
169
+
170
+ Parameters
171
+ ----------
172
+ trace_histories : list[list[str]]
173
+ List of tool execution sequences.
174
+ tool_names : list[str]
175
+ Available tool names.
176
+ step_columns : list[str]
177
+ Sankey column labels.
178
+
179
+ Returns
180
+ -------
181
+ tuple
182
+ (
183
+ row_base_positions,
184
+ max_row_heights,
185
+ total_height,
186
+ )
187
+ """
188
+ node_usage_stats: dict[
189
+ tuple[str, str],
190
+ dict[str, int],
191
+ ] = {}
192
+
193
+ for column_name in step_columns:
194
+ for tool_name in tool_names:
195
+ node_usage_stats[(column_name, tool_name)] = {
196
+ "outputs": 0,
197
+ "inputs": 0,
198
+ }
199
+
200
+ flow_data: Counter[Any] = Counter()
201
+
202
+ for sequence in trace_histories:
203
+ effective_steps = min(
204
+ len(sequence),
205
+ len(step_columns),
206
+ )
207
+
208
+ for index in range(effective_steps - 1):
209
+ flow_data[
210
+ (
211
+ step_columns[index],
212
+ sequence[index],
213
+ step_columns[index + 1],
214
+ sequence[index + 1],
215
+ )
216
+ ] += 1
217
+
218
+ if not flow_data:
219
+ # PSS: Added protection for empty datasets.
220
+ return {}, {}, 0
221
+
222
+ max_flow_value = max(flow_data.values())
223
+ min_flow_value = min(flow_data.values())
224
+
225
+ for column_name in step_columns:
226
+ for tool_name in tool_names:
227
+ output_count = sum(
228
+ count
229
+ for (
230
+ source_column,
231
+ source_tool,
232
+ _,
233
+ _,
234
+ ), count in flow_data.items()
235
+ if source_column == column_name and source_tool == tool_name
236
+ )
237
+
238
+ input_count = sum(
239
+ count
240
+ for (
241
+ _,
242
+ _,
243
+ target_column,
244
+ target_tool,
245
+ ), count in flow_data.items()
246
+ if target_column == column_name and target_tool == tool_name
247
+ )
248
+
249
+ node_usage_stats[(column_name, tool_name)] = {
250
+ "outputs": output_count,
251
+ "inputs": input_count,
252
+ }
253
+
254
+ max_row_heights: dict[str, float] = {}
255
+
256
+ for tool_name in tool_names:
257
+ cell_values = []
258
+
259
+ for column_name in step_columns:
260
+ flow_value = (
261
+ node_usage_stats[(column_name, tool_name)]["outputs"]
262
+ if column_name == step_columns[0]
263
+ else node_usage_stats[(column_name, tool_name)]["inputs"]
264
+ )
265
+
266
+ scaled_value = scale_range(
267
+ flow_value,
268
+ min_flow_value,
269
+ max_flow_value,
270
+ MIN_CONTAINER_HEIGHT,
271
+ MAX_CONTAINER_HEIGHT,
272
+ )
273
+
274
+ cell_values.append(scaled_value)
275
+
276
+ max_row_heights[tool_name] = max(cell_values)
277
+
278
+ row_base_positions: dict[str, float] = {}
279
+ accumulated_y = 0
280
+
281
+ for tool_name in reversed(tool_names):
282
+ row_base_positions[tool_name] = accumulated_y
283
+
284
+ accumulated_y += max_row_heights[tool_name] + ROW_GAP
285
+
286
+ return (
287
+ row_base_positions,
288
+ max_row_heights,
289
+ accumulated_y + 0,
290
+ )
291
+
292
+
293
+ def draw_sankey_on_axis(
294
+ axis: plt.Axes,
295
+ trace_sequences: list[list[str]],
296
+ tool_names: list[str],
297
+ step_columns: list[str],
298
+ color_palette: dict[str, str],
299
+ chart_title: str,
300
+ row_base_positions: dict[str, float],
301
+ max_row_heights: dict[str, float],
302
+ global_max_height: float,
303
+ font_color: str = "#1F2937",
304
+ ) -> None:
305
+ """Render a Sankey diagram on a Matplotlib axis."""
306
+ column_positions = {
307
+ column_name: index * 16 for index, column_name in enumerate(step_columns)
308
+ }
309
+
310
+ flow_data: Counter[Any] = Counter()
311
+
312
+ for sequence in trace_sequences:
313
+ effective_steps = min(
314
+ len(sequence),
315
+ len(step_columns),
316
+ )
317
+
318
+ for index in range(effective_steps - 1):
319
+ flow_data[
320
+ (
321
+ step_columns[index],
322
+ sequence[index],
323
+ step_columns[index + 1],
324
+ sequence[index + 1],
325
+ )
326
+ ] += 1
327
+
328
+ if not flow_data:
329
+ return
330
+
331
+ max_flow_value = max(flow_data.values())
332
+ min_flow_value = min(flow_data.values())
333
+
334
+ node_usage: dict[
335
+ tuple[str, str],
336
+ dict[str, float],
337
+ ] = {}
338
+
339
+ for column_name in step_columns:
340
+ for tool_name in tool_names:
341
+ output_count = sum(
342
+ count
343
+ for (
344
+ source_column,
345
+ source_tool,
346
+ _,
347
+ _,
348
+ ), count in flow_data.items()
349
+ if source_column == column_name and source_tool == tool_name
350
+ )
351
+
352
+ input_count = sum(
353
+ count
354
+ for (
355
+ _,
356
+ _,
357
+ target_column,
358
+ target_tool,
359
+ ), count in flow_data.items()
360
+ if target_column == column_name and target_tool == tool_name
361
+ )
362
+
363
+ container_height = (
364
+ output_count if column_name == step_columns[0] else input_count
365
+ )
366
+
367
+ container_height = scale_range(
368
+ container_height,
369
+ min_flow_value,
370
+ max_flow_value,
371
+ MIN_CONTAINER_HEIGHT,
372
+ MAX_CONTAINER_HEIGHT,
373
+ )
374
+
375
+ node_usage[(column_name, tool_name)] = {
376
+ "outputs": output_count,
377
+ "inputs": input_count,
378
+ "container_height": container_height,
379
+ }
380
+
381
+ node_coordinates: dict[
382
+ tuple[str, str],
383
+ dict[str, float],
384
+ ] = {}
385
+
386
+ for column_name in step_columns:
387
+ x_position = column_positions[column_name]
388
+
389
+ for tool_name in tool_names:
390
+ row_base_y = row_base_positions[tool_name] - 3
391
+
392
+ row_max_height = max_row_heights[tool_name]
393
+
394
+ container_height = node_usage[(column_name, tool_name)]["container_height"]
395
+
396
+ container_base_y = row_base_y + (row_max_height - container_height) / 2
397
+
398
+ usage_data = node_usage[(column_name, tool_name)]
399
+
400
+ is_used = usage_data["inputs"] > 0 or usage_data["outputs"] > 0
401
+
402
+ if is_used:
403
+ background_rect = plt.Rectangle(
404
+ (
405
+ x_position - NODE_WIDTH / 2,
406
+ container_base_y,
407
+ ),
408
+ NODE_WIDTH,
409
+ container_height,
410
+ color="#D1D5DB",
411
+ alpha=0.4,
412
+ zorder=1,
413
+ )
414
+
415
+ axis.add_patch(background_rect)
416
+
417
+ tool_color = color_palette.get(
418
+ tool_name,
419
+ "#999999",
420
+ )
421
+
422
+ colored_rect = plt.Rectangle(
423
+ (
424
+ x_position - NODE_WIDTH / 2,
425
+ container_base_y,
426
+ ),
427
+ NODE_WIDTH,
428
+ container_height,
429
+ color=tool_color,
430
+ alpha=0.85,
431
+ zorder=2,
432
+ )
433
+
434
+ axis.add_patch(colored_rect)
435
+
436
+ node_coordinates[(column_name, tool_name)] = {
437
+ "x": x_position,
438
+ "top_y": (container_base_y + container_height),
439
+ "height": (container_height if is_used else 0),
440
+ "output_accumulator": 0,
441
+ "input_accumulator": 0,
442
+ }
443
+
444
+ center_y = row_base_y + row_max_height / 2
445
+
446
+ if is_used:
447
+ axis.text(
448
+ x_position,
449
+ center_y,
450
+ tool_name,
451
+ ha="center",
452
+ va="center",
453
+ fontsize=9,
454
+ fontweight="bold",
455
+ color=font_color,
456
+ zorder=3,
457
+ )
458
+
459
+ for (
460
+ source_column,
461
+ source_tool,
462
+ target_column,
463
+ target_tool,
464
+ ), flow_value in flow_data.items():
465
+ source_key = (
466
+ source_column,
467
+ source_tool,
468
+ )
469
+
470
+ target_key = (
471
+ target_column,
472
+ target_tool,
473
+ )
474
+
475
+ source_node = node_coordinates[source_key]
476
+ target_node = node_coordinates[target_key]
477
+
478
+ total_source_flow = (
479
+ node_usage[source_key]["outputs"]
480
+ if source_column == step_columns[0]
481
+ else node_usage[source_key]["inputs"]
482
+ )
483
+
484
+ total_target_flow = node_usage[target_key]["inputs"]
485
+
486
+ source_flow_height = (
487
+ (flow_value / total_source_flow) * source_node["height"]
488
+ if total_source_flow > 0
489
+ else 0
490
+ )
491
+
492
+ target_flow_height = (
493
+ (flow_value / total_target_flow) * target_node["height"]
494
+ if total_target_flow > 0
495
+ else 0
496
+ )
497
+
498
+ source_y = source_node["top_y"] - source_node["output_accumulator"]
499
+
500
+ target_y = target_node["top_y"] - target_node["input_accumulator"]
501
+
502
+ flow_color = color_palette.get(
503
+ source_tool,
504
+ "#999999",
505
+ )
506
+
507
+ draw_bezier_flow(
508
+ axis=axis,
509
+ source_x=(source_node["x"] + NODE_WIDTH / 2),
510
+ target_x=(target_node["x"] - NODE_WIDTH / 2),
511
+ source_y=source_y,
512
+ source_height=source_flow_height,
513
+ target_y=target_y,
514
+ target_height=target_flow_height,
515
+ color=flow_color,
516
+ )
517
+
518
+ source_node["output_accumulator"] += source_flow_height
519
+
520
+ target_node["input_accumulator"] += target_flow_height
521
+
522
+ axis.set_xlim(
523
+ -NODE_WIDTH,
524
+ max(column_positions.values()) + NODE_WIDTH,
525
+ )
526
+
527
+ axis.set_ylim(
528
+ -5,
529
+ global_max_height + 5,
530
+ )
531
+
532
+ for (
533
+ column_name,
534
+ x_position,
535
+ ) in column_positions.items():
536
+ axis.text(
537
+ x_position,
538
+ global_max_height,
539
+ column_name,
540
+ ha="center",
541
+ va="bottom",
542
+ fontsize=12,
543
+ fontweight="bold",
544
+ color="#111827",
545
+ alpha=0.5,
546
+ )
547
+
548
+ axis.set_title(
549
+ chart_title,
550
+ fontsize=20,
551
+ fontweight="bold",
552
+ color="#919191",
553
+ pad=10,
554
+ )
555
+
556
+ axis.axis("off")
557
+
558
+
559
+ def plot_steps(
560
+ figure: plt.Figure,
561
+ planning_traces: list[list[str]],
562
+ validation_traces: list[list[str]],
563
+ color_map_name: str = "tab10",
564
+ font_color: str = "#1F2937",
565
+ ) -> None:
566
+ """Plot planning and validation Sankey diagrams."""
567
+ color_map = plt.get_cmap(color_map_name)
568
+
569
+ planning_steps = [
570
+ f"STEP-{index}"
571
+ for index in range(
572
+ 1,
573
+ 1 + max(len(trace) for trace in planning_traces),
574
+ )
575
+ ]
576
+
577
+ planning_tools = ["Logic"] + list(
578
+ {tool for trace in planning_traces for tool in trace}
579
+ )
580
+
581
+ planning_palette = {
582
+ tool_name: mcolors.to_hex(color_map(index))
583
+ for index, tool_name in enumerate(planning_tools)
584
+ }
585
+
586
+ (
587
+ planning_positions,
588
+ planning_heights,
589
+ planning_max_height,
590
+ ) = precompute_global_heights(
591
+ planning_traces,
592
+ planning_tools,
593
+ planning_steps,
594
+ )
595
+
596
+ planning_axis = figure.add_subplot(211)
597
+
598
+ draw_sankey_on_axis(
599
+ axis=planning_axis,
600
+ trace_sequences=planning_traces,
601
+ tool_names=planning_tools,
602
+ step_columns=planning_steps,
603
+ color_palette=planning_palette,
604
+ chart_title="PLANIFICATION",
605
+ row_base_positions=planning_positions,
606
+ max_row_heights=planning_heights,
607
+ global_max_height=planning_max_height,
608
+ font_color=font_color,
609
+ )
610
+
611
+ validation_steps = [
612
+ f"STEP-{index}"
613
+ for index in range(
614
+ 1,
615
+ 1 + max(len(trace) for trace in validation_traces),
616
+ )
617
+ ]
618
+
619
+ validation_tools = ["Logic"] + list(
620
+ {tool for trace in validation_traces for tool in trace}
621
+ )
622
+
623
+ validation_palette = {
624
+ tool_name: mcolors.to_hex(color_map(index))
625
+ for index, tool_name in enumerate(validation_tools)
626
+ }
627
+
628
+ (
629
+ validation_positions,
630
+ validation_heights,
631
+ validation_max_height,
632
+ ) = precompute_global_heights(
633
+ validation_traces,
634
+ validation_tools,
635
+ validation_steps,
636
+ )
637
+
638
+ validation_axis = figure.add_subplot(212)
639
+
640
+ draw_sankey_on_axis(
641
+ axis=validation_axis,
642
+ trace_sequences=validation_traces,
643
+ tool_names=validation_tools,
644
+ step_columns=validation_steps,
645
+ color_palette=validation_palette,
646
+ chart_title="VALIDATION",
647
+ row_base_positions=validation_positions,
648
+ max_row_heights=validation_heights,
649
+ global_max_height=validation_max_height,
650
+ font_color=font_color,
651
+ )
652
+
653
+
654
+ def filter_structures(
655
+ trace_data: list[list[str]],
656
+ threshold: float = 0.95,
657
+ ) -> list[list[str]]:
658
+ """Filter structures according to frequency.
659
+
660
+ Parameters
661
+ ----------
662
+ trace_data : list[list[str]]
663
+ Trace sequences.
664
+ threshold : float, optional
665
+ Preserve the most frequent trajectories until covering at least `threshold` of all traces.
666
+
667
+ Returns
668
+ -------
669
+ list[list[str]]
670
+ Filtered traces.
671
+ """
672
+ if not trace_data:
673
+ return []
674
+
675
+ structure_counts = Counter(tuple(trace_sequence) for trace_sequence in trace_data)
676
+
677
+ total_traces = len(trace_data)
678
+ sorted_structures = structure_counts.most_common()
679
+ selected_structures: set[tuple[str, ...]] = set()
680
+ cumulative_probability = 0.0
681
+
682
+ for structure, count in sorted_structures:
683
+ probability = count / total_traces
684
+ selected_structures.add(structure)
685
+ cumulative_probability += probability
686
+ if cumulative_probability >= threshold:
687
+ break
688
+
689
+ return [
690
+ trace_sequence
691
+ for trace_sequence in trace_data
692
+ if tuple(trace_sequence) in selected_structures
693
+ ]
694
+
695
+
696
+ def get_plot_data(
697
+ trace_files: list[str],
698
+ threshold: float = 0.95,
699
+ ) -> tuple[
700
+ list[list[str]],
701
+ list[list[str]],
702
+ ]:
703
+ """Load and preprocess Sankey trace data.
704
+
705
+ Parameters
706
+ ----------
707
+ trace_files : list[str]
708
+ Trace file names.
709
+ base_path : str
710
+ Directory path.
711
+ threshold : float, optional
712
+ Frequency filter threshold.
713
+
714
+ Returns
715
+ -------
716
+ tuple[list[list[str]], list[list[str]]]
717
+ Planning and validation traces.
718
+ """
719
+ traces: list[Any] = []
720
+
721
+ for trace_file in trace_files:
722
+
723
+ if not os.path.exists(trace_file):
724
+ raise FileNotFoundError(f"File not found: {trace_file}")
725
+
726
+ with open(trace_file, "rb") as file:
727
+ traces.append(pickle.load(file))
728
+
729
+ def normalize_tool_name(
730
+ tool_name: str,
731
+ ) -> str:
732
+ """Normalize tool names."""
733
+ return tool_name.replace("_", " ").capitalize()
734
+
735
+ planning_sequences: list[list[str]] = []
736
+ validation_sequences: list[list[str]] = []
737
+
738
+ for trace in traces:
739
+ planning_steps: list[str] = []
740
+ validation_steps: list[str] = []
741
+
742
+ for trace_item in trace["trace"]:
743
+ planning_steps.extend(
744
+ [
745
+ normalize_tool_name(item.tool)
746
+ for item in trace_item["solution_trace"]
747
+ ]
748
+ )
749
+
750
+ validation_steps.extend(
751
+ [
752
+ normalize_tool_name(item.tool)
753
+ for item in trace_item["verification_trace"]
754
+ ]
755
+ )
756
+
757
+ planning_sequences.append(planning_steps)
758
+
759
+ validation_sequences.append(validation_steps)
760
+
761
+ # PSS: Replaced filter/lambda with clearer comprehensions.
762
+ planning_sequences = [
763
+ sequence for sequence in planning_sequences if len(sequence) > 2
764
+ ]
765
+
766
+ validation_sequences = [
767
+ sequence for sequence in validation_sequences if len(sequence) > 2
768
+ ]
769
+
770
+ planning_sequences = filter_structures(
771
+ planning_sequences,
772
+ threshold,
773
+ )
774
+
775
+ validation_sequences = filter_structures(
776
+ validation_sequences,
777
+ threshold,
778
+ )
779
+
780
+ return (
781
+ planning_sequences,
782
+ validation_sequences,
783
+ )
784
+
785
+
786
+ def count_trace_frequencies(
787
+ trace_sequences: Iterable[Iterable[Any]],
788
+ ) -> tuple[list[str], list[float]]:
789
+ """Count repeated trace sequences and compute relative frequencies.
790
+
791
+ Parameters
792
+ ----------
793
+ trace_sequences : Iterable[Iterable[Any]]
794
+ Collection of trace sequences to analyze.
795
+
796
+ Returns
797
+ -------
798
+ tuple[list[str], list[float]]
799
+ A tuple containing:
800
+ - A list of trace labels.
801
+ - A list of relative frequencies in percentage.
802
+ """
803
+ trace_frequency_counter = Counter(
804
+ tuple(trace_sequence) for trace_sequence in trace_sequences
805
+ )
806
+
807
+ sorted_trace_frequencies = trace_frequency_counter.most_common()
808
+
809
+ trace_labels = [
810
+ f"Path {index + 1}" for index in range(len(sorted_trace_frequencies))
811
+ ]
812
+
813
+ absolute_frequencies = [
814
+ frequency_count for _, frequency_count in sorted_trace_frequencies
815
+ ]
816
+
817
+ total_frequency_count = sum(absolute_frequencies)
818
+
819
+ relative_frequencies = [
820
+ (100 * frequency_count / total_frequency_count)
821
+ for frequency_count in absolute_frequencies
822
+ ]
823
+
824
+ return (
825
+ trace_labels,
826
+ relative_frequencies,
827
+ )
828
+
829
+
830
+ def draw_trace_frequency_bars(
831
+ figure: Any,
832
+ trace_datasets: tuple[
833
+ Iterable[Iterable[Any]],
834
+ Iterable[Iterable[Any]],
835
+ ],
836
+ ) -> None:
837
+ """Draw bar charts for planning and validation trace frequencies.
838
+
839
+ Parameters
840
+ ----------
841
+ figure : Any
842
+ Figure-like object containing the subplot method.
843
+ trace_datasets : tuple[
844
+ Iterable[Iterable[Any]],
845
+ Iterable[Iterable[Any]],
846
+ ]
847
+ Tuple containing:
848
+ - Planning trace data.
849
+ - Validation trace data.
850
+
851
+ Returns
852
+ -------
853
+ None
854
+ """
855
+ if len(trace_datasets) != 2:
856
+ msg = "trace_datasets must contain exactly " "two datasets."
857
+ raise ValueError(msg)
858
+
859
+ planning_trace_sequences, validation_trace_sequences = trace_datasets
860
+
861
+ # PSS: Renamed subplot variables and functions
862
+ # to preserve Sankey/trace visualization context.
863
+
864
+ # Plot planning trace frequencies.
865
+ planning_axis = figure.add_subplot(121)
866
+ hide_axis_spines(planning_axis)
867
+
868
+ plt.grid(
869
+ True,
870
+ zorder=-1,
871
+ )
872
+
873
+ (
874
+ planning_trace_labels,
875
+ planning_trace_percentages,
876
+ ) = count_trace_frequencies(planning_trace_sequences)
877
+
878
+ plt.bar(
879
+ planning_trace_labels,
880
+ planning_trace_percentages,
881
+ zorder=99,
882
+ color="C0",
883
+ )
884
+
885
+ plt.title("Planning trace frequencies")
886
+ plt.ylabel("Percentage (%)")
887
+
888
+ # Plot validation trace frequencies.
889
+ validation_axis = figure.add_subplot(122)
890
+ hide_axis_spines(validation_axis)
891
+
892
+ plt.grid(
893
+ True,
894
+ zorder=-1,
895
+ )
896
+
897
+ (
898
+ validation_trace_labels,
899
+ validation_trace_percentages,
900
+ ) = count_trace_frequencies(validation_trace_sequences)
901
+
902
+ plt.bar(
903
+ validation_trace_labels,
904
+ validation_trace_percentages,
905
+ zorder=99,
906
+ color="C1",
907
+ )
908
+
909
+ plt.title("Validation trace frequencies")
910
+
911
+
912
+ def get_execution_metrics(trace_file_paths: list[str]) -> dict[str, list[int]]:
913
+ """Extract execution metrics from trace files.
914
+
915
+ Parameters
916
+ ----------
917
+ trace_file_paths : list[str]
918
+ List of file paths containing serialized trace data.
919
+
920
+ Returns
921
+ -------
922
+ dict[str, list[int]]
923
+ Dictionary containing:
924
+ - attempts: Number of attempts per task.
925
+ - solutions: Number of solution steps per task.
926
+ - verifications: Number of verification steps per task.
927
+ """
928
+ execution_metrics = {
929
+ "attempts": [],
930
+ "solutions": [],
931
+ "verifications": [],
932
+ }
933
+
934
+ processed_tracks = 0
935
+
936
+ for trace_file_path in trace_file_paths:
937
+ with open(trace_file_path, "rb") as file:
938
+ trace_data = pickle.load(file)
939
+
940
+ if not trace_data["trace"]:
941
+ print(f"Empty trace found: {trace_file_path}")
942
+ continue
943
+
944
+ processed_tracks += 1
945
+
946
+ execution_summary = trace_data["trace"][0]["execution_summary"]
947
+
948
+ solution_count = 0
949
+ verification_count = 0
950
+
951
+ for attempt_id in execution_summary:
952
+ attempt_data = execution_summary[attempt_id]
953
+
954
+ solution_count += len(
955
+ [
956
+ step_name
957
+ for step_name in attempt_data.keys()
958
+ if step_name.startswith("Solution")
959
+ ]
960
+ )
961
+
962
+ verification_count += len(
963
+ [
964
+ step_name
965
+ for step_name in attempt_data.keys()
966
+ if step_name.startswith("Verification")
967
+ ]
968
+ )
969
+
970
+ execution_metrics["attempts"].append(len(execution_summary))
971
+ execution_metrics["solutions"].append(solution_count)
972
+ execution_metrics["verifications"].append(verification_count)
973
+
974
+ return execution_metrics
975
+
976
+
977
+ def plot_execution_profile(figure, execution_metrics: dict[str, list[int]]) -> None:
978
+ """Plot execution complexity and retry metrics.
979
+
980
+ Parameters
981
+ ----------
982
+ execution_metrics : dict[str, list[int]]
983
+ Dictionary containing attempts, solutions, and verifications.
984
+
985
+ Returns
986
+ -------
987
+ None
988
+ """
989
+ metrics_dataframe = pd.DataFrame(execution_metrics)
990
+
991
+ primary_axis = figure.add_subplot(111)
992
+
993
+ # PSS: Renamed variables for clarity and improved readability.
994
+ metrics_dataframe[["solutions", "verifications"]].plot(
995
+ kind="bar",
996
+ stacked=True,
997
+ ax=primary_axis,
998
+ color=["#3498db", "#9b59b6"],
999
+ zorder=10,
1000
+ )
1001
+
1002
+ max_steps = (
1003
+ metrics_dataframe["solutions"] + metrics_dataframe["verifications"]
1004
+ ).max()
1005
+
1006
+ primary_axis.set_ylabel("Number of Steps")
1007
+ primary_axis.set_yticks(range(1, max_steps + 1))
1008
+ primary_axis.grid(True, axis="y", zorder=0)
1009
+
1010
+ secondary_axis = primary_axis.twinx()
1011
+
1012
+ secondary_axis.plot(
1013
+ metrics_dataframe.index,
1014
+ metrics_dataframe["attempts"],
1015
+ color="#e74c3c",
1016
+ marker="o",
1017
+ linewidth=2,
1018
+ label="Global Attempts",
1019
+ zorder=10,
1020
+ )
1021
+
1022
+ secondary_axis.set_ylabel("Global Attempts")
1023
+
1024
+ max_attempts = metrics_dataframe["attempts"].max() + 1
1025
+
1026
+ secondary_axis.set_yticks(range(1, max_attempts))
1027
+ secondary_axis.set_ylim(0, max_attempts)
1028
+
1029
+ primary_axis.set_xticklabels(
1030
+ [f"Task {task_index + 1}" for task_index in range(metrics_dataframe.shape[0])]
1031
+ )
1032
+
1033
+ plt.title("Execution Profile: Complexity vs Retries")
1034
+ # plt.tight_layout()