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.
- deer/__init__.py +36 -0
- deer/builtins/__init__.py +7 -0
- deer/builtins/python_manager/agent.py +29 -0
- deer/builtins/python_manager/tools.py +54 -0
- deer/core/__init__.py +1 -0
- deer/core/agent.py +463 -0
- deer/core/ui.py +31 -0
- deer/drivers/__init__.py +66 -0
- deer/drivers/base_driver.py +56 -0
- deer/drivers/gemini_driver.py +62 -0
- deer/drivers/ollama_driver.py +69 -0
- deer/executor/__init__.py +1 -0
- deer/executor/executor.py +168 -0
- deer/executor/logic.py +75 -0
- deer/executor/logic_secure.py +102 -0
- deer/main.py +71 -0
- deer/planner/__init__.py +1 -0
- deer/planner/planner.py +71 -0
- deer/prompts/__init__.py +6 -0
- deer/prompts/error_explain.py +30 -0
- deer/prompts/goal_improvement.py +65 -0
- deer/prompts/goal_validation.py +31 -0
- deer/prompts/humanizer.py +20 -0
- deer/prompts/planner.py +92 -0
- deer/prompts/response_improvement.py +23 -0
- deer/schema/__init__.py +2 -0
- deer/schema/io.py +40 -0
- deer/schema/plan.py +75 -0
- deer/tools/__init__.py +4 -0
- deer/tools/base.py +141 -0
- deer/tools/builtin/__init__.py +3 -0
- deer/tools/builtin/file_manager.py +136 -0
- deer/tools/builtin/git_manager.py +67 -0
- deer/tools/builtin/search_manager.py +123 -0
- deer/tools/decorators.py +127 -0
- deer/tools/registry.py +114 -0
- deer/tracing/__init__.py +2 -0
- deer/tracing/logging_config.py +20 -0
- deer/tracing/store.py +21 -0
- deer/utils/__init__.py +0 -0
- deer/utils/console.py +11 -0
- deer/utils/plots/__init__.py +7 -0
- deer/utils/plots/plot_traces.py +1034 -0
- deer/validator/__init__.py +1 -0
- deer/validator/plan_validator.py +23 -0
- deer/validator/rules.py +98 -0
- deer_agent_framework-0.0.dist-info/METADATA +163 -0
- deer_agent_framework-0.0.dist-info/RECORD +52 -0
- deer_agent_framework-0.0.dist-info/WHEEL +5 -0
- deer_agent_framework-0.0.dist-info/entry_points.txt +2 -0
- deer_agent_framework-0.0.dist-info/licenses/LICENSE +24 -0
- 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()
|