plot3 0.4.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.
plot3/masking.py ADDED
@@ -0,0 +1,494 @@
1
+ """R-style bare-name / backtick column masking for plot3 (Jupyter layer).
2
+
3
+ Mirrors tidy3 so ggplot2-style aesthetics work without string quotes::
4
+
5
+ aes(x=wt, y=mpg, colour=cyl)
6
+ aes(x=`First Name`, y=`Age (%)`)
7
+ facet_wrap(cyl)
8
+
9
+ Two-phase design (SolveIt / IPython)::
10
+
11
+ 1. **Source preparser** turns backticks into a sentinel::
12
+
13
+ `First Name` → __plot3_bt__("First Name")
14
+
15
+ 2. **AST transformer** resolves the sentinel (and bare names) inside
16
+ :func:`aes` / :func:`facet_wrap` as **column-name strings**::
17
+
18
+ aes(x=wt, y=mpg) → aes(x="wt", y="mpg")
19
+ aes(x=`First Name`, y=mpg) → aes(x="First Name", y="mpg")
20
+ facet_wrap(cyl) → facet_wrap("cyl")
21
+
22
+ Plain ``.py`` files are unchanged — keep ``aes(x="wt", y="mpg")`` there.
23
+
24
+ When tidy3 is also loaded, its backtick preparser may emit
25
+ ``__tidy3_bt__("name")``; this transformer treats that sentinel the same way
26
+ inside ``aes`` / ``facet_wrap`` so dual-stack notebooks keep working.
27
+ """
28
+
29
+ from __future__ import annotations
30
+
31
+ import ast
32
+ import builtins
33
+ import io
34
+ import re
35
+ import tokenize
36
+ from typing import Any, Iterable
37
+
38
+ # Sentinels (plot3 own + tidy3 compatibility).
39
+ BT_NAME = "__plot3_bt__"
40
+ TIDY3_BT_NAME = "__tidy3_bt__"
41
+ _BT_SENTINELS = frozenset({BT_NAME, TIDY3_BT_NAME})
42
+
43
+ # Calls whose arguments / keywords are column *selectors* (→ strings).
44
+ _SELECTOR_FUNCS = frozenset(
45
+ {
46
+ "aes",
47
+ "vars",
48
+ "facet_wrap",
49
+ "facet_grid",
50
+ "transition_time",
51
+ "transition_states",
52
+ "slider",
53
+ }
54
+ )
55
+
56
+ # Keywords that are plain labels / options — never rewrite as column names.
57
+ _PASSTHROUGH_KW = frozenset(
58
+ {
59
+ "title",
60
+ "subtitle",
61
+ "caption",
62
+ "label",
63
+ "labels",
64
+ "trans",
65
+ "limits",
66
+ "palette",
67
+ "option",
68
+ "scales",
69
+ "ncol",
70
+ "nrow",
71
+ "bins",
72
+ "binwidth",
73
+ "width",
74
+ "size",
75
+ "alpha",
76
+ "color", # const color on geoms is a colour code, not a column
77
+ "colour",
78
+ "linewidth",
79
+ "wireframe",
80
+ "levels",
81
+ "n",
82
+ "bw",
83
+ "kernel",
84
+ "trim",
85
+ "coef",
86
+ "varwidth",
87
+ "outlier",
88
+ "na_rm",
89
+ "hide",
90
+ "height",
91
+ "theme",
92
+ "kind",
93
+ "max_points",
94
+ "size_mode",
95
+ "fov",
96
+ "near",
97
+ "far",
98
+ "target",
99
+ "up",
100
+ "position",
101
+ "mark",
102
+ "at",
103
+ "stream",
104
+ "tlim",
105
+ "baseline",
106
+ }
107
+ )
108
+
109
+ _FORMULA_CALLS = frozenset({"geom_function", "geom_vector_field"})
110
+
111
+ _BT_RE = re.compile(r"`([^`\n]+)`")
112
+
113
+
114
+ def rewrite_formula_carets(source: str) -> str:
115
+ """Inside ``geom_function(...)``, turn ``^`` into ``**`` before Python parses.
116
+
117
+ ``2*x^3`` and ``(2*x)^3`` are the same tree once Python has parsed ``^``
118
+ as xor, so the rewrite has to happen on the text.
119
+ """
120
+ if "^" not in source or not any(name in source for name in _FORMULA_CALLS):
121
+ return source
122
+ try:
123
+ tokens = list(tokenize.generate_tokens(io.StringIO(source).readline))
124
+ except (tokenize.TokenError, SyntaxError):
125
+ return source
126
+ starts = [0]
127
+ for line in source.splitlines(keepends=True):
128
+ starts.append(starts[-1] + len(line))
129
+ hits: list[int] = []
130
+ depth = 0 # paren depth inside the current geom_function call
131
+ armed = False # just saw the name geom_function
132
+ for tok in tokens:
133
+ if depth == 0:
134
+ if tok.type == tokenize.NAME and tok.string in _FORMULA_CALLS:
135
+ armed = True
136
+ continue
137
+ if armed and tok.type == tokenize.OP and tok.string == "(":
138
+ depth = 1
139
+ armed = False
140
+ continue
141
+ if tok.type == tokenize.OP:
142
+ if tok.string in "([{":
143
+ depth += 1
144
+ elif tok.string in ")]}":
145
+ depth -= 1
146
+ elif tok.string == "^":
147
+ row, col = tok.start
148
+ hits.append(starts[row - 1] + col)
149
+ for pos in reversed(hits):
150
+ source = source[:pos] + "**" + source[pos + 1 :]
151
+ return source
152
+
153
+
154
+ def rewrite_backticks(source: str) -> str:
155
+ """Preparse backticks, and ``^`` inside ``geom_function``, before parsing.
156
+
157
+ `` `col name` `` → ``__plot3_bt__("col name")``. ``^`` → ``**`` only
158
+ inside a ``geom_function(...)`` call, where it means power.
159
+ """
160
+ source = rewrite_formula_carets(source)
161
+ if "`" not in source:
162
+ return source
163
+ return _BT_RE.sub(lambda m: f"{BT_NAME}({m.group(1)!r})", source)
164
+
165
+
166
+ def plot3_backtick_transform(lines: list[str]) -> list[str]:
167
+ """IPython input transformer: backticks, and ``^`` inside geom_function."""
168
+ if not lines:
169
+ return lines
170
+ src = "".join(lines)
171
+ if "`" not in src and "^" not in src:
172
+ return lines
173
+ out = rewrite_backticks(src)
174
+ if out == src:
175
+ return lines
176
+ if out.endswith("\n"):
177
+ return out.splitlines(keepends=True)
178
+ parts = [ln + "\n" for ln in out.splitlines()]
179
+ return parts or [out]
180
+
181
+
182
+ def _is_bt_call(node: ast.AST) -> bool:
183
+ return (
184
+ isinstance(node, ast.Call)
185
+ and isinstance(node.func, ast.Name)
186
+ and node.func.id in _BT_SENTINELS
187
+ and len(node.args) == 1
188
+ and isinstance(node.args[0], ast.Constant)
189
+ and isinstance(node.args[0].value, str)
190
+ )
191
+
192
+
193
+ def _string_const(name: str, old: ast.AST) -> ast.AST:
194
+ return ast.copy_location(ast.Constant(value=name), old)
195
+
196
+
197
+ class MaskSelectors(ast.NodeTransformer):
198
+ """Rewrite bare names / backtick sentinels to string constants."""
199
+
200
+ def __init__(self, known: set[str]):
201
+ self.known = known
202
+
203
+ def visit_Name(self, node: ast.Name) -> ast.AST:
204
+ if isinstance(node.ctx, ast.Load) and node.id not in self.known:
205
+ return _string_const(node.id, node)
206
+ return node
207
+
208
+ def visit_Call(self, node: ast.Call) -> ast.AST:
209
+ if _is_bt_call(node):
210
+ return _string_const(str(node.args[0].value), node)
211
+
212
+ # Nested calls (e.g. unlikely helpers): keep func name; mask args.
213
+ if not isinstance(node.func, ast.Name):
214
+ node.func = self.visit(node.func)
215
+ node.args = [self.visit(arg) for arg in node.args]
216
+ node.keywords = [
217
+ ast.keyword(arg=kw.arg, value=self.visit(kw.value))
218
+ for kw in node.keywords
219
+ ]
220
+ return node
221
+
222
+
223
+ def default_known_names(extra: Iterable[str] | None = None) -> set[str]:
224
+ """Names that must not become column strings (funcs, builtins, API)."""
225
+ known = set(dir(builtins))
226
+ known.update(_BT_SENTINELS)
227
+ known.update({"True", "False", "None"})
228
+ try:
229
+ import plot3 as p3
230
+
231
+ for name in getattr(p3, "__all__", ()):
232
+ if not str(name).startswith("_"):
233
+ known.add(name)
234
+ except Exception:
235
+ pass
236
+ # Common frame / helper names left alone when present in the user ns
237
+ # are added via extra from IPython.user_ns.
238
+ if extra:
239
+ known.update(extra)
240
+ return known
241
+
242
+
243
+ class Plot3MaskTransformer(ast.NodeTransformer):
244
+ """AST pass: bare names / backticks → column strings inside aes / facet_wrap."""
245
+
246
+ def __init__(self, known: set[str] | None = None):
247
+ self._known_static = known
248
+
249
+ def _known(self) -> set[str]:
250
+ if self._known_static is not None:
251
+ return set(self._known_static)
252
+ extra: set[str] = set()
253
+ try:
254
+ from IPython import get_ipython
255
+
256
+ ip = get_ipython()
257
+ if ip is not None and getattr(ip, "user_ns", None) is not None:
258
+ extra.update(ip.user_ns.keys())
259
+ except Exception:
260
+ pass
261
+ return default_known_names(extra)
262
+
263
+ def _mask_selector(self, node: ast.AST) -> ast.AST:
264
+ return MaskSelectors(self._known()).visit(node)
265
+
266
+ def visit_Call(self, node: ast.Call) -> ast.AST:
267
+ # Always recurse so nested aes(...) inside geom_point(aes(...)) is seen.
268
+ if not isinstance(node.func, ast.Name):
269
+ return self.generic_visit(node)
270
+
271
+ name = node.func.id
272
+ if name in _FORMULA_CALLS:
273
+ return self._visit_geom_function(node)
274
+ if name not in _SELECTOR_FUNCS:
275
+ return self.generic_visit(node)
276
+
277
+ if name == "aes":
278
+ # All positional + keyword values are column selectors.
279
+ node.args = [self._mask_selector(a) for a in node.args]
280
+ node.keywords = [
281
+ ast.keyword(arg=kw.arg, value=self._mask_selector(kw.value))
282
+ for kw in node.keywords
283
+ ]
284
+ return node
285
+
286
+ if name == "transition_time":
287
+ # The column is a selector. Parameter ranges (``a=(0, 2*pi)``)
288
+ # and ``frames=`` stay Python, so ``pi`` is not quoted.
289
+ node.args = [self._mask_selector(a) for a in node.args]
290
+ new_kws: list[ast.keyword] = []
291
+ for kw in node.keywords:
292
+ if kw.arg == "column":
293
+ new_kws.append(
294
+ ast.keyword(
295
+ arg=kw.arg, value=self._mask_selector(kw.value)
296
+ )
297
+ )
298
+ else:
299
+ new_kws.append(
300
+ ast.keyword(arg=kw.arg, value=self.visit(kw.value))
301
+ )
302
+ node.keywords = new_kws
303
+ return node
304
+
305
+ if name == "slider":
306
+ # Ranges and ``steps=`` stay Python, so ``pi`` is not quoted.
307
+ node.args = [self.visit(a) for a in node.args]
308
+ node.keywords = [
309
+ ast.keyword(arg=kw.arg, value=self.visit(kw.value))
310
+ for kw in node.keywords
311
+ ]
312
+ return node
313
+
314
+ if name == "transition_states":
315
+ # The frame column is a selector, positional or keyword.
316
+ node.args = [self._mask_selector(a) for a in node.args]
317
+ node.keywords = [
318
+ ast.keyword(arg=kw.arg, value=self._mask_selector(kw.value))
319
+ for kw in node.keywords
320
+ ]
321
+ return node
322
+
323
+ if name in {"facet_wrap", "facet_grid"}:
324
+ # facets / rows / cols are column selectors; ncol, nrow, and
325
+ # scales stay as-is.
326
+ node.args = [self._mask_selector(a) for a in node.args]
327
+ new_kws: list[ast.keyword] = []
328
+ for kw in node.keywords:
329
+ if kw.arg in (None, "facets") or (
330
+ kw.arg is not None and kw.arg not in _PASSTHROUGH_KW
331
+ ):
332
+ new_kws.append(
333
+ ast.keyword(arg=kw.arg, value=self._mask_selector(kw.value))
334
+ )
335
+ else:
336
+ new_kws.append(kw)
337
+ node.keywords = new_kws
338
+ return node
339
+
340
+ return self.generic_visit(node)
341
+
342
+ def _visit_geom_function(self, node: ast.Call) -> ast.AST:
343
+ """Stringify formula arguments (``x=``, ``y=``, ``z=``, ``f=``, ``dx=``, ``dy=``).
344
+
345
+ A parametric curve and a vector field quote every formula keyword.
346
+ A number such as ``z=1`` stays Python, so it is a coefficient and
347
+ not a second formula. ``a=``, ``xlim=``, ``n=``, ``mark=``, and
348
+ ``tlim=`` stay Python too. ``^`` is already ``**`` if the text hook ran.
349
+ """
350
+ known = self._known()
351
+ call_name = node.func.id if isinstance(node.func, ast.Name) else ""
352
+ axis_names = (
353
+ {"dx", "dy", "dz"}
354
+ if call_name == "geom_vector_field"
355
+ else {"x", "y", "z", "f"}
356
+ )
357
+ formula_done = False
358
+ new_args: list[ast.AST] = []
359
+ for index, arg in enumerate(node.args):
360
+ if index == 0 and not formula_done and _is_formula_expr(arg, known):
361
+ new_args.append(_formula_string(arg))
362
+ formula_done = True
363
+ else:
364
+ new_args.append(self.visit(arg))
365
+ expr_axes: list[ast.keyword] = []
366
+ numeric_axes: list[ast.keyword] = []
367
+ for kw in node.keywords:
368
+ if kw.arg in axis_names and _is_number_const(kw.value):
369
+ numeric_axes.append(kw)
370
+ elif kw.arg in axis_names and _is_formula_expr(kw.value, known):
371
+ expr_axes.append(kw)
372
+ # geom_function(y=2) is a constant formula. z=1 next to a real
373
+ # formula is a coefficient, so it stays a number.
374
+ quote_numeric = (
375
+ call_name == "geom_function"
376
+ and not formula_done
377
+ and not expr_axes
378
+ and len(numeric_axes) == 1
379
+ )
380
+ quoted = {id(kw) for kw in expr_axes}
381
+ if quote_numeric:
382
+ quoted.add(id(numeric_axes[0]))
383
+ new_kws: list[ast.keyword] = []
384
+ for kw in node.keywords:
385
+ if kw.arg == "mark" and _mark_is_bare(kw.value):
386
+ new_kws.append(
387
+ ast.keyword(arg="mark", value=_quote_mark(kw.value))
388
+ )
389
+ elif id(kw) in quoted:
390
+ new_kws.append(
391
+ ast.keyword(arg=kw.arg, value=_formula_string(kw.value))
392
+ )
393
+ else:
394
+ new_kws.append(
395
+ ast.keyword(arg=kw.arg, value=self.visit(kw.value))
396
+ )
397
+ node.args = new_args
398
+ node.keywords = new_kws
399
+ return node
400
+
401
+
402
+ def _is_number_const(node: ast.AST) -> bool:
403
+ """True for ``1`` and ``-1``, which are coefficients rather than formulas."""
404
+ if isinstance(node, ast.UnaryOp) and isinstance(node.op, (ast.UAdd, ast.USub)):
405
+ node = node.operand
406
+ return (
407
+ isinstance(node, ast.Constant)
408
+ and isinstance(node.value, (int, float))
409
+ and not isinstance(node.value, bool)
410
+ )
411
+
412
+
413
+ def _mark_is_bare(node: ast.AST) -> bool:
414
+ if isinstance(node, ast.Name):
415
+ return True
416
+ if isinstance(node, (ast.List, ast.Tuple)):
417
+ return bool(node.elts) and all(isinstance(elt, ast.Name) for elt in node.elts)
418
+ return False
419
+
420
+
421
+ def _quote_mark(node: ast.AST) -> ast.AST:
422
+ if isinstance(node, ast.Name):
423
+ return ast.copy_location(ast.Constant(value=node.id), node)
424
+ elts = [
425
+ ast.copy_location(ast.Constant(value=elt.id), elt)
426
+ if isinstance(elt, ast.Name)
427
+ else elt
428
+ for elt in node.elts
429
+ ]
430
+ copied = ast.List(elts=elts, ctx=ast.Load()) if isinstance(node, ast.List) else ast.Tuple(elts=elts, ctx=ast.Load())
431
+ return ast.copy_location(copied, node)
432
+
433
+
434
+ def _is_formula_expr(node: ast.AST, known: set[str]) -> bool:
435
+ """True when this argument is a formula to quote, not a Python value."""
436
+ if isinstance(node, ast.Lambda):
437
+ return False
438
+ if isinstance(node, ast.Constant) and isinstance(node.value, str):
439
+ return False
440
+ if isinstance(node, (ast.Attribute, ast.Subscript)):
441
+ return False
442
+ if isinstance(node, ast.Name) and node.id in known:
443
+ return False
444
+ if isinstance(node, ast.Call) and isinstance(node.func, ast.Attribute):
445
+ return False
446
+ return True
447
+
448
+
449
+ def _formula_string(node: ast.AST) -> ast.Constant:
450
+ # ``^`` left in the tree is xor, and ``2*x^3`` is already ``(2*x)^3``.
451
+ # Do not guess. The text hook should have rewritten it.
452
+ if any(isinstance(child, ast.BitXor) for child in ast.walk(node)):
453
+ raise ValueError('use ** or quote the formula: "y = x^2"')
454
+ text = ast.unparse(node)
455
+ return ast.copy_location(ast.Constant(value=text), node)
456
+
457
+
458
+ def apply_masking(
459
+ source: str,
460
+ *,
461
+ known: set[str] | None = None,
462
+ backticks: bool = True,
463
+ ) -> str:
464
+ """Apply backtick rewrite + AST masking; return unparsed source (tests)."""
465
+ text = rewrite_formula_carets(source)
466
+ if backticks:
467
+ text = rewrite_backticks(text)
468
+ tree = ast.parse(text)
469
+ tree = Plot3MaskTransformer(known=known or default_known_names()).visit(tree)
470
+ ast.fix_missing_locations(tree)
471
+ return ast.unparse(tree)
472
+
473
+
474
+ def is_mask_transformer(obj: Any) -> bool:
475
+ return isinstance(obj, Plot3MaskTransformer) or (
476
+ type(obj).__name__ == "Plot3MaskTransformer"
477
+ and getattr(type(obj), "__module__", "").startswith("plot3")
478
+ )
479
+
480
+
481
+ def is_backtick_transformer(obj: Any) -> bool:
482
+ return (
483
+ getattr(obj, "__module__", "") in {"plot3.masking", "plot3.jupyter"}
484
+ and getattr(obj, "__name__", "") == "plot3_backtick_transform"
485
+ )
486
+
487
+
488
+ def is_tidy3_backtick_transformer(obj: Any) -> bool:
489
+ """True if tidy3 already installed a backtick preparser (share it)."""
490
+ mod = getattr(obj, "__module__", "") or ""
491
+ name = getattr(obj, "__name__", "") or ""
492
+ return name.endswith("backtick_transform") and (
493
+ mod.startswith("tidy3") or name == "tidy3_backtick_transform"
494
+ )