diff --git a/docs/contributing.md b/docs/contributing.md index 2a1ea510..271bbed5 100644 --- a/docs/contributing.md +++ b/docs/contributing.md @@ -98,14 +98,14 @@ stale anchor fails it. `pixi run docs-serve` builds the site and serves it at The same construct passes through three layers, and each names it in full. The suffix says which layer: -| Layer | Suffix | Example | -| ------------------------------- | -------------------- | ------------------------------------------ | -| YAML block (`math_spec.model`) | `Block` | `VariableBlock`, `PiecewiseBlock` | -| Core AST (`math_spec.*_parser`) | `Node` | `VariableNode`, `UnresolvedComparisonNode` | -| Program (`math_spec.program`) | none / `Declaration` | `Variable`, `VariableDeclaration` | +| Layer | Suffix | Example | +| ------------------------------ | -------------------- | -------------------------------------- | +| YAML block (`math_spec.model`) | `Block` | `VariableBlock`, `PiecewiseBlock` | +| Syntax (`math_spec.*_parser`) | `Node` | `NameNode`, `UnresolvedComparisonNode` | +| Program (`math_spec.program`) | none / `Declaration` | `Variable`, `VariableDeclaration` | -A node names the operation, not the verb a file writes. One verb can lower to -two nodes, so the file's spelling cannot decide the name. +A node names the operation, not the verb a file writes. One verb can resolve +to two nodes, so the file's spelling cannot decide the name. | File verb | Node | What the node names | | ------------------ | ----------- | ------------------------------ | @@ -121,11 +121,10 @@ Nothing is abbreviated. Start with the grammar, which is usually free because `f(x, k=v)` already parses. Then declare the signature in `operators.BUILTINS`. It holds the number -of arguments and says which arguments name dimensions, and resolution, -validation and lowering all read it from there. Then write the dimension rule in -`dimensions.py`, the degree verdict in `degree.py`, the node it lowers to in -`program.py`, and the entry in the -[language reference](reference/language/operators.md). +of arguments and says which arguments name dimensions, and resolution reads it +from there. Then write the node in `program.py` and how resolution builds it, +the dimension rule in `dimensions.py`, the degree verdict in `degree.py`, and +the entry in the [language reference](reference/language/operators.md). ## Submitting changes diff --git a/docs/reference/language/errors.md b/docs/reference/language/errors.md index 0ea141d0..210d905b 100644 --- a/docs/reference/language/errors.md +++ b/docs/reference/language/errors.md @@ -50,13 +50,12 @@ variable with no constraint row. ## Which error you get -| | | -| ------------------------- | -------------------------------------------------------------------------------------------------------------------------------- | -| `MathSpecError` | The root. Everything below is an instance of it | -| `LanguageError` | Something in the model: a construct outside the language, a dimension set that does not compose, or a name that nothing declares | -| `SchemaError` | Something in the file: an unknown key, a malformed declaration, or a bad symbol table | -| `DimensionError` | Dimensions that disagree, such as a constraint whose expression does not equal its `dims` | -| `PiecewiseExpansionError` | A `piecewise:` block that cannot be expanded | +| | | +| ---------------- | -------------------------------------------------------------------------------------------------------------------------------- | +| `MathSpecError` | The root. Everything below is an instance of it | +| `LanguageError` | Something in the model: a construct outside the language, a dimension set that does not compose, or a name that nothing declares | +| `SchemaError` | Something in the file: an unknown key, a malformed declaration, or a bad symbol table | +| `DimensionError` | Dimensions that disagree, such as a constraint whose expression does not equal its `dims` | Every one of these is reproducible from the YAML alone. An engine that binds numbers or calls a solver adds its own errors below `MathSpecError`. diff --git a/docs/reference/reading.md b/docs/reference/reading.md index d5237914..14af2af1 100644 --- a/docs/reference/reading.md +++ b/docs/reference/reading.md @@ -111,17 +111,17 @@ refusal quotes. The engine, which has the numbers, runs each one and raises `assumption_message` where it fails: ```python -from math_spec.program import Holds, assumption_message +from math_spec.program import Assumption, assumption_message sorted(program.assumptions) # ['cost_is_never_negative', 'curve_complete', 'curve_curvature', 'curve_increasing'] -isinstance(program.assumptions['curve_increasing'], Holds) # True +isinstance(program.assumptions['curve_increasing'], Assumption) # True message = assumption_message('curve_increasing', program.assumptions['curve_increasing']) message # "assumption 'curve_increasing' does not hold for the data bound to 'bp_x' — piecewise 'curve': method: convex requires strictly increasing breakpoints in 'bp_x' along 'bp'" written = assumption_message('cost_is_never_negative', program.assumptions['cost_is_never_negative']) written # "assumption 'cost_is_never_negative' does not hold for the data bound to 'bp_y' — a negative cost is a gain the objective would chase" ``` -One kind stands in that mapping. A `Holds` carries a predicate as two masks — +One kind stands in that mapping. An `Assumption` carries a predicate as two masks — `predicate`, and the `where` it is checked under — and the sentence a refusal trails under `description`. What a `piecewise:` block's method implies about its breakpoints is written in the same language and stands beside what the @@ -138,6 +138,10 @@ node's operands, and `where_children()` walks a predicate's. `walk()` yields every node under an expression, parents first. `walk_regions()` yields each node with the `cases:` regions it stands inside, outermost first. +`Named` is the one node no program carries. A `Spec.resolved` tree holds it +where an `expressions:` entry is used, and lowering inlines the entry's body +there before the program is built, so `Expression` does not name it. + Every `where` arrives as a `Mask`. Its `.root` is the resolved predicate. The mask also answers four questions: diff --git a/schema/math-spec.schema.json b/schema/math-spec.schema.json index d9c8e533..7dec94e3 100644 --- a/schema/math-spec.schema.json +++ b/schema/math-spec.schema.json @@ -264,7 +264,7 @@ }, "MacroBlock": { "additionalProperties": false, - "description": "A parameterised expression template, defined in the YAML itself.\n\nLanguage, not code: formals (``args`` positional, ``kwargs`` keyword)\nshadow model names inside the template, and every call site expands into\ncore AST before either backend sees the expression.", + "description": "A parameterised expression template, defined in the YAML itself.\n\nLanguage, not code: formals (``args`` positional, ``kwargs`` keyword)\nshadow model names inside the template, and every call site expands in\nthe syntax tree before resolution reads the expression.", "properties": { "args": { "default": [], diff --git a/src/math_spec/__init__.py b/src/math_spec/__init__.py index af4b52d7..37513955 100644 --- a/src/math_spec/__init__.py +++ b/src/math_spec/__init__.py @@ -18,7 +18,6 @@ DimensionError, LanguageError, MathSpecError, - PiecewiseExpansionError, SchemaError, did_you_mean, schema_error, @@ -65,7 +64,6 @@ 'DimensionError', 'LanguageError', 'MathSpecError', - 'PiecewiseExpansionError', 'SchemaError', 'SosBlock', 'Spec', diff --git a/src/math_spec/_expression_parser.py b/src/math_spec/_expression_parser.py index 511468a8..891d2db9 100644 --- a/src/math_spec/_expression_parser.py +++ b/src/math_spec/_expression_parser.py @@ -2,10 +2,11 @@ # # SPDX-License-Identifier: MIT -"""The core AST every pass reads, and the pyparsing grammar that builds it — package-private. +"""The syntax tree the expression grammar builds, and the grammar — package-private. -Arithmetic nests anywhere; a comparison appears only at the top of a parsed -expression. +Only expansion and resolution read it: resolution rewrites it into the +:mod:`math_spec.program` vocabulary, which every pass after reads. Arithmetic +nests anywhere; a comparison appears only at the top of a parsed expression. """ from __future__ import annotations @@ -18,13 +19,10 @@ from math_spec._sealed import Sealed from math_spec.errors import SchemaError -from math_spec.operators import EDGE_WRAP if TYPE_CHECKING: from collections.abc import Callable, Iterable, Iterator, Mapping - from math_spec.program import Direction, Partition, Predicate - #: The relation a comparison may carry — the three an expression may be #: written with, which is what a constraint's sense is read off. ComparisonOperator = Literal['<=', '>=', '=='] @@ -60,57 +58,7 @@ def __str__(self) -> str: @dataclass(frozen=True) class NameNode: - """A bare name whose kind only the schema knows; resolution rewrites every one into a typed node.""" - - name: str - - def __str__(self) -> str: - return self.name - - -@dataclass(frozen=True) -class VariableNode: - """A resolved reference to a declared decision variable.""" - - name: str - - def __str__(self) -> str: - return self.name - - -@dataclass(frozen=True) -class ParameterNode: - """A resolved reference to a declared parameter.""" - - name: str - - def __str__(self) -> str: - return self.name - - -@dataclass(frozen=True) -class DualNode: - """A resolved ``dual(c)``: the row dual of the declared constraint *c*, a leaf. - - Constraints sit outside the flat namespace, so a bare name never resolves - to one; ``dual(c)`` is the one position that reads the constraint store. A - dual is a number only a solve produces, so the loader refuses this leaf - anywhere the math is built (:mod:`math_spec.validation`). - """ - - constraint: str - - def __str__(self) -> str: - return f'dual({self.constraint})' - - -@dataclass(frozen=True) -class DimensionNode: - """A resolved reference to a declared dimension. - - Only legal in operator kwarg *values* (``sum(x, over=generator)``), never as - a value in arithmetic — a dimension is a coordinate space, not data. - """ + """A bare name whose kind only the schema knows; resolution rewrites every one into a program node.""" name: str @@ -131,26 +79,6 @@ def __str__(self) -> str: return shown(self.names) -@dataclass(frozen=True) -class DirectionNode: - """A resolved ``by=`` on ``sum`` or ``at``: the relation, read in the :class:`Direction` the call names.""" - - direction: Direction - - def __str__(self) -> str: - return self.direction.name - - -@dataclass(frozen=True) -class PartitionNode: - """A resolved ``by=`` on ``shift`` or ``sum_back``: the relation, as the :class:`Partition` the call steps inside.""" - - partition: Partition - - def __str__(self) -> str: - return self.partition.name - - @dataclass(frozen=True) class KeywordNode: """A quoted closed keyword in a kwarg value — ``shift(..., edge='wrap')``. @@ -164,14 +92,6 @@ def __str__(self) -> str: return f"'{self.value}'" -@dataclass(frozen=True) -class EdgeNode: - """The resolved ``edge='wrap'``; a number in the same position stays a :class:`NumberNode`.""" - - def __str__(self) -> str: - return f"'{EDGE_WRAP}'" - - @dataclass(frozen=True) class UnaryOperatorNode: op: UnaryOperator @@ -211,87 +131,10 @@ def __str__(self) -> str: return f'{self.name}({", ".join(passed)})' -@dataclass(frozen=True) -class CaseArm: - """One region of a :class:`CasesNode`: where it applies, and the value there. - - ``when`` is ``None`` on the **last** arm and only there — the block's - ``otherwise:``, which is what makes the quantity total without anything - having to prove it. Every other arm's ``when`` is proved apart from every - other arm's. - """ - - label: str - when: Predicate | None - value: ArithmeticNode - - -def case_context(name: str, label: str | None) -> str: - """The context an error inside one arm of a cased expression is reported under. - - Args: - name: The named expression the arm belongs to. - label: The case's name, or ``None`` for the block's ``otherwise:``. - - Returns: - The context prefix an error message carries. - """ - where = 'otherwise' if label is None else f"case '{label}'" - return f"Named expression '{name}', {where}" - - -@dataclass(frozen=True) -class CasesNode: - """A value defined by region — a named expression's ``cases:``, inlined where its name stood. - - Exactly one arm applies at every coordinate, which :mod:`math_spec.exclusivity` - proves at load; the last arm is the block's ``otherwise:`` and carries no - ``when``. The arms are in file order. The frame is not carried here: it is - on the declaration. - """ - - name: str - arms: tuple[CaseArm, ...] - - def __str__(self) -> str: - """The name the file wrote, which is all an expression ever said: ``cases:`` is YAML and not syntax.""" - return self.name - - -@dataclass(frozen=True) -class DefinitionNode: - """A plain named expression's body, inlined where its name stood — carrying the name. - - The math is the body's: every pass reads through this node as if the body - stood here bare. The name is for the typesetter, which may print the - quantity under it and define it once, as a paper does. - """ - - name: str - body: ArithmeticNode - - def __str__(self) -> str: - """The name the file wrote, rather than the body inlined under it.""" - return self.name - - +#: Every arithmetic node the grammar builds. A name, a name list and a quoted +#: keyword are what resolution reads for their kind; the rest is structure. ArithmeticNode = ( - NumberNode - | NameNode - | NameListNode - | VariableNode - | ParameterNode - | DualNode - | DimensionNode - | DirectionNode - | PartitionNode - | EdgeNode - | KeywordNode - | UnaryOperatorNode - | BinaryOperatorNode - | FunctionCallNode - | CasesNode - | DefinitionNode + NumberNode | NameNode | NameListNode | KeywordNode | UnaryOperatorNode | BinaryOperatorNode | FunctionCallNode ) @@ -306,9 +149,7 @@ def __str__(self) -> str: return f'{self.left} {self.op} {self.right}' -#: A whole spec-side expression tree — parse output and the resolved tree alike. -#: Named apart from :data:`math_spec.program.Expression`, the lowered -#: vocabulary a consumer reads. +#: A whole parsed expression: arithmetic, or one comparison over it. ParsedNode = ArithmeticNode | ComparisonNode @@ -323,8 +164,7 @@ def operand(node: ArithmeticNode) -> str: Whoever writes a node into a larger text — an operator, a line of a dumped sum — asks this rather than restating when brackets are needed. - A leaf, a call and a named expression are self-delimiting, and an operator - node is not. The brackets go on every operator operand rather than only the + A leaf and a call are self-delimiting, and an operator node is not. The brackets go on every operator operand rather than only the ones precedence would regroup, because a node prints without knowing its parent: ``a + (b * c)`` keeps the tree where ``a + b * c`` would rely on the reader knowing which binds tighter. @@ -332,31 +172,11 @@ def operand(node: ArithmeticNode) -> str: return f'({node})' if isinstance(node, (UnaryOperatorNode, BinaryOperatorNode)) else str(node) -# Node groups - -#: A resolved reference the language admits only as an operator kwarg *value*: -#: ``sum(x, along=d)``, ``sum(x, by=l)``, ``shift(..., edge='wrap')``. None of -#: the three is data, so none may stand in arithmetic — which is why the passes -#: that walk a value position refuse them together. -KwargNode = DimensionNode | DirectionNode | PartitionNode | EdgeNode - -#: What resolution rewrites away: a bare name, whose kind only the schema -#: knows, and the two kwarg-only literals its kwarg consumes. Meeting one -#: downstream means the expression skipped :func:`~math_spec.resolution.resolve_expression`. -UnresolvedNode = NameNode | NameListNode | KeywordNode - -#: Every leaf — nothing below it to descend into. -LeafNode = NumberNode | VariableNode | ParameterNode | DualNode | KwargNode | UnresolvedNode - - def children(node: ParsedNode) -> tuple[ArithmeticNode, ...]: """The sub-expressions of *node* — the structural half of any walk. - Every pass that recurses the whole tree and acts only at certain leaves - goes through here, so a node added later reaches all of them. An - operator's kwargs are children too — a dimension or coordinate is an + An operator's kwargs are children too — a dimension or coordinate is an ordinary node in a kwarg value, which is what lets a macro bind a formal. - A case arm's ``when`` is not: it is a mask over the frame, not a value in it. """ if isinstance(node, UnaryOperatorNode): return (node.operand,) @@ -364,10 +184,6 @@ def children(node: ParsedNode) -> tuple[ArithmeticNode, ...]: return (node.left, node.right) if isinstance(node, FunctionCallNode): return (*node.args, *node.kwargs.values()) - if isinstance(node, CasesNode): - return tuple(arm.value for arm in node.arms) - if isinstance(node, DefinitionNode): - return (node.body,) return () @@ -383,12 +199,8 @@ def nodes(*roots: ParsedNode) -> Iterator[ParsedNode]: def with_children(node: ArithmeticNode, recurse: Callable[[ArithmeticNode], ArithmeticNode]) -> ArithmeticNode: - """*node* rebuilt with *recurse* applied to each of its :func:`children`; a leaf comes back as is. - - A case arm's ``when`` is a mask over the frame, not a value in it, and is - carried across unchanged. - """ - if isinstance(node, LeafNode): + """*node* rebuilt with *recurse* applied to each of its :func:`children`; a leaf comes back as is.""" + if isinstance(node, NumberNode | NameNode | NameListNode | KeywordNode): return node if isinstance(node, UnaryOperatorNode): return UnaryOperatorNode(node.op, recurse(node.operand)) @@ -400,10 +212,6 @@ def with_children(node: ArithmeticNode, recurse: Callable[[ArithmeticNode], Arit tuple(recurse(a) for a in node.args), {k: recurse(v) for k, v in node.kwargs.items()}, ) - if isinstance(node, CasesNode): - return CasesNode(node.name, tuple(CaseArm(a.label, a.when, recurse(a.value)) for a in node.arms)) - if isinstance(node, DefinitionNode): - return DefinitionNode(node.name, recurse(node.body)) assert_never(node) diff --git a/src/math_spec/degree.py b/src/math_spec/degree.py index 1ebfc3ad..facc5faa 100644 --- a/src/math_spec/degree.py +++ b/src/math_spec/degree.py @@ -23,36 +23,25 @@ from __future__ import annotations -from math_spec._expression_parser import ( - BinaryOperatorNode, - DualNode, - FunctionCallNode, - ParsedNode, - UnresolvedNode, - VariableNode, +from math_spec.errors import LanguageError +from math_spec.program import ( + Add, + Divide, + Dual, + Expression, + GroupSum, + Multiply, + Power, + Sum, + Variable, + WindowSum, + carries_variable, children, - nodes, + walk, ) -from math_spec.errors import LanguageError - - -def carries_variable(node: ParsedNode) -> bool: - """Whether *node* contains a decision variable, over the core AST. - - :func:`math_spec.program.carries_variable` answers the same question over a - program. An unresolved node reaching here is a resolution bug, so it is refused - rather than silently answered. - """ - for found in nodes(node): - if isinstance(found, UnresolvedNode): - msg = f'{found!r} reached the degree check. Expressions go through resolution.resolve_expression() first.' - raise AssertionError(msg) - if isinstance(found, VariableNode): - return True - return False -def _adds(node: ParsedNode) -> bool: +def _adds(node: Expression) -> bool: """Whether *node* adds anywhere inside it. Anywhere, not only at its head: every operator over a variable-free @@ -60,14 +49,14 @@ def _adds(node: ParsedNode) -> bool: under a ``sum`` or a product reaches the quotient as two factors just as one at the top does. """ - return any(isinstance(found, BinaryOperatorNode) and found.op in ('+', '-') for found in nodes(node)) + return any(isinstance(found, Add) for found in walk(node)) -def check_binary(node: BinaryOperatorNode, context: str, *, ceiling: int) -> None: +def check_binary(node: Multiply | Divide | Power, context: str, *, ceiling: int) -> None: """Check that *node* stays inside the degree its position allows. Args: - node: The product, quotient or sum to judge. + node: The product, quotient or power to judge. context: What to name in the message — the declaration being read. ceiling: The highest degree this position can honour — 2 in an objective or a constraint, 1 everywhere else. @@ -79,26 +68,29 @@ def check_binary(node: BinaryOperatorNode, context: str, *, ceiling: int) -> Non a variable or adding. """ where = f'{context}: ' if context else '' - if node.op == '**': + if isinstance(node, Power): if carries_variable(node): raise LanguageError(_a_variable_under_a_power_message(where)) - if _adds(node.left) or _adds(node.right): + if _adds(node.base) or _adds(node.exponent): raise LanguageError( f'{where}a base and an exponent must each be a single Constant/Parameter factor, ' f'not a sum — addition does not distribute over `**`, so `(1 + rate) ** period` is ' f'refused where `growth ** period` is not. Bind the factor itself.' ) - if node.op == '/' and carries_variable(node.right): - raise LanguageError( - f'{where}the divisor contains variables, which is not affine. ' - f'Divide by a parameter, or precompute the reciprocal as one.' - ) - if node.op == '/' and _adds(node.right): - raise LanguageError( - f'{where}a divisor must be a single Constant/Parameter factor, ' - f'not a sum — rewrite as multiplication by a precomputed parameter' - ) - if node.op != '*' or not (carries_variable(node.left) and carries_variable(node.right)): + return + if isinstance(node, Divide): + if carries_variable(node.divisor): + raise LanguageError( + f'{where}the divisor contains variables, which is not affine. ' + f'Divide by a parameter, or precompute the reciprocal as one.' + ) + if _adds(node.divisor): + raise LanguageError( + f'{where}a divisor must be a single Constant/Parameter factor, ' + f'not a sum — rewrite as multiplication by a precomputed parameter' + ) + return + if not (carries_variable(node.left) and carries_variable(node.right)): return if ceiling < 2: raise LanguageError(_degree_two_here_message(where)) @@ -107,7 +99,7 @@ def check_binary(node: BinaryOperatorNode, context: str, *, ceiling: int) -> Non _check_single_term_factor(node, where) -def _degree(node: ParsedNode) -> int: +def _degree(node: Expression) -> int: """The polynomial degree *node* stands for, counted structurally. A product adds its factors' degrees and a division keeps the dividend's @@ -117,12 +109,12 @@ def _degree(node: ParsedNode) -> int: what stops a cubic from reaching a consumer to be refused by whichever one happens to notice. """ - if isinstance(node, VariableNode): + if isinstance(node, Variable): return 1 - if isinstance(node, BinaryOperatorNode) and node.op == '*': + if isinstance(node, Multiply): return _degree(node.left) + _degree(node.right) - if isinstance(node, BinaryOperatorNode) and node.op == '/': - return _degree(node.left) + if isinstance(node, Divide): + return _degree(node.numerator) return max((_degree(child) for child in children(node)), default=0) @@ -157,7 +149,7 @@ def _degree_two_here_message(where: str) -> str: ) -def _check_single_term_factor(node: BinaryOperatorNode, where: str) -> None: +def _check_single_term_factor(node: Multiply, where: str) -> None: """Refuse a degree-2 product of two multi-term factors.""" if not (_multi_term(node.left) and _multi_term(node.right)): return @@ -170,38 +162,31 @@ def _check_single_term_factor(node: BinaryOperatorNode, where: str) -> None: ) -def _multi_term(node: ParsedNode) -> bool: +def _multi_term(node: Expression) -> bool: """Whether *node* stands for more than one variable term at a coordinate. A reduction does, and so does an addition of two variable-carrying operands; a product is multi-term exactly when one of its factors is, a coefficient not multiplying the count. Structural, so it needs no data. """ - return any(_joins_terms(found) for found in nodes(node)) + return any(_joins_terms(found) for found in walk(node)) -def _joins_terms(node: ParsedNode) -> bool: - """Whether *node* itself makes several terms of one: a reduction over a variable, or a sum of two variable-carrying sides.""" - if isinstance(node, FunctionCallNode): - return node.name in _REDUCTIONS and any(carries_variable(a) for a in node.args) - return ( - isinstance(node, BinaryOperatorNode) - and node.op in ('+', '-') - and carries_variable(node.left) - and carries_variable(node.right) - ) +def _joins_terms(node: Expression) -> bool: + """Whether *node* itself makes several terms of one: a reduction over a variable, or a sum of two variable-carrying sides. - -#: The operators that fold several coordinates onto one, and so turn a term -#: into a sum of terms. ``at`` and ``shift`` re-index and are not here: they -#: move a term, leaving one term where there was one. -_REDUCTIONS = frozenset({'sum', 'sum_back'}) + A pullback and a translation re-index and are not reductions: they move a + term, leaving one term where there was one. + """ + if isinstance(node, Sum | GroupSum | WindowSum): + return carries_variable(node.operand) + return isinstance(node, Add) and carries_variable(node.left) and carries_variable(node.right) -def check_expression(node: ParsedNode, context: str, *, ceiling: int = 1) -> None: - """What the math admits at one position: no ``dual()`` anywhere under *node*, then :func:`check_binary` everywhere in it. +def check_expression(node: Expression, context: str, *, ceiling: int = 1) -> None: + """What the math admits at one position: no dual anywhere under *node*, then :func:`check_binary` everywhere in it. - Asked of the *expanded* tree, so a dual or a product inlined through a + Asked of the resolved tree, so a dual or a product inlined through a macro or a named expression is caught alongside one written in place. What a plan node can represent is the consumer's question, not this one's. @@ -209,16 +194,16 @@ def check_expression(node: ParsedNode, context: str, *, ceiling: int = 1) -> Non LanguageError: A dual, which exists only after a solve; or what :func:`check_binary` refuses. """ - for found in nodes(node): - if isinstance(found, DualNode): + for found in walk(node): + if isinstance(found, Dual): raise LanguageError( f'{context}: a dual exists only after a solve; the math cannot read one — ' f'keep the entry that carries it out of constraints, the objective, bounds and where.' ) - if isinstance(found, BinaryOperatorNode): + if isinstance(found, Multiply | Divide | Power): check_binary(found, context, ceiling=ceiling) -def calls_dual(node: ParsedNode) -> bool: - """Whether a :class:`DualNode` stands anywhere in the resolved *node*.""" - return any(isinstance(found, DualNode) for found in nodes(node)) +def calls_dual(node: Expression) -> bool: + """Whether a :class:`~math_spec.program.Dual` stands anywhere under *node*.""" + return any(isinstance(found, Dual) for found in walk(node)) diff --git a/src/math_spec/dimensions.py b/src/math_spec/dimensions.py index f7faaf6b..7fd84e42 100644 --- a/src/math_spec/dimensions.py +++ b/src/math_spec/dimensions.py @@ -5,125 +5,113 @@ """Static dim-set checking — a type system whose type is a set of dim names. Every node's dim set is computable before any data is bound, so this pass runs -at load on the resolved AST. The per-node rules are the "Dim algebra" table in +at load on the resolved tree. The per-node rules are the "Dim algebra" table in ``docs/reference/language/expressions.md``; a constraint's two sides together must equal its ``dims``, and a where or a bound may not exceed the frame. """ from __future__ import annotations -import math -from typing import TYPE_CHECKING, NamedTuple, assert_never - -import math_spec.degree as degree -from math_spec._expression_parser import ( - ArithmeticNode, - BinaryOperatorNode, - CasesNode, - ComparisonNode, - DefinitionNode, - DimensionNode, - DirectionNode, - DualNode, - EdgeNode, - FunctionCallNode, - KwargNode, - NumberNode, - ParameterNode, - ParsedNode, - PartitionNode, - UnaryOperatorNode, - UnresolvedNode, - VariableNode, - case_context, - children, -) -from math_spec.errors import DimensionError -from math_spec.operators import BUILTINS +from typing import TYPE_CHECKING, assert_never + +from math_spec.errors import DimensionError, case_context +from math_spec.operators import AMOUNTS from math_spec.program import ( + Add, + Cases, + Constant, CountComparison, DimensionComparison, DimensionPosition, Direction, + Divide, + Dual, + Expression, ExpressionComparison, + GroupSum, Mask, + Multiply, + Named, + Negate, + Parameter, ParameterComparison, ParameterDefined, Partition, + Power, + Pullback, PulledBackPredicate, RelationComparison, RelationDefined, RelationPairComparison, + Sum, + Translate, TranslatedPredicate, + Variable, VariableDefined, + WindowSum, ) if TYPE_CHECKING: - from collections.abc import Callable - from math_spec.model import Spec from math_spec.resolution import Resolved -def dims_of( - node: ParsedNode, - schema: Spec, - context: str, -) -> frozenset[str]: +def dims_of(node: Expression, schema: Spec, context: str) -> frozenset[str]: """The dim set of a resolved expression, checking every rule on the way. Raises: DimensionError: On the first rule broken. """ - if isinstance(node, ComparisonNode): - return _dims(node.left, schema, context) | _dims(node.right, schema, context) - return _dims(node, schema, context) - - -def _dims( - node: ArithmeticNode, - schema: Spec, - context: str, -) -> frozenset[str]: - """The recursive worker under :func:`dims_of`. - - An operator has a rule of its own and a cased entry declares its frame; - every other branch carries the union of what is under it. - """ - if isinstance(node, NumberNode): + if isinstance(node, Constant): return frozenset() - if isinstance(node, ParameterNode): + if isinstance(node, Parameter): return frozenset(schema.parameters[node.name].dims) - if isinstance(node, VariableNode): + if isinstance(node, Variable): return frozenset(schema.variables[node.name].dims) - if isinstance(node, UnresolvedNode | KwargNode): - msg = f'{type(node).__name__} reached the dim checker; resolve the expression first.' - raise AssertionError(msg) - - if isinstance(node, DualNode): + if isinstance(node, Dual): return frozenset(schema.constraints[node.constraint].dims) - if isinstance(node, FunctionCallNode): - return _dims_call(node, schema, context) + if isinstance(node, Named): + return _named_dims(node, schema, context) + + if isinstance(node, Cases): + return frozenset().union(*(dims_of(region.value, schema, context) for region in node.regions)) - if isinstance(node, CasesNode): - return _cases_dims(node, schema) + if isinstance(node, Negate | Add | Multiply | Power | Divide): + return frozenset().union(*(dims_of(child, schema, context) for child in _operands(node))) - if isinstance(node, UnaryOperatorNode | BinaryOperatorNode | DefinitionNode): - return frozenset().union(*(_dims(child, schema, context) for child in children(node))) + inner = dims_of(node.operand, schema, context) + if isinstance(node, Sum): + return _sum_dims(node, inner, context) + if isinstance(node, GroupSum): + return _group_sum_dims(node, inner, context) + if isinstance(node, Pullback): + return _at_dims(node, inner, context) + if isinstance(node, Translate | WindowSum): + return _translation_dims(node, inner, schema, context) assert_never(node) -def _cases_dims(node: CasesNode, schema: Spec) -> frozenset[str]: - """The declared frame rather than the union of the arms. +def _operands(node: Negate | Add | Multiply | Power | Divide) -> tuple[Expression, ...]: + if isinstance(node, Negate): + return (node.operand,) + if isinstance(node, Add | Multiply): + return (node.left, node.right) + if isinstance(node, Power): + return (node.base, node.exponent) + return (node.numerator, node.divisor) - A narrower arm broadcasts, as a parameter with fewer dims does. - """ - return frozenset(schema.expressions[node.name].dims or ()) + +def _named_dims(node: Named, schema: Spec, context: str) -> frozenset[str]: + """A cased entry's declared frame rather than the union of its arms — a narrower arm broadcasts — and a plain entry's body.""" + declared = schema.expressions[node.name].dims + if declared is not None: + return frozenset(declared) + return dims_of(node.body, schema, context) def _not_carried(context: str, call: str, inner: frozenset[str], rewrite: str) -> str: @@ -134,34 +122,17 @@ def _not_carried(context: str, call: str, inner: frozenset[str], rewrite: str) - ) -def _dims_call(node: FunctionCallNode, schema: Spec, context: str) -> frozenset[str]: - """The dim rule of the operator *node* calls, applied to the dims its operand carries.""" - inner = _dims(node.args[0], schema, context) - return _CALL_RULES[node.name](node, inner, schema, context) - +def _sum_dims(node: Sum, inner: frozenset[str], context: str) -> frozenset[str]: + """``sum`` reduces each named dim away, so the operand carries every one.""" + for consumed in node.over: + if consumed not in inner: + raise DimensionError(_not_carried(context, f'sum(over={consumed})', inner, 'drop the sum, or fix the dim')) + return inner - frozenset(node.over) -def _sum_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spec, context: str) -> frozenset[str]: - """``sum`` reduces a dim away, or reads a relation: the consumed dim goes, the produced dims arrive, the joined stay.""" - by = node.kwargs.get('by') - if by is None and 'over' not in node.kwargs: - if not inner: - raise DimensionError( - f'{context}: sum() with no over= or by= sums every dim the operand ' - f'carries, and this one carries none — the expression is already a ' - f'scalar. Drop the sum.' - ) - return frozenset() - if by is None: - consumed = node.kwargs['over'] - assert isinstance(consumed, DimensionNode) - if consumed.name not in inner: - raise DimensionError( - _not_carried(context, f'sum(over={consumed.name})', inner, 'drop the sum, or fix the dim') - ) - return inner - {consumed.name} - assert isinstance(by, DirectionNode), 'resolution reads sum(by=) in a direction' - direction = by.direction +def _group_sum_dims(node: GroupSum, inner: frozenset[str], context: str) -> frozenset[str]: + """``sum`` through a relation: the consumed dim goes, the produced dims arrive, the joined stay.""" + direction = node.direction if missing := sorted(set(direction.consumed_dims) - inner): raise DimensionError( _not_carried( @@ -174,11 +145,9 @@ def _sum_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spec, conte return _read_dims(f'sum(by={direction.name})', direction, inner, context) -def _at_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spec, context: str) -> frozenset[str]: +def _at_dims(node: Pullback, inner: frozenset[str], context: str) -> frozenset[str]: """``at`` is the adjoint of ``sum(by=)``: it consumes the dims a sum produces and produces the ones it consumes.""" - by = node.kwargs['by'] - assert isinstance(by, DirectionNode), 'resolution reads at(by=) in a direction' - return pulled_back_dims(by.direction, inner, context, 'the expression') + return pulled_back_dims(node.direction, inner, context, 'the expression') def pulled_back_dims(direction: Direction, inner: frozenset[str], context: str, operand: str) -> frozenset[str]: @@ -198,26 +167,25 @@ def pulled_back_dims(direction: Direction, inner: frozenset[str], context: str, return _read_dims(f'at(by={direction.name})', direction, inner, context) -def _translation_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spec, context: str) -> frozenset[str]: - """``shift`` and ``sum_back`` keep every dim, and their amount, edge and partition are checked here.""" - over = node.kwargs['along'] - assert isinstance(over, DimensionNode) - if over.name not in inner: +#: The verb a file writes each translation with, which its refusals quote. +_VERBS: dict[type[Translate | WindowSum], str] = {Translate: 'shift', WindowSum: 'sum_back'} + + +def _translation_dims(node: Translate | WindowSum, inner: frozenset[str], schema: Spec, context: str) -> frozenset[str]: + """``shift`` and ``sum_back`` keep every dim, and a named amount and a partition are checked here.""" + verb = _VERBS[type(node)] + if node.along not in inner: raise DimensionError( _not_carried( context, - f'{node.name}(along={over.name})', + f'{verb}(along={node.along})', inner, - f'name a dim the operand carries, or drop the {node.name}', + f'name a dim the operand carries, or drop the {verb}', ) ) - _check_named_amount(node, over.name, inner, schema, context) - _check_amount_form(node, context) - _check_edge(node, context) - by = node.kwargs.get('by') - if by is not None: - assert isinstance(by, PartitionNode), "resolution reads a translation's by= as a partition" - _check_joined(f'{node.name}(along={over.name}, by={by.partition.name})', by.partition, inner, context) + _check_named_amount(node, verb, inner, schema, context) + if node.partition is not None: + _check_joined(f'{verb}(along={node.along}, by={node.partition.name})', node.partition, inner, context) return inner @@ -267,191 +235,38 @@ def _check_joined(call: str, use: Direction | Partition, inner: frozenset[str], ) -#: The dim rule of each built-in, by name. -_CALL_RULES: dict[str, Callable[[FunctionCallNode, frozenset[str], Spec, str], frozenset[str]]] = { - 'sum': _sum_dims, - 'at': _at_dims, - 'shift': _translation_dims, - 'sum_back': _translation_dims, -} - - -class _Amount(NamedTuple): - """What the errors of an operator that steps along an axis say about the amount it takes.""" - - #: The word for the amount. - noun: str - #: Why negating a named one at the call site is not what the caller means. - negated: str - #: What a named one that varies over the axis it steps along becomes. - varies: str - #: The least whole number a literal may be. - minimum: float - #: What a literal must be written as, after ``operator(kwarg=...)``. - form: str - - -_AMOUNTS = { - 'shift': _Amount( - 'offset', - 'A named offset carries its sign in its values, so that one row pointing backwards says ' - 'so where the data is read — negate the column instead.', - 'a permutation rather than a lag', - -math.inf, - 'must be a whole number, or the name of an integer parameter when the offset differs per ' - 'entity — a lead time, a transit time, a minimum up time.', - ), - 'sum_back': _Amount( - 'width', - 'A width counts positions and so has no direction; which way a window reaches is the ' - "operator's own name rather than the sign of its width.", - 'a different window at every position, which is no longer "the last n"', - 1, - 'needs a whole number of positions of at least 1, or the name of an integer parameter when ' - 'the window differs per entity. A width of 1 is the operand itself.', - ), -} - - -def _amount_of(node: FunctionCallNode) -> tuple[str, ArithmeticNode]: - """The kwarg an operator that steps along an axis takes its amount through, and the value written there.""" - (kwarg,) = BUILTINS[node.name].required_value_kwargs - return kwarg, node.kwargs[kwarg] - - -def _whole(node: ArithmeticNode, minimum: float) -> bool: - """Whether *node* is a literal whole number of at least *minimum*.""" - return isinstance(node, NumberNode) and int(node.value) == node.value and node.value >= minimum - - -def _check_amount_form(node: FunctionCallNode, context: str) -> None: - """An ``offset=`` or ``window=`` is a whole number in the operator's range, or a parameter name.""" - kwarg, amount = _amount_of(node) - if isinstance(amount, ParameterNode) or _whole(amount, _AMOUNTS[node.name].minimum): - return - raise DimensionError(f'{context}: {node.name}({kwarg}=...) {_AMOUNTS[node.name].form}') - - -def _check_edge(node: FunctionCallNode, context: str) -> None: - """What an ``edge=`` may say, and where saying nothing is an answer. - - Every rule here is decidable from the file — whether the operand carries a - variable, whether the offset is named, what the edge is written as — so a - file breaking one is refused at load rather than by whoever lowers it. - """ - edge = node.kwargs.get('edge') - if node.name == 'sum_back': - if edge is not None and not isinstance(edge, EdgeNode): - raise DimensionError( - f"{context}: sum_back(edge=...) takes 'wrap' or nothing. A window sums the terms " - f'it reaches, so a position before the first contributes nothing rather than a ' - f'fill value; add the constant to the expression if you want one.' - ) - return - - if isinstance(edge, EdgeNode): - return - fill = _edge_fill(edge, context) - has_var = degree.carries_variable(node.args[0]) - if has_var and fill is not None and fill != 0: - raise DimensionError( - f'{context}: shift(edge={fill:g}) over an expression containing a variable — only ' - f'fill=0 is representable there, since a vacated slot contributes no term. A nonzero ' - f'fill would be a constant standing where a term was; add that constant to the ' - f'expression instead.' - ) - offset = node.kwargs['offset'] - if fill is None and _vacates(offset) and not has_var: - raise DimensionError(_shift_over_data_message(context)) - if fill is None and isinstance(offset, ParameterNode): - raise DimensionError(f'{context}: {_named_offset_edge_message(offset.name)}') - - -def _vacates(offset: ArithmeticNode) -> bool: - """Whether a translation leaves anything behind. - - A literal zero step reaches every coordinate from itself, so there is no - vacated position for an ``edge=`` to answer for and the refusal below has - nothing to refuse. A *named* offset may be zero in the data and is not - known here, so it vacates until proved otherwise. - """ - return not (isinstance(offset, NumberNode) and offset.value == 0) - - -def _edge_fill(edge: ArithmeticNode | None, context: str) -> float | None: - """The number an ``edge=`` names, or ``None`` where it names nothing.""" - if edge is None: - return None - assert isinstance(edge, NumberNode), ( - f'{context}: resolution refuses an edge that is neither wrap nor a number first' - ) - return edge.value - - -def _named_offset_edge_message(name: str) -> str: - """Why a named offset must say what the vacated positions contribute. - - The absent edge propagates through a presence frame keyed by the translated - dimension alone, and a per-entity offset vacates a different slot for each - entity — which that frame cannot say. Refused rather than answered wrongly - (#850); the two edges that write their own answer are allowed. - """ - return ( - f'shift(offset={name}) leaves the vacated positions absent, which a ' - f'per-entity offset cannot say yet.\n' - f"Add edge='wrap' for a cyclic translation, or edge= for what the " - f'vacated positions contribute.' - ) - - -def _shift_over_data_message(context: str) -> str: - """The three ways out of a translation over data with no ``edge=``, the third being two things at once.""" - return ( - f'{context}: shift() over a variable-free expression leaves vacated positions with no ' - f'value, and inventing one is what silently pinned a bound to zero. Say which you mean:\n' - f" shift(x, along=d, offset=n, edge='wrap') the dimension really is cyclic\n" - f' shift(x, along=d, offset=n, edge=0) the vacated positions contribute zero\n' - f' ...and a where: excluding them the vacated rows should not exist at all\n' - f'A where: alone does not lift this — it is decided on the expression, before any mask ' - f'is read — and edge=0 alone leaves a row whose bound is that zero.' - ) - - -def _check_named_amount(node: FunctionCallNode, over: str, inner: frozenset[str], schema: Spec, context: str) -> None: +def _check_named_amount( + node: Translate | WindowSum, verb: str, inner: frozenset[str], schema: Spec, context: str +) -> None: """The rules that hold of an ``offset=`` or ``window=`` naming a parameter; a literal breaks none of them.""" - kwarg, amount = _amount_of(node) - words = _AMOUNTS[node.name] - if isinstance(amount, UnaryOperatorNode) and isinstance(amount.operand, ParameterNode): - raise DimensionError( - f'{context}: {node.name}({kwarg}={amount.op}{amount.operand.name}) negates a named {words.noun}. {words.negated}' - ) - if not isinstance(amount, ParameterNode): + kwarg, amount = ('offset', node.offset) if isinstance(node, Translate) else ('window', node.width) + if not isinstance(amount, str): return - declared = schema.parameters[amount.name] + words = AMOUNTS[verb] + declared = schema.parameters[amount] if declared.dtype != 'int': raise DimensionError( - f'{context}: {node.name}({kwarg}={amount.name}) counts positions along ' - f"'{over}', but '{amount.name}' is declared dtype: {declared.dtype}. A count of " + f'{context}: {verb}({kwarg}={amount}) counts positions along ' + f"'{node.along}', but '{amount}' is declared dtype: {declared.dtype}. A count of " f'positions is integral — declare it dtype: int, which binds only an integer ' f'column, so a fractional {words.noun} has nowhere to arrive from.' ) - if over in declared.dims: + if node.along in declared.dims: raise DimensionError( - f'{context}: {node.name}({kwarg}={amount.name}) steps along ' - f"'{over}', but '{amount.name}' is declared over {sorted(declared.dims)}, which " + f'{context}: {verb}({kwarg}={amount}) steps along ' + f"'{node.along}', but '{amount}' is declared over {sorted(declared.dims)}, which " f'carries it. A named {words.noun} that varies over the axis it steps along is {words.varies} ' - f"— declare '{amount.name}' over dims '{over}' is not one of." + f"— declare '{amount}' over dims '{node.along}' is not one of." ) - by = node.kwargs.get('by') groups = ( - frozenset(by.partition.dim(v) for v in by.partition.group) if isinstance(by, PartitionNode) else frozenset() + frozenset(node.partition.dim(v) for v in node.partition.group) if node.partition is not None else frozenset() ) if stray := sorted(frozenset(declared.dims) - inner - groups): raise DimensionError( - f'{context}: {node.name}({kwarg}={amount.name}) reads its {words.noun} at the coordinate it ' - f"steps from, but '{amount.name}' varies over {stray}, which that coordinate does not carry " + f'{context}: {verb}({kwarg}={amount}) reads its {words.noun} at the coordinate it ' + f"steps from, but '{amount}' varies over {stray}, which that coordinate does not carry " f'(dims {sorted(inner)}). A dim the coordinate does not have is no coordinate at all — ' - f"declare '{amount.name}' over dims the expression carries, or group by a relation into " + f"declare '{amount}' over dims the expression carries, or group by a relation into " f'one of {stray}, so that each group is reached by its own {words.noun}.' ) @@ -482,21 +297,22 @@ def check_schema(schema: Spec, resolved: Resolved) -> None: f'{sorted(frame)}.' ) - for ename, node in resolved.expressions.items(): - if not isinstance(node, CasesNode): + for ename, entry in resolved.expressions.items(): + if not isinstance(entry.body, Cases): continue - frame = frozenset(schema.expressions[ename].dims or []) - for arm in node.arms: - context = case_context(ename, None if arm.when is None else arm.label) - if arm.when is not None: - _check_where_dims(Mask(arm.when), frame, context) - _check_value_dims(arm.value, schema, frame, context) - - for cname, (expression, where) in resolved.constraints.items(): - frame = frozenset(schema.constraints[cname].dims) + block = schema.expressions[ename] + frame = frozenset(block.dims or []) + for region, label in zip(entry.body.regions, [*block.cases, None], strict=True): + context = case_context(ename, label) + if label is not None: + _check_where_dims(region.when, frame, context) + _check_value_dims(region.value, schema, frame, context) + + for cname, constraint in resolved.constraints.items(): + frame = frozenset(constraint.dims) context = f"Constraint '{cname}'" - _check_where_dims(where, frame, context) - got = dims_of(expression, schema, context) + _check_where_dims(constraint.where, frame, context) + got = dims_of(constraint.lhs, schema, context) | dims_of(constraint.rhs, schema, context) if got != frame: stray, missing = sorted(got - frame), sorted(frame - got) detail = ( @@ -512,7 +328,7 @@ def check_schema(schema: Spec, resolved: Resolved) -> None: if resolved.objective is not None: context = 'The objective' - got = dims_of(resolved.objective, schema, context) + got = dims_of(resolved.objective.expression, schema, context) if got: raise DimensionError( f'{context}: the expression carries dims {sorted(got)}, and an objective is one ' @@ -521,7 +337,7 @@ def check_schema(schema: Spec, resolved: Resolved) -> None: ) -def _check_value_dims(node: ArithmeticNode, schema: Spec, frame: frozenset[str], context: str) -> None: +def _check_value_dims(node: Expression, schema: Spec, frame: frozenset[str], context: str) -> None: """A region's value may only carry dims the frame does — the ``otherwise:`` included. A wider one would give the quantity dims its declaration does not, which is diff --git a/src/math_spec/errors.py b/src/math_spec/errors.py index 2c7ad7ad..97220d2b 100644 --- a/src/math_spec/errors.py +++ b/src/math_spec/errors.py @@ -60,10 +60,6 @@ class DimensionError(LanguageError): """A dim-set rule was violated. Raised at load time, before any data.""" -class PiecewiseExpansionError(LanguageError): - """A piecewise block references something that doesn't exist or collides.""" - - def did_you_mean(name: str, known: Iterable[str], *, label: str = 'Declared') -> str: """The repair clause for an unrecognised name: the near miss, or the set.""" candidates = sorted(known) @@ -97,3 +93,14 @@ def schema_error(exc: ValidationError) -> LanguageError: def prefixed(context: str, e: ValueError) -> str: """*e* under *context*, once — an expansion error already carries it.""" return str(e) if str(e).startswith(context) else f'{context}: {e}' + + +def case_context(name: str, label: str | None) -> str: + """The context an error inside one arm of a cased expression is reported under. + + Args: + name: The named expression the arm belongs to. + label: The case's name, or ``None`` for the block's ``otherwise:``. + """ + where = 'otherwise' if label is None else f"case '{label}'" + return f"Named expression '{name}', {where}" diff --git a/src/math_spec/exclusivity.py b/src/math_spec/exclusivity.py index 930453b7..13325fdc 100644 --- a/src/math_spec/exclusivity.py +++ b/src/math_spec/exclusivity.py @@ -21,17 +21,19 @@ from enum import Enum from typing import TYPE_CHECKING, Literal, assert_never -from math_spec._expression_parser import NumberNode, ParameterNode, UnaryOperatorNode from math_spec.program import ( And, BooleanLiteral, + Constant, CountComparison, DimensionComparison, DimensionPosition, ExpressionComparison, Mask, + Negate, Not, Or, + Parameter, ParameterComparison, ParameterDefined, PulledBackPredicate, @@ -46,9 +48,8 @@ if TYPE_CHECKING: from collections.abc import Iterable, Iterator, Mapping - from math_spec._expression_parser import ArithmeticNode from math_spec.model import DeclaredDtype - from math_spec.program import Predicate, PredicateOperator + from math_spec.program import Expression, Predicate, PredicateOperator #: The most cells one pair may multiply out to; a pair past it is several expressions. CELL_BUDGET = 8192 @@ -228,7 +229,7 @@ def _undecided(mask: Predicate) -> str | None: return None -def _expression_rewrite(node: ExpressionComparison[ArithmeticNode]) -> str: +def _expression_rewrite(node: ExpressionComparison) -> str: """Why a comparison of expressions is not decided, and what to write instead. A parameter against a literal is decided, and the same test with its sides @@ -238,13 +239,11 @@ def _expression_rewrite(node: ExpressionComparison[ArithmeticNode]) -> str: never reaches here, and a quoted label cannot stand on the left at all. """ left, right = node.left, node.right - number = isinstance(left, NumberNode) or ( - isinstance(left, UnaryOperatorNode) and isinstance(left.operand, NumberNode) - ) - if number and isinstance(right, ParameterNode): + number = _signed_literal(left) + if number is not None and isinstance(right, Parameter): return ( f'the literal is on the left, and a comparison is read as arithmetic there — write it as ' - f'the same test the other way round, {right.name} {_FLIPPED[node.op]} {left}' + f'the same test the other way round, {right.name} {_FLIPPED[node.op]} {number:g}' ) return ( 'it compares expressions, whose values only the data decides — compare one parameter against a ' @@ -252,6 +251,15 @@ def _expression_rewrite(node: ExpressionComparison[ArithmeticNode]) -> str: ) +def _signed_literal(node: Expression) -> float | None: + """The number *node* is, its sign folded in — ``None`` where it is not a literal.""" + if isinstance(node, Constant): + return node.value + if isinstance(node, Negate) and isinstance(node.operand, Constant): + return -node.operand.value + return None + + def _observe( node: TypedPredicate, subject: Subject, values: set[_Literal], dtypes: Mapping[str, DeclaredDtype] ) -> None: diff --git a/src/math_spec/expansion.py b/src/math_spec/expansion.py index 0af39a52..5200e9ab 100644 --- a/src/math_spec/expansion.py +++ b/src/math_spec/expansion.py @@ -2,7 +2,11 @@ # # SPDX-License-Identifier: MIT -"""Named sub-expressions and macros, expanded into the core AST before anything reads the expression.""" +"""Macro calls, expanded into the syntax tree before resolution reads it. + +A named expression is resolution's: it resolves the entry once and puts that +node where the name stood. +""" from __future__ import annotations @@ -25,43 +29,33 @@ def parse_and_expand(text: str, ns: Namespace, context: str) -> ParsedNode: - """Parse *text* and expand named sub-expressions and macros to core AST. + """Parse *text* and expand every macro call in it. Args: text: The expression as the file wrote it. - ns: Where names and macros are declared, and where a named expression is resolved. + ns: Where the macros are declared. context: What an error names. """ return expand(parse_expression(text), ns, context) @overload -def expand(node: ArithmeticNode, ns: Namespace, context: str, *, shadow: frozenset[str] = ...) -> ArithmeticNode: ... +def expand(node: ArithmeticNode, ns: Namespace, context: str) -> ArithmeticNode: ... @overload -def expand(node: ComparisonNode, ns: Namespace, context: str, *, shadow: frozenset[str] = ...) -> ComparisonNode: ... - +def expand(node: ComparisonNode, ns: Namespace, context: str) -> ComparisonNode: ... -def expand(node: ParsedNode, ns: Namespace, context: str, *, shadow: frozenset[str] = frozenset()) -> ParsedNode: - """Expand all named sub-expressions and macro calls under *node*. - A comparison stays a comparison and an arithmetic node stays arithmetic. A - named expression arrives as the node :meth:`Namespace.named` resolved it - to, once for every use. +def expand(node: ParsedNode, ns: Namespace, context: str) -> ParsedNode: + """Expand every macro call under *node*; a comparison stays a comparison and arithmetic stays arithmetic. Args: node: The parsed expression. - ns: Where names and macros are declared. + ns: Where the macros are declared. context: What an error names. - shadow: Names left as written even where a named expression has that - name — a template's formals, checked without a call to bind them. """ if isinstance(node, ComparisonNode): - return ComparisonNode( - node.op, - _expand(node.left, ns, context, (), shadow), - _expand(node.right, ns, context, (), shadow), - ) - return _expand(node, ns, context, (), shadow) + return ComparisonNode(node.op, _expand(node.left, ns, context, ()), _expand(node.right, ns, context, ())) + return _expand(node, ns, context, ()) def macro_signature(name: str, macro: MacroBlock) -> str: @@ -79,32 +73,16 @@ def parse_template(name: str, macro: MacroBlock, context: str) -> ArithmeticNode return body -def _expand( - node: ArithmeticNode, - ns: Namespace, - context: str, - stack: tuple[str, ...], - shadow: frozenset[str], -) -> ArithmeticNode: - if isinstance(node, NameNode) and node.name in ns.schema.expressions and node.name not in shadow: - return ns.named(node.name, context) - +def _expand(node: ArithmeticNode, ns: Namespace, context: str, stack: tuple[str, ...]) -> ArithmeticNode: if isinstance(node, FunctionCallNode) and node.name in ns.schema.macros: if node.name in stack: msg = f'{context}: circular macro reference: {" -> ".join([*stack, node.name])}' raise SchemaError(msg) - return _expand_macro(node, ns, context, stack, shadow) - - return with_children(node, lambda child: _expand(child, ns, context, stack, shadow)) + return _expand_macro(node, ns, context, stack) + return with_children(node, lambda child: _expand(child, ns, context, stack)) -def _expand_macro( - call: FunctionCallNode, - ns: Namespace, - context: str, - stack: tuple[str, ...], - shadow: frozenset[str], -) -> ArithmeticNode: +def _expand_macro(call: FunctionCallNode, ns: Namespace, context: str, stack: tuple[str, ...]) -> ArithmeticNode: """Call-by-value: arguments are expanded before substitution, and the substituted body is expanded again.""" macro = ns.schema.macros[call.name] signature = macro_signature(call.name, macro) @@ -123,12 +101,12 @@ def _expand_macro( raise SchemaError(msg) bindings = { - **{formal: _expand(arg, ns, context, stack, shadow) for formal, arg in zip(macro.args, call.args, strict=True)}, - **{formal: _expand(call.kwargs[formal], ns, context, stack, shadow) for formal in macro.kwargs}, + **{formal: _expand(arg, ns, context, stack) for formal, arg in zip(macro.args, call.args, strict=True)}, + **{formal: _expand(call.kwargs[formal], ns, context, stack) for formal in macro.kwargs}, } body = parse_template(call.name, macro, context) substituted = _substitute(body, bindings) - return _expand(substituted, ns, context, (*stack, call.name), shadow) + return _expand(substituted, ns, context, (*stack, call.name)) def _substitute(node: ArithmeticNode, bindings: dict[str, ArithmeticNode]) -> ArithmeticNode: diff --git a/src/math_spec/lowering.py b/src/math_spec/lowering.py index 396f8213..61aa81ca 100644 --- a/src/math_spec/lowering.py +++ b/src/math_spec/lowering.py @@ -4,61 +4,29 @@ """Lower a validated model to a :class:`~math_spec.program.Program`. -One lowering, on the language side: it reads the typed AST and emits -declarations with names resolved and shapes fixed, and reaches no consumer. A -construct with no lowering raises :class:`~math_spec.errors.LanguageError` -naming its rewrite. +One lowering, on the language side: it packages the declarations a model +resolved to, with every named expression inlined where the math reads it, and +reaches no consumer. A construct with no lowering raises +:class:`~math_spec.errors.LanguageError` naming its rewrite. """ from __future__ import annotations -from dataclasses import dataclass, replace +from dataclasses import replace from typing import TYPE_CHECKING, assert_never import math_spec.program as program -from math_spec._expression_parser import ( - ArithmeticNode, - BinaryOperatorNode, - CasesNode, - DefinitionNode, - DimensionNode, - DirectionNode, - DualNode, - EdgeNode, - FunctionCallNode, - KwargNode, - NumberNode, - ParameterNode, - PartitionNode, - UnaryOperatorNode, - UnresolvedNode, - VariableNode, -) -from math_spec.dimensions import dims_of from math_spec.errors import LanguageError from math_spec.piecewise import declaration_of from math_spec.validation import to_spec if TYPE_CHECKING: - from collections.abc import Callable, Mapping + from collections.abc import Mapping from pathlib import Path from math_spec.model import Spec -def _none_of(masks: list[program.Mask]) -> program.Mask: - """The region left over: where not one of *masks* holds. - - The ``otherwise`` arm's own mask, built rather than written. An empty list - cannot reach here — ``cases:`` carries at least one case — so there is no - vacuous truth to spell. - """ - remainder = ~masks[0] - for mask in masks[1:]: - remainder = remainder & ~mask - return remainder - - def to_program(spec: str | Path | Mapping[str, object] | Spec | program.Program) -> program.Program: """*spec* as a :class:`~math_spec.program.Program` — the public door. @@ -130,47 +98,34 @@ def lower_program(expanded: Spec) -> program.Program: lower, upper = _bound_expression(vdef.bounds.lower), _bound_expression(vdef.bounds.upper) variables[vname] = program.VariableDeclaration( tuple(vdef.dims), - where=_Lowering(expanded, f"variable '{vname}'").mask(resolved.variables[vname]), + where=_inlined_mask(resolved.variables[vname]), lower=lower, upper=upper, domain=domain, absence=vdef.absence, ) - constraints = {} - for cname, cdef in expanded.constraints.items(): - expression, where = resolved.constraints[cname] - lowering = _Lowering(expanded, f"constraint '{cname}'") - constraints[cname] = program.ConstraintDeclaration( - tuple(cdef.dims), - lhs=lowering.expr(expression.left), - sense=expression.op, - rhs=lowering.expr(expression.right), - where=lowering.mask(where), - ) - + constraints = { + cname: replace(c, lhs=inline(c.lhs), rhs=inline(c.rhs), where=_inlined_mask(c.where)) + for cname, c in resolved.constraints.items() + } objective = None - if (odef := expanded.objective) is not None: - assert resolved.objective is not None, 'validation resolves the objective the file declares' - objective = program.ObjectiveDeclaration( - odef.sense, - _Lowering(expanded, 'the objective').expr(resolved.objective), - ) + if resolved.objective is not None: + objective = replace(resolved.objective, expression=inline(resolved.objective.expression)) dimensions = {dname: program.DimensionDeclaration(ddef.dtype) for dname, ddef in expanded.dimensions.items()} sos = { - sname: program.SosDeclaration( - sdef.variable, - sdef.over, - sos_type=sdef.type, - ) + sname: program.SosDeclaration(sdef.variable, sdef.over, sos_type=sdef.type) for sname, sdef in expanded.sos.items() } - expressions: dict[str, program.ExpressionDeclaration] = {} - for name, ast in resolved.expressions.items(): - expressions[name] = program.ExpressionDeclaration( - _Lowering(expanded, f"named expression '{name}'").expr(ast), in_math=name in resolved.read_by_the_math - ) + expressions = { + name: program.ExpressionDeclaration(inline(entry), in_math=name in resolved.read_by_the_math) + for name, entry in resolved.expressions.items() + } + assumptions = { + name: replace(holds, predicate=inline_mask(holds.predicate), where=_inlined_mask(holds.where)) + for name, holds in resolved.assumptions.items() + } return program.Program( parameters=parameters, variables=variables, @@ -180,228 +135,61 @@ def lower_program(expanded: Spec) -> program.Program: relations=resolved.relations, sos=sos, piecewise={name: declaration_of(pw) for name, pw in expanded._expanded_piecewise.items()}, - assumptions=_assumptions(expanded), + assumptions=assumptions, expressions=expressions, ) -def _assumptions(expanded: Spec) -> dict[str, program.Assumption]: - """Everything the data has to satisfy, in the order the model states it. +def inline(node: program.Expression | program.Named) -> program.Expression: + """*node* with every :class:`~math_spec.program.Named` replaced by its body — the tree a program carries. - One mapping rather than two, because a consumer binding data checks them - all the same way and refuses in the same words. A curve's conditions are - already here: the expansion writes them into ``assumptions:``, and a load - derives the same text for a block the file still declares. + A region's ``when`` is inlined with its value, since a mask may compare + expressions that name an entry. """ - assumptions: dict[str, program.Assumption] = {} - for name, (holds, where, description) in expanded.resolved.assumptions.items(): - lowering = _Lowering(expanded, f"assumption '{name}'") - predicate = lowering.mask(holds) - assert predicate is not None, 'a predicate that admits every row was refused as deciding nothing' - assumptions[name] = program.Holds(predicate, lowering.mask(where), description) - return assumptions - - -# --------------------------------------------------------------------------- -# expression lowering -# --------------------------------------------------------------------------- - - -@dataclass(frozen=True) -class _Lowering: - """One expression walk, and the two things every step of it reads.""" - - schema: Spec - context: str - - def expr(self, node: ArithmeticNode) -> program.Expression: - """Rewrite one resolved core-AST expression as a program expression.""" - if isinstance(node, NumberNode): - return program.Constant(node.value) - - if isinstance(node, VariableNode): - return program.Variable(node.name) - - if isinstance(node, ParameterNode): - return program.Parameter(node.name) - - if isinstance(node, UnresolvedNode | KwargNode): - msg = f'{node!r} reached lowering. Expressions go through resolution.resolve_expression() first.' - raise AssertionError(msg) - - if isinstance(node, DualNode): - return program.Dual(node.constraint) - - if isinstance(node, UnaryOperatorNode): - inner = self.expr(node.operand) - return program.Negate(inner) if node.op == '-' else inner - - if isinstance(node, BinaryOperatorNode): - left = self.expr(node.left) - right = self.expr(node.right) - match node.op: - case '+': - return program.Add(left, right) - case '-': - return program.Add(left, program.Negate(right)) - case '*': - return program.Multiply(left, right) - case '/': - return program.Divide(left, right) - case '**': - return program.Power(left, right) - case _: # pragma: no cover — the parser admits no other operator - raise AssertionError(f'{self.context}: operator {node.op!r} reached lowering') - - if isinstance(node, FunctionCallNode): - return _CALLS[node.name](self, node) - - if isinstance(node, CasesNode): - return self._cases(node) - - if isinstance(node, DefinitionNode): - return self.expr(node.body) - - assert_never(node) - - def _cases(self, node: CasesNode) -> program.Cases: - """A cased expression, with every region carrying the mask it applies under. - - The ``otherwise`` arm carries no ``when`` in the file; here it carries - the negation of every other region's, so a consumer adds regions rather - than working out which one is left. The language proved the rest apart - before this ran, so the negation is exactly the remainder and the - regions stay disjoint and total. - - Every ``when`` arrives folded from resolution, and an arm that folded - to a literal was refused at load — so no literal reaches a region. - """ - stated = [program.Mask(self._predicate(arm.when)) for arm in node.arms if arm.when is not None] - regions = [] - for arm in node.arms: - when = program.Mask(self._predicate(arm.when)) if arm.when is not None else _none_of(stated) - regions.append(program.Region(when, self.expr(arm.value))) - return program.Cases(tuple(regions)) - - def mask(self, mask: program.Mask | None) -> program.Mask | None: - """*mask* with every comparison of expressions lowered, so a program's masks are program vocabulary throughout. - - 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 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.ExpressionComparison): - 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 | program.PulledBackPredicate): - 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): - return program.And(self._predicate(node.left), self._predicate(node.right)) - if isinstance(node, program.Or): - return program.Or(self._predicate(node.left), self._predicate(node.right)) + if isinstance(node, program.Named): + return inline(node.body) + if isinstance(node, program.Constant | program.Parameter | program.Variable | program.Dual): return node - - def sum(self, node: FunctionCallNode) -> program.Expression: - """``sum(x)``, ``sum(x, over=d)`` or ``sum(x, by=relation)``. - - Two program nodes under one surface verb: reducing a dim away and reducing it - *into* another are different relational shapes, so ``by=`` decides which - before anything else is read. - """ - by_node = node.kwargs.get('by') - operand = self.expr(node.args[0]) - if by_node is None and 'over' not in node.kwargs: - return program.Sum(operand, tuple(sorted(dims_of(node.args[0], self.schema, self.context)))) - if by_node is None: - consumed = node.kwargs['over'] - assert isinstance(consumed, DimensionNode), 'resolution refuses a over= that is not a dimension' - return program.Sum(operand, (consumed.name,)) - assert isinstance(by_node, DirectionNode), 'resolution reads sum(by=) in a direction' - return program.GroupSum(operand, direction=by_node.direction) - - def at(self, node: FunctionCallNode) -> program.Expression: - """``at(x, by=relation)`` — the adjoint of :meth:`sum`'s ``by=`` form.""" - by_node = node.kwargs['by'] - assert isinstance(by_node, DirectionNode), 'resolution reads at(by=) in a direction' - return program.Pullback(self.expr(node.args[0]), direction=by_node.direction) - - def sum_back(self, node: FunctionCallNode) -> program.Expression: - """``sum_back(x, along=d, window=w)`` — a trailing window along one dimension. - - *window* is an integer literal of at least one, or a parameter naming a - per-entity width, which the language holds to the two rules that make it - mean one thing before this is reached. - - ``by=`` names the relation the window stops at the edges of, and rides on - the node the way it rides on a translation — the dim rules have already - held it to one relation over the dimension stepped along. - """ - over_node = node.kwargs['along'] - assert isinstance(over_node, DimensionNode), 'resolution refuses an along= that is not a dimension' - operand = self.expr(node.args[0]) - wrap = isinstance(node.kwargs.get('edge'), EdgeNode) - return program.WindowSum( - operand, over_node.name, width=_amount(node.kwargs['window']), wrap=wrap, partition=_partition_of(node) - ) - - def shift(self, node: FunctionCallNode) -> program.Expression: - """``shift(x, along=d, offset=n)`` — the value at *t - offset* along one dim. - - What the vacated positions contribute is ``edge=``'s to say, and the - language has already held it to the keyword or a number. - """ - over_node = node.kwargs['along'] - assert isinstance(over_node, DimensionNode), 'resolution refuses an along= that is not a dimension' - operand = self.expr(node.args[0]) - edge = node.kwargs.get('edge') - return program.Translate( - operand, - over_node.name, - offset=_amount(node.kwargs['offset']), - wrap=isinstance(edge, EdgeNode), - fill=edge.value if isinstance(edge, NumberNode) else None, - partition=_partition_of(node), - ) - - -#: One lowering per name in the language's ``BUILTIN_NAMES``. -_CALLS: dict[str, Callable[[_Lowering, FunctionCallNode], program.Expression]] = { - 'sum': _Lowering.sum, - 'at': _Lowering.at, - 'sum_back': _Lowering.sum_back, - 'shift': _Lowering.shift, -} - - -def _amount(node: ArithmeticNode) -> int | str: - """A translation's offset or a window's width: a literal step count, or the parameter that holds one per entity.""" - if isinstance(node, ParameterNode): - return node.name - assert isinstance(node, NumberNode), 'an offset= or window= that is neither is refused at load' - return int(node.value) - - -def _partition_of(node: FunctionCallNode) -> program.Partition | None: - """The partition a translation steps inside, if the call names a relation. - - That it is a *single* relation, stepped *along the translated dimension*, is - checked with the other dim rules (``math_spec.dimensions``), where a model - is refused before any data is read. - """ - by_node = node.kwargs.get('by') - if by_node is None: - return None - assert isinstance(by_node, PartitionNode), "resolution reads a translation's by= as a partition" - return by_node.partition + if isinstance(node, program.Negate): + return program.Negate(inline(node.operand)) + if isinstance(node, program.Add): + return program.Add(inline(node.left), inline(node.right)) + if isinstance(node, program.Multiply): + return program.Multiply(inline(node.left), inline(node.right)) + if isinstance(node, program.Power): + return program.Power(inline(node.base), inline(node.exponent)) + if isinstance(node, program.Divide): + return program.Divide(inline(node.numerator), inline(node.divisor)) + if isinstance(node, program.Sum | program.GroupSum | program.Pullback | program.Translate | program.WindowSum): + return replace(node, operand=inline(node.operand)) + if isinstance(node, program.Cases): + return program.Cases(tuple(program.Region(inline_mask(r.when), inline(r.value)) for r in node.regions)) + assert_never(node) + + +def inline_mask(mask: program.Mask) -> program.Mask: + """*mask* with every named expression its comparisons read inlined, as :func:`inline` does for a tree.""" + return program.Mask(_inline_predicate(mask.root)) + + +def _inlined_mask(mask: program.Mask | None) -> program.Mask | None: + return None if mask is None else inline_mask(mask) + + +def _inline_predicate(node: program.Predicate) -> program.Predicate: + if isinstance(node, program.ExpressionComparison): + return replace(node, left=inline(node.left), right=inline(node.right)) + if isinstance(node, program.CountComparison): + return replace(node, predicate=inline_mask(node.predicate)) + if isinstance(node, program.TranslatedPredicate | program.PulledBackPredicate): + return replace(node, operand=inline_mask(node.operand)) + if isinstance(node, program.Not): + return program.Not(_inline_predicate(node.operand)) + if isinstance(node, program.And): + return program.And(_inline_predicate(node.left), _inline_predicate(node.right)) + if isinstance(node, program.Or): + return program.Or(_inline_predicate(node.left), _inline_predicate(node.right)) + return node def _bound_expression(value: float | str) -> program.Expression: diff --git a/src/math_spec/model.py b/src/math_spec/model.py index bfc28102..c756590f 100644 --- a/src/math_spec/model.py +++ b/src/math_spec/model.py @@ -107,9 +107,6 @@ def _reject_unknown_keys(cls, data: object) -> object: #: Which way an objective is optimised (the declaration rules). ObjectiveSense = Literal['minimize', 'maximize'] -#: The relation a link may pin its expression to the curve with. -LinkSign = ComparisonOperator - #: The order of special ordered set. SosType = Literal[1, 2] @@ -343,8 +340,8 @@ class MacroBlock(_StrictBlock): """A parameterised expression template, defined in the YAML itself. Language, not code: formals (``args`` positional, ``kwargs`` keyword) - shadow model names inside the template, and every call site expands into - core AST before either backend sees the expression. + shadow model names inside the template, and every call site expands in + the syntax tree before resolution reads the expression. """ _label: ClassVar[str] = 'a macro declaration' @@ -550,7 +547,7 @@ class PiecewiseLink(_StrictBlock): expression: str values: str - sign: LinkSign = '==' + sign: ComparisonOperator = '==' @model_validator(mode='before') @classmethod @@ -612,6 +609,11 @@ class PiecewiseBlock(_StrictBlock): points: str | None = None description: str | None = None + @property + def nominated(self) -> str | None: + """The block's own values parameter ``points:`` names, so the mask is derived from it — or ``None``.""" + return self.points if self.points in {link.values for link in self.links} else None + @property def curve(self) -> tuple[PiecewiseLink, PiecewiseLink]: """The two links as ``(x, y)``, the bounded one last. @@ -943,6 +945,8 @@ def _validate_references(self) -> Spec: *self._sos_shapes(), *self._sos_bounds(), *self._sos_emitted_names(), + *self._piecewise_references(), + *self._piecewise_emitted_names(), ] if errors: raise ValueError('\n'.join(errors)) @@ -1116,26 +1120,79 @@ def _sos_bounds(self) -> Iterator[str]: def _sos_emitted_names(self) -> Iterator[str]: """No name a set's expansion writes is one the file already declares.""" - declared: dict[str, Iterable[str]] = {'variable': self.variables, 'constraint': self.constraints} for sname, block in self.sos.items(): - for kind, names in Emitted.of(sname, block.type).by_kind: - yield from ( - f"Sos '{sname}': its expansion writes {kind} '{one}', which this file already declares. " - f'Rename one of them.' - for one in names - if one in declared[kind] + yield from self._collisions(f"Sos '{sname}'", Emitted.of(sname, block.type).by_kind) + + def _piecewise_references(self) -> Iterator[str]: + """A curve runs along a declared dimension through values parameters carrying it, gated by a binary, masked by a bool.""" + for name, pw in self.piecewise.items(): + context = f"piecewise '{name}'" + if pw.over not in self.dimensions: + yield undeclared_dimension('piecewise', name, pw.over) + for i, link in enumerate(pw.links): + if link.values not in self.parameters: + yield f"{context}: link {i} values references undeclared parameter '{link.values}'" + elif pw.over not in self.parameters[link.values].dims: + yield ( + f"{context}: link {i} values parameter '{link.values}' must carry dim " + f"'{pw.over}' (has {self.parameters[link.values].dims})" + ) + if (activity := pw.activity) is not None: + if activity not in self.variables: + yield ( + f"{context}: activity '{activity}' is not a declared variable. A gate is a binary variable; " + f'declare it, or drop activity: for weights that sum to 1.' + ) + elif self.variables[activity].domain != 'binary': + yield f"{context}: activity variable '{activity}' must be binary" + if (points := pw.points) is None or pw.nominated is not None: + continue + if points not in self.parameters: + yield f"{context}: points references undeclared parameter '{points}'" + elif (dtype := self.parameters[points].dtype) != 'bool': + yield ( + f"{context}: points parameter '{points}' is {dtype}, and a mask is a bool parameter — one " + f'saying, per breakpoint, whether the curve reaches it. Declare it dtype: bool.' + ) + elif pw.over not in self.parameters[points].dims: + yield ( + f"{context}: points parameter '{points}' must carry dim '{pw.over}' — " + f'it says how far each curve runs along it (has {self.parameters[points].dims})' ) + def _piecewise_emitted_names(self) -> Iterator[str]: + """No name a curve's expansion writes is one the file already declares.""" + from math_spec.piecewise import Emitted as EmittedCurve + + for name, pw in self.piecewise.items(): + yield from self._collisions(f"piecewise '{name}'", EmittedCurve.of(name, pw).by_kind) + + def _collisions(self, context: str, by_kind: Iterable[tuple[str, Iterable[str]]]) -> Iterator[str]: + """The refusal for each name *context*'s expansion writes that the file already declares, by kind.""" + declared: dict[str, Iterable[str]] = { + 'variable': self.variables, + 'constraint': self.constraints, + 'sos': self.sos, + 'assumption': self.assumptions, + } + for kind, names in by_kind: + yield from ( + f"{context}: its expansion writes {kind} '{one}', which this file already declares. Rename one of them." + for one in names + if one in declared[kind] + ) + @model_validator(mode='after') def _validate_expressions(self) -> Spec: """Every expression and where string — this file's own, and every one a curve emits. - A curve's expansion is a model in its own right, so validating it is - what holds the declarations it writes to the language; it runs first, - so a fault in a link is named against the link the file wrote. + This file's own first, so a fault in a link is named against the link + the file wrote, and the expansion reads the typed links rather than the + text again. A curve's expansion is a model in its own right, so + validating it is what holds the declarations it writes to the language. """ - self.expand('piecewise') _ = self.resolved + self.expand('piecewise') return self diff --git a/src/math_spec/operators.py b/src/math_spec/operators.py index 50dd48a7..cfa31e9e 100644 --- a/src/math_spec/operators.py +++ b/src/math_spec/operators.py @@ -10,8 +10,9 @@ from __future__ import annotations +import math from dataclasses import dataclass -from typing import TYPE_CHECKING, Literal +from typing import TYPE_CHECKING, Literal, NamedTuple if TYPE_CHECKING: from collections.abc import Iterable @@ -135,6 +136,44 @@ def kind_of( BUILTIN_NAMES = frozenset(BUILTINS) + +class Amount(NamedTuple): + """What the errors of an operator that steps along an axis say about the amount it takes.""" + + #: The word for the amount. + noun: str + #: Why negating a named one at the call site is not what the caller means. + negated: str + #: What a named one that varies over the axis it steps along becomes. + varies: str + #: The least whole number a literal may be. + minimum: float + #: What a literal must be written as, after ``operator(kwarg=...)``. + form: str + + +#: The amount each operator that steps along an axis takes, by operator name. +AMOUNTS: dict[str, Amount] = { + 'shift': Amount( + 'offset', + 'A named offset carries its sign in its values, so that one row pointing backwards says ' + 'so where the data is read — negate the column instead.', + 'a permutation rather than a lag', + -math.inf, + 'must be a whole number, or the name of an integer parameter when the offset differs per ' + 'entity — a lead time, a transit time, a minimum up time.', + ), + 'sum_back': Amount( + 'width', + 'A width counts positions and so has no direction; which way a window reaches is the ' + "operator's own name rather than the sign of its width.", + 'a different window at every position, which is no longer "the last n"', + 1, + 'needs a whole number of positions of at least 1, or the name of an integer parameter when ' + 'the window differs per entity. A width of 1 is the operand itself.', + ), +} + #: The one closed keyword an ``edge=`` accepts. Everything else in that #: position is a number: the value the vacated positions contribute. EDGE_WRAP = 'wrap' diff --git a/src/math_spec/piecewise.py b/src/math_spec/piecewise.py index f5378075..c00a0224 100644 --- a/src/math_spec/piecewise.py +++ b/src/math_spec/piecewise.py @@ -6,32 +6,27 @@ A block becomes ordinary affine declarations before anything reads the model, under names prefixed with the block's own; what each method emits is tabled in -``docs/reference/language/piecewise.md``. A link expression is judged before -expansion, so a refusal names the link the file wrote rather than an emitted -constraint. +``docs/reference/language/piecewise.md``. Every rule a block is held to is +decided at load, before this runs: the names it references in +:class:`~math_spec.model.Spec`, its links where every expression is typed, and +its frame in :func:`curve_frame`. """ from __future__ import annotations -from typing import TYPE_CHECKING, Literal, NamedTuple +from dataclasses import dataclass +from typing import TYPE_CHECKING, Literal -from math_spec._expression_parser import ComparisonNode -from math_spec.degree import check_expression +import math_spec.sos as sos from math_spec.dimensions import dims_of -from math_spec.errors import LanguageError, PiecewiseExpansionError -from math_spec.expansion import parse_and_expand -from math_spec.model import Curvature, PiecewiseBlock, PiecewiseMethod, Spec, undeclared_dimension +from math_spec.errors import DimensionError +from math_spec.model import AssumptionBlock, Curvature, PiecewiseBlock, PiecewiseMethod, Spec from math_spec.program import PiecewiseDeclaration -from math_spec.resolution import Namespace, resolve_expression -from math_spec.sos import Emitted, emit, section if TYPE_CHECKING: - from collections.abc import Iterable, Iterator + from collections.abc import Iterable - -def _nominated(pw: PiecewiseBlock) -> str | None: - """The block's own values parameter ``points:`` names, so the mask is derived from it — or ``None``.""" - return pw.points if pw.points in {link.values for link in pw.links} else None + from math_spec.program import Expression #: The suffix on the second gate row, where the gate variable does not exist. @@ -69,20 +64,7 @@ def declaration_of(pw: PiecewiseBlock) -> PiecewiseDeclaration: ) -class Assumed(NamedTuple): - """One condition a method puts on the numbers, as the language writes it. - - ``holds`` and ``where`` are where strings, resolved like any the file - wrote. ``description`` is the sentence a refusal quotes, which names the - method and the rewrite that takes a curve of any shape. - """ - - holds: str - where: str | None - description: str - - -def assumptions_of(block: str, pw: PiecewiseBlock) -> dict[str, Assumed]: +def assumptions_of(block: str, pw: PiecewiseBlock) -> dict[str, AssumptionBlock]: """What *block* assumes of its numbers, by the name the document prints and a refusal quotes. Every curve assumes its breakpoints are there: a missing parameter row is @@ -94,16 +76,17 @@ def assumptions_of(block: str, pw: PiecewiseBlock) -> dict[str, Assumed]: Read off the block rather than off an expansion, so a model states what it assumes whether or not its curves have been written out. Each condition is - a where string over the parameters the file declared: the expansion writes - them into ``assumptions:``, and a model that still declares the block - derives the same text at load. + an ``assumptions:`` entry over the parameters the file declared, its + ``description`` naming the method and the rewrite that takes a curve of any + shape: the expansion writes them into the model, and a model that still + declares the block resolves the same entries at load. """ d, mask = pw.over, pw.points - assumed: dict[str, Assumed] = {} - assumed[f'{block}_complete'] = Assumed( - ' AND '.join(dict.fromkeys(link.values for link in pw.links)), - mask, - f"piecewise '{block}': every breakpoint the curve runs through needs a row in " + assumed: dict[str, AssumptionBlock] = {} + assumed[f'{block}_complete'] = AssumptionBlock( + holds=' AND '.join(dict.fromkeys(link.values for link in pw.links)), + where=mask, + description=f"piecewise '{block}': every breakpoint the curve runs through needs a row in " f'{_quoted(link.values for link in pw.links)} — a missing row is read as a zero rather than as a ' f'shorter curve, so it sits the curve on the origin. ' + ( @@ -115,25 +98,23 @@ def assumptions_of(block: str, pw: PiecewiseBlock) -> dict[str, Assumed]: curvature = _curvature_required(pw) if curvature is not None: x, y = (link.values for link in pw.curve) - assumed[f'{block}_increasing'] = Assumed( - f'{_back(x, d, 1)} < {x}', - _neighbours(d, mask), - f"piecewise '{block}': method: {pw.method} requires strictly increasing breakpoints in '{x}' along '{d}'", + assumed[f'{block}_increasing'] = AssumptionBlock( + holds=f'{_back(x, d, 1)} < {x}', + where=_neighbours(d, mask), + description=f"piecewise '{block}': method: {pw.method} requires strictly increasing breakpoints in '{x}' along '{d}'", ) assumed[f'{block}_curvature'] = _bends(block, pw, x, y, curvature) if pw.method == 'lp': - assumed[f'{block}_breakpoints'] = Assumed( - f'count({mask or pw.curve[0].values}, over={d}) >= 2', - None, - f"piecewise '{block}': method: lp needs at least two breakpoints per curve — the method *is* its " + assumed[f'{block}_breakpoints'] = AssumptionBlock( + holds=f'count({mask or pw.curve[0].values}, over={d}) >= 2', + description=f"piecewise '{block}': method: lp needs at least two breakpoints per curve — the method *is* its " f'segment lines, so a curve with no segment states nothing and leaves the bounded link on its own ' f'bound. Use method: adjacency, sos2 or convex, which pin it to the points it does have.', ) if mask is not None: - assumed[f'{block}_contiguous'] = Assumed( - f'count({_edge(d, mask, "first")}, over={d}) == 1', - None, - f"piecewise '{block}': points: '{mask}' must mark a consecutive run of at least one breakpoint per " + assumed[f'{block}_contiguous'] = AssumptionBlock( + holds=f'count({_edge(d, mask, "first")}, over={d}) == 1', + description=f"piecewise '{block}': points: '{mask}' must mark a consecutive run of at least one breakpoint per " f'curve — {_GAP[pw.method]}.', ) return assumed @@ -188,7 +169,7 @@ def _interior(over: str, mask: str | None) -> str: return f'{mask} AND shift({mask}, along={over}, offset=1) AND shift({mask}, along={over}, offset=-1)' -def _bends(block: str, pw: PiecewiseBlock, x: str, y: str, curvature: Curvature) -> Assumed: +def _bends(block: str, pw: PiecewiseBlock, x: str, y: str, curvature: Curvature) -> AssumptionBlock: """The curve bends the way *curvature* says, as a comparison of the two slopes at each breakpoint. The slopes are compared as a cross-product rather than as two quotients, @@ -210,27 +191,122 @@ def _bends(block: str, pw: PiecewiseBlock, x: str, y: str, curvature: Curvature) ) if curvature == 'either': up, down = bend.format('>'), bend.format('<') - return Assumed( - f'count({up} AND {interior}, over={d}) == 0 OR count({down} AND {interior}, over={d}) == 0', - None, - description, + return AssumptionBlock( + holds=f'count({up} AND {interior}, over={d}) == 0 OR count({down} AND {interior}, over={d}) == 0', + description=description, ) - return Assumed(bend.format('<=' if curvature == 'convex' else '>='), interior, description) + return AssumptionBlock( + holds=bend.format('<=' if curvature == 'convex' else '>='), where=interior, description=description + ) -class _Block: - """One ``piecewise:`` block being expanded into the raw model it writes. +@dataclass(frozen=True) +class Emitted: + """Every name one block's expansion may write, spelled once for the emitter and the collision check. - Every name the expansion may write is spelled once here, so the emitters - and the collision check read the same table — ``set`` is the one a method - that states a set writes through :func:`math_spec.sos.emit`. ``mask`` is - the parameter masking the weights, or ``None`` for a whole curve: the - ``bool`` the file named, or one of the block's own values parameters, - which as a bare name in a ``where`` is true wherever it has a row. + The set a block states writes names of its own, and they are reserved + whichever method the block declares: which of the two write them is the + method's business, and a collision is the file's either way. + """ + + name: str + lam: str + convexity: str + set: sos.Emitted + chord: str + domain_lo: str + domain_hi: str + links: tuple[str, ...] + assumptions: tuple[str, ...] + + @classmethod + def of(cls, name: str, pw: PiecewiseBlock) -> Emitted: + """The names block *name* writes.""" + return cls( + name, + f'{name}_lam', + f'{name}_convexity', + sos.Emitted.of(name, 2), + f'{name}_chord', + f'{name}_domain_lo', + f'{name}_domain_hi', + tuple(f'{name}_link{i}' for i in range(len(pw.links))), + tuple(assumptions_of(name, pw)), + ) + + @property + def by_kind(self) -> tuple[tuple[str, tuple[str, ...]], ...]: + """Each name by the kind of declaration it would collide with.""" + return ( + ('variable', (self.lam, self.set.seg)), + ( + 'constraint', + ( + self.convexity, + self.convexity + _UNGATED, + self.set.pick, + self.set.link, + self.chord, + self.domain_lo, + self.domain_hi, + *self.links, + ), + ), + ('sos', (self.name,)), + ('assumption', self.assumptions), + ) + + +def curve_frame(schema: Spec, name: str, pw: PiecewiseBlock, links: Iterable[Expression]) -> tuple[str, ...]: + """The dimensions block *name* builds one curve per coordinate of: every one its links and its gate carry. + + In declaration order, because iterating a set would vary the emitted + ``dims`` — and every column index behind it — per process. *links* are the + block's link expressions typed, as + :attr:`~math_spec.resolution.Resolved.piecewise` holds them. Raises: - PiecewiseExpansionError: A block naming something that does not exist, - or emitting a name the file already declares. + DimensionError: A link or the gate carries the breakpoint dimension, or + a values or ``points:`` parameter varies along a dimension no link + expression carries. + """ + context = f"piecewise '{name}'" + carried = [(f'link {i} expression', dims_of(node, schema, f'{context} link {i}')) for i, node in enumerate(links)] + if pw.activity is not None: + carried.append(('activity', frozenset(schema.variables[pw.activity].dims))) + frame: list[str] = [] + for what, found in carried: + for d in (d for d in schema.dimensions if d in found): + if d == pw.over: + raise DimensionError(f"{context}: {what} already carries the breakpoint dim '{pw.over}'") + if d not in frame: + frame.append(d) + for i, link in enumerate(pw.links): + if stray := [d for d in schema.parameters[link.values].dims if d != pw.over and d not in frame]: + raise DimensionError( + f"{context}: link {i} values parameter '{link.values}' carries {stray}, which no link " + f'expression does — the block builds one curve per coordinate of {frame}, so a curve ' + f'varying along {stray} has nothing to vary against. Declare a link expression over ' + f"it, or drop it from '{link.values}'." + ) + if pw.points is not None and pw.nominated is None: + mask = schema.parameters[pw.points].dims + if stray := [d for d in mask if d != pw.over and d not in frame]: + raise DimensionError( + f"{context}: points parameter '{pw.points}' carries {stray}, which the links do not — " + f"a mask says which of the block's own coordinates exist, and cannot add coordinates" + ) + return tuple(frame) + + +class _Block: + """One ``piecewise:`` block being expanded into the raw model it writes. + + ``mask`` is the parameter masking the weights, or ``None`` for a whole + curve: the ``bool`` the file named, or one of the block's own values + parameters, which as a bare name in a ``where`` is true wherever it has a + row. Nothing here can fail: every rule a block is held to was decided when + *schema* loaded. """ def __init__(self, schema: Spec, raw: dict[str, object], name: str, pw: PiecewiseBlock) -> None: @@ -238,17 +314,9 @@ def __init__(self, schema: Spec, raw: dict[str, object], name: str, pw: Piecewis self.raw = raw self.name = name self.pw = pw - self.lam = f'{name}_lam' - self.convexity = f'{name}_convexity' - self.set = Emitted.of(name, 2) - self.chord = f'{name}_chord' - self.domain_lo = f'{name}_domain_lo' - self.domain_hi = f'{name}_domain_hi' - self.links = tuple(f'{name}_link{i}' for i in range(len(pw.links))) + self.emitted = Emitted.of(name, pw) self.mask = pw.points - self.ns = Namespace(schema) - self.context = f"piecewise '{name}'" - self.frame = self._validated_frame() + self.frame = curve_frame(schema, name, pw, schema.resolved.piecewise[name]) def expand(self) -> None: """Write the block's declarations into the raw model.""" @@ -265,25 +333,22 @@ def _assumptions(self) -> None: model that has been written out carries them as language rather than as something a consumer has to know to ask for. """ - assumptions = section(self.raw, 'assumptions') + assumptions = sos.section(self.raw, 'assumptions') for name, assumed in assumptions_of(self.name, self.pw).items(): - entry: dict[str, object] = {'holds': assumed.holds, 'description': assumed.description} - if assumed.where is not None: - entry['where'] = assumed.where - assumptions[name] = entry + assumptions[name] = assumed.model_dump() # -- emitters ---------------------------------------------------------- def _weight(self, name: str, **fields: object) -> None: """A variable over the frame and the breakpoint dim, masked as the block is.""" - section(self.raw, 'variables')[name] = { + sos.section(self.raw, 'variables')[name] = { 'dims': [*self.frame, self.pw.over], **({'where': self.mask} if self.mask else {}), **fields, } def _constraint(self, name: str, dims: list[str], expression: str, where: str | None = None) -> None: - section(self.raw, 'constraints')[name] = { + sos.section(self.raw, 'constraints')[name] = { 'dims': dims, **({'where': where} if where else {}), 'expression': expression, @@ -293,21 +358,23 @@ def _weights(self) -> None: """The convex-combination form: weights, their convexity, a row per link, and the method's restriction.""" d = self.pw.over self._weight( - self.lam, + self.emitted.lam, bounds={'lower': 0.0, 'upper': 1.0}, description='convex-combination weight on a breakpoint', ) gated = self._gate_rows() for suffix, where, rhs in gated: - self._constraint(self.convexity + suffix, list(self.frame), f'sum({self.lam}, over={d}) == {rhs}', where) - for cname, link in zip(self.links, self.pw.links, strict=True): + self._constraint( + self.emitted.convexity + suffix, list(self.frame), f'sum({self.emitted.lam}, over={d}) == {rhs}', where + ) + for cname, link in zip(self.emitted.links, self.pw.links, strict=True): self._constraint( cname, list(self.frame), - f'({link.expression}) {link.sign} sum({self.lam} * {link.values}, over={d})', + f'({link.expression}) {link.sign} sum({self.emitted.lam} * {link.values}, over={d})', ) if self.pw.method in ('sos2', 'adjacency'): - section(self.raw, 'sos')[self.name] = {'variable': self.lam, 'over': d, 'type': 2} + sos.section(self.raw, 'sos')[self.name] = {'variable': self.emitted.lam, 'over': d, 'type': 2} def _gate_rows(self) -> tuple[tuple[str, str | None, str], ...]: """What the weights sum to, as ``(name suffix, where, right-hand side)``. @@ -348,177 +415,18 @@ def _segment_lines(self) -> None: run = f'({x_link.values} - shift({x_link.values}, along={d}, offset=1, edge=0))' rise = f'({y_link.values} - shift({y_link.values}, along={d}, offset=1, edge=0))' self._constraint( - self.chord, + self.emitted.chord, [*self.frame, d], f'({y_link.expression}) * {run} {y_link.sign} ' f'{rise} * (({x_link.expression}) - {x_link.values}) + {y_link.values} * {run}', _neighbours(d, self.mask), ) - edges = ((self.domain_lo, '>=', 'first'), (self.domain_hi, '<=', 'last')) + edges = ((self.emitted.domain_lo, '>=', 'first'), (self.emitted.domain_hi, '<=', 'last')) for cname, sense, end in edges: self._constraint( cname, [*self.frame, d], f'({x_link.expression}) {sense} {x_link.values}', _edge(d, self.mask, end) ) - # -- checks ------------------------------------------------------------ - - def _emitted_by_kind(self) -> tuple[tuple[str, tuple[str, ...]], ...]: - """Every name this block may write, by the kind of declaration each would collide with. - - The set a block states writes names of its own, and they are reserved - whichever method the block declares: which of the two write them is the - method's business, and a collision is the file's either way. - """ - return ( - ('variable', (self.lam, self.set.seg)), - ( - 'constraint', - ( - self.convexity, - self.convexity + _UNGATED, - self.set.pick, - self.set.link, - self.chord, - self.domain_lo, - self.domain_hi, - *self.links, - ), - ), - ('sos', (self.name,)), - ('assumption', tuple(assumptions_of(self.name, self.pw))), - ) - - def _validated_frame(self) -> tuple[str, ...]: - """Check every name the block writes and infer its frame: the union of the links' and the gate's dims. - - A values parameter is checked against the frame in a second pass, since - the last link's expression widens the frame as readily as the first; left - to the emitted declarations the refusal would name ``_link0``, a - constraint the author never wrote. - """ - if self.pw.over not in self.schema.dimensions: - raise PiecewiseExpansionError(undeclared_dimension('piecewise', self.name, self.pw.over)) - frame: list[str] = [] - self._widen(frame, self._link_dims()) - self._widen(frame, self._activity_dims()) - self._values_fit(frame) - self._points_fit(frame) - self._nothing_collides() - return tuple(frame) - - def _widen(self, frame: list[str], dims: Iterable[tuple[str, frozenset[str]]]) -> None: - """Add each labelled dim set to *frame* in declaration order, refusing the breakpoint dim. - - Declaration order, because iterating a set would vary the emitted - ``dims`` — and every column index behind it — per process. - """ - for what, found in dims: - for d in (d for d in self.schema.dimensions if d in found): - if d == self.pw.over: - raise PiecewiseExpansionError( - f"{self.context}: {what} already carries the breakpoint dim '{self.pw.over}'" - ) - if d not in frame: - frame.append(d) - - def _link_dims(self) -> Iterator[tuple[str, frozenset[str]]]: - """Each link's expression dims, its values parameter checked to exist and to run along the breakpoint dim.""" - schema, pw = self.schema, self.pw - for i, link in enumerate(pw.links): - values = link.values - if values not in schema.parameters: - raise PiecewiseExpansionError( - f"{self.context}: link {i} values references undeclared parameter '{values}'" - ) - if pw.over not in schema.parameters[values].dims: - raise PiecewiseExpansionError( - f"{self.context}: link {i} values parameter '{values}' must carry dim " - f"'{pw.over}' (has {schema.parameters[values].dims})" - ) - yield f'link {i} expression', self._expr_dims(link.expression, f'{self.context} link {i}') - - def _activity_dims(self) -> Iterator[tuple[str, frozenset[str]]]: - """The gate's dims, if the block names one: a declared binary variable.""" - activity = self.pw.activity - if activity is None: - return - if activity not in self.schema.variables: - raise PiecewiseExpansionError( - f"{self.context}: activity '{activity}' is not a declared variable. A gate is a binary variable; " - f'declare it, or drop activity: for weights that sum to 1.' - ) - if self.schema.variables[activity].domain != 'binary': - raise PiecewiseExpansionError(f"{self.context}: activity variable '{activity}' must be binary") - yield 'activity', self._expr_dims(activity, f'{self.context} activity') - - def _values_fit(self, frame: list[str]) -> None: - """A values parameter varies along the frame and the breakpoint dim, and nothing else.""" - for i, link in enumerate(self.pw.links): - if stray := [d for d in self.schema.parameters[link.values].dims if d != self.pw.over and d not in frame]: - raise PiecewiseExpansionError( - f"{self.context}: link {i} values parameter '{link.values}' carries {stray}, which no link " - f'expression does — the block builds one curve per coordinate of {frame}, so a curve ' - f'varying along {stray} has nothing to vary against. Declare a link expression over ' - f"it, or drop it from '{link.values}'." - ) - - def _points_fit(self, frame: list[str]) -> None: - """A ``points:`` naming a parameter of its own is a bool mask along the breakpoint dim, inside the frame.""" - pw, ctx = self.pw, self.context - if pw.points is None or _nominated(pw) is not None: - return - if pw.points not in self.schema.parameters: - raise PiecewiseExpansionError(f"{ctx}: points references undeclared parameter '{pw.points}'") - if (dtype := self.schema.parameters[pw.points].dtype) != 'bool': - raise PiecewiseExpansionError( - f"{ctx}: points parameter '{pw.points}' is {dtype}, and a mask is a bool parameter — one " - f'saying, per breakpoint, whether the curve reaches it. Declare it dtype: bool.' - ) - mask = self.schema.parameters[pw.points].dims - if pw.over not in mask: - raise PiecewiseExpansionError( - f"{ctx}: points parameter '{pw.points}' must carry dim '{pw.over}' — " - f'it says how far each curve runs along it (has {mask})' - ) - if stray := [d for d in mask if d != pw.over and d not in frame]: - raise PiecewiseExpansionError( - f"{ctx}: points parameter '{pw.points}' carries {stray}, which the links do not — " - f"a mask says which of the block's own coordinates exist, and cannot add coordinates" - ) - - def _nothing_collides(self) -> None: - """No name the block writes is one the file already declares.""" - declared = { - 'variable': self.schema.variables, - 'constraint': self.schema.constraints, - 'sos': self.schema.sos, - 'assumption': self.schema.assumptions, - } - for kind, names in self._emitted_by_kind(): - for one in names: - if one in declared[kind]: - raise PiecewiseExpansionError( - f"{self.context}: emitted {kind} '{one}' collides with a declared {kind}" - ) - - def _expr_dims(self, text: str, ctx: str) -> frozenset[str]: - """Dims of an affine link expression, asked of ``dimensions`` before any declaration exists to carry it.""" - ast = parse_and_expand(text, self.ns, ctx) - if isinstance(ast, ComparisonNode): - raise PiecewiseExpansionError(f'{ctx}: link expressions must not contain a comparison, got {text!r}') - errors: list[str] = [] - resolved = resolve_expression(ast, self.ns, ctx, errors) - if resolved is None: - raise PiecewiseExpansionError('\n'.join(errors)) - assert not isinstance(resolved, ComparisonNode) - try: - check_expression(resolved, ctx) - return dims_of(resolved, self.schema, ctx) - except LanguageError as exc: - raise PiecewiseExpansionError( - f'{ctx}: link expression {text!r} is not a valid affine expression: {exc}' - ) from exc - def expand_piecewise(schema: Spec) -> Spec: """*schema* with every ``piecewise:`` block written out — *schema* itself where it declares none. @@ -527,10 +435,6 @@ def expand_piecewise(schema: Spec) -> Spec: ``method: sos2`` states, and then that set is written out here too: the binaries are what the method *is*, so the model that comes back carries no set of its own (:func:`math_spec.sos.emit` is where they are spelled). - - Raises: - PiecewiseExpansionError: A block naming something that does not exist, - or emitting a name the file already declares. """ if not schema.piecewise: return schema @@ -543,7 +447,7 @@ def expand_piecewise(schema: Spec) -> Spec: raw['piecewise'].clear() for name, pw in schema.piecewise.items(): if pw.method == 'adjacency': - emit(raw, name) + sos.emit(raw, name) expanded = Spec.model_validate(raw) expanded._expanded_piecewise = dict(schema.piecewise) return expanded diff --git a/src/math_spec/program.py b/src/math_spec/program.py index a1087608..4521e973 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -61,9 +61,9 @@ 'FanIn', 'Footprint', 'GroupSum', - 'Holds', 'Mask', 'Multiply', + 'Named', 'Negate', 'Not', 'ObjectiveDeclaration', @@ -347,6 +347,22 @@ class Cases: regions: tuple[Region, ...] +@dataclass(frozen=True) +class Named: + """A use of an ``expressions:`` entry, standing where its name was written, with the entry's body under it. + + Only a :attr:`~math_spec.model.Spec.resolved` tree holds one: it is what + lets the typesetter print the symbol where the name stood and define it + once, and what ``in_math`` is read off. Lowering inlines every one, so no + :class:`Program` carries it and :data:`Expression` does not name it. Every + use of one entry holds the one node resolution built for it, and a walk + steps through it. + """ + + name: str + body: Expression + + #: Every expression node, as one type — what a walk takes. The set is #: *closed*: nothing registers into it, so a consumer that walks it ends in #: ``assert_never`` and a node added without a branch is a type error at the @@ -396,6 +412,8 @@ def fan_in(expression: Expression) -> FanIn: def children(expression: Expression) -> tuple[Expression, ...]: """The sub-expressions of *expression* — what every walk recurses through.""" + if isinstance(expression, Named): + return (expression.body,) if isinstance(expression, Negate): return (expression.operand,) if isinstance(expression, (Add, Multiply)): @@ -543,7 +561,7 @@ class PiecewiseDeclaration: The expansion lowered the links into constraints over the file's own parameters, and emitted none. What the block assumes of its numbers is an - :data:`Assumption` like any other, under :attr:`Program.assumptions`; what + :class:`Assumption` like any other, under :attr:`Program.assumptions`; what is left here is the curve. Attributes: @@ -558,7 +576,7 @@ class PiecewiseDeclaration: @dataclass(frozen=True) -class Holds: +class Assumption: """A predicate the file states of its data, under the name it wrote in ``assumptions:``. ``predicate`` is true at every coordinate of its frame — the product of @@ -576,14 +594,6 @@ class Holds: description: str | None = None -#: One fact about the data a consumer has to check before it solves — the -#: file's own, and every one a ``piecewise:`` method implies, which the -#: expansion writes into ``assumptions:`` and a load derives for a block still -#: declared. The data decides whether each holds, so the language states the -#: condition and the consumer holding the numbers checks. -Assumption = Holds - - def assumption_message(name: str, assumption: Assumption) -> str: """The sentence a consumer raises when the data bound to *assumption*, called *name*, fails it. @@ -1053,22 +1063,18 @@ class ParameterComparison: @dataclass(frozen=True) -class ExpressionComparison[Side]: +class ExpressionComparison: """Compare two variable-free expressions, coordinate by coordinate — ``p_min <= 0.5 * p_max``. ``dims`` is every dim either side carries. A side whose value is absent at a coordinate — a parameter row missing, a translation that vacated it — makes the comparison false there, as a null does in every other comparison; under a summing operator the absent term is one fewer. - - In a program each side is an :data:`Expression`. Before lowering, the - readers of the file — the typesetter, the dim rules, the exclusivity - check — see the same node with its sides in the core syntax tree. """ - left: Side + left: Expression op: PredicateOperator - right: Side + right: Expression dims: tuple[str, ...] @@ -1212,7 +1218,6 @@ class Or: #: decide about them. TypedPredicate = ( ParameterComparison - # pyrefly: ignore[implicit-any-type-argument] # the union is an isinstance target, which takes no parameterized class | ExpressionComparison | ParameterDefined | VariableDefined diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index fea7482e..dcf68ee6 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -2,11 +2,12 @@ # # SPDX-License-Identifier: MIT -"""Name resolution — the pass that makes the core AST fully typed. +"""Name resolution — the pass that reads the syntax tree into the program vocabulary. -Parsers emit unresolved names; this module rewrites each into the typed node -its kind asks for, so the AST reaching a consumer holds none. The rules live in -the language reference. +The grammars emit bare names and calls; this module builds the +:mod:`math_spec.program` node each stands for, so every pass after — the dim +rules, the degree rules, the typesetter, lowering — reads one vocabulary. The +rules live in the language reference. """ from __future__ import annotations @@ -21,29 +22,15 @@ from math_spec._expression_parser import ( ArithmeticNode, BinaryOperatorNode, - CaseArm, - CasesNode, ComparisonNode, - DefinitionNode, - DimensionNode, - DirectionNode, - DualNode, - EdgeNode, FunctionCallNode, KeywordNode, - KwargNode, NameListNode, NameNode, NumberNode, - ParameterNode, - ParsedNode, - PartitionNode, UnaryOperatorNode, - VariableNode, - case_context, nodes, shown, - with_children, ) from math_spec._where_parser import ( ColumnNode, @@ -54,11 +41,12 @@ parse_where, ) from math_spec.dimensions import dims_of, pulled_back_dims -from math_spec.errors import DimensionError, LanguageError, SchemaError, did_you_mean, prefixed +from math_spec.errors import DimensionError, LanguageError, SchemaError, case_context, did_you_mean, prefixed from math_spec.exclusivity import overlapping from math_spec.expansion import expand, parse_and_expand from math_spec.model import NUMERIC_DTYPES from math_spec.operators import ( + AMOUNTS, BUILTINS, EDGE_WRAP, PARTITION_NAMES_ITS_GROUP, @@ -67,34 +55,58 @@ unknown_operator_message, ) from math_spec.program import ( + Add, And, + Assumption, BooleanLiteral, + Cases, + Constant, + ConstraintDeclaration, CountComparison, DimensionComparison, DimensionPosition, Direction, + Divide, + Dual, + Expression, ExpressionComparison, + GroupSum, Mask, + Multiply, + Named, + Negate, Not, + ObjectiveDeclaration, Or, + Parameter, ParameterComparison, ParameterDefined, Partition, + Power, Predicate, PredicateOperator, + Pullback, PulledBackPredicate, + Region, RelationComparison, RelationDeclaration, RelationDefined, RelationPairComparison, + Sum, + Translate, TranslatedPredicate, TypedPredicate, + Variable, VariableDefined, + WindowSum, + carries_variable, + walk, ) if TYPE_CHECKING: from collections.abc import Iterable, Mapping + from math_spec._expression_parser import ComparisonOperator from math_spec.model import DeclaredDtype, ExpressionBlock, Spec @@ -103,6 +115,10 @@ #: than over the stores it would otherwise have to try in order. DeclarationKind = Literal['variable', 'parameter', 'dimension', 'relation'] +#: An ``edge=`` as a translation carries it: whether it wraps, and the number +#: the vacated positions contribute where it does not. +_Edge = tuple[bool, float | None] + class Namespace: """The declared names of one schema, by kind — the whole of what a file may name, read once. @@ -154,11 +170,11 @@ def __init__(self, schema: Spec) -> None: } #: named expression -> its resolved node, or ``None``, and its refusals; #: filled the first time anything reads the name. - self._named: dict[str, tuple[CasesNode | DefinitionNode | None, tuple[str, ...]]] = {} + self._named: dict[str, tuple[Named | None, tuple[str, ...]]] = {} #: The named expressions being resolved, outermost first — a cycle's chain. self._loading: list[str] = [] - def named(self, name: str, context: str) -> CasesNode | DefinitionNode: + def named(self, name: str, context: str) -> Named: """The ``expressions:`` entry *name* as the node that stands where its name is written. Resolved under the entry's own context the first time it is asked @@ -177,7 +193,7 @@ def named(self, name: str, context: str) -> CasesNode | DefinitionNode: raise SchemaError(msg) return node - def named_entry(self, name: str) -> tuple[CasesNode | DefinitionNode | None, tuple[str, ...]]: + def named_entry(self, name: str) -> tuple[Named | None, tuple[str, ...]]: """The ``expressions:`` entry *name* resolved, or ``None``, with every refusal it earned.""" if name not in self._named: errors: list[str] = [] @@ -232,47 +248,27 @@ def unknown_constraint(self, name: str, context: str, *, formals: Iterable[str] ) -class ResolvedConstraint(NamedTuple): - """One constraint's typed halves: the comparison it states, and the mask it holds under.""" - - expression: ComparisonNode - where: Mask | None - - -class ResolvedAssumption(NamedTuple): - """One assumption's typed halves: the predicate it states, and the mask it is checked under. - - ``description`` is the sentence a refusal quotes where one was written or - a method implied one, and ``None`` where the name is the whole of what a - reader is told. - """ - - holds: Mask - where: Mask | None - description: str | None = None - - @dataclass(frozen=True) class Resolved: - """Every expression and where string of one schema, typed once at load. + """Every expression and where string of one schema, typed once at load, in the program's own vocabulary. :func:`~math_spec.validation.validate_expressions` builds it, and every reader after — the dim rules, lowering, the typesetter — walks these trees rather than parsing, expanding and resolving the text again. Each mapping is keyed as the schema's own section is. A ``where`` the file did not - write, or one every row passes, is ``None``. + write, or one every row passes, is ``None``. What a program does not carry + is here alone: every use of an ``expressions:`` entry stands as the + :class:`~math_spec.program.Named` node resolution built for it, which + lowering inlines. Attributes: - expressions: Each ``expressions:`` entry as the node its name expands - to — a plain entry a :class:`~math_spec._expression_parser.DefinitionNode` - carrying its name over its body, a cased one a - :class:`~math_spec._expression_parser.CasesNode` with every arm's - ``when`` typed. Every entry either names is inlined where it - stood, so a walk over one sees the whole chain. + expressions: Each ``expressions:`` entry as the node every use of it + holds — a plain entry's body, or a cased one's + :class:`~math_spec.program.Cases` with every region's mask typed + and the ``otherwise`` carrying the negation of the rest. variables: Each variable's ``where``. - constraints: Each constraint's comparison and ``where``. - objective: The objective's expression, ``None`` where the file - declares none. + constraints: Each constraint, as a program declares it. + objective: The objective, ``None`` where the file declares none. relations: Each relation's columns and key, as declared — the one copy, which every :class:`~math_spec.program.Direction` and :class:`~math_spec.program.Partition` in the trees holds. @@ -281,13 +277,13 @@ class Resolved: piecewise: Each ``piecewise:`` block's link expressions, in link order. """ - expressions: dict[str, CasesNode | DefinitionNode] + expressions: dict[str, Named] variables: dict[str, Mask | None] - constraints: dict[str, ResolvedConstraint] - objective: ArithmeticNode | None + constraints: dict[str, ConstraintDeclaration] + objective: ObjectiveDeclaration | None relations: dict[str, RelationDeclaration] - assumptions: dict[str, ResolvedAssumption] - piecewise: dict[str, tuple[ArithmeticNode, ...]] + assumptions: dict[str, Assumption] + piecewise: dict[str, tuple[Expression, ...]] @cached_property def read_by_the_math(self) -> frozenset[str]: @@ -300,11 +296,11 @@ def read_by_the_math(self) -> frozenset[str]: counts because it states rows, so the answer does not move when the curve is written out (:meth:`~math_spec.model.Spec.expand`). """ - roots: list[ParsedNode] = [constraint.expression for constraint in self.constraints.values()] + roots = [side for constraint in self.constraints.values() for side in (constraint.lhs, constraint.rhs)] if self.objective is not None: - roots.append(self.objective) + roots.append(self.objective.expression) roots.extend(link for links in self.piecewise.values() for link in links) - return frozenset(node.name for node in nodes(*roots) if isinstance(node, CasesNode | DefinitionNode)) + return frozenset(node.name for node in walk(*roots) if isinstance(node, Named)) # --------------------------------------------------------------------------- @@ -326,31 +322,44 @@ def mask_of(node: Predicate | None) -> Mask | None: return Mask(node) +def remainder(masks: Iterable[Mask]) -> Mask: + """The region left over: where not one of *masks* holds. + + The ``otherwise`` arm's own mask, built rather than written. ``cases:`` + carries at least one case, so there is no vacuous truth to spell. + """ + first, *rest = masks + left = ~first + for mask in rest: + left = left & ~mask + return left + + # --------------------------------------------------------------------------- # expressions # --------------------------------------------------------------------------- def resolve_expression( - node: ParsedNode, + node: ArithmeticNode, ns: Namespace, context: str, errors: list[str], *, formals: frozenset[str] = frozenset(), -) -> ParsedNode | None: - """Rewrite every ``NameNode`` under *node* to a typed node, checking operator call shapes on the way. - - A name in *formals* stays bare, so a macro template is checked by the - rules a call site is, before anything calls it. +) -> Expression | None: + """Build the program tree *node* stands for, checking every name and operator call shape on the way. Returns: - The typed tree, or ``None`` once anything failed — appending to - *errors* rather than raising, so a caller collecting problems across a - whole schema reports them together. + The tree, or ``None`` once anything failed — appending to *errors* + rather than raising, so a caller collecting problems across a whole + schema reports them together. Also ``None``, with nothing appended, + where a name in *formals* stands under *node*: a macro template is + checked by the rules a call site is before anything calls it, and + only the call site that binds its formals has a tree to build. """ before = len(errors) - resolved = _Resolver(ns, context, errors, formals=formals).expression(node) + resolved = _Resolver(ns, context, errors, formals=formals).arith(node) return None if len(errors) > before else resolved @@ -395,20 +404,106 @@ def resolve_where_text( return resolve_where(node, ns, context, errors, self_variable) -def _named(name: str, block: ExpressionBlock, ns: Namespace, errors: list[str]) -> CasesNode | DefinitionNode | None: - """One ``expressions:`` entry as the node its name expands to, or ``None`` once anything in it failed. +def resolve_expression_text( + text: str, ns: Namespace, context: str, errors: list[str], *, ceiling: int | None +) -> Expression | None: + """Parse, expand, resolve and degree-check one expression string that stands for a value. + + *ceiling* is the degree the position honours, and ``None`` for an + ``expressions:`` entry's body: what the math admits + (:func:`~math_spec.degree.check_expression`) is a rule about the position + that *reads* it, so it fires on the expanded tree of every objective and + piecewise link, and not where an entry is declared. A constraint is + :func:`resolve_constraint_text`'s. + + Returns: + The typed tree, or ``None`` once anything failed, the problem appended + to *errors*. + """ + ast = _parsed(text, ns, context, errors) + if ast is None: + return None + if isinstance(ast, ComparisonNode): + errors.append(f'{context}: expression must not contain a comparison operator.\nGot: {text!r}') + return None + resolved = resolve_expression(ast, ns, context, errors) + if resolved is None or ceiling is None: + return resolved + return None if _over_the_ceiling(resolved, context, errors, ceiling=ceiling) else resolved + + +def resolve_constraint_text( + text: str, ns: Namespace, context: str, errors: list[str] +) -> tuple[Expression, ComparisonOperator, Expression] | None: + """Parse, expand, resolve and degree-check one constraint string: exactly one comparison, a variable on a side (#1171). + + Returns: + The two sides and the sense between them, or ``None`` once anything + failed, the problem appended to *errors*. + """ + ast = _parsed(text, ns, context, errors) + if ast is None: + return None + if not isinstance(ast, ComparisonNode): + errors.append( + f'{context}: expression must contain exactly one comparison operator (<=, >=, ==).\nGot: {text!r}' + ) + return None + found = len(errors) + resolver = _Resolver(ns, context, errors) + left, right = resolver.arith(ast.left), resolver.arith(ast.right) + if len(errors) > found or left is None or right is None: + return None + if any(_over_the_ceiling(side, context, errors, ceiling=2) for side in (left, right)): + return None + if not (carries_variable(left) or carries_variable(right)): + errors.append( + f'{context}: neither side of the comparison carries a variable, so the row decides nothing.\n' + f'Got: {text!r}\n' + f'A constraint is a claim about a decision, and a comparison of numbers and parameters ' + f'is settled before the solve — no consumer builds a row for it. Name the variable it should ' + f'bound, or state the fact under `assumptions:`, where the consumer binding the data checks it.' + ) + return None + return left, ast.op, right + + +def _parsed(text: str, ns: Namespace, context: str, errors: list[str]) -> ComparisonNode | ArithmeticNode | None: + """*text* parsed and its macros expanded, or ``None`` with the refusal appended.""" + try: + return parse_and_expand(text, ns, context) + except ValueError as e: + errors.append(prefixed(context, e)) + return None + + +def _over_the_ceiling(node: Expression, context: str, errors: list[str], *, ceiling: int) -> bool: + """Whether *node* breaks the degree rules at *ceiling*, the refusal appended.""" + try: + degree.check_expression(node, context, ceiling=ceiling) + except LanguageError as e: + errors.append(str(e)) + return True + return False + + +def _named(name: str, block: ExpressionBlock, ns: Namespace, errors: list[str]) -> Named | None: + """One ``expressions:`` entry as the node every use of it holds, or ``None`` once anything in it failed. A cased entry's arms are checked one by one, so every fault is collected rather than the first, and proved apart only once all of them resolve. + The ``otherwise`` arm becomes the region left over, so a consumer adds + regions rather than working out which one is left; the language proved + the rest apart, so the regions are disjoint and total. """ context = f"Named expression '{name}'" if not block.cases: assert block.expression is not None - body = _value(block.expression, ns, context, errors) - return None if body is None else DefinitionNode(name, body) + body = resolve_expression_text(block.expression, ns, context, errors, ceiling=None) + return None if body is None else Named(name, body) found = len(errors) - arms: list[CaseArm] = [] + regions: list[Region] = [] masks: dict[str, Predicate] = {} for case_name, case in block.cases.items(): arm_context = case_context(name, case_name) @@ -417,30 +512,18 @@ def _named(name: str, block: ExpressionBlock, ns: Namespace, errors: list[str]) errors.append(_constant_arm(arm_context, value=when.value)) elif when is not None: masks[case_name] = when - value = _value(case.expression, ns, arm_context, errors) + value = resolve_expression_text(case.expression, ns, arm_context, errors, ceiling=None) if when is not None and value is not None: - arms.append(CaseArm(case_name, when, value)) + regions.append(Region(Mask(when), value)) assert block.otherwise is not None - fallback = _value(block.otherwise, ns, case_context(name, None), errors) + fallback = resolve_expression_text(block.otherwise, ns, case_context(name, None), errors, ceiling=None) if len(errors) > found or fallback is None: return None errors.extend(f'{context}: {problem}' for problem in overlapping(masks, ns.dtypes)) - return CasesNode(name, (*arms, CaseArm('otherwise', None, fallback))) - - -def _value(text: str, ns: Namespace, context: str, errors: list[str]) -> ArithmeticNode | None: - """One expression string that stands for a value, typed; ``None`` once anything in it failed.""" - try: - ast = parse_and_expand(text, ns, context) - except ValueError as e: - errors.append(prefixed(context, e)) - return None - if isinstance(ast, ComparisonNode): - errors.append(f'{context}: expression must not contain a comparison operator.\nGot: {text!r}') + if len(errors) > found: return None - resolved = resolve_expression(ast, ns, context, errors) - assert not isinstance(resolved, ComparisonNode), 'an arithmetic tree resolves to arithmetic' - return resolved + left_over = Region(remainder(region.when for region in regions), fallback) + return Named(name, Cases((*regions, left_over))) def _constant_arm(context: str, *, value: bool) -> str: @@ -464,12 +547,12 @@ def _constant_arm(context: str, *, value: bool) -> str: class _Resolver: """One resolution walk, and the three things every step of it reads. - A node that cannot be typed comes back unresolved with its refusal - appended to ``errors``; the public doors discard the tree once ``errors`` - grew, which is what lets a connective's children be typed as resolved. - ``self_variable`` is the variable whose own ``where`` is being read, which - may not ask whether it exists. ``formals`` are a macro template's formals, - which stay bare: a formal has no kind until a call site binds it. + A node that cannot be built comes back as ``None`` with its refusal + appended to ``errors``; every sibling is still read, so a declaration + with two faults reports both. ``self_variable`` is the variable whose own + ``where`` is being read, which may not ask whether it exists. ``formals`` + are a macro template's formals: a formal has no kind until a call site + binds it, so a node one stands under is ``None`` with nothing appended. """ ns: Namespace @@ -484,31 +567,23 @@ def _formal(self, value: ArithmeticNode) -> bool: # -- expressions ------------------------------------------------------- - def expression(self, node: ParsedNode) -> ParsedNode: - """Every ``NameNode`` under *node* typed; a comparison keeps its shape.""" - if isinstance(node, ComparisonNode): - return ComparisonNode(node.op, self._arith(node.left), self._arith(node.right)) - return self._arith(node) + def arith(self, node: ArithmeticNode) -> Expression | None: + """The program node *node* stands for, or ``None``. - def _arith(self, node: ArithmeticNode, *, amount: bool = False) -> ArithmeticNode: - """One arithmetic node typed. - - *amount* marks an ``offset=``/``window=`` value, whose dtype rule is - ``dimensions._check_named_amount``'s and stricter than "a number", so the - numeric check here stands aside for it. A quoted keyword or a name list in - arithmetic arrives through a macro formal bound to one. A named - expression arrives resolved, from :meth:`Namespace.named`, and passes. + A quoted keyword or a name list in arithmetic arrives through a macro + formal bound to one. """ - if isinstance( - node, NumberNode | VariableNode | ParameterNode | DualNode | KwargNode | CasesNode | DefinitionNode - ): - return node - if self._formal(node): - return node + if isinstance(node, NumberNode): + return Constant(node.value) if isinstance(node, NameNode): - return self._name(node, amount=amount) - if isinstance(node, UnaryOperatorNode | BinaryOperatorNode): - return with_children(node, self._arith) + return self._name(node) + if isinstance(node, UnaryOperatorNode): + operand = self.arith(node.operand) + if operand is None: + return None + return Negate(operand) if node.op == '-' else operand + if isinstance(node, BinaryOperatorNode): + return self._binary(node) if isinstance(node, FunctionCallNode): return self._call(node) if isinstance(node, KeywordNode): @@ -517,27 +592,60 @@ def _arith(self, node: ArithmeticNode, *, amount: bool = False) -> ArithmeticNod f"operator kwarg value such as shift(..., edge='wrap'). In an expression, quote " f'nothing — names resolve and numbers are written bare.' ) - return node + return None if isinstance(node, NameListNode): self.errors.append( f'{self.context}: {node} is a list of names, which is only legal as an operator ' f'kwarg value such as sum(x, by=[gen_bus, gen_tech]). In an expression, write the ' f'terms out and add them.' ) - return node + return None assert_never(node) - def _name(self, node: NameNode, *, amount: bool) -> ArithmeticNode: - """A bare name as the variable or parameter it declares; a dimension or relation is not a value.""" + def _binary(self, node: BinaryOperatorNode) -> Expression | None: + """A subtraction is an addition of the negation, so a program has one additive node.""" + left, right = self.arith(node.left), self.arith(node.right) + if left is None or right is None: + return None + match node.op: + case '+': + return Add(left, right) + case '-': + return Add(left, Negate(right)) + case '*': + return Multiply(left, right) + case '/': + return Divide(left, right) + case '**': + return Power(left, right) + case _: + assert_never(node.op) + + def _name(self, node: NameNode) -> Expression | None: + """A bare name as the variable, parameter or named expression it declares; a dimension or relation is not a value. + + A named expression arrives as the one node :meth:`Namespace.named` + built for it; the cast is the one place a + :class:`~math_spec.program.Named` enters a tree typed as a program's, + which lowering makes true. + """ + if node.name in self.formals: + return None + if node.name in self.ns.schema.expressions: + try: + return cast('Expression', self.ns.named(node.name, self.context)) + except SchemaError as e: + self.errors.append(str(e)) + return None match self.ns.kind(node.name): case 'variable': - return VariableNode(node.name) + return Variable(node.name) case 'parameter': dtype = self.ns.dtypes.get(node.name) - if not amount and dtype is not None and dtype not in NUMERIC_DTYPES: + if dtype is not None and dtype not in NUMERIC_DTYPES: self.errors.append(_not_a_number(node.name, dtype, self.context)) - return node - return ParameterNode(node.name) + return None + return Parameter(node.name) case 'dimension': self.errors.append( f"{self.context}: '{node.name}' is a dimension, and a dimension is " @@ -546,7 +654,7 @@ def _name(self, node: NameNode, *, amount: bool) -> ArithmeticNode: f'and in where-comparisons — to use its coordinates as data, ' f'declare a parameter over it.' ) - return node + return None case 'relation': self.errors.append( f"{self.context}: '{node.name}' is a relation, and a relation is structure " @@ -554,24 +662,28 @@ def _name(self, node: NameNode, *, amount: bool) -> ArithmeticNode: f'appears in a helper (sum(x, by={node.name})) and in a where — to ' f'carry numbers along this dimension, declare a parameter over it.' ) - return node + return None case _: self.errors.append(self.ns.unknown(node.name, self.context, allow_dims=False, formals=self.formals)) - return node + return None - def _call(self, node: FunctionCallNode) -> ArithmeticNode: - """An operator call: its shape checked, and each kwarg typed by the kind the operator declares for it.""" + def _call(self, node: FunctionCallNode) -> Expression | None: + """An operator call as the node it is: its shape checked, and each kwarg read by the kind the operator declares for it. + + Every argument is read even after one failed, so a call with two + faults reports both. A formal anywhere under the call builds nothing + and refuses nothing. + """ if node.name not in BUILTINS: self.errors.append(f'{self.context}: {unknown_operator_message(node.name)}') - return node + return None builtin = BUILTINS[node.name] shape_error = call_shape_error(node.name, len(node.args), node.kwargs) if shape_error is not None: self.errors.append(f'{self.context}: {shape_error}') if node.name == 'dual': - return node if shape_error is not None else self._dual(node) - args = tuple(self._arith(a) for a in node.args) - kwargs: dict[str, ArithmeticNode] = {} + return None if shape_error is not None else self._dual(node) + args = [self.arith(a) for a in node.args] with_relation = any(k in node.kwargs for k in builtin.relation_kwargs) roles = {k: v for k, v in node.kwargs.items() if builtin.kind_of(k, with_relation=with_relation) == 'role'} if roles and 'by' not in node.kwargs: @@ -579,99 +691,215 @@ def _call(self, node: FunctionCallNode) -> ArithmeticNode: f'{self.context}: {node.name}({", ".join(f"{k}=" for k in roles)}) names a column of a relation, ' f'and no by= names the relation. Write {builtin.usage}' ) + dims: dict[str, str | None] = {} + amounts: dict[str, int | str | None] = {} + edge: _Edge | None = None for key, value in node.kwargs.items(): match builtin.kind_of(key, with_relation=with_relation): case 'edge': - kwargs[key] = self._edge(value, node.name) + edge = self._edge(value, node.name) case 'dimension': - kwargs[key] = self._dim_ref(value, node.name, key) - case 'relation': - kwargs[key] = self._relation_ref(value, node.name, key, roles, node.kwargs.get('along')) - case 'role': - pass + dims[key] = self._dim_ref(value, node.name, key) case 'value': - kwargs[key] = self._amount(value, node.name, key) - case None: - pass # a keyword the operator does not declare; the shape error already named it - return FunctionCallNode(node.name, args, kwargs) + amounts[key] = self._amount(value, node.name, key) + case 'relation' | 'role' | None: + pass + read = None + if 'by' in node.kwargs and builtin.kind_of('by') == 'relation': + read = self._relation_ref(node.kwargs['by'], node.name, 'by', roles, dims.get('along')) + unread = ( + shape_error is not None + or not args + or args[0] is None + or None in dims.values() + or None in amounts.values() + or ('edge' in node.kwargs and edge is None) + or ('by' in node.kwargs and read is None) + ) + if unread: + return None + return self._built(node.name, cast('Expression', args[0]), dims, amounts, edge, read) - def _amount(self, value: ArithmeticNode, operator: str, key: str) -> ArithmeticNode: - """``offset=`` or ``window=``: a number or a parameter name, never an expression. + def _built( + self, + operator: str, + operand: Expression, + dims: Mapping[str, str | None], + amounts: Mapping[str, int | str | None], + edge: _Edge | None, + read: Direction | Partition | None, + ) -> Expression | None: + """The node *operator* builds from its read arguments, or ``None`` with the refusal appended.""" + if operator == 'sum': + if read is not None: + assert isinstance(read, Direction), 'a sum reads its relation in a direction' + return GroupSum(operand, read) + if (over := dims.get('over')) is not None: + return Sum(operand, (over,)) + return self._bare_sum(operand) + if operator == 'at': + assert isinstance(read, Direction), 'at reads its relation in a direction' + return Pullback(operand, read) + assert read is None or isinstance(read, Partition), 'a translation reads its relation as a partition' + along = dims['along'] + assert along is not None + wrap, fill = edge if edge is not None else (False, None) + if operator == 'shift': + offset = amounts['offset'] + assert offset is not None + if not self._edge_fits(operand, offset, wrap=wrap, fill=fill): + return None + return Translate(operand, along, offset, wrap=wrap, fill=fill, partition=read) + if fill is not None: + self.errors.append( + f"{self.context}: sum_back(edge=...) takes 'wrap' or nothing. A window sums the terms " + f'it reaches, so a position before the first contributes nothing rather than a ' + f'fill value; add the constant to the expression if you want one.' + ) + return None + width = amounts['window'] + assert width is not None + return WindowSum(operand, along, width, wrap=wrap, partition=read) - Closed so that :func:`math_spec.dimensions._check_named_amount` sees every - parameter an amount carries. + def _bare_sum(self, operand: Expression) -> Expression | None: + """``sum(x)`` with no ``over=`` or ``by=`` reduces every dim the operand carries, which it has to carry some of.""" + try: + inner = dims_of(operand, self.ns.schema, self.context) + except DimensionError as e: + self.errors.append(str(e)) + return None + if not inner: + self.errors.append( + f'{self.context}: sum() with no over= or by= sums every dim the operand ' + f'carries, and this one carries none — the expression is already a ' + f'scalar. Drop the sum.' + ) + return None + return Sum(operand, tuple(sorted(inner))) + + def _edge_fits(self, operand: Expression, offset: int | str, *, wrap: bool, fill: float | None) -> bool: + """What a ``shift``'s ``edge=`` may say, and where saying nothing is an answer. + + Every rule here is decidable from the file — whether the operand + carries a variable, whether the offset is named, what the edge is + written as — so a file breaking one is refused at load rather than by + whoever lowers it. """ + if wrap: + return True + has_var = carries_variable(operand) + if has_var and fill is not None and fill != 0: + self.errors.append( + f'{self.context}: shift(edge={fill:g}) over an expression containing a variable — only ' + f'fill=0 is representable there, since a vacated slot contributes no term. A nonzero ' + f'fill would be a constant standing where a term was; add that constant to the ' + f'expression instead.' + ) + return False + if fill is None and _vacates(offset) and not has_var: + self.errors.append(_shift_over_data_message(self.context)) + return False + if fill is None and isinstance(offset, str): + self.errors.append(f'{self.context}: {_named_offset_edge_message(offset)}') + return False + return True + + def _amount(self, value: ArithmeticNode, operator: str, key: str) -> int | str | None: + """``offset=`` or ``window=``: a whole number in the operator's range, or the name of a parameter. + + Closed so that :func:`math_spec.dimensions._check_named_amount` sees + every parameter an amount carries, and so that a program's + ``offset`` and ``width`` are the ``int | str`` they say. + """ + if self._formal(value): + return None + words = AMOUNTS[operator] if (literal := _literal(value)) is not None: - return literal - if not isinstance(_without_sign(value), NameNode): + if not (literal.value.is_integer() and literal.value >= words.minimum): + self.errors.append(f'{self.context}: {operator}({key}=...) {words.form}') + return None + return int(literal.value) + bare = _without_sign(value) + if not isinstance(bare, NameNode): self.errors.append( f'{self.context}: {operator}({key}=) takes a number or the name of an integer parameter. ' f'Precompute it as a parameter.' ) - return value - return self._arith(value, amount=True) + return None + if self._formal(bare): + return None + if self.ns.kind(bare.name) != 'parameter': + if self._name(bare) is not None: + self.errors.append(f'{self.context}: {operator}({key}=...) {words.form}') + return None + if isinstance(value, UnaryOperatorNode) and value.op == '-': + self.errors.append( + f'{self.context}: {operator}({key}=-{bare.name}) negates a named {words.noun}. {words.negated}' + ) + return None + return bare.name - def _edge(self, value: ArithmeticNode, operator: str) -> ArithmeticNode: + def _edge(self, value: ArithmeticNode, operator: str) -> _Edge | None: """``edge=``: the closed keyword ``wrap``, or a number to contribute; a name here is a typo.""" if self._formal(value): - return value + return None if isinstance(value, KeywordNode): if value.value == EDGE_WRAP: - return EdgeNode() + return True, None self.errors.append(f'{self.context}: {edge_error(operator, repr(value.value))}') - return value + return None if isinstance(value, NameNode): if value.name == EDGE_WRAP: self.errors.append( f'{self.context}: {operator}(edge={EDGE_WRAP}) is a bare name where a keyword belongs. ' f"Write edge='{EDGE_WRAP}', quoted." ) - return value + return None self.errors.append(f'{self.context}: {edge_error(operator, value.name)}') - return value + return None if (literal := _literal(value)) is None: self.errors.append( f"{self.context}: {operator}(edge=) is an expression, and an edge is the keyword '{EDGE_WRAP}' " f'or a number. Write the number itself.' ) - return value - return literal + return None + return False, literal.value - def _dim_ref(self, value: ArithmeticNode, operator: str, key: str) -> ArithmeticNode: + def _dim_ref(self, value: ArithmeticNode, operator: str, key: str) -> str | None: """An operator kwarg whose *value* must name a declared dimension.""" if self._formal(value): - return value + return None if not isinstance(value, NameNode): self.errors.append(f'{self.context}: {operator}({key}=...) must name a dimension.') - return value + return None if value.name not in self.ns.dimensions: self.errors.append( _undeclared_dim(self.context, operator, f'{key}={value.name}', value.name, self.ns, self.formals) ) - return value - return DimensionNode(value.name) + return None + return value.name - def _dual(self, node: FunctionCallNode) -> ArithmeticNode: - """``dual(c)`` typed to the leaf it is, its one argument the name of a declared constraint. + def _dual(self, node: FunctionCallNode) -> Dual | None: + """``dual(c)`` as the leaf it is, its one argument the name of a declared constraint. Constraints sit outside the flat namespace, so this store is consulted only here — a bare name in arithmetic never reaches it. A dual standing where the math is built is refused separately - (:mod:`math_spec.validation`); this pass only types the name. + (:mod:`math_spec.degree`); this pass only types the name. """ (value,) = node.args if self._formal(value): - return node + return None if not isinstance(value, NameNode): self.errors.append( f'{self.context}: dual() takes the name of a declared constraint, written bare — ' f'dual(). Name the constraint whose row dual you want.' ) - return node + return None if value.name not in self.ns.constraints: self.errors.append(self.ns.unknown_constraint(value.name, self.context, formals=self.formals)) - return node - return DualNode(value.name) + return None + return Dual(value.name) def _relation_ref( self, @@ -679,50 +907,48 @@ def _relation_ref( operator: str, key: str, roles: Mapping[str, ArithmeticNode], - over: ArithmeticNode | None, - ) -> ArithmeticNode: - """An operator's ``by=``, with the ``over=`` and ``into=`` that say which direction it is read in. + along: str | None, + ) -> Direction | Partition | None: + """An operator's ``by=`` as the direction or the partition the call reads its relation in. A relation carries its own dimensions, so the call names columns rather than dims: ``over=`` the column consumed, ``into=`` the column produced, every other key column joined on. A value column not named is not read, and a bare relation's columns are all key. One call addresses one table, so several columns of one table are a list and - several tables are not. + several tables are not. *along* is the dimension a translation steps + along, already read, or ``None`` where it was refused. """ names = names_in(value) if not names: self.errors.append(f'{self.context}: {operator}({key}=...) must name a relation.') - return value + return None if len(names) > 1: self.errors.append( f'{self.context}: {operator}({key}={shown(names)}) names {len(names)} relations, and one call ' f'reads one table. Declare one relation with the columns of all of them, or read them in turn, ' f'one call each.' ) - return value + return None name = names[0] if name in self.formals: - return value + return None if (problem := self._not_a_relation(name, operator, key)) is not None: self.errors.append(problem) - return value + return None if any(n in self.formals for v in roles.values() for n in names_in(v)): - return value + return None read = {k: self._role_name(v, operator, k) for k, v in roles.items()} if any(r is None for r in read.values()): - return value + return None named = {k: r for k, r in read.items() if r is not None} if operator in ('shift', 'sum_back'): if 'within' not in named: - return value # the call shape refused it already, with the wording that names the rewrite - over_dim = over.name if isinstance(over, NameNode | DimensionNode) and not self._formal(over) else None - partition = self._partition(name, operator, over_dim, named['within']) - return value if partition is None else PartitionNode(partition) + return None # the call shape refused it already, with the wording that names the rewrite + return self._partition(name, operator, along, named['within']) if not ({'over', 'into'} <= set(named)): - return value # the call shape refused it already, with the wording that names the rewrite - direction = self._direction(name, operator, named['over'], named['into']) - return value if direction is None else DirectionNode(direction) + return None # the call shape refused it already, with the wording that names the rewrite + return self._direction(name, operator, named['over'], named['into']) def _role_name(self, value: ArithmeticNode, operator: str, key: str) -> tuple[str, ...] | None: """``over=`` or ``into=`` as the column names it must be — one bare name, or a bracketed list of them.""" @@ -1013,14 +1239,14 @@ def _pulled_back(self, node: UnresolvedPredicateCallNode, mask: Mask) -> Predica found = len(self.errors) roles = {key: node.kwargs[key] for key in ('over', 'into')} by = self._relation_ref(node.kwargs['by'], 'at', 'by', roles, None) - if len(self.errors) > found or not isinstance(by, DirectionNode): + if len(self.errors) > found or not isinstance(by, Direction): return node try: - dims = pulled_back_dims(by.direction, mask.dims, context, 'the predicate') + dims = pulled_back_dims(by, mask.dims, context, 'the predicate') except DimensionError as refusal: self.errors.append(str(refusal)) return node - return PulledBackPredicate(mask, by.direction, tuple(sorted(dims))) + return PulledBackPredicate(mask, by, tuple(sorted(dims))) def _count(self, node: UnresolvedCountNode) -> Predicate | UnresolvedWhereNode: """``count(, over=) `` — how many coordinates the predicate admits. @@ -1107,9 +1333,7 @@ def _plain(self, node: UnresolvedComparisonNode) -> _Plain | None: return None return _Plain(name, node.op, value, quoted) - def _expression_comparison( - self, node: UnresolvedComparisonNode - ) -> ExpressionComparison[ArithmeticNode] | UnresolvedComparisonNode: + def _expression_comparison(self, node: UnresolvedComparisonNode) -> ExpressionComparison | UnresolvedComparisonNode: """``expression expression``: each side expanded, typed and held to what a mask may read. A side is read as an expression is — macros and named expressions @@ -1118,7 +1342,7 @@ def _expression_comparison( """ ns, context = self.ns, self.context found = len(self.errors) - sides = [] + sides: list[Expression] = [] for side in (node.left, node.right): if isinstance(side, ColumnNode | KeywordNode): self.errors.append(_not_arithmetic(context, side)) @@ -1134,12 +1358,14 @@ def _expression_comparison( except ValueError as e: self.errors.append(prefixed(context, e)) continue - sides.append(self._arith(expanded)) + if (resolved := self.arith(expanded)) is not None: + sides.append(resolved) if len(self.errors) > found: return node + assert len(sides) == 2, 'a side of a where builds or refuses, since a where holds no formal' dims: set[str] = set() for side in sides: - if degree.carries_variable(side): + if carries_variable(side): self.errors.append( f'{context}: a where compares expressions, and one side names a variable. A where mask ' f'is built before variables exist — it may test parameters and dimension coordinates only.' @@ -1491,18 +1717,13 @@ def _listed(items: list[str]) -> str: return f'{", ".join(quoted[:-1])} and {quoted[-1]}' -def _is_number(side: ArithmeticNode) -> bool: +def _is_number(side: Expression) -> 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)) + return all(isinstance(n, Constant | Negate | Add | Multiply | Divide | Power) for n in walk(side)) def _literal(value: ArithmeticNode) -> NumberNode | None: - """The number a literal names, its sign folded in — ``None`` where *value* is not one. - - Folded here so that every later reader of an ``offset=`` or ``edge=`` — - the dim rules, lowering, the typesetter — meets one signed number rather - than each peeling a unary minus of its own. - """ + """The number a literal names, its sign folded in — ``None`` where *value* is not one.""" if isinstance(value, NumberNode): return value if isinstance(value, UnaryOperatorNode) and isinstance(value.operand, NumberNode): @@ -1574,3 +1795,43 @@ def _relation_pair_error(context: str, node: _Plain, other: str, ns: Namespace, f'where they are over the same dimension.' ) return None + + +def _vacates(offset: int | str) -> bool: + """Whether a translation leaves anything behind. + + A literal zero step reaches every coordinate from itself, so there is no + vacated position for an ``edge=`` to answer for and the refusal has + nothing to refuse. A *named* offset may be zero in the data and is not + known here, so it vacates until proved otherwise. + """ + return offset != 0 + + +def _named_offset_edge_message(name: str) -> str: + """Why a named offset must say what the vacated positions contribute. + + The absent edge propagates through a presence frame keyed by the translated + dimension alone, and a per-entity offset vacates a different slot for each + entity — which that frame cannot say. Refused rather than answered wrongly + (#850); the two edges that write their own answer are allowed. + """ + return ( + f'shift(offset={name}) leaves the vacated positions absent, which a ' + f'per-entity offset cannot say yet.\n' + f"Add edge='wrap' for a cyclic translation, or edge= for what the " + f'vacated positions contribute.' + ) + + +def _shift_over_data_message(context: str) -> str: + """The three ways out of a translation over data with no ``edge=``, the third being two things at once.""" + return ( + f'{context}: shift() over a variable-free expression leaves vacated positions with no ' + f'value, and inventing one is what silently pinned a bound to zero. Say which you mean:\n' + f" shift(x, along=d, offset=n, edge='wrap') the dimension really is cyclic\n" + f' shift(x, along=d, offset=n, edge=0) the vacated positions contribute zero\n' + f' ...and a where: excluding them the vacated rows should not exist at all\n' + f'A where: alone does not lift this — it is decided on the expression, before any mask ' + f'is read — and edge=0 alone leaves a row whose bound is that zero.' + ) diff --git a/src/math_spec/typesetting/symbols.py b/src/math_spec/typesetting/symbols.py index 2d58fe30..b752d881 100644 --- a/src/math_spec/typesetting/symbols.py +++ b/src/math_spec/typesetting/symbols.py @@ -15,9 +15,10 @@ from pathlib import Path from typing import TYPE_CHECKING, cast -import math_spec.degree as degree from math_spec._yaml import read_yaml +from math_spec.degree import calls_dual from math_spec.errors import SchemaError, did_you_mean +from math_spec.program import carries_variable from math_spec.typesetting.format import NOTATIONS if TYPE_CHECKING: @@ -84,8 +85,8 @@ def chosen_expressions(schema: Spec) -> frozenset[str]: """ return frozenset( name - for name, node in schema.resolved.expressions.items() - if degree.carries_variable(node) or degree.calls_dual(node) + for name, entry in schema.resolved.expressions.items() + if carries_variable(entry.body) or calls_dual(entry.body) ) diff --git a/src/math_spec/typesetting/walk.py b/src/math_spec/typesetting/walk.py index b4037963..2f6c7aec 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -2,7 +2,7 @@ # # SPDX-License-Identifier: MIT -"""The walk: resolved AST → typeset lines. Written once, for every format. +"""The walk: resolved tree → typeset lines. Written once, for every format. Everything here is a decision about the *math* — where a bracket changes the reading, which dimension a reduction binds, that a mask belongs on the ∀ rather @@ -15,54 +15,56 @@ from dataclasses import dataclass, field, replace from typing import TYPE_CHECKING, Literal, assert_never -from math_spec._expression_parser import ( - ArithmeticNode, - BinaryOperator, - BinaryOperatorNode, - CasesNode, - DefinitionNode, - DimensionNode, - DirectionNode, - DualNode, - EdgeNode, - FunctionCallNode, - KwargNode, - NumberNode, - ParameterNode, - PartitionNode, - UnaryOperatorNode, - UnresolvedNode, - VariableNode, -) from math_spec.dimensions import dims_of +from math_spec.piecewise import curve_frame from math_spec.program import ( + Add, And, BooleanLiteral, + Cases, + Constant, CountComparison, DimensionComparison, DimensionPosition, Direction, + Divide, + Dual, + Expression, ExpressionComparison, + GroupSum, Mask, + Multiply, + Named, + Negate, Not, Or, + Parameter, ParameterComparison, ParameterDefined, + Partition, + Power, Predicate, PredicateOperator, + Pullback, PulledBackPredicate, RelationComparison, RelationDefined, RelationPairComparison, + Sum, + Translate, TranslatedPredicate, + Variable, VariableDefined, + WindowSum, ) +from math_spec.resolution import remainder from math_spec.typesetting.format import Entry, Line, OperatorName if TYPE_CHECKING: import datetime from collections.abc import Iterable, Mapping + from math_spec._expression_parser import BinaryOperator from math_spec.model import PiecewiseBlock, RelationBlock, SosBlock, Spec from math_spec.typesetting.format import Format from math_spec.typesetting.symbols import Symbols @@ -81,7 +83,6 @@ #: align on the way it aligns a constraint. AlignedComparison = ( ParameterComparison - # pyrefly: ignore[implicit-any-type-argument] # the union is an isinstance target, which takes no parameterized class | ExpressionComparison | CountComparison | DimensionComparison @@ -116,18 +117,6 @@ } -def _amount(node: ArithmeticNode) -> int | str: - """``shift``'s ``offset=``: a signed number, or the name of a parameter. - - A named offset is always backward — a negated one is refused at load, in - :func:`math_spec.dimensions.check_schema` — which the assert relies on. - """ - if isinstance(node, ParameterNode): - return node.name - assert isinstance(node, NumberNode), 'resolution folds a literal offset to one signed number' - return int(node.value) - - @dataclass(frozen=True) class _Step: """One translation of an index, and what stands where it vacated. @@ -221,12 +210,14 @@ def indexed(self, symbol: str, dims: list[str]) -> str: return self.walk.format.subscript(symbol, [self.subscript(d) for d in dims]) -def _unsigned(node: ArithmeticNode) -> ArithmeticNode | None: +def _unsigned(node: Expression) -> Expression | None: """*node* without its leading minus — on the node, or on the first factor of a product it heads — else ``None``.""" - if isinstance(node, UnaryOperatorNode) and node.op == '-': + if isinstance(node, Negate): return node.operand - if isinstance(node, BinaryOperatorNode) and node.op in ('*', '/') and (left := _unsigned(node.left)) is not None: - return BinaryOperatorNode(node.op, left, node.right) + if isinstance(node, Multiply | Divide): + first, second = (node.left, node.right) if isinstance(node, Multiply) else (node.numerator, node.divisor) + if (head := _unsigned(first)) is not None: + return Multiply(head, second) if isinstance(node, Multiply) else Divide(head, second) return None @@ -271,7 +262,7 @@ def _frame_of(self, name: str) -> list[str]: block = self.schema.expressions[name] if block.cases: return list(block.dims or ()) - return self._sorted(dims_of(self.schema.resolved.expressions[name], self.schema, f"expression '{name}'")) + return self._sorted(dims_of(self.schema.resolved.expressions[name].body, self.schema, f"expression '{name}'")) def _op(self, name: OperatorName) -> str: return self.format.operators[name] @@ -346,151 +337,150 @@ def _number(self, value: float) -> str: # -- arithmetic -------------------------------------------------------- - def _expression(self, node: ArithmeticNode, ctx: _Context, *, need: int = 0) -> str: + def _expression(self, node: Expression, ctx: _Context, *, need: int = 0) -> str: text, precedence = self._arithmetic(node, ctx) return self.format.parenthesise(text) if precedence < need else text - def _arithmetic(self, node: ArithmeticNode, ctx: _Context) -> tuple[str, int]: - """Render *node*, returning the text and the precedence it binds at. + def _arithmetic(self, node: Expression, ctx: _Context) -> tuple[str, int]: + """Render *node*, returning the text and the precedence it binds at.""" + if isinstance(node, Named): + if self.inline_expressions and not isinstance(node.body, Cases): + return self._arithmetic(node.body, ctx) + return ctx.indexed(self.symbols.name[node.name], self.frames[node.name]), _ATOM - A ``NameNode`` here means resolution was skipped, and a bare dimension - or coordinate in a value position is a language error caught long - before this module runs — so meeting either is an assertion, not a - rendering decision. - """ - if isinstance(node, NumberNode): + if isinstance(node, Constant): return self._number(node.value), _ATOM if node.value >= 0 else 1 - if isinstance(node, ParameterNode): + if isinstance(node, Parameter): return ctx.indexed(self.symbols.name[node.name], list(self.schema.parameters[node.name].dims)), _ATOM - if isinstance(node, VariableNode): + if isinstance(node, Variable): return ctx.indexed(self.symbols.name[node.name], list(self.schema.variables[node.name].dims)), _ATOM - if isinstance(node, UnaryOperatorNode): - if node.op == '+': - return self._arithmetic(node.operand, ctx) + if isinstance(node, Negate): text, precedence = self._arithmetic(node.operand, ctx) operand = self.format.parenthesise(text) if precedence < 2 else text return f'{self._op("minus")}{operand}', 2 - if isinstance(node, BinaryOperatorNode): + if isinstance(node, Add | Multiply | Divide | Power): return self._binary(node, ctx) - if isinstance(node, FunctionCallNode): - return self._call(node, ctx) + if isinstance(node, Sum): + return self._sum(node, ctx) - if isinstance(node, DefinitionNode) and self.inline_expressions: - return self._arithmetic(node.body, ctx) + if isinstance(node, GroupSum): + return self._group_sum(node, ctx) - if isinstance(node, CasesNode | DefinitionNode): - return ctx.indexed(self.symbols.name[node.name], self.frames[node.name]), _ATOM + if isinstance(node, Pullback): + return self._pullback(node, ctx) - if isinstance(node, DualNode): - return self._dual(node, ctx), _ATOM + if isinstance(node, Translate): + return self._translate(node, ctx) + + if isinstance(node, WindowSum): + return self._window_sum(node, ctx) - if isinstance(node, UnresolvedNode | KwargNode): - msg = f'{type(node).__name__} reached the typesetter; resolve the expression first.' - raise AssertionError(msg) + if isinstance(node, Cases): + return self.format.cases(self._arms(node, ctx)), _ATOM + + if isinstance(node, Dual): + return self._dual(node, ctx), _ATOM assert_never(node) - def _dual(self, node: DualNode, ctx: _Context) -> str: + def _dual(self, node: Dual, ctx: _Context) -> str: """λ subscripted by the constraint's symbol, then the indices of the constraint's own frame.""" - frame = self._sorted(dims_of(node, self.schema, 'a dual')) + frame = self._sorted(frozenset(self.schema.constraints[node.constraint].dims)) return self.format.subscript( self._op('dual'), [self.symbols.constraint[node.constraint], *(ctx.subscript(d) for d in frame)] ) - def _binary(self, node: BinaryOperatorNode, ctx: _Context) -> tuple[str, int]: + def _binary(self, node: Add | Multiply | Divide | Power, ctx: _Context) -> tuple[str, int]: """Render a binary operator, bracketing only where the reading demands. - Subtraction raises the requirement on its right operand by one: - ``a - (b - c)`` and ``a - (b + c)`` need the bracket; ``a - b*c`` - does not. A negation folds into the sign beside it — ``a + -b`` is - ``a - b`` and ``a - -b`` is ``a + b`` — and as a factor it is - bracketed, since ``a · -b`` is a spelling nobody reads. A power is - atomic to everything but another power, a stacked superscript being - ambiguous. + A subtraction arrives as an addition of a negation and prints as the + subtraction it was: ``a + -b`` is ``a - b`` and ``a - -b`` is ``a + b``, + the sign folding until the right operand carries none. Subtraction + raises the requirement on its right operand by one: ``a - (b - c)`` + and ``a - (b + c)`` need the bracket; ``a - b*c`` does not. A negated + factor is bracketed, since ``a · -b`` is a spelling nobody reads. A + power is atomic to everything but another power, a stacked + superscript being ambiguous. """ - if node.op == '/': - top = self._expression(node.left, ctx) - bottom = self._expression(node.right, ctx) + if isinstance(node, Divide): + top = self._expression(node.numerator, ctx) + bottom = self._expression(node.divisor, ctx) return self.format.fraction(top, bottom), _ATOM - if node.op == '**': - base = self._expression(node.left, ctx, need=_PRECEDENCE['**'] + 1) - return self.format.superscript(base, self._expression(node.right, ctx)), _PRECEDENCE['**'] - precedence = _PRECEDENCE[node.op] + if isinstance(node, Power): + base = self._expression(node.base, ctx, need=_PRECEDENCE['**'] + 1) + return self.format.superscript(base, self._expression(node.exponent, ctx)), _PRECEDENCE['**'] + op: BinaryOperator = '*' if isinstance(node, Multiply) else '+' + precedence = _PRECEDENCE[op] left = self._expression(node.left, ctx, need=precedence) - operand, op = node.right, node.op - if op in ('+', '-') and (unsigned := _unsigned(operand)) is not None: - operand, op = unsigned, '-' if op == '+' else '+' - negated_factor = op == '*' and isinstance(operand, UnaryOperatorNode) and operand.op == '-' + operand = node.right + if op == '+': + while (unsigned := _unsigned(operand)) is not None: + operand, op = unsigned, '-' if op == '+' else '+' + negated_factor = op == '*' and isinstance(operand, Negate) need = _ATOM if negated_factor else _PRECEDENCE[op] + (1 if op == '-' else 0) right = self._expression(operand, ctx, need=need) names: dict[BinaryOperator, OperatorName] = {'*': 'cdot', '+': 'plus', '-': 'minus'} return self.format.joined([left, right], self._op(names[op])), precedence - def _call(self, node: FunctionCallNode, ctx: _Context) -> tuple[str, int]: - """Render an operator: a translation at the leaves, or a summation. + def _sum(self, node: Sum, ctx: _Context) -> tuple[str, int]: + """A reduction over named dims: one dummy index per dim, in declaration order.""" + memberships = [] + inner = ctx + for d in self._sorted(frozenset(node.over)): + dummy, inner = inner.reducing(d) + memberships.append(self._membership(d, dummy)) + domain = self.format.joined(memberships, '') + return self.format.summation(domain, self._reduction_body(node.operand, inner)), _PRECEDENCE['+'] + + def _group_sum(self, node: GroupSum, ctx: _Context) -> tuple[str, int]: + """A sum through a relation: a dummy per consumed dim, and the row it joins on as the domain's condition.""" + direction = node.direction + dummies: dict[str, str] = {} + inner = ctx + for d in direction.consumed_dims: + dummies[d], inner = inner.reducing(d) + conditions = list(self._grouping(direction, dummies, ctx)) + domain = ( + f'{self.format.joined([self._membership(d, dummies[d]) for d in direction.consumed_dims], "")} ' + f'{self._op("such_that")} {self.format.joined(conditions, self._op("and"))}' + ) + return self.format.summation(domain, self._reduction_body(node.operand, inner)), _PRECEDENCE['+'] + + def _pullback(self, node: Pullback, ctx: _Context) -> tuple[str, int]: + """``at`` emits no operator of its own: it re-indexes the operand, so the read shows at the leaves.""" + return self._arithmetic(node.operand, self._pulled_back(node.direction, ctx)) + + def _translate(self, node: Translate, ctx: _Context) -> tuple[str, int]: + """``shift`` emits no operator of its own: it re-indexes the operand, so the translation shows at the leaves. - ``shift`` and ``at`` emit no operator of their own — they re-index the - operand, so the substitution shows at the leaves. A ``sum`` naming no - dim binds every dim its operand carries, and the domain has to say - which, since the call does not. + ``edge='wrap'`` and a number are the two policies that print a symbol + of their own; absent is the bare shift, whose vacated positions are + absent. """ - if node.name == 'shift': - dim = node.kwargs['along'] - assert isinstance(dim, DimensionNode) - step = self._step(_amount(node.kwargs['offset']), node.kwargs.get('edge')) - self.noticed.policies.add(step.policy) - step = replace(step, within=self._group(node.kwargs.get('by'), dim.name)) - return self._arithmetic(node.args[0], ctx.translated(dim.name, step)) - - if node.name == 'sum_back': - over = node.kwargs['along'] - assert isinstance(over, DimensionNode) - policy = 'wrap' if isinstance(node.kwargs.get('edge'), EdgeNode) else 'plain' - step = _Step(1, policy, within=self._group(node.kwargs.get('by'), over.name)) - self.noticed.policies.add(step.policy) - source, inner = ctx.reducing(over.name) - lag = f'{ctx.subscript(over.name)} {self._translation(step)} {source}' - domain = ( - f'{source} {self._op("in")} {self.symbols.set[over.name]} {self._op("such_that")} ' - f'0 {self._op("le")} {lag} {self._op("lt")} {self._width(node.kwargs["window"])}' - ) - body = self._reduction_body(node.args[0], inner) - return self.format.summation(domain, body), _PRECEDENCE['+'] - - if node.name == 'at': - by = node.kwargs['by'] - assert isinstance(by, DirectionNode) - return self._arithmetic(node.args[0], self._pulled_back(by.direction, ctx)) - - if (by := node.kwargs.get('by')) is not None: - assert isinstance(by, DirectionNode) - direction = by.direction - dummies: dict[str, str] = {} - inner = ctx - for d in direction.consumed_dims: - dummies[d], inner = inner.reducing(d) - conditions = list(self._grouping(direction, dummies, ctx)) - domain = ( - f'{self.format.joined([self._membership(d, dummies[d]) for d in direction.consumed_dims], "")} ' - f'{self._op("such_that")} {self.format.joined(conditions, self._op("and"))}' - ) - elif (consumed := node.kwargs.get('over')) is not None: - assert isinstance(consumed, DimensionNode) - dummy, inner = ctx.reducing(consumed.name) - domain = self._membership(consumed.name, dummy) - else: - memberships = [] - inner = ctx - for d in self._sorted(dims_of(node.args[0], self.schema, 'a sum')): - dummy, inner = inner.reducing(d) - memberships.append(self._membership(d, dummy)) - domain = self.format.joined(memberships, '') - return self.format.summation(domain, self._reduction_body(node.args[0], inner)), _PRECEDENCE['+'] + policy: TranslationPolicy = 'wrap' if node.wrap else 'edge' if node.fill is not None else 'plain' + fill = '' if node.fill is None else self._number(node.fill) + self.noticed.policies.add(policy) + step = _Step(node.offset, policy, fill, self._group(node.partition)) + return self._arithmetic(node.operand, ctx.translated(node.along, step)) + + def _window_sum(self, node: WindowSum, ctx: _Context) -> tuple[str, int]: + """``sum_back``: a sum over the positions behind the row, the lag written as a translation of the index.""" + policy: TranslationPolicy = 'wrap' if node.wrap else 'plain' + step = _Step(1, policy, within=self._group(node.partition)) + self.noticed.policies.add(step.policy) + source, inner = ctx.reducing(node.along) + lag = f'{ctx.subscript(node.along)} {self._translation(step)} {source}' + domain = ( + f'{source} {self._op("in")} {self.symbols.set[node.along]} {self._op("such_that")} ' + f'0 {self._op("le")} {lag} {self._op("lt")} {self._width(node.width)}' + ) + body = self._reduction_body(node.operand, inner) + return self.format.summation(domain, body), _PRECEDENCE['+'] def _pulled_back(self, direction: Direction, ctx: _Context) -> _Context: """*ctx* with each dimension *direction* consumes read at the relation, as ``at`` re-indexes a leaf.""" @@ -517,21 +507,19 @@ def _grouping(self, direction: Direction, dummies: Mapping[str, str], ctx: _Cont return [self._relation_row(direction.name, at)] return [f'{self._relation_read(direction.name, at, r)} {self._op("equal")} {at[r]}' for r in fixed] - def _group(self, by: ArithmeticNode | None, dim: str) -> str: + def _group(self, partition: Partition | None) -> str: """A ``by=`` as the superscript its translation operator carries. The bare index, not the subscript in force: the group is a property of the row being written, and a window whose operand is itself translated still asks which group *that row* is in. """ - if by is None: + if partition is None: return '' - assert isinstance(by, PartitionNode) - partition = by.partition at = {r: self.symbols.index[partition.dim(r)] for r in (partition.along, *partition.joined)} return self._tuple([self._relation_read(partition.name, at, r) for r in partition.group]) - def _width(self, node: ArithmeticNode) -> str: + def _width(self, width: int | str) -> str: """``sum_back``'s ``window=``: a number, or a parameter's own symbol. Unsubscripted where it is named, as a translation's named offset is: @@ -539,39 +527,21 @@ def _width(self, node: ArithmeticNode) -> str: where repeating them inside a summation's domain crowds out the condition that domain exists to state. """ - if isinstance(node, ParameterNode): - return self.symbols.name[node.name] - assert isinstance(node, NumberNode) - return self._number(node.value) - - def _step(self, by: int | str, edge: ArithmeticNode | None) -> _Step: - """Which of the three edge policies this ``shift`` asked for. - - ``edge='wrap'`` is the language's one keyword and arrives as an - :class:`EdgeNode`; a number in the same position stays a - :class:`NumberNode` and is the value the vacated positions contribute; - absent is the bare shift, whose vacated positions are absent. - """ - if isinstance(edge, EdgeNode): - return _Step(by, 'wrap') - if edge is None: - return _Step(by, 'plain') - assert isinstance(edge, NumberNode) - return _Step(by, 'edge', self._number(edge.value)) + if isinstance(width, str): + return self.symbols.name[width] + return self._number(float(width)) def _membership(self, dim: str, index: str | None = None) -> str: return f'{index or self.symbols.index[dim]} {self._op("in")} {self.symbols.set[dim]}' - def _reduction_body(self, node: ArithmeticNode, ctx: _Context) -> str: + def _reduction_body(self, node: Expression, ctx: _Context) -> str: """What sits to the right of a sum, bracketed only where it must be. A sum binds everything up to the next ``+`` or ``-`` at its own level, so an additive body needs the bracket and nothing else does — including a nested reduction, which is unambiguous. """ - additive = isinstance(node, UnaryOperatorNode) or ( - isinstance(node, BinaryOperatorNode) and node.op in ('+', '-') - ) + additive = isinstance(node, Negate | Add) return self._expression(node, ctx, need=2 if additive else 0) # -- where strings ----------------------------------------------------- @@ -731,9 +701,9 @@ def _objective(self) -> list[Line]: if block is None: return [] sense = self._op('minimize' if block.sense == 'minimize' else 'maximize') - node = self.schema.resolved.objective - assert node is not None, 'validation resolves the objective the file declares' - return [Line(label='', left=sense, right=self._expression(node, self._context()))] + objective = self.schema.resolved.objective + assert objective is not None, 'validation resolves the objective the file declares' + return [Line(label='', left=sense, right=self._expression(objective.expression, self._context()))] def _constraints(self) -> list[Line]: """Every constraint, then every curve. @@ -749,13 +719,13 @@ def _constraints(self) -> list[Line]: def _constraint(self, name: str) -> Line: block = self.schema.constraints[name] - node, where = self.schema.resolved.constraints[name] + constraint = self.schema.resolved.constraints[name] ctx = self._context(frame=block.dims) - condition = self._condition(ctx, where) + condition = self._condition(ctx, constraint.where) return Line( label=name, - left=self._expression(node.left, ctx), - right=f'{self._op(_PREDICATES[node.op])} {self._expression(node.right, ctx)}', + left=self._expression(constraint.lhs, ctx), + right=f'{self._op(_PREDICATES[constraint.sense])} {self._expression(constraint.rhs, ctx)}', condition=self._quantifier(list(block.dims), condition), ) @@ -784,13 +754,13 @@ def _defined(self) -> list[str]: def definition(self, name: str) -> Line: """The line defining one named expression, ``symbol = body`` over its frame.""" - node = self.schema.resolved.expressions[name] + entry = self.schema.resolved.expressions[name] frame = self.frames[name] ctx = self._context(frame) body = ( - self.format.cases(self._arms(node, ctx)) - if isinstance(node, CasesNode) - else self._expression(node.body, ctx) + self.format.cases(self._arms(entry.body, ctx)) + if isinstance(entry.body, Cases) + else self._expression(entry.body, ctx) ) return Line( label=name, @@ -818,21 +788,22 @@ def line(self, name: str) -> Line: return self._piecewise(name) return self._variable(name) - def _arms(self, node: CasesNode, ctx: _Context) -> list[tuple[str, str]]: - """Each arm as its value and the words saying where it applies. + def _arms(self, node: Cases, ctx: _Context) -> list[tuple[str, str]]: + """Each region as its value and the words saying where it applies. - Which arm is the fallback is a fact about the math, so the *walk* - chooses between "if" and "otherwise" and a Format only stacks the rows. + Which region is the fallback is a fact about the math, so the *walk* + chooses between "if" and "otherwise" and a Format only stacks the + rows: the last region is the ``otherwise`` where its mask is the + remainder of the others, which is how resolution builds it. """ - arms = [] - for arm in node.arms: - when = ( - self.format.prose('otherwise') - if arm.when is None - else f'{self.format.prose("if ")} {self._predicate(arm.when, ctx, need=_WHERE_PRECEDENCE["and"])}' - ) - arms.append((self._expression(arm.value, ctx), when)) - return arms + *stated, last = node.regions + arms = [(self._expression(region.value, ctx), self._arm_condition(region.when, ctx)) for region in stated] + left_over = bool(stated) and last.when == remainder(region.when for region in stated) + when = self.format.prose('otherwise') if left_over else self._arm_condition(last.when, ctx) + return [*arms, (self._expression(last.value, ctx), when)] + + def _arm_condition(self, when: Mask, ctx: _Context) -> str: + return f'{self.format.prose("if ")} {self._predicate(when.root, ctx, need=_WHERE_PRECEDENCE["and"])}' def _variables(self) -> list[Line]: """One line per variable, and one more for a set the variable carries. @@ -900,7 +871,8 @@ def _assumptions(self) -> list[Line]: def _assumption(self, name: str) -> Line: """One assumption: the predicate over the frame both its masks name, under its ``where``.""" - holds, where, _ = self.schema.resolved.assumptions[name] + assumption = self.schema.resolved.assumptions[name] + holds, where = assumption.predicate, assumption.where frame = self._sorted(holds.dims | (where.dims if where is not None else frozenset())) ctx = self._context(frame) if isinstance(holds.root, AlignedComparison): @@ -920,7 +892,7 @@ def _piecewise(self, name: str) -> Line: """ block = self.schema.piecewise[name] links = self.schema.resolved.piecewise[name] - frame = self._curve_frame(name, block, links) + frame = list(curve_frame(self.schema, name, block, links)) ctx = self._context([*frame, block.over]) locus = self._locus(block, ctx) bounded = next((i for i, link in enumerate(block.links) if link.sign != '=='), None) @@ -989,20 +961,6 @@ def _gate(self, block: PiecewiseBlock, ctx: _Context) -> str: [(symbol, f'{self.format.prose("if ")} {where}'), ('1', self.format.prose('otherwise'))] ) - def _curve_frame(self, name: str, block: PiecewiseBlock, links: tuple[ArithmeticNode, ...]) -> list[str]: - """The dimensions the block builds one curve per coordinate of: every one its links and its gate carry. - - The union the expansion takes its own frame from, and the expansion has - already held it to the rules — that no link carries the breakpoint - dimension among them (:mod:`math_spec.piecewise`). - """ - dims: frozenset[str] = frozenset() - for i, node in enumerate(links): - dims |= dims_of(node, self.schema, f"piecewise '{name}' link {i}") - if block.activity is not None: - dims |= frozenset(self.schema.variables[block.activity].dims) - return self._sorted(dims) - def _bound(self, ctx: _Context, value: float | str) -> str: if isinstance(value, str): return ctx.indexed(self.symbols.name[value], list(self.schema.parameters[value].dims)) diff --git a/src/math_spec/validation.py b/src/math_spec/validation.py index b876837d..1e47c205 100644 --- a/src/math_spec/validation.py +++ b/src/math_spec/validation.py @@ -7,36 +7,38 @@ from __future__ import annotations from collections.abc import Mapping -from typing import TYPE_CHECKING, Literal, overload - -import math_spec.degree as degree -from math_spec._expression_parser import ( - ArithmeticNode, - CasesNode, - ComparisonNode, - DefinitionNode, - ParsedNode, -) +from typing import TYPE_CHECKING + from math_spec._yaml import read_model from math_spec.dimensions import check_schema -from math_spec.errors import LanguageError, SchemaError, prefixed -from math_spec.expansion import expand, parse_and_expand, parse_template -from math_spec.model import AssumptionBlock, Spec -from math_spec.piecewise import assumptions_of -from math_spec.program import BooleanLiteral, Mask, VariableDefined +from math_spec.errors import SchemaError, prefixed +from math_spec.expansion import expand, parse_template +from math_spec.model import Spec +from math_spec.piecewise import assumptions_of, curve_frame +from math_spec.program import ( + Assumption, + BooleanLiteral, + ConstraintDeclaration, + Mask, + ObjectiveDeclaration, + VariableDefined, +) from math_spec.resolution import ( Namespace, Resolved, - ResolvedAssumption, - ResolvedConstraint, mask_of, + resolve_constraint_text, resolve_expression, + resolve_expression_text, resolve_where_text, ) if TYPE_CHECKING: from pathlib import Path + from math_spec.model import AssumptionBlock + from math_spec.program import Named + def to_spec(model: str | Path | Mapping[str, object] | Spec) -> Spec: """Load and validate a model definition — the language's front door. @@ -81,8 +83,9 @@ def validate_expressions(schema: Spec) -> Resolved: ``over=snapshot`` under a formal ``snapshot`` cannot say which it means; - every dim rule (``dimensions.check_schema``), once names resolve. - A ``piecewise:`` block's links are resolved here too, so the typesetter - reads the curve a file states without expanding it. + A ``piecewise:`` block's links are resolved here too, and its frame + checked, so the typesetter reads the curve a file states without expanding + it and the expansion reads the typed links. Returns: Every declaration's typed tree — what the dim rules, lowering and the @@ -100,7 +103,7 @@ def validate_expressions(schema: Spec) -> Resolved: context = f"Macro '{mname}'" formals = frozenset((*macro.args, *macro.kwargs)) try: - body_ast = expand(parse_template(mname, macro, context), ns, context, shadow=formals) + body_ast = expand(parse_template(mname, macro, context), ns, context) except ValueError as e: errors.append(prefixed(context, e)) continue @@ -112,7 +115,7 @@ def validate_expressions(schema: Spec) -> Resolved: ) resolve_expression(body_ast, ns, context, errors, formals=formals) - expressions: dict[str, CasesNode | DefinitionNode] = {} + expressions: dict[str, Named] = {} for ename in schema.expressions: node, refusals = ns.named_entry(ename) errors.extend(refusals) @@ -126,35 +129,34 @@ def validate_expressions(schema: Spec) -> Resolved: for vname, vdef in schema.variables.items() } - constraints: dict[str, ResolvedConstraint] = {} + constraints: dict[str, ConstraintDeclaration] = {} for cname, cdef in schema.constraints.items(): context = f"Constraint '{cname}'" where = resolve_where_text(cdef.where, ns, context, errors) - expression = _check_expression(cdef.expression, ns, context, errors, comparison=True, ceiling=2) - if expression is not None: - constraints[cname] = ResolvedConstraint(expression, mask_of(where)) + if (sides := resolve_constraint_text(cdef.expression, ns, context, errors)) is not None: + lhs, sense, rhs = sides + constraints[cname] = ConstraintDeclaration(tuple(cdef.dims), lhs, sense, rhs, mask_of(where)) objective = None if schema.objective is not None: - objective = _check_expression( - schema.objective.expression, ns, 'The objective', errors, comparison=False, ceiling=2 - ) + expression = resolve_expression_text(schema.objective.expression, ns, 'The objective', errors, ceiling=2) + if expression is not None: + objective = ObjectiveDeclaration(schema.objective.sense, expression) - assumptions: dict[str, ResolvedAssumption] = {} + assumptions: dict[str, Assumption] = {} for aname, adef in schema.assumptions.items(): if (assumption := _assumption(aname, adef, ns, errors)) is not None: assumptions[aname] = assumption for block, pw in schema.piecewise.items(): for aname, assumed in assumptions_of(block, pw).items(): - entry = AssumptionBlock(holds=assumed.holds, where=assumed.where, description=assumed.description) - if (assumption := _assumption(aname, entry, ns, errors)) is not None: + if (assumption := _assumption(aname, assumed, ns, errors)) is not None: assumptions[aname] = assumption piecewise = {} for pname, pdef in schema.piecewise.items(): links = [ - _check_expression(link.expression, ns, f"piecewise '{pname}' link {i}", errors, comparison=False, ceiling=1) + resolve_expression_text(link.expression, ns, f"piecewise '{pname}' link {i}", errors, ceiling=1) for i, link in enumerate(pdef.links) ] if all(link is not None for link in links): @@ -165,10 +167,12 @@ def validate_expressions(schema: Spec) -> Resolved: resolved = Resolved(expressions, variables, constraints, objective, ns.relations, assumptions, piecewise) check_schema(schema, resolved) + for pname, pdef in schema.piecewise.items(): + curve_frame(schema, pname, pdef, resolved.piecewise[pname]) return resolved -def _assumption(name: str, block: AssumptionBlock, ns: Namespace, errors: list[str]) -> ResolvedAssumption | None: +def _assumption(name: str, block: AssumptionBlock, ns: Namespace, errors: list[str]) -> Assumption | None: """One ``assumptions:`` entry typed, or ``None`` once anything in it failed. A predicate the connectives decide is refused: one that folds to true @@ -198,7 +202,7 @@ def _assumption(name: str, block: AssumptionBlock, ns: Namespace, errors: list[s if len(errors) > found: return None assert holds is not None, 'a where string that read to nothing appended an error' - return ResolvedAssumption(Mask(holds), mask_of(where), block.description) + return Assumption(Mask(holds), mask_of(where), block.description) def _decided_assumption(context: str, text: str, *, value: bool) -> str: @@ -222,77 +226,3 @@ def _decided_where(context: str, text: str, *, value: bool) -> str: f'{context}: the where {text!r} folds to false, so the assumption is checked on no row. ' f'Delete the entry, or write the where the data can satisfy.' ) - - -@overload -def _check_expression( - expression: str, - ns: Namespace, - context: str, - errors: list[str], - *, - comparison: Literal[True], - ceiling: int | None, -) -> ComparisonNode | None: ... -@overload -def _check_expression( - expression: str, - ns: Namespace, - context: str, - errors: list[str], - *, - comparison: Literal[False], - ceiling: int | None, -) -> ArithmeticNode | None: ... - - -def _check_expression( - expression: str, - ns: Namespace, - context: str, - errors: list[str], - *, - comparison: bool, - ceiling: int | None, -) -> ParsedNode | None: - """Parse, expand, resolve and degree-check one expression — nothing resolves once the shape is wrong, and a comparison must carry a variable (#1171). - - Returns the typed tree, or ``None`` once anything failed, the problem - appended to *errors*. ``ceiling`` is the degree the position honours, and - ``None`` for an ``expressions:`` entry's body: what the math admits - (:func:`~math_spec.degree.check_expression`) is a rule about the position - that *reads* it, so it fires on the expanded tree of every constraint, - objective, bound, where and piecewise link, and not where an entry is - declared. - """ - try: - ast = parse_and_expand(expression, ns, context) - except ValueError as e: - errors.append(prefixed(context, e)) - return None - if comparison and not isinstance(ast, ComparisonNode): - errors.append( - f'{context}: expression must contain exactly one comparison operator (<=, >=, ==).\nGot: {expression!r}' - ) - return None - if not comparison and isinstance(ast, ComparisonNode): - errors.append(f'{context}: expression must not contain a comparison operator.\nGot: {expression!r}') - return None - resolved = resolve_expression(ast, ns, context, errors) - if resolved is None or ceiling is None: - return resolved - try: - degree.check_expression(resolved, context, ceiling=ceiling) - except LanguageError as e: - errors.append(str(e)) - return None - if isinstance(resolved, ComparisonNode) and not degree.carries_variable(resolved): - errors.append( - f'{context}: neither side of the comparison carries a variable, so the row decides nothing.\n' - f'Got: {expression!r}\n' - f'A constraint is a claim about a decision, and a comparison of numbers and parameters ' - f'is settled before the solve — no consumer builds a row for it. Name the variable it should ' - f'bound, or state the fact under `assumptions:`, where the consumer binding the data checks it.' - ) - return None - return resolved diff --git a/tests/fixtures.py b/tests/fixtures.py index 05be8513..777b9751 100644 --- a/tests/fixtures.py +++ b/tests/fixtures.py @@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Any from math_spec import Spec +from math_spec._expression_parser import ComparisonNode from math_spec._yaml import parse_yaml, read_yaml from math_spec.errors import LanguageError from math_spec.expansion import parse_and_expand @@ -18,8 +19,7 @@ from math_spec.validation import to_spec if TYPE_CHECKING: - from math_spec._expression_parser import ParsedNode - from math_spec.program import Mask + from math_spec.program import Expression, Mask EXAMPLES = Path(__file__).resolve().parent.parent / 'examples' @@ -97,16 +97,30 @@ def raw_of(source: str | Path | dict[str, Any]) -> dict[str, Any]: return read_yaml(source) if isinstance(source, Path) else parse_yaml(source) -def expression_of(text: str, ns: Namespace, context: str) -> ParsedNode: - """Parse, expand and resolve one expression, raising every problem at once rather than collecting.""" +def expression_of(text: str, ns: Namespace, context: str) -> Expression: + """Parse, expand and resolve one expression into its program tree, raising every problem at once rather than collecting.""" errors: list[str] = [] - resolved = resolve_expression(parse_and_expand(text, ns, context), ns, context, errors) + ast = parse_and_expand(text, ns, context) + assert not isinstance(ast, ComparisonNode), 'a comparison is a constraint, which comparison_of reads' + resolved = resolve_expression(ast, ns, context, errors) if errors: raise LanguageError('\n'.join(errors)) assert resolved is not None return resolved +def comparison_of(text: str, ns: Namespace, context: str) -> tuple[Expression, str, Expression]: + """Parse, expand and resolve one comparison into its two program trees and the sense between them.""" + errors: list[str] = [] + ast = parse_and_expand(text, ns, context) + assert isinstance(ast, ComparisonNode), 'a value is an expression, which expression_of reads' + left, right = (resolve_expression(side, ns, context, errors) for side in (ast.left, ast.right)) + if errors: + raise LanguageError('\n'.join(errors)) + assert left is not None and right is not None + return left, ast.op, right + + def where_of(text: str | None, ns: Namespace, context: str, self_variable: str | None = None) -> Mask | None: """Parse and resolve one where string into the mask a declaration carries, raising every problem at once.""" errors: list[str] = [] diff --git a/tests/test_degree.py b/tests/test_degree.py index edb32039..3dfb61d2 100644 --- a/tests/test_degree.py +++ b/tests/test_degree.py @@ -13,8 +13,8 @@ import pytest from math_spec import LanguageError -from math_spec._expression_parser import NameNode -from math_spec.degree import calls_dual, carries_variable, check_binary, check_expression +from math_spec.degree import calls_dual, check_binary, check_expression +from math_spec.program import carries_variable from math_spec.resolution import Namespace from tests.fixtures import SMALL_MODEL, expression_of, schema_of @@ -106,11 +106,6 @@ def test_the_context_prefixes_the_sentence_and_an_empty_one_leaves_it_bare(conte check_binary(_ast('p * q'), context, ceiling=1) -def test_carries_variable_refuses_an_unresolved_name(): - with pytest.raises(AssertionError, match=r'resolution\.resolve_expression'): - carries_variable(NameNode('p')) - - def _dual_ast(text: str): schema = schema_of(SMALL_MODEL, **{'constraints.lim': {'dims': ['g'], 'expression': 'p <= c'}}) return expression_of(text, Namespace(schema), 'test') @@ -136,12 +131,12 @@ def test_calls_dual_finds_a_dual_wherever_it_stands(text, found): def test_calls_dual_finds_a_dual_inside_a_cased_arm(): - """`calls_dual` recurses through a `CasesNode` arm, not only the top node. + """`calls_dual` recurses through a region of a `Cases`, not only the top node. - The reference resolves straight to the `CasesNode` expansion.py builds, so - this also guards that `children()` walking its arm values reaches a dual a - non-recursive check — one that only inspected the node it was handed — - would miss. + The reference resolves to the `Named` node carrying the block, so this also + guards that the walk steps through it into the region values, reaching a + dual a non-recursive check — one that only inspected the node it was + handed — would miss. """ schema = schema_of( SMALL_MODEL, diff --git a/tests/test_dimensions.py b/tests/test_dimensions.py index 32ded946..01c5932d 100644 --- a/tests/test_dimensions.py +++ b/tests/test_dimensions.py @@ -11,6 +11,7 @@ import pytest from math_spec.dimensions import DimensionError, _check_where_dims, dims_of +from math_spec.errors import LanguageError from math_spec.program import Mask, RelationPairComparison from math_spec.resolution import Namespace from math_spec.validation import to_spec @@ -321,7 +322,8 @@ def test_a_bare_name_reaches_the_variable_a_dual_the_same_named_constraint(): ], ) def test_an_ill_dimensioned_expression_is_rejected(expr, match): - with pytest.raises(DimensionError, match=match): + """A rule on the operand's dims is the dim checker's; one on the form of an amount is resolution's, so the class is the language's.""" + with pytest.raises(LanguageError, match=match): _dims(expr) @@ -420,7 +422,7 @@ class TestTheEdgeRulesAreDecidedAtLoad: def _refused(self, expression: str) -> str: raw = override(self.BASE, **{'constraints.k.expression': expression}) - with pytest.raises(DimensionError) as caught: + with pytest.raises(LanguageError) as caught: to_spec(raw) return str(caught.value) diff --git a/tests/test_expansion.py b/tests/test_expansion.py index d0d2eaa9..5117cfbe 100644 --- a/tests/test_expansion.py +++ b/tests/test_expansion.py @@ -10,11 +10,12 @@ import pytest -from math_spec._expression_parser import ComparisonNode, DefinitionNode, with_children from math_spec.errors import LanguageError from math_spec.expansion import parse_and_expand +from math_spec.lowering import inline +from math_spec.program import Multiply, Named, Parameter, Sum, Variable from math_spec.resolution import Namespace -from tests.fixtures import DISPATCH_MODEL, SMALL_MODEL, expression_of, schema_of +from tests.fixtures import DISPATCH_MODEL, SMALL_MODEL, comparison_of, expression_of, schema_of WEIGHTED_SUM = { 'args': ['array', 'weights'], @@ -25,12 +26,19 @@ schema = partial(schema_of, DISPATCH_MODEL) -def _bodies(node): - """*node* with every named expression's body standing bare where its name was.""" - if isinstance(node, ComparisonNode): - return ComparisonNode(node.op, _bodies(node.left), _bodies(node.right)) - node = node.body if isinstance(node, DefinitionNode) else node - return with_children(node, _bodies) +def _resolved(text, ns): + """*text* as its program tree, or as its two sides and the sense between them where it compares.""" + if any(op in text for op in ('<=', '>=', '==')): + return comparison_of(text, ns, 'expression') + return expression_of(text, ns, 'expression') + + +def _bodies(resolved): + """*resolved* with every named expression's body standing bare where its name was.""" + if isinstance(resolved, tuple): + left, op, right = resolved + return inline(left), op, inline(right) + return inline(resolved) @pytest.mark.parametrize( @@ -105,20 +113,21 @@ def _bodies(node): ], ) def test_a_call_expands_to_core_ast(expressions, macros, call, want): - """The math a call expands to is what `want` spells; a plain named - expression's body arrives under the node carrying its name, which `_bodies` - reads through, as every pass does.""" + """The math a call expands to is what `want` spells; a named expression's + body arrives under the `Named` node carrying its name, which `_bodies` + inlines, as lowering does.""" ns = Namespace(schema(expressions=expressions, macros=macros)) - assert _bodies(expression_of(call, ns, 'expression')) == expression_of(want, ns, 'expression') + assert _bodies(_resolved(call, ns)) == _resolved(want, ns) def test_a_named_expression_arrives_under_the_node_carrying_its_name(): ns = Namespace(schema(expressions={'gen_cost': 'p * cost'})) - expanded = parse_and_expand('sum(gen_cost, over=generator)', ns, 'e') - assert expanded.args[0] == DefinitionNode('gen_cost', expression_of('p * cost', ns, 'e')), ( + resolved = expression_of('sum(gen_cost, over=generator)', ns, 'e') + assert resolved == Sum(Named('gen_cost', Multiply(Variable('p'), Parameter('cost'))), ('generator',)), ( 'the body is inlined resolved and the name kept, for the typesetter to define it once' ) - assert expanded.args[0] is parse_and_expand('gen_cost', ns, 'another use'), 'every use reads the one node' + assert isinstance(resolved, Sum) + assert resolved.operand is expression_of('gen_cost', ns, 'another use'), 'every use reads the one node' @pytest.mark.parametrize( diff --git a/tests/test_lowering.py b/tests/test_lowering.py index e0f2875f..ff7586c9 100644 --- a/tests/test_lowering.py +++ b/tests/test_lowering.py @@ -11,20 +11,20 @@ from __future__ import annotations from dataclasses import FrozenInstanceError -from typing import TYPE_CHECKING, get_args +from typing import get_args import pytest from math_spec import LanguageError, Spec, to_program -from math_spec._expression_parser import FunctionCallNode, NumberNode from math_spec._where_parser import parse_where from math_spec.exclusivity import overlapping -from math_spec.lowering import _Lowering, lower_program +from math_spec.lowering import lower_program from math_spec.piecewise import expand_piecewise from math_spec.program import ( QUADRATIC_POSITIONS, Add, And, + Assumption, BooleanLiteral, Cases, Constant, @@ -38,7 +38,6 @@ ExpressionComparison, Footprint, GroupSum, - Holds, Mask, Multiply, Negate, @@ -71,9 +70,6 @@ from math_spec.resolution import Namespace from tests.fixtures import DISPATCH_MODEL, EXAMPLES, SMALL_MODEL, expanded, expression_of, override, schema_of, where_of -if TYPE_CHECKING: - from math_spec._expression_parser import ArithmeticNode - DISPATCH_YAML = EXAMPLES / 'dispatch.yaml' #: The mask `examples/dispatch.yaml` puts on `dispatch`, as the plan carries it. @@ -109,12 +105,11 @@ ) -def resolved(text: str, schema: Spec) -> ArithmeticNode: - """Parse, expand and resolve — exactly what the lowering pass receives. +def resolved(text: str, schema: Spec) -> Expression: + """Parse, expand and resolve — the program tree a declaration holds. - A raw ``parse_expression`` result still holds ``NameNode``s, and lowering - asserts those never reach it. The ``'t'`` is the error-context label the - resolver stamps on refusals, not a dimension. + The ``'t'`` is the error-context label the resolver stamps on refusals, + not a dimension. """ return expression_of(text, Namespace(schema), 't') @@ -177,9 +172,9 @@ def test_a_file_with_no_objective_lowers_to_no_sense(): 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(dispatch_schema) - node = expression_of('shift(dispatch, along=snapshot, offset=-1, edge=+2)', ns, 't') - assert isinstance(node, FunctionCallNode) - assert (node.kwargs['offset'], node.kwargs['edge']) == (NumberNode(-1.0), NumberNode(2.0)) + node = expression_of('shift(dispatch, along=snapshot, offset=-1, edge=+0)', ns, 't') + assert isinstance(node, Translate) + assert (node.offset, node.fill) == (-1, 0.0) @pytest.mark.parametrize( @@ -503,7 +498,7 @@ def test_assumptions_carry_the_file_s_entries_and_the_curves_behind_them(): program = to_program(expanded(EXAMPLES / 'piecewise_lp.yaml', 'piecewise')) derived = [name for name in program.assumptions if name.startswith('cost_curve_')] - assert all(isinstance(a, Holds) for a in program.assumptions.values()), ( + assert all(isinstance(a, Assumption) for a in program.assumptions.values()), ( 'a method states its conditions in the language the file writes, so one kind stands in the mapping' ) assert derived == [ @@ -519,7 +514,7 @@ def test_an_assumption_lowers_both_of_its_masks(): program = to_program(override(SHAPES_MODEL, assumptions={'sound': {'holds': 'c <= 0.5 * k', 'where': 'flag'}})) assumption = program.assumptions['sound'] - assert assumption == Holds( + assert assumption == Assumption( Mask(ExpressionComparison(Parameter('c'), '<=', Multiply(Constant(0.5), Parameter('k')), ('g',))), Mask(ParameterDefined('flag', ('g',))), ), 'the arithmetic side is a program expression, and the where is the mask the file wrote' @@ -573,9 +568,8 @@ def test_a_mask_with_no_arithmetic_is_the_same_mask_after_lowering(dispatch_prog assert dispatch_program.variables['dispatch'].where == Mask(CAPACITY_POSITIVE) -def test_a_power_lowers_to_a_node_of_its_own(dispatch_schema): - lowered = _Lowering(dispatch_schema, 't').expr(resolved('cost ** cost', dispatch_schema)) - assert isinstance(lowered, Power), 'a variable-free power has a plan node of its own' +def test_a_power_resolves_to_a_node_of_its_own(dispatch_schema): + assert isinstance(resolved('cost ** cost', dispatch_schema), Power), 'a variable-free power has a node of its own' @pytest.mark.parametrize( @@ -643,10 +637,9 @@ def test_a_power_lowers_to_a_node_of_its_own(dispatch_schema): ), ], ) -def test_a_construct_lowers_to_its_node(shapes_schema, expression, expected): +def test_a_construct_resolves_to_its_node(shapes_schema, expression, expected): """Which node each surface construct becomes, and every field it arrives with.""" - lowered = _Lowering(shapes_schema, 't').expr(resolved(expression, shapes_schema)) - assert lowered == expected, 'the whole frozen node, so no field is asserted by omission' + assert resolved(expression, shapes_schema) == expected, 'the whole frozen node, so no field is asserted by omission' def test_a_partition_keeps_its_group_when_the_relation_gains_a_value_column(): diff --git a/tests/test_operators.py b/tests/test_operators.py index d749ad0b..3c8f6041 100644 --- a/tests/test_operators.py +++ b/tests/test_operators.py @@ -6,41 +6,17 @@ from __future__ import annotations -import pytest +from math_spec.operators import AMOUNTS, BUILTINS -from math_spec.degree import _REDUCTIONS -from math_spec.dimensions import _AMOUNTS, _CALL_RULES -from math_spec.lowering import _CALLS -from math_spec.operators import BUILTIN_NAMES, BUILTINS -#: Every operator that reaches lowering and the dim rules as a call. ``dual`` -#: resolves to a leaf of its own, so no table after resolution has a row for it. -CALLED = BUILTIN_NAMES - {'dual'} +def test_the_amount_words_cover_every_operator_taking_an_amount(): + """An operator added to `BUILTINS` alone fails as a `KeyError` inside resolution (#401). - -@pytest.mark.parametrize( - ('table', 'keys'), - [ - pytest.param(_CALLS, CALLED, id='lowering-has-one-rewrite-per-operator'), - pytest.param(_CALL_RULES, CALLED, id='the-dim-rules-have-one-rule-per-operator'), - pytest.param( - _AMOUNTS, - frozenset(name for name, builtin in BUILTINS.items() if builtin.required_value_kwargs), - id='the-amount-words-cover-every-operator-taking-an-amount', - ), - ], -) -def test_every_table_keyed_by_operator_agrees_with_the_closed_set(table, keys): - """An operator added to `BUILTINS` alone fails as a `KeyError` inside lowering (#401). - - Each table is keyed by operator name and read with `[]`, so the closed - set and every table must name the same operators — here, before a model - finds the missing row. + The table is keyed by operator name and read with `[]`, so the closed set + and the table must name the same operators — here, before a model finds + the missing row. """ - assert frozenset(table) == keys, ( - f'missing rows: {sorted(keys - set(table))}; stray rows: {sorted(set(table) - keys)}' + keys = frozenset(name for name, builtin in BUILTINS.items() if builtin.required_value_kwargs) + assert frozenset(AMOUNTS) == keys, ( + f'missing rows: {sorted(keys - set(AMOUNTS))}; stray rows: {sorted(set(AMOUNTS) - keys)}' ) - - -def test_a_reduction_is_an_operator(): - assert _REDUCTIONS <= BUILTIN_NAMES, f'not operators: {sorted(_REDUCTIONS - BUILTIN_NAMES)}' diff --git a/tests/test_parser.py b/tests/test_parser.py index c7593b32..0d4d50f3 100644 --- a/tests/test_parser.py +++ b/tests/test_parser.py @@ -19,22 +19,13 @@ import math_spec.program as program_module from math_spec._expression_parser import ( BinaryOperatorNode, - CasesNode, ComparisonNode, - DefinitionNode, - DimensionNode, - DirectionNode, - DualNode, - EdgeNode, FunctionCallNode, KeywordNode, NameListNode, NameNode, NumberNode, - ParameterNode, - PartitionNode, UnaryOperatorNode, - VariableNode, parse_expression, ) from math_spec._where_parser import ( @@ -46,11 +37,8 @@ from math_spec.program import ( And, BooleanLiteral, - Direction, Not, Or, - Partition, - RelationDeclaration, _conjuncts, ) @@ -538,37 +526,3 @@ def test_a_parsed_tree_prints_to_text_that_parses_to_the_same_tree(text): ) def test_a_node_prints_as_the_file_writes_it(text, printed): assert str(parse_expression(text)) == printed, 'the spelling is the one a file could be written with' - - -_ZONE_OF = RelationDeclaration((('u', 'unit'), ('zone', 'zone')), ('u',)) - - -@pytest.mark.parametrize( - ('node', 'printed'), - [ - pytest.param(VariableNode('p'), 'p', id='a-variable'), - pytest.param(ParameterNode('cost'), 'cost', id='a-parameter'), - pytest.param(DimensionNode('t'), 't', id='a-dimension'), - pytest.param(DualNode('budget'), 'dual(budget)', id='a-dual'), - pytest.param( - DirectionNode(Direction('zone_of', _ZONE_OF, ('u',), ('zone',), ())), - 'zone_of', - id='a-relation-read-in-a-direction', - ), - pytest.param( - PartitionNode(Partition('zone_of', _ZONE_OF, 'u', ('zone',), ())), - 'zone_of', - id='a-relation-stepped-along-as-a-partition', - ), - pytest.param(EdgeNode(), "'wrap'", id='a-resolved-edge'), - pytest.param(DefinitionNode('headroom', NameNode('p')), 'headroom', id='a-named-expression-prints-its-name'), - pytest.param(CasesNode('startup', ()), 'startup', id='and-so-does-a-cased-one'), - ], -) -def test_a_node_resolution_built_prints_the_name_the_file_wrote(node, printed): - """These eight never come out of the parser, so no round trip reaches them: - resolution rewrites a `NameNode` into each. A `DefinitionNode` and a - `CasesNode` stand where a name stood, and the name is what the file says at - that position — printing the inlined body would print an expression the - author never wrote.""" - assert str(node) == printed, 'a resolved node prints the text it was resolved from' diff --git a/tests/test_piecewise.py b/tests/test_piecewise.py index 76b82367..46601e11 100644 --- a/tests/test_piecewise.py +++ b/tests/test_piecewise.py @@ -14,10 +14,10 @@ import pytest from math_spec import CURVATURES -from math_spec.errors import LanguageError, PiecewiseExpansionError, SchemaError +from math_spec.errors import LanguageError, SchemaError from math_spec.lowering import lower_program, to_program from math_spec.piecewise import expand_piecewise -from math_spec.program import Holds, assumption_message +from math_spec.program import Assumption, assumption_message from tests.fixtures import DISPATCH_MODEL, expanded, override, raw_of, schema_of #: Larger than a minimal probe on purpose: a curve that exercises adjacency @@ -88,7 +88,7 @@ def test_an_emitted_set_may_not_collide_with_a_declared_one(): """The emitted-name rule, for the one declaration kind that is new.""" - with pytest.raises(PiecewiseExpansionError, match="emitted sos 'cost_curve' collides"): + with pytest.raises(SchemaError, match="writes sos 'cost_curve', which this file already declares"): schema_of(NONCONVEX_YAML, sos={'cost_curve': {'variable': 'p', 'over': 'snapshot', 'type': 1}}) @@ -247,7 +247,7 @@ def test_a_malformed_block_is_refused(model, patch, match): ) def test_a_link_outside_the_language_is_named_where_the_user_wrote_it(link_expression, message): """Lowering would catch these too, but naming ``cost_curve_link0`` — a declaration the user never wrote.""" - with pytest.raises(PiecewiseExpansionError, match=message) as exc: + with pytest.raises(SchemaError, match=message) as exc: schema_of(NONCONVEX_YAML, **{'piecewise.cost_curve.links': [[link_expression, 'bp_x'], ['op_cost', 'bp_y']]}) assert "piecewise 'cost_curve' link 0" in str(exc.value) @@ -261,7 +261,7 @@ def test_a_link_reading_a_nonlinear_entry_is_refused(): declaration: the entry-declaration relocation for the other math positions is `TestValidateExpressions.test_a_nonlinear_entry_is_refused_where_the_math_reads_it`. """ - with pytest.raises(PiecewiseExpansionError, match='the divisor contains variables') as exc: + with pytest.raises(SchemaError, match='the divisor contains variables') as exc: schema_of( NONCONVEX_YAML, **{ @@ -280,7 +280,7 @@ def test_a_link_reading_a_degree_two_product_entry_is_refused(): reads a named entry and rejects the product. A constraint and the objective accept degree 2, so they are not the refusing site here. """ - with pytest.raises(PiecewiseExpansionError, match='which is degree 2') as exc: + with pytest.raises(SchemaError, match='which is degree 2') as exc: schema_of( NONCONVEX_YAML, **{ @@ -307,7 +307,7 @@ def test_a_link_reading_a_dual_entry_is_refused(): a dual carries no variable — and hand lowering a leaf no piecewise expansion can build. """ - with pytest.raises(PiecewiseExpansionError, match='a dual exists only after a solve'): + with pytest.raises(SchemaError, match='a dual exists only after a solve'): schema_of( NONCONVEX_YAML, **{ @@ -327,7 +327,7 @@ def test_a_link_reading_a_dual_entry_is_refused(): ) def test_a_gate_that_is_not_a_variable_is_refused(activity, match): """Only a variable has a declaration to say what its absence means, and the block needs that answer.""" - with pytest.raises(PiecewiseExpansionError, match=match): + with pytest.raises(SchemaError, match=match): expand_piecewise(schema_of(GATED, **{'piecewise.cost_curve.activity': activity})) @@ -456,7 +456,7 @@ def test_a_block_assumes_of_its_data_what_the_method_implies(): 'cost_curve_breakpoints', 'cost_curve_contiguous', ], 'an lp curve with a mask assumes all five, each named after the block that implies it' - assert all(isinstance(a, Holds) for a in program.assumptions.values()), ( + assert all(isinstance(a, Assumption) for a in program.assumptions.values()), ( 'a method states its conditions in the same language the file does, so a consumer has one kind to read' ) assert program.assumptions['cost_curve_increasing'].predicate.names_read == frozenset({'bp_x'}), ( @@ -472,7 +472,9 @@ def test_a_block_assumes_of_its_data_what_the_method_implies(): def test_a_curves_conditions_cannot_collide_with_a_written_assumption(): """A condition a method states is a name the block emits, and a file writing it is the collision every emitted name is.""" - with pytest.raises(LanguageError, match="emitted assumption 'cost_curve_increasing' collides"): + with pytest.raises( + SchemaError, match="writes assumption 'cost_curve_increasing', which this file already declares" + ): expanded(override(LP, assumptions={'cost_curve_increasing': 'bp_x > 0'}), 'piecewise') diff --git a/tests/test_public_surface.py b/tests/test_public_surface.py index a7dbc0e8..28e54880 100644 --- a/tests/test_public_surface.py +++ b/tests/test_public_surface.py @@ -27,7 +27,7 @@ 'Spec', 'to_spec', 'program', 'to_program', # the error tree 'MathSpecError', 'LanguageError', 'SchemaError', 'DimensionError', - 'PiecewiseExpansionError', 'did_you_mean', 'schema_error', + 'did_you_mean', 'schema_error', # the verdicts a consumer asks for rather than re-deriving 'advice', 'Advice', # the closed operator set, and the wording of its refusals diff --git a/tests/test_validation.py b/tests/test_validation.py index 8e51e6a7..93bdd3f9 100644 --- a/tests/test_validation.py +++ b/tests/test_validation.py @@ -1997,8 +1997,9 @@ def record(*args, **kwargs): return record + doors = (resolution.resolve_expression, resolution.resolve_constraint_text, resolution.resolve_where_text) for module in (validation, resolution): - for door in (resolution.resolve_expression, resolution.resolve_where_text): + for door in doors: monkeypatch.setattr(module, door.__name__, recorded(door)) spec = to_spec( @@ -2019,8 +2020,8 @@ def record(*args, **kwargs): to_markdown(spec) assert sorted(seen) == [ - ('resolve_expression', "Constraint 'balance'"), - ('resolve_expression', "Constraint 'spare'"), + ('resolve_constraint_text', "Constraint 'balance'"), + ('resolve_constraint_text', "Constraint 'spare'"), ('resolve_expression', "Named expression 'headroom', case 'opening'"), ('resolve_expression', "Named expression 'headroom', otherwise"), ('resolve_expression', 'The objective'), diff --git a/tests/typesetting/test_golden.py b/tests/typesetting/test_golden.py index c1a01a32..c063dd01 100644 --- a/tests/typesetting/test_golden.py +++ b/tests/typesetting/test_golden.py @@ -15,9 +15,8 @@ import pytest -from math_spec._expression_parser import ArithmeticNode, ComparisonNode, DualNode, FunctionCallNode from math_spec.operators import BUILTIN_NAMES -from math_spec.program import Predicate +from math_spec.program import Dual, Expression, GroupSum, Named, Predicate, Pullback, Sum, Translate, WindowSum from math_spec.typesetting import FORMATS, to_latex, typeset, walk from math_spec.typesetting.format import OPERATOR_NAMES from math_spec.validation import to_spec @@ -125,43 +124,32 @@ def _rendered_trees() -> Iterator[object]: printed at all. """ resolved = to_spec(golden.MODEL).resolved - yield resolved.objective - for expression, mask in resolved.constraints.values(): - yield expression - if mask is not None: - yield mask.root + assert resolved.objective is not None + yield resolved.objective.expression + for constraint in resolved.constraints.values(): + yield constraint.lhs + yield constraint.rhs + if constraint.where is not None: + yield constraint.where.root for mask in resolved.variables.values(): if mask is not None: yield mask.root - for holds, where, _ in resolved.assumptions.values(): - yield holds.root - if where is not None: - yield where.root + for assumption in resolved.assumptions.values(): + yield assumption.predicate.root + if assumption.where is not None: + yield assumption.where.root yield from resolved.expressions.values() for links in resolved.piecewise.values(): yield from links -#: What resolution never hands the walk: the two nodes a where carries before -#: its sides are read, and the three an expression and a where carry before -#: names are resolved. The walk raises on each rather than rendering it, so a -#: fixture reaching one would be a bug in resolution rather than a case worth -#: committing output for. -UNRESOLVED = { - 'UnresolvedComparisonNode', - 'ColumnNode', - 'NameNode', - 'NameListNode', - 'KeywordNode', -} - -#: A dataclass the walk steps *through* rather than renders: an arm has no +#: A dataclass the walk steps *through* rather than renders: a region has no #: 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'} +CARRIERS = {'Region', 'Direction', 'Mask', 'Partition', 'RelationDeclaration'} def test_the_golden_model_carries_every_node_kind_the_walk_renders(): @@ -173,10 +161,10 @@ def test_the_golden_model_carries_every_node_kind_the_walk_renders(): `coverage` installed, and its failure names the construct rather than a line. """ kinds = {type(node).__name__ for tree in _rendered_trees() for node in _nodes(tree)} - CARRIERS - declared = {node.__name__ for node in (*get_args(Predicate), *get_args(ArithmeticNode), ComparisonNode)} - assert kinds == declared - UNRESOLVED, ( + declared = {node.__name__ for node in (*get_args(Predicate), *get_args(Expression), Named)} + assert kinds == declared, ( f'tests/typesetting/golden/model.yaml reaches {sorted(kinds - declared)} and misses ' - f'{sorted(declared - UNRESOLVED - kinds)}. Every node the walk renders needs a case here, ' + f'{sorted(declared - kinds)}. Every node the walk renders needs a case here, ' f'or its arm ships output nobody has read.' ) @@ -184,31 +172,27 @@ def test_the_golden_model_carries_every_node_kind_the_walk_renders(): def test_the_golden_model_calls_every_operator_in_the_language(): """``BUILTINS`` is the closed set, so a new operator lands with its case here. - ``dual`` resolves to its own leaf rather than staying a call, so it is - counted by that leaf. + Each operator resolves to the node it is, so the census counts the nodes + by the verb the file writes them with. """ + verbs = {Sum: 'sum', GroupSum: 'sum', Pullback: 'at', Translate: 'shift', WindowSum: 'sum_back', Dual: 'dual'} nodes = [node for tree in _rendered_trees() for node in _nodes(tree)] - calls = {node.name for node in nodes if isinstance(node, FunctionCallNode)} - calls |= {'dual' for node in nodes if isinstance(node, DualNode)} + calls = {verb for node in nodes for kind, verb in verbs.items() if isinstance(node, kind)} assert calls == BUILTIN_NAMES, ( f'tests/typesetting/golden/model.yaml never calls {sorted(BUILTIN_NAMES - calls)}. ' f'An operator with no case here renders untested.' ) -#: What the fixture cannot reach, by the source text of the line. The guards -#: are what the walk raises when resolution hands it something it types away, -#: so a model reaching one is a bug upstream. The absent objective is the arm a -#: *different* model takes — a file declares at most one — and +#: What the fixture cannot reach, by the source text of the line. A bare +#: ``Cases`` stands under the ``Named`` node resolution builds for its entry +#: and nowhere else, so the arm that would print one in place is the type's +#: closure rather than a case. The absent objective is the arm a *different* +#: model takes — a file declares at most one — and #: `test_a_model_with_no_objective_prints_the_rest` covers it. UNREACHABLE = { - 'if isinstance(node, UnresolvedNode | KwargNode):', - "msg = f'{type(node).__name__} reached the typesetter; resolve the expression first.'", - 'if isinstance(node, ExpressionComparison):', - "msg = 'a lowered comparison reached the typesetter; it prints the resolved tree, which lowering rebuilds.'", - 'if not isinstance(node, ComparisonNode):', - "msg = f'{context}: expected a comparison, got {type(node).__name__}'", - 'raise AssertionError(msg)', + 'if isinstance(node, Cases):', + 'return self.format.cases(self._arms(node, ctx)), _ATOM', 'assert_never(node)', 'assert_never(check)', 'if block is None:',