From b14777d5e866e122cfced0d6351bb103f1f4a22b Mon Sep 17 00:00:00 2001 From: Claude Date: Sun, 20 Sep 2026 22:54:11 +0000 Subject: [PATCH 1/2] feat(language): a where counts the coordinates a predicate admits, and reads one at a neighbour MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Closes #590 and #591. Two calls read a predicate where every other operator reads arithmetic: count(, over=) shift(, along=, offset=) `count` reduces the counted dimension away, so the count is one number per remaining coordinate and a claim about each group needs no word for the group. `shift` is false where the translation vacates and takes no `edge=`: the arithmetic form needs one because no number is neutral, and false is what a missing row already means in a mask. Both are where-only atoms rather than `BUILTINS` rows, as `position()` is. `count` is spelled in the grammar because only the grammar can decide to read its argument as a predicate; every other predicate-reading call stands where arithmetic cannot, so the comparison above it has already been tried. `CountComparison` and `TranslatedPredicate` join `Predicate` and carry their operand as a `Mask` — the wrapper a leaf holds a predicate in, where a connective holds a bare one. A walk recurses through the second and stops at the first, and `names_read` and `dims` see through both. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01DBU3ocfqkHe99mWmxijw64 --- docs/reference/language/expressions.md | 94 +++++++++++++++--- docs/reference/notation.md | 45 +++++++++ docs/reference/reading.md | 8 ++ src/math_spec/_where_parser.py | 85 ++++++++++++++-- src/math_spec/dimensions.py | 6 ++ src/math_spec/exclusivity.py | 21 ++++ src/math_spec/lowering.py | 12 ++- src/math_spec/program.py | 47 +++++++++ src/math_spec/resolution.py | 124 +++++++++++++++++++++++ src/math_spec/typesetting/format.py | 4 + src/math_spec/typesetting/latex.py | 3 + src/math_spec/typesetting/typst.py | 3 + src/math_spec/typesetting/walk.py | 14 +++ tests/test_lowering.py | 37 +++++++ tests/test_validation.py | 131 +++++++++++++++++++++++++ tests/typesetting/golden/latex.out | 3 + tests/typesetting/golden/markdown.out | 18 ++++ tests/typesetting/golden/model.yaml | 12 +++ tests/typesetting/golden/typst.out | 3 + tests/typesetting/test_golden.py | 11 ++- tests/typesetting/test_walk.py | 43 +++++++- 21 files changed, 691 insertions(+), 33 deletions(-) diff --git a/docs/reference/language/expressions.md b/docs/reference/language/expressions.md index bccc7924..a94255e1 100644 --- a/docs/reference/language/expressions.md +++ b/docs/reference/language/expressions.md @@ -114,30 +114,35 @@ A `where:` is a boolean mask, and true means "this coordinate exists". where_expr ::= atom | "NOT" where_expr | where_expr ("AND"|"OR") where_expr | "(" where_expr ")" atom ::= NAME | NAME COMPARATOR value | expression COMPARATOR expression - | POSITION COMPARATOR INTEGER | "True" | "False" + | POSITION COMPARATOR INTEGER | COUNT COMPARATOR INTEGER | TRANSLATED + | "True" | "False" COMPARATOR ::= "<=" | ">=" | "==" | "!=" | "<" | ">" value ::= NUMBER | QUOTED | NAME_OR_STRING expression ::= the arithmetic grammar above, with no variable and no dual in it POSITION ::= "position" "(" NAME [ "," "by" "=" NAME "," "within" "=" COLUMNS ] ")" +COUNT ::= "count" "(" where_expr "," "over" "=" NAME ")" +TRANSLATED ::= "shift" "(" where_expr "," "along" "=" NAME "," "offset" "=" INTEGER ")" COLUMNS ::= NAME | "[" NAME { "," NAME } "]" QUOTED ::= "'" chars "'" | '"' chars '"' ``` -| Written as | Names a… | Meaning | -| --------------------------------------- | -------------------------- | ----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | -| `name` (bare) | parameter | The value is defined here. A `bool` is its own answer. A `str` is defined wherever the table has a row. A number has to have a row and be finite | -| `name` (bare) | variable | The variable exists at this coordinate | -| `name` (bare) | relation | A row exists, read at the relation's key. A relation may be [partial](relations.md#the-data-contract), and this selects the labels that do map | -| `name` (bare) | dimension | A load error. It would be true everywhere | -| `name OP value` | parameter | Element-wise, and a null compares false | -| `name OP value` | dimension | A filter on the frame's own coordinate column | -| `name OP value`, `name.col OP value` | relation | A filter on a value column, read at the relation's key. Name the column where the key determines several | -| `name OP name`, `name.a OP name.b` | two relation columns | Legal where both relations are keyed over the same dimensions and both columns are over one dimension. `ends.bus0 != ends.bus1` excludes a self-loop | -| `expression OP expression` | arithmetic over parameters | Coordinate by coordinate, over every dimension either side carries ([arithmetic in a comparison](#arithmetic-in-a-comparison)). A side with no value at a coordinate compares false | -| `position(name) OP i` | dimension | Where the row sits along the dimension's own order. `0` is first, and a negative number counts from the end | -| `position(name, by=relation, within=c)` | dimension | The same, counted within each group the relation makes | -| `AND` `OR` `NOT` | — | Case-insensitive. `NOT` binds tighter than `AND`, and `AND` tighter than `OR` | -| `True` / `False` | — | `True` is the same as no `where`; `False` gives a declaration with no rows. A [case `when:`](named.md#the-rules-that-keep-the-cases-apart) may not fold to either | +| Written as | Names a… | Meaning | +| ----------------------------------------- | -------------------------- | ----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `name` (bare) | parameter | The value is defined here. A `bool` is its own answer. A `str` is defined wherever the table has a row. A number has to have a row and be finite | +| `name` (bare) | variable | The variable exists at this coordinate | +| `name` (bare) | relation | A row exists, read at the relation's key. A relation may be [partial](relations.md#the-data-contract), and this selects the labels that do map | +| `name` (bare) | dimension | A load error. It would be true everywhere | +| `name OP value` | parameter | Element-wise, and a null compares false | +| `name OP value` | dimension | A filter on the frame's own coordinate column | +| `name OP value`, `name.col OP value` | relation | A filter on a value column, read at the relation's key. Name the column where the key determines several | +| `name OP name`, `name.a OP name.b` | two relation columns | Legal where both relations are keyed over the same dimensions and both columns are over one dimension. `ends.bus0 != ends.bus1` excludes a self-loop | +| `expression OP expression` | arithmetic over parameters | Coordinate by coordinate, over every dimension either side carries ([arithmetic in a comparison](#arithmetic-in-a-comparison)). A side with no value at a coordinate compares false | +| `position(name) OP i` | dimension | Where the row sits along the dimension's own order. `0` is first, and a negative number counts from the end | +| `position(name, by=relation, within=c)` | dimension | The same, counted within each group the relation makes | +| `count(where_expr, over=name) OP i` | dimension | How many coordinates along the dimension the predicate admits ([counting what a predicate admits](#counting-what-a-predicate-admits)) | +| `shift(where_expr, along=name, offset=i)` | dimension | The predicate read `i` coordinates back, and false where that vacates | +| `AND` `OR` `NOT` | — | Case-insensitive. `NOT` binds tighter than `AND`, and `AND` tighter than `OR` | +| `True` / `False` | — | `True` is the same as no `where`; `False` gives a declaration with no rows. A [case `when:`](named.md#the-rules-that-keep-the-cases-apart) may not fold to either | The dimensions of the mask must not exceed the frame it sits in. A bare name that is not declared is a load error. @@ -147,6 +152,63 @@ that is not declared is a load error. A bare parameter name is true wherever the table has a row, and a row holding `0.0` is a row. Where you mean non-zero, write `where: "inflow != 0"`. +### Counting what a predicate admits + +`count(, over=)` is how many coordinates along that +dimension the predicate is true at. It is the one place a predicate is read as +a number, and it is compared against a whole number: + +```yaml +dimensions: + bp: { dtype: int } + generator: { dtype: str } +parameters: + bp_x: { dims: [generator, bp] } + points: { dims: [generator, bp], dtype: bool } +variables: + p: + dims: [generator] + where: "count(points, over=bp) >= 2" + bounds: { lower: 0 } +constraints: + cap: + dims: [generator] + expression: p <= 1 +objective: + sense: minimize + expression: sum(p, over=generator) +``` + +$$\lvert \{ b \in \mathcal{B} \thinspace : \thinspace \mathrm{points}_{g,b} \} \rvert \ge 2 \qquad \forall\thinspace g \in \mathcal{G}$$ + +The dimension counted over is **removed**, as a `sum(over=)` removes it, so +what is left is one number per remaining coordinate — one per generator above. +The count therefore states a fact about each group without naming the group. +Counting along a dimension the predicate does not read is a load error. + +The comparison takes a whole number on the right. A count is a number of +coordinates, so a fraction and a parameter are both load errors. + +### Reading a predicate at the previous coordinate + +`shift(, along=, offset=)` reads the predicate +`offset` coordinates back. It is **false** where the translation vacates, and +it takes no `edge=`: the arithmetic `shift` needs one because no number is +neutral, and false is what a missing row already means in a mask. + +The two together name the start of a run — a coordinate the mask admits whose +neighbour before it the mask does not: + +```yaml +where: "count(points AND NOT shift(points, along=bp, offset=1), over=bp) == 1" +``` + +That reads: the marked breakpoints are one consecutive run. + +A negative `offset` reads forwards. `by=`, `within=` and `edge='wrap'` are not +in this form; where you need a grouped or cyclic translation, compare the +arithmetic one instead. + ### The right-hand side of a comparison A bare name on the right is read as a string label when the model does not diff --git a/docs/reference/notation.md b/docs/reference/notation.md index e3395367..56f433f7 100644 --- a/docs/reference/notation.md +++ b/docs/reference/notation.md @@ -747,6 +747,51 @@ covered: \sum_{t \in \mathcal{T},\ g \in \mathcal{G}} p_{t,g} \le \mathrm{budget} \qquad \text{where } \sum_{g \in \mathcal{G}} \mathrm{p}^{\mathrm{max}}_{g} \ge \mathrm{budget} ``` +#### `counted` + +a count of the coordinates a predicate admits, which reduces one dim away + +```yaml +counted: + dims: [bus] + where: "count(tech_cap > 0, over=technology) >= 2" + expression: theta <= budget +``` + +```math +\theta_{b} \le \mathrm{budget} \qquad \forall\, b \in \mathcal{B} \,:\, \lvert \{ e \in \mathcal{E} \,:\, \mathrm{tech\_cap}_{b,e} > 0 \} \rvert \ge 2 +``` + +#### `counted_here` + +the same count along a dim the frame carries, so the set takes a primed dummy + +```yaml +counted_here: + dims: [bus, technology] + where: "count(tech_cap > 0, over=technology) >= 2" + expression: theta <= tech_cap +``` + +```math +\theta_{b} \le \mathrm{tech\_cap}_{b,e} \qquad \forall\, b \in \mathcal{B},\ e \in \mathcal{E} \,:\, \lvert \{ e' \in \mathcal{E} \,:\, \mathrm{tech\_cap}_{b,e'} > 0 \} \rvert \ge 2 +``` + +#### `run_start` + +a predicate read one coordinate back, which is false where the translation vacates + +```yaml +run_start: + dims: [snapshot, bus] + where: "load AND NOT shift(load, along=snapshot, offset=1)" + expression: slack <= load +``` + +```math +\mathit{slack}_{t} \le \mathrm{load}_{t,b} \qquad \forall\, t \in \mathcal{T},\ b \in \mathcal{B} \,:\, \mathrm{load}_{t,b} \text{ is defined} \wedge \neg \left( \mathrm{load}_{t - 1,b} \text{ is defined} \right) +``` + #### `capped` an expressions: entry on a side, read by the name the file gave it diff --git a/docs/reference/reading.md b/docs/reference/reading.md index de490c89..65397ba1 100644 --- a/docs/reference/reading.md +++ b/docs/reference/reading.md @@ -107,6 +107,14 @@ A name compared against a literal does not arrive this way. `p_max > 5` is a both mask the same coordinates. Match both where you read a comparison over parameters. +Two predicates read another predicate rather than a declaration. A +`CountComparison` carries the mask it counts and the dimension it counts away; +a `TranslatedPredicate` carries the mask it reads at a neighbouring +coordinate. Each holds that mask as a `Mask`, where a connective holds a bare +predicate: the walk recurses through a connective and stops at these, so read +the field where you need what is inside. `.names_read` and `.dims` already see +through both. + A predicate you build yourself answers the same four questions: wrap it in `Mask`, or build it there with `~`, `&` and `|`. A mask folds as it is built, so a boolean literal stands at a mask's root or nowhere. A `Region`'s `when` diff --git a/src/math_spec/_where_parser.py b/src/math_spec/_where_parser.py index eea60243..4b18cfd8 100644 --- a/src/math_spec/_where_parser.py +++ b/src/math_spec/_where_parser.py @@ -14,13 +14,14 @@ from __future__ import annotations -from dataclasses import dataclass +from dataclasses import dataclass, field from functools import lru_cache from typing import TYPE_CHECKING, cast, get_args import pyparsing as pp from math_spec._expression_parser import ARITHMETIC, NAME, ArithmeticNode, children, parse_text +from math_spec._sealed import Sealed from math_spec.program import ( And, BooleanLiteral, @@ -33,7 +34,7 @@ ) if TYPE_CHECKING: - from collections.abc import Callable + from collections.abc import Callable, Mapping # --------------------------------------------------------------------------- # AST nodes @@ -72,6 +73,40 @@ class QuotedNode: value: str +@dataclass(frozen=True) +class UnresolvedPredicateCallNode: + """``([, …])`` — an operator reading a predicate rather than arithmetic. + + The expression grammar cannot carry this shape: its call rule takes + arithmetic arguments, and a predicate is not arithmetic. So the where + grammar reads it, and resolution decides which operator the name is and + what the kwargs mean. ``kwargs`` is held and hashed as + :class:`~math_spec._expression_parser.FunctionCallNode` holds its own. + """ + + name: str + operand: Predicate | UnresolvedWhereNode + kwargs: Mapping[str, ArithmeticNode] = field(default_factory=dict, hash=False) + + def __post_init__(self) -> None: + object.__setattr__(self, 'kwargs', Sealed(self.kwargs)) + + +@dataclass(frozen=True) +class UnresolvedCountNode: + """``count(, over=) `` — a count against a literal. + + Its own node rather than an :class:`UnresolvedComparisonNode` with a call + on the left: every other comparison has arithmetic on both sides, and + widening that one to carry a predicate would widen every reader of a side + with it. + """ + + call: UnresolvedPredicateCallNode + op: PredicateOperator + value: ArithmeticNode + + @dataclass(frozen=True) class UnresolvedComparisonNode: """``side side`` before the sides are read; ``resolution.py`` decides what each is. @@ -86,9 +121,9 @@ class UnresolvedComparisonNode: right: ArithmeticNode | ColumnNode | QuotedNode -#: What resolution rewrites away on the where side — the two nodes whose -#: leaves are still names the schema has not been asked about. -UnresolvedWhereNode = UnresolvedNameNode | UnresolvedComparisonNode +#: What resolution rewrites away on the where side — the nodes whose leaves +#: are still names the schema has not been asked about. +UnresolvedWhereNode = UnresolvedNameNode | UnresolvedComparisonNode | UnresolvedPredicateCallNode | UnresolvedCountNode #: Every node a parsed where string is built of: the connectives and literals, #: the unresolved leaves, and the arithmetic and the two side nodes under a @@ -122,14 +157,42 @@ def _build_where_grammar() -> pp.ParserElement: ) comparator = pp.one_of(list(get_args(PredicateOperator))) - comparison = ((column | ARITHMETIC) + comparator + (quoted | column | ARITHMETIC)).set_parse_action( + kwarg = (name + pp.Suppress('=') + (quoted | ARITHMETIC)).set_parse_action(lambda t: (t[0], t[1])) + + def _call(head: pp.ParserElement) -> pp.ParserElement: + """``([, …])`` — the one shape whose operand is a predicate.""" + return ( + head + pp.Suppress('(') + where_expr + pp.ZeroOrMore(pp.Suppress(',') + kwarg) + pp.Suppress(')') + # pyrefly: ignore[implicit-any-lambda] + ).set_parse_action(lambda t: UnresolvedPredicateCallNode(t[0], t[1], dict(t[2:]))) + + # `count` is spelled here rather than left to resolution, as `position` is, + # because only the grammar can decide to read its argument as a predicate. + # Every other predicate-taking call stands where arithmetic cannot, so the + # comparison above it has already been tried and the name is free. + count_call = _call(pp.CaselessKeyword('count')) + predicate_call = _call(name.copy()) + + count_comparison = (count_call + comparator + ARITHMETIC).set_parse_action( # pyrefly: ignore[implicit-any-lambda] - lambda t: UnresolvedComparisonNode(t[0], t[1], t[2]) + lambda t: UnresolvedCountNode(t[0], t[1], t[2]) ) + comparison = ( + (column | ARITHMETIC) + comparator + (quoted | column | ARITHMETIC) + # pyrefly: ignore[implicit-any-lambda] + ).set_parse_action(lambda t: UnresolvedComparisonNode(t[0], t[1], t[2])) # pyrefly: ignore[implicit-any-lambda] existence = name.copy().set_parse_action(lambda t: UnresolvedNameNode(t[0])) - atom = true_lit | false_lit | comparison | existence | (pp.Suppress('(') + where_expr + pp.Suppress(')')) + atom = ( + true_lit + | false_lit + | count_comparison + | comparison + | predicate_call + | existence + | (pp.Suppress('(') + where_expr + pp.Suppress(')')) + ) NOT = pp.CaselessKeyword('NOT').suppress() # pyrefly: ignore[implicit-any-lambda] @@ -191,7 +254,11 @@ def _named_rewrite(text: str, loc: int) -> str | None: def _nested(node: _ParsedWhere) -> tuple[_ParsedWhere, ...]: - """What a where string nests through: a connective's operands, and the arithmetic on a comparison's sides.""" + """What a where string nests through: a connective's operands, a comparison's sides, and a call's predicate.""" + if isinstance(node, UnresolvedCountNode): + return (node.call, node.value) + if isinstance(node, UnresolvedPredicateCallNode): + return (node.operand, *node.kwargs.values()) if isinstance(node, UnresolvedComparisonNode): return (node.left, node.right) if isinstance(node, ArithmeticNode): diff --git a/src/math_spec/dimensions.py b/src/math_spec/dimensions.py index d533f894..5a5f5cbd 100644 --- a/src/math_spec/dimensions.py +++ b/src/math_spec/dimensions.py @@ -42,6 +42,7 @@ from math_spec.operators import BUILTINS from math_spec.program import ( ArithmeticComparison, + CountComparison, DimensionComparison, DimensionPosition, Direction, @@ -53,6 +54,7 @@ RelationComparison, RelationDefined, RelationPairComparison, + TranslatedPredicate, VariableDefined, ) @@ -551,6 +553,10 @@ def _check_where_dims( leaf = f"where-relation '{atom.name}'" case ArithmeticComparison() | ExpressionComparison(): leaf = 'a where-comparison of expressions' + case CountComparison(): + leaf = f"a where-count over '{atom.over}'" + case TranslatedPredicate(): + leaf = f"a where-predicate translated along '{atom.along}'" case _: assert_never(atom) raise DimensionError( diff --git a/src/math_spec/exclusivity.py b/src/math_spec/exclusivity.py index 2106e044..2fd99c4c 100644 --- a/src/math_spec/exclusivity.py +++ b/src/math_spec/exclusivity.py @@ -25,6 +25,7 @@ And, ArithmeticComparison, BooleanLiteral, + CountComparison, DimensionComparison, DimensionPosition, ExpressionComparison, @@ -36,6 +37,7 @@ RelationComparison, RelationDefined, RelationPairComparison, + TranslatedPredicate, TypedPredicate, VariableDefined, ) @@ -205,6 +207,18 @@ def _observe( 'literal, or precompute the test as a boolean parameter and test that' ) raise Undecidable(msg) + if isinstance(node, CountComparison): + msg = ( + 'it counts the coordinates a predicate admits, which only the data decides — test a parameter ' + 'against a literal, or precompute the count as a parameter and test that' + ) + raise Undecidable(msg) + if isinstance(node, TranslatedPredicate): + msg = ( + 'it reads a predicate at a neighbouring coordinate, and which rows that admits only the data ' + 'decides — test this row, or precompute the neighbour as a boolean parameter and test that' + ) + raise Undecidable(msg) if isinstance(node, DimensionPosition): values.add(node.position) elif isinstance(node, RelationPairComparison): @@ -244,6 +258,10 @@ def _subject_of(node: TypedPredicate) -> Subject: return Subject('relation_pair', name, other) case ArithmeticComparison() | ExpressionComparison(): return Subject('expression', 'a comparison of expressions') + case CountComparison(): + return Subject('expression', 'a count of the coordinates a predicate admits') + case TranslatedPredicate(): + return Subject('expression', 'a predicate read at a neighbouring coordinate') case _: assert_never(node) @@ -438,6 +456,9 @@ def _atom(node: TypedPredicate, cell: dict[Subject, Cell], grid: _Grid) -> bool: case ArithmeticComparison() | ExpressionComparison(): msg = 'a comparison of expressions is refused as undecidable before any cell is read' raise AssertionError(msg) + case CountComparison() | TranslatedPredicate(): + msg = 'a predicate read as a count or at a neighbour is refused as undecidable before any cell is read' + raise AssertionError(msg) case DimensionPosition(op=op, position=position): return _compare(value, op, position) case ParameterComparison(op=op, value=literal) | RelationComparison(op=op, value=literal): diff --git a/src/math_spec/lowering.py b/src/math_spec/lowering.py index 66b48e5a..5cf01292 100644 --- a/src/math_spec/lowering.py +++ b/src/math_spec/lowering.py @@ -12,7 +12,7 @@ from __future__ import annotations -from dataclasses import dataclass +from dataclasses import dataclass, replace from typing import TYPE_CHECKING, assert_never import math_spec.program as program @@ -255,11 +255,19 @@ def mask(self, mask: program.Mask | None) -> program.Mask | None: Every other predicate node is already the program's own and passes through; a mask holding none comes back equal to the one handed in. """ - return None if mask is None else program.Mask(self._predicate(mask.root)) + return None if mask is None else self._mask(mask) + + def _mask(self, mask: program.Mask) -> program.Mask: + """*mask* rebuilt — the one a leaf carries is rebuilt the same way as the one a declaration does.""" + return program.Mask(self._predicate(mask.root)) def _predicate(self, node: program.Predicate) -> program.Predicate: if isinstance(node, program.ArithmeticComparison): return program.ExpressionComparison(self.expr(node.left), node.op, self.expr(node.right), node.dims) + if isinstance(node, program.CountComparison): + return replace(node, predicate=self._mask(node.predicate)) + if isinstance(node, program.TranslatedPredicate): + return replace(node, operand=self._mask(node.operand)) if isinstance(node, program.Not): return program.Not(self._predicate(node.operand)) if isinstance(node, program.And): diff --git a/src/math_spec/program.py b/src/math_spec/program.py index bf4fa3ec..8315f442 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -52,6 +52,7 @@ 'ConstraintDeclaration', 'ConstraintSense', 'Contiguous', + 'CountComparison', 'Curved', 'Derivation', 'DimensionComparison', @@ -101,6 +102,7 @@ 'SosDeclaration', 'Sum', 'Translate', + 'TranslatedPredicate', 'TypedPredicate', 'Variable', 'VariableAbsence', @@ -1242,6 +1244,41 @@ class RelationDefined: dims: tuple[str, ...] = () +@dataclass(frozen=True) +class CountComparison: + """How many coordinates *predicate* admits along *over*, against a literal — ``count(points, over=bp) >= 2``. + + The count is one number per coordinate of ``dims``, which is every dim + *predicate* reads minus *over*, so a claim about each curve is written + without saying "each curve". A predicate a leaf reads arrives as a + :class:`Mask`, where a connective's operand is a bare :data:`Predicate`: + a walk recurses through the second and stops at the first. + """ + + predicate: Mask + over: str + op: PredicateOperator + value: float + dims: tuple[str, ...] + + +@dataclass(frozen=True) +class TranslatedPredicate: + """*operand* read at a neighbouring coordinate — ``shift(points, along=bp, offset=1)``. + + False where the translation vacates, and there is no ``edge=`` to state. + The arithmetic translation needs one because no number is neutral and + inventing one changes the answer; false is what a missing row already + means in a mask, so the predicate form has the value the language already + gives it. + """ + + operand: Mask + along: str + offset: int + dims: tuple[str, ...] + + @dataclass(frozen=True) class Not: operand: Predicate @@ -1277,6 +1314,8 @@ class Or: | RelationComparison | RelationPairComparison | RelationDefined + | CountComparison + | TranslatedPredicate | Not | And | Or @@ -1296,6 +1335,8 @@ class Or: | RelationComparison | RelationPairComparison | RelationDefined + | CountComparison + | TranslatedPredicate ) #: The boolean connectives — the only where nodes carrying other where nodes, @@ -1356,6 +1397,8 @@ def _atom_dims(atom: TypedPredicate) -> frozenset[str]: | ArithmeticComparison() | ParameterDefined() | VariableDefined() + | CountComparison() + | TranslatedPredicate() ): return frozenset(atom.dims) case DimensionComparison(): @@ -1389,6 +1432,10 @@ def _atom_names(atom: TypedPredicate) -> frozenset[str]: case ArithmeticComparison(): msg = 'a resolved mask is asked what it reads; lowering rebuilds it first, and the program mask answers.' raise AssertionError(msg) + case CountComparison(): + return atom.predicate.names_read + case TranslatedPredicate(): + return atom.operand.names_read case DimensionComparison() | DimensionPosition(): return frozenset() case _: diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index d93b4727..9c58bf64 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -49,7 +49,9 @@ ColumnNode, QuotedNode, UnresolvedComparisonNode, + UnresolvedCountNode, UnresolvedNameNode, + UnresolvedPredicateCallNode, UnresolvedWhereNode, parse_where, ) @@ -69,6 +71,7 @@ And, ArithmeticComparison, BooleanLiteral, + CountComparison, DimensionComparison, DimensionPosition, Direction, @@ -84,6 +87,7 @@ RelationDeclaration, RelationDefined, RelationPairComparison, + TranslatedPredicate, TypedPredicate, VariableDefined, ) @@ -730,6 +734,10 @@ def where(self, node: Predicate | UnresolvedWhereNode) -> Predicate | Unresolved return self._where_name(node) if isinstance(node, UnresolvedComparisonNode): return self._comparison(node) + if isinstance(node, UnresolvedPredicateCallNode): + return self._predicate_call(node) + if isinstance(node, UnresolvedCountNode): + return self._count(node) if isinstance(node, Not): return Not(self._child(node.operand)) if isinstance(node, And): @@ -780,6 +788,92 @@ def _where_name(self, node: UnresolvedNameNode) -> Predicate | UnresolvedWhereNo return VariableDefined(node.name, ns.leaf_dims[node.name]) return node + def _predicate_call(self, node: UnresolvedPredicateCallNode) -> Predicate | UnresolvedWhereNode: + """``shift(, along=, offset=)`` — the one operator that reads a predicate and answers one. + + ``count`` answers a number, so it stands on a comparison's side and + :meth:`_count` reads it there. Anything else naming a predicate is + refused here rather than resolved into arithmetic it cannot be. + + An operand that failed to resolve is handed straight back: resolution + collects problems rather than raising, and asking an unresolved + predicate for its dims asserts instead of refusing. + """ + context, found = self.context, len(self.errors) + if node.name != 'shift': + self.errors.append( + f"{context}: '{node.name}()' does not read a predicate. `shift` reads one and answers one, " + f'`count` reads one and answers a number, and every other operator reads arithmetic. ' + f'Compare the predicate, or name a parameter carrying it.' + ) + return node + operand = self._child(node.operand) + if len(self.errors) > found: + return node + if (refusal := _kwargs_error(context, 'shift', node.kwargs, required=('along', 'offset'))) is not None: + self.errors.append(refusal) + return node + along = node.kwargs['along'] + offset = _literal(node.kwargs['offset']) + if not isinstance(along, NameNode) or self.ns.kind(along.name) != 'dimension': + self.errors.append( + f'{context}: shift(, along=) names the dimension the predicate is read back along. ' + f'Name a declared dimension.' + ) + return node + if offset is None or not offset.value.is_integer(): + self.errors.append( + f'{context}: shift(, offset=) counts whole coordinates back along ' + f"'{along.name}'. Write an integer." + ) + return node + mask = Mask(operand) + if along.name not in mask.dims: + self.errors.append( + f"{context}: shift(, along='{along.name}') reads the predicate back along a dimension " + f'it does not carry — it reads {_listed(sorted(mask.dims))}. Translate it along one of those.' + ) + return node + return TranslatedPredicate(mask, along.name, int(offset.value), tuple(sorted(mask.dims))) + + def _count(self, node: UnresolvedCountNode) -> Predicate | UnresolvedWhereNode: + """``count(, over=) `` — how many coordinates the predicate admits. + + The reduction leaves every dim but ``over``, so the count is one + number per remaining coordinate and a claim about each group needs no + word for the group. + """ + context, found = self.context, len(self.errors) + operand = self._child(node.call.operand) + if len(self.errors) > found: + return node + if (refusal := _kwargs_error(context, 'count', node.call.kwargs, required=('over',))) is not None: + self.errors.append(refusal) + return node + over = node.call.kwargs['over'] + if not isinstance(over, NameNode) or self.ns.kind(over.name) != 'dimension': + self.errors.append( + f'{context}: count(, over=) names the dimension the coordinates are counted along. ' + f'Name a declared dimension.' + ) + return node + value = _literal(node.value) + if value is None or not value.value.is_integer(): + self.errors.append( + f'{context}: a count is a whole number of coordinates, so it is compared against one. ' + f'Write count(…, over={over.name}) {node.op} .' + ) + return node + mask = Mask(operand) + if over.name not in mask.dims: + self.errors.append( + f"{context}: count(, over='{over.name}') counts along a dimension the predicate does " + f'not carry — it reads {_listed(sorted(mask.dims))}. Count along one of those.' + ) + return node + dims = tuple(sorted(mask.dims - {over.name})) + return CountComparison(mask, over.name, node.op, value.value, dims) + def _comparison(self, node: UnresolvedComparisonNode) -> Predicate | UnresolvedWhereNode: """``side side``, read for what each side is. @@ -1141,6 +1235,36 @@ def _position_shape(call: FunctionCallNode) -> tuple[str, str | None, tuple[str, return call.args[0].name, by.name if by is not None else None, into +def _kwargs_error( + context: str, name: str, kwargs: Mapping[str, ArithmeticNode], required: tuple[str, ...] +) -> str | None: + """Why *kwargs* is not what *name* takes over a predicate, or ``None`` where it is. + + A predicate-reading call takes exactly the keywords named here. The + arithmetic forms of these operators take more — an ``edge=``, a ``by=`` — + and each is refused rather than ignored, since a predicate answers the + vacated coordinate itself and a grouped form has nobody asking for it yet. + """ + missing = [key for key in required if key not in kwargs] + if missing: + return f'{context}: {name}() needs {_listed([f"{key}=" for key in missing])}.' + if extra := sorted(set(kwargs) - set(required)): + return ( + f'{context}: {name}() does not take {_listed([f"{key}=" for key in extra])}. ' + f'It takes {_listed([f"{key}=" for key in required])}, and nothing else: a predicate is false ' + f'where a translation vacates, so there is no edge to state.' + ) + return None + + +def _listed(items: list[str]) -> str: + """``'a'``, ``'a' and 'b'``, ``'a', 'b' and 'c'`` — one rule, so every message reads the same.""" + quoted = [f"'{item}'" for item in items] + if len(quoted) <= 1: + return quoted[0] if quoted else 'nothing' + return f'{", ".join(quoted[:-1])} and {quoted[-1]}' + + def _literal(value: ArithmeticNode) -> NumberNode | None: """The number a literal names, its sign folded in — ``None`` where *value* is not one. diff --git a/src/math_spec/typesetting/format.py b/src/math_spec/typesetting/format.py index a0bedfb5..80612b69 100644 --- a/src/math_spec/typesetting/format.py +++ b/src/math_spec/typesetting/format.py @@ -223,6 +223,10 @@ def cardinality(self, inner: str) -> str: def fraction(self, numerator: str, denominator: str) -> str: ... + def set_of(self, members: str, condition: str) -> str: + """A set by comprehension: ``{ k ∈ K : condition }``.""" + ... + def summation(self, domain: str, body: str) -> str: ... def cases(self, arms: list[tuple[str, str]]) -> str: diff --git a/src/math_spec/typesetting/latex.py b/src/math_spec/typesetting/latex.py index d09c3ff2..3e3296f5 100644 --- a/src/math_spec/typesetting/latex.py +++ b/src/math_spec/typesetting/latex.py @@ -98,6 +98,9 @@ def cardinality(self, inner: str) -> str: def fraction(self, numerator: str, denominator: str) -> str: return rf'\frac{{{numerator}}}{{{denominator}}}' + def set_of(self, members: str, condition: str) -> str: + return rf'\{{ {members} {self.operators["such_that"]} {condition} \}}' + def cases(self, arms: list[tuple[str, str]]) -> str: rows = self.cases_row.join(f'{value} & {condition}' for value, condition in arms) return rf'\begin{{cases}} {rows} \end{{cases}}' diff --git a/src/math_spec/typesetting/typst.py b/src/math_spec/typesetting/typst.py index 55c34c7f..d83a1aeb 100644 --- a/src/math_spec/typesetting/typst.py +++ b/src/math_spec/typesetting/typst.py @@ -106,6 +106,9 @@ def cardinality(self, inner: str) -> str: def fraction(self, numerator: str, denominator: str) -> str: return f'frac({numerator}, {denominator})' + def set_of(self, members: str, condition: str) -> str: + return f'{{{members} {self.operators["such_that"]} {condition}}}' + def cases(self, arms: list[tuple[str, str]]) -> str: return 'cases({})'.format(self.cases_row.join(f'{value} & {condition}' for value, condition in arms)) diff --git a/src/math_spec/typesetting/walk.py b/src/math_spec/typesetting/walk.py index 2413c3c6..efe8f5f3 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -39,6 +39,7 @@ And, ArithmeticComparison, BooleanLiteral, + CountComparison, DimensionComparison, DimensionPosition, Direction, @@ -53,6 +54,7 @@ RelationComparison, RelationDefined, RelationPairComparison, + TranslatedPredicate, VariableDefined, ) from math_spec.typesetting.format import Entry, Glossary, Line, OperatorName @@ -620,6 +622,18 @@ def _where(self, node: Predicate, ctx: _Context) -> tuple[str, int]: right = self._value_read(node.other, node.other_column, ctx) return f'{left} {self._op(_PREDICATES[node.op])} {right}', comparison + if isinstance(node, CountComparison): + index, inner = ctx.reducing(node.over) + counted = self.format.set_of( + self._membership(node.over, index), self._predicate(node.predicate.root, inner) + ) + size = self.format.cardinality(counted) + return f'{size} {self._op(_PREDICATES[node.op])} {self._number(node.value)}', comparison + + if isinstance(node, TranslatedPredicate): + moved = ctx.translated(node.along, _Step(node.offset, 'plain')) + return self._where(node.operand.root, moved) + if isinstance(node, RelationDefined): return self._relation_row(node.name, self._frame_key(node.name, ctx)), comparison diff --git a/tests/test_lowering.py b/tests/test_lowering.py index e0e8c282..7a8c6548 100644 --- a/tests/test_lowering.py +++ b/tests/test_lowering.py @@ -28,6 +28,7 @@ BooleanLiteral, Cases, Constant, + CountComparison, DimensionComparison, DimensionDeclaration, Direction, @@ -430,6 +431,42 @@ def test_a_comparison_of_expressions_lowers_to_program_expressions_on_both_sides ) +def test_a_predicate_a_leaf_carries_is_lowered_like_any_other_mask(): + """A comparison of expressions inside a count is rebuilt too, so a program mask is program vocabulary throughout.""" + program = to_program( + override( + SHAPES_MODEL, + **{'constraints.w': {'dims': ['g'], 'where': 'count(c <= 0.5 * k, over=g) >= 2', 'expression': 'p <= c'}}, + ) + ) + mask = program.constraints['w'].where + assert mask is not None and isinstance(mask.root, CountComparison) + assert mask.root.predicate.root == ExpressionComparison( + Parameter('c'), '<=', Multiply(Constant(0.5), Parameter('k')), ('g',) + ), 'the counted predicate is rebuilt, not handed through with the resolved comparison still in it' + assert mask.names_read == frozenset({'c', 'k'}), 'what the counted predicate reads is data the consumer binds' + + +def test_a_translated_predicate_keeps_what_it_reads_in_reach(): + """A walk that asks a mask what it names has to see through the translation, or the column is silently dropped.""" + program = to_program( + override( + SHAPES_MODEL, + **{ + 'constraints.w': { + 'dims': ['g'], + 'where': 'flag AND NOT shift(flag, along=g, offset=1)', + 'expression': 'p <= c', + } + }, + ) + ) + mask = program.constraints['w'].where + assert mask is not None + assert mask.names_read == frozenset({'flag'}), 'the translated half reads the same column as the plain one' + assert sorted(mask.dims) == ['g'] + + def test_a_mask_with_no_arithmetic_is_the_same_mask_after_lowering(dispatch_program): """Every other predicate node is already the program's own, so lowering hands it through unchanged.""" assert dispatch_program.variables['dispatch'].where == Mask(CAPACITY_POSITIVE) diff --git a/tests/test_validation.py b/tests/test_validation.py index df7ac848..2f7d1007 100644 --- a/tests/test_validation.py +++ b/tests/test_validation.py @@ -708,6 +708,137 @@ def test_a_case_comparing_expressions_is_refused_as_undecidable(self): assert 'precompute the test as a boolean parameter' in message +class TestAPredicateIsAnOperand: + """``count`` and ``shift`` over a predicate — the two calls that read one rather than arithmetic. + + Everything else in the language takes arithmetic, so the grammar reads + these shapes itself and resolution decides what each name is. + """ + + @pytest.mark.parametrize( + 'where', + [ + pytest.param('count(flag, over=g) >= 2', id='a-bare-mask'), + pytest.param('count(c > 0, over=g) == 0', id='a-comparison'), + pytest.param('count(flag AND NOT shift(flag, along=g, offset=1), over=g) == 1', id='a-run-start'), + pytest.param('count(c > 0, over=g) == 0 OR count(c < 0, over=g) == 0', id='two-counts-under-or'), + pytest.param('shift(flag, along=g, offset=1)', id='a-translated-mask'), + pytest.param('shift(flag, along=g, offset=-1)', id='a-translation-forwards'), + pytest.param('count(shift(flag, along=g, offset=1), over=g) >= 1', id='a-translation-under-a-count'), + ], + ) + def test_a_shape_the_language_admits(self, where): + mask = where_of(where, Namespace(_schema()), 'probe') + assert mask is not None, 'the predicate decides some rows, so it is a mask rather than nothing' + + @pytest.mark.parametrize( + ('where', 'fragments'), + [ + pytest.param( + 'count(flag, by=lk) >= 2', + ("count() needs 'over='",), + id='a-count-with-no-over', + ), + pytest.param( + 'count(flag, over=g, by=lk) >= 2', + ("does not take 'by='", "It takes 'over='"), + id='a-count-with-a-keyword-it-lacks', + ), + pytest.param( + 'count(flag, over=c) >= 2', + ('names the dimension the coordinates are counted along',), + id='a-count-over-a-parameter', + ), + pytest.param( + 'count(flag, over=h) >= 2', + ('counts along a dimension the predicate does not carry', "it reads 'g'"), + id='a-count-over-a-dim-the-predicate-lacks', + ), + pytest.param( + 'count(flag, over=g) >= 2.5', + ('a count is a whole number of coordinates',), + id='a-count-against-a-fraction', + ), + pytest.param( + 'count(flag, over=g) >= k', + ('a count is a whole number of coordinates',), + id='a-count-against-a-parameter', + ), + pytest.param( + 'shift(flag, along=g, offset=1, edge=0)', + ("does not take 'edge='", 'a predicate is false where a translation vacates'), + id='a-translated-predicate-with-an-edge', + ), + pytest.param( + 'shift(flag, along=g)', + ("shift() needs 'offset='",), + id='a-translation-with-no-offset', + ), + pytest.param( + 'shift(flag, along=h, offset=1)', + ('reads the predicate back along a dimension it does not carry',), + id='a-translation-along-a-dim-the-predicate-lacks', + ), + pytest.param( + 'shift(flag, along=g, offset=0.5)', + ('counts whole coordinates back',), + id='a-translation-by-a-fraction', + ), + pytest.param( + 'sum_back(flag, along=g, window=2)', + ("'sum_back()' does not read a predicate", '`count` reads one and answers a number'), + id='an-operator-that-reads-arithmetic', + ), + ], + ) + def test_a_shape_the_language_refuses(self, where, fragments): + with pytest.raises(LanguageError) as caught: + where_of(where, Namespace(_schema()), 'probe') + for fragment in fragments: + assert fragment in str(caught.value) + + @pytest.mark.parametrize( + 'where', + [ + pytest.param('count(nope, over=g) >= 2', id='under-a-count'), + pytest.param('shift(nope, along=g, offset=1)', id='under-a-translation'), + ], + ) + def test_a_name_the_operand_does_not_declare_is_reported_rather_than_walked(self, where): + """The operand is asked for its dims, and a walk over an unresolved node asserts rather than refusing (#590). + + Resolution collects problems instead of raising, so a failed operand + comes back unresolved and the count had walked it anyway. + """ + with pytest.raises(LanguageError) as caught: + where_of(where, Namespace(_schema()), 'probe') + assert "'nope' not found" in str(caught.value) + + def test_a_count_is_undecidable_in_a_case_when(self): + """Two cases are proved apart with no data, and how many coordinates a mask admits is the data's to say.""" + message = _refusal( + expressions={ + 'pick': { + 'dims': ['g'], + 'cases': { + 'many': {'when': 'count(flag, over=g) >= 2', 'expression': '1'}, + 'some': {'when': 'c > 0', 'expression': '2'}, + }, + 'otherwise': '0', + } + }, + constraints={'cap': {'dims': ['g'], 'expression': 'p <= pick'}}, + ) + assert 'it counts the coordinates a predicate admits, which only the data decides' in message + + def test_a_count_reduces_the_dim_it_counts_along_away(self): + """The count is one number per remaining coordinate, so a claim about each group needs no word for the group.""" + mask = where_of('count(q, over=h) >= 2', Namespace(_schema()), 'probe') + assert mask is not None + assert sorted(mask.dims) == ['g'], "'q' is read over g and h, and h is counted away" + assert mask.names_read == frozenset({'q'}), 'a consumer binds what the counted predicate reads' + + class TestRulesDecidedWithoutData: """Every refusal the schema or the resolver makes with no data bound, one row each.""" diff --git a/tests/typesetting/golden/latex.out b/tests/typesetting/golden/latex.out index 3c711fa9..5620e9ff 100644 --- a/tests/typesetting/golden/latex.out +++ b/tests/typesetting/golden/latex.out @@ -118,6 +118,9 @@ \text{margin} && p_{t,g} & \le \mathrm{p}^{\mathrm{max}}_{g} && \forall\, t \in \mathcal{T},\ g \in \mathcal{G} \,:\, \mathrm{p}^{\mathrm{max}}_{g} - \mathrm{p}^{\mathrm{min}}_{g} > \frac{\mathrm{cost}_{g}}{2} \\ \text{ramped} && \mathit{slack}_{t} & \le \mathrm{load}_{t,b} && \forall\, t \in \mathcal{T},\ b \in \mathcal{B} \,:\, \mathrm{load}_{t,b} - \mathrm{load}_{t \boxminus_{0} 1,b} \le \mathrm{zone\_cap}_{\mathrm{zone\_of}(b)} \wedge \mathrm{pos}(t) > 0 \\ \text{covered} && \sum_{t \in \mathcal{T},\ g \in \mathcal{G}} p_{t,g} & \le \mathrm{budget} && \text{where } \sum_{g \in \mathcal{G}} \mathrm{p}^{\mathrm{max}}_{g} \ge \mathrm{budget} \\ +\text{counted} && \theta_{b} & \le \mathrm{budget} && \forall\, b \in \mathcal{B} \,:\, \lvert \{ e \in \mathcal{E} \,:\, \mathrm{tech\_cap}_{b,e} > 0 \} \rvert \ge 2 \\ +\text{counted\_here} && \theta_{b} & \le \mathrm{tech\_cap}_{b,e} && \forall\, b \in \mathcal{B},\ e \in \mathcal{E} \,:\, \lvert \{ e' \in \mathcal{E} \,:\, \mathrm{tech\_cap}_{b,e'} > 0 \} \rvert \ge 2 \\ +\text{run\_start} && \mathit{slack}_{t} & \le \mathrm{load}_{t,b} && \forall\, t \in \mathcal{T},\ b \in \mathcal{B} \,:\, \mathrm{load}_{t,b} \text{ is defined} \wedge \neg \left( \mathrm{load}_{t - 1,b} \text{ is defined} \right) \\ \text{capped} && p_{t,g} & \le \mathrm{p}^{\mathrm{max}}_{g} && \forall\, t \in \mathcal{T},\ g \in \mathcal{G} \,:\, \mathrm{spend}^{\mathrm{cap}}_{g} > 0 \vee \neg \mathrm{is\_flexible}_{g} \end{align} diff --git a/tests/typesetting/golden/markdown.out b/tests/typesetting/golden/markdown.out index 225d22ac..b96a7226 100644 --- a/tests/typesetting/golden/markdown.out +++ b/tests/typesetting/golden/markdown.out @@ -323,6 +323,24 @@ p_{t,g} \le \mathrm{p}^{\mathrm{max}}_{g} \qquad \forall\, t \in \mathcal{T},\ g \sum_{t \in \mathcal{T},\ g \in \mathcal{G}} p_{t,g} \le \mathrm{budget} \qquad \text{where } \sum_{g \in \mathcal{G}} \mathrm{p}^{\mathrm{max}}_{g} \ge \mathrm{budget} ``` +**`counted`** + +```math +\theta_{b} \le \mathrm{budget} \qquad \forall\, b \in \mathcal{B} \,:\, \lvert \{ e \in \mathcal{E} \,:\, \mathrm{tech\_cap}_{b,e} > 0 \} \rvert \ge 2 +``` + +**`counted_here`** + +```math +\theta_{b} \le \mathrm{tech\_cap}_{b,e} \qquad \forall\, b \in \mathcal{B},\ e \in \mathcal{E} \,:\, \lvert \{ e' \in \mathcal{E} \,:\, \mathrm{tech\_cap}_{b,e'} > 0 \} \rvert \ge 2 +``` + +**`run_start`** + +```math +\mathit{slack}_{t} \le \mathrm{load}_{t,b} \qquad \forall\, t \in \mathcal{T},\ b \in \mathcal{B} \,:\, \mathrm{load}_{t,b} \text{ is defined} \wedge \neg \left( \mathrm{load}_{t - 1,b} \text{ is defined} \right) +``` + **`capped`** ```math diff --git a/tests/typesetting/golden/model.yaml b/tests/typesetting/golden/model.yaml index 43de710a..5302d214 100644 --- a/tests/typesetting/golden/model.yaml +++ b/tests/typesetting/golden/model.yaml @@ -241,6 +241,18 @@ constraints: dims: [] where: "sum(p_max, over=generator) >= budget" expression: sum(p) <= budget + counted: # a count of the coordinates a predicate admits, which reduces one dim away + dims: [bus] + where: "count(tech_cap > 0, over=technology) >= 2" + expression: theta <= budget + counted_here: # the same count along a dim the frame carries, so the set takes a primed dummy + dims: [bus, technology] + where: "count(tech_cap > 0, over=technology) >= 2" + expression: theta <= tech_cap + run_start: # a predicate read one coordinate back, which is false where the translation vacates + dims: [snapshot, bus] + where: "load AND NOT shift(load, along=snapshot, offset=1)" + expression: slack <= load capped: # an expressions: entry on a side, read by the name the file gave it dims: [snapshot, generator] where: "spend_cap > 0 OR NOT is_flexible" diff --git a/tests/typesetting/golden/typst.out b/tests/typesetting/golden/typst.out index 97bf944c..34b6afd9 100644 --- a/tests/typesetting/golden/typst.out +++ b/tests/typesetting/golden/typst.out @@ -105,6 +105,9 @@ $ upright("budgeted") & italic("spend")_(t) & <= upright("budget") & forall t in upright("margin") & p_(t,g) & <= upright("p")^(upright("max"))_(g) & forall t in cal(T), g in cal(G) colon upright("p")^(upright("max"))_(g) - upright("p")^(upright("min"))_(g) > frac(upright("cost")_(g), 2) \ upright("ramped") & italic("slack")_(t) & <= upright("load")_(t,b) & forall t in cal(T), b in cal(B) colon upright("load")_(t,b) - upright("load")_(t minus.square_(0) 1,b) <= upright("zone_cap")_(upright("zone_of")(b)) and upright("pos")(t) > 0 \ upright("covered") & sum_(t in cal(T), g in cal(G)) p_(t,g) & <= upright("budget") & upright("where ") sum_(g in cal(G)) upright("p")^(upright("max"))_(g) >= upright("budget") \ + upright("counted") & theta_(b) & <= upright("budget") & forall b in cal(B) colon abs({e in cal(E) colon upright("tech_cap")_(b,e) > 0}) >= 2 \ + upright("counted_here") & theta_(b) & <= upright("tech_cap")_(b,e) & forall b in cal(B), e in cal(E) colon abs({e' in cal(E) colon upright("tech_cap")_(b,e') > 0}) >= 2 \ + upright("run_start") & italic("slack")_(t) & <= upright("load")_(t,b) & forall t in cal(T), b in cal(B) colon upright("load")_(t,b) upright(" is defined") and not (upright("load")_(t - 1,b) upright(" is defined")) \ upright("capped") & p_(t,g) & <= upright("p")^(upright("max"))_(g) & forall t in cal(T), g in cal(G) colon upright("spend")^(upright("cap"))_(g) > 0 or not upright("is_flexible")_(g) $ == Definitions diff --git a/tests/typesetting/test_golden.py b/tests/typesetting/test_golden.py index a23e7220..81d0b982 100644 --- a/tests/typesetting/test_golden.py +++ b/tests/typesetting/test_golden.py @@ -150,11 +150,12 @@ def _rendered_trees() -> Iterator[object]: } #: A dataclass the walk steps *through* rather than renders: an arm has no -#: branch of its own — its ``when`` and ``value`` do — and a direction and the -#: relation it reads are the facts a node carries rather than nodes. None is a -#: member of any node union, so they are subtracted from what the tree walk -#: finds rather than added to what the vocabulary declares. -CARRIERS = {'CaseArm', 'Direction', 'Partition', 'RelationDeclaration'} +#: branch of its own — its ``when`` and ``value`` do — a direction and the +#: relation it reads are the facts a node carries rather than nodes, and a +#: ``Mask`` is the wrapper a leaf carries a predicate in. None is a member of +#: any node union, so they are subtracted from what the tree walk finds rather +#: than added to what the vocabulary declares. +CARRIERS = {'CaseArm', 'Direction', 'Mask', 'Partition', 'RelationDeclaration'} def test_the_golden_model_carries_every_node_kind_the_walk_renders(): diff --git a/tests/typesetting/test_walk.py b/tests/typesetting/test_walk.py index 8653c5f3..3471d9ec 100644 --- a/tests/typesetting/test_walk.py +++ b/tests/typesetting/test_walk.py @@ -13,7 +13,7 @@ from math_spec.errors import LanguageError from math_spec.piecewise import expand_piecewise -from math_spec.typesetting import FORMATS, SymbolTable, to_latex, typeset +from math_spec.typesetting import FORMATS, SymbolTable, to_latex, typeset, typeset_declaration from math_spec.typesetting.format import OPERATOR_NAMES from math_spec.typesetting.symbols import Symbols, _derive_name_symbol, chosen_expressions from math_spec.validation import to_spec @@ -835,3 +835,44 @@ def test_a_comparison_of_expressions_prints_as_the_arithmetic_it_is(name: Format p_max = fmt.subscript(fmt.superscript(fmt.upright('p'), fmt.upright('max')), ['g']) cost = fmt.subscript(fmt.upright('cost'), ['g']) assert f'{cost} {fmt.operators["le"]} {fmt.fraction(p_max, "2")}' in text + + +@EVERY_FORMAT +def test_a_count_prints_as_the_size_of_the_set_the_predicate_admits(name: FormatName, fmt: Format): + """A count is a cardinality over a set by comprehension, which is how a paper writes one.""" + model = override(DISPATCH_MODEL, **{'constraints.balance.where': 'count(p_max > 0, over=generator) >= 2'}) + text = typeset(model, name, legend=False) + p_max = fmt.subscript(fmt.superscript(fmt.upright('p'), fmt.upright('max')), ['g']) + counted = fmt.set_of( + f'g {fmt.operators["in"]} {"\\mathcal{G}" if name == "latex" or name == "markdown" else "cal(G)"}', + f'{p_max} {fmt.operators["gt"]} 0', + ) + assert f'{fmt.cardinality(counted)} {fmt.operators["ge"]} 2' in text + + +@EVERY_FORMAT +def test_a_translated_predicate_prints_at_the_index_it_reads(name: FormatName, fmt: Format): + """The translation shows at the leaf, as it does for arithmetic — it emits no operator of its own.""" + model = override( + DISPATCH_MODEL, **{'constraints.balance.where': 'load AND NOT shift(load, along=snapshot, offset=1)'} + ) + text = typeset(model, name, legend=False) + assert f'{fmt.subscript(fmt.upright("load"), ["t"])} {fmt.prose(" is defined")}' in text + moved = fmt.subscript(fmt.upright('load'), [f't {fmt.operators["minus"]} 1']) + assert f'{moved} {fmt.prose(" is defined")}' in text, 'the translated half reads one coordinate back' + + +def test_a_count_along_a_dim_the_frame_carries_takes_a_primed_dummy(): + """The set's index would otherwise shadow the frame's, and the two stand for different coordinates.""" + model = override( + DISPATCH_MODEL, + **{ + 'constraints.balance': { + 'dims': ['snapshot', 'generator'], + 'where': 'count(p_max > 0, over=generator) >= 2', + 'expression': 'p <= p_max', + } + }, + ) + line = typeset_declaration(model, 'balance', 'latex') + assert r"g' \in \mathcal{G}" in line, 'the counted dimension is quantified already, so the set takes a fresh index' From 2199a95a03429fc73205e52eb6f64e68280f2b8c Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 22 Sep 2026 08:30:04 +0000 Subject: [PATCH 2/2] fix(language): a count comparison the data cannot decide, and a keyword given twice, are refused at load MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit count(…) >= 0 and count(…) < 0 are settled by a count never being negative, and loaded. A keyword written twice in count() or shift() over a predicate was last-wins, where the arithmetic grammar refuses it. A bare count(), a count on the right of its comparison, and a keyword count() lacks each got a sentence about something else; each now names its rewrite. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_019tCWoetBmjF1LpbTbjQY29 --- docs/reference/language/expressions.md | 8 ++++-- src/math_spec/_where_parser.py | 29 ++++++++++++++----- src/math_spec/resolution.py | 36 ++++++++++++++++++++++-- tests/test_validation.py | 39 ++++++++++++++++++++++++-- 4 files changed, 98 insertions(+), 14 deletions(-) diff --git a/docs/reference/language/expressions.md b/docs/reference/language/expressions.md index 800ae2ea..207a608c 100644 --- a/docs/reference/language/expressions.md +++ b/docs/reference/language/expressions.md @@ -140,8 +140,8 @@ QUOTED ::= "'" chars "'" | '"' chars '"' | `expression OP expression` | arithmetic over parameters | Coordinate by coordinate, over every dimension either side carries ([arithmetic in a comparison](#arithmetic-in-a-comparison)). A side with no value at a coordinate compares false | | `position(name) OP i` | dimension | Where the row sits along the dimension's own order. `0` is first, and a negative number counts from the end | | `position(name, by=relation, within=c)` | dimension | The same, counted within each group the relation makes | -| `count(where_expr, over=name) OP i` | dimension | How many coordinates along the dimension the predicate admits ([counting what a predicate admits](#counting-what-a-predicate-admits)) | -| `shift(where_expr, along=name, offset=i)` | dimension | The predicate read `i` coordinates back, and false where that vacates | +| `count(where_expr, over=name) OP i` | a predicate | How many coordinates along the dimension the predicate admits ([counting what a predicate admits](#counting-what-a-predicate-admits)) | +| `shift(where_expr, along=name, offset=i)` | a predicate | The predicate read `i` coordinates back, and false where that vacates | | `AND` `OR` `NOT` | — | Case-insensitive. `NOT` binds tighter than `AND`, and `AND` tighter than `OR` | | `True` / `False` | — | `True` is the same as no `where`; `False` gives a declaration with no rows. A [case `when:`](named.md#the-rules-that-keep-the-cases-apart) may not fold to either | @@ -188,7 +188,9 @@ The count therefore states a fact about each group without naming the group. Counting along a dimension the predicate does not read is a load error. The comparison takes a whole number on the right. A count is a number of -coordinates, so a fraction and a parameter are both load errors. +coordinates, so a fraction and a parameter are both load errors, and so is a +comparison a count can never fail or never meet: `>= 0`, `< 0`, or any +negative number. ### Reading a predicate at the previous coordinate diff --git a/src/math_spec/_where_parser.py b/src/math_spec/_where_parser.py index 4b18cfd8..6c99399e 100644 --- a/src/math_spec/_where_parser.py +++ b/src/math_spec/_where_parser.py @@ -22,6 +22,7 @@ from math_spec._expression_parser import ARITHMETIC, NAME, ArithmeticNode, children, parse_text from math_spec._sealed import Sealed +from math_spec.errors import SchemaError from math_spec.program import ( And, BooleanLiteral, @@ -136,6 +137,18 @@ class UnresolvedComparisonNode: # --------------------------------------------------------------------------- +def _predicate_call(tokens: pp.ParseResults) -> UnresolvedPredicateCallNode: + """The call node, with a keyword given twice refused as the arithmetic grammar refuses it.""" + name, operand, *pairs = tokens + kwargs: dict[str, ArithmeticNode] = {} + for key, value in pairs: + if key in kwargs: + msg = f'{name}({key}=) is given twice. A keyword names one value; drop one of them.' + raise SchemaError(msg) + kwargs[key] = value + return UnresolvedPredicateCallNode(name, operand, kwargs) + + def _build_where_grammar() -> pp.ParserElement: """Build the pyparsing grammar for where strings. @@ -160,16 +173,18 @@ def _build_where_grammar() -> pp.ParserElement: kwarg = (name + pp.Suppress('=') + (quoted | ARITHMETIC)).set_parse_action(lambda t: (t[0], t[1])) def _call(head: pp.ParserElement) -> pp.ParserElement: - """``([, …])`` — the one shape whose operand is a predicate.""" + """``([, …])`` — the one shape whose operand is a predicate. + + ``count`` is spelled in the grammar rather than left to resolution, as + ``position`` is, because only the grammar can decide to read its + argument as a predicate. Every other predicate-taking call stands where + arithmetic cannot, so the comparison above it has already been tried + and the name is free. + """ return ( head + pp.Suppress('(') + where_expr + pp.ZeroOrMore(pp.Suppress(',') + kwarg) + pp.Suppress(')') - # pyrefly: ignore[implicit-any-lambda] - ).set_parse_action(lambda t: UnresolvedPredicateCallNode(t[0], t[1], dict(t[2:]))) + ).set_parse_action(_predicate_call) - # `count` is spelled here rather than left to resolution, as `position` is, - # because only the grammar can decide to read its argument as a predicate. - # Every other predicate-taking call stands where arithmetic cannot, so the - # comparison above it has already been tried and the name is free. count_call = _call(pp.CaselessKeyword('count')) predicate_call = _call(name.copy()) diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index bfb622c3..dbb7633e 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -819,6 +819,12 @@ def _predicate_call(self, node: UnresolvedPredicateCallNode) -> Predicate | Unre predicate for its dims asserts instead of refusing. """ context, found = self.context, len(self.errors) + if node.name == 'count': + self.errors.append( + f'{context}: count() answers a number, and a where is a predicate. Compare it: ' + f'count(, over=) .' + ) + return node if node.name != 'shift': self.errors.append( f"{context}: '{node.name}()' does not read a predicate. `shift` reads one and answers one, " @@ -883,6 +889,12 @@ def _count(self, node: UnresolvedCountNode) -> Predicate | UnresolvedWhereNode: f'Write count(…, over={over.name}) {node.op} .' ) return node + if (decided := _decided_count(node.op, value.value)) is not None: + self.errors.append( + f'{context}: count(…, over={over.name}) {node.op} {value} holds at {decided} coordinate, because ' + f'a count is never negative. Delete the comparison, or write the bound it means.' + ) + return node mask = Mask(operand) if over.name not in mask.dims: self.errors.append( @@ -948,6 +960,12 @@ def _expression_comparison(self, node: UnresolvedComparisonNode) -> ArithmeticCo if isinstance(side, ColumnNode | QuotedNode): self.errors.append(_not_arithmetic(context, side)) continue + if any(isinstance(n, FunctionCallNode) and n.name == 'count' for n in nodes(side)): + self.errors.append( + f'{context}: count() stands on the left of its comparison, and reads a predicate rather than ' + f'arithmetic. Write count(, over=) .' + ) + continue try: expanded = expand(side, ns.schema, context) except ValueError as e: @@ -1275,14 +1293,28 @@ def _kwargs_error( if missing: return f'{context}: {name}() needs {_listed([f"{key}=" for key in missing])}.' if extra := sorted(set(kwargs) - set(required)): + edge = ' A predicate is false where a translation vacates, so there is no edge to state.' return ( f'{context}: {name}() does not take {_listed([f"{key}=" for key in extra])}. ' - f'It takes {_listed([f"{key}=" for key in required])}, and nothing else: a predicate is false ' - f'where a translation vacates, so there is no edge to state.' + f'It takes {_listed([f"{key}=" for key in required])}, and nothing else.' + f'{edge if "edge" in extra else ""}' ) return None +def _decided_count(op: str, value: float) -> str | None: + """Whether comparing a count with *op* against *value* is settled by the count never being negative. + + Returns ``'every'`` where the comparison always holds, ``'no'`` where it + never does, and ``None`` where the data decides. + """ + if value < 0: + return 'every' if op in ('>', '>=', '!=') else 'no' + if value == 0 and op in ('>=', '<'): + return 'every' if op == '>=' else 'no' + return None + + def _listed(items: list[str]) -> str: """``'a'``, ``'a' and 'b'``, ``'a', 'b' and 'c'`` — one rule, so every message reads the same.""" quoted = [f"'{item}'" for item in items] diff --git a/tests/test_validation.py b/tests/test_validation.py index e028c915..70d03f43 100644 --- a/tests/test_validation.py +++ b/tests/test_validation.py @@ -762,9 +762,24 @@ def test_a_shape_the_language_admits(self, where): ), pytest.param( 'count(flag, over=g, by=lk) >= 2', - ("does not take 'by='", "It takes 'over='"), + ("does not take 'by='", "It takes 'over=', and nothing else."), id='a-count-with-a-keyword-it-lacks', ), + pytest.param( + 'count(flag, over=g, over=h) >= 2', + ('count(over=) is given twice',), + id='a-count-with-a-keyword-given-twice', + ), + pytest.param( + 'count(flag, over=g)', + ('count() answers a number', 'count(, over=) '), + id='a-count-standing-as-a-predicate', + ), + pytest.param( + '2 <= count(flag, over=g)', + ('count() stands on the left of its comparison',), + id='a-count-on-the-right', + ), pytest.param( 'count(flag, over=c) >= 2', ('names the dimension the coordinates are counted along',), @@ -785,11 +800,31 @@ def test_a_shape_the_language_admits(self, where): ('a count is a whole number of coordinates',), id='a-count-against-a-parameter', ), + pytest.param( + 'count(flag, over=g) >= -1', + ('a count is never negative',), + id='a-count-against-a-negative-number', + ), + pytest.param( + 'count(flag, over=g) >= 0', + ('holds at every coordinate', 'a count is never negative'), + id='a-count-at-least-zero', + ), + pytest.param( + 'count(flag, over=g) < 0', + ('holds at no coordinate', 'a count is never negative'), + id='a-count-below-zero', + ), pytest.param( 'shift(flag, along=g, offset=1, edge=0)', - ("does not take 'edge='", 'a predicate is false where a translation vacates'), + ("does not take 'edge='", 'A predicate is false where a translation vacates'), id='a-translated-predicate-with-an-edge', ), + pytest.param( + 'shift(flag, along=g, offset=1, offset=2)', + ('shift(offset=) is given twice',), + id='a-translation-with-a-keyword-given-twice', + ), pytest.param( 'shift(flag, along=g)', ("shift() needs 'offset='",),