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/instrument.py ADDED
@@ -0,0 +1,614 @@
1
+ """Source instrumentation: trace what operator overloading cannot see.
2
+
3
+ Traced values are subclasses of builtins, so most operations reach their
4
+ overridden methods. A few do not, because CPython handles them in C without
5
+ consulting the subclass:
6
+
7
+ * a plain ``float``/``Decimal`` on the left of a traced ``int``
8
+ (``0.8 * income``) and the matching comparisons (``0.5 < size``);
9
+ * conversions that parse or read the raw value: ``float(text)``,
10
+ ``Decimal(text)``, ``int(text)``, ``math.sqrt(x)``;
11
+ * methods of *plain* strings given traced arguments: ``", ".join(parts)``,
12
+ ``"{:,.2f}".format(x)``, and f-strings;
13
+ * indexing a plain list with a traced ``int`` (``table[size]``) and
14
+ ``range(n)``.
15
+
16
+ Values computed that way silently lose their lineage. For code you execute
17
+ from source (such as generated workflow code), :func:`exec_source` rewrites
18
+ the syntax tree so these operations are routed through RedRoot first, the
19
+ way pytest rewrites ``assert`` statements. Instrumented code behaves exactly
20
+ like the original; with no active trace it only adds a function call per
21
+ operation.
22
+
23
+ Usage::
24
+
25
+ namespace = redroot.instrument.exec_source(generated_code)
26
+ trace, result = redroot.run(namespace["main"], extracted, input_root="ext")
27
+ """
28
+
29
+ from __future__ import annotations
30
+
31
+ import ast
32
+ import bisect
33
+ import collections
34
+ import heapq
35
+ import itertools
36
+ import math
37
+ import operator
38
+ import statistics
39
+ from collections.abc import Callable
40
+ from decimal import Decimal
41
+ from types import CodeType
42
+ from typing import Any
43
+
44
+ from redroot import _core
45
+ from redroot._core import Traced, observe_value
46
+ from redroot.ops import (
47
+ BINARY_OPS,
48
+ COMPARISON_OPS,
49
+ DECIMAL_METHODS,
50
+ MATH_FUNCTIONS,
51
+ STATISTICS_FUNCTIONS,
52
+ STR_METHODS,
53
+ fstring,
54
+ )
55
+
56
+ __all__ = ["RUNTIME_NAME", "compile_source", "exec_source", "instrument"]
57
+
58
+ RUNTIME_NAME = "__redroot__"
59
+ """Global name under which instrumented code finds the runtime helpers."""
60
+
61
+ # --------------------------------------------------------------------------
62
+ # Runtime helpers called by instrumented code
63
+ # --------------------------------------------------------------------------
64
+
65
+ # Pure callables recorded as one node when given traced arguments: id -> op.
66
+ _CALLS: dict[int, str] = {
67
+ id(fn): name
68
+ for fn, name in [
69
+ (int, "int"),
70
+ (float, "float"),
71
+ (str, "str"),
72
+ (bool, "bool"),
73
+ (Decimal, "decimal"),
74
+ (round, "round"),
75
+ (abs, "abs"),
76
+ (min, "min"),
77
+ (max, "max"),
78
+ (sum, "sum"),
79
+ (len, "len"),
80
+ (divmod, "divmod"),
81
+ (pow, "pow"),
82
+ (format, "format"),
83
+ (range, "range"),
84
+ (repr, "repr"),
85
+ (ascii, "ascii"),
86
+ (sorted, "sorted"),
87
+ (bisect.bisect_left, "bisect.bisect_left"),
88
+ (bisect.bisect_right, "bisect.bisect_right"),
89
+ (heapq.nsmallest, "heapq.nsmallest"),
90
+ (heapq.nlargest, "heapq.nlargest"),
91
+ (math.floor, "floor"),
92
+ (math.ceil, "ceil"),
93
+ (math.trunc, "trunc"),
94
+ *((getattr(math, name), f"math.{name}") for name in MATH_FUNCTIONS),
95
+ *((getattr(statistics, name), f"statistics.{name}") for name in STATISTICS_FUNCTIONS),
96
+ ]
97
+ }
98
+ # Position of the iterable argument to materialize so it can be inspected
99
+ # (and is not consumed twice); min/max only take one when called with one.
100
+ _ITERABLE_ARG: dict[int, int] = {
101
+ id(sum): 0,
102
+ id(sorted): 0,
103
+ id(math.fsum): 0,
104
+ id(math.prod): 0,
105
+ id(heapq.nsmallest): 1,
106
+ id(heapq.nlargest): 1,
107
+ **{id(getattr(statistics, name)): 0 for name in STATISTICS_FUNCTIONS},
108
+ }
109
+ _SINGLE_ITERABLE = {id(min), id(max)}
110
+ _METHODS: dict[type, tuple[str, frozenset[str]]] = {
111
+ str: ("str", frozenset(STR_METHODS)),
112
+ Decimal: ("decimal", frozenset(DECIMAL_METHODS)),
113
+ }
114
+ # Containers whose methods only store, fetch or compare values through the
115
+ # values' own (traced) methods.
116
+ _CONTAINERS: tuple[type, ...] = (
117
+ list,
118
+ dict,
119
+ set,
120
+ frozenset,
121
+ tuple,
122
+ collections.deque,
123
+ collections.defaultdict,
124
+ collections.OrderedDict,
125
+ collections.Counter,
126
+ )
127
+ # Builtins that pass values through, or depend only on types, so a call that
128
+ # returns no traced value has not lost lineage.
129
+ _TRANSPARENT: set[int] = {
130
+ id(fn)
131
+ for fn in (
132
+ print,
133
+ isinstance,
134
+ issubclass,
135
+ callable,
136
+ hasattr,
137
+ getattr,
138
+ setattr,
139
+ delattr,
140
+ id,
141
+ iter,
142
+ next,
143
+ enumerate,
144
+ zip,
145
+ reversed,
146
+ map,
147
+ filter,
148
+ list,
149
+ tuple,
150
+ dict,
151
+ set,
152
+ frozenset,
153
+ any,
154
+ all,
155
+ vars,
156
+ dir,
157
+ object,
158
+ type,
159
+ slice,
160
+ )
161
+ }
162
+ _TRANSPARENT_MODULES = frozenset(
163
+ {"logging", "warnings", "copy", "itertools", "collections", "dataclasses"}
164
+ )
165
+ _INSTRUMENTED_FILES: set[str] = set()
166
+
167
+
168
+ def _has_traced(obj: Any, depth: int = 2) -> bool:
169
+ if isinstance(obj, Traced):
170
+ return True
171
+ if depth and type(obj) in (list, tuple):
172
+ return any(_has_traced(item, depth - 1) for item in obj)
173
+ if depth and type(obj) is dict:
174
+ return any(_has_traced(item, depth - 1) for item in obj.values())
175
+ return False
176
+
177
+
178
+ def _pure_callable(func: Any) -> str | None:
179
+ """The op name of a callable RedRoot can record as one node, if any."""
180
+ op = _CALLS.get(id(func))
181
+ if op is not None:
182
+ return op
183
+ owner = getattr(func, "__objclass__", None) # unbound method, e.g. str.lower
184
+ method = _METHODS.get(owner) if isinstance(owner, type) else None
185
+ if method is not None and getattr(func, "__name__", None) in method[1]:
186
+ return f"{method[0]}.{func.__name__}"
187
+ return None
188
+
189
+
190
+ def _wrap_callable(func: Any) -> Any:
191
+ """Route calls of a known pure callable through :func:`call`."""
192
+ if func is None or _pure_callable(func) is None:
193
+ return func
194
+
195
+ def routed(*args: Any, **kwargs: Any) -> Any:
196
+ return call(func, *args, **kwargs)
197
+
198
+ return routed
199
+
200
+
201
+ def _is_instrumented(func: Any) -> bool:
202
+ code = getattr(func, "__code__", None) or getattr(
203
+ getattr(func, "__func__", None), "__code__", None
204
+ )
205
+ return code is not None and code.co_filename in _INSTRUMENTED_FILES
206
+
207
+
208
+ def _is_transparent(func: Any) -> bool:
209
+ if id(func) in _TRANSPARENT or hasattr(func, "__redroot_op__") or _is_instrumented(func):
210
+ return True
211
+ if isinstance(getattr(func, "__self__", None), _CONTAINERS):
212
+ return True
213
+ module = getattr(func, "__module__", None) or ""
214
+ return module.partition(".")[0] in _TRANSPARENT_MODULES
215
+
216
+
217
+ def _keeps_lineage(result: Any, trace: Any, first_new_node: int, depth: int = 2) -> bool:
218
+ """Whether ``result`` carries traced values that link back to the trace's inputs."""
219
+ if isinstance(result, Traced):
220
+ node = result._rt_node
221
+ # A value rebuilt from a plain number (e.g. type(x)(...) in a library)
222
+ # becomes a fresh anonymous input: that is lost lineage, not kept.
223
+ fresh_input = node.kind == "leaf" and node.key is None and node.id >= first_new_node
224
+ return node.trace is trace and not fresh_input
225
+ if not depth:
226
+ return False
227
+ if type(result) in (list, tuple, set, frozenset) or isinstance(result, tuple):
228
+ return any(_keeps_lineage(item, trace, first_new_node, depth - 1) for item in result)
229
+ if type(result) is dict:
230
+ return any(
231
+ _keeps_lineage(item, trace, first_new_node, depth - 1) for item in result.values()
232
+ )
233
+ attributes = getattr(result, "__dict__", None)
234
+ if isinstance(attributes, dict):
235
+ return any(
236
+ _keeps_lineage(item, trace, first_new_node, depth - 1) for item in attributes.values()
237
+ )
238
+ return False
239
+
240
+
241
+ def _sort_in_place(items: list[Any], key: Any = None, reverse: bool = False) -> None:
242
+ if key is None and _has_traced(items, depth=1):
243
+ items[:] = _core.call(
244
+ "sorted", sorted, (list(items),), {"reverse": reverse} if reverse else None
245
+ )
246
+ else:
247
+ items.sort(key=_wrap_callable(key), reverse=reverse)
248
+
249
+
250
+ def _opaque_call(func: Any, args: tuple[Any, ...], kwargs: dict[str, Any]) -> Any:
251
+ """Call code RedRoot cannot see into; guard its traced inputs if lineage is lost.
252
+
253
+ The guard cannot be re-evaluated, so any change to those inputs means
254
+ re-execution: sound, if conservative.
255
+ """
256
+ trace = _core.active_trace()
257
+ if trace is None or _is_transparent(func):
258
+ return func(*args, **kwargs)
259
+ refs, _, linked = _core.encode(args, trace)
260
+ kw_refs, _, kw_linked = _core.encode(kwargs, trace)
261
+ if not (linked or kw_linked):
262
+ return func(*args, **kwargs)
263
+ name = getattr(func, "__qualname__", None) or type(func).__qualname__
264
+ op = f"call:{getattr(func, '__module__', None) or ''}.{name}".replace(":.", ":")
265
+ first_new_node = len(trace.nodes)
266
+ try:
267
+ result = func(*args, **kwargs)
268
+ except Exception as exc:
269
+ _core.record_raise(trace, op, tuple(refs), kw_refs or None, exc, _core.NOT_REPLAYABLE)
270
+ raise
271
+ if not _keeps_lineage(result, trace, first_new_node):
272
+ _core.add_node(
273
+ trace,
274
+ op,
275
+ tuple(refs),
276
+ kw_refs or None,
277
+ _core.unwrap_deep(result),
278
+ _core.GUARD,
279
+ fn=_core.NOT_REPLAYABLE,
280
+ )
281
+ return result
282
+
283
+
284
+ def call(func: Any, /, *args: Any, **kwargs: Any) -> Any:
285
+ """Call ``func`` from instrumented code, recording what overloading cannot see."""
286
+ func_id = id(func)
287
+ op = _CALLS.get(func_id)
288
+ if op is not None:
289
+ if kwargs.get("key") is not None:
290
+ # A custom key may close over traced values: run natively so the
291
+ # comparisons of its results are recorded as guards.
292
+ kwargs["key"] = _wrap_callable(kwargs["key"])
293
+ return func(*args, **kwargs)
294
+ position = _ITERABLE_ARG.get(
295
+ func_id, 0 if func_id in _SINGLE_ITERABLE and len(args) == 1 else -1
296
+ )
297
+ if 0 <= position < len(args) and type(args[position]) not in (list, tuple):
298
+ args = (*args[:position], list(args[position]), *args[position + 1 :])
299
+ if _has_traced(args) or _has_traced(kwargs):
300
+ return _core.call(op, func, args, kwargs or None)
301
+ return func(*args, **kwargs)
302
+ if func is map or func is filter:
303
+ if args:
304
+ args = (_wrap_callable(args[0]), *args[1:])
305
+ return func(*args, **kwargs)
306
+ if func is type and len(args) == 1 and not kwargs and isinstance(args[0], Traced):
307
+ return args[0]._base # what type() returns for the plain value
308
+ receiver = getattr(func, "__self__", None)
309
+ method = _METHODS.get(type(receiver))
310
+ if method is not None and func.__name__ in method[1]:
311
+ if func.__name__ == "join" and args and type(args[0]) not in (list, tuple):
312
+ args = (list(args[0]),)
313
+ if _has_traced(args) or _has_traced(kwargs):
314
+ name = func.__name__
315
+ unbound = getattr(type(receiver), name)
316
+ return _core.call(f"{method[0]}.{name}", unbound, (receiver, *args), kwargs or None)
317
+ return func(*args, **kwargs)
318
+ if type(receiver) is list and getattr(func, "__name__", None) == "sort":
319
+ return _sort_in_place(receiver, *args, **kwargs)
320
+ op = _pure_callable(func) # unbound methods: str.upper(x)
321
+ if op is not None and (_has_traced(args) or _has_traced(kwargs)):
322
+ return _core.call(op, func, args, kwargs or None)
323
+ if not (_has_traced(args) or _has_traced(kwargs)):
324
+ return func(*args, **kwargs)
325
+ return _opaque_call(func, args, kwargs)
326
+
327
+
328
+ def binop(op: str, left: Any, right: Any) -> Any:
329
+ """``left <op> right``, recorded if either operand is traced."""
330
+ fn = BINARY_OPS[op][0]
331
+ if isinstance(left, Traced) or isinstance(right, Traced):
332
+ return _core.binary(op, fn, left, right)
333
+ if op == "mod" and type(left) is str and _has_traced(right):
334
+ return _core.call(op, fn, (left, right)) # "%s" % (traced, ...)
335
+ return fn(left, right)
336
+
337
+
338
+ _INPLACE: dict[str, Callable[[Any, Any], Any]] = {
339
+ "add": operator.iadd,
340
+ "sub": operator.isub,
341
+ "mul": operator.imul,
342
+ "truediv": operator.itruediv,
343
+ "floordiv": operator.ifloordiv,
344
+ "mod": operator.imod,
345
+ "pow": operator.ipow,
346
+ "lshift": operator.ilshift,
347
+ "rshift": operator.irshift,
348
+ "and": operator.iand,
349
+ "or": operator.ior,
350
+ "xor": operator.ixor,
351
+ }
352
+
353
+
354
+ def ibinop(op: str, left: Any, right: Any) -> Any:
355
+ """``left <op>= right``: in place for mutable ``left``, recorded for traced values."""
356
+ if isinstance(left, Traced) or (
357
+ isinstance(right, Traced) and not hasattr(type(left), f"__i{op}__")
358
+ ):
359
+ return _core.binary(op, BINARY_OPS[op][0], left, right)
360
+ return _INPLACE[op](left, right)
361
+
362
+
363
+ def compare(op: str, left: Any, right: Any) -> Any:
364
+ """``left <op> right`` for a single comparison, recorded as a guard."""
365
+ if op in ("in", "not in"):
366
+ if isinstance(left, Traced) and not isinstance(right, Traced):
367
+ left = observe_value("observe", _identity, left) # membership in a plain container
368
+ result = left in right
369
+ return not result if op == "not in" else result
370
+ if isinstance(left, Traced) or isinstance(right, Traced):
371
+ return observe_value(op, COMPARISON_OPS[op][0], left, right)
372
+ return COMPARISON_OPS[op][0](left, right)
373
+
374
+
375
+ def _identity(value: Any) -> Any:
376
+ return value
377
+
378
+
379
+ def _plain_key(key: Any) -> Any:
380
+ if isinstance(key, Traced):
381
+ return observe_value("observe", _identity, key)
382
+ if type(key) is slice and _has_traced((key.start, key.stop, key.step)):
383
+ return slice(_plain_key(key.start), _plain_key(key.stop), _plain_key(key.step))
384
+ if type(key) is tuple and _has_traced(key, depth=1):
385
+ return tuple(_plain_key(k) for k in key)
386
+ return key
387
+
388
+
389
+ def key(value: Any) -> Any:
390
+ """A subscript key for a store or delete: traced keys become guarded plain keys."""
391
+ return _plain_key(value)
392
+
393
+
394
+ def getitem(container: Any, index: Any) -> Any:
395
+ """``container[index]``; a traced index into a plain container is guarded."""
396
+ if not isinstance(container, Traced):
397
+ index = _plain_key(index)
398
+ return container[index]
399
+
400
+
401
+ def fstring_(*parts: Any) -> Any:
402
+ """Build an f-string, recording it if any interpolated value is traced."""
403
+ if _has_traced(parts, depth=8): # values may be containers of traced values
404
+ return _core.call("fstring", fstring, parts)
405
+ return fstring(*parts)
406
+
407
+
408
+ class _Runtime:
409
+ """Namespace object bound to :data:`RUNTIME_NAME` in instrumented code."""
410
+
411
+ call = staticmethod(call)
412
+ binop = staticmethod(binop)
413
+ ibinop = staticmethod(ibinop)
414
+ compare = staticmethod(compare)
415
+ getitem = staticmethod(getitem)
416
+ key = staticmethod(key)
417
+ fstring = staticmethod(fstring_)
418
+ slice = slice
419
+
420
+
421
+ RUNTIME = _Runtime()
422
+
423
+ # --------------------------------------------------------------------------
424
+ # AST transformation
425
+ # --------------------------------------------------------------------------
426
+
427
+ _BINOP_NAMES: dict[type[ast.operator], str] = {
428
+ ast.Add: "add",
429
+ ast.Sub: "sub",
430
+ ast.Mult: "mul",
431
+ ast.Div: "truediv",
432
+ ast.FloorDiv: "floordiv",
433
+ ast.Mod: "mod",
434
+ ast.Pow: "pow",
435
+ ast.LShift: "lshift",
436
+ ast.RShift: "rshift",
437
+ ast.BitAnd: "and",
438
+ ast.BitOr: "or",
439
+ ast.BitXor: "xor",
440
+ }
441
+ _COMPARE_NAMES: dict[type[ast.cmpop], str] = {
442
+ ast.Lt: "lt",
443
+ ast.LtE: "le",
444
+ ast.Gt: "gt",
445
+ ast.GtE: "ge",
446
+ ast.Eq: "eq",
447
+ ast.NotEq: "ne",
448
+ ast.In: "in",
449
+ ast.NotIn: "not in",
450
+ }
451
+ # Calls that depend on the calling frame must not be wrapped.
452
+ _FRAME_SENSITIVE = frozenset(
453
+ {"super", "locals", "vars", "globals", "eval", "exec", "dir", "breakpoint"}
454
+ )
455
+ _CONVERSIONS = {-1: None, 115: "s", 114: "r", 97: "a"}
456
+
457
+
458
+ def _runtime(attr: str) -> ast.expr:
459
+ return ast.Attribute(ast.Name(RUNTIME_NAME, ast.Load()), attr, ast.Load())
460
+
461
+
462
+ def _rt_call(attr: str, *args: ast.expr) -> ast.Call:
463
+ return ast.Call(_runtime(attr), list(args), [])
464
+
465
+
466
+ def _is_constant(node: ast.expr) -> bool:
467
+ return isinstance(node, ast.Constant)
468
+
469
+
470
+ class _Instrumenter(ast.NodeTransformer):
471
+ def __init__(self) -> None:
472
+ self._temps = itertools.count()
473
+
474
+ def _temp(self) -> str:
475
+ return f"__redroot_tmp{next(self._temps)}"
476
+
477
+ # Annotations are left alone: rewriting them gains nothing.
478
+ def visit_arg(self, node: ast.arg) -> ast.arg:
479
+ return node
480
+
481
+ def visit_AnnAssign(self, node: ast.AnnAssign) -> ast.AnnAssign:
482
+ node.target = self.visit(node.target)
483
+ if node.value is not None:
484
+ node.value = self.visit(node.value)
485
+ return node
486
+
487
+ def _visit_function(self, node: ast.FunctionDef | ast.AsyncFunctionDef) -> ast.AST:
488
+ returns, node.returns = node.returns, None
489
+ self.generic_visit(node)
490
+ node.returns = returns
491
+ return node
492
+
493
+ visit_FunctionDef = _visit_function
494
+ visit_AsyncFunctionDef = _visit_function
495
+
496
+ def visit_BinOp(self, node: ast.BinOp) -> ast.expr:
497
+ self.generic_visit(node)
498
+ name = _BINOP_NAMES.get(type(node.op))
499
+ if name is None or (_is_constant(node.left) and _is_constant(node.right)):
500
+ return node
501
+ return ast.copy_location(_rt_call("binop", ast.Constant(name), node.left, node.right), node)
502
+
503
+ def visit_Compare(self, node: ast.Compare) -> ast.expr:
504
+ self.generic_visit(node)
505
+ if len(node.ops) != 1:
506
+ return node # chained comparisons keep their overloading-based tracing
507
+ name = _COMPARE_NAMES.get(type(node.ops[0]))
508
+ if name is None or (_is_constant(node.left) and _is_constant(node.comparators[0])):
509
+ return node
510
+ call = _rt_call("compare", ast.Constant(name), node.left, node.comparators[0])
511
+ return ast.copy_location(call, node)
512
+
513
+ def visit_Call(self, node: ast.Call) -> ast.expr:
514
+ self.generic_visit(node)
515
+ if isinstance(node.func, ast.Name) and node.func.id in _FRAME_SENSITIVE:
516
+ return node
517
+ call = ast.Call(_runtime("call"), [node.func, *node.args], node.keywords)
518
+ return ast.copy_location(call, node)
519
+
520
+ def visit_JoinedStr(self, node: ast.JoinedStr) -> ast.expr:
521
+ parts: list[ast.expr] = []
522
+ for value in node.values:
523
+ if isinstance(value, ast.FormattedValue):
524
+ spec: ast.expr = (
525
+ self.visit(value.format_spec)
526
+ if value.format_spec is not None
527
+ else ast.Constant("")
528
+ )
529
+ conversion = ast.Constant(_CONVERSIONS.get(value.conversion))
530
+ parts.append(ast.Tuple([self.visit(value.value), conversion, spec], ast.Load()))
531
+ else:
532
+ parts.append(value)
533
+ return ast.copy_location(_rt_call("fstring", *parts), node)
534
+
535
+ def _key(self, index: ast.expr) -> ast.expr:
536
+ if isinstance(index, ast.Slice):
537
+ bounds = [
538
+ b if b is not None else ast.Constant(None)
539
+ for b in (index.lower, index.upper, index.step)
540
+ ]
541
+ index = ast.Call(_runtime("slice"), bounds, [])
542
+ return index
543
+
544
+ def visit_Subscript(self, node: ast.Subscript) -> ast.expr:
545
+ self.generic_visit(node)
546
+ index = self._key(node.slice)
547
+ if isinstance(node.ctx, ast.Load):
548
+ return ast.copy_location(_rt_call("getitem", node.value, index), node)
549
+ node.slice = _rt_call("key", index)
550
+ return node
551
+
552
+ def visit_AugAssign(self, node: ast.AugAssign) -> ast.AST | list[ast.stmt]:
553
+ name = _BINOP_NAMES.get(type(node.op))
554
+ if name is None:
555
+ return self.generic_visit(node)
556
+ value = self.visit(node.value)
557
+ target = node.target
558
+ prelude: list[ast.stmt] = []
559
+ if isinstance(target, ast.Name):
560
+ load: ast.expr = ast.Name(target.id, ast.Load())
561
+ store: ast.expr = ast.Name(target.id, ast.Store())
562
+ elif isinstance(target, ast.Attribute):
563
+ obj = self._temp()
564
+ prelude.append(ast.Assign([ast.Name(obj, ast.Store())], self.visit(target.value)))
565
+ load = ast.Attribute(ast.Name(obj, ast.Load()), target.attr, ast.Load())
566
+ store = ast.Attribute(ast.Name(obj, ast.Load()), target.attr, ast.Store())
567
+ elif isinstance(target, ast.Subscript):
568
+ obj, idx = self._temp(), self._temp()
569
+ prelude.append(ast.Assign([ast.Name(obj, ast.Store())], self.visit(target.value)))
570
+ index = _rt_call("key", self._key(self.visit(target.slice)))
571
+ prelude.append(ast.Assign([ast.Name(idx, ast.Store())], index))
572
+ load = ast.Subscript(ast.Name(obj, ast.Load()), ast.Name(idx, ast.Load()), ast.Load())
573
+ store = ast.Subscript(ast.Name(obj, ast.Load()), ast.Name(idx, ast.Load()), ast.Store())
574
+ else: # pragma: no cover - the grammar allows no other targets
575
+ return self.generic_visit(node)
576
+ assign = ast.Assign([store], _rt_call("ibinop", ast.Constant(name), load, value))
577
+ statements = [*prelude, assign]
578
+ for statement in statements:
579
+ ast.copy_location(statement, node)
580
+ return statements
581
+
582
+
583
+ def instrument(source: str, filename: str = "<workflow>") -> ast.Module:
584
+ """Parse ``source`` and return its instrumented syntax tree."""
585
+ tree: ast.Module = _Instrumenter().visit(ast.parse(source, filename=filename))
586
+ return ast.fix_missing_locations(tree)
587
+
588
+
589
+ def compile_source(source: str, filename: str = "<workflow>") -> CodeType:
590
+ """Compile ``source`` with instrumentation, for use with ``exec``.
591
+
592
+ The code expects :data:`RUNTIME_NAME` in its globals; prefer
593
+ :func:`exec_source`, which provides it.
594
+ """
595
+ _INSTRUMENTED_FILES.add(filename)
596
+ # dont_inherit: never leak this module's __future__ flags into the workflow.
597
+ return compile(instrument(source, filename), filename, "exec", dont_inherit=True)
598
+
599
+
600
+ def exec_source(
601
+ source: str,
602
+ namespace: dict[str, Any] | None = None,
603
+ *,
604
+ filename: str = "<workflow>",
605
+ ) -> dict[str, Any]:
606
+ """Execute instrumented ``source`` and return its namespace.
607
+
608
+ Functions defined by the source stay instrumented whenever they are
609
+ called later, e.g. by :func:`redroot.run`.
610
+ """
611
+ namespace = {} if namespace is None else namespace
612
+ namespace[RUNTIME_NAME] = RUNTIME
613
+ exec(compile_source(source, filename), namespace) # noqa: S102 - executing the caller's code is the purpose
614
+ return namespace