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/expr.py ADDED
@@ -0,0 +1,1557 @@
1
+ """Parse math formulas for ``geom_function``.
2
+
3
+ The quoted string is the main form (``"y = 2x + 2"``). A raw LaTeX string
4
+ (``r"\\frac{\\sin x}{x}"`` or ``"$xy$"``) is translated into that same text
5
+ first. A small normaliser turns math notation into Python, ``ast`` parses
6
+ it, and a whitelist walk rejects anything that is not arithmetic or a known
7
+ function call. Evaluation uses that checked tree with NumPy functions — not
8
+ a raw ``eval`` of the user string.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import ast
14
+ import io
15
+ import tokenize
16
+ from dataclasses import dataclass
17
+ from typing import Any, Callable
18
+
19
+ import numpy as np
20
+
21
+ from plot3 import special as _special
22
+
23
+ from plot3.latexin import (
24
+ is_latex,
25
+ latex_to_source,
26
+ reject_eaten_backslashes,
27
+ unwrap_latex,
28
+ )
29
+ from plot3.mathtext import formula_texts
30
+
31
+ # Single letters that may be plot variables without a keyword value.
32
+ # Other single letters (a, b, k, …) are coefficients and must be passed in.
33
+ _PLOT_LETTERS = set("xyztuvwrs")
34
+
35
+ _BINOPS = (ast.Add, ast.Sub, ast.Mult, ast.Div, ast.FloorDiv, ast.Mod, ast.Pow)
36
+ _UNARY = (ast.UAdd, ast.USub)
37
+ _CMPOPS = (ast.Lt, ast.LtE, ast.Gt, ast.GtE, ast.Eq, ast.NotEq)
38
+ _CMP_TEXT = {
39
+ ast.Lt: "<",
40
+ ast.LtE: "<=",
41
+ ast.Gt: ">",
42
+ ast.GtE: ">=",
43
+ ast.Eq: "==",
44
+ ast.NotEq: "!=",
45
+ }
46
+
47
+
48
+ class ExprError(ValueError):
49
+ """A formula plot3 will not draw, with a message aimed at the call site."""
50
+
51
+
52
+ def _is_number(value: Any) -> bool:
53
+ if isinstance(value, (bool, np.bool_)):
54
+ return False
55
+ return isinstance(value, (int, float, np.integer, np.floating))
56
+
57
+
58
+ def _fmt_num(value: Any) -> str:
59
+ number = float(value)
60
+ if number == int(number) and abs(number) < 1e15:
61
+ return str(int(number))
62
+ return f"{number:.12g}"
63
+
64
+
65
+ def math_namespace() -> dict[str, Any]:
66
+ """NumPy callables and constants a formula may use by name."""
67
+ return {
68
+ "sin": np.sin,
69
+ "cos": np.cos,
70
+ "tan": np.tan,
71
+ "asin": np.arcsin,
72
+ "acos": np.arccos,
73
+ "atan": np.arctan,
74
+ "arcsin": np.arcsin,
75
+ "arccos": np.arccos,
76
+ "arctan": np.arctan,
77
+ "sinh": np.sinh,
78
+ "cosh": np.cosh,
79
+ "tanh": np.tanh,
80
+ "exp": np.exp,
81
+ "log": np.log,
82
+ "ln": np.log,
83
+ "log10": np.log10,
84
+ "sqrt": np.sqrt,
85
+ "cbrt": np.cbrt,
86
+ "abs": np.abs,
87
+ "floor": np.floor,
88
+ "ceil": np.ceil,
89
+ "pi": float(np.pi),
90
+ "e": float(np.e),
91
+ "nan": float("nan"),
92
+ "where": np.where,
93
+ # gamma, beta, erf, and R-style densities (dnorm, dbeta, pnorm, ...).
94
+ **_special.FUNCTIONS,
95
+ }
96
+
97
+
98
+ _MATH = math_namespace()
99
+ _MATH_FUNCS = {name for name, value in _MATH.items() if callable(value)}
100
+ _MATH_CONSTS = {name for name, value in _MATH.items() if not callable(value)}
101
+ # Function names that are also Greek letters people use as coefficients.
102
+ _GREEK_FUNCS = frozenset({"gamma", "beta"})
103
+
104
+
105
+ @dataclass(frozen=True)
106
+ class Formula:
107
+ """A checked formula ready to sample.
108
+
109
+ ``mode`` is ``explicit`` (``y = f(x)``), ``implicit`` (``F = 0``),
110
+ ``callable`` (a Python function), ``parametric`` (``x = cos(t), y = sin(t)``),
111
+ ``inequality`` (``y > x^2``), or ``field`` (``dx = -y, dy = x``).
112
+ ``variables`` are the free plot variables in appearance order.
113
+ ``namespace`` holds constants, numeric parameters, and callables — not
114
+ the sampled arrays. ``relation`` is ``>``, ``>=``, ``<``, or ``<=`` for
115
+ an inequality. ``components`` are the named pieces of a parametric
116
+ curve or a vector field.
117
+ """
118
+
119
+ label: str
120
+ mode: str
121
+ code: Any
122
+ dependent: str | None
123
+ variables: tuple[str, ...]
124
+ namespace: dict[str, Any]
125
+ fn: Callable[..., Any] | None = None
126
+ fn_args: tuple[str, ...] = ()
127
+ latex: str = ""
128
+ pretty: str = ""
129
+ legend_latex: str = ""
130
+ legend_pretty: str = ""
131
+ caption_latex: str = ""
132
+ caption_pretty: str = ""
133
+ # Coefficients such as ``a`` in ``y = a x^2``. Empty unless parsing was
134
+ # asked to wait: a following transition_time or slider may bind them.
135
+ pending: tuple[str, ...] = ()
136
+ relation: str = ""
137
+ components: tuple[tuple[str, Any], ...] = ()
138
+ parameter: str = ""
139
+ # Right-hand side (or inequality boundary) for a symbolic derivative.
140
+ body: Any = None
141
+
142
+ def _repr_latex_(self) -> str:
143
+ """Notebook display of the parsed formula, not the raw input text."""
144
+ body = self.latex or self.label
145
+ return f"${body}$"
146
+
147
+
148
+ def parse_formula(
149
+ expr: Any,
150
+ params: dict[str, Any] | None = None,
151
+ *,
152
+ defer_missing: bool = False,
153
+ role: str = "curve",
154
+ ) -> Formula:
155
+ """Turn a string, callable, or sympy-like object into a :class:`Formula`.
156
+
157
+ ``defer_missing`` records unbound coefficients on ``Formula.pending``
158
+ instead of raising. ``geom_function`` uses that so a later
159
+ ``transition_time(a=(0, 3))`` or ``slider(a=(0, 3))`` can still bind
160
+ them. A direct ``parse_formula`` call keeps raising immediately.
161
+ """
162
+ bound = dict(params or {})
163
+ if isinstance(expr, str):
164
+ return _parse_math(expr, bound, defer_missing=defer_missing, role=role)
165
+ if callable(expr):
166
+ return _parse_callable(expr)
167
+ if getattr(expr, "free_symbols", None) is not None:
168
+ return _parse_sympyish(expr, bound, defer_missing=defer_missing)
169
+ raise TypeError(
170
+ "geom_function() expects a formula string or a callable, "
171
+ f"got {type(expr).__name__}"
172
+ )
173
+
174
+
175
+ def evaluate(formula: Formula, variables: dict[str, np.ndarray]) -> np.ndarray:
176
+ """Evaluate ``formula`` on array values for its free variables."""
177
+ # Invalid samples (log of a negative, divide by zero) become NaN quietly.
178
+ # A warning there pulls in ``__import__``, which the empty builtins hide.
179
+ with np.errstate(all="ignore"):
180
+ if formula.mode == "callable":
181
+ if formula.fn is None:
182
+ raise ExprError("geom_function() callable is missing")
183
+ args = [variables[name] for name in formula.fn_args]
184
+ value = formula.fn(*args)
185
+ else:
186
+ if formula.code is None:
187
+ raise ExprError("geom_function() formula is missing")
188
+ env = {"__builtins__": {}}
189
+ env.update(formula.namespace)
190
+ env.update(variables)
191
+ try:
192
+ value = eval(formula.code, env) # noqa: S307
193
+ except ExprError:
194
+ raise
195
+ except Exception as exc:
196
+ raise ExprError(f"could not evaluate formula: {exc}") from exc
197
+ return np.asarray(value, dtype=np.float64)
198
+
199
+
200
+ def _parse_callable(fn: Callable[..., Any]) -> Formula:
201
+ import inspect
202
+
203
+ try:
204
+ signature = inspect.signature(fn)
205
+ except (TypeError, ValueError):
206
+ arg_names = ("x",)
207
+ else:
208
+ arg_names = []
209
+ for param in signature.parameters.values():
210
+ if param.kind in (
211
+ inspect.Parameter.VAR_POSITIONAL,
212
+ inspect.Parameter.VAR_KEYWORD,
213
+ inspect.Parameter.KEYWORD_ONLY,
214
+ ):
215
+ continue
216
+ if param.default is not inspect.Parameter.empty:
217
+ continue
218
+ arg_names.append(param.name)
219
+ if not arg_names:
220
+ arg_names = ["x"]
221
+ if len(arg_names) > 2:
222
+ listed = ", ".join(arg_names)
223
+ raise ExprError(
224
+ f"too many free variables ({listed}): at most 2. "
225
+ "A function plot takes one argument (a curve) or two (a surface)"
226
+ )
227
+ label = f"f({', '.join(arg_names)})"
228
+ return Formula(
229
+ label=label,
230
+ mode="callable",
231
+ code=None,
232
+ dependent=None,
233
+ variables=tuple(arg_names),
234
+ namespace={},
235
+ fn=_wrap_user_callable(label, fn),
236
+ fn_args=tuple(arg_names),
237
+ **_callable_math(label),
238
+ )
239
+
240
+
241
+ def _parse_sympyish(
242
+ expr: Any, params: dict[str, Any], *, defer_missing: bool = False
243
+ ) -> Formula:
244
+ """Duck-type sympy: ``Equality`` via our parser, other exprs via lambdify."""
245
+ lhs = getattr(expr, "lhs", None)
246
+ rhs = getattr(expr, "rhs", None)
247
+ if lhs is not None and rhs is not None and type(expr).__name__ == "Equality":
248
+ return _parse_math(
249
+ f"({lhs}) = ({rhs})", params, defer_missing=defer_missing
250
+ )
251
+
252
+ try:
253
+ import sympy
254
+ except ImportError:
255
+ sympy = None
256
+ if sympy is not None and isinstance(expr, sympy.Basic):
257
+ symbols = list(expr.free_symbols)
258
+ symbols.sort(
259
+ key=lambda sym: (
260
+ sym.name not in _PLOT_LETTERS,
261
+ sym.name,
262
+ )
263
+ )
264
+ fn = sympy.lambdify(symbols, expr, modules="numpy")
265
+ names = tuple(sym.name for sym in symbols)
266
+ if len(names) > 2:
267
+ listed = ", ".join(names)
268
+ raise ExprError(
269
+ f"too many free variables ({listed}): at most 2. "
270
+ f"Pass parameters as keywords, e.g. {names[-1]}=1"
271
+ )
272
+ return Formula(
273
+ label=str(expr),
274
+ mode="callable",
275
+ code=None,
276
+ dependent=None,
277
+ variables=names,
278
+ namespace={},
279
+ fn=_wrap_user_callable(str(expr), fn),
280
+ fn_args=names,
281
+ **_callable_math(str(expr)),
282
+ )
283
+ return _parse_math(str(expr), params, defer_missing=defer_missing)
284
+
285
+
286
+ def _evaluated_tree(source: str, params: dict[str, Any] | None = None) -> ast.AST:
287
+ """The tree ``evaluate`` runs. Implicit equations are ``lhs - rhs``."""
288
+ sink: dict[str, ast.AST] = {}
289
+ _parse_math(source, dict(params or {}), sink)
290
+ return sink["tree"]
291
+
292
+
293
+ def _parse_math(
294
+ source: str,
295
+ params: dict[str, Any],
296
+ _sink: dict[str, ast.AST] | None = None,
297
+ *,
298
+ defer_missing: bool = False,
299
+ role: str = "curve",
300
+ ) -> Formula:
301
+ try:
302
+ reject_eaten_backslashes(source)
303
+ except ValueError as exc:
304
+ raise ExprError(str(exc)) from exc
305
+ original = source.strip()
306
+ if not original:
307
+ raise ExprError(
308
+ 'geom_function() needs a formula, for example geom_function("y = 2x")'
309
+ )
310
+ numbers, callables = _split_bindings(params)
311
+ func_names = set(_MATH_FUNCS)
312
+ func_names.update(callables)
313
+ # Names that are values, so ``a(x+1)`` means ``a*(x+1)`` rather than a call.
314
+ value_names = set(numbers) | set(_MATH_CONSTS)
315
+ user_latex = None
316
+ plain = original
317
+ if is_latex(original):
318
+ user_latex = unwrap_latex(original)
319
+ if not user_latex:
320
+ raise ExprError(
321
+ 'geom_function() needs a formula, for example geom_function("y = 2x")'
322
+ )
323
+ try:
324
+ # e^{x} is exp(x) unless the user passed a value for e.
325
+ plain = latex_to_source(user_latex, e_is_constant="e" not in numbers)
326
+ except ValueError as exc:
327
+ raise ExprError(str(exc)) from exc
328
+ pieces = _split_top(plain, ",")
329
+ if len(pieces) >= 2 and all(_has_top_equals(piece) for piece in pieces):
330
+ return _parse_components(
331
+ pieces,
332
+ original,
333
+ user_latex,
334
+ numbers,
335
+ callables,
336
+ func_names,
337
+ value_names,
338
+ defer_missing=defer_missing,
339
+ role=role,
340
+ sink=_sink,
341
+ )
342
+ lhs_src, rhs_src, _ignored = _normalize(plain, func_names, value_names)
343
+ rhs_tree = _parse_side(rhs_src, original)
344
+ lhs_tree = _parse_side(lhs_src, original) if lhs_src is not None else None
345
+ _validate(rhs_tree, func_names)
346
+ if lhs_tree is not None:
347
+ _validate(lhs_tree, func_names)
348
+
349
+ rhs_values = _value_names(rhs_tree)
350
+ lhs_values = _value_names(lhs_tree) if lhs_tree is not None else []
351
+ # Appearance order, rhs first so "y = x + t" lists x before anything on the left.
352
+ ordered: list[str] = []
353
+ for name in rhs_values + lhs_values:
354
+ if name not in ordered:
355
+ ordered.append(name)
356
+
357
+ namespace: dict[str, Any] = {}
358
+ for name, fn in _MATH.items():
359
+ if name in _MATH_FUNCS:
360
+ namespace[name] = fn
361
+ for name, value in numbers.items():
362
+ if name in _MATH_FUNCS and name not in _GREEK_FUNCS:
363
+ raise ExprError(
364
+ f"'{name}' is a math function; pass a callable to replace it"
365
+ )
366
+ # beta=2 or gamma=0.5: the Greek coefficient, not the function.
367
+ namespace[name] = float(value)
368
+ for name, value in _MATH.items():
369
+ if name in _MATH_CONSTS and name not in numbers:
370
+ namespace[name] = value
371
+ for name, fn in callables.items():
372
+ namespace[name] = _wrap_user_callable(name, fn)
373
+
374
+ lone_lhs = isinstance(lhs_tree, ast.Name) and lhs_tree.id not in rhs_values
375
+ # The output name on the left (``v = 9.8 t``) is an axis label, not a
376
+ # coefficient that still needs a value.
377
+ skip_names = {lhs_tree.id} if lone_lhs else set()
378
+
379
+ free: list[str] = []
380
+ pending: list[str] = []
381
+ for name in ordered:
382
+ if name in skip_names:
383
+ continue
384
+ if name in numbers or name in _MATH_CONSTS:
385
+ continue
386
+ if name in _GREEK_FUNCS and name not in callables:
387
+ raise ExprError(
388
+ f"'{name}' is the {name.capitalize()} function here. Pass "
389
+ f"{name}=2 at the end to use it as a coefficient, or call {name}(...)"
390
+ )
391
+ if name in callables or name in _MATH_FUNCS:
392
+ raise ExprError(f"'{name}' is a function; call it as {name}(...)")
393
+ hint = _juxtaposition_hint(name)
394
+ if hint is not None:
395
+ raise ExprError(f"unknown name {name!r}: did you mean {hint}?")
396
+ # A coefficient is not a plot variable. Recording it on ``pending``
397
+ # (instead of ``free``) keeps ``y = a x^2`` a curve, not a surface.
398
+ if len(name) == 1 and name not in _PLOT_LETTERS:
399
+ if defer_missing:
400
+ pending.append(name)
401
+ continue
402
+ raise ExprError(_missing_param_message(name))
403
+ free.append(name)
404
+
405
+ if lhs_tree is None and isinstance(rhs_tree, ast.Compare):
406
+ return _parse_inequality(
407
+ rhs_tree,
408
+ original,
409
+ user_latex,
410
+ numbers,
411
+ callables,
412
+ namespace,
413
+ pending,
414
+ defer_missing=defer_missing,
415
+ sink=_sink,
416
+ )
417
+
418
+ if lhs_tree is None:
419
+ mode = "explicit"
420
+ if len(free) == 0:
421
+ dependent = "y"
422
+ elif len(free) == 1:
423
+ dependent = "x" if free[0] == "y" else "y"
424
+ else:
425
+ dependent = "z" if "z" not in free else "f"
426
+ code_tree = rhs_tree
427
+ elif lone_lhs:
428
+ mode = "explicit"
429
+ dependent = lhs_tree.id
430
+ code_tree = rhs_tree
431
+ else:
432
+ mode = "implicit"
433
+ dependent = None
434
+ code_tree = ast.BinOp(left=lhs_tree, op=ast.Sub(), right=rhs_tree)
435
+ # Names only on the left still count (already in ``free``).
436
+
437
+ if len(free) > 2:
438
+ listed = ", ".join(free)
439
+ raise ExprError(
440
+ f"too many free variables ({listed}): at most 2. "
441
+ f"Pass parameters as keywords, e.g. {free[-1]}=1"
442
+ )
443
+ if mode == "implicit" and len(free) != 2:
444
+ listed = ", ".join(free) if free else "none"
445
+ raise ExprError(
446
+ f"implicit equation needs 2 free variables, got ({listed})"
447
+ )
448
+
449
+ if user_latex is not None:
450
+ label = user_latex
451
+ elif lhs_tree is None:
452
+ label = original if "=" in original else f"{dependent} = {original}"
453
+ else:
454
+ label = original
455
+
456
+ expr_node = ast.fix_missing_locations(ast.Expression(body=code_tree))
457
+ code = compile(expr_node, "<geom_function>", "eval")
458
+ texts = formula_texts(
459
+ lhs_tree,
460
+ rhs_tree,
461
+ mode=mode,
462
+ dependent=dependent,
463
+ parameters=numbers,
464
+ )
465
+ if user_latex is not None:
466
+ texts = _keep_user_latex(texts, user_latex)
467
+ if _sink is not None:
468
+ _sink["tree"] = code_tree
469
+ return Formula(
470
+ label=label,
471
+ mode=mode,
472
+ code=code,
473
+ dependent=dependent,
474
+ variables=tuple(free),
475
+ namespace=namespace,
476
+ pending=tuple(pending),
477
+ body=code_tree,
478
+ **texts,
479
+ )
480
+
481
+
482
+ def _split_top(source: str, sep: str) -> list[str]:
483
+ """Split on ``sep`` that is not inside brackets."""
484
+ parts: list[str] = []
485
+ depth = 0
486
+ start = 0
487
+ for index, char in enumerate(source):
488
+ if char in "([{":
489
+ depth += 1
490
+ elif char in ")]}":
491
+ depth = max(0, depth - 1)
492
+ elif char == sep and depth == 0:
493
+ piece = source[start:index].strip()
494
+ if piece:
495
+ parts.append(piece)
496
+ start = index + 1
497
+ tail = source[start:].strip()
498
+ if tail:
499
+ parts.append(tail)
500
+ return parts
501
+
502
+
503
+ def _has_top_equals(source: str) -> bool:
504
+ """True when ``source`` has one assignment ``=``, ignoring ``<=`` and ``==``."""
505
+ depth = 0
506
+ count = 0
507
+ index = 0
508
+ while index < len(source):
509
+ char = source[index]
510
+ if char in "([{":
511
+ depth += 1
512
+ elif char in ")]}":
513
+ depth = max(0, depth - 1)
514
+ elif char == "=" and depth == 0:
515
+ prev = source[index - 1] if index else ""
516
+ nxt = source[index + 1] if index + 1 < len(source) else ""
517
+ if prev in "<>!" or nxt == "=":
518
+ index += 1
519
+ continue
520
+ count += 1
521
+ index += 1
522
+ return count == 1
523
+
524
+
525
+ def _bind_namespace(
526
+ numbers: dict[str, float], callables: dict[str, Callable[..., Any]]
527
+ ) -> dict[str, Any]:
528
+ namespace: dict[str, Any] = {}
529
+ for name, fn in _MATH.items():
530
+ if name in _MATH_FUNCS:
531
+ namespace[name] = fn
532
+ for name, value in numbers.items():
533
+ if name in _MATH_FUNCS and name not in _GREEK_FUNCS:
534
+ raise ExprError(
535
+ f"'{name}' is a math function; pass a callable to replace it"
536
+ )
537
+ # beta=2 or gamma=0.5: the Greek coefficient, not the function.
538
+ namespace[name] = float(value)
539
+ for name, value in _MATH.items():
540
+ if name in _MATH_CONSTS and name not in numbers:
541
+ namespace[name] = value
542
+ for name, fn in callables.items():
543
+ namespace[name] = _wrap_user_callable(name, fn)
544
+ return namespace
545
+
546
+
547
+ def _classify_names(
548
+ ordered: list[str],
549
+ numbers: dict[str, float],
550
+ callables: dict[str, Callable[..., Any]],
551
+ defer_missing: bool,
552
+ skip: set[str],
553
+ ) -> tuple[list[str], list[str]]:
554
+ free: list[str] = []
555
+ pending: list[str] = []
556
+ for name in ordered:
557
+ if name in skip or name in free or name in pending:
558
+ continue
559
+ if name in numbers or name in _MATH_CONSTS:
560
+ continue
561
+ if name in _GREEK_FUNCS and name not in callables:
562
+ raise ExprError(
563
+ f"'{name}' is the {name.capitalize()} function here. Pass "
564
+ f"{name}=2 at the end to use it as a coefficient, or call {name}(...)"
565
+ )
566
+ if name in callables or name in _MATH_FUNCS:
567
+ raise ExprError(f"'{name}' is a function; call it as {name}(...)")
568
+ hint = _juxtaposition_hint(name)
569
+ if hint is not None:
570
+ raise ExprError(f"unknown name {name!r}: did you mean {hint}?")
571
+ if len(name) == 1 and name not in _PLOT_LETTERS:
572
+ if defer_missing:
573
+ pending.append(name)
574
+ continue
575
+ raise ExprError(_missing_param_message(name))
576
+ free.append(name)
577
+ return free, pending
578
+
579
+
580
+ def _compile_tree(tree: ast.AST):
581
+ expr = ast.fix_missing_locations(ast.Expression(body=tree))
582
+ return compile(expr, "<geom_function>", "eval")
583
+
584
+
585
+ def _pack_texts(latex: str, pretty: str, numbers: dict[str, float]) -> dict[str, str]:
586
+ caption_l, caption_p = latex, pretty
587
+ if numbers:
588
+ bits = [f"{name} = {_fmt_num(value)}" for name, value in numbers.items()]
589
+ caption_l = latex + " \\quad (" + ",\\ ".join(bits) + ")"
590
+ caption_p = pretty + " (" + ", ".join(bits) + ")"
591
+ return {
592
+ "latex": latex,
593
+ "pretty": pretty,
594
+ "legend_latex": latex,
595
+ "legend_pretty": pretty,
596
+ "caption_latex": caption_l,
597
+ "caption_pretty": caption_p,
598
+ }
599
+
600
+
601
+ def _parse_components(
602
+ pieces: list[str],
603
+ original: str,
604
+ user_latex: str | None,
605
+ numbers: dict[str, float],
606
+ callables: dict[str, Callable[..., Any]],
607
+ func_names: set[str],
608
+ value_names: set[str],
609
+ *,
610
+ defer_missing: bool,
611
+ role: str,
612
+ sink: dict[str, ast.AST] | None,
613
+ ) -> Formula:
614
+ named: list[tuple[str, ast.AST]] = []
615
+ for piece in pieces:
616
+ lhs_src, rhs_src, _ignored = _normalize(piece, func_names, value_names)
617
+ if lhs_src is None:
618
+ raise ExprError(
619
+ 'each piece needs a name on the left, for example '
620
+ '"x = cos(t), y = sin(t)"'
621
+ )
622
+ lhs_tree = _parse_side(lhs_src, original)
623
+ rhs_tree = _parse_side(rhs_src, original)
624
+ _validate(lhs_tree, func_names)
625
+ _validate(rhs_tree, func_names)
626
+ if not isinstance(lhs_tree, ast.Name):
627
+ raise ExprError(
628
+ 'each piece needs a name on the left, for example '
629
+ '"x = cos(t), y = sin(t)"'
630
+ )
631
+ named.append((lhs_tree.id, rhs_tree))
632
+ lhs_names = [name for name, _tree in named]
633
+ if len(set(lhs_names)) != len(lhs_names):
634
+ raise ExprError("each output can be assigned only once")
635
+ ordered: list[str] = []
636
+ for _name, tree in named:
637
+ for found in _value_names(tree):
638
+ if found not in ordered:
639
+ ordered.append(found)
640
+ free, pending = _classify_names(
641
+ ordered, numbers, callables, defer_missing, set(lhs_names)
642
+ )
643
+ namespace = _bind_namespace(numbers, callables)
644
+ fieldish = all(name in {"dx", "dy", "dz"} for name in lhs_names)
645
+ if fieldish:
646
+ if role != "field":
647
+ raise ExprError(
648
+ 'this is a vector field. Use geom_vector_field("dx = -y, dy = x")'
649
+ )
650
+ return _finish_field(
651
+ named, original, user_latex, numbers, namespace, free, pending, sink
652
+ )
653
+ if role == "field":
654
+ raise ExprError(
655
+ 'geom_vector_field() needs dx and dy, for example "dx = -y, dy = x"'
656
+ )
657
+ return _finish_parametric(
658
+ named, original, user_latex, numbers, namespace, free, pending, sink
659
+ )
660
+
661
+
662
+ def _finish_parametric(
663
+ named, original, user_latex, numbers, namespace, free, pending, sink
664
+ ) -> Formula:
665
+ if len(free) != 1:
666
+ listed = ", ".join(free) if free else "none"
667
+ raise ExprError(
668
+ f"a parametric curve needs one parameter, got ({listed}). "
669
+ 'For example geom_function("x = cos(t), y = sin(t)")'
670
+ )
671
+ if not {"x", "y"} <= {name for name, _tree in named}:
672
+ raise ExprError(
673
+ 'a parametric curve needs x and y, for example '
674
+ '"x = cos(t), y = sin(t)"'
675
+ )
676
+ latex_bits = []
677
+ pretty_bits = []
678
+ compiled = []
679
+ for name, tree in named:
680
+ bit = formula_texts(
681
+ ast.Name(id=name, ctx=ast.Load()),
682
+ tree,
683
+ mode="explicit",
684
+ dependent=name,
685
+ parameters=numbers,
686
+ )
687
+ latex_bits.append(bit["latex"])
688
+ pretty_bits.append(bit["pretty"])
689
+ compiled.append((name, _compile_tree(tree)))
690
+ texts = _pack_texts(", ".join(latex_bits), ", ".join(pretty_bits), {})
691
+ # Parameters are already inside each piece's caption. Keep one joint caption.
692
+ if numbers:
693
+ texts = _pack_texts(texts["latex"], texts["pretty"], numbers)
694
+ if user_latex is not None:
695
+ texts = _keep_user_latex(texts, user_latex)
696
+ if sink is not None:
697
+ sink["tree"] = named[0][1]
698
+ return Formula(
699
+ label=original,
700
+ mode="parametric",
701
+ code=compiled[0][1],
702
+ dependent=None,
703
+ variables=tuple(free),
704
+ namespace=namespace,
705
+ pending=tuple(pending),
706
+ components=tuple(compiled),
707
+ parameter=free[0],
708
+ **texts,
709
+ )
710
+
711
+
712
+ def _finish_field(
713
+ named, original, user_latex, numbers, namespace, free, pending, sink
714
+ ) -> Formula:
715
+ names = {name for name, _tree in named}
716
+ if "dz" in names:
717
+ raise ExprError("geom_vector_field() draws a plane field of dx and dy")
718
+ if names != {"dx", "dy"}:
719
+ raise ExprError(
720
+ 'geom_vector_field() needs dx and dy, for example "dx = -y, dy = x"'
721
+ )
722
+ if len(free) > 2:
723
+ listed = ", ".join(free)
724
+ raise ExprError(
725
+ f"too many free variables ({listed}): at most 2. "
726
+ f"Pass parameters as keywords, e.g. {free[-1]}=1"
727
+ )
728
+ by_name = {name: tree for name, tree in named}
729
+ compiled = tuple(
730
+ (name, _compile_tree(by_name[name])) for name in ("dx", "dy")
731
+ )
732
+ latex_bits = []
733
+ pretty_bits = []
734
+ for name in ("dx", "dy"):
735
+ bit = formula_texts(
736
+ ast.Name(id=name, ctx=ast.Load()),
737
+ by_name[name],
738
+ mode="explicit",
739
+ dependent=name,
740
+ parameters=numbers,
741
+ )
742
+ latex_bits.append(bit["latex"])
743
+ pretty_bits.append(bit["pretty"])
744
+ texts = _pack_texts(", ".join(latex_bits), ", ".join(pretty_bits), numbers)
745
+ if user_latex is not None:
746
+ texts = _keep_user_latex(texts, user_latex)
747
+ if sink is not None:
748
+ sink["tree"] = by_name["dx"]
749
+ return Formula(
750
+ label=original,
751
+ mode="field",
752
+ code=compiled[0][1],
753
+ dependent=None,
754
+ variables=tuple(free),
755
+ namespace=namespace,
756
+ pending=tuple(pending),
757
+ components=compiled,
758
+ **texts,
759
+ )
760
+
761
+
762
+ def _flip_relation(relation: str) -> str:
763
+ return {">": "<", ">=": "<=", "<": ">", "<=": ">="}[relation]
764
+
765
+
766
+ def _relation_marks(relation: str) -> tuple[str, str]:
767
+ return {
768
+ ">": (">", ">"),
769
+ ">=": ("\\ge", "≥"),
770
+ "<": ("<", "<"),
771
+ "<=": ("\\le", "≤"),
772
+ }[relation]
773
+
774
+
775
+ def _parse_inequality(
776
+ tree: ast.Compare,
777
+ original: str,
778
+ user_latex: str | None,
779
+ numbers: dict[str, float],
780
+ callables: dict[str, Callable[..., Any]],
781
+ namespace: dict[str, Any],
782
+ pending: list[str],
783
+ *,
784
+ defer_missing: bool,
785
+ sink: dict[str, ast.AST] | None,
786
+ ) -> Formula:
787
+ del defer_missing
788
+ relation = _CMP_TEXT[type(tree.ops[0])]
789
+ if relation in {"==", "!="}:
790
+ raise ExprError(
791
+ "use = for an equation. Shade a region with >, <, >=, or <="
792
+ )
793
+ left, right = tree.left, tree.comparators[0]
794
+
795
+ def axis_name(node: ast.AST) -> str | None:
796
+ if isinstance(node, ast.Name) and node.id in {"x", "y"}:
797
+ return node.id
798
+ return None
799
+
800
+ left_axis = axis_name(left)
801
+ right_axis = axis_name(right)
802
+ boundary: ast.AST | None = None
803
+ dependent: str | None = None
804
+ if left_axis and left_axis not in _value_names(right):
805
+ dependent = left_axis
806
+ boundary = right
807
+ elif right_axis and right_axis not in _value_names(left):
808
+ dependent = right_axis
809
+ boundary = left
810
+ relation = _flip_relation(relation)
811
+ if boundary is not None and dependent is not None:
812
+ ordered = []
813
+ for found in _value_names(boundary):
814
+ if found not in ordered:
815
+ ordered.append(found)
816
+ free, boundary_pending = _classify_names(
817
+ ordered, numbers, callables, True, set()
818
+ )
819
+ # A coefficient on the boundary was already recorded, or just now.
820
+ pending_names = list(dict.fromkeys([*pending, *boundary_pending]))
821
+ free = [
822
+ name
823
+ for name in free
824
+ if name not in pending_names and name != dependent
825
+ ]
826
+ if len(free) > 1:
827
+ listed = ", ".join(free)
828
+ raise ExprError(
829
+ f"too many free variables ({listed}): at most 2. "
830
+ f"Pass parameters as keywords, e.g. {free[-1]}=1"
831
+ )
832
+ raw = formula_texts(
833
+ ast.Name(id=dependent, ctx=ast.Load()),
834
+ boundary,
835
+ mode="explicit",
836
+ dependent=dependent,
837
+ parameters=numbers,
838
+ )
839
+ mark_l, mark_p = _relation_marks(relation)
840
+ texts = _pack_texts(
841
+ raw["latex"].replace(" = ", f" {mark_l} ", 1),
842
+ raw["pretty"].replace(" = ", f" {mark_p} ", 1),
843
+ {},
844
+ )
845
+ texts["caption_latex"] = raw["caption_latex"].replace(" = ", f" {mark_l} ", 1)
846
+ texts["caption_pretty"] = raw["caption_pretty"].replace(" = ", f" {mark_p} ", 1)
847
+ texts["legend_latex"] = texts["latex"]
848
+ texts["legend_pretty"] = texts["pretty"]
849
+ if user_latex is not None:
850
+ texts = _keep_user_latex(texts, user_latex)
851
+ if sink is not None:
852
+ sink["tree"] = boundary
853
+ return Formula(
854
+ label=original,
855
+ mode="inequality",
856
+ code=_compile_tree(boundary),
857
+ dependent=dependent,
858
+ variables=tuple(free),
859
+ namespace=namespace,
860
+ pending=tuple(pending_names),
861
+ relation=relation,
862
+ body=boundary,
863
+ **texts,
864
+ )
865
+ ordered = []
866
+ for found in _value_names(left) + _value_names(right):
867
+ if found not in ordered:
868
+ ordered.append(found)
869
+ free, region_pending = _classify_names(
870
+ ordered, numbers, callables, True, set()
871
+ )
872
+ pending_names = list(dict.fromkeys([*pending, *region_pending]))
873
+ free = [name for name in free if name not in pending_names]
874
+ if len(free) != 2:
875
+ listed = ", ".join(free) if free else "none"
876
+ raise ExprError(
877
+ "shade y > f(x), or a region in x and y such as "
878
+ f"x^2 + y^2 < 1 (got {listed})"
879
+ )
880
+ if relation in {">", ">="}:
881
+ signed = ast.BinOp(left=left, op=ast.Sub(), right=right)
882
+ else:
883
+ signed = ast.BinOp(left=right, op=ast.Sub(), right=left)
884
+ texts = _pack_texts(original, original, numbers)
885
+ if user_latex is not None:
886
+ texts = _keep_user_latex(texts, user_latex)
887
+ if sink is not None:
888
+ sink["tree"] = signed
889
+ return Formula(
890
+ label=original,
891
+ mode="inequality",
892
+ code=_compile_tree(signed),
893
+ dependent=None,
894
+ variables=tuple(free),
895
+ namespace=namespace,
896
+ pending=tuple(pending_names),
897
+ relation=relation,
898
+ body=signed,
899
+ **texts,
900
+ )
901
+
902
+
903
+ def differentiate(formula: Formula, var: str) -> ast.AST | None:
904
+ """Symbolic derivative of an explicit formula, or None if it is not one."""
905
+ tree = getattr(formula, "body", None)
906
+ if tree is None or formula.mode not in {"explicit", "inequality"}:
907
+ return None
908
+ if formula.mode == "inequality" and not formula.dependent:
909
+ return None
910
+ try:
911
+ return _simplify(_diff(tree, var))
912
+ except ExprError:
913
+ return None
914
+
915
+
916
+ def _diff(node: ast.AST, var: str) -> ast.AST:
917
+ if isinstance(node, ast.Constant) and isinstance(node.value, (int, float)):
918
+ return ast.Constant(value=0.0)
919
+ if isinstance(node, ast.Name):
920
+ return ast.Constant(value=1.0 if node.id == var else 0.0)
921
+ if isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.UAdd):
922
+ return _diff(node.operand, var)
923
+ if isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.USub):
924
+ return ast.UnaryOp(op=ast.USub(), operand=_diff(node.operand, var))
925
+ if isinstance(node, ast.BinOp) and isinstance(node.op, ast.Add):
926
+ return ast.BinOp(left=_diff(node.left, var), op=ast.Add(), right=_diff(node.right, var))
927
+ if isinstance(node, ast.BinOp) and isinstance(node.op, ast.Sub):
928
+ return ast.BinOp(left=_diff(node.left, var), op=ast.Sub(), right=_diff(node.right, var))
929
+ if isinstance(node, ast.BinOp) and isinstance(node.op, ast.Mult):
930
+ return ast.BinOp(
931
+ left=ast.BinOp(
932
+ left=_diff(node.left, var), op=ast.Mult(), right=node.right
933
+ ),
934
+ op=ast.Add(),
935
+ right=ast.BinOp(
936
+ left=node.left, op=ast.Mult(), right=_diff(node.right, var)
937
+ ),
938
+ )
939
+ if isinstance(node, ast.BinOp) and isinstance(node.op, ast.Div):
940
+ numer = ast.BinOp(
941
+ left=ast.BinOp(
942
+ left=_diff(node.left, var), op=ast.Mult(), right=node.right
943
+ ),
944
+ op=ast.Sub(),
945
+ right=ast.BinOp(
946
+ left=node.left, op=ast.Mult(), right=_diff(node.right, var)
947
+ ),
948
+ )
949
+ denom = ast.BinOp(left=node.right, op=ast.Pow(), right=ast.Constant(value=2))
950
+ return ast.BinOp(left=numer, op=ast.Div(), right=denom)
951
+ if isinstance(node, ast.BinOp) and isinstance(node.op, ast.Pow):
952
+ return _diff_pow(node, var)
953
+ if isinstance(node, ast.Call) and isinstance(node.func, ast.Name):
954
+ return _diff_call(node, var)
955
+ raise ExprError("no symbolic derivative")
956
+
957
+
958
+ def _diff_pow(node: ast.BinOp, var: str) -> ast.AST:
959
+ base, exp = node.left, node.right
960
+ if isinstance(exp, ast.Constant) and isinstance(exp.value, (int, float)):
961
+ power = ast.BinOp(
962
+ left=base,
963
+ op=ast.Pow(),
964
+ right=ast.Constant(value=exp.value - 1),
965
+ )
966
+ return ast.BinOp(
967
+ left=ast.BinOp(
968
+ left=ast.Constant(value=exp.value), op=ast.Mult(), right=power
969
+ ),
970
+ op=ast.Mult(),
971
+ right=_diff(base, var),
972
+ )
973
+ # a^u and the general case share u^v * (v' ln u + v u' / u).
974
+ log_base = ast.Call(func=ast.Name(id="log", ctx=ast.Load()), args=[base], keywords=[])
975
+ left = ast.BinOp(left=_diff(exp, var), op=ast.Mult(), right=log_base)
976
+ right = ast.BinOp(
977
+ left=exp,
978
+ op=ast.Mult(),
979
+ right=ast.BinOp(left=_diff(base, var), op=ast.Div(), right=base),
980
+ )
981
+ factor = ast.BinOp(left=left, op=ast.Add(), right=right)
982
+ return ast.BinOp(left=node, op=ast.Mult(), right=factor)
983
+
984
+
985
+ def _diff_call(node: ast.Call, var: str) -> ast.AST:
986
+ name = node.func.id
987
+ if len(node.args) != 1:
988
+ raise ExprError("no symbolic derivative")
989
+ arg = node.args[0]
990
+ inner = _diff(arg, var)
991
+ if name == "sin":
992
+ outer = ast.Call(func=ast.Name(id="cos", ctx=ast.Load()), args=[arg], keywords=[])
993
+ elif name == "cos":
994
+ outer = ast.UnaryOp(
995
+ op=ast.USub(),
996
+ operand=ast.Call(
997
+ func=ast.Name(id="sin", ctx=ast.Load()), args=[arg], keywords=[]
998
+ ),
999
+ )
1000
+ elif name == "tan":
1001
+ outer = ast.BinOp(
1002
+ left=ast.Constant(value=1),
1003
+ op=ast.Div(),
1004
+ right=ast.BinOp(
1005
+ left=ast.Call(
1006
+ func=ast.Name(id="cos", ctx=ast.Load()), args=[arg], keywords=[]
1007
+ ),
1008
+ op=ast.Pow(),
1009
+ right=ast.Constant(value=2),
1010
+ ),
1011
+ )
1012
+ elif name == "exp":
1013
+ outer = ast.Call(func=ast.Name(id="exp", ctx=ast.Load()), args=[arg], keywords=[])
1014
+ elif name in {"log", "ln", "log10"}:
1015
+ denom = arg
1016
+ if name == "log10":
1017
+ denom = ast.BinOp(
1018
+ left=arg,
1019
+ op=ast.Mult(),
1020
+ right=ast.Call(
1021
+ func=ast.Name(id="log", ctx=ast.Load()),
1022
+ args=[ast.Constant(value=10)],
1023
+ keywords=[],
1024
+ ),
1025
+ )
1026
+ outer = ast.BinOp(left=ast.Constant(value=1), op=ast.Div(), right=denom)
1027
+ elif name == "sqrt":
1028
+ outer = ast.BinOp(
1029
+ left=ast.Constant(value=1),
1030
+ op=ast.Div(),
1031
+ right=ast.BinOp(
1032
+ left=ast.Constant(value=2), op=ast.Mult(), right=node
1033
+ ),
1034
+ )
1035
+ elif name == "asin":
1036
+ outer = ast.BinOp(
1037
+ left=ast.Constant(value=1),
1038
+ op=ast.Div(),
1039
+ right=ast.Call(
1040
+ func=ast.Name(id="sqrt", ctx=ast.Load()),
1041
+ args=[
1042
+ ast.BinOp(
1043
+ left=ast.Constant(value=1),
1044
+ op=ast.Sub(),
1045
+ right=ast.BinOp(left=arg, op=ast.Pow(), right=ast.Constant(value=2)),
1046
+ )
1047
+ ],
1048
+ keywords=[],
1049
+ ),
1050
+ )
1051
+ elif name == "acos":
1052
+ positive = _diff_call(
1053
+ ast.Call(func=ast.Name(id="asin", ctx=ast.Load()), args=[arg], keywords=[]),
1054
+ var,
1055
+ )
1056
+ return ast.UnaryOp(op=ast.USub(), operand=positive)
1057
+ elif name == "atan":
1058
+ outer = ast.BinOp(
1059
+ left=ast.Constant(value=1),
1060
+ op=ast.Div(),
1061
+ right=ast.BinOp(
1062
+ left=ast.Constant(value=1),
1063
+ op=ast.Add(),
1064
+ right=ast.BinOp(left=arg, op=ast.Pow(), right=ast.Constant(value=2)),
1065
+ ),
1066
+ )
1067
+ else:
1068
+ raise ExprError("no symbolic derivative")
1069
+ return ast.BinOp(left=outer, op=ast.Mult(), right=inner)
1070
+
1071
+
1072
+ def _is_number_node(node: ast.AST, value: float | None = None) -> bool:
1073
+ if not isinstance(node, ast.Constant) or isinstance(node.value, bool):
1074
+ return False
1075
+ if not isinstance(node.value, (int, float)):
1076
+ return False
1077
+ if value is None:
1078
+ return True
1079
+ return float(node.value) == float(value)
1080
+
1081
+
1082
+ def _simplify(node: ast.AST) -> ast.AST:
1083
+ """Fold constants and drop factors of 0 and 1. The tree stays exact."""
1084
+ if isinstance(node, ast.UnaryOp):
1085
+ operand = _simplify(node.operand)
1086
+ if isinstance(node.op, ast.UAdd):
1087
+ return operand
1088
+ if isinstance(node.op, ast.USub):
1089
+ if isinstance(operand, ast.UnaryOp) and isinstance(operand.op, ast.USub):
1090
+ return operand.operand
1091
+ if _is_number_node(operand):
1092
+ return ast.Constant(value=-float(operand.value))
1093
+ return ast.UnaryOp(op=ast.USub(), operand=operand)
1094
+ return node
1095
+ if isinstance(node, ast.BinOp):
1096
+ left = _simplify(node.left)
1097
+ right = _simplify(node.right)
1098
+ if isinstance(node.op, ast.Add):
1099
+ if _is_number_node(left, 0):
1100
+ return right
1101
+ if _is_number_node(right, 0):
1102
+ return left
1103
+ if _is_number_node(left) and _is_number_node(right):
1104
+ return ast.Constant(value=float(left.value) + float(right.value))
1105
+ if isinstance(node.op, ast.Sub):
1106
+ if _is_number_node(right, 0):
1107
+ return left
1108
+ if _is_number_node(left) and _is_number_node(right):
1109
+ return ast.Constant(value=float(left.value) - float(right.value))
1110
+ if isinstance(node.op, ast.Mult):
1111
+ if _is_number_node(left, 0) or _is_number_node(right, 0):
1112
+ return ast.Constant(value=0.0)
1113
+ if _is_number_node(left, 1):
1114
+ return right
1115
+ if _is_number_node(right, 1):
1116
+ return left
1117
+ if _is_number_node(left) and _is_number_node(right):
1118
+ return ast.Constant(value=float(left.value) * float(right.value))
1119
+ if isinstance(node.op, ast.Div):
1120
+ if _is_number_node(left, 0):
1121
+ return ast.Constant(value=0.0)
1122
+ if _is_number_node(right, 1):
1123
+ return left
1124
+ if isinstance(node.op, ast.Pow):
1125
+ if _is_number_node(right, 1):
1126
+ return left
1127
+ if _is_number_node(right, 0):
1128
+ return ast.Constant(value=1.0)
1129
+ if _is_number_node(left) and _is_number_node(right):
1130
+ return ast.Constant(value=float(left.value) ** float(right.value))
1131
+ return ast.BinOp(left=left, op=node.op, right=right)
1132
+ if isinstance(node, ast.Call):
1133
+ return ast.Call(
1134
+ func=node.func,
1135
+ args=[_simplify(arg) for arg in node.args],
1136
+ keywords=[],
1137
+ )
1138
+ return node
1139
+
1140
+
1141
+ def _keep_user_latex(texts: dict[str, str], user_latex: str) -> dict[str, str]:
1142
+ """Show the LaTeX the user wrote, not a regenerated string.
1143
+
1144
+ Numeric parameters stay in the caption. The legend keeps their source
1145
+ so a pasted ``\\frac`` is not rewritten.
1146
+ """
1147
+ symbolic = texts["latex"]
1148
+ caption = texts["caption_latex"]
1149
+ suffix = caption[len(symbolic) :] if caption.startswith(symbolic) else ""
1150
+ out = dict(texts)
1151
+ out["latex"] = user_latex
1152
+ out["legend_latex"] = user_latex
1153
+ out["legend_pretty"] = texts["pretty"]
1154
+ out["caption_latex"] = user_latex + suffix
1155
+ return out
1156
+
1157
+
1158
+ def _callable_math(label: str) -> dict[str, str]:
1159
+ """A lambda has no tree. The signature is the whole display."""
1160
+ return {
1161
+ "latex": label,
1162
+ "pretty": label,
1163
+ "legend_latex": label,
1164
+ "legend_pretty": label,
1165
+ "caption_latex": label,
1166
+ "caption_pretty": label,
1167
+ }
1168
+
1169
+
1170
+ def _split_bindings(
1171
+ params: dict[str, Any],
1172
+ ) -> tuple[dict[str, float], dict[str, Callable[..., Any]]]:
1173
+ numbers: dict[str, float] = {}
1174
+ callables: dict[str, Callable[..., Any]] = {}
1175
+ for name, value in params.items():
1176
+ if _is_number(value):
1177
+ numbers[name] = float(value)
1178
+ elif callable(value):
1179
+ callables[name] = value
1180
+ else:
1181
+ raise TypeError(
1182
+ f"geom_function() keyword {name}={value!r} must be a number "
1183
+ "or a function"
1184
+ )
1185
+ return numbers, callables
1186
+
1187
+
1188
+ def _missing_param_message(name: str) -> str:
1189
+ found = _notebook_number(name)
1190
+ if found is not None:
1191
+ return (
1192
+ f"'{name}' has no value. Your notebook has {name} = {_fmt_num(found)}. "
1193
+ f"Use geom_function(..., {name}={name})"
1194
+ )
1195
+ return (
1196
+ f"'{name}' has no value. Pass it at the end: "
1197
+ f"geom_function(..., {name}=2)"
1198
+ )
1199
+
1200
+
1201
+ def _missing_param_build_message(name: str) -> str:
1202
+ """Same hint as parse time, plus the transition that can animate it.
1203
+
1204
+ Raised from ``build_spec``, once the figure's layers are known. A
1205
+ notebook still shows it on the cell that displays the plot.
1206
+ """
1207
+ found = _notebook_number(name)
1208
+ animate = (
1209
+ f", or animate it: + transition_time({name}=(0, 3))"
1210
+ f", or drag it: + slider({name}=(0, 3))"
1211
+ )
1212
+ if found is not None:
1213
+ return (
1214
+ f"'{name}' has no value. Your notebook has {name} = {_fmt_num(found)}. "
1215
+ f"Use geom_function(..., {name}={name}){animate}"
1216
+ )
1217
+ return (
1218
+ f"'{name}' has no value. Pass it at the end: "
1219
+ f"geom_function(..., {name}=2){animate}"
1220
+ )
1221
+
1222
+
1223
+ def _notebook_number(name: str) -> Any:
1224
+ """A numeric value from the notebook namespace, for a better error only."""
1225
+ try:
1226
+ from IPython import get_ipython
1227
+ except Exception:
1228
+ return None
1229
+ shell = get_ipython()
1230
+ if shell is None:
1231
+ return None
1232
+ user_ns = getattr(shell, "user_ns", None)
1233
+ if not isinstance(user_ns, dict) or name not in user_ns:
1234
+ return None
1235
+ value = user_ns[name]
1236
+ if _is_number(value):
1237
+ return value
1238
+ return None
1239
+
1240
+
1241
+ def _juxtaposition_hint(name: str) -> str | None:
1242
+ """``xy`` → ``x*y`` when every letter is itself a plot variable."""
1243
+ if len(name) < 2 or not name.isalpha():
1244
+ return None
1245
+ if all(ch in _PLOT_LETTERS for ch in name):
1246
+ return "*".join(name)
1247
+ return None
1248
+
1249
+
1250
+ def _value_names(tree: ast.AST | None) -> list[str]:
1251
+ """Names used as values, in source order. Call targets are skipped."""
1252
+ if tree is None:
1253
+ return []
1254
+ found: list[str] = []
1255
+
1256
+ def walk(node: ast.AST) -> None:
1257
+ if isinstance(node, ast.Call):
1258
+ for arg in node.args:
1259
+ walk(arg)
1260
+ return
1261
+ if isinstance(node, ast.Name):
1262
+ found.append(node.id)
1263
+ return
1264
+ for child in ast.iter_child_nodes(node):
1265
+ walk(child)
1266
+
1267
+ walk(tree)
1268
+ return found
1269
+
1270
+
1271
+ def _validate(tree: ast.AST, func_names: set[str]) -> None:
1272
+ allowed_ops = _BINOPS + _UNARY
1273
+ for node in ast.walk(tree):
1274
+ if isinstance(node, (ast.Expression, ast.Load)):
1275
+ continue
1276
+ if isinstance(node, allowed_ops):
1277
+ continue
1278
+ if isinstance(node, ast.BinOp):
1279
+ if not isinstance(node.op, _BINOPS):
1280
+ raise ExprError("only + - * / // % and power are allowed")
1281
+ continue
1282
+ if isinstance(node, ast.UnaryOp):
1283
+ if not isinstance(node.op, _UNARY):
1284
+ raise ExprError("only + and - are allowed as signs")
1285
+ continue
1286
+ if isinstance(node, ast.Name):
1287
+ continue
1288
+ if isinstance(node, ast.Constant):
1289
+ if isinstance(node.value, bool) or not isinstance(node.value, (int, float)):
1290
+ raise ExprError("only numbers are allowed as constants")
1291
+ continue
1292
+ if isinstance(node, ast.Call):
1293
+ if node.keywords or any(isinstance(arg, ast.Starred) for arg in node.args):
1294
+ raise ExprError("function keyword arguments are not allowed")
1295
+ if not isinstance(node.func, ast.Name):
1296
+ raise ExprError("only plain function calls are allowed")
1297
+ fname = node.func.id
1298
+ if fname not in func_names:
1299
+ raise ExprError(
1300
+ f"unknown function {fname!r}: pass it as "
1301
+ f"geom_function(..., {fname}={fname})"
1302
+ )
1303
+ if fname == "where" and len(node.args) != 3:
1304
+ raise ExprError(
1305
+ "where() needs 3 arguments: where(x < 0, 0, x^2)"
1306
+ )
1307
+ continue
1308
+ if isinstance(node, ast.Compare):
1309
+ if len(node.ops) != 1 or not isinstance(node.ops[0], _CMPOPS):
1310
+ raise ExprError("only one comparison is allowed")
1311
+ continue
1312
+ if isinstance(node, _CMPOPS):
1313
+ continue
1314
+ raise ExprError(f"not allowed in a formula: {type(node).__name__}")
1315
+
1316
+
1317
+ def _parse_side(source: str, original: str) -> ast.AST:
1318
+ try:
1319
+ return ast.parse(source, mode="eval").body
1320
+ except SyntaxError as exc:
1321
+ column = exc.offset or 1
1322
+ token = ""
1323
+ if exc.text and exc.offset:
1324
+ token = exc.text[exc.offset - 1 : exc.offset]
1325
+ if not token:
1326
+ token = original.strip()[-1:] or "?"
1327
+ raise ExprError(
1328
+ f"syntax error at {token!r} (column {column})"
1329
+ ) from exc
1330
+
1331
+
1332
+ def _normalize(
1333
+ source: str, func_names: set[str], value_names: set[str]
1334
+ ) -> tuple[str | None, str, str]:
1335
+ """Return ``(lhs, rhs, label)`` with implicit ``*`` and ``^`` → ``**``.
1336
+
1337
+ ``lhs`` is ``None`` when the formula has no ``=``.
1338
+ """
1339
+ try:
1340
+ raw = list(tokenize.generate_tokens(io.StringIO(source).readline))
1341
+ except tokenize.TokenizeError as exc:
1342
+ raise ExprError(f"syntax error at {exc}") from exc
1343
+
1344
+ tokens = []
1345
+ for tok in raw:
1346
+ if tok.type in (
1347
+ tokenize.ENCODING,
1348
+ tokenize.ENDMARKER,
1349
+ tokenize.NEWLINE,
1350
+ tokenize.NL,
1351
+ tokenize.COMMENT,
1352
+ tokenize.INDENT,
1353
+ tokenize.DEDENT,
1354
+ ):
1355
+ continue
1356
+ if tok.type == tokenize.ERRORTOKEN:
1357
+ raise ExprError(
1358
+ f"syntax error at {tok.string!r} (column {tok.start[1] + 1})"
1359
+ )
1360
+ tokens.append(tok)
1361
+ if not tokens:
1362
+ raise ExprError(
1363
+ 'geom_function() needs a formula, for example geom_function("y = 2x")'
1364
+ )
1365
+ # ``{ }`` groups like parentheses, and ``x_{0}`` is the name x_0.
1366
+ # ``x^{2}`` then works in plain text as well as in LaTeX.
1367
+ tokens = _fold_groups(tokens)
1368
+
1369
+ last = tokens[-1]
1370
+ if last.type == tokenize.OP and last.string not in {")", "}"}:
1371
+ raise ExprError(
1372
+ f"syntax error at {last.string!r} (column {last.start[1] + 1})"
1373
+ )
1374
+
1375
+ depth = 0
1376
+ eq_at: int | None = None
1377
+ for index, tok in enumerate(tokens):
1378
+ if tok.string == "(":
1379
+ depth += 1
1380
+ elif tok.string == ")":
1381
+ depth -= 1
1382
+ if depth < 0:
1383
+ raise ExprError(
1384
+ f"syntax error at ')' (column {tok.start[1] + 1})"
1385
+ )
1386
+ elif tok.string == "=" and depth == 0:
1387
+ if eq_at is not None:
1388
+ raise ExprError(
1389
+ f"syntax error at '=' (column {tok.start[1] + 1})"
1390
+ )
1391
+ eq_at = index
1392
+ if depth != 0:
1393
+ raise ExprError("syntax error at '(' (column 1)")
1394
+
1395
+ if eq_at is not None:
1396
+ lhs_tokens = tokens[:eq_at]
1397
+ rhs_tokens = tokens[eq_at + 1 :]
1398
+ if not lhs_tokens or not rhs_tokens:
1399
+ bad = tokens[eq_at]
1400
+ raise ExprError(
1401
+ f"syntax error at '=' (column {bad.start[1] + 1})"
1402
+ )
1403
+ return (
1404
+ _render(lhs_tokens, func_names, value_names),
1405
+ _render(rhs_tokens, func_names, value_names),
1406
+ source,
1407
+ )
1408
+ return None, _render(tokens, func_names, value_names), source
1409
+
1410
+
1411
+ def _fold_groups(tokens: list[tokenize.TokenInfo]) -> list[tokenize.TokenInfo]:
1412
+ """Turn plain ``{ }`` into grouping parentheses. ``x_{0}`` stays one name."""
1413
+ out: list[tokenize.TokenInfo] = []
1414
+ index = 0
1415
+ while index < len(tokens):
1416
+ tok = tokens[index]
1417
+ # Python tokenizes ``x_`` as one name, so ``x_{0}`` is NAME ``x_`` then ``{0}``.
1418
+ name_then_brace = (
1419
+ tok.string == "{"
1420
+ and out
1421
+ and out[-1].type == tokenize.NAME
1422
+ and out[-1].string.endswith("_")
1423
+ )
1424
+ op_then_brace = (
1425
+ tok.string == "_"
1426
+ and out
1427
+ and out[-1].type == tokenize.NAME
1428
+ and index + 1 < len(tokens)
1429
+ and tokens[index + 1].string == "{"
1430
+ )
1431
+ if name_then_brace or op_then_brace:
1432
+ brace_at = index if name_then_brace else index + 1
1433
+ inner, nxt = _brace_inner(tokens, brace_at)
1434
+ if _plain_subscript(inner):
1435
+ tail = "".join(part.string for part in inner)
1436
+ prefix = out[-1].string if name_then_brace else out[-1].string + "_"
1437
+ out[-1] = out[-1]._replace(string=prefix + tail)
1438
+ index = nxt
1439
+ continue
1440
+ if tok.string == "{":
1441
+ out.append(tok._replace(string="("))
1442
+ elif tok.string == "}":
1443
+ out.append(tok._replace(string=")"))
1444
+ else:
1445
+ out.append(tok)
1446
+ index += 1
1447
+ return out
1448
+
1449
+
1450
+ def _brace_inner(
1451
+ tokens: list[tokenize.TokenInfo], start: int
1452
+ ) -> tuple[list[tokenize.TokenInfo], int]:
1453
+ """``tokens[start]`` is ``{``. Return the inside and the index after ``}``."""
1454
+ depth = 0
1455
+ inner: list[tokenize.TokenInfo] = []
1456
+ for index in range(start, len(tokens)):
1457
+ tok = tokens[index]
1458
+ if tok.string == "{":
1459
+ depth += 1
1460
+ if depth > 1:
1461
+ inner.append(tok)
1462
+ continue
1463
+ if tok.string == "}":
1464
+ depth -= 1
1465
+ if depth == 0:
1466
+ return inner, index + 1
1467
+ inner.append(tok)
1468
+ continue
1469
+ inner.append(tok)
1470
+ raise ExprError("syntax error at '{' (column 1)")
1471
+
1472
+
1473
+ def _plain_subscript(tokens: list[tokenize.TokenInfo]) -> bool:
1474
+ if not tokens:
1475
+ return False
1476
+ for tok in tokens:
1477
+ if tok.type not in {tokenize.NAME, tokenize.NUMBER} and tok.string != "_":
1478
+ return False
1479
+ text = "".join(tok.string for tok in tokens)
1480
+ return bool(text) and all(char.isalnum() or char == "_" for char in text)
1481
+
1482
+
1483
+ def _star_before_paren(name: str, func_names: set[str], value_names: set[str]) -> bool:
1484
+ """True when ``name(`` is a product, not a function call.
1485
+
1486
+ Known functions stay calls. An unknown word (``foo(x)``) stays a call so
1487
+ the validator can say it is an unknown function. A coefficient or plot
1488
+ variable (``a(x+1)``, ``x(y+1)``) is multiplication.
1489
+ """
1490
+ if name in func_names:
1491
+ return False
1492
+ if len(name) == 1 or name in value_names:
1493
+ return True
1494
+ return False
1495
+
1496
+
1497
+ def _render(
1498
+ tokens: list[tokenize.TokenInfo],
1499
+ func_names: set[str],
1500
+ value_names: set[str],
1501
+ ) -> str:
1502
+ """Join tokens, inserting ``*`` where math writes juxtaposition."""
1503
+ parts: list[str] = []
1504
+ prev_end = False
1505
+ prev_name: str | None = None
1506
+ for tok in tokens:
1507
+ is_name = tok.type == tokenize.NAME
1508
+ is_num = tok.type == tokenize.NUMBER
1509
+ is_lpar = tok.string == "("
1510
+ starts_value = is_name or is_num or is_lpar
1511
+ if prev_end and starts_value:
1512
+ keep_call = (
1513
+ prev_name is not None
1514
+ and is_lpar
1515
+ and not _star_before_paren(prev_name, func_names, value_names)
1516
+ )
1517
+ if not keep_call:
1518
+ parts.append("*")
1519
+ parts.append("**" if tok.string == "^" else tok.string)
1520
+ prev_end = is_name or is_num or tok.string == ")"
1521
+ prev_name = tok.string if is_name else None
1522
+ return "".join(parts)
1523
+
1524
+
1525
+ def _wrap_user_callable(name: str, fn: Callable[..., Any]) -> Callable[..., Any]:
1526
+ """Call ``fn`` on arrays; fall back to ``np.vectorize`` once if it cannot."""
1527
+ state: dict[str, Any] = {"warned": False, "vectorized": None}
1528
+
1529
+ def vectorized(*args: Any) -> Any:
1530
+ import warnings
1531
+
1532
+ if not state["warned"]:
1533
+ warnings.warn(
1534
+ f"geom_function: {name}() does not accept arrays; "
1535
+ "using np.vectorize (slower)",
1536
+ UserWarning,
1537
+ stacklevel=4,
1538
+ )
1539
+ state["warned"] = True
1540
+ if state["vectorized"] is None:
1541
+ state["vectorized"] = np.vectorize(fn, otypes=[np.float64])
1542
+ return state["vectorized"](*args)
1543
+
1544
+ def wrapped(*args: Any) -> Any:
1545
+ if state["vectorized"] is not None:
1546
+ return state["vectorized"](*args)
1547
+ try:
1548
+ out = fn(*args)
1549
+ except Exception:
1550
+ return vectorized(*args)
1551
+ if args and np.ndim(args[0]) > 0:
1552
+ arr = np.asarray(out)
1553
+ if getattr(arr, "shape", None) != np.shape(args[0]):
1554
+ return vectorized(*args)
1555
+ return out
1556
+
1557
+ return wrapped