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/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)
|