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/ops.py ADDED
@@ -0,0 +1,332 @@
1
+ """The operation registry.
2
+
3
+ Every node in a trace names the operation that produced it. The registry maps
4
+ that name to the plain Python function that computes it. The *same* function
5
+ is used when the value is first computed and whenever the graph re-evaluates
6
+ the node, so re-evaluation reproduces traced execution exactly.
7
+
8
+ Built-in operations (arithmetic, comparisons, conversions, ``str`` and
9
+ ``Decimal`` methods) are registered on import. User functions are registered
10
+ with :func:`redroot.traced`.
11
+ """
12
+
13
+ from __future__ import annotations
14
+
15
+ import bisect
16
+ import heapq
17
+ import math
18
+ import operator
19
+ import statistics
20
+ from collections.abc import Callable, Mapping
21
+ from dataclasses import dataclass
22
+ from decimal import Decimal
23
+ from typing import Any, Literal
24
+
25
+ Style = Literal["infix", "prefix", "call", "method", "getitem", "contains"]
26
+
27
+
28
+ @dataclass(frozen=True, slots=True)
29
+ class OpSpec:
30
+ """How to compute, and how to display, one kind of operation.
31
+
32
+ Attributes:
33
+ name: The name recorded on nodes.
34
+ fn: Computes the operation from plain argument values.
35
+ replayable: ``False`` for operations whose result cannot be recomputed
36
+ (e.g. LLM calls, web lookups). Nodes that depend on them become
37
+ stale instead of being re-evaluated when their inputs change.
38
+ symbol: The operator symbol (infix/prefix styles) or method name
39
+ (method style) used by :meth:`redroot.Trace.explain`.
40
+ style: How :meth:`redroot.Trace.explain` renders the operation.
41
+ """
42
+
43
+ name: str
44
+ fn: Callable[..., Any]
45
+ replayable: bool = True
46
+ symbol: str | None = None
47
+ style: Style = "call"
48
+
49
+
50
+ _REGISTRY: dict[str, OpSpec] = {}
51
+
52
+
53
+ def register_op(
54
+ name: str,
55
+ fn: Callable[..., Any],
56
+ *,
57
+ replayable: bool = True,
58
+ symbol: str | None = None,
59
+ style: Style = "call",
60
+ ) -> OpSpec:
61
+ """Register (or replace) the operation ``name``.
62
+
63
+ Operations must be pure: given equal arguments they must return equal
64
+ results and have no side effects, otherwise re-evaluation is meaningless.
65
+ Mark anything else ``replayable=False``.
66
+ """
67
+ spec = OpSpec(name, fn, replayable, symbol, style)
68
+ _REGISTRY[name] = spec
69
+ return spec
70
+
71
+
72
+ def get_op(name: str) -> OpSpec | None:
73
+ """Return the registered operation ``name``, or ``None``."""
74
+ return _REGISTRY.get(name)
75
+
76
+
77
+ def registered_ops() -> Mapping[str, OpSpec]:
78
+ """Return a read-only snapshot of the registry."""
79
+ return dict(_REGISTRY)
80
+
81
+
82
+ class Opaque:
83
+ """Stands in for a constant that could not be serialized.
84
+
85
+ Evaluating a node with an opaque argument fails, which marks the node
86
+ stale rather than silently producing a wrong value.
87
+ """
88
+
89
+ __slots__ = ("description",)
90
+
91
+ def __init__(self, description: str) -> None:
92
+ self.description = description
93
+
94
+ def __repr__(self) -> str:
95
+ return f"<opaque {self.description}>"
96
+
97
+
98
+ def shape(value: Any) -> Any:
99
+ """The structural signature of a container: its keys or its length."""
100
+ return tuple(value) if isinstance(value, dict) else len(value)
101
+
102
+
103
+ def observe(value: Any) -> Any:
104
+ """Identity; records that a value was consumed in a way we cannot model."""
105
+ return value
106
+
107
+
108
+ def fstring(*parts: Any) -> str:
109
+ """Rebuild an f-string from literal parts and ``(value, conversion, spec)`` triples."""
110
+ out: list[str] = []
111
+ for part in parts:
112
+ if isinstance(part, str):
113
+ out.append(part)
114
+ continue
115
+ value, conversion, spec = part
116
+ if conversion == "r":
117
+ value = repr(value)
118
+ elif conversion == "s":
119
+ value = str(value)
120
+ elif conversion == "a":
121
+ value = ascii(value)
122
+ out.append(format(value, spec))
123
+ return "".join(out)
124
+
125
+
126
+ BINARY_OPS: dict[str, tuple[Callable[[Any, Any], Any], str]] = {
127
+ "add": (operator.add, "+"),
128
+ "sub": (operator.sub, "-"),
129
+ "mul": (operator.mul, "*"),
130
+ "truediv": (operator.truediv, "/"),
131
+ "floordiv": (operator.floordiv, "//"),
132
+ "mod": (operator.mod, "%"),
133
+ "pow": (pow, "**"),
134
+ "lshift": (operator.lshift, "<<"),
135
+ "rshift": (operator.rshift, ">>"),
136
+ "and": (operator.and_, "&"),
137
+ "or": (operator.or_, "|"),
138
+ "xor": (operator.xor, "^"),
139
+ }
140
+
141
+ COMPARISON_OPS: dict[str, tuple[Callable[[Any, Any], Any], str]] = {
142
+ "lt": (operator.lt, "<"),
143
+ "le": (operator.le, "<="),
144
+ "gt": (operator.gt, ">"),
145
+ "ge": (operator.ge, ">="),
146
+ "eq": (operator.eq, "=="),
147
+ "ne": (operator.ne, "!="),
148
+ }
149
+
150
+ UNARY_OPS: dict[str, tuple[Callable[[Any], Any], str]] = {
151
+ "neg": (operator.neg, "-"),
152
+ "pos": (operator.pos, "+"),
153
+ "invert": (operator.invert, "~"),
154
+ }
155
+
156
+ # Builtins applied as functions: name -> implementation.
157
+ FUNCTION_OPS: dict[str, Callable[..., Any]] = {
158
+ "abs": abs,
159
+ "round": round,
160
+ "divmod": divmod,
161
+ "floor": math.floor,
162
+ "ceil": math.ceil,
163
+ "trunc": math.trunc,
164
+ "bool": bool,
165
+ "int": int,
166
+ "float": float,
167
+ "complex": complex,
168
+ "str": str,
169
+ "decimal": Decimal,
170
+ "format": format,
171
+ "hash": hash,
172
+ "len": len,
173
+ "index": operator.index,
174
+ "shape": shape,
175
+ "observe": observe,
176
+ "fstring": fstring,
177
+ "min": min,
178
+ "max": max,
179
+ "sum": sum,
180
+ "range": range,
181
+ "repr": repr,
182
+ "ascii": ascii,
183
+ "sorted": sorted,
184
+ "bisect.bisect_left": bisect.bisect_left,
185
+ "bisect.bisect_right": bisect.bisect_right,
186
+ "heapq.nsmallest": heapq.nsmallest,
187
+ "heapq.nlargest": heapq.nlargest,
188
+ }
189
+
190
+ STATISTICS_FUNCTIONS: tuple[str, ...] = (
191
+ "fmean",
192
+ "geometric_mean",
193
+ "harmonic_mean",
194
+ "mean",
195
+ "median",
196
+ "median_high",
197
+ "median_low",
198
+ "mode",
199
+ "pstdev",
200
+ "pvariance",
201
+ "stdev",
202
+ "variance",
203
+ )
204
+ FUNCTION_OPS.update(
205
+ {f"statistics.{name}": getattr(statistics, name) for name in STATISTICS_FUNCTIONS}
206
+ )
207
+
208
+ MATH_FUNCTIONS: tuple[str, ...] = (
209
+ "copysign",
210
+ "exp",
211
+ "fabs",
212
+ "fsum",
213
+ "hypot",
214
+ "isclose",
215
+ "isfinite",
216
+ "isinf",
217
+ "isnan",
218
+ "log",
219
+ "log10",
220
+ "log2",
221
+ "pow",
222
+ "prod",
223
+ "sqrt",
224
+ )
225
+ FUNCTION_OPS.update({f"math.{name}": getattr(math, name) for name in MATH_FUNCTIONS})
226
+
227
+ STR_METHODS: tuple[str, ...] = (
228
+ # str -> str
229
+ "capitalize",
230
+ "casefold",
231
+ "center",
232
+ "expandtabs",
233
+ "format",
234
+ "format_map",
235
+ "join",
236
+ "ljust",
237
+ "lower",
238
+ "lstrip",
239
+ "removeprefix",
240
+ "removesuffix",
241
+ "replace",
242
+ "rjust",
243
+ "rstrip",
244
+ "strip",
245
+ "swapcase",
246
+ "title",
247
+ "translate",
248
+ "upper",
249
+ "zfill",
250
+ # str -> containers / ints / bools
251
+ "count",
252
+ "endswith",
253
+ "find",
254
+ "index",
255
+ "isalnum",
256
+ "isalpha",
257
+ "isascii",
258
+ "isdecimal",
259
+ "isdigit",
260
+ "isidentifier",
261
+ "islower",
262
+ "isnumeric",
263
+ "isprintable",
264
+ "isspace",
265
+ "istitle",
266
+ "isupper",
267
+ "partition",
268
+ "rfind",
269
+ "rindex",
270
+ "rpartition",
271
+ "rsplit",
272
+ "split",
273
+ "splitlines",
274
+ "startswith",
275
+ "encode",
276
+ )
277
+
278
+ DECIMAL_METHODS: tuple[str, ...] = (
279
+ "adjusted",
280
+ "as_integer_ratio",
281
+ "as_tuple",
282
+ "compare",
283
+ "copy_abs",
284
+ "copy_negate",
285
+ "copy_sign",
286
+ "exp",
287
+ "fma",
288
+ "is_finite",
289
+ "is_infinite",
290
+ "is_nan",
291
+ "is_signed",
292
+ "is_zero",
293
+ "ln",
294
+ "log10",
295
+ "max",
296
+ "min",
297
+ "normalize",
298
+ "quantize",
299
+ "scaleb",
300
+ "sqrt",
301
+ "to_integral",
302
+ "to_integral_exact",
303
+ "to_integral_value",
304
+ )
305
+
306
+ FLOAT_METHODS: tuple[str, ...] = ("as_integer_ratio", "is_integer")
307
+ INT_METHODS: tuple[str, ...] = ("as_integer_ratio", "bit_count", "bit_length")
308
+
309
+
310
+ def _register_builtins() -> None:
311
+ for name, (fn, symbol) in BINARY_OPS.items():
312
+ register_op(name, fn, symbol=symbol, style="infix")
313
+ for name, (fn, symbol) in COMPARISON_OPS.items():
314
+ register_op(name, fn, symbol=symbol, style="infix")
315
+ for name, (unary, symbol) in UNARY_OPS.items():
316
+ register_op(name, unary, symbol=symbol, style="prefix")
317
+ for name, fn in FUNCTION_OPS.items():
318
+ register_op(name, fn)
319
+ register_op("decimal", Decimal, symbol="Decimal")
320
+ register_op("getitem", operator.getitem, style="getitem")
321
+ register_op("contains", operator.contains, symbol="in", style="contains")
322
+ for prefix, owner, methods in (
323
+ ("str", str, STR_METHODS),
324
+ ("decimal", Decimal, DECIMAL_METHODS),
325
+ ("float", float, FLOAT_METHODS),
326
+ ("int", int, INT_METHODS),
327
+ ):
328
+ for method in methods:
329
+ register_op(f"{prefix}.{method}", getattr(owner, method), symbol=method, style="method")
330
+
331
+
332
+ _register_builtins()
redroot/paths.py ADDED
@@ -0,0 +1,162 @@
1
+ """Semantic keys for values inside nested data.
2
+
3
+ A key names one value by its position in a nested structure of dicts and
4
+ lists, prefixed by a *root* that says which structure it belongs to::
5
+
6
+ ext:bank.line_items[3].amount
7
+ out:form1["line 12"]
8
+ out:tables[0][2]
9
+
10
+ Dict keys that are ASCII identifiers use dot notation. Other keys and list indexes use
11
+ brackets holding a JSON literal, so ``[3]`` is the integer 3 (a list index or
12
+ an ``int`` dict key) and ``["3"]`` is the string ``"3"``. Keys round-trip:
13
+ ``parse_key(format_key(root, path)) == (root, path)``.
14
+ """
15
+
16
+ from __future__ import annotations
17
+
18
+ import copy
19
+ import json
20
+ import operator
21
+ import re
22
+ from collections.abc import Mapping, Sequence
23
+ from typing import Any
24
+
25
+ from redroot._core import rebuild_tuple, unwrap
26
+
27
+ __all__ = ["apply_edits", "child_key", "format_key", "get_path", "parse_key", "set_path"]
28
+
29
+ Segment = str | int | float | bool | None
30
+ Path = tuple[Segment, ...]
31
+
32
+ _ROOT = re.compile(r"[A-Za-z_][A-Za-z0-9_\-]*")
33
+ _NAME = re.compile(r"[A-Za-z_][A-Za-z0-9_]*")
34
+
35
+
36
+ def _check_root(root: str) -> str:
37
+ if not _ROOT.fullmatch(root):
38
+ raise ValueError(f"invalid root {root!r}: expected a name such as 'ext' or 'out'")
39
+ return root
40
+
41
+
42
+ def _segment(segment: Any, first: bool) -> str:
43
+ segment = unwrap(segment)
44
+ if isinstance(segment, str) and _NAME.fullmatch(segment):
45
+ return segment if first else f".{segment}"
46
+ if segment is None or isinstance(segment, (str, int, float, bool)):
47
+ return f"[{json.dumps(segment)}]"
48
+ # Keys that are not JSON scalars (tuples, enums...) are named by their repr.
49
+ return f"[{json.dumps(repr(segment))}]"
50
+
51
+
52
+ def child_key(parent: str, segment: Any) -> str:
53
+ """Return the key of ``segment`` inside the value named ``parent``.
54
+
55
+ >>> child_key("ext", "bank"), child_key("ext:bank", 0)
56
+ ('ext:bank', 'ext:bank[0]')
57
+ """
58
+ if ":" not in parent: # a bare root
59
+ return f"{parent}:{_segment(segment, first=True)}"
60
+ return parent + _segment(segment, first=False)
61
+
62
+
63
+ def format_key(root: str, path: Sequence[Any] = ()) -> str:
64
+ """Build the key of the value at ``path`` under ``root``.
65
+
66
+ >>> format_key("ext", ["bank", "line_items", 3, "amount"])
67
+ 'ext:bank.line_items[3].amount'
68
+ """
69
+ _check_root(root)
70
+ if not path:
71
+ return root
72
+ return f"{root}:" + "".join(_segment(s, first=i == 0) for i, s in enumerate(path))
73
+
74
+
75
+ _DECODER = json.JSONDecoder()
76
+
77
+
78
+ def parse_key(key: str) -> tuple[str, Path]:
79
+ """Split a key into its root and path.
80
+
81
+ >>> parse_key('out:form1["line 12"]')
82
+ ('out', ('form1', 'line 12'))
83
+ """
84
+ root, sep, rest = key.partition(":")
85
+ _check_root(root)
86
+ if not sep:
87
+ return root, ()
88
+ path: list[Segment] = []
89
+ i = 0
90
+ while i < len(rest):
91
+ if rest[i] == "[":
92
+ try:
93
+ value, end = _DECODER.raw_decode(rest, i + 1)
94
+ except json.JSONDecodeError as exc:
95
+ raise ValueError(f"invalid key {key!r}: bad bracket segment") from exc
96
+ if end >= len(rest) or rest[end] != "]":
97
+ raise ValueError(f"invalid key {key!r}: unclosed bracket")
98
+ if isinstance(value, (list, dict)):
99
+ raise ValueError(f"invalid key {key!r}: brackets must hold a scalar")
100
+ path.append(value)
101
+ i = end + 1
102
+ continue
103
+ if rest[i] == ".":
104
+ if not path:
105
+ raise ValueError(f"invalid key {key!r}: unexpected '.'")
106
+ i += 1
107
+ elif path:
108
+ raise ValueError(f"invalid key {key!r}: expected '.' or '['")
109
+ match = _NAME.match(rest, i)
110
+ if match is None:
111
+ raise ValueError(f"invalid key {key!r}: expected a name at position {i}")
112
+ path.append(match.group())
113
+ i = match.end()
114
+ if not path:
115
+ raise ValueError(f"invalid key {key!r}: empty path")
116
+ return root, tuple(path)
117
+
118
+
119
+ def get_path(data: Any, path: Sequence[Segment]) -> Any:
120
+ """Return the value at ``path`` inside ``data``."""
121
+ for segment in path:
122
+ data = data[segment]
123
+ return data
124
+
125
+
126
+ def set_path(data: Any, path: Sequence[Segment], value: Any) -> None:
127
+ """Replace the value at ``path`` inside ``data`` (in place)."""
128
+ if not path:
129
+ raise ValueError("cannot replace the root value in place")
130
+ parent = get_path(data, path[:-1])
131
+ if isinstance(parent, tuple):
132
+ raise TypeError(f"cannot assign into a tuple at {list(path[:-1])}")
133
+ parent[path[-1]] = value
134
+
135
+
136
+ def apply_edits(data: Any, edits: Mapping[str, Any], *, root: str) -> Any:
137
+ """Return a deep copy of ``data`` with each edit applied.
138
+
139
+ Edits are keyed like :meth:`redroot.Trace.propagate` edits. This is how
140
+ to build the inputs for re-executing a workflow after a review.
141
+ """
142
+ edited = copy.deepcopy(data)
143
+ for key, value in edits.items():
144
+ key_root, path = parse_key(key)
145
+ if key_root != root:
146
+ raise ValueError(f"edit {key!r} does not belong to root {root!r}")
147
+ edited = _replace(edited, path, value)
148
+ return edited
149
+
150
+
151
+ def _replace(data: Any, path: Sequence[Segment], value: Any) -> Any:
152
+ """Return ``data`` with the value at ``path`` replaced, rebuilding tuples."""
153
+ if not path:
154
+ return value
155
+ head = path[0]
156
+ child = _replace(data[head], path[1:], value)
157
+ if isinstance(data, tuple):
158
+ items = list(data)
159
+ items[operator.index(head)] = child # type: ignore[arg-type]
160
+ return rebuild_tuple(data, items)
161
+ data[head] = child
162
+ return data