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..55cb31f20 100644 --- a/src/lpspec/linopy/where.py +++ b/src/lpspec/linopy/where.py @@ -99,6 +99,13 @@ 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) + 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): arr = dataset[node.name] @@ -147,7 +154,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_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() diff --git a/tests/test_resolution_parity.py b/tests/test_resolution_parity.py index ae44f4db4..e6d15a985 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,18 @@ 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', + #: 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 @@ -96,6 +107,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', }