From eb59c73fa79b6809d333ab6f58b47fe8a8a7fa47 Mon Sep 17 00:00:00 2001 From: Claude Date: Sat, 19 Sep 2026 18:13:01 +0000 Subject: [PATCH 1/3] feat(language): a where may compare arithmetic over parameters on both lanes MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit math-spec#566 lets either side of a where comparison be arithmetic over parameters — `p_max > cost`, `0.5 * p_max > 0`, `sum(p_max, over=generator) > 0`. Both lanes now build it, so what the language accepts is what this package builds. A side is compiled the way a constant side of a constraint is: the relational lane folds its constant fragments over the comparison's dims and joins the value column onto the carrier, the linopy lane evaluates it to a DataArray. A coordinate no fragment reaches stays null and the comparison reads false, which is what every other atom over a missing value does. Two architectural decisions to review rather than take on trust: Hard rule 0 moves. A mask now reads an expression, and a cases expression already read a mask, so the recursion is the language's own grammar and no ordering of a lane's two modules removes it. DELIBERATE_LAZY_IMPORTS was empty and that emptiness was the claim; it now holds two entries, one per lane, and architecture.md says so. The alternative was threading an evaluator callback through masked() at seven call sites in two modules, which buys a parameter and a detour for every future reader. NEVER_LOWERED excuses ArithmeticComparisonNode from both dispatch guards. It is in program.WhereNode but lower_program rewrites every one into an ExpressionComparisonNode, so no lowered mask holds one and neither lane can dispatch on it. The exclusion is a claim about upstream, so test_no_never_lowered_node_survives_lowering checks it rather than trusting it, and the two assert_never pragmas name the same reason. Whether the union should carry a node lowering always replaces is a question for #566. Coverage moved: 'p_max > cost' leaves the refusal sweep, where it asserted "compares two parameters", and joins ACCEPTED with the two arithmetic cases, so all three are swept for lane agreement on rows and status. pytest -q -n auto: 4023 passed, 251 skipped, 1 xfailed. ruff check, ruff format and pyrefly clean, the last with 6 suppressed rather than 4. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01NHXujoUG77G2SoDbhau5mS --- docs/about/architecture.md | 12 ++-- src/lpspec/linopy/where.py | 8 ++- .../relational/engines/polars/predicates.py | 32 ++++++++- tests/test_architecture.py | 65 +++++++++++++++++-- tests/test_resolution_parity.py | 14 +++- 5 files changed, 118 insertions(+), 13 deletions(-) diff --git a/docs/about/architecture.md b/docs/about/architecture.md index eb90e5e7e..cfe1da46c 100644 --- a/docs/about/architecture.md +++ b/docs/about/architecture.md @@ -269,10 +269,14 @@ rule. A cost phrased as a rule makes one implementation's choice load-bearing in the language's rulebook. 0. **The layers are ordered, and imports prove it.** Every module imports only - downward, at module level, with **no exception at all**. - `DELIBERATE_LAZY_IMPORTS` in `tests/test_architecture.py` is empty, and an - undeclared in-function import fails the build. A lazy import here is a cycle - to remove, not to defer. + downward, at module level, and every exception is declared in + `DELIBERATE_LAZY_IMPORTS` in `tests/test_architecture.py` with the cycle it + breaks. An undeclared in-function import fails the build, and a declared one + that disappears fails it too. There are two, one per lane, and they are one + cycle: a `where` may compare arithmetic over parameters, so a mask reads an + expression, while a `cases` expression reads a mask. That recursion is the + language's own grammar, so no ordering of a lane's two modules removes it. + Every other lazy import is a cycle to remove, not to defer. 1. **Core AST is the whole language, and the language is upstream.** Both lanes consume only core AST. Macros and `piecewise:` are expanded away before dispatch, and so is a named expression unless it states `cases:`. That one diff --git a/src/lpspec/linopy/where.py b/src/lpspec/linopy/where.py index 2243cc7d8..6be884d3e 100644 --- a/src/lpspec/linopy/where.py +++ b/src/lpspec/linopy/where.py @@ -99,6 +99,12 @@ def evaluate(child: program.WhereNode) -> xr.DataArray: if isinstance(node, program.VariableDefinedNode): return absence.present(ctx.model, node.name) + if isinstance(node, program.ExpressionComparisonNode): + from lpspec.linopy.builder import _eval # mask ↔ expression recursion + + left, right = _eval(node.left, ctx), _eval(node.right, ctx) + return _PREDICATE_OPS[node.op](left, right).fillna(value=False).astype(bool) + if isinstance(node, (program.ParameterComparisonNode, program.DimensionComparisonNode)): if isinstance(node, program.ParameterComparisonNode): arr = dataset[node.name] @@ -147,7 +153,7 @@ def evaluate(child: program.WhereNode) -> xr.DataArray: if isinstance(node, program.OrNode): return evaluate(node.left) | evaluate(node.right) - assert_never(node) + assert_never(node) # pyrefly: ignore[bad-argument-type] — ArithmeticComparisonNode is in the union and lowering always replaces it (NEVER_LOWERED) def _defined(arr: xr.DataArray, dtype: str) -> xr.DataArray: diff --git a/src/lpspec/relational/engines/polars/predicates.py b/src/lpspec/relational/engines/polars/predicates.py index 1a2e30ecf..0ae5cde3f 100644 --- a/src/lpspec/relational/engines/polars/predicates.py +++ b/src/lpspec/relational/engines/polars/predicates.py @@ -6,7 +6,8 @@ built first and the frame read after. A closed vocabulary of its own — comparisons against a parameter, a dimension -label, a position along a dimension, a relation, and the three connectives. It +label, a position along a dimension, a relation, arithmetic over parameters, +and the three connectives. It takes the :class:`~lpspec.relational.engines.polars.scope.Scope` as an argument and holds nothing. :func:`masked` is the product a declaration is instantiated over, cut by its mask: the one place the two meet. @@ -24,6 +25,7 @@ from math_spec import program from lpspec.errors import DataError, position_out_of_range_message, short_groups_message +from lpspec.relational.engines.polars.fragments import join_on from lpspec.relational.engines.polars.relations import GROUP_RANK, GROUP_SIZE, Grouping if TYPE_CHECKING: @@ -204,7 +206,33 @@ def join_relation(relation: str, dims: tuple[str, ...], column: str | None) -> s ), ) + def _side(expression: program.ExpressionNode, dims: tuple[str, ...]) -> str: + """*expression* joined onto the carrier as one value column, and that column's name. + + A side of a comparison is arithmetic over parameters — the language + refuses a variable and a ``dual()`` there — so it compiles to constant + fragments alone, and they are added over the comparison's dims the way + a constant side of a constraint is. A coordinate no fragment reaches + stays null so the comparison reads false, which is what every other + atom over a missing value does. + """ + from lpspec.relational.engines.polars.compiler import PolarsCompiler # mask ↔ expression recursion + + def attach(frame: pl.LazyFrame, alias: str) -> pl.LazyFrame: + compiled = PolarsCompiler(scope).expression(expression, 'a where clause') + assert not compiled.terms and not compiled.quads, ( + 'a where side is arithmetic over parameters, which compiles to constants alone' + ) + product = frame.select(*dims).unique() if dims else frame.select(pl.lit(0).alias('__one__')).head(1) + added = PolarsCompiler(scope).added(compiled.consts, product, fill=False) + return join_on(frame, added.rename({'cval': alias}), dims, 'left') + + return carrier.once(f'__where expression {expression!r}__', attach) + def walk(p: program.WhereNode) -> pl.Expr: + if isinstance(p, program.ExpressionComparisonNode): + left, right = (_side(p.left, p.dims), _side(p.right, p.dims)) + return falsy_if_null(_COLUMN_COMPARISONS[p.op](pl.col(left), pl.col(right))) if isinstance(p, program.ParameterComparisonNode): return _compare(pl.col(join_param(p.name)), p.op, p.value) if isinstance(p, program.DimensionComparisonNode): @@ -249,7 +277,7 @@ def walk(p: program.WhereNode) -> pl.Expr: return walk(p.left) | walk(p.right) if isinstance(p, program.NotNode): return ~falsy_if_null(walk(p.operand)) - assert_never(p) + assert_never(p) # pyrefly: ignore[bad-argument-type] — ArithmeticComparisonNode is in the union and lowering always replaces it (NEVER_LOWERED) condition = walk(mask.root) return carrier.frame, condition diff --git a/tests/test_architecture.py b/tests/test_architecture.py index 85db9d3a7..ea1c81b46 100644 --- a/tests/test_architecture.py +++ b/tests/test_architecture.py @@ -696,6 +696,46 @@ def test_no_sink_reaches_a_sibling(): ) +#: Nodes in the language's unions that ``lower_program`` always replaces, so +#: no lowered plan holds one and neither lane can dispatch on it. +#: ``ArithmeticComparisonNode`` is the resolved form of a where comparison +#: between expressions; lowering rewrites every one into an +#: ``ExpressionComparisonNode`` over program expressions. +NEVER_LOWERED = frozenset({'ArithmeticComparisonNode'}) + + +def test_no_never_lowered_node_survives_lowering(): + """``NEVER_LOWERED`` is a claim about upstream, so it is checked rather than trusted. + + An entry that lowering stopped replacing would silently excuse both lanes + from a node they now have to build. + """ + import dataclasses + + from math_spec import to_program + + from tests.fixtures import DISPATCH_SPEC, override + + seen: set[str] = set() + + def walk(node): + if node is None: + return + seen.add(type(node).__name__) + for field in dataclasses.fields(node): + child = getattr(node, field.name) + if dataclasses.is_dataclass(child): + walk(child) + + for where in ('p_max > cost', '0.5 * p_max > 0', 'NOT p_max > cost'): + mask = to_program(override(DISPATCH_SPEC, **{'variables.p.where': where})).variables['p'].where + walk(None if mask is None else mask.root) + assert seen & NEVER_LOWERED == set(), ( + f'{sorted(seen & NEVER_LOWERED)} reached a lowered mask — it is no longer excused from either lane' + ) + assert 'ExpressionComparisonNode' in seen, 'the probe must reach the node lowering replaces them with' + + def test_every_plan_node_is_handled_by_the_compiler(): """Two-tier economy: a primitive is not done until the engine consumes it. @@ -724,7 +764,11 @@ def test_every_plan_node_is_handled_by_the_compiler(): ] for qualifier, union, module in walkers: source = module.read_text() - unhandled = [c.__name__ for c in get_args(union) if f'{qualifier}.{c.__name__}' not in source] + unhandled = [ + c.__name__ + for c in get_args(union) + if c.__name__ not in NEVER_LOWERED and f'{qualifier}.{c.__name__}' not in source + ] assert not unhandled, f'{qualifier} nodes unknown to {module.name}: {unhandled}' @@ -991,7 +1035,9 @@ class and cannot collide. from math_spec import program - declared = {node.__name__ for union in (program.ExpressionNode, program.WhereNode) for node in get_args(union)} + declared = { + node.__name__ for union in (program.ExpressionNode, program.WhereNode) for node in get_args(union) + } - NEVER_LOWERED assert declared, 'no plan node classes found — the census has nothing to run over' def dispatched_on(*paths: Path) -> set[str]: @@ -1056,9 +1102,18 @@ def test_every_module_is_documented_somewhere(): #: Every in-function ``lpspec`` import in the package, with the cycle it -#: breaks. Empty, and that is the claim: the layers are ordered with no -#: exception at all, so a lazy import is only ever a leftover. -DELIBERATE_LAZY_IMPORTS: dict[tuple[str, str], str] = {} +#: breaks. The two here are one cycle, once per lane: a ``where`` may compare +#: arithmetic over parameters, so a mask reads an expression, and a ``cases`` +#: expression reads a mask. The recursion is the language's own grammar, so no +#: ordering of these two modules removes it and each lane's mask walk reaches +#: its expression evaluator inside the one handler that needs it. +DELIBERATE_LAZY_IMPORTS: dict[tuple[str, str], str] = { + ('linopy/where.py', 'lpspec.linopy.builder'): 'a where side is an expression, and a cases expression holds a mask', + ( + 'relational/engines/polars/predicates.py', + 'lpspec.relational.engines.polars.compiler', + ): 'a where side is an expression, and a cases expression holds a mask', +} def test_lazy_intra_package_imports_are_all_declared(): diff --git a/tests/test_resolution_parity.py b/tests/test_resolution_parity.py index ae44f4db4..f8cd953e2 100644 --- a/tests/test_resolution_parity.py +++ b/tests/test_resolution_parity.py @@ -25,7 +25,6 @@ ('where', 'match'), [ pytest.param('typo_name > 0', "'typo_name' not found", id='a-name-nothing-declares'), - pytest.param('p_max > cost', 'compares two parameters', id='two-parameters-compared'), pytest.param('generator == snapshot', 'compares against dimension', id='a-dimension-on-the-right'), pytest.param('nonexistent', "'nonexistent' not found", id='a-bare-name-nothing-declares'), pytest.param('snapshot', 'bare dimension name is true at every coordinate', id='a-bare-dimension-name'), @@ -79,6 +78,14 @@ def test_both_lanes_refuse_a_comparison_that_carries_no_variable(tmp_path, dispa #: The one position a literal survives to: alone, and false. `True` alone #: is no mask at all and arrives as `None`. 'False', + #: Two parameters compared, which math-spec 0.0.0-alpha.105 still refused + #: with "compares two parameters, which is not in the language". + 'p_max > cost', + #: Arithmetic on a side, and a reduction on one — both sides of a + #: comparison are expressions now, so a side carries what a constant side + #: of a constraint carries. + '0.5 * p_max > 0', + 'sum(p_max, over=generator) > 0', ] #: Predicates this sweep cannot host, with where they are checked instead. The @@ -96,6 +103,11 @@ def test_both_lanes_refuse_a_comparison_that_carries_no_variable(tmp_path, dispa 'RelationComparisonNode': 'tests/test_label_coords.py::test_a_where_reads_a_relation', 'RelationPairComparisonNode': 'tests/test_label_coords.py::test_a_relation_where_agrees_with_the_oracle', 'RelationDefinedNode': 'tests/test_label_coords.py::test_a_where_reads_a_relation', + #: The resolved form of an arithmetic comparison, which `lower_program` + #: replaces with `ExpressionComparisonNode` over program expressions. It is + #: in `program.WhereNode` but reaches no lowered mask, so no lane dispatches + #: on it and the sweep above cannot host it. + 'ArithmeticComparisonNode': 'math-spec tests/test_lowering.py — never lowered into a program mask', } From 3714f5bf92ae51b99c94f729083f46c109a18172 Mon Sep 17 00:00:00 2001 From: Claude Date: Sat, 19 Sep 2026 18:14:54 +0000 Subject: [PATCH 2/3] test(engine): a where side holding a variable is caught rather than compiled as a constant The constants-only guard in _side survived deletion with the whole suite green, because lower_program refuses a variable on a where side before the engine is reached and no model can carry one there. The probe hands compile_predicate the node directly, so the guard is exercised by the one caller that can reach it. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01NHXujoUG77G2SoDbhau5mS --- tests/test_compiler.py | 22 ++++++++++++++++++++++ 1 file changed, 22 insertions(+) diff --git a/tests/test_compiler.py b/tests/test_compiler.py index 8b4a983c9..002b985f1 100644 --- a/tests/test_compiler.py +++ b/tests/test_compiler.py @@ -491,3 +491,25 @@ def test_a_zero_edge_writes_its_rows_like_any_other_fill(): bare = program.Translate(program.Parameter('load'), 'snapshot', 1, wrap=False, fill=None) vacated = q.expression(bare, 'test').consts[0].frame.collect() assert vacated['snapshot'].to_list() == [1, 2], 'a bare shift vacates, and that is a gap on purpose' + + +def test_a_where_side_holding_a_variable_is_caught_rather_than_compiled_as_a_constant(): + """The guard the suite survives deleting, because the language gets there first. + + `lower_program` only builds an `ExpressionComparisonNode` from a side it + has already refused a variable and a `dual()` on, so no model reaches + `_side` with a variable term — the guard is the invariant that stays true + if that refusal ever moves. Handed the node directly, it says which + assumption broke instead of folding a variable's coefficient into a value + column and masking on it. + """ + from lpspec.relational.engines.polars.predicates import compile_predicate + + mask = program.Mask( + program.ExpressionComparisonNode(program.Variable('p'), '>', program.Constant(0.0), ('snapshot', 'generator')) + ) + frame = pl.LazyFrame(schema={'snapshot': pl.Int64, 'generator': pl.String}) + + with pytest.raises(AssertionError, match='arithmetic over parameters, which compiles to constants alone'): + carrier, _ = compile_predicate(compiler().scope, frame, mask, ('snapshot', 'generator')) + carrier.frame.collect() From f23d81486e3bdecbcc1a51d58221c4f9bfc4ef41 Mon Sep 17 00:00:00 2001 From: Claude Date: Sat, 19 Sep 2026 19:15:01 +0000 Subject: [PATCH 3/3] fix(language): a where comparing two constants builds on both lanes rather than crashing the eager one `where: "1 > 0"` reached the eager lane as a Python bool, because both sides evaluate to scalars and the comparison of two scalars carries no dimension. The handler called `.fillna` on it and raised AttributeError, while the relational lane built the model. One language, two answers. The comparison is wrapped as a DataArray before the null fill, which is the 0-dimensional mask `evaluate_where` already returns for the no-mask case, so callers keep combining with `&` and `|` without case analysis. Found by sweeping the expression language across a where side rather than by a report. '1 > 0' joins ACCEPTED, so the lanes are held to agreeing on it the way they are on every other predicate; it fails there without the fix. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01NHXujoUG77G2SoDbhau5mS --- src/lpspec/linopy/where.py | 3 ++- tests/test_resolution_parity.py | 4 ++++ 2 files changed, 6 insertions(+), 1 deletion(-) diff --git a/src/lpspec/linopy/where.py b/src/lpspec/linopy/where.py index 6be884d3e..55cb31f20 100644 --- a/src/lpspec/linopy/where.py +++ b/src/lpspec/linopy/where.py @@ -103,7 +103,8 @@ def evaluate(child: program.WhereNode) -> xr.DataArray: from lpspec.linopy.builder import _eval # mask ↔ expression recursion left, right = _eval(node.left, ctx), _eval(node.right, ctx) - return _PREDICATE_OPS[node.op](left, right).fillna(value=False).astype(bool) + compared = xr.DataArray(_PREDICATE_OPS[node.op](left, right)) + return compared.fillna(value=False).astype(bool) if isinstance(node, (program.ParameterComparisonNode, program.DimensionComparisonNode)): if isinstance(node, program.ParameterComparisonNode): diff --git a/tests/test_resolution_parity.py b/tests/test_resolution_parity.py index f8cd953e2..e6d15a985 100644 --- a/tests/test_resolution_parity.py +++ b/tests/test_resolution_parity.py @@ -86,6 +86,10 @@ def test_both_lanes_refuse_a_comparison_that_carries_no_variable(tmp_path, dispa #: of a constraint carries. '0.5 * p_max > 0', 'sum(p_max, over=generator) > 0', + #: Both sides constant, so the comparison is one scalar against another and + #: carries no dimension. The eager lane read the result as an array and got + #: a plain bool, which has no `fillna`. + '1 > 0', ] #: Predicates this sweep cannot host, with where they are checked instead. The