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/__init__.py
ADDED
|
@@ -0,0 +1,71 @@
|
|
|
1
|
+
"""RedRoot: value-level lineage for Python, in both directions.
|
|
2
|
+
|
|
3
|
+
Record a computation once, then ask where any output came from, and what
|
|
4
|
+
every output becomes when an input changes, without running it again::
|
|
5
|
+
|
|
6
|
+
import redroot as rr
|
|
7
|
+
|
|
8
|
+
with rr.Trace() as trace:
|
|
9
|
+
data = trace.track({"a": 1200, "b": 300}, root="ext")
|
|
10
|
+
trace.collect({"total": data["a"] + data["b"]}, root="out")
|
|
11
|
+
|
|
12
|
+
trace.sources("out:total") # [ext:a, ext:b]
|
|
13
|
+
trace.propagate({"ext:a": 900}).outputs # out:total -> 1200, exactly
|
|
14
|
+
|
|
15
|
+
See the README for the full guide.
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
from redroot._core import Node, Traced, active_trace, node_of, unwrap
|
|
19
|
+
from redroot.functions import annotate, derive, traced, traced_llm
|
|
20
|
+
from redroot.paths import apply_edits, format_key, parse_key
|
|
21
|
+
from redroot.propagation import (
|
|
22
|
+
Mismatch,
|
|
23
|
+
OutputChange,
|
|
24
|
+
Propagation,
|
|
25
|
+
Reason,
|
|
26
|
+
ReasonKind,
|
|
27
|
+
Status,
|
|
28
|
+
Verification,
|
|
29
|
+
)
|
|
30
|
+
from redroot.trace import Output, Trace, run, track
|
|
31
|
+
from redroot.types import TracedDecimal, TracedFloat, TracedInt, TracedStr
|
|
32
|
+
|
|
33
|
+
__version__ = "0.2.0"
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def is_traced(value: object) -> bool:
|
|
37
|
+
"""Whether ``value`` is a traced value."""
|
|
38
|
+
return isinstance(value, Traced)
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
__all__ = [
|
|
42
|
+
"Mismatch",
|
|
43
|
+
"Node",
|
|
44
|
+
"Output",
|
|
45
|
+
"OutputChange",
|
|
46
|
+
"Propagation",
|
|
47
|
+
"Reason",
|
|
48
|
+
"ReasonKind",
|
|
49
|
+
"Status",
|
|
50
|
+
"Trace",
|
|
51
|
+
"Traced",
|
|
52
|
+
"TracedDecimal",
|
|
53
|
+
"TracedFloat",
|
|
54
|
+
"TracedInt",
|
|
55
|
+
"TracedStr",
|
|
56
|
+
"Verification",
|
|
57
|
+
"__version__",
|
|
58
|
+
"active_trace",
|
|
59
|
+
"annotate",
|
|
60
|
+
"apply_edits",
|
|
61
|
+
"derive",
|
|
62
|
+
"format_key",
|
|
63
|
+
"is_traced",
|
|
64
|
+
"node_of",
|
|
65
|
+
"parse_key",
|
|
66
|
+
"run",
|
|
67
|
+
"traced",
|
|
68
|
+
"traced_llm",
|
|
69
|
+
"track",
|
|
70
|
+
"unwrap",
|
|
71
|
+
]
|
redroot/_core.py
ADDED
|
@@ -0,0 +1,556 @@
|
|
|
1
|
+
"""Recording machinery shared by the traced types and :class:`~redroot.Trace`.
|
|
2
|
+
|
|
3
|
+
This module is internal. It holds the hot paths that run on every traced
|
|
4
|
+
operation, so it favours plain functions and ``__slots__`` over abstraction.
|
|
5
|
+
|
|
6
|
+
Invariants (see ``AGENTS.md``):
|
|
7
|
+
|
|
8
|
+
* ``Node.value`` is always a *plain* value, never a traced one.
|
|
9
|
+
* ``Node.args`` hold, in call order, either a :class:`Node` (a traced operand
|
|
10
|
+
recorded in the same trace) or a plain constant. Nothing is dropped.
|
|
11
|
+
* Library code never calls ``int()``/``float()``/``str()``/``hash()``/``==``
|
|
12
|
+
on a traced value: those are recorded as guards. Use :func:`unwrap`.
|
|
13
|
+
"""
|
|
14
|
+
|
|
15
|
+
from __future__ import annotations
|
|
16
|
+
|
|
17
|
+
import decimal
|
|
18
|
+
import math
|
|
19
|
+
from collections.abc import Callable, Iterator
|
|
20
|
+
from contextvars import ContextVar
|
|
21
|
+
from decimal import Decimal
|
|
22
|
+
from typing import TYPE_CHECKING, Any, Final
|
|
23
|
+
|
|
24
|
+
from redroot.ops import shape
|
|
25
|
+
|
|
26
|
+
if TYPE_CHECKING:
|
|
27
|
+
from redroot.trace import Trace
|
|
28
|
+
|
|
29
|
+
LEAF: Final = "leaf"
|
|
30
|
+
"""Kind of an input value."""
|
|
31
|
+
OP: Final = "op"
|
|
32
|
+
"""Kind of a value computed from other values."""
|
|
33
|
+
GUARD: Final = "guard"
|
|
34
|
+
"""Kind of a value that left tracking (a comparison outcome, an ``int()``...).
|
|
35
|
+
|
|
36
|
+
Execution may have depended on it in ways the graph cannot see, so if it
|
|
37
|
+
changes on re-evaluation the graph can no longer be trusted.
|
|
38
|
+
"""
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
class Node:
|
|
42
|
+
"""One recorded value in a :class:`~redroot.Trace`.
|
|
43
|
+
|
|
44
|
+
Attributes:
|
|
45
|
+
id: Position in the trace. Operands always have smaller ids, so id
|
|
46
|
+
order is a topological order.
|
|
47
|
+
op: Name of the operation (see :mod:`redroot.ops`); ``"leaf"`` for
|
|
48
|
+
inputs.
|
|
49
|
+
args: Operands in call order: :class:`Node` references or constants.
|
|
50
|
+
May nest inside lists, tuples and dicts.
|
|
51
|
+
kwargs: Keyword operands, or ``None``.
|
|
52
|
+
value: The plain value computed when the trace was recorded.
|
|
53
|
+
kind: ``"leaf"``, ``"op"`` or ``"guard"``.
|
|
54
|
+
key: Semantic id of a leaf, e.g. ``"ext:bank.line_items[3].amount"``.
|
|
55
|
+
meta: Free-form metadata (model name, source document...), or ``None``.
|
|
56
|
+
trace: The trace that owns the node; ``None`` if detached.
|
|
57
|
+
ctx: Index of the ``decimal`` context the node was computed under, in
|
|
58
|
+
the trace's contexts, or ``None`` if no ``Decimal`` was involved.
|
|
59
|
+
fn: The function that computed the node, when it is not the
|
|
60
|
+
registered one (``@traced`` functions), or :data:`NOT_REPLAYABLE`.
|
|
61
|
+
"""
|
|
62
|
+
|
|
63
|
+
__slots__ = ("args", "ctx", "fn", "id", "key", "kind", "kwargs", "meta", "op", "trace", "value")
|
|
64
|
+
|
|
65
|
+
def __init__(
|
|
66
|
+
self,
|
|
67
|
+
op: str,
|
|
68
|
+
args: tuple[Any, ...],
|
|
69
|
+
kwargs: dict[str, Any] | None,
|
|
70
|
+
value: Any,
|
|
71
|
+
kind: str,
|
|
72
|
+
*,
|
|
73
|
+
key: str | None = None,
|
|
74
|
+
meta: dict[str, Any] | None = None,
|
|
75
|
+
) -> None:
|
|
76
|
+
self.id = -1
|
|
77
|
+
self.op = op
|
|
78
|
+
self.args = args
|
|
79
|
+
self.kwargs = kwargs
|
|
80
|
+
self.value = value
|
|
81
|
+
self.kind = kind
|
|
82
|
+
self.key = key
|
|
83
|
+
self.meta = meta
|
|
84
|
+
self.trace: Trace | None = None
|
|
85
|
+
self.ctx: int | None = None
|
|
86
|
+
self.fn: Callable[..., Any] | None = None
|
|
87
|
+
|
|
88
|
+
def __repr__(self) -> str:
|
|
89
|
+
label = self.key if self.key is not None else self.op
|
|
90
|
+
return f"<Node #{self.id} {self.kind} {label}={self.value!r}>"
|
|
91
|
+
|
|
92
|
+
@property
|
|
93
|
+
def is_leaf(self) -> bool:
|
|
94
|
+
return self.kind == LEAF
|
|
95
|
+
|
|
96
|
+
@property
|
|
97
|
+
def is_guard(self) -> bool:
|
|
98
|
+
return self.kind == GUARD
|
|
99
|
+
|
|
100
|
+
def operands(self) -> Iterator[Node]:
|
|
101
|
+
"""Yield the nodes this node was computed from, in argument order."""
|
|
102
|
+
yield from iter_nodes(self.args)
|
|
103
|
+
if self.kwargs:
|
|
104
|
+
yield from iter_nodes(self.kwargs)
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def iter_nodes(obj: Any) -> Iterator[Node]:
|
|
108
|
+
"""Yield every :class:`Node` referenced in a (possibly nested) argument."""
|
|
109
|
+
kind = type(obj)
|
|
110
|
+
if kind is Node:
|
|
111
|
+
yield obj
|
|
112
|
+
elif kind is list or isinstance(obj, tuple):
|
|
113
|
+
for item in obj:
|
|
114
|
+
yield from iter_nodes(item)
|
|
115
|
+
elif kind is dict:
|
|
116
|
+
for item in obj.values():
|
|
117
|
+
yield from iter_nodes(item)
|
|
118
|
+
|
|
119
|
+
|
|
120
|
+
def rebuild_tuple(template: tuple[Any, ...], items: Iterator[Any] | list[Any]) -> tuple[Any, ...]:
|
|
121
|
+
"""Build a tuple of ``template``'s type (plain or named) from ``items``."""
|
|
122
|
+
if type(template) is tuple:
|
|
123
|
+
return tuple(items)
|
|
124
|
+
return type(template)._make(items) # type: ignore[attr-defined, no-any-return]
|
|
125
|
+
|
|
126
|
+
|
|
127
|
+
class Raised:
|
|
128
|
+
"""The recorded outcome of an operation that raised an exception.
|
|
129
|
+
|
|
130
|
+
Operations on traced values that raise are recorded as guards with this
|
|
131
|
+
value, so code that catches the exception (``try``/``except``) is guarded:
|
|
132
|
+
if an edit makes the operation stop raising, or raise something else,
|
|
133
|
+
the guard flips.
|
|
134
|
+
"""
|
|
135
|
+
|
|
136
|
+
__slots__ = ("type_name",)
|
|
137
|
+
|
|
138
|
+
def __init__(self, type_name: str) -> None:
|
|
139
|
+
self.type_name = type_name
|
|
140
|
+
|
|
141
|
+
@classmethod
|
|
142
|
+
def of(cls, exc: BaseException) -> Raised:
|
|
143
|
+
kind = type(exc)
|
|
144
|
+
module = "" if kind.__module__ == "builtins" else f"{kind.__module__}."
|
|
145
|
+
return cls(module + kind.__qualname__)
|
|
146
|
+
|
|
147
|
+
def __eq__(self, other: object) -> bool:
|
|
148
|
+
return isinstance(other, Raised) and other.type_name == self.type_name
|
|
149
|
+
|
|
150
|
+
def __hash__(self) -> int:
|
|
151
|
+
return hash(self.type_name)
|
|
152
|
+
|
|
153
|
+
def __repr__(self) -> str:
|
|
154
|
+
return f"<raised {self.type_name}>"
|
|
155
|
+
|
|
156
|
+
|
|
157
|
+
def _not_replayable(*_args: Any, **_kwargs: Any) -> Any:
|
|
158
|
+
raise LookupError("this operation cannot be re-evaluated")
|
|
159
|
+
|
|
160
|
+
|
|
161
|
+
NOT_REPLAYABLE: Final = _not_replayable
|
|
162
|
+
"""Marks a node (``Node.fn``) whose operation must never be re-run."""
|
|
163
|
+
|
|
164
|
+
|
|
165
|
+
# --------------------------------------------------------------------------
|
|
166
|
+
# Active trace
|
|
167
|
+
# --------------------------------------------------------------------------
|
|
168
|
+
|
|
169
|
+
|
|
170
|
+
class Observer:
|
|
171
|
+
"""Stands in for the active trace while a ``@traced`` function runs.
|
|
172
|
+
|
|
173
|
+
Nothing is recorded, but every node of the trace whose value the function
|
|
174
|
+
touches is collected, so the call's true inputs are known even when they
|
|
175
|
+
reach it through closures, globals, sets or dict views.
|
|
176
|
+
"""
|
|
177
|
+
|
|
178
|
+
__slots__ = ("touched", "trace")
|
|
179
|
+
|
|
180
|
+
def __init__(self, trace: Trace) -> None:
|
|
181
|
+
self.trace = trace
|
|
182
|
+
self.touched: dict[int, Node] = {}
|
|
183
|
+
|
|
184
|
+
def touch_node(self, node: Node | None) -> None:
|
|
185
|
+
if node is not None and node.trace is self.trace:
|
|
186
|
+
self.touched[node.id] = node
|
|
187
|
+
|
|
188
|
+
def touch(self, *values: Any) -> None:
|
|
189
|
+
for value in values:
|
|
190
|
+
if isinstance(value, Traced):
|
|
191
|
+
node = value._rt_node
|
|
192
|
+
if node.trace is self.trace:
|
|
193
|
+
self.touched[node.id] = node
|
|
194
|
+
elif type(value) in (list, tuple, dict):
|
|
195
|
+
self.touch(*(value.values() if type(value) is dict else value))
|
|
196
|
+
|
|
197
|
+
|
|
198
|
+
_active: ContextVar[Trace | Observer | None] = ContextVar("redroot_active_trace", default=None)
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
def active_trace() -> Trace | None:
|
|
202
|
+
"""Return the trace recording in the current context, if any."""
|
|
203
|
+
trace = _active.get()
|
|
204
|
+
return None if type(trace) is Observer else trace # type: ignore[return-value]
|
|
205
|
+
|
|
206
|
+
|
|
207
|
+
# --------------------------------------------------------------------------
|
|
208
|
+
# Traced values
|
|
209
|
+
# --------------------------------------------------------------------------
|
|
210
|
+
|
|
211
|
+
|
|
212
|
+
class Traced:
|
|
213
|
+
"""Base class of every traced value type.
|
|
214
|
+
|
|
215
|
+
Each instance carries the :class:`Node` describing how it was produced.
|
|
216
|
+
"""
|
|
217
|
+
|
|
218
|
+
__slots__ = ()
|
|
219
|
+
_rt_node: Node
|
|
220
|
+
_base: type[Any]
|
|
221
|
+
"""The builtin type this traced type stands for."""
|
|
222
|
+
|
|
223
|
+
|
|
224
|
+
def unwrap(value: Any) -> Any:
|
|
225
|
+
"""Return the plain value behind ``value`` (``value`` itself if untraced).
|
|
226
|
+
|
|
227
|
+
Never triggers guards, so it is safe to use anywhere.
|
|
228
|
+
"""
|
|
229
|
+
return value._rt_node.value if isinstance(value, Traced) else value
|
|
230
|
+
|
|
231
|
+
|
|
232
|
+
def node_of(value: Any) -> Node | None:
|
|
233
|
+
"""Return the :class:`Node` behind a traced value, or ``None``."""
|
|
234
|
+
return value._rt_node if isinstance(value, Traced) else None
|
|
235
|
+
|
|
236
|
+
|
|
237
|
+
def same_value(a: Any, b: Any) -> bool:
|
|
238
|
+
"""Whether two plain values are interchangeable as results.
|
|
239
|
+
|
|
240
|
+
Stricter than ``==``: types must match (``1 != 1.0 != True``), ``-0.0``
|
|
241
|
+
differs from ``0.0``, ``Decimal("1.0")`` differs from ``Decimal("1.00")``
|
|
242
|
+
(they print differently), and NaN equals NaN.
|
|
243
|
+
"""
|
|
244
|
+
kind = type(a)
|
|
245
|
+
if kind is not type(b):
|
|
246
|
+
return False
|
|
247
|
+
if kind is float:
|
|
248
|
+
if a != a: # NaN
|
|
249
|
+
return bool(b != b)
|
|
250
|
+
return bool(a == b and math.copysign(1.0, a) == math.copysign(1.0, b))
|
|
251
|
+
if kind is Decimal:
|
|
252
|
+
return bool(a.as_tuple() == b.as_tuple())
|
|
253
|
+
if kind is list or kind is tuple or isinstance(a, tuple):
|
|
254
|
+
return len(a) == len(b) and all(same_value(x, y) for x, y in zip(a, b, strict=True))
|
|
255
|
+
if kind is dict:
|
|
256
|
+
return list(a) == list(b) and all(same_value(a[k], b[k]) for k in a)
|
|
257
|
+
try:
|
|
258
|
+
result = a == b
|
|
259
|
+
except Exception:
|
|
260
|
+
return False
|
|
261
|
+
return result if isinstance(result, bool) else False
|
|
262
|
+
|
|
263
|
+
|
|
264
|
+
# Exact plain type -> factory building the traced counterpart around a node.
|
|
265
|
+
_WRAPPERS: dict[type, Callable[[Any, Node], Any]] = {}
|
|
266
|
+
|
|
267
|
+
|
|
268
|
+
def register_wrapper(plain_type: type, factory: Callable[[Any, Node], Any]) -> None:
|
|
269
|
+
"""Declare which traced type wraps results of ``plain_type``."""
|
|
270
|
+
_WRAPPERS[plain_type] = factory
|
|
271
|
+
|
|
272
|
+
|
|
273
|
+
# --------------------------------------------------------------------------
|
|
274
|
+
# Recording
|
|
275
|
+
# --------------------------------------------------------------------------
|
|
276
|
+
|
|
277
|
+
|
|
278
|
+
def add_node(
|
|
279
|
+
trace: Trace,
|
|
280
|
+
op: str,
|
|
281
|
+
args: tuple[Any, ...],
|
|
282
|
+
kwargs: dict[str, Any] | None,
|
|
283
|
+
value: Any,
|
|
284
|
+
kind: str,
|
|
285
|
+
*,
|
|
286
|
+
key: str | None = None,
|
|
287
|
+
meta: dict[str, Any] | None = None,
|
|
288
|
+
ctx: int | None = None,
|
|
289
|
+
fn: Callable[..., Any] | None = None,
|
|
290
|
+
) -> Node:
|
|
291
|
+
"""Append a node to ``trace`` and return it."""
|
|
292
|
+
node = Node(op, args, kwargs, value, kind, key=key, meta=meta)
|
|
293
|
+
nodes = trace._nodes
|
|
294
|
+
node.id = len(nodes)
|
|
295
|
+
node.trace = trace
|
|
296
|
+
node.ctx = ctx
|
|
297
|
+
node.fn = fn
|
|
298
|
+
nodes.append(node)
|
|
299
|
+
return node
|
|
300
|
+
|
|
301
|
+
|
|
302
|
+
def decimal_context_id(trace: Trace) -> int:
|
|
303
|
+
"""Intern the current ``decimal`` context in ``trace`` and return its index."""
|
|
304
|
+
c = decimal.getcontext()
|
|
305
|
+
fingerprint = (id(c), c.prec, c.rounding, c.Emin, c.Emax, c.capitals, c.clamp)
|
|
306
|
+
index = trace._context_ids.get(fingerprint)
|
|
307
|
+
if index is None:
|
|
308
|
+
index = trace._context_ids[fingerprint] = len(trace._contexts)
|
|
309
|
+
trace._contexts.append(c.copy())
|
|
310
|
+
return index
|
|
311
|
+
|
|
312
|
+
|
|
313
|
+
def _uses_decimal(values: Any) -> bool:
|
|
314
|
+
return any(type(v) is Decimal for v in values)
|
|
315
|
+
|
|
316
|
+
|
|
317
|
+
def _is_namedtuple(value: Any) -> bool:
|
|
318
|
+
return isinstance(value, tuple) and hasattr(type(value), "_fields")
|
|
319
|
+
|
|
320
|
+
|
|
321
|
+
def is_projectable(value: Any) -> bool:
|
|
322
|
+
"""Whether ``value`` is a container whose items are traced individually."""
|
|
323
|
+
kind = type(value)
|
|
324
|
+
return kind is list or kind is tuple or kind is dict or _is_namedtuple(value)
|
|
325
|
+
|
|
326
|
+
|
|
327
|
+
def record_result(
|
|
328
|
+
trace: Trace,
|
|
329
|
+
op: str,
|
|
330
|
+
args: tuple[Any, ...],
|
|
331
|
+
kwargs: dict[str, Any] | None,
|
|
332
|
+
value: Any,
|
|
333
|
+
meta: dict[str, Any] | None = None,
|
|
334
|
+
ctx: int | None = None,
|
|
335
|
+
fn: Callable[..., Any] | None = None,
|
|
336
|
+
) -> Any:
|
|
337
|
+
"""Record ``value`` as the result of ``op(*args, **kwargs)``.
|
|
338
|
+
|
|
339
|
+
Returns the traced form of ``value``:
|
|
340
|
+
|
|
341
|
+
* a scalar with a traced counterpart (``int``, ``float``, ``Decimal``,
|
|
342
|
+
``str``) is wrapped;
|
|
343
|
+
* a ``list``, ``tuple`` or ``dict`` is rebuilt from per-item projections,
|
|
344
|
+
plus a guard on its shape (length or keys);
|
|
345
|
+
* anything else (``bool``, ``None``, objects) is returned unchanged and
|
|
346
|
+
recorded as a guard, since it leaves tracking.
|
|
347
|
+
"""
|
|
348
|
+
kind = type(value)
|
|
349
|
+
if isinstance(value, Traced) or kind is list or kind is dict or isinstance(value, tuple):
|
|
350
|
+
value = unwrap_deep(value) # e.g. a function returned one of its traced arguments
|
|
351
|
+
factory = _WRAPPERS.get(type(value))
|
|
352
|
+
if factory is not None:
|
|
353
|
+
return factory(
|
|
354
|
+
value, add_node(trace, op, args, kwargs, value, OP, meta=meta, ctx=ctx, fn=fn)
|
|
355
|
+
)
|
|
356
|
+
if is_projectable(value):
|
|
357
|
+
node = add_node(trace, op, args, kwargs, value, OP, meta=meta, ctx=ctx, fn=fn)
|
|
358
|
+
return project(trace, node)
|
|
359
|
+
add_node(trace, op, args, kwargs, value, GUARD, meta=meta, ctx=ctx, fn=fn)
|
|
360
|
+
return value
|
|
361
|
+
|
|
362
|
+
|
|
363
|
+
def record_raise(
|
|
364
|
+
trace: Trace,
|
|
365
|
+
op: str,
|
|
366
|
+
args: tuple[Any, ...],
|
|
367
|
+
kwargs: dict[str, Any] | None,
|
|
368
|
+
exc: BaseException,
|
|
369
|
+
fn: Callable[..., Any] | None = None,
|
|
370
|
+
) -> None:
|
|
371
|
+
"""Record that ``op(*args, **kwargs)`` raised ``exc``, as a guard."""
|
|
372
|
+
ctx = decimal_context_id(trace) if _uses_decimal(unwrap_deep(args)) else None
|
|
373
|
+
add_node(trace, op, args, kwargs, Raised.of(exc), GUARD, ctx=ctx, fn=fn)
|
|
374
|
+
|
|
375
|
+
|
|
376
|
+
def project(trace: Trace, node: Node) -> Any:
|
|
377
|
+
"""Return a copy of the container ``node.value`` whose items are traced."""
|
|
378
|
+
value = node.value
|
|
379
|
+
add_node(trace, "shape", (node,), None, shape(value), GUARD)
|
|
380
|
+
if type(value) is dict:
|
|
381
|
+
return {k: record_result(trace, "getitem", (node, k), None, v) for k, v in value.items()}
|
|
382
|
+
items = [record_result(trace, "getitem", (node, i), None, v) for i, v in enumerate(value)]
|
|
383
|
+
if type(value) is list:
|
|
384
|
+
return items
|
|
385
|
+
return rebuild_tuple(value, items)
|
|
386
|
+
|
|
387
|
+
|
|
388
|
+
def encode(obj: Any, trace: Trace) -> tuple[Any, Any, bool]:
|
|
389
|
+
"""Split a call argument into ``(reference, plain, linked)``.
|
|
390
|
+
|
|
391
|
+
``reference`` is what gets stored in ``Node.args``: traced values from
|
|
392
|
+
``trace`` become their :class:`Node`; lists, tuples and dicts are walked.
|
|
393
|
+
``plain`` is the same structure with every traced value unwrapped.
|
|
394
|
+
``linked`` says whether any part of ``obj`` belongs to ``trace``.
|
|
395
|
+
"""
|
|
396
|
+
if isinstance(obj, Traced):
|
|
397
|
+
node = obj._rt_node
|
|
398
|
+
if node.trace is trace:
|
|
399
|
+
return node, node.value, True
|
|
400
|
+
return node.value, node.value, False
|
|
401
|
+
kind = type(obj)
|
|
402
|
+
if kind is list or isinstance(obj, tuple):
|
|
403
|
+
refs: list[Any] = []
|
|
404
|
+
plains: list[Any] = []
|
|
405
|
+
linked = changed = False
|
|
406
|
+
for item in obj:
|
|
407
|
+
ref, plain, item_linked = encode(item, trace)
|
|
408
|
+
refs.append(ref)
|
|
409
|
+
plains.append(plain)
|
|
410
|
+
linked = linked or item_linked
|
|
411
|
+
changed = changed or plain is not item
|
|
412
|
+
if not changed:
|
|
413
|
+
return obj, obj, False
|
|
414
|
+
if kind is list:
|
|
415
|
+
return refs, plains, linked
|
|
416
|
+
return rebuild_tuple(obj, refs), rebuild_tuple(obj, plains), linked
|
|
417
|
+
if kind is dict:
|
|
418
|
+
ref_d: dict[Any, Any] = {}
|
|
419
|
+
plain_d: dict[Any, Any] = {}
|
|
420
|
+
linked = changed = False
|
|
421
|
+
for k, item in obj.items():
|
|
422
|
+
ref, plain, item_linked = encode(item, trace)
|
|
423
|
+
ref_d[k] = ref
|
|
424
|
+
plain_d[k] = plain
|
|
425
|
+
linked = linked or item_linked
|
|
426
|
+
changed = changed or plain is not item
|
|
427
|
+
if not changed:
|
|
428
|
+
return obj, obj, False
|
|
429
|
+
return ref_d, plain_d, linked
|
|
430
|
+
return obj, obj, False
|
|
431
|
+
|
|
432
|
+
|
|
433
|
+
def unwrap_deep(obj: Any) -> Any:
|
|
434
|
+
"""Unwrap traced values inside lists, tuples and dicts."""
|
|
435
|
+
if isinstance(obj, Traced):
|
|
436
|
+
return obj._rt_node.value
|
|
437
|
+
kind = type(obj)
|
|
438
|
+
if kind is list:
|
|
439
|
+
return [unwrap_deep(item) for item in obj]
|
|
440
|
+
if isinstance(obj, tuple):
|
|
441
|
+
return rebuild_tuple(obj, [unwrap_deep(item) for item in obj])
|
|
442
|
+
if kind is dict:
|
|
443
|
+
return {k: unwrap_deep(v) for k, v in obj.items()}
|
|
444
|
+
return obj
|
|
445
|
+
|
|
446
|
+
|
|
447
|
+
def call(
|
|
448
|
+
op: str,
|
|
449
|
+
fn: Callable[..., Any],
|
|
450
|
+
args: tuple[Any, ...],
|
|
451
|
+
kwargs: dict[str, Any] | None = None,
|
|
452
|
+
*,
|
|
453
|
+
meta: dict[str, Any] | None = None,
|
|
454
|
+
) -> Any:
|
|
455
|
+
"""Apply ``fn`` to the plain form of the arguments and record the call.
|
|
456
|
+
|
|
457
|
+
Records nothing (and returns the plain result) unless a trace is active
|
|
458
|
+
and at least one argument belongs to it.
|
|
459
|
+
"""
|
|
460
|
+
trace = _active.get()
|
|
461
|
+
if trace is None or isinstance(trace, Observer):
|
|
462
|
+
if trace is not None:
|
|
463
|
+
trace.touch(args, kwargs or {})
|
|
464
|
+
return fn(*unwrap_deep(args), **unwrap_deep(kwargs or {}))
|
|
465
|
+
refs, plains, linked = encode(args, trace)
|
|
466
|
+
kw_refs: dict[str, Any] | None = None
|
|
467
|
+
kw_plains: dict[str, Any] = {}
|
|
468
|
+
if kwargs:
|
|
469
|
+
kw_refs, kw_plains, kw_linked = encode(kwargs, trace)
|
|
470
|
+
linked = linked or kw_linked
|
|
471
|
+
if not linked:
|
|
472
|
+
return fn(*plains, **kw_plains)
|
|
473
|
+
try:
|
|
474
|
+
result = fn(*plains, **kw_plains)
|
|
475
|
+
except Exception as exc:
|
|
476
|
+
record_raise(trace, op, tuple(refs), kw_refs, exc)
|
|
477
|
+
raise
|
|
478
|
+
ctx = (
|
|
479
|
+
decimal_context_id(trace)
|
|
480
|
+
if type(result) is Decimal or _uses_decimal(plains) or _uses_decimal(kw_plains.values())
|
|
481
|
+
else None
|
|
482
|
+
)
|
|
483
|
+
return record_result(trace, op, tuple(refs), kw_refs, result, meta, ctx)
|
|
484
|
+
|
|
485
|
+
|
|
486
|
+
def binary(op: str, fn: Callable[[Any, Any], Any], left: Any, right: Any) -> Any:
|
|
487
|
+
"""Fast path of :func:`call` for two scalar operands."""
|
|
488
|
+
if isinstance(left, Traced):
|
|
489
|
+
left_node: Node | None = left._rt_node
|
|
490
|
+
left = left_node.value # type: ignore[union-attr]
|
|
491
|
+
else:
|
|
492
|
+
left_node = None
|
|
493
|
+
if isinstance(right, Traced):
|
|
494
|
+
right_node: Node | None = right._rt_node
|
|
495
|
+
right = right_node.value # type: ignore[union-attr]
|
|
496
|
+
else:
|
|
497
|
+
right_node = None
|
|
498
|
+
trace = _active.get()
|
|
499
|
+
if trace is None:
|
|
500
|
+
return fn(left, right)
|
|
501
|
+
if isinstance(trace, Observer):
|
|
502
|
+
trace.touch_node(left_node)
|
|
503
|
+
trace.touch_node(right_node)
|
|
504
|
+
return fn(left, right)
|
|
505
|
+
left_ref, right_ref, linked = left, right, False
|
|
506
|
+
if left_node is not None and left_node.trace is trace:
|
|
507
|
+
left_ref, linked = left_node, True
|
|
508
|
+
if right_node is not None and right_node.trace is trace:
|
|
509
|
+
right_ref, linked = right_node, True
|
|
510
|
+
if not linked:
|
|
511
|
+
return fn(left, right)
|
|
512
|
+
try:
|
|
513
|
+
result = fn(left, right)
|
|
514
|
+
except Exception as exc:
|
|
515
|
+
record_raise(trace, op, (left_ref, right_ref), None, exc)
|
|
516
|
+
raise
|
|
517
|
+
ctx = (
|
|
518
|
+
decimal_context_id(trace)
|
|
519
|
+
if type(result) is Decimal or type(left) is Decimal or type(right) is Decimal
|
|
520
|
+
else None
|
|
521
|
+
)
|
|
522
|
+
return record_result(trace, op, (left_ref, right_ref), None, result, None, ctx)
|
|
523
|
+
|
|
524
|
+
|
|
525
|
+
def observe_value(op: str, fn: Callable[..., Any], *operands: Any) -> Any:
|
|
526
|
+
"""Like :func:`binary` but always records a guard and returns a plain result.
|
|
527
|
+
|
|
528
|
+
Used for comparisons, truth tests and conversions that must return a
|
|
529
|
+
builtin type (``__bool__``, ``__int__``, ``__hash__``...).
|
|
530
|
+
"""
|
|
531
|
+
trace = _active.get()
|
|
532
|
+
plains = [
|
|
533
|
+
operand._rt_node.value if isinstance(operand, Traced) else operand for operand in operands
|
|
534
|
+
]
|
|
535
|
+
if trace is None:
|
|
536
|
+
return fn(*plains)
|
|
537
|
+
if isinstance(trace, Observer):
|
|
538
|
+
trace.touch(*operands)
|
|
539
|
+
return fn(*plains)
|
|
540
|
+
refs: list[Any] = []
|
|
541
|
+
linked = False
|
|
542
|
+
for operand, plain in zip(operands, plains, strict=True):
|
|
543
|
+
if isinstance(operand, Traced) and operand._rt_node.trace is trace:
|
|
544
|
+
refs.append(operand._rt_node)
|
|
545
|
+
linked = True
|
|
546
|
+
else:
|
|
547
|
+
refs.append(plain)
|
|
548
|
+
if not linked:
|
|
549
|
+
return fn(*plains)
|
|
550
|
+
try:
|
|
551
|
+
result = fn(*plains)
|
|
552
|
+
except Exception as exc:
|
|
553
|
+
record_raise(trace, op, tuple(refs), None, exc)
|
|
554
|
+
raise
|
|
555
|
+
add_node(trace, op, tuple(refs), None, result, GUARD)
|
|
556
|
+
return result
|