redroot 0.2.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.
- redroot/__init__.py +71 -0
- redroot/_core.py +556 -0
- redroot/cli.py +98 -0
- redroot/functions.py +240 -0
- redroot/instrument.py +614 -0
- redroot/ops.py +332 -0
- redroot/paths.py +162 -0
- redroot/propagation.py +447 -0
- redroot/py.typed +0 -0
- redroot/serialization.py +245 -0
- redroot/trace.py +474 -0
- redroot/types.py +385 -0
- redroot/validation.py +221 -0
- redroot/visualizer/__init__.py +11 -0
- redroot/visualizer/data.py +63 -0
- redroot/visualizer/graphviz.py +44 -0
- redroot/visualizer/networkx.py +36 -0
- redroot/visualizer/web/__init__.py +1 -0
- redroot/visualizer/web/server.py +52 -0
- redroot/visualizer/web/static/app.js +125 -0
- redroot/visualizer/web/static/index.html +38 -0
- redroot/visualizer/web/static/style.css +49 -0
- redroot-0.2.0.dist-info/METADATA +309 -0
- redroot-0.2.0.dist-info/RECORD +27 -0
- redroot-0.2.0.dist-info/WHEEL +4 -0
- redroot-0.2.0.dist-info/entry_points.txt +2 -0
- redroot-0.2.0.dist-info/licenses/LICENSE +21 -0
redroot/types.py
ADDED
|
@@ -0,0 +1,385 @@
|
|
|
1
|
+
"""Traced counterparts of ``int``, ``float``, ``Decimal`` and ``str``.
|
|
2
|
+
|
|
3
|
+
A traced value *is* an instance of its builtin type (``isinstance(x, int)``
|
|
4
|
+
holds, JSON encoders and C APIs accept it) and behaves exactly like it. While a
|
|
5
|
+
:class:`~redroot.Trace` is active, every operation on it records a node;
|
|
6
|
+
outside a trace, operations simply return plain builtin values.
|
|
7
|
+
|
|
8
|
+
What each operation records:
|
|
9
|
+
|
|
10
|
+
* arithmetic, ``round()``, ``math.floor()``, ``format()``, ``str()``, ``str``
|
|
11
|
+
and ``Decimal`` methods -> a re-evaluable node, returned as a traced value;
|
|
12
|
+
* comparisons, truth tests, ``int()``/``float()``, ``hash()``, ``len()`` ->
|
|
13
|
+
a *guard*: the plain result is returned, and the trace remembers what it
|
|
14
|
+
was so that re-evaluation can tell whether it would change.
|
|
15
|
+
|
|
16
|
+
Constructing ``TracedX(value, key=None, meta=None)`` converts ``value`` like
|
|
17
|
+
``X(value)``. Inside an active trace, converting a value traced in that trace
|
|
18
|
+
records a re-evaluable conversion; anything else becomes a new input, named
|
|
19
|
+
``key`` if given. Outside a trace the value is detached: it behaves like a
|
|
20
|
+
plain value and records nothing.
|
|
21
|
+
"""
|
|
22
|
+
|
|
23
|
+
from __future__ import annotations
|
|
24
|
+
|
|
25
|
+
import math
|
|
26
|
+
import operator
|
|
27
|
+
from collections.abc import Callable
|
|
28
|
+
from decimal import Decimal
|
|
29
|
+
from typing import TYPE_CHECKING, Any
|
|
30
|
+
|
|
31
|
+
from redroot._core import (
|
|
32
|
+
LEAF,
|
|
33
|
+
OP,
|
|
34
|
+
Node,
|
|
35
|
+
Observer,
|
|
36
|
+
Traced,
|
|
37
|
+
_active,
|
|
38
|
+
add_node,
|
|
39
|
+
binary,
|
|
40
|
+
call,
|
|
41
|
+
decimal_context_id,
|
|
42
|
+
observe_value,
|
|
43
|
+
record_raise,
|
|
44
|
+
register_wrapper,
|
|
45
|
+
unwrap,
|
|
46
|
+
)
|
|
47
|
+
from redroot.ops import (
|
|
48
|
+
BINARY_OPS,
|
|
49
|
+
COMPARISON_OPS,
|
|
50
|
+
DECIMAL_METHODS,
|
|
51
|
+
FLOAT_METHODS,
|
|
52
|
+
INT_METHODS,
|
|
53
|
+
STR_METHODS,
|
|
54
|
+
UNARY_OPS,
|
|
55
|
+
observe,
|
|
56
|
+
register_op,
|
|
57
|
+
)
|
|
58
|
+
|
|
59
|
+
if TYPE_CHECKING:
|
|
60
|
+
from pydantic import GetCoreSchemaHandler
|
|
61
|
+
from pydantic_core import CoreSchema
|
|
62
|
+
|
|
63
|
+
__all__ = ["TracedDecimal", "TracedFloat", "TracedInt", "TracedStr"]
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def _construct(
|
|
67
|
+
cls: type[Any],
|
|
68
|
+
base: type[Any],
|
|
69
|
+
cast_op: str,
|
|
70
|
+
value: Any,
|
|
71
|
+
key: str | None,
|
|
72
|
+
meta: dict[str, Any] | None,
|
|
73
|
+
) -> Any:
|
|
74
|
+
"""Shared constructor logic.
|
|
75
|
+
|
|
76
|
+
``TracedX(v)`` behaves like ``X(v)``. Inside an active trace, converting
|
|
77
|
+
a value traced in that trace records a re-evaluable cast; anything else
|
|
78
|
+
creates a new leaf (an input). Outside a trace the value is detached.
|
|
79
|
+
"""
|
|
80
|
+
trace = _active.get()
|
|
81
|
+
source = value._rt_node if isinstance(value, Traced) else None
|
|
82
|
+
is_cast = trace is not None and key is None and source is not None and source.trace is trace
|
|
83
|
+
try:
|
|
84
|
+
plain = base(unwrap(value))
|
|
85
|
+
except Exception as exc:
|
|
86
|
+
if is_cast:
|
|
87
|
+
record_raise(trace, cast_op, (source,), None, exc) # type: ignore[arg-type] # is_cast implies a Trace
|
|
88
|
+
raise
|
|
89
|
+
if trace is None or isinstance(trace, Observer):
|
|
90
|
+
if trace is not None:
|
|
91
|
+
trace.touch_node(source)
|
|
92
|
+
node = Node(LEAF, (), None, plain, LEAF, key=key, meta=meta)
|
|
93
|
+
elif is_cast:
|
|
94
|
+
ctx = decimal_context_id(trace) if base is Decimal else None
|
|
95
|
+
node = add_node(trace, cast_op, (source,), None, plain, OP, meta=meta, ctx=ctx)
|
|
96
|
+
else:
|
|
97
|
+
node = trace._add_leaf(plain, key, meta)
|
|
98
|
+
obj: Any = base.__new__(cls, plain)
|
|
99
|
+
obj._rt_node = node
|
|
100
|
+
return obj
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
def _factory(cls: type[Any], base: type[Any]) -> Callable[[Any, Node], Any]:
|
|
104
|
+
def wrap(value: Any, node: Node) -> Any:
|
|
105
|
+
obj: Any = base.__new__(cls, value)
|
|
106
|
+
obj._rt_node = node
|
|
107
|
+
return obj
|
|
108
|
+
|
|
109
|
+
return wrap
|
|
110
|
+
|
|
111
|
+
|
|
112
|
+
class _TracedScalar(Traced):
|
|
113
|
+
"""Behaviour shared by every traced type: repr, copying, pickling, Pydantic."""
|
|
114
|
+
|
|
115
|
+
__slots__ = ()
|
|
116
|
+
_base: type[Any]
|
|
117
|
+
|
|
118
|
+
def __repr__(self) -> str:
|
|
119
|
+
return repr(self._rt_node.value)
|
|
120
|
+
|
|
121
|
+
def __copy__(self) -> Any:
|
|
122
|
+
return self
|
|
123
|
+
|
|
124
|
+
def __deepcopy__(self, memo: dict[int, Any]) -> Any:
|
|
125
|
+
return self
|
|
126
|
+
|
|
127
|
+
def __reduce__(self) -> tuple[Any, ...]:
|
|
128
|
+
# Lineage does not survive serialization: unpickle as the plain value.
|
|
129
|
+
return (self._base, (self._rt_node.value,))
|
|
130
|
+
|
|
131
|
+
@classmethod
|
|
132
|
+
def __get_pydantic_core_schema__(
|
|
133
|
+
cls, source_type: Any, handler: GetCoreSchemaHandler
|
|
134
|
+
) -> CoreSchema:
|
|
135
|
+
from pydantic_core import core_schema
|
|
136
|
+
|
|
137
|
+
schemas: dict[type, Callable[[], CoreSchema]] = {
|
|
138
|
+
int: core_schema.int_schema,
|
|
139
|
+
float: core_schema.float_schema,
|
|
140
|
+
Decimal: core_schema.decimal_schema,
|
|
141
|
+
str: core_schema.str_schema,
|
|
142
|
+
}
|
|
143
|
+
from_plain = core_schema.no_info_after_validator_function(cls, schemas[cls._base]())
|
|
144
|
+
return core_schema.json_or_python_schema(
|
|
145
|
+
json_schema=from_plain,
|
|
146
|
+
# Already-traced values pass through untouched, keeping their lineage.
|
|
147
|
+
python_schema=core_schema.union_schema(
|
|
148
|
+
[core_schema.is_instance_schema(cls), from_plain]
|
|
149
|
+
),
|
|
150
|
+
serialization=core_schema.plain_serializer_function_ser_schema(unwrap),
|
|
151
|
+
)
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
class TracedInt(_TracedScalar, int):
|
|
155
|
+
"""An ``int`` that records how it was computed."""
|
|
156
|
+
|
|
157
|
+
# int/str are variable-size: instances keep the node in __dict__.
|
|
158
|
+
_base = int
|
|
159
|
+
|
|
160
|
+
def __new__(
|
|
161
|
+
cls, value: Any = 0, *, key: str | None = None, meta: dict[str, Any] | None = None
|
|
162
|
+
) -> TracedInt:
|
|
163
|
+
"""Convert ``value`` like ``int(value)`` (see the module docstring)."""
|
|
164
|
+
return _construct(cls, int, "int", value, key, meta) # type: ignore[no-any-return]
|
|
165
|
+
|
|
166
|
+
|
|
167
|
+
class TracedFloat(_TracedScalar, float):
|
|
168
|
+
"""A ``float`` that records how it was computed."""
|
|
169
|
+
|
|
170
|
+
__slots__ = ("_rt_node",)
|
|
171
|
+
_base = float
|
|
172
|
+
|
|
173
|
+
def __new__(
|
|
174
|
+
cls, value: Any = 0.0, *, key: str | None = None, meta: dict[str, Any] | None = None
|
|
175
|
+
) -> TracedFloat:
|
|
176
|
+
"""Convert ``value`` like ``float(value)`` (see the module docstring)."""
|
|
177
|
+
return _construct(cls, float, "float", value, key, meta) # type: ignore[no-any-return]
|
|
178
|
+
|
|
179
|
+
|
|
180
|
+
class TracedDecimal(_TracedScalar, Decimal):
|
|
181
|
+
"""A ``decimal.Decimal`` that records how it was computed."""
|
|
182
|
+
|
|
183
|
+
__slots__ = ("_rt_node",)
|
|
184
|
+
_base = Decimal
|
|
185
|
+
|
|
186
|
+
def __new__(
|
|
187
|
+
cls, value: Any = "0", *, key: str | None = None, meta: dict[str, Any] | None = None
|
|
188
|
+
) -> TracedDecimal:
|
|
189
|
+
"""Convert ``value`` like ``Decimal(value)`` (see the module docstring)."""
|
|
190
|
+
return _construct(cls, Decimal, "decimal", value, key, meta) # type: ignore[no-any-return]
|
|
191
|
+
|
|
192
|
+
|
|
193
|
+
class TracedStr(_TracedScalar, str):
|
|
194
|
+
"""A ``str`` that records how it was computed."""
|
|
195
|
+
|
|
196
|
+
# int/str are variable-size: instances keep the node in __dict__.
|
|
197
|
+
_base = str
|
|
198
|
+
|
|
199
|
+
def __new__(
|
|
200
|
+
cls, value: Any = "", *, key: str | None = None, meta: dict[str, Any] | None = None
|
|
201
|
+
) -> TracedStr:
|
|
202
|
+
"""Convert ``value`` like ``str(value)`` (see the module docstring)."""
|
|
203
|
+
return _construct(cls, str, "str", value, key, meta) # type: ignore[no-any-return]
|
|
204
|
+
|
|
205
|
+
|
|
206
|
+
# --------------------------------------------------------------------------
|
|
207
|
+
# Method installation. Methods are attached after class creation so that
|
|
208
|
+
# static type checkers keep seeing the builtin signatures: a TracedInt is
|
|
209
|
+
# typed as an int, which is exactly what it behaves like.
|
|
210
|
+
# --------------------------------------------------------------------------
|
|
211
|
+
|
|
212
|
+
|
|
213
|
+
def _binary_pair(op: str) -> tuple[Callable[..., Any], Callable[..., Any]]:
|
|
214
|
+
fn = BINARY_OPS[op][0]
|
|
215
|
+
|
|
216
|
+
def forward(self: Any, other: Any) -> Any:
|
|
217
|
+
return binary(op, fn, self, other)
|
|
218
|
+
|
|
219
|
+
def reflected(self: Any, other: Any) -> Any:
|
|
220
|
+
return binary(op, fn, other, self)
|
|
221
|
+
|
|
222
|
+
return forward, reflected
|
|
223
|
+
|
|
224
|
+
|
|
225
|
+
def _pow(self: Any, other: Any, modulo: Any = None) -> Any:
|
|
226
|
+
if modulo is None:
|
|
227
|
+
return binary("pow", pow, self, other)
|
|
228
|
+
return call("pow", pow, (self, other, modulo))
|
|
229
|
+
|
|
230
|
+
|
|
231
|
+
def _rpow(self: Any, other: Any, modulo: Any = None) -> Any:
|
|
232
|
+
if modulo is None:
|
|
233
|
+
return binary("pow", pow, other, self)
|
|
234
|
+
return call("pow", pow, (other, self, modulo))
|
|
235
|
+
|
|
236
|
+
|
|
237
|
+
def _comparison(op: str) -> Callable[..., Any]:
|
|
238
|
+
fn = COMPARISON_OPS[op][0]
|
|
239
|
+
|
|
240
|
+
def compare(self: Any, other: Any) -> Any:
|
|
241
|
+
return observe_value(op, fn, self, other)
|
|
242
|
+
|
|
243
|
+
return compare
|
|
244
|
+
|
|
245
|
+
|
|
246
|
+
def _guard(op: str, fn: Callable[[Any], Any]) -> Callable[..., Any]:
|
|
247
|
+
def method(self: Any) -> Any:
|
|
248
|
+
return observe_value(op, fn, self)
|
|
249
|
+
|
|
250
|
+
return method
|
|
251
|
+
|
|
252
|
+
|
|
253
|
+
def _hash(self: Any) -> int:
|
|
254
|
+
# Guard the whole value, not just its hash: -1 and -2 hash alike, yet a
|
|
255
|
+
# dict keyed by them differs.
|
|
256
|
+
return hash(observe_value("observe", observe, self))
|
|
257
|
+
|
|
258
|
+
|
|
259
|
+
def _function(op: str, fn: Callable[..., Any]) -> Callable[..., Any]:
|
|
260
|
+
def method(self: Any, *args: Any) -> Any:
|
|
261
|
+
return call(op, fn, (self, *args))
|
|
262
|
+
|
|
263
|
+
return method
|
|
264
|
+
|
|
265
|
+
|
|
266
|
+
def _round(self: Any, ndigits: Any = None) -> Any:
|
|
267
|
+
if ndigits is None:
|
|
268
|
+
return call("round", round, (self,))
|
|
269
|
+
return call("round", round, (self, ndigits))
|
|
270
|
+
|
|
271
|
+
|
|
272
|
+
def _method(prefix: str, name: str, owner: type[Any]) -> Callable[..., Any]:
|
|
273
|
+
op = f"{prefix}.{name}"
|
|
274
|
+
fn = getattr(owner, name)
|
|
275
|
+
|
|
276
|
+
def method(self: Any, *args: Any, **kwargs: Any) -> Any:
|
|
277
|
+
return call(op, fn, (self, *args), kwargs or None)
|
|
278
|
+
|
|
279
|
+
method.__name__ = name
|
|
280
|
+
method.__doc__ = fn.__doc__
|
|
281
|
+
return method
|
|
282
|
+
|
|
283
|
+
|
|
284
|
+
def _install(cls: type[Any], name: str, func: Callable[..., Any]) -> None:
|
|
285
|
+
if not name.startswith("__"):
|
|
286
|
+
func.__qualname__ = f"{cls.__name__}.{name}"
|
|
287
|
+
setattr(cls, name, func)
|
|
288
|
+
|
|
289
|
+
|
|
290
|
+
def _install_number(
|
|
291
|
+
cls: type[Any], binary_ops: tuple[str, ...], unary_ops: tuple[str, ...]
|
|
292
|
+
) -> None:
|
|
293
|
+
for op in binary_ops:
|
|
294
|
+
forward, reflected = _binary_pair(op)
|
|
295
|
+
_install(cls, f"__{op}__", forward)
|
|
296
|
+
_install(cls, f"__r{op}__", reflected)
|
|
297
|
+
_install(cls, "__pow__", _pow)
|
|
298
|
+
_install(cls, "__rpow__", _rpow)
|
|
299
|
+
for op in unary_ops:
|
|
300
|
+
_install(cls, f"__{op}__", _function(op, UNARY_OPS[op][0]))
|
|
301
|
+
for op in COMPARISON_OPS:
|
|
302
|
+
_install(cls, f"__{op}__", _comparison(op))
|
|
303
|
+
_install(cls, "__abs__", _function("abs", abs))
|
|
304
|
+
_install(cls, "__round__", _round)
|
|
305
|
+
_install(cls, "__floor__", _function("floor", math.floor))
|
|
306
|
+
_install(cls, "__ceil__", _function("ceil", math.ceil))
|
|
307
|
+
_install(cls, "__trunc__", _function("trunc", math.trunc))
|
|
308
|
+
_install(cls, "__divmod__", _function("divmod", divmod))
|
|
309
|
+
_install(cls, "__rdivmod__", lambda self, other: call("divmod", divmod, (other, self)))
|
|
310
|
+
_install(cls, "__format__", _function("format", format))
|
|
311
|
+
_install(cls, "__str__", _function("str", str))
|
|
312
|
+
_install(cls, "__bool__", _guard("bool", bool))
|
|
313
|
+
_install(cls, "__int__", _guard("int", int))
|
|
314
|
+
_install(cls, "__float__", _guard("float", float))
|
|
315
|
+
_install(cls, "__complex__", _guard("complex", complex))
|
|
316
|
+
_install(cls, "__hash__", _hash)
|
|
317
|
+
|
|
318
|
+
|
|
319
|
+
_ARITHMETIC = ("add", "sub", "mul", "truediv", "floordiv", "mod")
|
|
320
|
+
|
|
321
|
+
_install_number(
|
|
322
|
+
TracedInt, (*_ARITHMETIC, "lshift", "rshift", "and", "or", "xor"), ("neg", "pos", "invert")
|
|
323
|
+
)
|
|
324
|
+
_install(TracedInt, "__index__", _guard("index", operator.index))
|
|
325
|
+
_install_number(TracedFloat, _ARITHMETIC, ("neg", "pos"))
|
|
326
|
+
_install_number(TracedDecimal, _ARITHMETIC, ("neg", "pos"))
|
|
327
|
+
|
|
328
|
+
for _prefix, _cls, _owner, _names in (
|
|
329
|
+
("int", TracedInt, int, INT_METHODS),
|
|
330
|
+
("float", TracedFloat, float, FLOAT_METHODS),
|
|
331
|
+
("decimal", TracedDecimal, Decimal, DECIMAL_METHODS),
|
|
332
|
+
("str", TracedStr, str, STR_METHODS),
|
|
333
|
+
):
|
|
334
|
+
for _name in _names:
|
|
335
|
+
_install(_cls, _name, _method(_prefix, _name, _owner))
|
|
336
|
+
|
|
337
|
+
|
|
338
|
+
def _getslice(value: Any, start: Any, stop: Any, step: Any) -> Any:
|
|
339
|
+
return value[start:stop:step]
|
|
340
|
+
|
|
341
|
+
|
|
342
|
+
register_op("getslice", _getslice)
|
|
343
|
+
|
|
344
|
+
|
|
345
|
+
def _str_getitem(self: Any, key: Any) -> Any:
|
|
346
|
+
if isinstance(key, slice):
|
|
347
|
+
return call("getslice", _getslice, (self, key.start, key.stop, key.step))
|
|
348
|
+
return call("getitem", operator.getitem, (self, key))
|
|
349
|
+
|
|
350
|
+
|
|
351
|
+
def _str_iter(self: Any) -> Any:
|
|
352
|
+
# Iterating exposes every character: treat the whole string as observed.
|
|
353
|
+
return iter(observe_value("observe", lambda v: v, self))
|
|
354
|
+
|
|
355
|
+
|
|
356
|
+
def _str_join(self: Any, iterable: Any) -> Any:
|
|
357
|
+
return call("str.join", str.join, (self, list(iterable)))
|
|
358
|
+
|
|
359
|
+
|
|
360
|
+
for _op in ("add", "mul"):
|
|
361
|
+
_forward, _reflected = _binary_pair(_op)
|
|
362
|
+
_install(TracedStr, f"__{_op}__", _forward)
|
|
363
|
+
_install(TracedStr, f"__r{_op}__", _reflected)
|
|
364
|
+
_install(TracedStr, "__mod__", lambda self, other: call("mod", operator.mod, (self, other)))
|
|
365
|
+
_install(TracedStr, "__rmod__", lambda self, other: call("mod", operator.mod, (other, self)))
|
|
366
|
+
for _op in COMPARISON_OPS:
|
|
367
|
+
_install(TracedStr, f"__{_op}__", _comparison(_op))
|
|
368
|
+
_install(TracedStr, "__getitem__", _str_getitem)
|
|
369
|
+
_install(
|
|
370
|
+
TracedStr,
|
|
371
|
+
"__contains__",
|
|
372
|
+
lambda self, item: observe_value("contains", operator.contains, self, item),
|
|
373
|
+
)
|
|
374
|
+
_install(TracedStr, "__len__", _guard("len", len))
|
|
375
|
+
_install(TracedStr, "__bool__", _guard("bool", bool))
|
|
376
|
+
_install(TracedStr, "__hash__", _hash)
|
|
377
|
+
_install(TracedStr, "__iter__", _str_iter)
|
|
378
|
+
_install(TracedStr, "__format__", _function("format", format))
|
|
379
|
+
_install(TracedStr, "__str__", lambda self: self)
|
|
380
|
+
_install(TracedStr, "join", _str_join)
|
|
381
|
+
|
|
382
|
+
register_wrapper(int, _factory(TracedInt, int))
|
|
383
|
+
register_wrapper(float, _factory(TracedFloat, float))
|
|
384
|
+
register_wrapper(Decimal, _factory(TracedDecimal, Decimal))
|
|
385
|
+
register_wrapper(str, _factory(TracedStr, str))
|
redroot/validation.py
ADDED
|
@@ -0,0 +1,221 @@
|
|
|
1
|
+
"""Checks that tell you whether tracing is trustworthy for a workflow.
|
|
2
|
+
|
|
3
|
+
Run these on real recorded workflows before relying on propagation:
|
|
4
|
+
|
|
5
|
+
* :func:`coverage` - how many outputs are linked to inputs, how many can be
|
|
6
|
+
recomputed, how many inputs feed guards.
|
|
7
|
+
* :func:`check_identity` - re-evaluating every node from its recorded
|
|
8
|
+
operands reproduces the recorded value exactly.
|
|
9
|
+
* :func:`check_perturbation` - after an edit, propagated outputs agree with a
|
|
10
|
+
real re-execution wherever propagation claimed to know the answer.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from __future__ import annotations
|
|
14
|
+
|
|
15
|
+
import decimal
|
|
16
|
+
from collections import deque
|
|
17
|
+
from collections.abc import Callable, Mapping
|
|
18
|
+
from dataclasses import dataclass
|
|
19
|
+
from typing import Any
|
|
20
|
+
|
|
21
|
+
from redroot._core import GUARD, LEAF, Node, same_value
|
|
22
|
+
from redroot.paths import apply_edits
|
|
23
|
+
from redroot.propagation import Propagation, Verification, evaluate_outcome, node_function
|
|
24
|
+
from redroot.trace import Trace, run
|
|
25
|
+
|
|
26
|
+
__all__ = [
|
|
27
|
+
"Coverage",
|
|
28
|
+
"IdentityMismatch",
|
|
29
|
+
"IdentityReport",
|
|
30
|
+
"PerturbationReport",
|
|
31
|
+
"check_identity",
|
|
32
|
+
"check_perturbation",
|
|
33
|
+
"coverage",
|
|
34
|
+
]
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
@dataclass(frozen=True, slots=True)
|
|
38
|
+
class Coverage:
|
|
39
|
+
"""How much of a trace's output the graph can explain and recompute."""
|
|
40
|
+
|
|
41
|
+
outputs: int
|
|
42
|
+
linked: int
|
|
43
|
+
"""Outputs computed from at least one keyed input."""
|
|
44
|
+
replayable: int
|
|
45
|
+
"""Linked outputs whose every ancestor operation can be re-evaluated."""
|
|
46
|
+
inputs: int
|
|
47
|
+
untracked_inputs: int
|
|
48
|
+
guards: int
|
|
49
|
+
guarded_inputs: tuple[str, ...]
|
|
50
|
+
"""Input keys that feed at least one guard: edits to them may force re-execution."""
|
|
51
|
+
|
|
52
|
+
@property
|
|
53
|
+
def unlinked(self) -> int:
|
|
54
|
+
"""Outputs not backed by a traced value."""
|
|
55
|
+
return self.outputs - self.linked
|
|
56
|
+
|
|
57
|
+
@property
|
|
58
|
+
def linked_ratio(self) -> float:
|
|
59
|
+
"""Share of outputs that are linked (1.0 if there are none)."""
|
|
60
|
+
return self.linked / self.outputs if self.outputs else 1.0
|
|
61
|
+
|
|
62
|
+
@property
|
|
63
|
+
def replayable_ratio(self) -> float:
|
|
64
|
+
"""Share of outputs that can be recomputed (1.0 if there are none)."""
|
|
65
|
+
return self.replayable / self.outputs if self.outputs else 1.0
|
|
66
|
+
|
|
67
|
+
def __str__(self) -> str:
|
|
68
|
+
return "\n".join(
|
|
69
|
+
[
|
|
70
|
+
f"outputs ........... {self.outputs}",
|
|
71
|
+
f" linked .......... {self.linked} ({self.linked_ratio:.0%})",
|
|
72
|
+
f" replayable ...... {self.replayable} ({self.replayable_ratio:.0%})",
|
|
73
|
+
f"inputs ............ {self.inputs} (+{self.untracked_inputs} untracked)",
|
|
74
|
+
f" feeding guards .. {len(self.guarded_inputs)}",
|
|
75
|
+
f"guards ............ {self.guards}",
|
|
76
|
+
]
|
|
77
|
+
)
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
def coverage(trace: Trace) -> Coverage:
|
|
81
|
+
"""Measure how well ``trace`` covers its outputs."""
|
|
82
|
+
replayable: dict[int, bool] = {}
|
|
83
|
+
keyed: dict[int, bool] = {} # computed from at least one keyed input
|
|
84
|
+
for node in trace.nodes:
|
|
85
|
+
if node.kind == LEAF:
|
|
86
|
+
replayable[node.id] = True
|
|
87
|
+
keyed[node.id] = node.key is not None
|
|
88
|
+
continue
|
|
89
|
+
operands = list(node.operands())
|
|
90
|
+
replayable[node.id] = node_function(node) is not None and all(
|
|
91
|
+
replayable[operand.id] for operand in operands
|
|
92
|
+
)
|
|
93
|
+
keyed[node.id] = any(keyed[operand.id] for operand in operands)
|
|
94
|
+
|
|
95
|
+
# Values built only from anonymous leaves (e.g. rebuilt inside a library
|
|
96
|
+
# with type(x)(...)) are not linked to the inputs, whatever their type.
|
|
97
|
+
linked = [
|
|
98
|
+
out.node for out in trace.outputs.values() if out.node is not None and keyed[out.node.id]
|
|
99
|
+
]
|
|
100
|
+
return Coverage(
|
|
101
|
+
outputs=len(trace.outputs),
|
|
102
|
+
linked=len(linked),
|
|
103
|
+
replayable=sum(1 for node in linked if replayable[node.id]),
|
|
104
|
+
inputs=len(trace.inputs),
|
|
105
|
+
untracked_inputs=len(trace.untracked_inputs),
|
|
106
|
+
guards=sum(1 for node in trace.nodes if node.kind == GUARD),
|
|
107
|
+
guarded_inputs=_guarded_inputs(trace),
|
|
108
|
+
)
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
def _guarded_inputs(trace: Trace) -> tuple[str, ...]:
|
|
112
|
+
seen: set[int] = set()
|
|
113
|
+
queue = deque(node for node in trace.nodes if node.kind == GUARD)
|
|
114
|
+
keys: list[tuple[int, str]] = []
|
|
115
|
+
while queue:
|
|
116
|
+
node = queue.popleft()
|
|
117
|
+
for operand in node.operands():
|
|
118
|
+
if operand.id in seen:
|
|
119
|
+
continue
|
|
120
|
+
seen.add(operand.id)
|
|
121
|
+
if operand.kind == LEAF:
|
|
122
|
+
if operand.key is not None:
|
|
123
|
+
keys.append((operand.id, operand.key))
|
|
124
|
+
else:
|
|
125
|
+
queue.append(operand)
|
|
126
|
+
return tuple(key for _, key in sorted(keys))
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
@dataclass(frozen=True, slots=True)
|
|
130
|
+
class IdentityMismatch:
|
|
131
|
+
"""A node whose re-evaluation disagrees with its recorded value."""
|
|
132
|
+
|
|
133
|
+
node: Node
|
|
134
|
+
recorded: Any
|
|
135
|
+
recomputed: Any
|
|
136
|
+
error: str | None = None
|
|
137
|
+
|
|
138
|
+
|
|
139
|
+
@dataclass(frozen=True, slots=True)
|
|
140
|
+
class IdentityReport:
|
|
141
|
+
"""Result of :func:`check_identity`."""
|
|
142
|
+
|
|
143
|
+
checked: int
|
|
144
|
+
skipped: int
|
|
145
|
+
"""Nodes that are not replayable (e.g. LLM calls)."""
|
|
146
|
+
mismatches: tuple[IdentityMismatch, ...]
|
|
147
|
+
|
|
148
|
+
@property
|
|
149
|
+
def ok(self) -> bool:
|
|
150
|
+
"""Whether every node re-evaluated to its recorded value."""
|
|
151
|
+
return not self.mismatches
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def check_identity(trace: Trace) -> IdentityReport:
|
|
155
|
+
"""Re-evaluate every node from its recorded operands and compare.
|
|
156
|
+
|
|
157
|
+
Any mismatch means an operation did not behave like its registered
|
|
158
|
+
implementation (or is not deterministic), so propagation through it
|
|
159
|
+
cannot be trusted.
|
|
160
|
+
"""
|
|
161
|
+
checked = skipped = 0
|
|
162
|
+
mismatches: list[IdentityMismatch] = []
|
|
163
|
+
with decimal.localcontext(trace.decimal_context):
|
|
164
|
+
for node in trace.nodes:
|
|
165
|
+
if node.kind == LEAF:
|
|
166
|
+
continue
|
|
167
|
+
if node_function(node) is None:
|
|
168
|
+
skipped += 1
|
|
169
|
+
continue
|
|
170
|
+
checked += 1
|
|
171
|
+
try:
|
|
172
|
+
recomputed = evaluate_outcome(node)
|
|
173
|
+
except Exception as exc:
|
|
174
|
+
mismatches.append(
|
|
175
|
+
IdentityMismatch(node, node.value, None, f"{type(exc).__name__}: {exc}")
|
|
176
|
+
)
|
|
177
|
+
continue
|
|
178
|
+
if not same_value(recomputed, node.value):
|
|
179
|
+
mismatches.append(IdentityMismatch(node, node.value, recomputed))
|
|
180
|
+
return IdentityReport(checked, skipped, tuple(mismatches))
|
|
181
|
+
|
|
182
|
+
|
|
183
|
+
@dataclass(frozen=True, slots=True)
|
|
184
|
+
class PerturbationReport:
|
|
185
|
+
"""Result of :func:`check_perturbation`."""
|
|
186
|
+
|
|
187
|
+
trace: Trace
|
|
188
|
+
"""The trace of the original run."""
|
|
189
|
+
propagation: Propagation
|
|
190
|
+
"""What the graph predicted for the edit."""
|
|
191
|
+
retrace: Trace
|
|
192
|
+
"""The trace of the re-execution on edited inputs."""
|
|
193
|
+
verification: Verification
|
|
194
|
+
|
|
195
|
+
@property
|
|
196
|
+
def consistent(self) -> bool:
|
|
197
|
+
"""Whether every exact prediction matched the re-execution."""
|
|
198
|
+
return self.verification.ok
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
def check_perturbation(
|
|
202
|
+
workflow: Callable[[Any], Any],
|
|
203
|
+
inputs: Any,
|
|
204
|
+
edits: Mapping[str, Any],
|
|
205
|
+
*,
|
|
206
|
+
input_root: str = "in",
|
|
207
|
+
output_root: str = "out",
|
|
208
|
+
) -> PerturbationReport:
|
|
209
|
+
"""Edit inputs, then compare graph propagation with a real re-execution.
|
|
210
|
+
|
|
211
|
+
The workflow runs twice under tracing: on ``inputs``, then on ``inputs``
|
|
212
|
+
with ``edits`` applied. Outputs propagation reports as ``UPDATED``,
|
|
213
|
+
``UNCHANGED`` or ``UNLINKED`` must match the re-execution; ``STALE``
|
|
214
|
+
outputs are expected to need it.
|
|
215
|
+
"""
|
|
216
|
+
trace, _ = run(workflow, inputs, input_root=input_root, output_root=output_root)
|
|
217
|
+
propagation = trace.propagate(edits)
|
|
218
|
+
edited = apply_edits(inputs, edits, root=input_root)
|
|
219
|
+
retrace, _ = run(workflow, edited, input_root=input_root, output_root=output_root)
|
|
220
|
+
verification = propagation.verify(retrace.output_values())
|
|
221
|
+
return PerturbationReport(trace, propagation, retrace, verification)
|
|
@@ -0,0 +1,11 @@
|
|
|
1
|
+
"""Exports and viewers for traces.
|
|
2
|
+
|
|
3
|
+
``to_networkx`` and ``to_graphviz`` need ``redroot[viz]``; the web viewer
|
|
4
|
+
(``redroot visualize trace.json``) needs ``redroot[web]``.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from redroot.visualizer.data import graph_data
|
|
8
|
+
from redroot.visualizer.graphviz import to_graphviz
|
|
9
|
+
from redroot.visualizer.networkx import to_networkx
|
|
10
|
+
|
|
11
|
+
__all__ = ["graph_data", "to_graphviz", "to_networkx"]
|
|
@@ -0,0 +1,63 @@
|
|
|
1
|
+
"""A display-oriented view of a trace, shared by the exporters and the web viewer."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from typing import Any
|
|
6
|
+
|
|
7
|
+
from redroot._core import Node
|
|
8
|
+
from redroot.trace import Trace
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def _short(value: Any, limit: int) -> str:
|
|
12
|
+
text = repr(value)
|
|
13
|
+
return text if len(text) <= limit else text[: limit - 1] + "…"
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def node_label(node: Node, limit: int = 24) -> str:
|
|
17
|
+
"""A one-line label: the key for inputs, ``op = value`` otherwise."""
|
|
18
|
+
if node.key is not None:
|
|
19
|
+
return node.key
|
|
20
|
+
return f"{node.op} = {_short(node.value, limit)}"
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def graph_data(trace: Trace, *, label_limit: int = 24) -> dict[str, Any]:
|
|
24
|
+
"""Return nodes, edges and outputs of ``trace`` as JSON-compatible data.
|
|
25
|
+
|
|
26
|
+
Edges point from an operand to the node computed from it.
|
|
27
|
+
"""
|
|
28
|
+
nodes = [
|
|
29
|
+
{
|
|
30
|
+
"id": node.id,
|
|
31
|
+
"op": node.op,
|
|
32
|
+
"kind": node.kind,
|
|
33
|
+
"key": node.key,
|
|
34
|
+
"label": node_label(node, label_limit),
|
|
35
|
+
"value": _short(node.value, 500),
|
|
36
|
+
"meta": {k: _short(v, 200) for k, v in (node.meta or {}).items()},
|
|
37
|
+
}
|
|
38
|
+
for node in trace.nodes
|
|
39
|
+
]
|
|
40
|
+
edges = [
|
|
41
|
+
{"source": operand_id, "target": node.id}
|
|
42
|
+
for node in trace.nodes
|
|
43
|
+
for operand_id in dict.fromkeys(operand.id for operand in node.operands())
|
|
44
|
+
]
|
|
45
|
+
outputs = []
|
|
46
|
+
for key, out in trace.outputs.items():
|
|
47
|
+
linked = out.node is not None
|
|
48
|
+
outputs.append(
|
|
49
|
+
{
|
|
50
|
+
"key": key,
|
|
51
|
+
"node": out.node.id if out.node is not None else None,
|
|
52
|
+
"value": _short(out.value, 500),
|
|
53
|
+
"expression": trace.explain(key) if linked else None,
|
|
54
|
+
"sources": [s.key or f"#{s.id}" for s in trace.sources(key)] if linked else [],
|
|
55
|
+
}
|
|
56
|
+
)
|
|
57
|
+
return {
|
|
58
|
+
"id": trace.id,
|
|
59
|
+
"name": trace.name,
|
|
60
|
+
"nodes": nodes,
|
|
61
|
+
"edges": edges,
|
|
62
|
+
"outputs": outputs,
|
|
63
|
+
}
|