diff --git a/docs/reference/language/expressions.md b/docs/reference/language/expressions.md index 57b3e5fa..207a608c 100644 --- a/docs/reference/language/expressions.md +++ b/docs/reference/language/expressions.md @@ -115,30 +115,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` | 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 | The dimensions of the mask must not exceed the frame it sits in. A bare name that is not declared is a load error. @@ -148,6 +153,65 @@ 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, 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 + +`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 c2f4a61c..77dda36c 100644 --- a/docs/reference/notation.md +++ b/docs/reference/notation.md @@ -764,6 +764,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 88bec8d5..e5fa2189 100644 --- a/docs/reference/reading.md +++ b/docs/reference/reading.md @@ -134,6 +134,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..6c99399e 100644 --- a/src/math_spec/_where_parser.py +++ b/src/math_spec/_where_parser.py @@ -14,13 +14,15 @@ 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.errors import SchemaError from math_spec.program import ( And, BooleanLiteral, @@ -33,7 +35,7 @@ ) if TYPE_CHECKING: - from collections.abc import Callable + from collections.abc import Callable, Mapping # --------------------------------------------------------------------------- # AST nodes @@ -72,6 +74,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 +122,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 @@ -101,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. @@ -122,14 +170,44 @@ 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. + + ``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(')') + ).set_parse_action(_predicate_call) + + 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 +269,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 63e0403c..b4f48d4f 100644 --- a/src/math_spec/exclusivity.py +++ b/src/math_spec/exclusivity.py @@ -26,6 +26,7 @@ And, ArithmeticComparison, BooleanLiteral, + CountComparison, DimensionComparison, DimensionPosition, ExpressionComparison, @@ -37,6 +38,7 @@ RelationComparison, RelationDefined, RelationPairComparison, + TranslatedPredicate, TypedPredicate, VariableDefined, ) @@ -260,6 +262,18 @@ def _observe( """ if isinstance(node, ArithmeticComparison | ExpressionComparison): raise Undecidable(_expression_rewrite(node)) + 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): @@ -299,6 +313,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) @@ -493,6 +511,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 226f9c4c..30dceb62 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 @@ -273,11 +273,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 af43ca40..21141139 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -53,6 +53,7 @@ 'ConstraintDeclaration', 'ConstraintSense', 'Contiguous', + 'CountComparison', 'Curved', 'Derivation', 'DimensionComparison', @@ -103,6 +104,7 @@ 'SosDeclaration', 'Sum', 'Translate', + 'TranslatedPredicate', 'TypedPredicate', 'Variable', 'VariableAbsence', @@ -1285,6 +1287,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 @@ -1320,6 +1357,8 @@ class Or: | RelationComparison | RelationPairComparison | RelationDefined + | CountComparison + | TranslatedPredicate | Not | And | Or @@ -1339,6 +1378,8 @@ class Or: | RelationComparison | RelationPairComparison | RelationDefined + | CountComparison + | TranslatedPredicate ) #: The boolean connectives — the only where nodes carrying other where nodes, @@ -1399,6 +1440,8 @@ def _atom_dims(atom: TypedPredicate) -> frozenset[str]: | ArithmeticComparison() | ParameterDefined() | VariableDefined() + | CountComparison() + | TranslatedPredicate() ): return frozenset(atom.dims) case DimensionComparison(): @@ -1432,6 +1475,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 119c471f..dbb7633e 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, ) @@ -749,6 +753,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): @@ -799,6 +807,104 @@ 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 == '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, " + 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 + 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( + 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. @@ -854,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: @@ -1167,6 +1279,50 @@ 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)): + 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.' + 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] + if len(quoted) <= 1: + return quoted[0] if quoted else 'nothing' + return f'{", ".join(quoted[:-1])} and {quoted[-1]}' + + def _is_number(side: ArithmeticNode) -> bool: """Whether *side* is arithmetic over literals alone — a value the language can fold, and a where may not test.""" return all(isinstance(n, NumberNode | UnaryOperatorNode | BinaryOperatorNode) for n in nodes(side)) diff --git a/src/math_spec/typesetting/walk.py b/src/math_spec/typesetting/walk.py index 75a5c52e..160e1a2f 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -43,6 +43,7 @@ BooleanLiteral, Check, Contiguous, + CountComparison, Curved, DimensionComparison, DimensionPosition, @@ -59,6 +60,7 @@ RelationComparison, RelationDefined, RelationPairComparison, + TranslatedPredicate, VariableDefined, ) from math_spec.typesetting.format import Entry, Glossary, Line, OperatorName @@ -86,6 +88,7 @@ AlignedComparison = ( ParameterComparison | ArithmeticComparison + | CountComparison | DimensionComparison | DimensionPosition | RelationComparison @@ -606,6 +609,10 @@ def _where(self, node: Predicate, ctx: _Context) -> tuple[str, int]: msg = 'a lowered comparison reached the typesetter; it prints the resolved tree, which lowering rebuilds.' raise AssertionError(msg) + 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 @@ -654,6 +661,12 @@ def sides(self, node: AlignedComparison, ctx: _Context) -> tuple[str, str]: elif isinstance(node, RelationPairComparison): left = self._value_read(node.name, node.column, ctx) right = self._value_read(node.other, node.other_column, ctx) + elif 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) + ) + left, right = self.format.cardinality(counted), self._number(node.value) else: assert_never(node) return left, f'{self._op(_PREDICATES[node.op])} {right}' diff --git a/tests/test_lowering.py b/tests/test_lowering.py index 99185599..7c495c8b 100644 --- a/tests/test_lowering.py +++ b/tests/test_lowering.py @@ -28,6 +28,7 @@ BooleanLiteral, Cases, Constant, + CountComparison, DimensionComparison, DimensionDeclaration, Direction, @@ -432,6 +433,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_assumptions_carry_the_file_s_entries_and_the_curves_behind_them(): """One mapping holds every fact about the data, so a consumer binding it has one loop and one refusal. diff --git a/tests/test_validation.py b/tests/test_validation.py index b1790d02..70d03f43 100644 --- a/tests/test_validation.py +++ b/tests/test_validation.py @@ -729,6 +729,172 @@ def test_a_lone_case_comparing_expressions_is_refused_too(self): assert "case 'wide' cannot be told apart before the data arrives: it compares expressions" 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=', 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',), + 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( + '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'), + 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='",), + 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 80745fd2..16a30d21 100644 --- a/tests/typesetting/golden/latex.out +++ b/tests/typesetting/golden/latex.out @@ -130,6 +130,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} \\ \text{fuel\_curve\_chord} && \mathit{fuel}_{t} \cdot \left( \mathrm{bp\_x}_{a} - \mathrm{bp\_x}_{a \boxminus_{0} 1} \right) & \ge \left( \mathrm{bp\_y}_{a} - \mathrm{bp\_y}_{a \boxminus_{0} 1} \right) \cdot \left( p^{\mathrm{bp}}_{t} - \mathrm{bp\_x}_{a} \right) + \mathrm{bp\_y}_{a} \cdot \left( \mathrm{bp\_x}_{a} - \mathrm{bp\_x}_{a \boxminus_{0} 1} \right) && \forall\, t \in \mathcal{T},\ a \in \mathcal{A} \,:\, \mathrm{fuel}^{\mathrm{curve,points}}_{a} \wedge \neg \mathrm{fuel}^{\mathrm{curve,starts}}_{a} \\ \text{fuel\_curve\_domain\_lo} && p^{\mathrm{bp}}_{t} & \ge \mathrm{bp\_x}_{a} && \forall\, t \in \mathcal{T},\ a \in \mathcal{A} \,:\, \mathrm{fuel}^{\mathrm{curve,starts}}_{a} \\ diff --git a/tests/typesetting/golden/markdown.out b/tests/typesetting/golden/markdown.out index 2b26c904..69dcfa7a 100644 --- a/tests/typesetting/golden/markdown.out +++ b/tests/typesetting/golden/markdown.out @@ -335,6 +335,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 15004aab..86789caa 100644 --- a/tests/typesetting/golden/model.yaml +++ b/tests/typesetting/golden/model.yaml @@ -273,6 +273,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 9771b28f..04aba50e 100644 --- a/tests/typesetting/golden/typst.out +++ b/tests/typesetting/golden/typst.out @@ -117,6 +117,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) \ upright("fuel_curve_chord") & italic("fuel")_(t) dot (upright("bp_x")_(a) - upright("bp_x")_(a minus.square_(0) 1)) & >= (upright("bp_y")_(a) - upright("bp_y")_(a minus.square_(0) 1)) dot (p^(upright("bp"))_(t) - upright("bp_x")_(a)) + upright("bp_y")_(a) dot (upright("bp_x")_(a) - upright("bp_x")_(a minus.square_(0) 1)) & forall t in cal(T), a in cal(A) colon upright("fuel")^(upright("curve,points"))_(a) and not upright("fuel")^(upright("curve,starts"))_(a) \ upright("fuel_curve_domain_lo") & p^(upright("bp"))_(t) & >= upright("bp_x")_(a) & forall t in cal(T), a in cal(A) colon upright("fuel")^(upright("curve,starts"))_(a) \ diff --git a/tests/typesetting/test_golden.py b/tests/typesetting/test_golden.py index 5229a1c4..e28d3f7a 100644 --- a/tests/typesetting/test_golden.py +++ b/tests/typesetting/test_golden.py @@ -154,11 +154,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 ef87e10e..8d306192 100644 --- a/tests/typesetting/test_walk.py +++ b/tests/typesetting/test_walk.py @@ -837,6 +837,47 @@ def test_a_comparison_of_expressions_prints_as_the_arithmetic_it_is(name: Format 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' + + @EVERY_FORMAT def test_an_assumption_prints_under_its_own_heading(name: FormatName, fmt: Format): """What the data is held to prints with the math, because a reader checking it reads the same document."""