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/trace.py ADDED
@@ -0,0 +1,474 @@
1
+ """The :class:`Trace`: one recorded execution and its value graph."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import dataclasses
6
+ import decimal
7
+ import sys
8
+ import uuid
9
+ from collections import deque
10
+ from collections.abc import Callable, Mapping, Sequence
11
+ from contextvars import Token
12
+ from dataclasses import dataclass
13
+ from types import MappingProxyType, TracebackType
14
+ from typing import TYPE_CHECKING, Any, TypeVar
15
+
16
+ from redroot._core import (
17
+ _WRAPPERS,
18
+ GUARD,
19
+ LEAF,
20
+ Node,
21
+ Traced,
22
+ _active,
23
+ active_trace,
24
+ add_node,
25
+ unwrap,
26
+ )
27
+ from redroot.ops import get_op
28
+ from redroot.paths import child_key, format_key
29
+
30
+ if TYPE_CHECKING:
31
+ from redroot.propagation import Propagation
32
+
33
+ __all__ = ["Output", "Trace", "run", "track"]
34
+
35
+ T = TypeVar("T")
36
+
37
+
38
+ @dataclass(frozen=True, slots=True)
39
+ class Output:
40
+ """A value the workflow produced, bound to the node that computed it.
41
+
42
+ Attributes:
43
+ key: Semantic id, e.g. ``"out:form1.line_12"``.
44
+ node: The node behind the value, or ``None`` if the value was not
45
+ traced ("not linked to the inputs").
46
+ value: The plain value at collection time.
47
+ """
48
+
49
+ key: str
50
+ node: Node | None
51
+ value: Any
52
+
53
+ @property
54
+ def linked(self) -> bool:
55
+ """Whether the value is linked to traced inputs."""
56
+ return self.node is not None
57
+
58
+
59
+ class Trace:
60
+ """A recording of one execution, as a graph from inputs to outputs.
61
+
62
+ Use it as a context manager. Inside the block, operations on traced
63
+ values are recorded::
64
+
65
+ with Trace("run-42") as trace:
66
+ data = trace.track(extracted, root="ext")
67
+ result = workflow(data)
68
+ trace.collect(result, root="out")
69
+
70
+ trace.sources("out:form1.line_12") # which inputs produced it
71
+ trace.propagate({"ext:a": 900}) # what changes if an input does
72
+
73
+ A trace is not thread-safe: record into it from one thread at a time.
74
+ """
75
+
76
+ def __init__(self, name: str | None = None, *, id: str | None = None) -> None:
77
+ self.id = id or uuid.uuid4().hex
78
+ self.name = name
79
+ self.decimal_context: decimal.Context = decimal.getcontext().copy()
80
+ """Decimal context used for re-evaluation; captured when recording ends."""
81
+ self._nodes: list[Node] = []
82
+ self._leaves: dict[str, Node] = {}
83
+ self._untracked: dict[str, Any] = {}
84
+ self._outputs: dict[str, Output] = {}
85
+ self._tokens: list[Token[Any]] = []
86
+ # decimal contexts nodes were computed under (see Node.ctx)
87
+ self._contexts: list[decimal.Context] = []
88
+ self._context_ids: dict[tuple[Any, ...], int] = {}
89
+
90
+ def __repr__(self) -> str:
91
+ name = f" {self.name!r}" if self.name else ""
92
+ return (
93
+ f"<Trace{name} nodes={len(self._nodes)} inputs={len(self._leaves)} "
94
+ f"outputs={len(self._outputs)} guards={len(self.guards)}>"
95
+ )
96
+
97
+ # ------------------------------------------------------------------
98
+ # Activation
99
+ # ------------------------------------------------------------------
100
+
101
+ def __enter__(self) -> Trace:
102
+ self._tokens.append(_active.set(self))
103
+ return self
104
+
105
+ def __exit__(
106
+ self,
107
+ exc_type: type[BaseException] | None,
108
+ exc: BaseException | None,
109
+ tb: TracebackType | None,
110
+ ) -> None:
111
+ _active.reset(self._tokens.pop())
112
+ self.decimal_context = decimal.getcontext().copy()
113
+
114
+ @property
115
+ def active(self) -> bool:
116
+ """Whether this trace is recording in the current context."""
117
+ return active_trace() is self
118
+
119
+ # ------------------------------------------------------------------
120
+ # Contents
121
+ # ------------------------------------------------------------------
122
+
123
+ @property
124
+ def nodes(self) -> Sequence[Node]:
125
+ """Every node, in recording (topological) order. Do not mutate."""
126
+ return self._nodes
127
+
128
+ @property
129
+ def decimal_contexts(self) -> Sequence[decimal.Context]:
130
+ """The ``decimal`` contexts that ``Decimal`` nodes were computed under."""
131
+ return tuple(self._contexts)
132
+
133
+ @property
134
+ def guards(self) -> list[Node]:
135
+ """Nodes whose values left tracking (comparisons, conversions...)."""
136
+ return [node for node in self._nodes if node.kind == GUARD]
137
+
138
+ @property
139
+ def inputs(self) -> Mapping[str, Node]:
140
+ """Keyed input (leaf) nodes."""
141
+ return MappingProxyType(self._leaves)
142
+
143
+ @property
144
+ def untracked_inputs(self) -> Mapping[str, Any]:
145
+ """Inputs passed to :meth:`track` that could not be traced (``None``, ``bool``...)."""
146
+ return MappingProxyType(self._untracked)
147
+
148
+ @property
149
+ def outputs(self) -> Mapping[str, Output]:
150
+ """Collected outputs by key."""
151
+ return MappingProxyType(self._outputs)
152
+
153
+ def output_values(self) -> dict[str, Any]:
154
+ """Plain value of every output, by key."""
155
+ return {key: out.value for key, out in self._outputs.items()}
156
+
157
+ # ------------------------------------------------------------------
158
+ # Inputs
159
+ # ------------------------------------------------------------------
160
+
161
+ def _add_leaf(self, value: Any, key: str | None, meta: dict[str, Any] | None) -> Node:
162
+ if key is not None:
163
+ self._check_new_key(key)
164
+ node = add_node(self, LEAF, (), None, value, LEAF, key=key, meta=meta)
165
+ if key is not None:
166
+ self._leaves[key] = node
167
+ return node
168
+
169
+ def _check_new_key(self, key: str) -> None:
170
+ if key in self._leaves or key in self._untracked or key in self._outputs:
171
+ raise ValueError(f"key {key!r} is already used in this trace")
172
+
173
+ def leaf(
174
+ self, value: Any, key: str | None = None, *, meta: dict[str, Any] | None = None
175
+ ) -> Any:
176
+ """Record ``value`` as an input and return its traced form.
177
+
178
+ Args:
179
+ value: An ``int``, ``float``, ``Decimal`` or ``str``.
180
+ key: Semantic id used to refer to the input, e.g. for edits.
181
+ meta: Free-form metadata stored on the node.
182
+ """
183
+ plain = unwrap(value)
184
+ factory = _WRAPPERS.get(type(plain))
185
+ if factory is None:
186
+ raise TypeError(f"cannot trace a value of type {type(plain).__name__}")
187
+ return factory(plain, self._add_leaf(plain, key, meta))
188
+
189
+ def track(self, data: T, root: str = "in") -> T:
190
+ """Return a copy of ``data`` whose scalar values are traced inputs.
191
+
192
+ Dicts, lists and tuples are copied; every ``int``, ``float``,
193
+ ``Decimal`` and ``str`` inside becomes an input keyed by its path,
194
+ e.g. ``"ext:bank.line_items[3].amount"``. Other values (``None``,
195
+ ``bool``, objects) are kept as-is and listed in
196
+ :attr:`untracked_inputs`: editing them requires re-execution.
197
+ """
198
+ return self._track(data, format_key(root)) # type: ignore[no-any-return]
199
+
200
+ def _track(self, value: Any, key: str) -> Any:
201
+ kind = type(value)
202
+ if kind is dict:
203
+ return {k: self._track(v, child_key(key, k)) for k, v in value.items()}
204
+ if kind is list:
205
+ return [self._track(v, child_key(key, i)) for i, v in enumerate(value)]
206
+ if kind is tuple:
207
+ return tuple(self._track(v, child_key(key, i)) for i, v in enumerate(value))
208
+ plain = unwrap(value)
209
+ factory = _WRAPPERS.get(type(plain))
210
+ if factory is not None:
211
+ return factory(plain, self._add_leaf(plain, key, None))
212
+ self._check_new_key(key)
213
+ self._untracked[key] = value
214
+ return value
215
+
216
+ # ------------------------------------------------------------------
217
+ # Outputs
218
+ # ------------------------------------------------------------------
219
+
220
+ def output(self, key: str, value: Any) -> Output:
221
+ """Record ``value`` as the output ``key``."""
222
+ self._check_new_key(key)
223
+ node = value._rt_node if isinstance(value, Traced) else None
224
+ if node is not None and node.trace is not self:
225
+ node = None
226
+ out = Output(key, node, unwrap(value))
227
+ self._outputs[key] = out
228
+ return out
229
+
230
+ def collect(self, data: Any, root: str = "out") -> None:
231
+ """Record every scalar value inside ``data`` as an output.
232
+
233
+ Walks mappings, lists, tuples, dataclasses and Pydantic models; each
234
+ leaf value becomes an output keyed by its path under ``root``.
235
+ """
236
+ self._collect(data, format_key(root))
237
+
238
+ def _collect(self, value: Any, key: str) -> None:
239
+ if isinstance(value, Traced):
240
+ self.output(key, value)
241
+ elif isinstance(value, Mapping):
242
+ for k, v in value.items():
243
+ self._collect(v, child_key(key, k))
244
+ elif isinstance(value, (list, tuple)):
245
+ for i, v in enumerate(value):
246
+ self._collect(v, child_key(key, i))
247
+ elif _is_pydantic_model(value):
248
+ for name in type(value).model_fields:
249
+ self._collect(getattr(value, name), child_key(key, name))
250
+ elif dataclasses.is_dataclass(value) and not isinstance(value, type):
251
+ for field in dataclasses.fields(value):
252
+ self._collect(getattr(value, field.name), child_key(key, field.name))
253
+ else:
254
+ self.output(key, value)
255
+
256
+ # ------------------------------------------------------------------
257
+ # Navigation
258
+ # ------------------------------------------------------------------
259
+
260
+ def node(self, ref: Any) -> Node:
261
+ """Resolve ``ref`` to one of this trace's nodes.
262
+
263
+ ``ref`` may be a :class:`Node`, a traced value, or the key of an
264
+ input or of a linked output.
265
+ """
266
+ if isinstance(ref, Node):
267
+ node = ref
268
+ elif isinstance(ref, Traced):
269
+ node = ref._rt_node
270
+ elif isinstance(ref, str):
271
+ if ref in self._leaves:
272
+ return self._leaves[ref]
273
+ out = self._outputs.get(ref)
274
+ if out is None:
275
+ raise KeyError(f"no input or output named {ref!r}")
276
+ if out.node is None:
277
+ raise LookupError(f"output {ref!r} is not linked to any traced input")
278
+ return out.node
279
+ else:
280
+ raise TypeError(f"cannot resolve a node from {type(ref).__name__}")
281
+ if node.trace is not self:
282
+ raise ValueError("node belongs to a different trace")
283
+ return node
284
+
285
+ def ancestors(self, ref: Any) -> list[Node]:
286
+ """``ref``'s node and every node it was computed from, in id order."""
287
+ start = self.node(ref)
288
+ seen: dict[int, Node] = {start.id: start}
289
+ queue = deque([start])
290
+ while queue:
291
+ for operand in queue.popleft().operands():
292
+ if operand.id not in seen:
293
+ seen[operand.id] = operand
294
+ queue.append(operand)
295
+ return [seen[i] for i in sorted(seen)]
296
+
297
+ def sources(self, ref: Any) -> list[Node]:
298
+ """The input nodes ``ref`` was computed from.
299
+
300
+ This answers "which extracted fields produced this value?".
301
+ """
302
+ return [node for node in self.ancestors(ref) if node.kind == LEAF]
303
+
304
+ def descendants(self, ref: Any) -> list[Node]:
305
+ """``ref``'s node and every node computed from it, in id order."""
306
+ start = self.node(ref)
307
+ reached = {start.id}
308
+ found = [start]
309
+ for node in self._nodes[start.id + 1 :]:
310
+ if any(operand.id in reached for operand in node.operands()):
311
+ reached.add(node.id)
312
+ found.append(node)
313
+ return found
314
+
315
+ def dependents(self, ref: Any) -> list[str]:
316
+ """Keys of the outputs computed from ``ref``.
317
+
318
+ This answers "which outputs use this extracted field?".
319
+ """
320
+ reached = {node.id for node in self.descendants(ref)}
321
+ return [
322
+ key
323
+ for key, out in self._outputs.items()
324
+ if out.node is not None and out.node.id in reached
325
+ ]
326
+
327
+ def explain(self, ref: Any, *, max_depth: int = 12) -> str:
328
+ """Render how ``ref`` was computed as an expression over the inputs.
329
+
330
+ >>> trace.explain("out:total") # doctest: +SKIP
331
+ '(ext:a + ext:b) * 0.8'
332
+
333
+ Sub-expressions deeper than ``max_depth`` are shown as ``#<node id>``.
334
+ """
335
+ return _render(self.node(ref), max_depth, top=True)
336
+
337
+ # ------------------------------------------------------------------
338
+ # Propagation and serialization
339
+ # ------------------------------------------------------------------
340
+
341
+ def propagate(self, edits: Mapping[str, Any]) -> Propagation:
342
+ """Recompute every output after changing some inputs.
343
+
344
+ See :func:`redroot.propagation.propagate`.
345
+ """
346
+ from redroot.propagation import propagate
347
+
348
+ return propagate(self, edits)
349
+
350
+ def to_dict(self) -> dict[str, Any]:
351
+ """Serialize to JSON-compatible data. See :mod:`redroot.serialization`."""
352
+ from redroot.serialization import trace_to_dict
353
+
354
+ return trace_to_dict(self)
355
+
356
+ def to_json(self, **kwargs: Any) -> str:
357
+ """Serialize to a JSON string; ``kwargs`` go to :func:`json.dumps`."""
358
+ import json
359
+
360
+ return json.dumps(self.to_dict(), **kwargs)
361
+
362
+ @classmethod
363
+ def from_dict(cls, data: Mapping[str, Any]) -> Trace:
364
+ """Rebuild a trace from :meth:`to_dict` output."""
365
+ from redroot.serialization import trace_from_dict
366
+
367
+ return trace_from_dict(data)
368
+
369
+ @classmethod
370
+ def from_json(cls, text: str) -> Trace:
371
+ """Rebuild a trace from :meth:`to_json` output."""
372
+ import json
373
+
374
+ return cls.from_dict(json.loads(text))
375
+
376
+
377
+ def _is_pydantic_model(value: Any) -> bool:
378
+ pydantic = sys.modules.get("pydantic")
379
+ return pydantic is not None and isinstance(value, pydantic.BaseModel)
380
+
381
+
382
+ def _render(node: Node, depth: int, top: bool = False) -> str:
383
+ if node.kind == LEAF:
384
+ return node.key if node.key is not None else repr(node.value)
385
+ if depth <= 0:
386
+ return f"#{node.id}"
387
+ args = [_render_arg(arg, depth - 1) for arg in node.args]
388
+ args += [f"{k}={_render_arg(v, depth - 1)}" for k, v in (node.kwargs or {}).items()]
389
+ if node.op == "fstring":
390
+ return _render_fstring(node.args, depth - 1)
391
+ if node.op == "getslice" and len(args) == 4:
392
+ start, stop, step = (
393
+ "" if a is None else t for a, t in zip(node.args[1:], args[1:], strict=True)
394
+ )
395
+ return f"{args[0]}[{start}:{stop}{':' + step if step else ''}]"
396
+ spec = get_op(node.op)
397
+ style = spec.style if spec is not None else "call"
398
+ symbol = spec.symbol if spec is not None else None
399
+ text: str
400
+ if style == "infix" and len(args) == 2:
401
+ text = f"{args[0]} {symbol} {args[1]}"
402
+ elif style == "contains" and len(args) == 2:
403
+ text = f"{args[1]} in {args[0]}"
404
+ elif style == "prefix" and len(args) == 1:
405
+ return f"{symbol}{args[0]}"
406
+ elif style == "getitem" and len(args) == 2:
407
+ return f"{args[0]}[{args[1]}]"
408
+ elif style == "method" and args:
409
+ return f"{args[0]}.{symbol}({', '.join(args[1:])})"
410
+ else:
411
+ return f"{symbol or node.op}({', '.join(args)})"
412
+ return text if top else f"({text})"
413
+
414
+
415
+ def _render_fstring(parts: tuple[Any, ...], depth: int) -> str:
416
+ out = []
417
+ for part in parts:
418
+ if isinstance(part, str):
419
+ out.append(part.replace("{", "{{").replace("}", "}}"))
420
+ continue
421
+ value, conversion, spec = part
422
+ conv = f"!{conversion}" if conversion else ""
423
+ spec_text = f":{_render_arg(spec, depth) if type(spec) is Node else spec}" if spec else ""
424
+ out.append(f"{{{_render_arg(value, depth)}{conv}{spec_text}}}")
425
+ return 'f"' + "".join(out).replace('"', '\\"') + '"'
426
+
427
+
428
+ def _render_arg(arg: Any, depth: int) -> str:
429
+ kind = type(arg)
430
+ if kind is Node:
431
+ return _render(arg, depth)
432
+ if kind is list:
433
+ return "[" + ", ".join(_render_arg(a, depth) for a in arg) + "]"
434
+ if kind is tuple:
435
+ inner = ", ".join(_render_arg(a, depth) for a in arg)
436
+ return f"({inner},)" if len(arg) == 1 else f"({inner})"
437
+ if kind is dict:
438
+ return "{" + ", ".join(f"{k!r}: {_render_arg(v, depth)}" for k, v in arg.items()) + "}"
439
+ return repr(arg)
440
+
441
+
442
+ # ----------------------------------------------------------------------
443
+ # Module-level helpers
444
+ # ----------------------------------------------------------------------
445
+
446
+
447
+ def track(data: T, root: str = "in") -> T:
448
+ """Trace ``data`` as inputs of the active trace. See :meth:`Trace.track`."""
449
+ trace = active_trace()
450
+ if trace is None:
451
+ raise RuntimeError("redroot.track() needs an active Trace (use 'with Trace():')")
452
+ return trace.track(data, root)
453
+
454
+
455
+ def run(
456
+ workflow: Callable[[Any], Any],
457
+ inputs: Any,
458
+ *,
459
+ input_root: str = "in",
460
+ output_root: str = "out",
461
+ name: str | None = None,
462
+ ) -> tuple[Trace, Any]:
463
+ """Run ``workflow(inputs)`` under a new trace.
464
+
465
+ Inputs are traced under ``input_root`` and every value in the returned
466
+ result is collected as an output under ``output_root``.
467
+
468
+ Returns:
469
+ The trace and the workflow's result.
470
+ """
471
+ with Trace(name) as trace:
472
+ result = workflow(trace.track(inputs, input_root))
473
+ trace.collect(result, output_root)
474
+ return trace, result