diff --git a/docs/reference/language/expressions.md b/docs/reference/language/expressions.md index 28f029e2..db43727c 100644 --- a/docs/reference/language/expressions.md +++ b/docs/reference/language/expressions.md @@ -161,8 +161,8 @@ A `where:` is a boolean mask, and true means "this coordinate exists". ```text where_expr ::= atom | "NOT" where_expr | where_expr ("AND"|"OR") where_expr | "(" where_expr ")" -atom ::= NAME | NAME COMPARATOR value | POSITION COMPARATOR INTEGER - | "True" | "False" +atom ::= NAME | NAME COMPARATOR value | NAME "in" "[" [ value { "," value } ] "]" + | POSITION COMPARATOR INTEGER | "True" | "False" COMPARATOR ::= "<=" | ">=" | "==" | "!=" | "<" | ">" value ::= NUMBER | QUOTED | NAME_OR_STRING POSITION ::= "position" "(" NAME [ "," "by" "=" NAME ] ")" @@ -176,6 +176,7 @@ QUOTED ::= "'" chars "'" | '"' chars '"' | `name` (bare) | dimension | load error: it is true everywhere, so it reads as a condition and is not one. Compare it instead | | `name OP value` | parameter | element-wise; a null compares false. The right-hand side is a literal number, or a bare name read as a string coordinate | | `name OP value` | dimension | a filter on the frame's own coordinate column | +| `name in [v, …]` | parameter, dimension, lookup | keeps the rows whose value is one of a **set of literals** — the same three left-hand sides a comparison takes, one `isin` filter instead of an `==` OR chain. Each element is dtype-checked like `==`, so a date is quoted per element | | `name` (bare) | lookup | defined: the label maps somewhere. A lookup may be [partial](dimensions.md#lookups), and this is how a declaration asks for the labels that do map | | `name OP value` | lookup | a filter on the lookup's column of its `over` dimension's index — which therefore has to be in the frame. A null value is **false**, whatever the comparator | | `name OP name` | two lookups | the one comparison whose both sides are structure. Legal only where both map out of the **same** dimension _and_ into the **same** one — `from != to` excludes a self-loop | @@ -221,7 +222,19 @@ dates: a `datetime` dimension compared to a number is compared against the **epoch**, so `snapshot > 0` would silently mean "after 1970-01-01". That is a load error naming the fix. A datetime boundary is a quoted ISO date — `snapshot > '2030-01-01'`, or `'2030-01-01T06:00'` with a time. Calendar -arithmetic, resampling and timezone conversion stay data prep. +arithmetic, resampling and timezone conversion stay data prep. Membership +checks each element the same way, so a set of dates is a list of quoted ISO +strings and a `str` column takes quoted labels; a `float` column may be tested +against a set of floats, exact equality and all, exactly as `==` permits it. + +**A membership list is literals only, and it is not empty.** `carrier in []` +matches nothing, which is `where: "False"` said obscurely — a load error names +that rewrite. A repeated element selects nothing extra and is a load error too. +A **declared** name among the elements is a near miss the same way a comparison's +right-hand side is: selecting by data on the right is +[data-driven membership](https://github.com/energy-models/math-spec/issues/258), +so the message points there, or to precomputing the test as a `bool` parameter. +Negation is the existing `NOT` — there is no `not in`. **`position(dim)` converts a dimension to where the row sits along it**, so a boundary clause survives the index being relabelled: diff --git a/docs/reference/notation.md b/docs/reference/notation.md index 340aa7c6..3b91b3b0 100644 --- a/docs/reference/notation.md +++ b/docs/reference/notation.md @@ -426,6 +426,19 @@ northern: $$\mathit{slack}_{t} \le \mathrm{load}_{t,b} \qquad \forall\thinspace t \in \mathcal{T},\enspace b \in \mathcal{B} \thinspace:\thinspace \mathrm{zone\_of}(b) = \text{'}\mathrm{north}\text{'} \wedge \mathrm{zone\_of}(b) \neq \mathrm{area\_of}(b) \wedge \mathrm{zone\_of}(b) \text{ is defined}$$ +#### `selected` + +set membership per kind: a dimension's coordinates, a lookup's values, a parameter's numbers + +```yaml +selected: + foreach: [snapshot, generator, bus] + where: "snapshot in [0, 3] OR zone_of in ['north', 'south'] OR min_up in [2, 3]" + expression: p <= load +``` + +$$p_{t,g} \le \mathrm{load}_{t,b} \qquad \forall\thinspace t \in \mathcal{T},\enspace g \in \mathcal{G},\enspace b \in \mathcal{B} \thinspace:\thinspace t \in \{0,\enspace 3\} \vee \mathrm{zone\_of}(b) \in \{\text{'}\mathrm{north}\text{'},\enspace \text{'}\mathrm{south}\text{'}\} \vee \mathrm{min\_up}_{g} \in \{2,\enspace 3\}$$ + #### `efficiency` a Greek-named parameter, which is given — so the convention wins and it prints as the word diff --git a/src/math_spec/dimensions.py b/src/math_spec/dimensions.py index 5a0640a4..a63580c1 100644 --- a/src/math_spec/dimensions.py +++ b/src/math_spec/dimensions.py @@ -48,9 +48,11 @@ from math_spec.resolution import Namespace, expression_of, where_of from math_spec.where_parser import ( DimensionComparisonNode, + DimensionMembershipNode, DimensionPositionNode, ParameterComparisonNode, ParameterDefinedNode, + ParameterMembershipNode, VariableDefinedNode, WhereNode, _atom_dims, @@ -483,7 +485,7 @@ def _check_where_dims( for atom in atoms(node): if not (outside := sorted(_atom_dims(atom, name_dims) - frame)): continue - if isinstance(atom, (ParameterDefinedNode, ParameterComparisonNode)): + if isinstance(atom, (ParameterDefinedNode, ParameterComparisonNode, ParameterMembershipNode)): raise DimensionError( f"{context}: where-parameter '{atom.name}' has dims " f'{outside} outside the frame {sorted(frame)}. Reducing ' @@ -496,7 +498,7 @@ def _check_where_dims( f'reducing over an unlisted dim would silently widen it — say which ' f'reduction you mean.' ) - if isinstance(atom, (DimensionComparisonNode, DimensionPositionNode)): + if isinstance(atom, (DimensionComparisonNode, DimensionMembershipNode, DimensionPositionNode)): raise DimensionError( f"{context}: where-comparison on dimension '{atom.name}', which is not in the frame {sorted(frame)}." ) diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index 5c2993b5..dceeea0d 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -14,7 +14,7 @@ import datetime import re -from typing import TYPE_CHECKING, assert_never +from typing import TYPE_CHECKING, assert_never, cast from math_spec.errors import LanguageError from math_spec.expansion import parse_and_expand @@ -49,16 +49,20 @@ AndNode, BooleanLiteralNode, DimensionComparisonNode, + DimensionMembershipNode, DimensionPositionNode, LookupComparisonNode, LookupDefinedNode, + LookupMembershipNode, LookupPairComparisonNode, NotNode, OrNode, ParameterComparisonNode, ParameterDefinedNode, + ParameterMembershipNode, TypedPredicateNode, UnresolvedComparisonNode, + UnresolvedMembershipNode, UnresolvedNameNode, UnresolvedPositionNode, VariableDefinedNode, @@ -616,28 +620,44 @@ def resolve_where( def _typed_literal( - node: UnresolvedComparisonNode, + name: str, + value: float | str, + quoted: bool, dtype: str, context: str, errors: list[str], ) -> float | str | datetime.date | None: - """The comparison's literal, checked against the declared dtype. + """One literal, checked against the declared dtype of the name it is tested against. + + The one home for the dtype rule a where-comparison and a where-membership + both run: a comparison passes its single value, a membership each element + of its list. Getting it wrong is silent: polars reads a datetime column + against an integer as an epoch offset, so ``snapshot > 0`` drops every + coordinate before 1970 without a word (#460). Returns ``None`` once it has + recorded an error, so the caller leaves the node unresolved. + + Args: + name: The declared name the literal is tested against. + value: The literal — a number, or a string label. + quoted: Whether it arrived in quotes, which shapes the rewrite the + message names. + dtype: The declared dtype of *name*. + context: Where a message locates itself. + errors: Collected problems, appended to on a mismatch. - Getting it wrong is silent: polars reads a datetime column against an - integer as an epoch offset, so ``snapshot > 0`` drops every coordinate - before 1970 without a word (#460). Returns ``None`` once it has recorded - an error, so the caller leaves the node unresolved. + Returns: + The literal in the dtype's own type — a :class:`datetime.date` for a + datetime dimension — or ``None`` on a mismatch. """ - value = node.value text = isinstance(value, str) if dtype == 'datetime': if not text: errors.append( - f"{context}: '{node.name}' is a datetime dimension, so comparing it to " - f'{value!r} compares against the epoch — {node.name} > 0 means "after ' - f'1970-01-01", not what it looks like. Quote an ISO date instead: ' - f"{node.name} {node.op} '2030-01-01'." + f"{context}: '{name}' is a datetime dimension, so comparing it to " + f'{value!r} compares against the epoch — {name} > 0 means "after ' + f'1970-01-01", not what it looks like. Quote an ISO date instead, ' + f"e.g. '2030-01-01'." ) return None try: @@ -648,22 +668,21 @@ def _typed_literal( ) except ValueError: errors.append( - f"{context}: '{node.name}' is a datetime dimension and {value!r} is not an " + f"{context}: '{name}' is a datetime dimension and {value!r} is not an " f"ISO date. Write '2030-01-01' or '2030-01-01T06:00'." ) return None if dtype == 'str' and not text: errors.append( - f"{context}: '{node.name}' has dtype 'str', so comparing it to the number " - f'{value!r} matches no label. Quote it if it is one: {node.name} {node.op} ' - f"'{value:g}'." + f"{context}: '{name}' has dtype 'str', so comparing it to the number " + f'{value!r} matches no label. Quote it if it is one, e.g. {f"{value:g}"!r}.' ) return None if dtype in ('int', 'float', 'bool') and text: + fix = 'Drop the quotes if it is a number.' if quoted else 'Write it as a number if it is one.' errors.append( - f"{context}: '{node.name}' has dtype '{dtype}', so comparing it to the string " - f'{value!r} matches nothing. Drop the quotes if it is a number.' + f"{context}: '{name}' has dtype '{dtype}', so comparing it to the string {value!r} matches nothing. {fix}" ) return None return value @@ -729,6 +748,98 @@ def _lookup_pair_error(context: str, node: UnresolvedComparisonNode, other: str, return None +def _variable_where_error(context: str, name: str) -> str: + """Why a variable may not sit where a where reads a value — the one home for the sentence a comparison and a membership share.""" + return ( + f"{context}: where references variable '{name}'. A where mask is built before " + f'variables exist — it may test parameters and dimension coordinates only.' + ) + + +def _declared_element_error(context: str, name: str, value: str, kind: str) -> str: + """Why a membership list may not name a declaration among its literals.""" + return ( + f"{context}: '{name} in [...]' names {value!r}, a declared " + f'{kind}, but a membership list takes literal labels only. A declared name on the ' + f'right is data-driven membership (#258); quote {value!r} to keep it a fixed label, ' + f'or precompute the test as a bool parameter and mask on that.' + ) + + +def _literal_repr(value: float | str | datetime.date) -> str: + """A typed where-literal as a refusal shows it — numbers via ``:g``, dates as ISO, labels quoted.""" + if isinstance(value, str): + return repr(value) + if isinstance(value, datetime.date): + return value.isoformat() + return f'{value:g}' + + +def _first_duplicate( + values: tuple[float | str | datetime.date, ...], +) -> float | str | datetime.date | None: + """The first element that repeats an earlier one, or ``None`` — a repeat selects nothing extra.""" + seen: set[float | str | datetime.date] = set() + for value in values: + if value in seen: + return value + seen.add(value) + return None + + +def _resolve_membership(node: UnresolvedMembershipNode, ns: Namespace, context: str, errors: list[str]) -> WhereNode: + """Type ``name in [l1, l2, …]`` — the set form of a where-comparison. + + The list carries literals only: empty selects no row and is refused for the + always-false mask it hides; a repeat says nothing; a declared name among the + elements is data-driven membership, which is #258. The left-hand side is the + same three kinds a comparison takes, each element checked against its dtype + by the one :func:`_typed_literal` a comparison runs. + """ + if not node.elements: + errors.append( + f"{context}: '{node.name} in []' matches nothing — an always-false mask is a " + f'declaration with no rows. Write where: "False" if that is what is meant.' + ) + return node + + for value, quoted in node.elements: + if not quoted and isinstance(value, str) and (element_kind := ns.kind(value)) is not None: + errors.append(_declared_element_error(context, node.name, value, element_kind)) + return node + + kind = ns.kind(node.name) + typed: list[float | str | datetime.date] = [] + if kind in ('parameter', 'dimension', 'lookup'): + dtype = ns.dtypes[node.name] + for value, quoted in node.elements: + one = _typed_literal(node.name, value, quoted, dtype, context, errors) + if one is None: + return node + typed.append(one) + if (duplicate := _first_duplicate(tuple(typed))) is not None: + errors.append( + f"{context}: '{node.name} in [...]' lists {_literal_repr(duplicate)} more than " + f'once, which selects nothing extra. Drop the duplicate.' + ) + return node + + match kind: + case 'parameter': + assert not any(isinstance(value, datetime.date) for value in typed) + return ParameterMembershipNode(node.name, cast('tuple[float | str, ...]', tuple(typed))) + case 'dimension': + return DimensionMembershipNode(node.name, tuple(typed)) + case 'lookup': + return LookupMembershipNode(node.name, ns.over_of(node.name), tuple(typed)) + case 'variable': + errors.append(_variable_where_error(context, node.name)) + return node + case _: + errors.append(ns._unknown(node.name, context, allow_dims=True)) + return node + + def _resolve_position(node: UnresolvedPositionNode, ns: Namespace, context: str, errors: list[str]) -> WhereNode: """Type ``position(dim[, by=lookup]) i``. @@ -818,7 +929,7 @@ def _resolve_where( kind = ns.kind(node.name) if kind in ('parameter', 'dimension', 'lookup'): - typed = _typed_literal(node, ns.dtypes[node.name], context, errors) + typed = _typed_literal(node.name, node.value, node.quoted, ns.dtypes[node.name], context, errors) if typed is None: return node value = typed @@ -832,16 +943,15 @@ def _resolve_where( case 'lookup': return LookupComparisonNode(node.name, ns.over_of(node.name), node.op, value) case 'variable': - errors.append( - f"{context}: where references variable '{node.name}'. A where " - f'mask is built before variables exist — it may test parameters ' - f'and dimension coordinates only.' - ) + errors.append(_variable_where_error(context, node.name)) return node case _: errors.append(ns._unknown(node.name, context, allow_dims=True)) return node + if isinstance(node, UnresolvedMembershipNode): + return _resolve_membership(node, ns, context, errors) + if isinstance(node, NotNode): return NotNode(_resolve_where(node.operand, ns, context, errors, self_variable)) if isinstance(node, AndNode): diff --git a/src/math_spec/typesetting/format.py b/src/math_spec/typesetting/format.py index f32d7cb7..a43c78ba 100644 --- a/src/math_spec/typesetting/format.py +++ b/src/math_spec/typesetting/format.py @@ -164,6 +164,10 @@ def cardinality(self, inner: str) -> str: """How many members a set has: ``|T|``. A fence, so not an infix entry in :data:`OPERATOR_NAMES`.""" ... + def set_braces(self, inner: str) -> str: + """A set literal's braces around its elements: ``{a, b}``. A fence, like :meth:`cardinality`.""" + ... + def fraction(self, numerator: str, denominator: str) -> str: ... def summation(self, domain: str, body: str) -> str: ... diff --git a/src/math_spec/typesetting/latex.py b/src/math_spec/typesetting/latex.py index afe6e324..ce09cfc3 100644 --- a/src/math_spec/typesetting/latex.py +++ b/src/math_spec/typesetting/latex.py @@ -106,6 +106,9 @@ def parenthesise(self, inner: str) -> str: def cardinality(self, inner: str) -> str: return rf'\lvert {inner} \rvert' + def set_braces(self, inner: str) -> str: + return rf'\{{{inner}\}}' + def fraction(self, numerator: str, denominator: str) -> str: return rf'\frac{{{numerator}}}{{{denominator}}}' diff --git a/src/math_spec/typesetting/typst.py b/src/math_spec/typesetting/typst.py index 90e768d8..1454ba4e 100644 --- a/src/math_spec/typesetting/typst.py +++ b/src/math_spec/typesetting/typst.py @@ -108,6 +108,9 @@ def parenthesise(self, inner: str) -> str: def cardinality(self, inner: str) -> str: return f'abs({inner})' + def set_braces(self, inner: str) -> str: + return f'{{{inner}}}' + def fraction(self, numerator: str, denominator: str) -> str: return f'frac({numerator}, {denominator})' diff --git a/src/math_spec/typesetting/walk.py b/src/math_spec/typesetting/walk.py index 2905c1f8..b25d3f96 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -40,14 +40,17 @@ AndNode, BooleanLiteralNode, DimensionComparisonNode, + DimensionMembershipNode, DimensionPositionNode, LookupComparisonNode, LookupDefinedNode, + LookupMembershipNode, LookupPairComparisonNode, NotNode, OrNode, ParameterComparisonNode, ParameterDefinedNode, + ParameterMembershipNode, UnresolvedWhereNode, VariableDefinedNode, WhereNode, @@ -477,11 +480,21 @@ def _where(self, node: WhereNode, ctx: _Context) -> tuple[str, int]: left = ctx.indexed(self.symbols.name[node.name], dims) return f'{left} {self.op(_PREDICATES[node.op])} {self.literal(node.value)}', 2 + if isinstance(node, ParameterMembershipNode): + dims = list(self.schema.parameters[node.name].dims) + left = ctx.indexed(self.symbols.name[node.name], dims) + return f'{left} {self.op("in")} {self.set_of(node.values)}', 2 + if isinstance(node, DimensionComparisonNode): if isinstance(node.value, (int, float)): self.numeric_coordinates.add(node.name) return f'{ctx.subscript(node.name)} {self.op(_PREDICATES[node.op])} {self.literal(node.value)}', 2 + if isinstance(node, DimensionMembershipNode): + if any(isinstance(value, (int, float)) for value in node.values): + self.numeric_coordinates.add(node.name) + return f'{ctx.subscript(node.name)} {self.op("in")} {self.set_of(node.values)}', 2 + if isinstance(node, DimensionPositionNode): grouping = None if node.by is None else self.lookup(node.by, ctx.subscript(node.name)) place = self.position(ctx.subscript(node.name), grouping) @@ -492,6 +505,10 @@ def _where(self, node: WhereNode, ctx: _Context) -> tuple[str, int]: applied = self.lookup(node.name, ctx.subscript(node.over)) return f'{applied} {self.op(_PREDICATES[node.op])} {self.literal(node.value)}', 2 + if isinstance(node, LookupMembershipNode): + applied = self.lookup(node.name, ctx.subscript(node.over)) + return f'{applied} {self.op("in")} {self.set_of(node.values)}', 2 + if isinstance(node, LookupPairComparisonNode): index = ctx.subscript(node.over) left = self.lookup(node.name, index) @@ -522,6 +539,10 @@ def _where(self, node: WhereNode, ctx: _Context) -> tuple[str, int]: def literal(self, value: float | str | datetime.date) -> str: return self.number(value) if isinstance(value, (int, float)) else self.format.quoted(str(value)) + def set_of(self, values: tuple[float | str | datetime.date, ...]) -> str: + """The set a membership tests against — ``{v1, v2, …}``, each element a :meth:`literal`.""" + return self.format.set_braces(self.format.joined([self.literal(v) for v in values], '')) + def position(self, index: str, grouping: str | None) -> str: """``position(dim)`` applied to the row, *grouping* as a subscript — as an argument it read as a second position.""" self.positions.add('grouped' if grouping is not None else 'plain') diff --git a/src/math_spec/where_parser.py b/src/math_spec/where_parser.py index 66daa8eb..8e668faa 100644 --- a/src/math_spec/where_parser.py +++ b/src/math_spec/where_parser.py @@ -57,6 +57,20 @@ class UnresolvedComparisonNode: quoted: bool = False +@dataclass(frozen=True) +class UnresolvedMembershipNode: + """``name in [l1, l2, …]`` before the name is checked. ``resolution.py`` types it. + + Each element carries the scalar comparison's ``quoted`` flag, so resolution + refuses a bare word that names a declaration exactly as the scalar + right-hand side does — a set of literals is the whole construct; a name + among them is data-driven membership, which is #258. + """ + + name: str + elements: tuple[tuple[float | str, bool], ...] + + @dataclass(frozen=True) class UnresolvedPositionNode: """``position(dim) i`` before the name is checked. @@ -100,6 +114,14 @@ class ParameterComparisonNode: value: float | str +@dataclass(frozen=True) +class ParameterMembershipNode: + """Keep the rows where a parameter's value is one of a set of literals.""" + + name: str + values: tuple[float | str, ...] + + @dataclass(frozen=True) class DimensionComparisonNode: """Compare a dimension's own coordinates against a literal.""" @@ -109,6 +131,14 @@ class DimensionComparisonNode: value: float | str | datetime.date +@dataclass(frozen=True) +class DimensionMembershipNode: + """Keep the coordinates a dimension's own labels put in a set of literals.""" + + name: str + values: tuple[float | str | datetime.date, ...] + + @dataclass(frozen=True) class DimensionPositionNode: """Compare where a row sits along a dimension against a position — ``position(snapshot) == 0``. @@ -140,6 +170,19 @@ class LookupComparisonNode: value: float | str | datetime.date +@dataclass(frozen=True) +class LookupMembershipNode: + """Keep the rows a lookup's values put in a set of literals — ``period_of in [2030, 2040]``. + + ``over`` is the dimension the lookup maps out of, copied off the + declaration during resolution so every consumer reads it here. + """ + + name: str + over: str + values: tuple[float | str | datetime.date, ...] + + @dataclass(frozen=True) class LookupPairComparisonNode: """Compare two lookups over one dimension — ``from != to``. @@ -191,13 +234,17 @@ class OrNode: BooleanLiteralNode | UnresolvedNameNode | UnresolvedComparisonNode + | UnresolvedMembershipNode | UnresolvedPositionNode | DimensionPositionNode | ParameterDefinedNode | VariableDefinedNode | ParameterComparisonNode + | ParameterMembershipNode | DimensionComparisonNode + | DimensionMembershipNode | LookupComparisonNode + | LookupMembershipNode | LookupPairComparisonNode | LookupDefinedNode | NotNode @@ -209,18 +256,21 @@ class OrNode: #: left-hand side is still a name the schema has not been asked about. The #: expression side has :data:`~math_spec.expression_parser.UnresolvedNode` for #: the same reason, and a pass meeting either ran before resolution. -UnresolvedWhereNode = UnresolvedNameNode | UnresolvedComparisonNode | UnresolvedPositionNode +UnresolvedWhereNode = UnresolvedNameNode | UnresolvedComparisonNode | UnresolvedMembershipNode | UnresolvedPositionNode #: Every predicate resolution has typed: it names a declaration and the kind is #: settled. Resolution passes these straight through, having nothing left to #: decide about them. TypedPredicateNode = ( ParameterComparisonNode + | ParameterMembershipNode | ParameterDefinedNode | VariableDefinedNode | DimensionComparisonNode + | DimensionMembershipNode | DimensionPositionNode | LookupComparisonNode + | LookupMembershipNode | LookupPairComparisonNode | LookupDefinedNode ) @@ -298,11 +348,11 @@ def _atom_dims(atom: TypedPredicateNode, name_dims: Mapping[str, Sequence[str]]) it. """ match atom: - case ParameterComparisonNode() | ParameterDefinedNode() | VariableDefinedNode(): + case ParameterComparisonNode() | ParameterMembershipNode() | ParameterDefinedNode() | VariableDefinedNode(): return frozenset(name_dims.get(atom.name, ())) - case DimensionComparisonNode() | DimensionPositionNode(): + case DimensionComparisonNode() | DimensionMembershipNode() | DimensionPositionNode(): return frozenset({atom.name}) - case LookupComparisonNode() | LookupPairComparisonNode() | LookupDefinedNode(): + case LookupComparisonNode() | LookupMembershipNode() | LookupPairComparisonNode() | LookupDefinedNode(): return frozenset({atom.over}) case _: assert_never(atom) @@ -328,11 +378,29 @@ def _position_comparison(tokens: pp.ParseResults) -> UnresolvedPositionNode: ) +def _literal_token(token: Any) -> tuple[float | str, bool]: + """One literal token as ``(value, quoted)`` — a number, a quoted label, or a bare word. + + The one home for turning a raw grammar token into a typed literal plus its + quoted flag; a comparison's single right-hand side and a membership's list + both read one token at a time through it. + """ + if isinstance(token, _Quoted): + return str(token), True + return (token if isinstance(token, float) else str(token)), False + + def _comparison(tokens: pp.ParseResults) -> UnresolvedComparisonNode: """``name literal`` off the tokens the grammar captured, the quoted marker turned into a flag.""" name, op, value = tokens - quoted = isinstance(value, _Quoted) - return UnresolvedComparisonNode(str(name), cast('PredicateOperator', op), str(value) if quoted else value, quoted) + value, quoted = _literal_token(value) + return UnresolvedComparisonNode(str(name), cast('PredicateOperator', op), value, quoted) + + +def _membership(tokens: pp.ParseResults) -> UnresolvedMembershipNode: + """``name in [l1, l2, …]`` off the tokens the grammar captured; the list may be empty for resolution to refuse.""" + name, *elements = tokens + return UnresolvedMembershipNode(str(name), tuple(_literal_token(e) for e in elements)) def _build_where_grammar() -> pp.ParserElement: @@ -366,16 +434,26 @@ def _build_where_grammar() -> pp.ParserElement: position_comparison = (position_call + comparator + position).set_parse_action(_position_comparison) comparison = (name + comparator + (number | quoted | name)).set_parse_action(_comparison) + + IN = pp.Suppress(pp.CaselessKeyword('in')) + element = number | quoted | name.copy() + membership = ( + name + IN + pp.Suppress('[') + pp.Optional(pp.DelimitedList(element)) + pp.Suppress(']') + ).set_parse_action(_membership) + # pyrefly: ignore[implicit-any-lambda] existence = name.copy().set_parse_action(lambda t: UnresolvedNameNode(t[0])) # `position_comparison` leads: it starts with a keyword that `existence` # would otherwise take for a bare name, and `comparison` for a parameter. + # `membership` precedes `comparison` and `existence`, which would take its + # name for a whole atom and leave `in [...]` for the parse to choke on. # See `DimensionPositionNode` for why it converts on the left (#32). atom = ( true_lit | false_lit | position_comparison + | membership | comparison | existence | (pp.Suppress('(') + where_expr + pp.Suppress(')')) diff --git a/tests/test_dimensions.py b/tests/test_dimensions.py index e7063d43..4beb3270 100644 --- a/tests/test_dimensions.py +++ b/tests/test_dimensions.py @@ -204,6 +204,21 @@ def test_an_outer_product_is_legal_and_carries_both_dim_sets(): "where-comparison on dimension 'snapshot'", id='where-comparison-on-a-dim-outside-the-frame', ), + pytest.param( + {'variables.cap': {'foreach': ['generator'], 'where': 'load in [0, 1]'}}, + r"where-parameter 'load' has dims \['bus', 'snapshot'\]", + id='where-membership-on-a-parameter-outside-the-frame', + ), + pytest.param( + {'variables.cap': {'foreach': ['generator'], 'where': 'snapshot in [0, 1]'}}, + "where-comparison on dimension 'snapshot'", + id='where-membership-on-a-dim-outside-the-frame', + ), + pytest.param( + {'variables.cap': {'foreach': ['generator'], 'where': "snap_bus in ['x']"}}, + "where-comparison on lookup 'snap_bus', which is over dimension 'snapshot'", + id='where-membership-on-a-lookup-outside-the-frame', + ), pytest.param( {'variables.cap': {'foreach': ['generator'], 'bounds': {'lower': 0, 'upper': 'load'}}}, r"bounds.upper parameter 'load' has dims \['bus', 'snapshot'\]", diff --git a/tests/test_lowering.py b/tests/test_lowering.py index 492fd192..00f5fe0a 100644 --- a/tests/test_lowering.py +++ b/tests/test_lowering.py @@ -57,8 +57,11 @@ from math_spec.where_parser import ( BooleanLiteralNode, DimensionComparisonNode, + DimensionMembershipNode, + LookupMembershipNode, ParameterComparisonNode, ParameterDefinedNode, + ParameterMembershipNode, ) from tests.fixtures import DISPATCH_MODEL, SMALL_MODEL, override, schema_of @@ -167,6 +170,27 @@ def test_where_lowering(dispatch_schema, where, expected): assert where_of(where, Namespace.of(dispatch_schema), 't') == expected +@pytest.mark.parametrize( + ('where', 'expected'), + [ + pytest.param('c in [1.5, 2.5]', ParameterMembershipNode('c', (1.5, 2.5)), id='over-a-parameter'), + pytest.param("g in ['g1', 'g2']", DimensionMembershipNode('g', ('g1', 'g2')), id='over-a-dimension'), + pytest.param("lk in ['h1']", LookupMembershipNode('lk', 'g', ('h1',)), id='over-a-groupable-lookup'), + pytest.param("tag in ['north']", LookupMembershipNode('tag', 'g', ('north',)), id='over-a-label-space'), + ], +) +def test_a_membership_mask_reaches_the_program_unchanged(where, expected): + """A `where` atom *is* the program's predicate — no lowering case rewrites one. + + That is what lets a new atom (#254) cost nothing in `lowering.py`, and it is + the half no other suite covers: resolution proves the node is built and the + typesetter proves it prints, while only this asks whether it survives the + pass in between. + """ + program = lower_program(expand_piecewise(schema_of(SMALL_MODEL, **{'variables.p.where': where}))) + assert program.variables['p'].where == expected, 'the atom the file wrote, not a rewrite of it' + + def test_a_literal_amount_resolves_to_one_signed_number(dispatch_schema): """`offset=-1` parses as a unary minus over `1`; after resolution it is `-1`, for every reader alike.""" ns = Namespace.of(dispatch_schema) diff --git a/tests/test_parser.py b/tests/test_parser.py index b7425071..3ce1fe88 100644 --- a/tests/test_parser.py +++ b/tests/test_parser.py @@ -27,6 +27,7 @@ NotNode, OrNode, UnresolvedComparisonNode, + UnresolvedMembershipNode, UnresolvedNameNode, UnresolvedPositionNode, parse_where, @@ -242,3 +243,63 @@ def test_an_unrelated_parse_failure_says_nothing_about_positions(): with pytest.raises(SchemaError) as excinfo: parse_where('p_max >') assert 'position()' not in str(excinfo.value) + + +@pytest.mark.parametrize( + ('text', 'elements'), + [ + ("carrier in ['ccgt', 'ocgt']", (('ccgt', True), ('ocgt', True))), + ('carrier in ["ccgt"]', (('ccgt', True),)), + ("carrier in ['combined-cycle']", (('combined-cycle', True),)), + ("carrier in ['it\\'s']", (("it's", True),)), + ('period in [2030, 2040]', ((2030.0, False), (2040.0, False))), + ('offset in [-1, -2]', ((-1.0, False), (-2.0, False))), + ('period in [1.5, 2e3]', ((1.5, False), (2000.0, False))), + ('carrier in [wind]', (('wind', False),)), + ('carrier IN [wind]', (('wind', False),)), + ('carrier in []', ()), + ], + ids=[ + 'labels', + 'double-quotes', + 'hyphen', + 'escaped-quote', + 'integers', + 'negatives', + 'floats', + 'bare-word', + 'upper-case-in', + 'empty', + ], +) +def test_a_membership_parses_to_its_elements(text, elements): + """The parser keeps each element's `quoted` flag; whether a bare word names + a declaration is resolution's problem, and an empty list parses so + resolution can word the refusal.""" + node = parse_where(text) + assert isinstance(node, UnresolvedMembershipNode) + assert node.elements == elements + + +@pytest.mark.parametrize( + ('text', 'tree'), + [ + ("NOT carrier in ['a']", NotNode(UnresolvedMembershipNode('carrier', (('a', True),)))), + ( + "carrier in ['a'] AND p > 0", + AndNode( + UnresolvedMembershipNode('carrier', (('a', True),)), + UnresolvedComparisonNode('p', '>', 0.0, quoted=False), + ), + ), + ( + "(carrier in ['a']) OR b", + OrNode(UnresolvedMembershipNode('carrier', (('a', True),)), UnresolvedNameNode('b')), + ), + ], + ids=['under-not', 'under-and', 'parenthesised-under-or'], +) +def test_a_membership_takes_the_connectives(text, tree): + """A membership is an atom like a comparison, so `NOT`/`AND`/`OR` and + parentheses bind around it the same way.""" + assert parse_where(text) == tree diff --git a/tests/test_validation.py b/tests/test_validation.py index 6dbf26a2..45c34419 100644 --- a/tests/test_validation.py +++ b/tests/test_validation.py @@ -274,6 +274,95 @@ def test_a_named_amount_keeps_its_own_sentence(self): ) +class TestMembership: + """`name in [...]` — the set form of a where-comparison, and how it refuses (#254).""" + + @pytest.mark.parametrize( + ('patch', 'fragments'), + [ + pytest.param( + {'variables.p.where': 'g in []'}, + ('matches nothing', 'where: "False"'), + id='an-empty-list-is-an-always-false-mask', + ), + pytest.param( + {'variables.p.where': "g in ['a', 'a']"}, + ('more than once', 'Drop the duplicate'), + id='a-repeated-element-selects-nothing-extra', + ), + pytest.param( + {'variables.p.where': 'g in [c]'}, + ("names 'c'", 'data-driven membership (#258)', 'bool parameter'), + id='a-declared-name-among-the-elements', + ), + pytest.param( + {'variables.p.where': 'p in [1]'}, + ('built before variables exist',), + id='a-variable-on-the-left', + ), + pytest.param( + {'variables.p.where': 'nope in [1]'}, + ("'nope' not found",), + id='an-unknown-name-on-the-left', + ), + pytest.param( + {'variables.p.where': 'g in [3]'}, + ('matches no label',), + id='a-number-against-a-str-dimension', + ), + pytest.param( + {'variables.p.where': 'lk in [3]'}, + ('matches no label',), + id="a-number-against-a-lookup's-str-target", + ), + pytest.param( + {'variables.p.where': "flag in ['x']"}, + ('matches nothing',), + id='a-label-against-a-bool-parameter', + ), + pytest.param( + {'dimensions.g': {'dtype': 'datetime'}, 'variables.p.where': 'g in [0, 1]'}, + ('compares against the epoch',), + id='a-number-against-a-datetime-dimension-reads-as-the-epoch', + ), + pytest.param( + {'variables.p.where': 'c in [0, 0]'}, + ('lists 0 more', 'Drop the duplicate'), + id='a-numeric-duplicate-renders-via-g-not-as-0-point-0', + ), + pytest.param( + { + 'dimensions.g': {'dtype': 'datetime'}, + 'variables.p.where': "g in ['2030-01-01T00:00', '2030-01-01T00:00:00']", + }, + ('more than once', 'Drop the duplicate'), + id='two-spellings-of-one-instant-are-a-duplicate-once-typed', + ), + ], + ) + def test_a_membership_the_data_cannot_decide_is_refused(self, patch, fragments): + with pytest.raises(LanguageError) as exc: + _schema(**patch) + for fragment in fragments: + assert fragment in str(exc.value) + + @pytest.mark.parametrize( + 'patch', + [ + pytest.param({'variables.p.where': 'c in [1.5, 2.5]'}, id='floats-against-a-float-parameter'), + pytest.param({'variables.p.where': "g in ['wind', 'solar']"}, id='labels-against-a-str-dimension'), + pytest.param({'variables.p.where': 'flag in [1]'}, id='a-number-against-a-bool-parameter'), + pytest.param( + {'dimensions.g': {'dtype': 'datetime'}, 'variables.p.where': "g in ['2030-01-01', '2040-01-01']"}, + id='iso-dates-against-a-datetime-dimension', + ), + ], + ) + def test_a_well_typed_membership_loads(self, patch): + """Float membership is allowed: the OR chain it stands in for permits exact float equality already (#254).""" + _schema(**patch) + + class TestVersion: """`version:` is refused when unknown, and does nothing else (#67).""" diff --git a/tests/typesetting/golden/latex.out b/tests/typesetting/golden/latex.out index bd0090ca..21bfedcb 100644 --- a/tests/typesetting/golden/latex.out +++ b/tests/typesetting/golden/latex.out @@ -91,6 +91,7 @@ \text{first} && \mathit{on}_{t,g} & = 1 && \forall\, t \in \mathcal{T},\ g \in \mathcal{G} \,:\, \mathrm{pos}(t) = 0 \vee \mathrm{pos}_{\mathrm{season\_of}(t)}(t) = 0 \\ \text{last} && \mathit{on}_{t,g} & = 0 && \forall\, t \in \mathcal{T},\ g \in \mathcal{G} \,:\, \mathrm{pos}(t) = \lvert \mathcal{T} \rvert - 1 \vee \mathrm{pos}_{\mathrm{season\_of}(t)}(t) = \lvert \mathcal{T}_{\mathrm{season\_of}(t)} \rvert - 1 \\ \text{northern} && \mathit{slack}_{t} & \le \mathrm{load}_{t,b} && \forall\, t \in \mathcal{T},\ b \in \mathcal{B} \,:\, \mathrm{zone\_of}(b) = \text{'}\mathrm{north}\text{'} \wedge \mathrm{zone\_of}(b) \neq \mathrm{area\_of}(b) \wedge \mathrm{zone\_of}(b) \text{ is defined} \\ +\text{selected} && p_{t,g} & \le \mathrm{load}_{t,b} && \forall\, t \in \mathcal{T},\ g \in \mathcal{G},\ b \in \mathcal{B} \,:\, t \in \{0,\ 3\} \vee \mathrm{zone\_of}(b) \in \{\text{'}\mathrm{north}\text{'},\ \text{'}\mathrm{south}\text{'}\} \vee \mathrm{min\_up}_{g} \in \{2,\ 3\} \\ \text{efficiency} && p_{t,g} & \le \mathrm{eta}_{g} \cdot \mathrm{p}^{\mathrm{max}}_{g} && \forall\, t \in \mathcal{T},\ g \in \mathcal{G} \\ \text{ceiling} && \theta_{b} & \le \infty && \forall\, b \in \mathcal{B} \\ \text{always} && \mathit{spill}_{t} & \ge 0 && \forall\, t \in \mathcal{T} \\ diff --git a/tests/typesetting/golden/markdown.out b/tests/typesetting/golden/markdown.out index f7eb4410..8bc76b29 100644 --- a/tests/typesetting/golden/markdown.out +++ b/tests/typesetting/golden/markdown.out @@ -157,6 +157,10 @@ $$\mathit{on}_{t,g} = 0 \qquad \forall\thinspace t \in \mathcal{T},\enspace g \i $$\mathit{slack}_{t} \le \mathrm{load}_{t,b} \qquad \forall\thinspace t \in \mathcal{T},\enspace b \in \mathcal{B} \thinspace:\thinspace \mathrm{zone\_of}(b) = \text{'}\mathrm{north}\text{'} \wedge \mathrm{zone\_of}(b) \neq \mathrm{area\_of}(b) \wedge \mathrm{zone\_of}(b) \text{ is defined}$$ +**`selected`** + +$$p_{t,g} \le \mathrm{load}_{t,b} \qquad \forall\thinspace t \in \mathcal{T},\enspace g \in \mathcal{G},\enspace b \in \mathcal{B} \thinspace:\thinspace t \in \{0,\enspace 3\} \vee \mathrm{zone\_of}(b) \in \{\text{'}\mathrm{north}\text{'},\enspace \text{'}\mathrm{south}\text{'}\} \vee \mathrm{min\_up}_{g} \in \{2,\enspace 3\}$$ + **`efficiency`** $$p_{t,g} \le \mathrm{eta}_{g} \cdot \mathrm{p}^{\mathrm{max}}_{g} \qquad \forall\thinspace t \in \mathcal{T},\enspace g \in \mathcal{G}$$ diff --git a/tests/typesetting/golden/model.yaml b/tests/typesetting/golden/model.yaml index eec16c95..c334aa8f 100644 --- a/tests/typesetting/golden/model.yaml +++ b/tests/typesetting/golden/model.yaml @@ -161,6 +161,10 @@ constraints: foreach: [snapshot, bus] where: "zone_of == 'north' AND zone_of != area_of AND zone_of" expression: slack <= load + selected: # set membership per kind: a dimension's coordinates, a lookup's values, a parameter's numbers + foreach: [snapshot, generator, bus] + where: "snapshot in [0, 3] OR zone_of in ['north', 'south'] OR min_up in [2, 3]" + expression: p <= load efficiency: # a Greek-named parameter, which is given — so the convention wins and it prints as the word foreach: [snapshot, generator] expression: p <= eta * p_max diff --git a/tests/typesetting/golden/typst.out b/tests/typesetting/golden/typst.out index a4cc2645..e30067f6 100644 --- a/tests/typesetting/golden/typst.out +++ b/tests/typesetting/golden/typst.out @@ -80,6 +80,7 @@ $ upright("balance") & sum_(g in cal(G) colon upright("gen_bus")(g) = b) p_(t,g) upright("first") & italic("on")_(t,g) & = 1 & forall t in cal(T), g in cal(G) colon upright("pos")(t) = 0 or upright("pos")_(upright("season_of")(t))(t) = 0 \ upright("last") & italic("on")_(t,g) & = 0 & forall t in cal(T), g in cal(G) colon upright("pos")(t) = abs(cal(T)) - 1 or upright("pos")_(upright("season_of")(t))(t) = abs(cal(T)_(upright("season_of")(t))) - 1 \ upright("northern") & italic("slack")_(t) & <= upright("load")_(t,b) & forall t in cal(T), b in cal(B) colon upright("zone_of")(b) = upright("'north'") and upright("zone_of")(b) != upright("area_of")(b) and upright("zone_of")(b) upright(" is defined") \ + upright("selected") & p_(t,g) & <= upright("load")_(t,b) & forall t in cal(T), g in cal(G), b in cal(B) colon t in {0, 3} or upright("zone_of")(b) in {upright("'north'"), upright("'south'")} or upright("min_up")_(g) in {2, 3} \ upright("efficiency") & p_(t,g) & <= upright("eta")_(g) dot upright("p")^(upright("max"))_(g) & forall t in cal(T), g in cal(G) \ upright("ceiling") & theta_(b) & <= infinity & forall b in cal(B) \ upright("always") & italic("spill")_(t) & >= 0 & forall t in cal(T) \ diff --git a/tests/typesetting/test_golden.py b/tests/typesetting/test_golden.py index 6abf4404..c11dd541 100644 --- a/tests/typesetting/test_golden.py +++ b/tests/typesetting/test_golden.py @@ -140,6 +140,7 @@ def _rendered_trees() -> Iterator[object]: UNRESOLVED = { 'UnresolvedNameNode', 'UnresolvedComparisonNode', + 'UnresolvedMembershipNode', 'UnresolvedPositionNode', 'NameNode', 'NameListNode', diff --git a/tests/typesetting/test_walk.py b/tests/typesetting/test_walk.py index 91ae13c4..964f90b1 100644 --- a/tests/typesetting/test_walk.py +++ b/tests/typesetting/test_walk.py @@ -367,6 +367,41 @@ def test_a_dimension_compared_against_a_number_says_what_its_coordinates_are(fmt assert 'coordinates)' not in typeset(_selected('position(snapshot) == 0'), fmt) +def _member(where: str) -> dict[str, Any]: + """One variable masked by *where*, with a dimension, a lookup and a parameter to test membership against.""" + return { + 'dimensions': {'g': {'dtype': 'str'}}, + 'lookups': {'kind_of': {'over': 'g', 'dtype': 'str'}}, + 'parameters': {'size': {'dims': ['g'], 'dtype': 'int'}}, + 'variables': {'p': {'foreach': ['g'], 'where': where, 'bounds': {'lower': 0}}}, + 'objective': {'sense': 'minimize', 'expression': 'sum(p, over=g)'}, + } + + +@EVERY_FORMAT +@pytest.mark.parametrize( + ('where', 'rendered'), + [ + pytest.param("g in ['wind', 'solar']", ['wind', 'solar'], id='a-dimension'), + pytest.param("kind_of in ['thermal']", ['thermal'], id='a-lookup'), + pytest.param('size in [2, 3]', [2, 3], id='a-parameter'), + ], +) +def test_a_membership_renders_as_a_set(fmt: Format, where: str, rendered: list[Any]): + """`name in [...]` prints `∈ {v1, v2, …}` in every format per kind, its elements as literals.""" + text = typeset(_member(where), fmt, legend=False) + elements = [str(v) if isinstance(v, int) else fmt.quoted(v) for v in rendered] + inner = fmt.set_braces(fmt.joined(elements, '')) + assert f'{fmt.operators["in"]} {inner}' in text + + +@EVERY_FORMAT +def test_a_membership_over_a_numeric_dimension_says_what_its_coordinates_are(fmt: Format): + """A numeric coordinate in a set can be taken for a position, so the legend places it — the `==` case one construct over.""" + text = typeset(_selected('snapshot in [0, 3]'), fmt) + assert f'({fmt.mono("int")} coordinates)' in text + + @EVERY_FORMAT def test_a_description_is_joined_to_its_name_by_a_dash_the_format_renders(fmt: Format): """``---`` is TeX's em-dash ligature and Typst's, and nothing in Markdown.