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/cli.py ADDED
@@ -0,0 +1,98 @@
1
+ """Command-line tools for traces saved with ``Trace.to_json()``."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ import json
7
+ import sys
8
+ from collections.abc import Sequence
9
+
10
+ from redroot.trace import Trace
11
+
12
+
13
+ def _load(path: str) -> Trace:
14
+ try:
15
+ with open(path, encoding="utf-8") as f:
16
+ return Trace.from_dict(json.load(f))
17
+ except FileNotFoundError:
18
+ raise SystemExit(f"error: file '{path}' not found") from None
19
+ except json.JSONDecodeError:
20
+ raise SystemExit(f"error: '{path}' is not valid JSON") from None
21
+ except (KeyError, TypeError, ValueError) as exc:
22
+ raise SystemExit(f"error: '{path}' is not a RedRoot trace: {exc}") from None
23
+
24
+
25
+ def _visualize(args: argparse.Namespace) -> int:
26
+ from redroot.visualizer.web.server import run_server
27
+
28
+ run_server(_load(args.trace), port=args.port, open_browser=not args.no_browser)
29
+ return 0
30
+
31
+
32
+ def _explain(args: argparse.Namespace) -> int:
33
+ trace = _load(args.trace)
34
+ keys = args.keys or list(trace.outputs)
35
+ for key in keys:
36
+ out = trace.outputs.get(key)
37
+ if out is not None and out.node is None:
38
+ print(f"{key} = {out.value!r} (not linked to any traced input)")
39
+ continue
40
+ try:
41
+ node = trace.node(key)
42
+ except KeyError:
43
+ print(f"error: no input or output named {key!r}", file=sys.stderr)
44
+ return 1
45
+ print(f"{key} = {node.value!r}")
46
+ print(f" computed as: {trace.explain(key)}")
47
+ sources = ", ".join(s.key or f"#{s.id}" for s in trace.sources(key))
48
+ print(f" from inputs: {sources or '-'}")
49
+ return 0
50
+
51
+
52
+ def _coverage(args: argparse.Namespace) -> int:
53
+ from redroot.validation import check_identity, coverage
54
+
55
+ trace = _load(args.trace)
56
+ print(coverage(trace))
57
+ report = check_identity(trace)
58
+ status = "ok" if report.ok else f"{len(report.mismatches)} MISMATCHES"
59
+ print(
60
+ f"identity .......... {status} ({report.checked} checked, {report.skipped} not replayable)"
61
+ )
62
+ return 0 if report.ok else 1
63
+
64
+
65
+ def build_parser() -> argparse.ArgumentParser:
66
+ """Build the ``redroot`` argument parser."""
67
+ parser = argparse.ArgumentParser(prog="redroot", description="Inspect RedRoot traces.")
68
+ sub = parser.add_subparsers(dest="command")
69
+
70
+ viz = sub.add_parser("visualize", help="open a trace in the web viewer")
71
+ viz.add_argument("trace", help="trace JSON file")
72
+ viz.add_argument("--port", type=int, default=8050, help="port to serve on (default 8050)")
73
+ viz.add_argument("--no-browser", action="store_true", help="do not open a browser")
74
+ viz.set_defaults(handler=_visualize)
75
+
76
+ explain = sub.add_parser("explain", help="show how outputs were computed")
77
+ explain.add_argument("trace", help="trace JSON file")
78
+ explain.add_argument("keys", nargs="*", help="input or output keys (default: every output)")
79
+ explain.set_defaults(handler=_explain)
80
+
81
+ cov = sub.add_parser("coverage", help="report coverage and check the trace re-evaluates")
82
+ cov.add_argument("trace", help="trace JSON file")
83
+ cov.set_defaults(handler=_coverage)
84
+ return parser
85
+
86
+
87
+ def main(argv: Sequence[str] | None = None) -> int:
88
+ """Run the CLI and return its exit code."""
89
+ parser = build_parser()
90
+ args = parser.parse_args(argv)
91
+ if args.command is None:
92
+ parser.print_help()
93
+ return 0
94
+ return int(args.handler(args))
95
+
96
+
97
+ if __name__ == "__main__":
98
+ sys.exit(main())
redroot/functions.py ADDED
@@ -0,0 +1,240 @@
1
+ """Recording whole function calls as single nodes.
2
+
3
+ Some steps are better seen as one opaque operation than as the arithmetic
4
+ inside them: a tax-table lookup, a zip-to-county resolver, an LLM call. A
5
+ :func:`traced` function runs on plain values with tracing suspended, and its
6
+ result is recorded as one node whose operands are the call's arguments.
7
+
8
+ Deterministic functions are re-called during propagation. Non-deterministic
9
+ ones (``deterministic=False``, e.g. web lookups or LLM calls) keep their
10
+ recorded result while their inputs are unchanged; if an input changes, the
11
+ values derived from them are reported stale.
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ import functools
17
+ from collections.abc import Callable, Iterable
18
+ from contextvars import ContextVar
19
+ from decimal import Decimal
20
+ from typing import Any, TypeVar, cast, overload
21
+
22
+ from redroot._core import (
23
+ NOT_REPLAYABLE,
24
+ Node,
25
+ Observer,
26
+ Traced,
27
+ _active,
28
+ active_trace,
29
+ decimal_context_id,
30
+ encode,
31
+ iter_nodes,
32
+ record_raise,
33
+ record_result,
34
+ unwrap,
35
+ unwrap_deep,
36
+ )
37
+ from redroot.ops import get_op, register_op
38
+
39
+ __all__ = ["annotate", "derive", "traced", "traced_llm"]
40
+
41
+ F = TypeVar("F", bound=Callable[..., Any])
42
+
43
+ _call_meta: ContextVar[dict[str, Any] | None] = ContextVar("redroot_call_meta", default=None)
44
+
45
+
46
+ @overload
47
+ def traced(fn: F, /) -> F: ...
48
+
49
+
50
+ @overload
51
+ def traced(
52
+ *,
53
+ name: str | None = None,
54
+ deterministic: bool = True,
55
+ meta: dict[str, Any] | None = None,
56
+ ) -> Callable[[F], F]: ...
57
+
58
+
59
+ def traced(
60
+ fn: F | None = None,
61
+ /,
62
+ *,
63
+ name: str | None = None,
64
+ deterministic: bool = True,
65
+ meta: dict[str, Any] | None = None,
66
+ ) -> F | Callable[[F], F]:
67
+ """Record each call of the decorated function as one node.
68
+
69
+ Args:
70
+ fn: The function, when used as a bare ``@traced`` decorator.
71
+ name: Operation name stored on nodes; defaults to
72
+ ``"<module>.<qualname>"``. To re-evaluate a deserialized trace,
73
+ the function must be importable and registered under this name.
74
+ deterministic: Whether calling again with equal arguments returns an
75
+ equal result. Only deterministic functions are re-evaluated.
76
+ meta: Metadata stored on every node this function records.
77
+
78
+ Usage::
79
+
80
+ @redroot.traced
81
+ def national_standard(household_size: int) -> Decimal: ...
82
+
83
+ @redroot.traced(name="geo.county", deterministic=False)
84
+ def county_for_zip(zip_code: str) -> str: ...
85
+ """
86
+
87
+ def decorate(func: F) -> F:
88
+ module = getattr(func, "__module__", None)
89
+ op_name = name or (f"{module}.{func.__qualname__}" if module else func.__qualname__)
90
+ register_op(op_name, func, replayable=deterministic)
91
+ # Nodes keep the exact function that computed them, so helpers that
92
+ # share a name (e.g. in two generated workflows) never get mixed up.
93
+ node_fn = func if deterministic else NOT_REPLAYABLE
94
+
95
+ @functools.wraps(func)
96
+ def wrapper(*args: Any, **kwargs: Any) -> Any:
97
+ trace = _active.get()
98
+ if trace is None:
99
+ return func(*args, **kwargs)
100
+ if isinstance(trace, Observer): # inside another traced call
101
+ trace.touch(args, kwargs)
102
+ return func(*args, **kwargs)
103
+ refs, _, linked = encode(args, trace)
104
+ kw_refs, _, kw_linked = encode(kwargs, trace)
105
+ # The function is one opaque step: it runs on the original
106
+ # arguments (so mutations behave as usual) with recording
107
+ # suspended, while an observer notes every traced value it reads.
108
+ observer = Observer(trace)
109
+ trace_token = _active.set(observer)
110
+ meta_token = _call_meta.set({})
111
+ try:
112
+ result = func(*args, **kwargs)
113
+ except Exception as exc:
114
+ if linked or kw_linked or observer.touched:
115
+ record_raise(trace, op_name, tuple(refs), kw_refs or None, exc, node_fn)
116
+ raise
117
+ finally:
118
+ annotations = _call_meta.get()
119
+ _call_meta.reset(meta_token)
120
+ _active.reset(trace_token)
121
+
122
+ operand_ids = {node.id for node in iter_nodes((refs, kw_refs))}
123
+ hidden = {**_nested_nodes((args, kwargs), trace), **observer.touched}
124
+ extra = [node for node_id, node in sorted(hidden.items()) if node_id not in operand_ids]
125
+ if not (linked or kw_linked or extra):
126
+ return result
127
+ fn: Callable[..., Any] = node_fn
128
+ if extra:
129
+ # Inputs reached the call where they cannot be replayed from
130
+ # (closures, sets, dict views...): depend on them, never replay.
131
+ fn = NOT_REPLAYABLE
132
+ kw_refs = {**kw_refs, "__depends_on__": extra}
133
+ node_meta = {**(meta or {}), **(annotations or {})} or None
134
+ ctx = decimal_context_id(trace) if _involves_decimal(refs, result) else None
135
+ return record_result(
136
+ trace,
137
+ op_name,
138
+ tuple(_sanitize(r) for r in refs),
139
+ kw_refs or None,
140
+ result,
141
+ node_meta,
142
+ ctx,
143
+ fn,
144
+ )
145
+
146
+ wrapper.__redroot_op__ = op_name # type: ignore[attr-defined]
147
+ return wrapper # type: ignore[return-value]
148
+
149
+ if fn is not None:
150
+ return decorate(fn)
151
+ return decorate
152
+
153
+
154
+ def traced_llm(model: str, task: str | None = None, **meta: Any) -> Callable[[F], F]:
155
+ """Record calls to a function that queries an LLM.
156
+
157
+ LLM calls are not deterministic, so they are never re-run: if their inputs
158
+ change, values derived from them are reported stale. The model name and
159
+ any ``meta`` are stored on the node; call :func:`annotate` inside the
160
+ function to add per-call data such as token counts.
161
+
162
+ Args:
163
+ model: Model identifier, stored as ``meta["model"]``.
164
+ task: Operation name; defaults to the function's qualified name.
165
+ **meta: More metadata stored on every node.
166
+ """
167
+ return traced(name=task, deterministic=False, meta={"model": model, **meta})
168
+
169
+
170
+ _VIEW_TYPES: tuple[type, ...] = (
171
+ set,
172
+ frozenset,
173
+ type({}.keys()),
174
+ type({}.values()),
175
+ type({}.items()),
176
+ )
177
+
178
+
179
+ def _nested_nodes(obj: Any, trace: Any, found: dict[int, Node] | None = None) -> dict[int, Node]:
180
+ """Nodes of ``trace`` inside ``obj``, including inside sets and dict views."""
181
+ found = {} if found is None else found
182
+ if isinstance(obj, Traced):
183
+ node = obj._rt_node
184
+ if node.trace is trace:
185
+ found[node.id] = node
186
+ elif type(obj) is dict:
187
+ for item in obj.values():
188
+ _nested_nodes(item, trace, found)
189
+ elif type(obj) is list or isinstance(obj, (tuple, *_VIEW_TYPES)):
190
+ for item in cast("Iterable[Any]", obj):
191
+ _nested_nodes(item, trace, found)
192
+ return found
193
+
194
+
195
+ def _sanitize(ref: Any) -> Any:
196
+ """Replace a set or dict view argument by a plain snapshot for storage."""
197
+ if isinstance(ref, _VIEW_TYPES):
198
+ items = [unwrap_deep(item) for item in cast("Iterable[Any]", ref)]
199
+ return type(ref)(items) if isinstance(ref, (set, frozenset)) else items
200
+ return ref
201
+
202
+
203
+ def _involves_decimal(refs: Any, result: Any) -> bool:
204
+ if type(unwrap(result)) is Decimal:
205
+ return True
206
+ return any(type(node.value) is Decimal for node in iter_nodes(refs))
207
+
208
+
209
+ def annotate(**meta: Any) -> None:
210
+ """Attach metadata to the node recorded for the current traced call.
211
+
212
+ Does nothing outside a recording call (e.g. during re-evaluation).
213
+ """
214
+ current = _call_meta.get()
215
+ if current is not None:
216
+ current.update(meta)
217
+
218
+
219
+ def derive(
220
+ value: Any, inputs: Iterable[Any], op: str, *, meta: dict[str, Any] | None = None
221
+ ) -> Any:
222
+ """Record that ``value`` was computed from ``inputs`` outside RedRoot's view.
223
+
224
+ Use this for steps that cannot be decorated, e.g. a result returned by an
225
+ external service. Unless ``op`` names a registered replayable operation,
226
+ the node cannot be re-evaluated and dependants become stale when
227
+ ``inputs`` change.
228
+
229
+ Returns:
230
+ ``value`` in traced form (or unchanged outside an active trace).
231
+ """
232
+ trace = active_trace()
233
+ if trace is None:
234
+ return value
235
+ refs, _, linked = encode(tuple(inputs), trace)
236
+ if not linked:
237
+ return value
238
+ spec = get_op(op)
239
+ fn = None if spec is not None and spec.replayable else NOT_REPLAYABLE
240
+ return record_result(trace, op, refs, None, unwrap_deep(value), meta, fn=fn)