rootfig 0.1.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.
@@ -0,0 +1,462 @@
1
+ """Parsing, validation and evaluation of expression strings."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import ast
6
+ import difflib
7
+ import re
8
+ from collections.abc import Collection, Iterable, Mapping
9
+ from dataclasses import dataclass, field
10
+ from types import CodeType
11
+ from typing import Any, Final
12
+
13
+ import awkward as ak
14
+ import numpy as np
15
+
16
+ from rootfig.errors import ExpressionError, MissingBranchError
17
+ from rootfig.expressions.functions import CONSTANTS, FUNCTIONS
18
+
19
+ __all__ = ["Expression", "ExpressionLike", "evaluate", "parse"]
20
+
21
+
22
+ # --------------------------------------------------------------------------------------
23
+ # Parsing and validation
24
+ # --------------------------------------------------------------------------------------
25
+
26
+ _FUNCTION_PREFIX: Final = "__rootfig_fn_"
27
+ _BACKTICK_PREFIX: Final = "__rootfig_bt_"
28
+ _NOT_NAME: Final = "__rootfig_not"
29
+ _BACKTICK_RE: Final = re.compile(r"`([^`]*)`")
30
+
31
+ _ALLOWED_BINOPS: Final[dict[type[ast.operator], str]] = {
32
+ ast.Add: "+",
33
+ ast.Sub: "-",
34
+ ast.Mult: "*",
35
+ ast.Div: "/",
36
+ ast.FloorDiv: "//",
37
+ ast.Mod: "%",
38
+ ast.Pow: "**",
39
+ ast.BitAnd: "&",
40
+ ast.BitOr: "|",
41
+ ast.BitXor: "^",
42
+ }
43
+ _ALLOWED_UNARYOPS: Final[tuple[type[ast.unaryop], ...]] = (ast.USub, ast.UAdd, ast.Invert)
44
+ _ALLOWED_CMPOPS: Final[tuple[type[ast.cmpop], ...]] = (
45
+ ast.Eq,
46
+ ast.NotEq,
47
+ ast.Lt,
48
+ ast.LtE,
49
+ ast.Gt,
50
+ ast.GtE,
51
+ )
52
+
53
+
54
+ def _mangle_backticks(text: str) -> tuple[str, dict[str, str]]:
55
+ """Replace backtick-quoted names with valid identifiers; return the reverse mapping."""
56
+ mapping: dict[str, str] = {}
57
+
58
+ def repl(match: re.Match[str]) -> str:
59
+ original = match.group(1)
60
+ if not original:
61
+ msg = "empty backticks `` in expression"
62
+ raise ExpressionError(msg)
63
+ for mangled, name in mapping.items():
64
+ if name == original:
65
+ return mangled
66
+ mangled = f"{_BACKTICK_PREFIX}{len(mapping)}"
67
+ mapping[mangled] = original
68
+ return mangled
69
+
70
+ return _BACKTICK_RE.sub(repl, text), mapping
71
+
72
+
73
+ class _Rewriter(ast.NodeTransformer):
74
+ """Validate the AST and rewrite Python-only constructs to element-wise ones.
75
+
76
+ Collects the referenced names into ``self.names`` (in order of first
77
+ appearance) and the called functions into ``self.functions``.
78
+ """
79
+
80
+ def __init__(self, text: str, backticks: dict[str, str]) -> None:
81
+ self.text = text
82
+ self.backticks = backticks
83
+ self.names: list[str] = []
84
+ self.functions: list[str] = []
85
+
86
+ # -- helpers -------------------------------------------------------------------
87
+
88
+ def _fail(self, node: ast.AST, what: str) -> ExpressionError:
89
+ segment = ast.get_source_segment(self.text, node) or type(node).__name__
90
+ return ExpressionError(f"{what} is not allowed in expressions: {segment!r}")
91
+
92
+ def _display_name(self, ident: str) -> str:
93
+ return self.backticks.get(ident, ident)
94
+
95
+ # -- allowed nodes ---------------------------------------------------------------
96
+
97
+ def visit_Expression(self, node: ast.Expression) -> ast.AST:
98
+ node.body = self.visit(node.body)
99
+ return node
100
+
101
+ def visit_Name(self, node: ast.Name) -> ast.AST:
102
+ if not isinstance(node.ctx, ast.Load):
103
+ raise self._fail(node, "assignment")
104
+ display = self._display_name(node.id)
105
+ if display not in self.names:
106
+ self.names.append(display)
107
+ return node
108
+
109
+ def visit_Attribute(self, node: ast.Attribute) -> ast.AST:
110
+ """``Collection.field.sub`` is one dotted branch name, as in EDM4hep/podio files."""
111
+ parts: list[str] = []
112
+ base: ast.expr = node
113
+ while isinstance(base, ast.Attribute):
114
+ if not isinstance(base.ctx, ast.Load):
115
+ raise self._fail(node, "assignment")
116
+ parts.append(base.attr)
117
+ base = base.value
118
+ if not isinstance(base, ast.Name):
119
+ raise self._fail(node, "attribute access on this construct")
120
+ dotted = ".".join([self._display_name(base.id), *reversed(parts)])
121
+ mangled = self._mangle(dotted)
122
+ if dotted not in self.names:
123
+ self.names.append(dotted)
124
+ return ast.copy_location(ast.Name(id=mangled, ctx=ast.Load()), node)
125
+
126
+ def _mangle(self, name: str) -> str:
127
+ """Register ``name`` (not a valid identifier) and return its stand-in identifier."""
128
+ for mangled, original in self.backticks.items():
129
+ if original == name:
130
+ return mangled
131
+ mangled = f"{_BACKTICK_PREFIX}{len(self.backticks)}"
132
+ self.backticks[mangled] = name
133
+ return mangled
134
+
135
+ def visit_Constant(self, node: ast.Constant) -> ast.AST:
136
+ if isinstance(node.value, bool | int | float):
137
+ return node
138
+ raise self._fail(node, f"a {type(node.value).__name__} literal")
139
+
140
+ def visit_BinOp(self, node: ast.BinOp) -> ast.AST:
141
+ if type(node.op) not in _ALLOWED_BINOPS:
142
+ raise self._fail(node, f"operator {type(node.op).__name__}")
143
+ node.left = self.visit(node.left)
144
+ node.right = self.visit(node.right)
145
+ return node
146
+
147
+ def visit_UnaryOp(self, node: ast.UnaryOp) -> ast.AST:
148
+ operand = self.visit(node.operand)
149
+ if isinstance(node.op, ast.Not):
150
+ # ``not x`` as a logical negation: ``~`` would turn a Python ``True`` into ``-2``
151
+ # and an integer array into bit patterns instead of booleans.
152
+ func = ast.copy_location(ast.Name(id=_NOT_NAME, ctx=ast.Load()), node)
153
+ return ast.copy_location(ast.Call(func=func, args=[operand], keywords=[]), node)
154
+ if not isinstance(node.op, _ALLOWED_UNARYOPS):
155
+ raise self._fail(node, f"operator {type(node.op).__name__}")
156
+ node.operand = operand
157
+ return node
158
+
159
+ def visit_BoolOp(self, node: ast.BoolOp) -> ast.AST:
160
+ op: ast.operator = ast.BitAnd() if isinstance(node.op, ast.And) else ast.BitOr()
161
+ values = [self.visit(v) for v in node.values]
162
+ result: ast.expr = values[0]
163
+ for value in values[1:]:
164
+ result = ast.copy_location(ast.BinOp(left=result, op=op, right=value), node)
165
+ return result
166
+
167
+ def visit_Compare(self, node: ast.Compare) -> ast.AST:
168
+ for op in node.ops:
169
+ if not isinstance(op, _ALLOWED_CMPOPS):
170
+ raise self._fail(node, f"comparison {type(op).__name__}")
171
+ operands = [self.visit(node.left), *(self.visit(c) for c in node.comparators)]
172
+ comparisons = [
173
+ ast.copy_location(ast.Compare(left=lhs, ops=[op], comparators=[rhs]), node)
174
+ for lhs, op, rhs in zip(operands[:-1], node.ops, operands[1:], strict=True)
175
+ ]
176
+ result: ast.expr = comparisons[0]
177
+ for comparison in comparisons[1:]:
178
+ result = ast.copy_location(
179
+ ast.BinOp(left=result, op=ast.BitAnd(), right=comparison), node
180
+ )
181
+ return result
182
+
183
+ def visit_Call(self, node: ast.Call) -> ast.AST:
184
+ if not isinstance(node.func, ast.Name):
185
+ raise self._fail(node.func, "calling anything but a known function by name")
186
+ name = node.func.id
187
+ if name not in FUNCTIONS:
188
+ suggestions = difflib.get_close_matches(name, FUNCTIONS, n=3)
189
+ hint = f" Did you mean {', '.join(suggestions)}?" if suggestions else ""
190
+ msg = f"unknown function {name!r}.{hint} Available functions: " + ", ".join(
191
+ sorted(FUNCTIONS)
192
+ )
193
+ raise ExpressionError(msg)
194
+ if name not in self.functions:
195
+ self.functions.append(name)
196
+ node.func = ast.copy_location(
197
+ ast.Name(id=f"{_FUNCTION_PREFIX}{name}", ctx=ast.Load()), node
198
+ )
199
+ node.args = [self.visit(a) for a in node.args]
200
+ for keyword in node.keywords:
201
+ if keyword.arg is None:
202
+ raise self._fail(node, "** argument unpacking")
203
+ keyword.value = self.visit(keyword.value)
204
+ return node
205
+
206
+ def visit_Subscript(self, node: ast.Subscript) -> ast.AST:
207
+ node.value = self.visit(node.value)
208
+ node.slice = self.visit(node.slice)
209
+ return node
210
+
211
+ def visit_Slice(self, node: ast.Slice) -> ast.AST:
212
+ for attr in ("lower", "upper", "step"):
213
+ value = getattr(node, attr)
214
+ if value is not None:
215
+ setattr(node, attr, self.visit(value))
216
+ return node
217
+
218
+ def visit_Tuple(self, node: ast.Tuple) -> ast.AST:
219
+ node.elts = [self.visit(e) for e in node.elts]
220
+ return node
221
+
222
+ def visit_IfExp(self, node: ast.IfExp) -> ast.AST:
223
+ # ``a if cond else b`` is ambiguous for arrays; steer users to where().
224
+ raise self._fail(node, "the conditional expression (use where(cond, a, b))")
225
+
226
+ def generic_visit(self, node: ast.AST) -> ast.AST:
227
+ raise self._fail(node, f"the construct {type(node).__name__}")
228
+
229
+
230
+ # --------------------------------------------------------------------------------------
231
+ # Public objects
232
+ # --------------------------------------------------------------------------------------
233
+
234
+
235
+ @dataclass(frozen=True)
236
+ class Expression:
237
+ """A parsed, validated expression ready to be evaluated on arrays.
238
+
239
+ Create instances with :func:`parse` (or pass strings anywhere rootfig
240
+ accepts an expression; they are parsed on the fly).
241
+
242
+ Attributes
243
+ ----------
244
+ text
245
+ The original expression string.
246
+ names
247
+ Names referenced by the expression in order of first appearance. These
248
+ are candidate branch names; a name that is not a branch may still
249
+ resolve to a constant (``pi``, ``e``, ``inf``, ``nan``).
250
+ functions
251
+ Names of the functions called by the expression.
252
+ """
253
+
254
+ text: str
255
+ names: tuple[str, ...]
256
+ functions: tuple[str, ...]
257
+ _code: CodeType = field(repr=False, compare=False)
258
+ _backticks: Mapping[str, str] = field(repr=False, compare=False, default_factory=dict)
259
+
260
+ def __str__(self) -> str:
261
+ return self.text
262
+
263
+ @property
264
+ def is_trivial(self) -> bool:
265
+ """True if the expression is a bare name (a single branch, no computation)."""
266
+ return (
267
+ len(self.names) == 1
268
+ and not self.functions
269
+ and self.text.strip()
270
+ in (
271
+ self.names[0],
272
+ f"`{self.names[0]}`",
273
+ )
274
+ )
275
+
276
+ def required_branches(self, available: Collection[str]) -> list[str]:
277
+ """Return the referenced names that must be read from ``available`` branches.
278
+
279
+ Names that are not available but match a constant are skipped. Any
280
+ other unknown name raises :class:`~rootfig.errors.MissingBranchError`
281
+ with close-match suggestions.
282
+ """
283
+ required: list[str] = []
284
+ for name in self.names:
285
+ if name in available:
286
+ required.append(name)
287
+ elif name not in CONSTANTS:
288
+ suggestions = difflib.get_close_matches(name, list(available), n=3)
289
+ raise MissingBranchError(
290
+ name,
291
+ available=sorted(available),
292
+ suggestions=suggestions,
293
+ context=f"expression {self.text!r}",
294
+ )
295
+ return required
296
+
297
+ def evaluate(
298
+ self, arrays: Mapping[str, Any] | ak.Array, *, length: int | None = None
299
+ ) -> ak.Array:
300
+ """Evaluate the expression using ``arrays`` to resolve branch names.
301
+
302
+ Parameters
303
+ ----------
304
+ arrays
305
+ Either a mapping from branch name to array, or an Awkward record
306
+ array whose fields are the branches.
307
+ length
308
+ Number of events a constant expression (``"1"``, ``"True"``) is
309
+ broadcast to when ``arrays`` holds nothing to take the length from.
310
+
311
+ Returns
312
+ -------
313
+ awkward.Array
314
+ The result, converted to an Awkward array if the expression produced
315
+ a NumPy array or a scalar.
316
+ """
317
+ lookup = _as_mapping(arrays)
318
+ self.required_branches(list(lookup.keys()))
319
+ namespace: dict[str, Any] = {
320
+ f"{_FUNCTION_PREFIX}{name}": fn for name, fn in FUNCTIONS.items()
321
+ }
322
+ namespace[_NOT_NAME] = np.logical_not
323
+ for mangled, original in self._backticks.items():
324
+ namespace[mangled] = (
325
+ _bind(lookup[original]) if original in lookup else CONSTANTS[original]
326
+ )
327
+ # A name may occur both quoted and unquoted (``x + `x```); bind both spellings.
328
+ for name in self.names:
329
+ if name.isidentifier():
330
+ namespace[name] = _bind(lookup[name]) if name in lookup else CONSTANTS[name]
331
+ try:
332
+ result = eval(self._code, {"__builtins__": {}}, namespace) # validated AST
333
+ except ExpressionError:
334
+ raise
335
+ except Exception as exc:
336
+ msg = f"failed to evaluate expression {self.text!r}: {type(exc).__name__}: {exc}"
337
+ raise ExpressionError(msg) from exc
338
+ return _as_awkward(result, self.text, lookup, length)
339
+
340
+
341
+ ExpressionLike = str | Expression
342
+ """Anything accepted where an expression is expected."""
343
+
344
+
345
+ def parse(expression: ExpressionLike) -> Expression:
346
+ """Parse and validate an expression string.
347
+
348
+ Raises
349
+ ------
350
+ ExpressionError
351
+ If the text is not a single Python expression or uses a disallowed
352
+ construct (attribute access, unknown functions, string literals, ...).
353
+ """
354
+ if isinstance(expression, Expression):
355
+ return expression
356
+ if not isinstance(expression, str): # runtime guard for untyped callers
357
+ msg = f"expression must be a string, got {type(expression).__name__}" # type: ignore[unreachable]
358
+ raise ExpressionError(msg)
359
+ text = expression.strip()
360
+ if not text:
361
+ msg = "expression is empty"
362
+ raise ExpressionError(msg)
363
+ mangled, backticks = _mangle_backticks(text)
364
+ try:
365
+ tree = ast.parse(mangled, mode="eval")
366
+ except SyntaxError as exc:
367
+ msg = f"invalid syntax in expression {text!r}: {exc.msg}"
368
+ raise ExpressionError(msg) from exc
369
+ rewriter = _Rewriter(mangled, backticks)
370
+ tree = rewriter.visit(tree)
371
+ ast.fix_missing_locations(tree)
372
+ code = compile(tree, "<rootfig expression>", "eval")
373
+ return Expression(
374
+ text=text,
375
+ names=tuple(rewriter.names),
376
+ functions=tuple(rewriter.functions),
377
+ _code=code,
378
+ _backticks=backticks,
379
+ )
380
+
381
+
382
+ def evaluate(
383
+ expression: ExpressionLike, arrays: Mapping[str, Any] | ak.Array, *, length: int | None = None
384
+ ) -> ak.Array:
385
+ """Parse (if needed) and evaluate ``expression`` on ``arrays``.
386
+
387
+ This is a convenience wrapper around :func:`parse` and
388
+ :meth:`Expression.evaluate`; ``length`` sizes constant expressions when
389
+ ``arrays`` is empty.
390
+
391
+ Examples
392
+ --------
393
+ >>> import awkward as ak
394
+ >>> arrays = {"pt": ak.Array([[10.0, 30.0], [], [50.0]])}
395
+ >>> evaluate("pt > 20", arrays).tolist()
396
+ [[False, True], [], [True]]
397
+ >>> evaluate("count(pt)", arrays).tolist()
398
+ [2, 0, 1]
399
+ """
400
+ return parse(expression).evaluate(arrays, length=length)
401
+
402
+
403
+ # --------------------------------------------------------------------------------------
404
+ # Internal helpers
405
+ # --------------------------------------------------------------------------------------
406
+
407
+
408
+ def _as_mapping(arrays: object) -> Mapping[str, Any]:
409
+ if isinstance(arrays, Mapping):
410
+ return arrays
411
+ if isinstance(arrays, ak.Array):
412
+ if not arrays.fields:
413
+ msg = "expected a record array with named fields, got an array without fields"
414
+ raise ExpressionError(msg)
415
+ return {name: arrays[name] for name in arrays.fields}
416
+ msg = f"cannot interpret {type(arrays).__name__} as a mapping of branch arrays"
417
+ raise ExpressionError(msg)
418
+
419
+
420
+ def _bind(value: Any) -> Any:
421
+ """Present fixed-size (regular) dimensions as variable-length lists before evaluation.
422
+
423
+ Regular arrays broadcast like NumPy matrices, so ``x >= 50 and met > 10`` with a
424
+ ``float x[3]`` branch would fail (or silently align the wrong axis); as lists of
425
+ objects they follow rootfig's per-event/per-object rules.
426
+ """
427
+ if isinstance(value, np.ndarray) and value.ndim > 1:
428
+ value = ak.Array(value)
429
+ if isinstance(value, ak.Array) and ak.to_layout(value).purelist_depth > 1:
430
+ return ak.from_regular(value, axis=None)
431
+ return value
432
+
433
+
434
+ def _as_awkward(
435
+ result: Any, text: str, lookup: Mapping[str, Any], length: int | None = None
436
+ ) -> ak.Array:
437
+ if isinstance(result, ak.Array):
438
+ return result
439
+ if isinstance(result, np.ndarray):
440
+ return ak.Array(result)
441
+ if isinstance(result, bool | int | float | np.generic):
442
+ # Broadcast a scalar over events, e.g. weight="1.5" or selection="True".
443
+ if length is None:
444
+ length = _common_length(lookup.values())
445
+ if length is None:
446
+ msg = (
447
+ f"expression {text!r} is a constant and no branch arrays were available to "
448
+ "determine the number of events"
449
+ )
450
+ raise ExpressionError(msg)
451
+ return ak.Array(np.full(length, result))
452
+ msg = f"expression {text!r} produced an unsupported result of type {type(result).__name__}"
453
+ raise ExpressionError(msg)
454
+
455
+
456
+ def _common_length(arrays: Iterable[Any]) -> int | None:
457
+ for array in arrays:
458
+ try:
459
+ return len(array)
460
+ except TypeError:
461
+ continue
462
+ return None
@@ -0,0 +1,67 @@
1
+ """Histogram construction, normalisation, ratios and statistics."""
2
+
3
+ from rootfig.histograms.build import Histogram, as_weight_storage, fill
4
+ from rootfig.histograms.cutflow import Cutflow, CutflowStep, CutflowTable, cutflow
5
+ from rootfig.histograms.efficiency import Efficiency, Profile, ProfileStatistic, efficiency, profile
6
+ from rootfig.histograms.normalize import (
7
+ NormalizeSpec,
8
+ normalization_label,
9
+ normalize,
10
+ normalize_hist,
11
+ )
12
+ from rootfig.histograms.pipeline import (
13
+ build_histograms,
14
+ build_histograms_2d,
15
+ combined_selection,
16
+ combined_weight,
17
+ load_columns,
18
+ read_arrays,
19
+ source_length,
20
+ )
21
+ from rootfig.histograms.ratio import (
22
+ SIGNIFICANCE_KINDS,
23
+ Ratio,
24
+ RatioUncertainty,
25
+ SignificanceKind,
26
+ compatible_binning,
27
+ ratio,
28
+ significance,
29
+ )
30
+ from rootfig.histograms.stats import Summary, correlation_matrix, describe_table, summarize
31
+
32
+ __all__ = [
33
+ "SIGNIFICANCE_KINDS",
34
+ "Cutflow",
35
+ "CutflowStep",
36
+ "CutflowTable",
37
+ "Efficiency",
38
+ "Histogram",
39
+ "NormalizeSpec",
40
+ "Profile",
41
+ "ProfileStatistic",
42
+ "Ratio",
43
+ "RatioUncertainty",
44
+ "SignificanceKind",
45
+ "Summary",
46
+ "as_weight_storage",
47
+ "build_histograms",
48
+ "build_histograms_2d",
49
+ "combined_selection",
50
+ "combined_weight",
51
+ "compatible_binning",
52
+ "correlation_matrix",
53
+ "cutflow",
54
+ "describe_table",
55
+ "efficiency",
56
+ "fill",
57
+ "load_columns",
58
+ "normalization_label",
59
+ "normalize",
60
+ "normalize_hist",
61
+ "profile",
62
+ "ratio",
63
+ "read_arrays",
64
+ "significance",
65
+ "source_length",
66
+ "summarize",
67
+ ]