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/__init__.py +301 -0
- plot3/__version__.py +1 -0
- plot3/aesexpr.py +271 -0
- plot3/build.py +3948 -0
- plot3/calculus.py +1179 -0
- plot3/compose.py +285 -0
- plot3/contour.py +476 -0
- plot3/craft.py +142 -0
- plot3/encode.py +68 -0
- plot3/expr.py +1557 -0
- plot3/flip.py +245 -0
- plot3/function.py +1301 -0
- plot3/geoms.py +2558 -0
- plot3/ggplot.py +713 -0
- plot3/io.py +76 -0
- plot3/jupyter.py +514 -0
- plot3/latexin.py +616 -0
- plot3/masking.py +494 -0
- plot3/mathtext.py +842 -0
- plot3/payload.py +216 -0
- plot3/remote.py +220 -0
- plot3/scales.py +387 -0
- plot3/scaling.py +636 -0
- plot3/special.py +407 -0
- plot3/stat2d.py +1539 -0
- plot3/static.py +3760 -0
- plot3/stats3d.py +462 -0
- plot3/table.py +775 -0
- plot3/themes.py +104 -0
- plot3/viewer.py +3354 -0
- plot3-0.4.0.dist-info/METADATA +504 -0
- plot3-0.4.0.dist-info/RECORD +35 -0
- plot3-0.4.0.dist-info/WHEEL +5 -0
- plot3-0.4.0.dist-info/licenses/LICENSE +21 -0
- plot3-0.4.0.dist-info/top_level.txt +1 -0
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
|
+
)
|