From 449424e74503cb325ca2759e87ebde25862736a4 Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 22 Sep 2026 10:15:22 +0000 Subject: [PATCH 01/18] docs(about): a relation is an indicator, and a sum through it is a contraction Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01SWBcNGLjNH2i4AqRsfyxaN --- docs/about/relations-as-linear-maps.md | 194 +++++++++++++++++++++++++ mkdocs.yml | 1 + 2 files changed, 195 insertions(+) create mode 100644 docs/about/relations-as-linear-maps.md diff --git a/docs/about/relations-as-linear-maps.md b/docs/about/relations-as-linear-maps.md new file mode 100644 index 00000000..89618352 --- /dev/null +++ b/docs/about/relations-as-linear-maps.md @@ -0,0 +1,194 @@ + + +# Relations as linear maps + +This page says what a [relation](../reference/language/relations.md) is in the +language of linear algebra. It then shows that the dimension rule on that page +and the join an engine runs are one computation. Read it if "the dimensions a call +consumes, produces and joins on" reads as bookkeeping and you want the math it +stands for. + +```yaml +dimensions: + snapshot: { dtype: int } + generator: { dtype: str } + zone: { dtype: str } +relations: + gen_zone: { key: [generator, snapshot], values: zone } +parameters: + zone_cap: { dims: [snapshot, zone] } +variables: + p: { dims: [snapshot, generator] } +constraints: + zonal: + dims: [snapshot, zone] + expression: sum(p, by=gen_zone, over=generator, into=zone) <= zone_cap + pulled: + dims: [snapshot, generator] + where: gen_zone + expression: p <= at(zone_cap, by=gen_zone, over=zone, into=generator) +objective: + sense: minimize + expression: sum(p) +``` + +## A relation is an indicator + +A relation with columns over the dimensions $`D_1, \dots, D_n`$ is a set of +rows, so it is a subset $`R \subseteq D_1 \times \dots \times D_n`$. Its +**indicator** $`\mathbf{1}_R`$ is $`1`$ at a row of the table and $`0`$ +everywhere else. `key:` is a claim about the shape of that set: one row per key +tuple. So a relation with `key: K` and `values: V` is the graph of a function +$`f: K \to V`$, and + +```math +\mathbf{1}_R(k, v) = [\, f(k) = v \,]. +``` + +The function is partial where a key tuple has no row. A bare relation is a +subset and nothing more. Above, `gen_zone` is the graph of +$`f: \mathcal{G} \times \mathcal{T} \to \mathcal{Z}`$. The +[data contract](../reference/language/relations.md#the-data-contract) makes it +one: the loader checks one row per key tuple when the data binds. + +## A sum is a contraction + +`sum(p, by=gen_zone, over=generator, into=zone)` is the product of two arrays, +summed over the one index they share and the call names: + +```math +y_{t,z} = \sum_{g} \mathbf{1}_R(g, t, z) \cdot p_{t,g} = \sum_{g \,:\, f(g,\,t) = z} p_{t,g} +``` + +The right-hand form is what the typesetter +[prints](../reference/notation.md). The left-hand form is a tensor contraction, +and each kind of dimension on the relations page is one position an index can +take in it: + +| The relations page says | In the formula | +| ----------------------- | ----------------------------------------------------------- | +| consumed (`over=`) | $`g`$ is on both factors and summed. It leaves. | +| produced (`into=`) | $`z`$ is on the indicator alone and not summed. It arrives. | +| joined on | $`t`$ is on both factors and not summed. It stays. | +| passes through | on the operand alone and not summed. It stays. | + +**The result carries the free indices.** Those are the indices of the operand +and the indicator together, less the summed one, which is +`(dims(x) − consumed) ∪ produced`. That is the rule in the [expressions +reference](../reference/language/expressions.md#how-dimensions-combine), +and `_read_dims` in `src/math_spec/dimensions.py` computes it as +`(inner - consumed) | produced`. The three refusals beside it are the three +things the formula needs: + +- **The operand carries every consumed dimension**, or there is nothing to sum. +- **The operand carries every joined dimension.** Otherwise $`t`$ would sit on + the indicator alone, which is the produced position, and the call names a + produced column with `into=`. +- **The operand carries no produced dimension.** Otherwise $`z`$ would sit on + both factors and not be summed. That ties the two occurrences together + instead of adding an axis. The language makes you write that tie outside the + operator: `load * sum(p, by=gen_bus, over=generator, into=bus)`. + +So "the indices an operation carries, adds and takes away" is the free-index +rule of a summation convention, applied to one product. Nothing on the relations +page is a separate rule. + +**The joined dimension makes the matrix block-diagonal.** At each $`t`$, +$`M_t[z, g] = \mathbf{1}_R(g, t, z)`$ is a $`|\mathcal{Z}| \times |\mathcal{G}|`$ +matrix of zeros and ones, and $`y_t = M_t\, p_t`$. A relation keyed by one +column has one block. + +## A read is the transpose + +`at(zone_cap, by=gen_zone, over=zone, into=generator)` reads the same table with +the other index bound: + +```math +w_{t,g} = \sum_{z} \mathbf{1}_R(g, t, z) \cdot \mathrm{zone\_cap}_{t,z} = \mathrm{zone\_cap}_{t,\, f(g,\,t)} +``` + +Same indicator, same contraction, so the same free-index rule gives the result's +dimensions. As a matrix it is $`M_t^{\mathsf{T}}`$. Because $`R`$ is the graph +of $`f`$, the sum over $`z`$ has exactly one term where $`(g, t)`$ is in the +domain of $`f`$, and none elsewhere. So the contraction is the composition +$`\mathrm{zone\_cap} \circ f`$, which is a pullback. + +**That one term is the whole difference between `at` and `sum`**, and +resolution decides it from the key alone. A call lands on the columns it names +in `into=` and the columns it joins on. Where those hold the whole key, each +output coordinate meets at most one row, and the call is a read. `_direction` +in `src/math_spec/resolution.py` names this `single_valued`. A `sum` that is +single-valued adds up nothing and is refused toward `at`. An `at` that is not +would have several terms and is refused toward `sum`. The +[relations page](../reference/language/relations.md#aggregates-and-reads) +quotes both messages. + +**The two are adjoint.** For `gen_bus: { key: generator, values: bus }`, +$`x`$ over generators and $`y`$ over buses, + +```math +\langle M x, y \rangle = \sum_{b} y_b \sum_{g} \mathbf{1}_R(g, b)\, x_g = \sum_{g} x_g \sum_{b} \mathbf{1}_R(g, b)\, y_b = \langle x, M^{\mathsf{T}} y \rangle, +``` + +which is what the program means by calling `Pullback` the adjoint of +`GroupSum`. A bare relation has the same matrix without the functional claim. +A column of $`M`$ may hold several ones, so the sum fans out and there is no +read. That is why `at` through a bare relation is refused. + +## The join and the aggregate + +The language fixes the formula, and an engine decides how to evaluate it. A +table stores $`\mathbf{1}_R`$ as its support: the rows where it is $`1`$. +Multiplying a matrix stored that way by a vector takes two steps. + +1. **Pair each row of the table with each row of the operand that agrees on + every shared index**, here $`g`$ and $`t`$. That is an inner equi-join on + the consumed and joined columns. Each pair is one nonzero product + $`1 \cdot p_{t,g}`$, relabelled by the produced column $`z`$. +2. **Add the pairs that agree on the free indices**, here $`(t, z)`$. That is a + group-by on the result's dimensions with a sum. + +lpspec, the reference engine, runs exactly these two steps. `walk_join` in +`src/lpspec/relational/engines/polars/relations.py` is step 1. It runs one inner +join on the consumed and joined dimensions. A select then drops the consumed +dimensions and renames the landing column to the produced one. Step 2 is the +terminal aggregate in `assembly.py`, a `group_by` over the row's coordinates +with `sum`, run once per constraint after every term has landed. Until then a +sum of linear terms is a list of terms, and adding is concatenation. `at` runs +step 1 against the same table and needs no step 2: single-valued means no two +pairs land on one coordinate. Where several $`(g, t)`$ share one $`z`$ the +join fans out, which is a column of $`M^{\mathsf{T}}`$ holding several ones. + +So the two pictures are one. The join is the multiplication by an entry of +$`\mathbf{1}_R`$, which is $`1`$ or absent. The group-by is the $`\sum`$. + +## Where the built model departs from the map + +Two positions of the formula have no variable to build, and there the model a +consumer builds is not the matrix. + +- **An empty fibre is the empty sum.** A zone no generator maps to at $`t`$ has + $`y_{t,z} = 0`$, and the row reads $`0 \le \mathrm{zone\_cap}_{t,z}`$. It + names no variable, so an engine [does not build + it](../reference/language/absence.md#rows-with-no-variable-terms) and reports + the omission. With `>=` the omitted row would have been infeasible. +- **Off the domain there is no value.** $`\mathrm{zone\_cap} \circ f`$ is + undefined where $`f`$ is, so `at` is absent there and [absence + spreads](../reference/language/absence.md#how-absence-travels) to the row. + `where: gen_zone` on `pulled` writes the domain of $`f`$ on the page, so a + reader sees which rows exist without opening the data. + +## Partitions and tests + +The other two uses of a relation do not contract against $`\mathbf{1}_R`$. + +- **A partition steps inside a fibre.** `shift(x, along=snapshot, offset=1, +by=season_of, within=season)` reads the neighbour $`t'`$ of $`t`$ with + $`f(t') = f(t)`$. The fibres of $`f`$ partition the axis, and the frame does + not change. +- **A test is the indicator itself.** A relation's name in a `where` evaluates + $`\mathbf{1}_R`$ at the frame's own coordinate, and keeps the coordinate where + it is $`1`$. diff --git a/mkdocs.yml b/mkdocs.yml index ccc32aff..29ca8113 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -74,6 +74,7 @@ nav: - The limits: about/limits.md - What counts as language: about/what-counts-as-language.md - What counts as public API: about/what-counts-as-public-api.md + - Relations as linear maps: about/relations-as-linear-maps.md - Contributing: contributing.md - Changelog: CHANGELOG.md From 54497645b14e576c938c70700fa842ab656b5da1 Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 22 Sep 2026 11:25:09 +0000 Subject: [PATCH 02/18] feat(language): a relation call is described as the columns it joins on and groups by The program's `Direction` is now `Join`, with `joined` and `grouped` roles in place of `consumed`, `produced` and `joined`; `Pullback` is `Lookup`, a join with no group-by; a partition's `group` is `grouped`. Every refusal, docstring and page says join and group-by in place of consume, produce and land on. The YAML surface is unchanged. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01SWBcNGLjNH2i4AqRsfyxaN --- docs/about/relations-as-linear-maps.md | 113 ++++++++++---------- docs/about/what-counts-as-language.md | 4 +- docs/contributing.md | 4 +- docs/examples/operators.md | 10 +- docs/reference/language/expressions.md | 24 ++--- docs/reference/language/operators.md | 13 +-- docs/reference/language/relations.md | 87 ++++++++------- docs/reference/notation.md | 18 ++-- docs/reference/reading.md | 2 +- examples/operators/at.yaml | 2 +- examples/operators/sum_by_column_lists.yaml | 4 +- examples/operators/sum_by_columns.yaml | 4 +- src/math_spec/_expression_parser.py | 14 +-- src/math_spec/advice.py | 14 +-- src/math_spec/boundedness.py | 4 +- src/math_spec/dimensions.py | 96 ++++++++--------- src/math_spec/exclusivity.py | 2 +- src/math_spec/lowering.py | 18 ++-- src/math_spec/operators.py | 4 +- src/math_spec/program.py | 102 +++++++++++------- src/math_spec/resolution.py | 65 +++++------ src/math_spec/separability.py | 12 +-- src/math_spec/typesetting/walk.py | 74 ++++++------- tests/test_advice.py | 2 +- tests/test_dimensions.py | 34 +++--- tests/test_lowering.py | 68 ++++++------ tests/test_parser.py | 8 +- tests/test_validation.py | 40 +++---- tests/typesetting/golden/latex.out | 6 +- tests/typesetting/golden/markdown.out | 6 +- tests/typesetting/golden/model.yaml | 8 +- tests/typesetting/golden/typst.out | 6 +- tests/typesetting/test_golden.py | 2 +- tests/typesetting/test_walk.py | 4 +- 34 files changed, 459 insertions(+), 415 deletions(-) diff --git a/docs/about/relations-as-linear-maps.md b/docs/about/relations-as-linear-maps.md index 89618352..40125db6 100644 --- a/docs/about/relations-as-linear-maps.md +++ b/docs/about/relations-as-linear-maps.md @@ -6,10 +6,9 @@ SPDX-License-Identifier: CC-BY-4.0 # Relations as linear maps This page says what a [relation](../reference/language/relations.md) is in the -language of linear algebra. It then shows that the dimension rule on that page -and the join an engine runs are one computation. Read it if "the dimensions a call -consumes, produces and joins on" reads as bookkeeping and you want the math it -stands for. +language of linear algebra. It then shows that the join and group-by on that +page are one computation with the sum a paper prints. Read it if "joined on" +and "grouped by" read as database words and you want the math they stand for. ```yaml dimensions: @@ -54,7 +53,7 @@ $`f: \mathcal{G} \times \mathcal{T} \to \mathcal{Z}`$. The [data contract](../reference/language/relations.md#the-data-contract) makes it one: the loader checks one row per key tuple when the data binds. -## A sum is a contraction +## A join and group-by is a contraction `sum(p, by=gen_zone, over=generator, into=zone)` is the product of two arrays, summed over the one index they share and the call names: @@ -65,46 +64,45 @@ y_{t,z} = \sum_{g} \mathbf{1}_R(g, t, z) \cdot p_{t,g} = \sum_{g \,:\, f(g,\,t) The right-hand form is what the typesetter [prints](../reference/notation.md). The left-hand form is a tensor contraction, -and each kind of dimension on the relations page is one position an index can -take in it: +and the two questions the relations page asks of a column are the two +positions an index can take in it: -| The relations page says | In the formula | -| ----------------------- | ----------------------------------------------------------- | -| consumed (`over=`) | $`g`$ is on both factors and summed. It leaves. | -| produced (`into=`) | $`z`$ is on the indicator alone and not summed. It arrives. | -| joined on | $`t`$ is on both factors and not summed. It stays. | -| passes through | on the operand alone and not summed. It stays. | +| The relations page says | In the formula | +| ------------------------ | ----------------------------------------------------------- | +| joined on, `over=` | $`g`$ is on both factors and summed. It leaves. | +| grouped by, `into=` | $`z`$ is on the indicator alone and not summed. It arrives. | +| joined on and grouped by | $`t`$ is on both factors and not summed. It stays. | +| neither | a value column not in the formula. It is not read. | **The result carries the free indices.** Those are the indices of the operand and the indicator together, less the summed one, which is -`(dims(x) − consumed) ∪ produced`. That is the rule in the [expressions -reference](../reference/language/expressions.md#how-dimensions-combine), -and `_read_dims` in `src/math_spec/dimensions.py` computes it as -`(inner - consumed) | produced`. The three refusals beside it are the three -things the formula needs: - -- **The operand carries every consumed dimension**, or there is nothing to sum. -- **The operand carries every joined dimension.** Otherwise $`t`$ would sit on - the indicator alone, which is the produced position, and the call names a - produced column with `into=`. -- **The operand carries no produced dimension.** Otherwise $`z`$ would sit on - both factors and not be summed. That ties the two occurrences together - instead of adding an axis. The language makes you write that tie outside the +`(dims(x) − joined) ∪ grouped`. That is the rule in the [expressions +reference](../reference/language/expressions.md#how-dimensions-combine), and +`_join_dims` in `src/math_spec/dimensions.py` computes it. The three refusals +beside it are the three things the formula needs: + +- **The operand carries every column joined on**, or there is nothing to match. +- **The operand carries every unnamed key column.** Otherwise $`t`$ would sit + on the indicator alone, which is the grouped position, and the call names a + grouped column with `into=`. +- **The operand carries no column grouped by.** Otherwise $`z`$ would sit on + both factors and not be summed. That matches the two occurrences instead of + adding an axis. The language makes you write that match outside the operator: `load * sum(p, by=gen_bus, over=generator, into=bus)`. -So "the indices an operation carries, adds and takes away" is the free-index -rule of a summation convention, applied to one product. Nothing on the relations -page is a separate rule. +So "joined on" and "grouped by" are the two positions of the summation +convention, applied to one product. Nothing on the relations page is a +separate rule. -**The joined dimension makes the matrix block-diagonal.** At each $`t`$, +**The unnamed key column makes the matrix block-diagonal.** At each $`t`$, $`M_t[z, g] = \mathbf{1}_R(g, t, z)`$ is a $`|\mathcal{Z}| \times |\mathcal{G}|`$ matrix of zeros and ones, and $`y_t = M_t\, p_t`$. A relation keyed by one column has one block. -## A read is the transpose +## A join alone is the transpose -`at(zone_cap, by=gen_zone, over=zone, into=generator)` reads the same table with -the other index bound: +`at(zone_cap, by=gen_zone, over=zone, into=generator)` joins the same table +with the other index bound: ```math w_{t,g} = \sum_{z} \mathbf{1}_R(g, t, z) \cdot \mathrm{zone\_cap}_{t,z} = \mathrm{zone\_cap}_{t,\, f(g,\,t)} @@ -117,14 +115,14 @@ domain of $`f`$, and none elsewhere. So the contraction is the composition $`\mathrm{zone\_cap} \circ f`$, which is a pullback. **That one term is the whole difference between `at` and `sum`**, and -resolution decides it from the key alone. A call lands on the columns it names -in `into=` and the columns it joins on. Where those hold the whole key, each -output coordinate meets at most one row, and the call is a read. `_direction` -in `src/math_spec/resolution.py` names this `single_valued`. A `sum` that is -single-valued adds up nothing and is refused toward `at`. An `at` that is not -would have several terms and is refused toward `sum`. The -[relations page](../reference/language/relations.md#aggregates-and-reads) -quotes both messages. +resolution decides it from the key alone. Where the columns a call groups by +hold the whole key, every group is one row, and the group-by adds nothing. +`_join` in `src/math_spec/resolution.py` names this `one_row_per_group`. A +`sum` with one row per group adds up nothing and is refused toward `at`. An +`at` with several rows per group would have several terms and is refused toward +`sum`. The +[relations page](../reference/language/relations.md#joins-and-group-bys) +quotes the message. **The two are adjoint.** For `gen_bus: { key: generator, values: bus }`, $`x`$ over generators and $`y`$ over buses, @@ -133,10 +131,10 @@ $`x`$ over generators and $`y`$ over buses, \langle M x, y \rangle = \sum_{b} y_b \sum_{g} \mathbf{1}_R(g, b)\, x_g = \sum_{g} x_g \sum_{b} \mathbf{1}_R(g, b)\, y_b = \langle x, M^{\mathsf{T}} y \rangle, ``` -which is what the program means by calling `Pullback` the adjoint of -`GroupSum`. A bare relation has the same matrix without the functional claim. -A column of $`M`$ may hold several ones, so the sum fans out and there is no -read. That is why `at` through a bare relation is refused. +which is why the program's `Lookup` is the same `Join` as its `GroupSum`, read +without the group-by. A bare relation has the same matrix without the +functional claim. A column of $`M`$ may hold several ones, so the sum fans out +and no group is one row. That is why `at` through a bare relation is refused. ## The join and the aggregate @@ -146,21 +144,22 @@ Multiplying a matrix stored that way by a vector takes two steps. 1. **Pair each row of the table with each row of the operand that agrees on every shared index**, here $`g`$ and $`t`$. That is an inner equi-join on - the consumed and joined columns. Each pair is one nonzero product - $`1 \cdot p_{t,g}`$, relabelled by the produced column $`z`$. + the columns joined on. Each pair is one nonzero product + $`1 \cdot p_{t,g}`$, relabelled by the column grouped by, $`z`$. 2. **Add the pairs that agree on the free indices**, here $`(t, z)`$. That is a group-by on the result's dimensions with a sum. -lpspec, the reference engine, runs exactly these two steps. `walk_join` in +lpspec, the reference engine, runs exactly these two steps. `join_relation` in `src/lpspec/relational/engines/polars/relations.py` is step 1. It runs one inner -join on the consumed and joined dimensions. A select then drops the consumed -dimensions and renames the landing column to the produced one. Step 2 is the -terminal aggregate in `assembly.py`, a `group_by` over the row's coordinates -with `sum`, run once per constraint after every term has landed. Until then a -sum of linear terms is a list of terms, and adding is concatenation. `at` runs -step 1 against the same table and needs no step 2: single-valued means no two -pairs land on one coordinate. Where several $`(g, t)`$ share one $`z`$ the -join fans out, which is a column of $`M^{\mathsf{T}}`$ holding several ones. +join on the dimensions joined on. A select then drops the dimensions not +grouped by and renames the landing column to the dimension grouped by. Step 2 +is the terminal aggregate in `assembly.py`, a `group_by` over the row's +coordinates with `sum`, run once per constraint after every term has landed. +Until then a sum of linear terms is a list of terms, and adding is +concatenation. `at` runs step 1 against the same table and needs no step 2: +one row per group means no two pairs land on one coordinate. Where several +$`(g, t)`$ share one $`z`$ the join fans out, which is a column of +$`M^{\mathsf{T}}`$ holding several ones. So the two pictures are one. The join is the multiplication by an entry of $`\mathbf{1}_R`$, which is $`1`$ or absent. The group-by is the $`\sum`$. @@ -170,7 +169,7 @@ $`\mathbf{1}_R`$, which is $`1`$ or absent. The group-by is the $`\sum`$. Two positions of the formula have no variable to build, and there the model a consumer builds is not the matrix. -- **An empty fibre is the empty sum.** A zone no generator maps to at $`t`$ has +- **An empty group is the empty sum.** A zone no generator maps to at $`t`$ has $`y_{t,z} = 0`$, and the row reads $`0 \le \mathrm{zone\_cap}_{t,z}`$. It names no variable, so an engine [does not build it](../reference/language/absence.md#rows-with-no-variable-terms) and reports diff --git a/docs/about/what-counts-as-language.md b/docs/about/what-counts-as-language.md index 8e4ca26e..24a51246 100644 --- a/docs/about/what-counts-as-language.md +++ b/docs/about/what-counts-as-language.md @@ -29,8 +29,8 @@ Four rules follow from the test: in the renderer. - The set of operators is fixed. A tool cannot add a `roll` that the others do not know. -- Each operator has one rule for the dimensions it produces. `sum(p, by=gen_bus, over=generator, into=bus)` - lands on `bus` for every program. +- Each operator has one rule for the dimensions of its result. `sum(p, by=gen_bus, over=generator, into=bus)` + groups by `bus` for every program. - Degree is decided when the file loads. Whether `x * y` is allowed does not depend on which engine builds the model. diff --git a/docs/contributing.md b/docs/contributing.md index 2a1ea510..6e9d5fe3 100644 --- a/docs/contributing.md +++ b/docs/contributing.md @@ -110,8 +110,8 @@ two nodes, so the file's spelling cannot decide the name. | File verb | Node | What the node names | | ------------------ | ----------- | ------------------------------ | | `sum(over=)` | `Sum` | dims removed from the result | -| `sum(by=)` | `GroupSum` | a sum through a relation | -| `at(by=)` | `Pullback` | a read through a relation | +| `sum(by=)` | `GroupSum` | a join and group-by | +| `at(by=)` | `Lookup` | a join with no group-by | | `shift(along=)` | `Translate` | a re-index along one dimension | | `sum_back(along=)` | `WindowSum` | a sum over a trailing window | diff --git a/docs/examples/operators.md b/docs/examples/operators.md index 1679d0fa..90167717 100644 --- a/docs/examples/operators.md +++ b/docs/examples/operators.md @@ -112,8 +112,8 @@ $`\sum_{g \in \mathcal{G} \,:\, \mathrm{gen\_bus}(g) = b} p_{t,g} \le \mathrm{li ```yaml description: >- A call that names its ends — `sum(array, by=relation, over=a, into=b)` - consumes column `a` and lands on column `b`, and the other key column is - joined on, so each zone's total is taken per period. + joins on column `a` and groups by column `b`, and the other key column is + joined on and kept, so each zone's total is taken per period. dimensions: generator: { dtype: str } @@ -148,8 +148,8 @@ $`\sum_{g \in \mathcal{G} \,:\, \mathrm{zone\_of}(g,\ e) = z} p_{g,e} \ge \mathr ```yaml description: >- A call with several columns at each end — `sum(array, by=relation, over=[a, …], into=[b, …])` - consumes both key columns at once and lands on the product of both value - columns in one join. + joins on both key columns at once and groups by both value columns in one + join. dimensions: generator: { dtype: str } @@ -184,7 +184,7 @@ $`\sum_{g \in \mathcal{G},\ e \in \mathcal{E} \,:\, \mathrm{slot\_of.bus}(g,\ e) ```yaml description: >- - The adjoint of the membership reduction — `at(array, by=relation, over=a, into=b)` reads one + The same join with no group-by — `at(array, by=relation, over=a, into=b)` reads one coarse value once per fine label pointing at it. dimensions: diff --git a/docs/reference/language/expressions.md b/docs/reference/language/expressions.md index 15bf21e9..44ec6ad8 100644 --- a/docs/reference/language/expressions.md +++ b/docs/reference/language/expressions.md @@ -82,18 +82,18 @@ after a variable. The objective has no name at all. The dimension set of every expression is known before any data binds: -| Node | Dim set | Error | -| -------------------------------- | --------------------------------- | -------------------------------------------------------------------------------- | -| number | `{}` | | -| parameter / variable | its `dims` | | -| `-x`, `+x` | `dims(x)` | | -| `a + b`, `a * b`, `a / b` | `dims(a) ∪ dims(b)` | | -| `sum(x)` | `{}` | error if `dims(x)` is already empty | -| `sum(x, over=d)` | `dims(x) − {d}` | error if `d ∉ dims(x)` | -| `sum(x, by=l, over=a, into=b)` | `(dims(x) − consumed) ∪ produced` | the refusals under [how a relation is used](relations.md#how-a-relation-is-used) | -| `at(x, by=l, over=a, into=b)` | `(dims(x) − consumed) ∪ produced` | the same | -| `shift(x, along=d, offset=n)` | `dims(x)` | error if `d ∉ dims(x)` | -| `sum_back(x, along=d, window=n)` | `dims(x)` | error if `d ∉ dims(x)` | +| Node | Dim set | Error | +| -------------------------------- | ------------------------------ | -------------------------------------------------------------------------------- | +| number | `{}` | | +| parameter / variable | its `dims` | | +| `-x`, `+x` | `dims(x)` | | +| `a + b`, `a * b`, `a / b` | `dims(a) ∪ dims(b)` | | +| `sum(x)` | `{}` | error if `dims(x)` is already empty | +| `sum(x, over=d)` | `dims(x) − {d}` | error if `d ∉ dims(x)` | +| `sum(x, by=l, over=a, into=b)` | `(dims(x) − joined) ∪ grouped` | the refusals under [how a relation is used](relations.md#how-a-relation-is-used) | +| `at(x, by=l, over=a, into=b)` | `(dims(x) − joined) ∪ grouped` | the same | +| `shift(x, along=d, offset=n)` | `dims(x)` | error if `d ∉ dims(x)` | +| `sum_back(x, along=d, window=n)` | `dims(x)` | error if `d ∉ dims(x)` | A binary operator takes the **union** of the two dimension sets, so an outer product is allowed. The declaration's own dimensions are its **frame**, and a diff --git a/docs/reference/language/operators.md b/docs/reference/language/operators.md index eeec1f3d..870da5a1 100644 --- a/docs/reference/language/operators.md +++ b/docs/reference/language/operators.md @@ -15,7 +15,7 @@ in a reported expression, are all of them. A composition of them goes in | `sum(array)` | Every dimension that `array` carries collapses. The result is a scalar | | `sum(array, over=dim)` | `dim` collapses. `array` must carry `dim` | | `sum(array, by=relation, over=a, into=b)` | Column `a` collapses onto column `b`. The other key columns are joined on, so the array carries them and the result keeps them | -| `sum(array, by=relation, over=[a, …], into=[b, …])` | The same with several columns on either side: consumed together, landed on a product | +| `sum(array, by=relation, over=[a, …], into=[b, …])` | The same with several columns on either side: joined on together, grouped by a product | | `at(array, by=relation, over=a, into=b)` | Column `a` is replaced by column `b`, one value per coordinate. Either may be a list | | `shift(array, along=dim, offset=n)` | The value `n` positions earlier along `dim`. The vacated edge is **absent** | | `shift(array, along=dim, offset=n, edge='wrap')` | The value `n` positions earlier, counted cyclically, so nothing is vacated | @@ -41,8 +41,8 @@ result is a scalar. An operand that is already scalar, and an `over=` naming a dimension the operand does not carry, are both errors. -`sum(x, by=l, over=a, into=b)` sums through a [relation](relations.md), -consuming column `a` and landing the result on column `b`. A nodal balance is +`sum(x, by=l, over=a, into=b)` sums through a [relation](relations.md): it +joins on column `a` and groups by column `b`. A nodal balance is one `sum(by=)` per kind of component: ```yaml @@ -78,9 +78,10 @@ is null belongs to no group. ## `at` -`at(x, by=l, over=a, into=b)` reads the relation the other way. It consumes a -value column and produces the key, so it reads one coarse value once for each -fine label that points at it ([reads](relations.md#aggregates-and-reads)). +`at(x, by=l, over=a, into=b)` joins the relation the other way, with no +group-by. It joins on a value column and groups by the key, so every group is +one row, and it reads one coarse value once for each fine label that points at +it ([joins](relations.md#joins-and-group-bys)). `at` reads a variable as readily as a parameter. One decision taken per bus, read once by every line that touches the bus, is `at(decision, by=line_bus, over=bus, into=line)`. diff --git a/docs/reference/language/relations.md b/docs/reference/language/relations.md index 63ab8950..4e7480b7 100644 --- a/docs/reference/language/relations.md +++ b/docs/reference/language/relations.md @@ -72,15 +72,16 @@ column per declared column, named after it. ## How a relation is used -The declaration fixes no direction. A call names the columns it reads, and a -key column it names at neither end is **joined on**. +The declaration fixes no direction. A call **joins** the operand to the +relation on the columns they share, and **groups** what the join produces. The +call names the columns that decide both. -| kind | what it does | written as | -| --------- | ------------------------------------------- | ----------------------------------------------------- | -| aggregate | many rows of the operand collapse onto one | `sum(x, by=l, over=a, into=b)` | -| read | one row's value becomes a coordinate | `at(x, by=l, over=a, into=b)` | -| partition | the frame stays, and its rows are grouped | `shift`, `sum_back`, `position` with `by=l, within=c` | -| test | a row's presence keeps or cuts a coordinate | the relation's name in a `where` | +| kind | what it does | written as | +| ----------------- | ---------------------------------------------------------------------------- | ----------------------------------------------------- | +| join and group by | rows of the operand match rows of the relation, and each group is added up | `sum(x, by=l, over=a, into=b)` | +| join | each row of the operand matches one row of the relation, and reads its value | `at(x, by=l, over=a, into=b)` | +| partition | the frame stays, and its rows are grouped | `shift`, `sum_back`, `position` with `by=l, within=c` | +| test | a row's presence keeps or cuts a coordinate | the relation's name in a `where` | Four rules hold for every use: @@ -90,37 +91,50 @@ Four rules hold for every use: 3. **The key is fixed.** To change it, declare a new relation. 4. **A dimension the relation does not name passes through** to the result. -### Aggregates and reads - -`over=` names the columns consumed and `into=` the columns produced, and either -may be a list. With `zone_of: { key: [generator, period], values: zone }`, the -bare `connection: { key: [generator, bus] }` and `p` over `[generator, period]`: - -| call | consumes | joins on | produces | result | -| -------------------------------------------------- | ----------- | ----------- | ----------- | --------------------- | -| `sum(p, by=zone_of, over=generator, into=zone)` | `generator` | `period` | `zone` | `[zone, period]` | -| `sum(p, by=zone_of, over=period, into=zone)` | `period` | `generator` | `zone` | `[generator, zone]` | -| `sum(p, by=connection, over=generator, into=bus)` | `generator` | nothing | `bus` | `[bus, period]` | -| `at(price, by=zone_of, over=zone, into=generator)` | `zone` | `period` | `generator` | `[generator, period]` | - -- **The result is the operand, less the consumed dimensions, plus the produced - ones.** The operand carries every dimension consumed or joined on, and none - that the call lands on. `sum(load * p, by=gen_bus, over=generator, into=bus)` - is refused; write `load * sum(p, by=gen_bus, over=generator, into=bus)`. -- **A read finds one row per coordinate, and a sum finds many.** The columns a - read lands on and joins on hold the whole key. A sum leaves a key column out, - and each call is refused in the other's case: +### Joins and group-bys + +Each column of the relation is either **joined on** or not, and either +**grouped by** or not. `over=` names the columns joined on and not grouped by. +`into=` names the columns grouped by and not joined on. Either may be a list. + +| column of the relation | joined on | grouped by | in the result | +| ----------------------------- | --------- | ---------- | ------------- | +| `over=` | yes | no | leaves | +| `into=` | no | yes | arrives | +| a key column the call omits | yes | yes | stays | +| a value column the call omits | no | no | is not read | + +With `zone_of: { key: [generator, period], values: zone }`, the bare +`connection: { key: [generator, bus] }` and `p` over `[generator, period]`: + +| call | joins on | groups by | result | +| -------------------------------------------------- | ------------------- | ------------------- | --------------------- | +| `sum(p, by=zone_of, over=generator, into=zone)` | `generator, period` | `zone, period` | `[zone, period]` | +| `sum(p, by=zone_of, over=period, into=zone)` | `period, generator` | `zone, generator` | `[generator, zone]` | +| `sum(p, by=connection, over=generator, into=bus)` | `generator` | `bus` | `[bus, period]` | +| `at(price, by=zone_of, over=zone, into=generator)` | `zone, period` | `generator, period` | `[generator, period]` | + +- **The result is the operand, less the columns joined on, plus the columns + grouped by.** The operand carries every column joined on, and none that the + call groups by: a column the operand carries would be joined on, not grouped + by. `sum(load * p, by=gen_bus, over=generator, into=bus)` is refused; write + `load * sum(p, by=gen_bus, over=generator, into=bus)`. +- **A sum adds up several rows per group, and a read finds one.** Where the + columns grouped by hold the whole key, every group is one row. That is a join + with no group-by, which is `at`. A `sum` there is refused toward `at`, and an + `at` whose groups hold several rows is refused toward `sum`: ```text - Constraint 'cap': sum(by=zone_of): this sum lands on the key ['generator', 'period'], so each coordinate has one term and nothing is added up — that is a read, which is at()'s. Write at(..., by=zone_of, over=['zone'], into=['generator']), or sum toward a value column. + Constraint 'cap': sum(by=zone_of): the columns this sum groups by, ['generator', 'period'], hold the whole key ['generator', 'period'], so every group is one row and nothing is added up — that is a join with no group-by, which is at()'s. Write at(..., by=zone_of, over=['zone'], into=['generator']), or sum toward a value column. ``` -- **A sum consumes at least one key column, and lands on any column it does not - consume.** Either end may name a value column beside a key one. `connection` - has only key columns, and the sum above consumes one and lands on the other. -- **A read consumes value columns, and lands on the key.** Its result carries - every key column — named in `into=`, or joined on — and whatever else the - operand carries that the read does not consume. +- **A sum joins on at least one key column, and groups by any column it does + not join on.** Either end may name a value column beside a key one. + `connection` has only key columns, and the sum above joins on one and groups + by the other. +- **A read joins on value columns, and groups by the key.** Its result carries + every key column, named in `into=` or left unnamed, and whatever else the + operand carries that the read does not join on. - **`over=` and `into=` name different columns**, and neither names two columns over one dimension. @@ -128,7 +142,8 @@ bare `connection: { key: [generator, bus] }` and `p` over `[generator, period]`: `shift(x, along=d, by=l, within=c)`, `sum_back(x, along=d, by=l, within=c)` and `position(d, by=l, within=c)` step along the key column over `d`, join on -the other key columns, and group by the value columns `within=` names. The +the other key columns, and partition the rows by the value columns `within=` +names. The frame does not change. `within=` is written whenever `by=` is. It may name two columns over one dimension, may not name a key column, and a bare relation partitions nothing. diff --git a/docs/reference/notation.md b/docs/reference/notation.md index 9ccdb9ff..02598cb2 100644 --- a/docs/reference/notation.md +++ b/docs/reference/notation.md @@ -51,7 +51,7 @@ relations: zone_of: { key: bus, values: zone } area_of: { key: bus, values: zone } # a second map into the same set, to compare against season_of: { key: snapshot, values: season } - gen_zone: { key: [generator, snapshot], values: zone } # a map keyed by two dimensions: a call consumes one and joins on the other + gen_zone: { key: [generator, snapshot], values: zone } # a map keyed by two dimensions: a call sums one away and joins on the other rep_of: { key: snapshot, values: { rep: snapshot } } # a map into its own dimension: the representative snapshot connection: { key: [generator, bus] } # a bare relation, with no value columns: many-to-many, read only by sum with both ends named gen_bt: { key: generator, values: [bus, technology] } # one table with two value columns, read to both at once @@ -366,12 +366,12 @@ seasonal_window: \sum_{t' \in \mathcal{T} \,:\, 0 \le t -^{\mathrm{season\_of}(t)} t' < 3} \mathit{on}_{t',g} \le \mathit{units}_{g} \qquad \forall\, t \in \mathcal{T},\ g \in \mathcal{G} ``` -#### `pullback` +#### `lookup` at(), which re-indexes through a relation instead of an offset ```yaml -pullback: +lookup: dims: [snapshot, bus] expression: spill <= at(zone_cap, by=zone_of, over=zone, into=bus) ``` @@ -394,12 +394,12 @@ grouped_once: \sum_{g \in \mathcal{G} \,:\, \mathrm{gen\_bt.bus}(g) = b \wedge \mathrm{gen\_bt.technology}(g) = e} p_{t,g} \le \mathrm{tech\_cap}_{b,e} \qquad \forall\, t \in \mathcal{T},\ b \in \mathcal{B},\ e \in \mathcal{E} ``` -#### `pulled_back_once` +#### `looked_up_once` -its adjoint, reading one slot through two columns of one table +the same table joined the other way, reading one slot through two columns ```yaml -pulled_back_once: +looked_up_once: dims: [generator] expression: units <= at(tech_cap, by=gen_bt, over=[bus, technology], into=generator) ``` @@ -508,12 +508,12 @@ zonal_membership: \sum_{g \in \mathcal{G} \,:\, \mathrm{gen\_zone}(g,\ t) \text{ is defined}} \mathit{units}_{g} \le \mathrm{budget} \qquad \forall\, t \in \mathcal{T} ``` -#### `zonal_pullback` +#### `zonal_lookup` -its adjoint, reading the slot the row's own snapshot puts the generator in +the same table joined the other way, reading the slot the row's own snapshot puts the generator in ```yaml -zonal_pullback: +zonal_lookup: dims: [snapshot, generator] where: "gen_zone == 'north' AND position(generator, by=gen_zone, within=zone) == 0" expression: p <= at(spill * zone_cap, by=gen_zone, into=generator, over=zone) diff --git a/docs/reference/reading.md b/docs/reference/reading.md index ebb4bb3e..65c002f6 100644 --- a/docs/reference/reading.md +++ b/docs/reference/reading.md @@ -138,7 +138,7 @@ Every declared axis has an entry. A coupling that a `piecewise:` expansion introduced is named under the declaration the expansion emitted. - `coupled` names each declaration that ties the whole axis together: a sum - over the axis in a constraint, a grouping that consumes the axis, a wrapped + over the axis in a constraint, a grouping that sums the axis away, a wrapped shift, or a set. After the dash, each entry names the one change that would remove the tie. - `undecided` lists each read whose reach only the data can say, as a `Reach`: diff --git a/examples/operators/at.yaml b/examples/operators/at.yaml index 7d5eb2d2..3cf8adbe 100644 --- a/examples/operators/at.yaml +++ b/examples/operators/at.yaml @@ -3,7 +3,7 @@ # SPDX-License-Identifier: MIT description: >- - The adjoint of the membership reduction — `at(array, by=relation, over=a, into=b)` reads one + The same join with no group-by — `at(array, by=relation, over=a, into=b)` reads one coarse value once per fine label pointing at it. dimensions: diff --git a/examples/operators/sum_by_column_lists.yaml b/examples/operators/sum_by_column_lists.yaml index 98f70b7f..787ac2ac 100644 --- a/examples/operators/sum_by_column_lists.yaml +++ b/examples/operators/sum_by_column_lists.yaml @@ -4,8 +4,8 @@ description: >- A call with several columns at each end — `sum(array, by=relation, over=[a, …], into=[b, …])` - consumes both key columns at once and lands on the product of both value - columns in one join. + joins on both key columns at once and groups by both value columns in one + join. dimensions: generator: { dtype: str } diff --git a/examples/operators/sum_by_columns.yaml b/examples/operators/sum_by_columns.yaml index 683fcd3c..8c9fc3fd 100644 --- a/examples/operators/sum_by_columns.yaml +++ b/examples/operators/sum_by_columns.yaml @@ -4,8 +4,8 @@ description: >- A call that names its ends — `sum(array, by=relation, over=a, into=b)` - consumes column `a` and lands on column `b`, and the other key column is - joined on, so each zone's total is taken per period. + joins on column `a` and groups by column `b`, and the other key column is + joined on and kept, so each zone's total is taken per period. dimensions: generator: { dtype: str } diff --git a/src/math_spec/_expression_parser.py b/src/math_spec/_expression_parser.py index ba740563..32afc92b 100644 --- a/src/math_spec/_expression_parser.py +++ b/src/math_spec/_expression_parser.py @@ -23,7 +23,7 @@ if TYPE_CHECKING: from collections.abc import Callable, Iterator, Mapping - from math_spec.program import Direction, Partition, Predicate + from math_spec.program import Join, 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. @@ -132,13 +132,13 @@ def __str__(self) -> str: @dataclass(frozen=True) -class DirectionNode: - """A resolved ``by=`` on ``sum`` or ``at``: the relation, read in the :class:`Direction` the call names.""" +class JoinNode: + """A resolved ``by=`` on ``sum`` or ``at``: the relation, as the :class:`Join` the call names.""" - direction: Direction + join: Join def __str__(self) -> str: - return self.direction.name + return self.join.name @dataclass(frozen=True) @@ -283,7 +283,7 @@ def __str__(self) -> str: | ParameterNode | DualNode | DimensionNode - | DirectionNode + | JoinNode | PartitionNode | EdgeNode | KeywordNode @@ -338,7 +338,7 @@ def operand(node: ArithmeticNode) -> str: #: ``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 +KwargNode = DimensionNode | JoinNode | 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 diff --git a/src/math_spec/advice.py b/src/math_spec/advice.py index ee3406e5..aef61352 100644 --- a/src/math_spec/advice.py +++ b/src/math_spec/advice.py @@ -14,7 +14,7 @@ from math_spec.boundedness import unbounded_notes from math_spec.errors import Advice from math_spec.lowering import to_program -from math_spec.program import GroupSum, Pullback, walk +from math_spec.program import GroupSum, Lookup, walk if TYPE_CHECKING: from collections.abc import Mapping @@ -50,7 +50,7 @@ def _never_an_axis(program: Program) -> list[Advice]: reached: set[str] = set() for declaration in (*program.parameters.values(), *program.variables.values(), *program.constraints.values()): reached.update(declaration.dims) - reached |= _produced_axes(program) + reached |= _grouped_axes(program) reached |= {dim for lk in program.relations.values() for dim in lk.dims} return [ @@ -66,14 +66,14 @@ def _never_an_axis(program: Program) -> list[Advice]: ] -def _produced_axes(program: Program) -> set[str]: +def _grouped_axes(program: Program) -> set[str]: """The axes the expressions create beyond what any declaration indexes. - ``sum(by=)`` lands on its target and ``at()`` spreads onto its fine - dimension: either way, the dims the direction produces. + ``sum(by=)`` groups onto its target and ``at()`` spreads onto its fine + dimension: either way, the dims the join groups by and did not join on. """ axes: set[str] = set() for node in walk(*program.roots): - if isinstance(node, GroupSum | Pullback): - axes.update(node.direction.produced_dims) + if isinstance(node, GroupSum | Lookup): + axes.update(node.join.added_dims) return axes diff --git a/src/math_spec/boundedness.py b/src/math_spec/boundedness.py index 1102abe8..994d2646 100644 --- a/src/math_spec/boundedness.py +++ b/src/math_spec/boundedness.py @@ -25,11 +25,11 @@ Dual, Expression, GroupSum, + Lookup, Multiply, Negate, Parameter, Power, - Pullback, Sum, Translate, Variable, @@ -162,7 +162,7 @@ def _record_signs(node: Expression, sign: Sign, signs: dict[str, Sign]) -> None: _record_signs(node.base, None, signs) _record_signs(node.exponent, None, signs) return - if isinstance(node, Sum | GroupSum | Pullback | Translate | WindowSum | Cases): + if isinstance(node, Sum | GroupSum | Lookup | Translate | WindowSum | Cases): for child in children(node): _record_signs(child, sign, signs) return diff --git a/src/math_spec/dimensions.py b/src/math_spec/dimensions.py index 72db88a2..e8ba48d1 100644 --- a/src/math_spec/dimensions.py +++ b/src/math_spec/dimensions.py @@ -23,10 +23,10 @@ ComparisonNode, DefinitionNode, DimensionNode, - DirectionNode, DualNode, EdgeNode, FunctionCallNode, + JoinNode, KwargNode, NumberNode, ParameterNode, @@ -43,7 +43,7 @@ from math_spec.program import ( DimensionComparison, DimensionPosition, - Direction, + Join, Mask, ParameterComparison, ParameterDefined, @@ -137,7 +137,7 @@ def _dims_call(node: FunctionCallNode, schema: Spec, context: str) -> frozenset[ 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.""" + """``sum`` reduces a dim away, or joins a relation and groups: the dims joined on go, the dims grouped by arrive.""" by = node.kwargs.get('by') if by is None and 'over' not in node.kwargs: if not inner: @@ -148,41 +148,41 @@ def _sum_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spec, conte ) return frozenset() if by is None: - consumed = node.kwargs['over'] - assert isinstance(consumed, DimensionNode) - if consumed.name not in inner: + summed = node.kwargs['over'] + assert isinstance(summed, DimensionNode) + if summed.name not in inner: raise DimensionError( - _not_carried(context, f'sum(over={consumed.name})', inner, 'drop the sum, or fix the dim') + _not_carried(context, f'sum(over={summed.name})', inner, 'drop the sum, or fix the dim') ) - return inner - {consumed.name} + return inner - {summed.name} - assert isinstance(by, DirectionNode), 'resolution reads sum(by=) in a direction' - direction = by.direction - if missing := sorted(set(direction.consumed_dims) - inner): + assert isinstance(by, JoinNode), 'resolution reads sum(by=) as a join' + join = by.join + if missing := sorted(set(join.dropped_dims) - inner): raise DimensionError( _not_carried( context, - f'sum(by={direction.name}) consumes {missing}, the dims it reads from,', + f'sum(by={join.name}) joins on {missing} to sum it away,', inner, 'drop the sum, or fix the dim', ) ) - return _read_dims(f'sum(by={direction.name})', direction, inner, context) + return _join_dims(f'sum(by={join.name})', join, inner, context) def _at_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spec, context: str) -> frozenset[str]: - """``at`` is the adjoint of ``sum(by=)``: it consumes the dims a sum produces and produces the ones it consumes.""" + """``at`` is the join of ``sum(by=)`` with no group-by: it joins on the dims a sum groups by and groups by the ones a sum joins on.""" by = node.kwargs['by'] - assert isinstance(by, DirectionNode), 'resolution reads at(by=) in a direction' - direction = by.direction - if absent := sorted(set(direction.consumed_dims) - inner): + assert isinstance(by, JoinNode), 'resolution reads at(by=) as a join' + join = by.join + if absent := sorted(set(join.dropped_dims) - inner): raise DimensionError( - f'{context}: at(by={direction.name}) reads through ' + f'{context}: at(by={join.name}) joins on ' f'{absent}, which the expression does not carry (dims ' - f'{sorted(inner)}). A pullback needs the coarse dims to read *from* — ' - f'sum is the direction that produces them.' + f'{sorted(inner)}). A lookup joins the operand on the columns it reads at — ' + f'sum is the call that groups by them.' ) - return _read_dims(f'at(by={direction.name})', direction, inner, context) + return _join_dims(f'at(by={join.name})', join, inner, context) def _translation_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spec, context: str) -> frozenset[str]: @@ -208,48 +208,48 @@ def _translation_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spe return inner -def _read_dims(call: str, direction: Direction, inner: frozenset[str], context: str) -> frozenset[str]: - """The dims after a relation is read in *direction*: the consumed go, the produced arrive, the joined stay. +def _join_dims(call: str, join: Join, inner: frozenset[str], context: str) -> frozenset[str]: + """The dims after *join*: the operand's, less the dims joined on, plus the dims grouped by. - The dims a call lands on are its own to bring, so the operand does not - already carry one. Where it does, the call would tie the operand's axis to - the one it produces rather than adding it, and it reads the same either + A column grouped by and not joined on brings its dim, so the operand does + not already carry it. Where it does, the join would match the operand on + that column instead of grouping by it, and the call reads the same either way. A relation into its own dimension is not that case: there the dim - landed on is the dim just consumed, so every factor is read at the - coordinate the sum runs over, and nothing is tied. + grouped by is the dim just joined on and summed away, so every factor is + read at the coordinate the sum runs over, and nothing is matched twice. """ - consumed, produced = set(direction.consumed_dims), set(direction.produced_dims) - if clash := sorted((produced & inner) - consumed): + added, dropped = set(join.added_dims), set(join.dropped_dims) + if clash := sorted((added & inner) - dropped): raise DimensionError( - f'{context}: {call} lands on {clash}, which the expression already carries.\n' - f'A call brings the dims it lands on, so that reading it tells you what it ' - f'adds. Move the factor carrying {clash} outside the operator, or read to a column ' - f'over another dimension.' + f'{context}: {call} groups by {clash}, which the expression already carries.\n' + f'A join on a column the operand carries matches it rather than grouping by it, so a ' + f'call brings the dims it groups by. Move the factor carrying {clash} outside the operator, ' + f'or group by a column over another dimension.' ) - _check_joined(call, direction, inner, context) - return (inner - consumed) | produced + _check_joined(call, join, inner, context) + return (inner - set(join.joined_dims)) | set(join.grouped_dims) -def _check_joined(call: str, use: Direction | Partition, inner: frozenset[str], context: str) -> None: - """The columns a call joins on are read at their dimensions, so the operand carries every one, each once. +def _check_joined(call: str, use: Join | Partition, inner: frozenset[str], context: str) -> None: + """The columns a call joins on are matched at their dimensions, so the operand carries every one, each once. - A joined dimension the call also consumes is the same ambiguity as two - joined columns over one dimension: the operand's one coordinate would - have to be read as both. A partition consumes nothing. + Two joined columns over one dimension would match the operand's one + coordinate twice: the ``over=`` column and an unnamed key column, or two + unnamed key columns. A partition joins on the key columns it does not + step along. """ dims = use.joined_dims if missing := sorted(set(dims) - inner): raise DimensionError( f'{context}: {call} joins on {missing} (columns {[r for r in use.joined if use.dim(r) in missing]} ' - f"of '{use.name}'), which the expression does not carry (dims {sorted(inner)}). A relation is " - f'read between two of its columns and joined at the others — index the operand by them, or ' - f'read it between different columns.' + f"of '{use.name}'), which the expression does not carry (dims {sorted(inner)}). A join matches " + f'the operand on every key column the call does not name — index the operand by them, or ' + f'name them in the call.' ) - consumed = use.consumed_dims if isinstance(use, Direction) else () - if twice := sorted({d for d in dims if dims.count(d) > 1 or d in consumed}): + if twice := sorted({d for d in dims if dims.count(d) > 1}): raise DimensionError( f"{context}: {call} joins '{use.name}' on {twice} through more than one column, and the operand " - f'carries each dimension once. Read between different columns, or use a relation whose joined ' + f'carries each dimension once. Join on distinct dimensions, or use a relation whose key ' f'columns are over distinct dimensions.' ) @@ -431,7 +431,7 @@ def _check_named_amount(node: FunctionCallNode, over: str, inner: frozenset[str] ) by = node.kwargs.get('by') groups = ( - frozenset(by.partition.dim(v) for v in by.partition.group) if isinstance(by, PartitionNode) else frozenset() + frozenset(by.partition.dim(v) for v in by.partition.grouped) if isinstance(by, PartitionNode) else frozenset() ) if stray := sorted(frozenset(declared.dims) - inner - groups): raise DimensionError( diff --git a/src/math_spec/exclusivity.py b/src/math_spec/exclusivity.py index 985fcf43..be50f6af 100644 --- a/src/math_spec/exclusivity.py +++ b/src/math_spec/exclusivity.py @@ -229,7 +229,7 @@ def _subject_of(node: TypedPredicate) -> Subject: case DimensionPosition(name=name, partition=partition): if partition is None: return Subject('rank', name) - return Subject('rank', name, partition.name, partition.group) + return Subject('rank', name, partition.name, partition.grouped) case RelationDefined(name=name) | RelationComparison(name=name): return Subject('relation', name) case RelationPairComparison(name=name, other=other): diff --git a/src/math_spec/lowering.py b/src/math_spec/lowering.py index 850ef854..b7d0d439 100644 --- a/src/math_spec/lowering.py +++ b/src/math_spec/lowering.py @@ -22,10 +22,10 @@ CasesNode, DefinitionNode, DimensionNode, - DirectionNode, DualNode, EdgeNode, FunctionCallNode, + JoinNode, KwargNode, NumberNode, ParameterNode, @@ -261,17 +261,17 @@ def sum(self, node: FunctionCallNode) -> program.Expression: 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) + summed = node.kwargs['over'] + assert isinstance(summed, DimensionNode), 'resolution refuses a over= that is not a dimension' + return program.Sum(operand, (summed.name,)) + assert isinstance(by_node, JoinNode), 'resolution reads sum(by=) as a join' + return program.GroupSum(operand, join=by_node.join) def at(self, node: FunctionCallNode) -> program.Expression: - """``at(x, by=relation)`` — the adjoint of :meth:`sum`'s ``by=`` form.""" + """``at(x, by=relation)`` — the join of :meth:`sum`'s ``by=`` form with no group-by.""" 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) + assert isinstance(by_node, JoinNode), 'resolution reads at(by=) as a join' + return program.Lookup(self.expr(node.args[0]), join=by_node.join) def sum_back(self, node: FunctionCallNode) -> program.Expression: """``sum_back(x, along=d, window=w)`` — a trailing window along one dimension. diff --git a/src/math_spec/operators.py b/src/math_spec/operators.py index 50dd48a7..b3db75d8 100644 --- a/src/math_spec/operators.py +++ b/src/math_spec/operators.py @@ -39,11 +39,11 @@ class Builtin: dimension_kwargs: tuple[str, ...] = () relation_kwargs: tuple[str, ...] = () #: Kwargs naming a column of the relation ``by=`` names — ``over=`` and - #: ``into=`` — which resolution folds into the direction it is read in. + #: ``into=`` — which resolution folds into the join the call makes. role_kwargs: tuple[str, ...] = () #: Kwargs naming a dimension on their own and a column of the relation where #: ``by=`` names one. ``sum(x, over=generator)`` reduces the dimension - #: away; ``sum(x, by=l, over=c)`` names the column the call consumes. + #: away; ``sum(x, by=l, over=c)`` names the column the call joins on and sums away. #: One meaning — what leaves the frame — read in the namespace ``by=`` #: decides. dimension_or_role_kwargs: tuple[str, ...] = () diff --git a/src/math_spec/program.py b/src/math_spec/program.py index a3d1aa94..203fb998 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -55,7 +55,6 @@ 'DimensionDeclaration', 'DimensionDtype', 'DimensionPosition', - 'Direction', 'Divide', 'Dual', 'Expression', @@ -65,7 +64,9 @@ 'Footprint', 'GroupSum', 'Increasing', + 'Join', 'LastOf', + 'Lookup', 'Mask', 'MaskOf', 'Multiply', @@ -85,7 +86,6 @@ 'Predicate', 'PredicateOperator', 'Program', - 'Pullback', 'QuadraticPosition', 'Reach', 'Region', @@ -246,29 +246,30 @@ class Sum: @dataclass(frozen=True) class GroupSum: - """Sum ``operand`` through a relation: the dims ``direction`` consumes go, the dims it produces arrive, the dims it joins on stay. + """Sum ``operand`` through a relation: a join on ``join.joined``, then a group-by on ``join.grouped`` with a sum. - The join keys on the consumed columns and every joined column, and the - operand carries every dim consumed or joined on. + The operand carries every dim joined on. The result drops the dims joined + on and not grouped by, keeps the ones both joined on and grouped by, and + gains the ones grouped by and not joined on. """ operand: Expression - direction: Direction + join: Join @dataclass(frozen=True) -class Pullback: - """Read ``operand`` through a relation — the adjoint of :class:`GroupSum`. - - The dims ``direction`` consumes go and the dims it produces arrive, one - value per coordinate because the read takes value columns at a key the - result fixes, which the loader checks. The join fans out, many - produced tuples sharing one consumed tuple — at each coordinate of the - joined columns, which the operand carries and the result keeps. +class Lookup: + """Read ``operand`` through a relation: the join of :class:`GroupSum` with no group-by. + + The grouped columns hold the relation's whole key, which the loader + checks, so each row of the result meets one row of the relation and + reads one value. The join fans out where several key tuples share the + values joined on, at each coordinate of the columns both joined on and + grouped by. """ operand: Expression - direction: Direction + join: Join @dataclass(frozen=True) @@ -372,7 +373,7 @@ class Cases: | Divide | Sum | GroupSum - | Pullback + | Lookup | Translate | WindowSum | Cases @@ -391,7 +392,7 @@ def fan_in(expression: Expression) -> FanIn: return 'one-to-many' if isinstance( expression, - (Constant, Parameter, Variable, Dual, Negate, Add, Multiply, Power, Divide, Pullback, Translate, Cases), + (Constant, Parameter, Variable, Dual, Negate, Add, Multiply, Power, Divide, Lookup, Translate, Cases), ): return 'one-to-one' assert_never(expression) @@ -407,7 +408,7 @@ def children(expression: Expression) -> tuple[Expression, ...]: return (expression.numerator, expression.divisor) if isinstance(expression, Power): return (expression.base, expression.exponent) - if isinstance(expression, (Sum, GroupSum, Pullback, Translate, WindowSum)): + if isinstance(expression, (Sum, GroupSum, Lookup, Translate, WindowSum)): return (expression.operand,) if isinstance(expression, Cases): return tuple(region.value for region in expression.regions) @@ -461,59 +462,84 @@ def dim(self, role: str) -> str: @dataclass(frozen=True) -class Direction: - """One relation as one call reads it — which columns are consumed, which produced, which joined on. +class Join: + """One relation as one call joins it: the columns joined on, and the columns grouped by. The declaration fixes no direction; the call does, and this is the one it named. ``name`` is the relation's, as :attr:`Program.relations` keys it. - ``consumed``, ``produced`` and ``joined`` are *roles* — column names of - ``relation``, which binds every role to its dimension and names the key. - ``joined`` is the key roles the call did not name (every role, for a bare - relation): the join keys on them, and a value role left unnamed is not - read. + ``joined`` and ``grouped`` are *roles* — column names of ``relation``, + which binds every role to its dimension and names the key. ``joined`` is + every column the join matches the operand on: the ``over=`` columns, then + every key column the call did not name. ``grouped`` is every column of the + relation the result keeps: the ``into=`` columns, then the same unnamed + key columns. A column in neither is not read, so a relation may gain a + value column without changing what a call means. + + A column both joined on and grouped by stays in the frame. One joined on + and not grouped by is :attr:`dropped`, and one grouped by and not joined + on is :attr:`added`, so the frame after the call is the operand's dims, + less the dims joined on, plus the dims grouped by. """ name: str relation: RelationDeclaration - consumed: tuple[str, ...] - produced: tuple[str, ...] joined: tuple[str, ...] + grouped: tuple[str, ...] def dim(self, role: str) -> str: """The dimension *role* is bound to.""" return self.relation.dim(role) @property - def consumed_dims(self) -> tuple[str, ...]: - return tuple(self.dim(role) for role in self.consumed) + def dropped(self) -> tuple[str, ...]: + """The roles joined on and not grouped by: what the call sums away or looks up.""" + return tuple(role for role in self.joined if role not in self.grouped) + + @property + def added(self) -> tuple[str, ...]: + """The roles grouped by and not joined on: what the call brings into the frame.""" + return tuple(role for role in self.grouped if role not in self.joined) @property - def produced_dims(self) -> tuple[str, ...]: - return tuple(self.dim(role) for role in self.produced) + def kept(self) -> tuple[str, ...]: + """The roles both joined on and grouped by: the key columns the call did not name.""" + return tuple(role for role in self.joined if role in self.grouped) @property def joined_dims(self) -> tuple[str, ...]: return tuple(self.dim(role) for role in self.joined) + @property + def grouped_dims(self) -> tuple[str, ...]: + return tuple(self.dim(role) for role in self.grouped) + + @property + def dropped_dims(self) -> tuple[str, ...]: + return tuple(self.dim(role) for role in self.dropped) + + @property + def added_dims(self) -> tuple[str, ...]: + return tuple(self.dim(role) for role in self.added) + @dataclass(frozen=True) class Partition: - """One relation as a partition steps along it — the key column stepped along, the group columns, and the key columns joined on. + """One relation as a partition steps along it — the key column stepped along, the columns partitioned by, and the key columns joined on. ``name`` is the relation's, as :attr:`Program.relations` keys it. - ``along``, ``group`` and ``joined`` are *roles* — column names of + ``along``, ``grouped`` and ``joined`` are *roles* — column names of ``relation``, which binds every role to its dimension and names the key. ``along`` is the one key column over the dimension stepped along, and - the frame keeps it. ``group`` is the value columns ``within=`` named, - read at the row's key. ``joined`` is the other key columns, whose - dimensions the frame carries. Nothing is consumed and nothing is - produced: the frame does not change. + the frame keeps it. ``grouped`` is the value columns ``within=`` named, + read at the row's key: the partition's group. ``joined`` is the other key + columns, whose dimensions the frame carries. No column is dropped and + none is added: the frame does not change. """ name: str relation: RelationDeclaration along: str - group: tuple[str, ...] + grouped: tuple[str, ...] joined: tuple[str, ...] def dim(self, role: str) -> str: diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index efbb2a8f..d5f416e5 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -25,10 +25,10 @@ ComparisonNode, DefinitionNode, DimensionNode, - DirectionNode, DualNode, EdgeNode, FunctionCallNode, + JoinNode, KeywordNode, KwargNode, NameListNode, @@ -68,7 +68,7 @@ BooleanLiteral, DimensionComparison, DimensionPosition, - Direction, + Join, Mask, Not, Or, @@ -223,7 +223,7 @@ class Resolved: objective: The objective's expression, ``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 + copy, which every :class:`~math_spec.program.Join` and :class:`~math_spec.program.Partition` in the trees holds. """ @@ -579,12 +579,13 @@ def _relation_ref( 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. + """An operator's ``by=``, with the ``over=`` and ``into=`` that say how the relation is joined and grouped. 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 + than dims: ``over=`` the columns joined on and summed away, ``into=`` + the columns grouped by, every other key column joined on and kept. 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. """ @@ -616,8 +617,8 @@ def _relation_ref( return value if partition is None else PartitionNode(partition) 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) + join = self._join(name, operator, named['over'], named['into']) + return value if join is None else JoinNode(join) 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.""" @@ -630,22 +631,22 @@ def _role_name(self, value: ArithmeticNode, operator: str, key: str) -> tuple[st ) return None - def _direction( + def _join( self, name: str, operator: str, from_roles: tuple[str, ...], into_roles: tuple[str, ...], - ) -> Direction | None: - """Which direction ``sum`` or ``at`` reads relation *name* in, between the columns the call named. + ) -> Join | None: + """How ``sum`` or ``at`` joins relation *name*, between the columns the call named. Both ends arrive written: the call shape refuses a call that leaves one unsaid, so that a relation may gain a value column without - changing what this call means. ``at`` needs the read single-valued - and ``sum`` needs it not: a sum that lands on the key has one term - per coordinate and adds up nothing, which is a read, so it is - refused toward ``at``. A read lands on key columns and nothing else, - because a column outside the key is one no coordinate of the read + changing what this call means. A lookup needs every group to be one + row and a sum needs it not: a sum whose grouped columns hold the whole + key has one row per group and adds up nothing, which is a lookup, so + it is refused toward ``at``. A lookup groups by key columns and nothing + else, because a column outside the key is one no row of the join fixes. """ ns, context = self.ns, self.context @@ -674,29 +675,29 @@ def _direction( if not forward and (outside := [r for r in into_roles if r not in shape.key]): self.errors.append( f"{context}: {call}: into={list(into_roles)} names {outside}, which the key of '{name}' does not " - f'hold. A read lands on the key it reads at, {list(shape.key)}, and a column outside that key ' - f'arrives as a dimension the read never fixes. Land on the key, or sum toward {outside}.' + f'hold. A lookup reads one row per key, {list(shape.key)}, and a column outside that key ' + f'arrives as a dimension no row of the join fixes. Group by the key, or sum toward {outside}.' ) return None - joined = tuple(r for r in shape.key if r not in from_roles and r not in into_roles) - single_valued = set(shape.key) <= {*into_roles, *joined} - direction = Direction(name, shape, from_roles, into_roles, joined) - if not forward and not single_valued: + kept = tuple(r for r in shape.key if r not in from_roles and r not in into_roles) + join = Join(name, shape, (*from_roles, *kept), (*into_roles, *kept)) + one_row_per_group = set(shape.key) <= set(join.grouped) + if not forward and not one_row_per_group: self.errors.append( - f"{context}: {call}: at reads one value per coordinate, and '{name}' is not single-valued in " - f'{list(from_roles)} at the columns the call lands on ({[*into_roles, *joined]}) — its key is ' - f'{list(shape.key)}. Key the table by the columns the call lands on, or read the other way.' + f"{context}: {call}: at reads one row per group, and grouping '{name}' by {list(join.grouped)} " + f'leaves several rows in a group — its key is {list(shape.key)}. Key the table by the columns ' + f'the call groups by, or sum instead.' ) return None - if forward and single_valued: + if forward and one_row_per_group: self.errors.append( - f'{context}: {call}: this sum lands on the key {list(shape.key)}, so each coordinate has one ' - f"term and nothing is added up — that is a read, which is at()'s. Write " - f'at(..., by={name}, over={list(from_roles)}, into={list(into_roles)}), or sum toward ' - f'a value column.' + f'{context}: {call}: the columns this sum groups by, {list(join.grouped)}, hold the whole key ' + f'{list(shape.key)}, so every group is one row and nothing is added up — that is a join with no ' + f"group-by, which is at()'s. Write at(..., by={name}, over={list(from_roles)}, " + f'into={list(into_roles)}), or sum toward a value column.' ) return None - return direction + return join def _known_roles(self, name: str, call: str, roles: tuple[str, ...], kwarg: str) -> bool: """Whether every role *kwarg* names is a column of relation *name*, each once; the refusal otherwise.""" diff --git a/src/math_spec/separability.py b/src/math_spec/separability.py index 2c559b28..c61fca43 100644 --- a/src/math_spec/separability.py +++ b/src/math_spec/separability.py @@ -12,8 +12,8 @@ Cases, DimensionPosition, GroupSum, + Lookup, Mask, - Pullback, Reach, Separability, Sum, @@ -107,16 +107,16 @@ def waits_on(dimension: str, label: str, name: str, kind: Literal['offset', 'par f'sums over {dimension} — a rolling sum_back(window=n) windows, a total over the horizon does not', ) elif isinstance(node, GroupSum): - for dimension in node.direction.consumed_dims: + for dimension in node.join.dropped_dims: report( 'coupled', dimension, label, - f'groups {dimension} into {", ".join(node.direction.produced_dims)} — window that dimension instead, or cut only at the group edges', + f'groups {dimension} into {", ".join(node.join.added_dims)} — window that dimension instead, or cut only at the group edges', ) - elif isinstance(node, Pullback): - for dimension in node.direction.consumed_dims: - waits_on(dimension, label, node.direction.name, 'coordinate') + elif isinstance(node, Lookup): + for dimension in node.join.dropped_dims: + waits_on(dimension, label, node.join.name, 'coordinate') elif isinstance(node, (Translate, WindowSum)): dimension = node.along if node.wrap: diff --git a/src/math_spec/typesetting/walk.py b/src/math_spec/typesetting/walk.py index 32a2ac2a..90cad9b7 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -22,10 +22,10 @@ CasesNode, DefinitionNode, DimensionNode, - DirectionNode, DualNode, EdgeNode, FunctionCallNode, + JoinNode, KwargNode, NumberNode, ParameterNode, @@ -40,7 +40,7 @@ BooleanLiteral, DimensionComparison, DimensionPosition, - Direction, + Join, Mask, Not, Or, @@ -154,8 +154,8 @@ class _Context: walk: Walk offsets: dict[str, tuple[_Step, ...]] = field(default_factory=dict) - #: dim -> the rendered subscript that replaces its index, as ``at`` re-indexes a leaf. - pullbacks: dict[str, str] = field(default_factory=dict) + #: dim -> the rendered subscript that replaces its index, as ``at`` looks a leaf up through a relation. + lookups: dict[str, str] = field(default_factory=dict) #: Every dimension whose index is in use here — the frame, then one entry #: per reduction entered — so a reduction over one takes a fresh dummy. bound: tuple[str, ...] = () @@ -164,10 +164,10 @@ def translated(self, dim: str, step: _Step) -> _Context: steps = self.offsets.get(dim, ()) merged = steps[-1].merged(step) if steps else None steps = (*steps[:-1], merged) if merged is not None else (*steps, step) - return _Context(self.walk, {**self.offsets, dim: steps}, self.pullbacks, self.bound) + return _Context(self.walk, {**self.offsets, dim: steps}, self.lookups, self.bound) - def pulled_back(self, dim: str, rendered: str) -> _Context: - return _Context(self.walk, self.offsets, {**self.pullbacks, dim: rendered}, self.bound) + def looked_up(self, dim: str, rendered: str) -> _Context: + return _Context(self.walk, self.offsets, {**self.lookups, dim: rendered}, self.bound) def reducing(self, dim: str) -> tuple[str, _Context]: """A dummy index for a reduction over *dim*, and the context its body reads under. @@ -178,18 +178,18 @@ def reducing(self, dim: str) -> tuple[str, _Context]: """ primes = "'" * self.bound.count(dim) dummy = f'{self.walk.symbols.index[dim]}{primes}' - body = _Context(self.walk, self.offsets, {**self.pullbacks, dim: dummy}, (*self.bound, dim)) + body = _Context(self.walk, self.offsets, {**self.lookups, dim: dummy}, (*self.bound, dim)) return dummy, body def subscript(self, dim: str) -> str: - """The index for *dim* here: its pullback if it has one, then every translation. + """The index for *dim* here: its lookup if it has one, then every translation. - A pullback is a base like any other rather than a stopping point. + A lookup is a base like any other rather than a stopping point. ``at`` and ``shift`` both re-index the leaf and the leaf has one subscript, so a reading that showed only whichever ran last dropped the other operator out of the equation. """ - text = self.pullbacks.get(dim, self.walk.symbols.index[dim]) + text = self.lookups.get(dim, self.walk.symbols.index[dim]) translated = False for step in self.offsets.get(dim, ()): if step.by == 0: @@ -447,30 +447,30 @@ def _call(self, node: FunctionCallNode, ctx: _Context) -> tuple[str, int]: if node.name == 'at': by = node.kwargs['by'] - assert isinstance(by, DirectionNode) + assert isinstance(by, JoinNode) outer = ctx - direction = by.direction - at = {r: outer.subscript(direction.dim(r)) for r in (*direction.produced, *direction.joined)} - for read in direction.consumed: - ctx = ctx.pulled_back(direction.dim(read), self._relation_read(direction.name, at, read)) + join = by.join + at = {r: outer.subscript(join.dim(r)) for r in join.grouped} + for read in join.dropped: + ctx = ctx.looked_up(join.dim(read), self._relation_read(join.name, at, read)) return self._arithmetic(node.args[0], ctx) if (by := node.kwargs.get('by')) is not None: - assert isinstance(by, DirectionNode) - direction = by.direction + assert isinstance(by, JoinNode) + join = by.join dummies: dict[str, str] = {} inner = ctx - for d in direction.consumed_dims: + for d in join.dropped_dims: dummies[d], inner = inner.reducing(d) - conditions = list(self._grouping(direction, dummies, ctx)) + conditions = list(self._grouping(join, dummies, ctx)) domain = ( - f'{self.format.joined([self._membership(d, dummies[d]) for d in direction.consumed_dims], "")} ' + f'{self.format.joined([self._membership(d, dummies[d]) for d in join.dropped_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) + elif (summed := node.kwargs.get('over')) is not None: + assert isinstance(summed, DimensionNode) + dummy, inner = ctx.reducing(summed.name) + domain = self._membership(summed.name, dummy) else: memberships = [] inner = ctx @@ -480,23 +480,23 @@ def _call(self, node: FunctionCallNode, ctx: _Context) -> tuple[str, int]: domain = self.format.joined(memberships, '') return self.format.summation(domain, self._reduction_body(node.args[0], inner)), _PRECEDENCE['+'] - def _grouping(self, direction: Direction, dummies: Mapping[str, str], ctx: _Context) -> list[str]: - """The conditions a grouped sum's domain carries for one direction: what it fixes of the row it joins on. + def _grouping(self, join: Join, dummies: Mapping[str, str], ctx: _Context) -> list[str]: + """The conditions a grouped sum's domain carries for one join: what it fixes of the row it joins on. - A direction fixes its relation's key either way, so each value column - it fixes — consumed or produced, one lookup at the key either way — is - that column read there. One that fixes none of them asks only that the + A join fixes its relation's key either way, so each value column it + names — joined on or grouped by, one read at the key either way — is + that column read there. One that names none of them asks only that the row is there, because a value column it does not touch is not read, and a bare relation has no value column to read at all. """ at = { - **{r: dummies[direction.dim(r)] for r in direction.consumed}, - **{r: ctx.subscript(direction.dim(r)) for r in (*direction.joined, *direction.produced)}, + **{r: dummies[join.dim(r)] for r in join.dropped}, + **{r: ctx.subscript(join.dim(r)) for r in join.grouped}, } - fixed = [r for r in self.schema.relations[direction.name].value_roles if r in at] + fixed = [r for r in self.schema.relations[join.name].value_roles if r in at] if not fixed: - 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] + return [self._relation_row(join.name, at)] + return [f'{self._relation_read(join.name, at, r)} {self._op("equal")} {at[r]}' for r in fixed] def _group(self, by: ArithmeticNode | None, dim: str) -> str: """A ``by=`` as the superscript its translation operator carries. @@ -510,7 +510,7 @@ def _group(self, by: ArithmeticNode | None, dim: str) -> str: 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]) + return self._tuple([self._relation_read(partition.name, at, r) for r in partition.grouped]) def _width(self, node: ArithmeticNode) -> str: """``sum_back``'s ``window=``: a number, or a parameter's own symbol. @@ -595,7 +595,7 @@ def _where(self, node: Predicate, ctx: _Context) -> tuple[str, int]: grouping = ( None if node.partition is None - else self._tuple([self._value_read(node.partition.name, c, ctx) for c in node.partition.group]) + else self._tuple([self._value_read(node.partition.name, c, ctx) for c in node.partition.grouped]) ) place = self._position(ctx.subscript(node.name), grouping) ordinal = self._ordinal(node.name, node.position, grouping) diff --git a/tests/test_advice.py b/tests/test_advice.py index 37232fd1..49019c22 100644 --- a/tests/test_advice.py +++ b/tests/test_advice.py @@ -52,7 +52,7 @@ def test_a_dimension_nothing_reaches_is_named(): ) def test_a_dimension_something_reaches_is_in_use(patch): assert not advice(override(TARGET_ONLY, **patch)), ( - 'a dimension a relation targets, a declaration indexes or a grouping lands on is in use' + 'a dimension a relation targets, a declaration indexes or a grouping groups by is in use' ) diff --git a/tests/test_dimensions.py b/tests/test_dimensions.py index 8e781381..30be7970 100644 --- a/tests/test_dimensions.py +++ b/tests/test_dimensions.py @@ -111,7 +111,7 @@ def namespace() -> Namespace: pytest.param( 'sum(p, by=gen_zone, over=generator, into=zone)', {'snapshot', 'zone'}, - id='a-two-key-relation-consumes-the-key-it-names-and-keeps-the-other', + id='a-two-key-relation-sums-away-the-key-it-names-and-keeps-the-other', ), pytest.param( 'sum(p, by=gen_zone, over=snapshot, into=zone)', @@ -121,7 +121,7 @@ def namespace() -> Namespace: pytest.param( 'at(zone_load, by=gen_zone, into=generator, over=zone)', {'snapshot', 'generator'}, - id='its-pullback-keeps-the-joined-key-too', + id='its-lookup-keeps-the-joined-key-too', ), pytest.param( "shift(p, along=generator, offset=1, edge='wrap', by=gen_zone, within=zone)", @@ -161,7 +161,7 @@ def namespace() -> Namespace: pytest.param( 'sum(p, by=gen_zone, over=[generator, snapshot], into=zone)', {'zone'}, - id='a-from-list-consumes-two-key-columns-at-once', + id='an-over-list-joins-on-two-key-columns-at-once', ), pytest.param( 'sum(p, by=gen_bz, into=bus, over=generator)', @@ -179,7 +179,7 @@ def namespace() -> Namespace: id='a-map-into-its-own-dimension-keeps-the-frame', ), pytest.param( - 'at(p, by=rep_of, over=rep, into=snapshot)', {'snapshot', 'generator'}, id='and-so-does-its-pullback' + 'at(p, by=rep_of, over=rep, into=snapshot)', {'snapshot', 'generator'}, id='and-so-does-its-lookup' ), pytest.param( "shift(p, along=snapshot, offset=1, edge='wrap', by=rep_of, within=rep)", @@ -260,7 +260,7 @@ def test_a_bare_name_reaches_the_variable_a_dual_the_same_named_constraint(): ), pytest.param( 'sum(load, by=gen_bus, over=generator, into=bus)', - r"sum\(by=gen_bus\) consumes \['generator'\], the dims it reads from", + r"sum\(by=gen_bus\) joins on \['generator'\] to sum it away", id='sum-requires-the-grouped-dim', ), pytest.param( @@ -311,7 +311,7 @@ def test_a_bare_name_reaches_the_variable_a_dual_the_same_named_constraint(): pytest.param( 'at(zone_cap, by=gen_zone, into=generator, over=zone)', r"at\(by=gen_zone\) joins on \['snapshot'\]", - id='a-pullback-needs-the-keys-it-joins-on', + id='a-lookup-needs-the-keys-it-joins-on', ), pytest.param( "shift(cost, along=generator, offset=1, edge='wrap', by=gen_zone, within=zone)", @@ -331,24 +331,24 @@ def test_an_ill_dimensioned_expression_is_rejected(expr, match): pytest.param( 'sum(p, by=diag, over=k, into=z)', {'key': {'k': 'generator', 'j': 'generator', 'z': 'zone'}}, - id='a-sum-consuming-a-column-over-the-dimension-it-joins-on', + id='a-sum-summing-away-a-column-over-a-dimension-it-also-joins-on', ), pytest.param( 'at(load, by=diag, over=rep, into=generator)', {'key': ['snapshot', 'generator'], 'values': {'rep': 'snapshot'}}, - id='a-read-consuming-a-column-over-the-dimension-it-joins-on', + id='a-lookup-reading-a-column-over-a-dimension-it-also-joins-on', ), ], ) -def test_a_joined_column_is_not_also_consumed(expr, diag): - """The operand carries one coordinate per dimension, so a column consumed and a column joined on cannot share one. - - The `at` case passed: the check asked whether a joined dimension was - *produced*, which the landing check already refuses, and not whether it - was consumed. `at(load, by=diag, over=rep, into=generator)` then read - `rep` at the operand's snapshot and joined on the key's snapshot at the - same coordinate, and landed on `[bus, generator]` with the joined - dimension gone. +def test_two_joined_columns_do_not_share_a_dimension(expr, diag): + """The operand carries one coordinate per dimension, so the `over=` column and an unnamed key column cannot share one. + + The `at` case passed: the check asked whether an unnamed key dimension + was also grouped by, which the grouping check already refuses, and not + whether it was also the one read. `at(load, by=diag, over=rep, into=generator)` + then read `rep` at the operand's snapshot and joined on the key's snapshot + at the same coordinate, and came out over `[bus, generator]` with the + joined dimension gone. """ with pytest.raises(DimensionError, match=r"joins 'diag' on \[.*\] through more than one column"): _dims_with(expr, **{'relations.diag': diag}) diff --git a/tests/test_lowering.py b/tests/test_lowering.py index e198898b..ce0da678 100644 --- a/tests/test_lowering.py +++ b/tests/test_lowering.py @@ -30,12 +30,13 @@ Constant, DimensionComparison, DimensionDeclaration, - Direction, Divide, Dual, Expression, Footprint, GroupSum, + Join, + Lookup, Mask, Multiply, Negate, @@ -47,7 +48,6 @@ Partition, Power, Program, - Pullback, Region, RelationDeclaration, Sum, @@ -84,15 +84,15 @@ 'constraints': {'c': {'dims': [], 'expression': 'sum(p, over=g) >= 1'}}, } -#: `lk` as `sum` reads it: key consumed, value produced, nothing joined. +#: `lk` as `sum` joins it: joined on the key, grouped by the value, no key column left unnamed. LK = RelationDeclaration((('g', 'g'), ('h', 'h')), ('g',)) LK2 = RelationDeclaration((('g', 'g'), ('z', 'z')), ('g',)) -LK_DIRECTION = Direction('lk', LK, ('g',), ('h',), ()) +LK_JOIN = Join('lk', LK, ('g',), ('h',)) AT_BUS = RelationDeclaration((('g', 'g'), ('bus', 'bus')), ('g',)) #: `fixtures.SMALL_MODEL` plus a second relation and a per-entity #: offset. Which node a construct becomes is mostly a claim about the dim it -#: consumes and the dim it lands on, and stating that needs a third dimension +#: joins on and the dim it groups by, and stating that needs a third dimension #: and two relations over one of them. SHAPES_MODEL = override( SMALL_MODEL, @@ -409,17 +409,17 @@ def test_a_power_lowers_to_a_node_of_its_own(dispatch_schema): @pytest.mark.parametrize( ('expression', 'expected'), [ - pytest.param('sum(q)', Sum(Variable('q'), ('g', 'h')), id='a-bare-sum-consumes-every-dim-the-operand-carries'), - pytest.param('sum(q, over=h)', Sum(Variable('q'), ('h',)), id='an-over-consumes-the-dim-it-names'), + pytest.param('sum(q)', Sum(Variable('q'), ('g', 'h')), id='a-bare-sum-sums-away-every-dim-the-operand-carries'), + pytest.param('sum(q, over=h)', Sum(Variable('q'), ('h',)), id='an-over-sums-away-the-dim-it-names'), pytest.param( 'sum(p, by=lk, over=g, into=h)', - GroupSum(Variable('p'), direction=LK_DIRECTION), - id='a-grouped-sum-names-the-dim-it-consumes-and-the-one-it-lands-on', + GroupSum(Variable('p'), join=LK_JOIN), + id='a-grouped-sum-names-the-dim-it-joins-on-and-the-one-it-groups-by', ), pytest.param( 'at(r, by=lk, over=h, into=g)', - Pullback(Variable('r'), direction=Direction('lk', LK, ('h',), ('g',), ())), - id='a-pullback-reads-the-same-table-back', + Lookup(Variable('r'), join=Join('lk', LK, ('h',), ('g',))), + id='a-lookup-joins-the-same-table-the-other-way', ), pytest.param( "shift(p, along=g, offset=1, edge='wrap')", @@ -499,7 +499,7 @@ def test_a_partition_keeps_its_group_when_the_relation_gains_a_value_column(): }, } ) - grouping[str(values)] = _partition_of(program.constraints['k']).group + grouping[str(values)] = _partition_of(program.constraints['k']).grouped assert grouping == {'day': ('day',), "['day', 'week']": ('day',)}, ( 'the group is the columns the call named, on both calendars' ) @@ -512,8 +512,8 @@ def _partition_of(row): return partition -def test_a_relation_lowers_with_the_direction_each_call_names(): - """Every node reading a relation carries its columns, its key and the direction, so a consumer joins on the right columns.""" +def test_a_relation_lowers_with_the_join_each_call_names(): + """Every node reading a relation carries its columns, its key and the join, so a consumer joins on the right columns.""" program = to_program( { 'dimensions': {'snapshot': {'dtype': 'int'}, 'generator': {}, 'zone': {}}, @@ -548,30 +548,32 @@ def test_a_relation_lowers_with_the_direction_each_call_names(): assert program.relations == {'zone_of': declared}, 'the relation sits once in the program, under its name' zonal = program.constraints['zonal'].lhs assert zonal == GroupSum( - Variable('p'), direction=Direction('zone_of', declared, ('generator',), ('zone',), ('snapshot',)) - ), 'a grouped sum names the column it consumes, the one it produces and the one it joins on' + Variable('p'), join=Join('zone_of', declared, ('generator', 'snapshot'), ('zone', 'snapshot')) + ), ( + 'a grouped sum joins on the over= column and the unnamed key column, and groups by the into= column and that key column' + ) assert isinstance(zonal, GroupSum) - assert (zonal.direction.consumed_dims, zonal.direction.produced_dims, zonal.direction.joined_dims) == ( + assert (zonal.join.dropped_dims, zonal.join.added_dims, zonal.join.kept) == ( ('generator',), ('zone',), ('snapshot',), - ), 'the dims a consumer reads are read off the direction' - assert zonal.direction.relation is program.relations['zone_of'], ( - 'the direction holds the one declaration the program holds, not an equal copy built again' + ), 'the dims a consumer reads are read off the join: dropped, added, and the key columns kept' + assert zonal.join.relation is program.relations['zone_of'], ( + 'the join holds the one declaration the program holds, not an equal copy built again' ) assert program.constraints['history'].lhs == GroupSum( - Variable('p'), direction=Direction('zone_of', declared, ('snapshot',), ('zone',), ('generator',)) - ), 'the same table read from its other key column' + Variable('p'), join=Join('zone_of', declared, ('snapshot', 'generator'), ('zone', 'generator')) + ), 'the same table joined on its other key column' priced = program.constraints['priced'].rhs - assert priced == Pullback( - Parameter('price'), direction=Direction('zone_of', declared, ('zone',), ('generator',), ('snapshot',)) - ), 'and its adjoint consumes the value column and produces the key column' - assert isinstance(priced, Pullback) - assert (priced.direction.consumed_dims, priced.direction.produced_dims, priced.direction.joined_dims) == ( + assert priced == Lookup( + Parameter('price'), join=Join('zone_of', declared, ('zone', 'snapshot'), ('generator', 'snapshot')) + ), 'and the lookup joins on the value column and groups by the key column' + assert isinstance(priced, Lookup) + assert (priced.join.dropped_dims, priced.join.added_dims, priced.join.kept) == ( ('zone',), ('generator',), ('snapshot',), - ), 'an at consumes the coarse dims, produces the fine, and joins on the rest of the key' + ), 'an at joins on the coarse dims, groups by the fine, and keeps the rest of the key' p_where = program.variables['p'].where assert p_where is not None assert [(type(a).__name__, a.dims) for a in p_where.atoms] == [ @@ -590,13 +592,13 @@ def test_a_binary_variable_lowers_to_a_binary_domain(): assert program.variables['dispatch'].domain == 'binary' -def test_a_divisor_under_a_pullback_is_still_named(): +def test_a_divisor_under_a_lookup_is_still_named(): """`children` has to descend through every node, or a refusal loses its name.""" quotient = Divide(Variable('x'), Parameter('rate')) component_of = RelationDeclaration((('flow', 'flow'), ('component', 'component')), ('flow',)) - pulled = Pullback(quotient, direction=Direction('component_of', component_of, ('component',), ('flow',), ())) + pulled = Lookup(quotient, join=Join('component_of', component_of, ('component',), ('flow',))) - assert divisor_parameters(pulled) == frozenset({'rate'}), 'the walk descends through `Pullback`' + assert divisor_parameters(pulled) == frozenset({'rate'}), 'the walk descends through `Lookup`' assert divisor_parameters(Sum(pulled, ('flow',))) == frozenset({'rate'}), 'and through a `Sum` over it' @@ -669,8 +671,8 @@ def test_walk_is_the_node_column_of_walk_regions(): Power(Parameter('c'), Constant(2.0)): 'one-to-one', Divide(Variable('p'), Parameter('c')): 'one-to-one', Sum(Variable('p'), ('g',)): 'many-to-one', - GroupSum(Variable('p'), direction=Direction('at_bus', AT_BUS, ('g',), ('bus',), ())): 'many-to-one', - Pullback(Variable('p'), direction=Direction('at_bus', AT_BUS, ('bus',), ('g',), ())): 'one-to-one', + GroupSum(Variable('p'), join=Join('at_bus', AT_BUS, ('g',), ('bus',))): 'many-to-one', + Lookup(Variable('p'), join=Join('at_bus', AT_BUS, ('bus',), ('g',))): 'one-to-one', Translate(Variable('p'), 't', offset=1, wrap=False, fill=0.0): 'one-to-one', WindowSum(Variable('p'), 't', width=2, wrap=False): 'one-to-many', Cases((Region(Mask(ParameterDefined('c', ('g',))), Variable('p')),)): 'one-to-one', diff --git a/tests/test_parser.py b/tests/test_parser.py index d0d73806..1db9619a 100644 --- a/tests/test_parser.py +++ b/tests/test_parser.py @@ -23,10 +23,10 @@ ComparisonNode, DefinitionNode, DimensionNode, - DirectionNode, DualNode, EdgeNode, FunctionCallNode, + JoinNode, NameListNode, NameNode, NumberNode, @@ -47,7 +47,7 @@ from math_spec.program import ( And, BooleanLiteral, - Direction, + Join, Not, Or, Partition, @@ -554,9 +554,9 @@ def test_a_node_prints_as_the_file_writes_it(text, printed): 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',), ())), + JoinNode(Join('zone_of', _ZONE_OF, ('u',), ('zone',))), 'zone_of', - id='a-relation-read-in-a-direction', + id='a-relation-as-a-call-joins-it', ), pytest.param( PartitionNode(Partition('zone_of', _ZONE_OF, 'u', ('zone',), ())), diff --git a/tests/test_validation.py b/tests/test_validation.py index 8efff1c3..18a52d0d 100644 --- a/tests/test_validation.py +++ b/tests/test_validation.py @@ -736,8 +736,8 @@ class TestRulesDecidedWithoutData: 'variables.q.dims': ['g', 'h', 'z'], 'objective': {'expression': 'sum(sum(q, by=lk, over=g, into=h))'}, }, - ("sum(by=lk) lands on ['h'], which the expression already carries",), - id='landing-on-a-dim-the-operand-carries', + ("sum(by=lk) groups by ['h'], which the expression already carries",), + id='grouping-by-a-dim-the-operand-carries', ), pytest.param( {'objective': {'expression': 'sum(sum(p, by=lk, over=z, into=h))'}}, @@ -836,8 +836,8 @@ class TestRulesDecidedWithoutData: 'objective': {'expression': 'sum(at(r, by=rel, over=h, into=g))'}, }, ( - "at reads one value per coordinate, and 'rel' is not single-valued in ['h'] at the columns " - "the call lands on (['g'])", + "at reads one row per group, and grouping 'rel' by ['g'] leaves several rows in a group", + "its key is ['g', 'h']", ), id='at-through-a-bare-relation', ), @@ -849,10 +849,10 @@ class TestRulesDecidedWithoutData: }, ( "into=['z'] names ['z'], which the key of 'lz' does not hold", - "A read lands on the key it reads at, ['g']", - "Land on the key, or sum toward ['z']", + "A lookup reads one row per key, ['g']", + "Group by the key, or sum toward ['z']", ), - id='a-read-landing-on-a-value-column', + id='a-lookup-grouping-by-a-value-column', ), pytest.param( { @@ -861,16 +861,16 @@ class TestRulesDecidedWithoutData: 'objective': {'expression': 'sum(at(r, by=lz, over=h, into=[g, z]))'}, }, ("into=['g', 'z'] names ['z'], which the key of 'lz' does not hold",), - id='a-read-landing-on-the-key-and-a-value-column', + id='a-lookup-grouping-by-the-key-and-a-value-column', ), pytest.param( {'objective': {'expression': 'sum(sum(q, by=lk, over=h, into=g))'}}, ( - "this sum lands on the key ['g']", - 'that is a read, which is', + "the columns this sum groups by, ['g'], hold the whole key ['g']", + 'that is a join with no group-by, which is', "at(..., by=lk, over=['h'], into=['g'])", ), - id='a-sum-that-lands-on-the-key-is-a-read', + id='a-sum-grouped-by-the-whole-key-is-a-lookup', ), pytest.param( { @@ -1042,8 +1042,8 @@ class TestRulesDecidedWithoutData: ), pytest.param( {'objective': {'expression': 'sum(at(c, by=lk, over=h, into=g))'}}, - ("at(by=lk) reads through ['h'], which the expression does not carry (dims ['g'])",), - id='a-read-whose-operand-lacks-the-column-it-reads-through', + ("at(by=lk) joins on ['h'], which the expression does not carry (dims ['g'])",), + id='a-lookup-whose-operand-lacks-the-column-it-joins-on', ), pytest.param( { @@ -1095,15 +1095,15 @@ def test_a_rule_decided_without_data(self, patch, fragments): for fragment in fragments: assert fragment in message - def test_the_at_a_sum_landing_on_the_key_names_is_one_the_language_takes(self): + def test_the_at_a_sum_grouped_by_the_whole_key_names_is_one_the_language_takes(self): """A refusal that names a call is holding out a rewrite, so the rewrite has to load. - Was: the message swapped the direction's ends, answering a sum refused - for landing on the key with `at(..., over=, into=)` — which - `at` refuses in turn, for reading a column that is not single-valued - at the one the operand fixes. Both operators take `over=` as the - column the read consumes, so the rewrite is the author's own spelling - with `at` in place of `sum`. + Was: the message swapped the join's ends, answering a sum refused + for grouping by the whole key with `at(..., over=, into=)` — + which `at` refuses in turn, for grouping in a way that leaves several + rows per group. Both operators take `over=` as the columns joined on + and dropped, so the rewrite is the author's own spelling with `at` in + place of `sum`. Reading the call out of the message rather than restating it is the point: a fragment can agree with a message that names a call nothing diff --git a/tests/typesetting/golden/latex.out b/tests/typesetting/golden/latex.out index abeef120..4f96a3d7 100644 --- a/tests/typesetting/golden/latex.out +++ b/tests/typesetting/golden/latex.out @@ -91,9 +91,9 @@ \text{window} && \sum_{t' \in \mathcal{T} \,:\, 0 \le t - t' < 3} \mathit{on}_{t',g} & \le \mathit{units}_{g} && \forall\, t \in \mathcal{T},\ g \in \mathcal{G} \\ \text{history} && \sum_{t' \in \mathcal{T} \,:\, 0 \le t \ominus t' < \mathrm{min\_up}} \mathit{on}_{t',g} & \le \mathit{units}_{g} && \forall\, t \in \mathcal{T},\ g \in \mathcal{G} \\ \text{seasonal\_window} && \sum_{t' \in \mathcal{T} \,:\, 0 \le t -^{\mathrm{season\_of}(t)} t' < 3} \mathit{on}_{t',g} & \le \mathit{units}_{g} && \forall\, t \in \mathcal{T},\ g \in \mathcal{G} \\ -\text{pullback} && \mathit{spill}_{t} & \le \mathrm{zone\_cap}_{\mathrm{zone\_of}(b)} && \forall\, t \in \mathcal{T},\ b \in \mathcal{B} \\ +\text{lookup} && \mathit{spill}_{t} & \le \mathrm{zone\_cap}_{\mathrm{zone\_of}(b)} && \forall\, t \in \mathcal{T},\ b \in \mathcal{B} \\ \text{grouped\_once} && \sum_{g \in \mathcal{G} \,:\, \mathrm{gen\_bt.bus}(g) = b \wedge \mathrm{gen\_bt.technology}(g) = e} p_{t,g} & \le \mathrm{tech\_cap}_{b,e} && \forall\, t \in \mathcal{T},\ b \in \mathcal{B},\ e \in \mathcal{E} \\ -\text{pulled\_back\_once} && \mathit{units}_{g} & \le \mathrm{tech\_cap}_{\mathrm{gen\_bt.bus}(g),\mathrm{gen\_bt.technology}(g)} && \forall\, g \in \mathcal{G} \\ +\text{looked\_up\_once} && \mathit{units}_{g} & \le \mathrm{tech\_cap}_{\mathrm{gen\_bt.bus}(g),\mathrm{gen\_bt.technology}(g)} && \forall\, g \in \mathcal{G} \\ \text{within\_bus} && \mathit{units}_{g} & \le \mathit{units}_{g \boxminus_{0}^{\mathrm{gen\_bt.bus}(g)} 1} && \forall\, g \in \mathcal{G} \,:\, \mathrm{pos}_{\left( \mathrm{gen\_bt.bus}(g),\ \mathrm{gen\_bt.technology}(g) \right)}(g) = 0 \\ \text{relational} && \sum_{g \in \mathcal{G} \,:\, \left( g,\ b \right) \in \mathrm{connection}} p_{t,g} & \le \mathrm{load}_{t,b} && \forall\, t \in \mathcal{T},\ b \in \mathcal{B} \\ \text{connected} && p_{t,g} & \le \mathrm{load}_{t,b} && \forall\, t \in \mathcal{T},\ g \in \mathcal{G},\ b \in \mathcal{B} \,:\, \left( g,\ b \right) \in \mathrm{connection} \\ @@ -101,7 +101,7 @@ \text{zonal} && \sum_{g \in \mathcal{G} \,:\, \mathrm{gen\_zone}(g,\ t) = z} p_{t,g} & \le \mathrm{zone\_cap}_{z} && \forall\, t \in \mathcal{T},\ z \in \mathcal{Z} \\ \text{zonal\_history} && \sum_{t \in \mathcal{T} \,:\, \mathrm{gen\_zone}(g,\ t) = z} p_{t,g} & \le \mathrm{zone\_cap}_{z} && \forall\, g \in \mathcal{G},\ z \in \mathcal{Z} \\ \text{zonal\_membership} && \sum_{g \in \mathcal{G} \,:\, \mathrm{gen\_zone}(g,\ t) \text{ is defined}} \mathit{units}_{g} & \le \mathrm{budget} && \forall\, t \in \mathcal{T} \\ -\text{zonal\_pullback} && p_{t,g} & \le \mathit{spill}_{t} \cdot \mathrm{zone\_cap}_{\mathrm{gen\_zone}(g,\ t)} && \forall\, t \in \mathcal{T},\ g \in \mathcal{G} \,:\, \mathrm{gen\_zone}(g,\ t) = \text{'}\mathrm{north}\text{'} \wedge \mathrm{pos}_{\mathrm{gen\_zone}(g,\ t)}(g) = 0 \\ +\text{zonal\_lookup} && p_{t,g} & \le \mathit{spill}_{t} \cdot \mathrm{zone\_cap}_{\mathrm{gen\_zone}(g,\ t)} && \forall\, t \in \mathcal{T},\ g \in \mathcal{G} \,:\, \mathrm{gen\_zone}(g,\ t) = \text{'}\mathrm{north}\text{'} \wedge \mathrm{pos}_{\mathrm{gen\_zone}(g,\ t)}(g) = 0 \\ \text{arithmetic} && \sum_{g \in \mathcal{G}} \left( \frac{p_{t,g}}{2} - \mathrm{cost}_{g} + 10^{-5} \cdot p_{t,g} + 2.5 \times 10^{-7} \cdot \mathrm{cost}_{g} + 0.5 \cdot p_{t,g} \right) & \ge -\left( \sum_{g \in \mathcal{G}} p_{t,g} \right) \cdot \left( -3 \right) && \forall\, t \in \mathcal{T} \\ \text{total} && \sum_{t \in \mathcal{T},\ g \in \mathcal{G}} p_{t,g} & \le \mathrm{budget} \\ \text{scalar} && \mathit{units}_{g} & \le \mathrm{budget} && \forall\, g \in \mathcal{G} \,:\, \mathrm{cost}_{g} \text{ is defined} \\ diff --git a/tests/typesetting/golden/markdown.out b/tests/typesetting/golden/markdown.out index c4759172..5dce49eb 100644 --- a/tests/typesetting/golden/markdown.out +++ b/tests/typesetting/golden/markdown.out @@ -166,7 +166,7 @@ p_{t,g} \le p_{t \boxminus_{0}^{\mathrm{season\_of}(t)} 1,g} \qquad \forall\, t \sum_{t' \in \mathcal{T} \,:\, 0 \le t -^{\mathrm{season\_of}(t)} t' < 3} \mathit{on}_{t',g} \le \mathit{units}_{g} \qquad \forall\, t \in \mathcal{T},\ g \in \mathcal{G} ``` -**`pullback`** +**`lookup`** ```math \mathit{spill}_{t} \le \mathrm{zone\_cap}_{\mathrm{zone\_of}(b)} \qquad \forall\, t \in \mathcal{T},\ b \in \mathcal{B} @@ -178,7 +178,7 @@ p_{t,g} \le p_{t \boxminus_{0}^{\mathrm{season\_of}(t)} 1,g} \qquad \forall\, t \sum_{g \in \mathcal{G} \,:\, \mathrm{gen\_bt.bus}(g) = b \wedge \mathrm{gen\_bt.technology}(g) = e} p_{t,g} \le \mathrm{tech\_cap}_{b,e} \qquad \forall\, t \in \mathcal{T},\ b \in \mathcal{B},\ e \in \mathcal{E} ``` -**`pulled_back_once`** +**`looked_up_once`** ```math \mathit{units}_{g} \le \mathrm{tech\_cap}_{\mathrm{gen\_bt.bus}(g),\mathrm{gen\_bt.technology}(g)} \qquad \forall\, g \in \mathcal{G} @@ -226,7 +226,7 @@ p_{t,g} \le \mathrm{load}_{t,b} \qquad \forall\, t \in \mathcal{T},\ g \in \math \sum_{g \in \mathcal{G} \,:\, \mathrm{gen\_zone}(g,\ t) \text{ is defined}} \mathit{units}_{g} \le \mathrm{budget} \qquad \forall\, t \in \mathcal{T} ``` -**`zonal_pullback`** +**`zonal_lookup`** ```math p_{t,g} \le \mathit{spill}_{t} \cdot \mathrm{zone\_cap}_{\mathrm{gen\_zone}(g,\ t)} \qquad \forall\, t \in \mathcal{T},\ g \in \mathcal{G} \,:\, \mathrm{gen\_zone}(g,\ t) = \text{'}\mathrm{north}\text{'} \wedge \mathrm{pos}_{\mathrm{gen\_zone}(g,\ t)}(g) = 0 diff --git a/tests/typesetting/golden/model.yaml b/tests/typesetting/golden/model.yaml index 321491ab..2dc71dab 100644 --- a/tests/typesetting/golden/model.yaml +++ b/tests/typesetting/golden/model.yaml @@ -25,7 +25,7 @@ relations: zone_of: { key: bus, values: zone } area_of: { key: bus, values: zone } # a second map into the same set, to compare against season_of: { key: snapshot, values: season } - gen_zone: { key: [generator, snapshot], values: zone } # a map keyed by two dimensions: a call consumes one and joins on the other + gen_zone: { key: [generator, snapshot], values: zone } # a map keyed by two dimensions: a call sums one away and joins on the other rep_of: { key: snapshot, values: { rep: snapshot } } # a map into its own dimension: the representative snapshot connection: { key: [generator, bus] } # a bare relation, with no value columns: many-to-many, read only by sum with both ends named gen_bt: { key: generator, values: [bus, technology] } # one table with two value columns, read to both at once @@ -146,13 +146,13 @@ constraints: seasonal_window: # a window partitioned by a relation: the group rides on the operator dims: [snapshot, generator] expression: sum_back(on, along=snapshot, window=3, by=season_of, within=season) <= units - pullback: # at(), which re-indexes through a relation instead of an offset + lookup: # at(), which re-indexes through a relation instead of an offset dims: [snapshot, bus] expression: spill <= at(zone_cap, by=zone_of, over=zone, into=bus) grouped_once: # one table read to two value columns: the domain carries a condition per column dims: [snapshot, bus, technology] expression: sum(p, by=gen_bt, into=[bus, technology], over=generator) <= tech_cap - pulled_back_once: # its adjoint, reading one slot through two columns of one table + looked_up_once: # the same table joined the other way, reading one slot through two columns dims: [generator] expression: units <= at(tech_cap, by=gen_bt, over=[bus, technology], into=generator) within_bus: # a partition grouped by one named value column of a two-value table, and a position within both @@ -178,7 +178,7 @@ constraints: zonal_membership: # the same table read between its two key columns: no value column is read, so the domain asks only that the row is there dims: [snapshot] expression: sum(units, by=gen_zone, over=generator, into=snapshot) <= budget - zonal_pullback: # its adjoint, reading the slot the row's own snapshot puts the generator in + zonal_lookup: # the same table joined the other way, reading the slot the row's own snapshot puts the generator in dims: [snapshot, generator] where: "gen_zone == 'north' AND position(generator, by=gen_zone, within=zone) == 0" expression: p <= at(spill * zone_cap, by=gen_zone, into=generator, over=zone) diff --git a/tests/typesetting/golden/typst.out b/tests/typesetting/golden/typst.out index ef848e20..bbff8cdc 100644 --- a/tests/typesetting/golden/typst.out +++ b/tests/typesetting/golden/typst.out @@ -78,9 +78,9 @@ $ upright("budgeted") & italic("spend")_(t) & <= upright("budget") & forall t in upright("window") & sum_(t' in cal(T) colon 0 <= t - t' < 3) italic("on")_(t',g) & <= italic("units")_(g) & forall t in cal(T), g in cal(G) \ upright("history") & sum_(t' in cal(T) colon 0 <= t minus.o t' < upright("min_up")) italic("on")_(t',g) & <= italic("units")_(g) & forall t in cal(T), g in cal(G) \ upright("seasonal_window") & sum_(t' in cal(T) colon 0 <= t -^(upright("season_of")(t)) t' < 3) italic("on")_(t',g) & <= italic("units")_(g) & forall t in cal(T), g in cal(G) \ - upright("pullback") & italic("spill")_(t) & <= upright("zone_cap")_(upright("zone_of")(b)) & forall t in cal(T), b in cal(B) \ + upright("lookup") & italic("spill")_(t) & <= upright("zone_cap")_(upright("zone_of")(b)) & forall t in cal(T), b in cal(B) \ upright("grouped_once") & sum_(g in cal(G) colon upright("gen_bt.bus")(g) = b and upright("gen_bt.technology")(g) = e) p_(t,g) & <= upright("tech_cap")_(b,e) & forall t in cal(T), b in cal(B), e in cal(E) \ - upright("pulled_back_once") & italic("units")_(g) & <= upright("tech_cap")_(upright("gen_bt.bus")(g),upright("gen_bt.technology")(g)) & forall g in cal(G) \ + upright("looked_up_once") & italic("units")_(g) & <= upright("tech_cap")_(upright("gen_bt.bus")(g),upright("gen_bt.technology")(g)) & forall g in cal(G) \ upright("within_bus") & italic("units")_(g) & <= italic("units")_(g minus.square_(0)^(upright("gen_bt.bus")(g)) 1) & forall g in cal(G) colon upright("pos")_((upright("gen_bt.bus")(g), upright("gen_bt.technology")(g)))(g) = 0 \ upright("relational") & sum_(g in cal(G) colon (g, b) in upright("connection")) p_(t,g) & <= upright("load")_(t,b) & forall t in cal(T), b in cal(B) \ upright("connected") & p_(t,g) & <= upright("load")_(t,b) & forall t in cal(T), g in cal(G), b in cal(B) colon (g, b) in upright("connection") \ @@ -88,7 +88,7 @@ $ upright("budgeted") & italic("spend")_(t) & <= upright("budget") & forall t in upright("zonal") & sum_(g in cal(G) colon upright("gen_zone")(g, t) = z) p_(t,g) & <= upright("zone_cap")_(z) & forall t in cal(T), z in cal(Z) \ upright("zonal_history") & sum_(t in cal(T) colon upright("gen_zone")(g, t) = z) p_(t,g) & <= upright("zone_cap")_(z) & forall g in cal(G), z in cal(Z) \ upright("zonal_membership") & sum_(g in cal(G) colon upright("gen_zone")(g, t) upright(" is defined")) italic("units")_(g) & <= upright("budget") & forall t in cal(T) \ - upright("zonal_pullback") & p_(t,g) & <= italic("spill")_(t) dot upright("zone_cap")_(upright("gen_zone")(g, t)) & forall t in cal(T), g in cal(G) colon upright("gen_zone")(g, t) = upright("'north'") and upright("pos")_(upright("gen_zone")(g, t))(g) = 0 \ + upright("zonal_lookup") & p_(t,g) & <= italic("spill")_(t) dot upright("zone_cap")_(upright("gen_zone")(g, t)) & forall t in cal(T), g in cal(G) colon upright("gen_zone")(g, t) = upright("'north'") and upright("pos")_(upright("gen_zone")(g, t))(g) = 0 \ upright("arithmetic") & sum_(g in cal(G)) (frac(p_(t,g), 2) - upright("cost")_(g) + 10^(-5) dot p_(t,g) + 2.5 times 10^(-7) dot upright("cost")_(g) + 0.5 dot p_(t,g)) & >= -(sum_(g in cal(G)) p_(t,g)) dot (-3) & forall t in cal(T) \ upright("total") & sum_(t in cal(T), g in cal(G)) p_(t,g) & <= upright("budget") \ upright("scalar") & italic("units")_(g) & <= upright("budget") & forall g in cal(G) colon upright("cost")_(g) upright(" is defined") \ diff --git a/tests/typesetting/test_golden.py b/tests/typesetting/test_golden.py index 1966268e..d692addf 100644 --- a/tests/typesetting/test_golden.py +++ b/tests/typesetting/test_golden.py @@ -152,7 +152,7 @@ def _rendered_trees() -> Iterator[object]: #: relation it reads are the facts a node carries rather than nodes. None is a #: member of any node union, so they are subtracted from what the tree walk #: finds rather than added to what the vocabulary declares. -CARRIERS = {'CaseArm', 'Direction', 'Partition', 'RelationDeclaration'} +CARRIERS = {'CaseArm', 'Join', 'Partition', 'RelationDeclaration'} def test_the_golden_model_carries_every_node_kind_the_walk_renders(): diff --git a/tests/typesetting/test_walk.py b/tests/typesetting/test_walk.py index c1c4cc34..d2705b96 100644 --- a/tests/typesetting/test_walk.py +++ b/tests/typesetting/test_walk.py @@ -176,7 +176,7 @@ def test_a_fill_and_a_group_take_the_operators_two_slots(name: FormatName, fmt: @EVERY_FORMAT -def test_a_translation_under_a_pullback_survives_it(name: FormatName, fmt: Format): +def test_a_translation_under_a_lookup_survives_it(name: FormatName, fmt: Format): """``at`` and ``shift`` both re-index at the leaf, and the leaf has one subscript. Whoever wrote it last used to win: ``at(shift(cap, along=period, offset=1, @@ -201,7 +201,7 @@ def test_a_translation_under_a_pullback_survives_it(name: FormatName, fmt: Forma } text = typeset(model, name, legend=False) assert fmt.operators['edge_minus'] in text, 'the shift under the at was dropped from the subscript' - assert fmt.apply(fmt.upright('period_of'), 't') in text, 'the pullback itself was dropped' + assert fmt.apply(fmt.upright('period_of'), 't') in text, 'the lookup itself was dropped' @EVERY_FORMAT From 0b24fc2a60c76004723033d3552daa02937c8036 Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 22 Sep 2026 11:27:35 +0000 Subject: [PATCH 03/18] feat(program): a join names the dims of the key columns it keeps Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01SWBcNGLjNH2i4AqRsfyxaN --- src/math_spec/program.py | 4 ++++ tests/fixtures/every_program_node.yaml | 2 +- tests/test_piecewise.py | 2 +- 3 files changed, 6 insertions(+), 2 deletions(-) diff --git a/src/math_spec/program.py b/src/math_spec/program.py index 203fb998..9a46f5e5 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -521,6 +521,10 @@ def dropped_dims(self) -> tuple[str, ...]: def added_dims(self) -> tuple[str, ...]: return tuple(self.dim(role) for role in self.added) + @property + def kept_dims(self) -> tuple[str, ...]: + return tuple(self.dim(role) for role in self.kept) + @dataclass(frozen=True) class Partition: diff --git a/tests/fixtures/every_program_node.yaml b/tests/fixtures/every_program_node.yaml index a371b705..ab3a5994 100644 --- a/tests/fixtures/every_program_node.yaml +++ b/tests/fixtures/every_program_node.yaml @@ -38,7 +38,7 @@ constraints: grouped: dims: [t, zone] expression: "sum(p, by=zone_of, over=g, into=zone) - q <= load" - pulled_back: + looked_up: dims: [t, g] expression: "p - at(q, by=zone_of, over=zone, into=g) <= 0" translated: diff --git a/tests/test_piecewise.py b/tests/test_piecewise.py index d951b26a..16329602 100644 --- a/tests/test_piecewise.py +++ b/tests/test_piecewise.py @@ -344,7 +344,7 @@ def test_a_link_reading_a_dual_entry_is_refused(): @pytest.mark.parametrize( ('activity', 'match'), [ - pytest.param('at(u_unit, by=unit_of)', 'is not a declared variable', id='a-pullback-through-a-relation'), + pytest.param('at(u_unit, by=unit_of)', 'is not a declared variable', id='a-lookup-through-a-relation'), pytest.param('shift(u, along=snapshot, offset=1)', 'is not a declared variable', id='a-shifted-gate'), pytest.param('u * 2', 'is not a declared variable', id='an-arithmetic-gate'), ], From e73d24a39d778b1368e461d303bf9b2f2673211e Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 22 Sep 2026 11:39:48 +0000 Subject: [PATCH 04/18] docs(language): the last relation prose says join and group-by Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01SWBcNGLjNH2i4AqRsfyxaN --- docs/about/relations-as-linear-maps.md | 4 ++-- docs/examples/operators.md | 2 +- docs/howto/declare-a-column.md | 2 +- docs/reference/notation.md | 4 ++-- examples/operators/sum_by.yaml | 2 +- schema/math-spec.schema.json | 2 +- src/math_spec/model.py | 4 ++-- src/math_spec/operators.py | 2 +- src/math_spec/program.py | 2 +- tests/test_dimensions.py | 20 ++++++++++---------- tests/test_lowering.py | 6 +++--- tests/test_separability.py | 6 +++--- tests/typesetting/golden/model.yaml | 4 ++-- tests/typesetting/test_golden.py | 2 +- 14 files changed, 31 insertions(+), 31 deletions(-) diff --git a/docs/about/relations-as-linear-maps.md b/docs/about/relations-as-linear-maps.md index 40125db6..98d45d93 100644 --- a/docs/about/relations-as-linear-maps.md +++ b/docs/about/relations-as-linear-maps.md @@ -25,7 +25,7 @@ constraints: zonal: dims: [snapshot, zone] expression: sum(p, by=gen_zone, over=generator, into=zone) <= zone_cap - pulled: + looked_up: dims: [snapshot, generator] where: gen_zone expression: p <= at(zone_cap, by=gen_zone, over=zone, into=generator) @@ -177,7 +177,7 @@ consumer builds is not the matrix. - **Off the domain there is no value.** $`\mathrm{zone\_cap} \circ f`$ is undefined where $`f`$ is, so `at` is absent there and [absence spreads](../reference/language/absence.md#how-absence-travels) to the row. - `where: gen_zone` on `pulled` writes the domain of $`f`$ on the page, so a + `where: gen_zone` on `looked_up` writes the domain of $`f`$ on the page, so a reader sees which rows exist without opening the data. ## Partitions and tests diff --git a/docs/examples/operators.md b/docs/examples/operators.md index 90167717..6f466621 100644 --- a/docs/examples/operators.md +++ b/docs/examples/operators.md @@ -75,7 +75,7 @@ $`\sum_{g \in \mathcal{G}} p_{t,g} \le \mathrm{limit}_{t} \qquad \forall\, t \in ```yaml description: >- - The membership reduction — `sum(array, by=relation, over=a, into=b)` lands the result on the + The membership reduction — `sum(array, by=relation, over=a, into=b)` groups the result by the column the relation is read to, which is what makes topology data rather than structure. diff --git a/docs/howto/declare-a-column.md b/docs/howto/declare-a-column.md index 12c8c173..38cd5a59 100644 --- a/docs/howto/declare-a-column.md +++ b/docs/howto/declare-a-column.md @@ -13,7 +13,7 @@ what the math does with the column. | The column… | is declared as | because | | ------------------------------------------------------------------------------------------------------------------------------------- | --------------------------------------- | ------------------------------------------------------------------------------------------------------------------------- | -| is an axis: something is indexed by it, or an aggregation lands terms on it | a `dimension` | its members are the coordinate set every table over it is reindexed onto | +| is an axis: something is indexed by it, or a grouping groups terms by it | a `dimension` | its members are the coordinate set every table over it is reindexed onto | | has one value per member of a dimension, or per tuple of several — a generator's bus, a line's two ends, a generator's zone by period | a `relation` with that `key` | it is a map every operator reads, and its values are checked against the dimensions they name | | relates members of two dimensions many-to-many, with nothing to weigh — which buses a generator may connect to | a bare `relation`, with no `values:` | `sum` reads it with both ends named, and a bare `where` tests it. Nothing reads it, because there is no one value to read | | relates members of two dimensions many-to-many, with a weight per pair — a link's efficiency to each bus, a cycle's lines | a `parameter` over both | the weight is the data, its row set is the relation, and the aggregation is `sum(w * x, over=a)` | diff --git a/docs/reference/notation.md b/docs/reference/notation.md index 02598cb2..db2d45cd 100644 --- a/docs/reference/notation.md +++ b/docs/reference/notation.md @@ -468,7 +468,7 @@ representative: #### `zonal` -a grouping through a two-key map, consuming one key: the condition reads the other, and the row keeps it +a grouping through a two-key map, summing one key away: the condition reads the other, and the row keeps it ```yaml zonal: @@ -482,7 +482,7 @@ zonal: #### `zonal_history` -the same table consuming its other key +the same table summing its other key away ```yaml zonal_history: diff --git a/examples/operators/sum_by.yaml b/examples/operators/sum_by.yaml index e02cf26c..4eb1eebd 100644 --- a/examples/operators/sum_by.yaml +++ b/examples/operators/sum_by.yaml @@ -3,7 +3,7 @@ # SPDX-License-Identifier: MIT description: >- - The membership reduction — `sum(array, by=relation, over=a, into=b)` lands the result on the + The membership reduction — `sum(array, by=relation, over=a, into=b)` groups the result by the column the relation is read to, which is what makes topology data rather than structure. diff --git a/schema/math-spec.schema.json b/schema/math-spec.schema.json index 450866fc..123afc07 100644 --- a/schema/math-spec.schema.json +++ b/schema/math-spec.schema.json @@ -450,7 +450,7 @@ }, "RelationBlock": { "additionalProperties": false, - "description": "A named relation between dimensions: the columns a row is keyed by, and the columns that key determines.\n\nEach side is a dimension, a list of them, or a mapping of column name to\ndimension where two columns share one. ``key:`` is the claim the language\nchecks at bind: one row per key tuple, so every ``values:`` column is a\nfunction of it. A relation with no ``values:`` is **bare** \u2014 every column is\nin its key, a row is its own identity, and nothing reads it::\n\n relations:\n gen_bus: {key: generator, values: bus}\n gen_bt: {key: [generator], values: [bus, technology]}\n zone_of: {key: [generator, period], values: zone}\n ends: {key: line, values: {bus0: bus, bus1: bus}}\n connection: {key: [generator, bus]}\n\nAn operator reads the table in the direction the call names\n(``over=``, ``into=``), joining on the other key columns; the\ndeclaration fixes no direction. The map itself is data, and arrives at bind\ntime under the relation's name, one column per role.", + "description": "A named relation between dimensions: the columns a row is keyed by, and the columns that key determines.\n\nEach side is a dimension, a list of them, or a mapping of column name to\ndimension where two columns share one. ``key:`` is the claim the language\nchecks at bind: one row per key tuple, so every ``values:`` column is a\nfunction of it. A relation with no ``values:`` is **bare** \u2014 every column is\nin its key, a row is its own identity, and nothing reads it::\n\n relations:\n gen_bus: {key: generator, values: bus}\n gen_bt: {key: [generator], values: [bus, technology]}\n zone_of: {key: [generator, period], values: zone}\n ends: {key: line, values: {bus0: bus, bus1: bus}}\n connection: {key: [generator, bus]}\n\nAn operator joins the table on the columns ``over=`` names and every\nother key column, and groups by the columns ``into=`` names; the\ndeclaration fixes no direction. The map itself is data, and arrives at bind\ntime under the relation's name, one column per role.", "properties": { "description": { "anyOf": [ diff --git a/src/math_spec/model.py b/src/math_spec/model.py index 91b9a91c..728a910d 100644 --- a/src/math_spec/model.py +++ b/src/math_spec/model.py @@ -182,8 +182,8 @@ class RelationBlock(_StrictBlock): ends: {key: line, values: {bus0: bus, bus1: bus}} connection: {key: [generator, bus]} - An operator reads the table in the direction the call names - (``over=``, ``into=``), joining on the other key columns; the + An operator joins the table on the columns ``over=`` names and every + other key column, and groups by the columns ``into=`` names; the declaration fixes no direction. The map itself is data, and arrives at bind time under the relation's name, one column per role. """ diff --git a/src/math_spec/operators.py b/src/math_spec/operators.py index b3db75d8..7d34eea2 100644 --- a/src/math_spec/operators.py +++ b/src/math_spec/operators.py @@ -92,7 +92,7 @@ def kind_of( #: The closed operator set. ``by=`` is the one keyword that addresses a relation, #: and a relation carries its own dimensions, so no sibling kwarg restates them. #: On ``shift`` and ``sum_back`` it partitions the axis the operator steps along: it -#: says which rows are neighbours, not which group a term lands in, and +#: says which rows are neighbours, not which group a term is added to, and #: ``within=`` names the value columns the group is made of, on every call #: that names a ``by=``. BUILTINS: dict[str, Builtin] = { diff --git a/src/math_spec/program.py b/src/math_spec/program.py index 9a46f5e5..ffe3ef04 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -863,7 +863,7 @@ class Separability: is pointwise; a ``shift`` of ``-2`` is ``2``. coupled: Each declaration that ties the axis together, to what ties it and the one modelling change that would not: a sum over the axis - in a constraint, a grouping that consumes it, a wrapped + in a constraint, a grouping that sums it away, a wrapped translation, a set. No window satisfies these, and no rewrite here would keep the model's meaning, so the remedy is named rather than applied. diff --git a/tests/test_dimensions.py b/tests/test_dimensions.py index 30be7970..a8f4574e 100644 --- a/tests/test_dimensions.py +++ b/tests/test_dimensions.py @@ -141,7 +141,7 @@ def namespace() -> Namespace: pytest.param( "shift(p, along=generator, offset=1, edge='wrap', by=pair, within=[b0, b1])", {'snapshot', 'generator'}, - id='a-partition-grouped-by-two-columns-over-one-dimension-lands-nothing', + id='a-partition-grouped-by-two-columns-over-one-dimension-adds-nothing', ), pytest.param( 'sum(p, by=gen_bus, over=generator, into=bus)', @@ -203,7 +203,7 @@ def _dims_with(expr: str, **overrides) -> frozenset[str]: pytest.param( 'sum(p, by=connection, over=generator, into=bus)', {'snapshot', 'bus'}, - id='a-sum-through-a-bare-relation-lands-on-a-key-column', + id='a-sum-through-a-bare-relation-groups-by-a-key-column', ), pytest.param( 'sum(load, by=connection, over=bus, into=generator)', @@ -213,22 +213,22 @@ def _dims_with(expr: str, **overrides) -> frozenset[str]: ], ) def test_a_bare_relation_is_summed_between_its_key_columns(expr, expected): - """A bare relation holds no value column, so the column a sum lands on is a key column.""" + """A bare relation holds no value column, so the column a sum groups by is a key column.""" assert _dims_with(expr, **{'relations.connection': {'key': ['generator', 'bus']}}) == expected -def test_a_read_carries_the_whole_key_and_what_the_operand_brings_beside_it(): - """A read lands on the key however the call splits it, and a dim the operand carries and the read does not consume rides along. +def test_a_lookup_carries_the_whole_key_and_what_the_operand_brings_beside_it(): + """A lookup groups by the key however the call splits it, and a dim the operand carries and the lookup does not join on rides along. - `gen_bz` is keyed by `generator` alone, so the key is produced whole; the - operand's `snapshot` is neither consumed nor part of the key, and the + `gen_bz` is keyed by `generator` alone, so the whole key is grouped by; the + operand's `snapshot` is neither joined on nor part of the key, and the result keeps it. """ assert _dims('at(zone_load, by=gen_bz, over=zone, into=generator)') == {'generator', 'snapshot'} -def test_a_sum_consumes_a_key_column_and_a_value_column_together(): - """A sum's consumed end is not one kind of column: it needs one key column, and may name a value column beside it.""" +def test_a_sum_joins_on_a_key_column_and_a_value_column_together(): + """The columns a sum joins on and sums away are not one kind: it needs one key column, and may name a value column beside it.""" assert _dims('sum(p * load, by=gen_bz, over=[generator, bus], into=zone)') == {'snapshot', 'zone'} @@ -251,7 +251,7 @@ def test_a_bare_name_reaches_the_variable_a_dual_the_same_named_constraint(): pytest.param( 'sum(p, over=bus)', r'sum\(over=bus\) but the expression has dims', - id='sum-consuming-an-absent-dim-is-an-error-not-a-noop', + id='sum-over-an-absent-dim-is-an-error-not-a-noop', ), pytest.param( 'sum(sum(p))', diff --git a/tests/test_lowering.py b/tests/test_lowering.py index ce0da678..c08ca860 100644 --- a/tests/test_lowering.py +++ b/tests/test_lowering.py @@ -596,10 +596,10 @@ def test_a_divisor_under_a_lookup_is_still_named(): """`children` has to descend through every node, or a refusal loses its name.""" quotient = Divide(Variable('x'), Parameter('rate')) component_of = RelationDeclaration((('flow', 'flow'), ('component', 'component')), ('flow',)) - pulled = Lookup(quotient, join=Join('component_of', component_of, ('component',), ('flow',))) + looked_up = Lookup(quotient, join=Join('component_of', component_of, ('component',), ('flow',))) - assert divisor_parameters(pulled) == frozenset({'rate'}), 'the walk descends through `Lookup`' - assert divisor_parameters(Sum(pulled, ('flow',))) == frozenset({'rate'}), 'and through a `Sum` over it' + assert divisor_parameters(looked_up) == frozenset({'rate'}), 'the walk descends through `Lookup`' + assert divisor_parameters(Sum(looked_up, ('flow',))) == frozenset({'rate'}), 'and through a `Sum` over it' def test_a_divisor_under_a_power_is_still_named(): diff --git a/tests/test_separability.py b/tests/test_separability.py index 961e3a07..c8ad7c5f 100644 --- a/tests/test_separability.py +++ b/tests/test_separability.py @@ -212,7 +212,7 @@ def test_the_lookahead_is_the_widest_reach_of_any_block(): assert verdict.ahead == 5, 'one window must see past its last row as far as any block reads' -def test_a_grouping_that_consumes_the_axis_couples_it(): +def test_a_grouping_that_sums_the_axis_away_couples_it(): program = ms.to_program( { **BASE, @@ -220,7 +220,7 @@ def test_a_grouping_that_consumes_the_axis_couples_it(): } ) verdict = program.separability['u'] - assert not verdict.windowable, 'the grouping consumes u, so a window of u is a different sum' + assert not verdict.windowable, 'the grouping sums u away, so a window of u is a different sum' def test_every_declared_axis_has_a_verdict_and_nothing_else_does(): @@ -247,7 +247,7 @@ def test_a_reduction_over_several_axes_couples_every_one_of_them(): so the verdict for each of them has to say so — a walk that read only the first would call the rest windowable.""" program = ms.to_program({**BASE, 'constraints': {'all': {'dims': [], 'expression': 'sum(p) <= budget'}}}) - assert not program.separability['h'].windowable, 'the reduction consumes h' + assert not program.separability['h'].windowable, 'the reduction sums h away' assert not program.separability['u'].windowable, 'and u, in the same node' diff --git a/tests/typesetting/golden/model.yaml b/tests/typesetting/golden/model.yaml index 2dc71dab..c7408841 100644 --- a/tests/typesetting/golden/model.yaml +++ b/tests/typesetting/golden/model.yaml @@ -169,10 +169,10 @@ constraints: representative: # a map into its own dimension, read both ways: the frame is unchanged and the index is primed dims: [snapshot] expression: sum(spill, by=rep_of, over=snapshot, into=rep) <= at(spill, by=rep_of, over=rep, into=snapshot) - zonal: # a grouping through a two-key map, consuming one key: the condition reads the other, and the row keeps it + zonal: # a grouping through a two-key map, summing one key away: the condition reads the other, and the row keeps it dims: [snapshot, zone] expression: sum(p, by=gen_zone, over=generator, into=zone) <= zone_cap - zonal_history: # the same table consuming its other key + zonal_history: # the same table summing its other key away dims: [generator, zone] expression: sum(p, by=gen_zone, over=snapshot, into=zone) <= zone_cap zonal_membership: # the same table read between its two key columns: no value column is read, so the domain asks only that the row is there diff --git a/tests/typesetting/test_golden.py b/tests/typesetting/test_golden.py index d692addf..ffe6f084 100644 --- a/tests/typesetting/test_golden.py +++ b/tests/typesetting/test_golden.py @@ -148,7 +148,7 @@ def _rendered_trees() -> Iterator[object]: } #: A dataclass the walk steps *through* rather than renders: an arm has no -#: branch of its own — its ``when`` and ``value`` do — and a direction and the +#: branch of its own — its ``when`` and ``value`` do — and a join and the #: relation it reads are the facts a node carries rather than nodes. None is a #: member of any node union, so they are subtracted from what the tree walk #: finds rather than added to what the vocabulary declares. From 53ba4a2660e88661092b340ff117ec28217eeac5 Mon Sep 17 00:00:00 2001 From: Felix <117816358+FBumann@users.noreply.github.com> Date: Tue, 22 Sep 2026 14:18:04 +0200 Subject: [PATCH 05/18] feat(program): a sum through a relation is a sum over a join (#607) `GroupSum` and `Lookup` are one node, `Join`, and `sum(by=)` lowers to the existing `Sum` over it, its `over` the dims the join drops; `at` lowers to the bare `Join`. The columns a call names move to `JoinColumns`. The YAML surface is unchanged. Claude-Session: https://claude.ai/code/session_01SWBcNGLjNH2i4AqRsfyxaN Co-authored-by: Claude --- docs/about/relations-as-linear-maps.md | 4 +- docs/contributing.md | 14 +++---- src/math_spec/_expression_parser.py | 8 ++-- src/math_spec/advice.py | 6 +-- src/math_spec/boundedness.py | 5 +-- src/math_spec/dimensions.py | 10 ++--- src/math_spec/lowering.py | 13 +++--- src/math_spec/program.py | 56 +++++++++++--------------- src/math_spec/resolution.py | 12 +++--- src/math_spec/separability.py | 30 ++++++++------ src/math_spec/typesetting/walk.py | 8 ++-- tests/test_lowering.py | 54 ++++++++++++------------- tests/test_parser.py | 4 +- tests/typesetting/test_golden.py | 12 +++--- 14 files changed, 116 insertions(+), 120 deletions(-) diff --git a/docs/about/relations-as-linear-maps.md b/docs/about/relations-as-linear-maps.md index 98d45d93..a3df2727 100644 --- a/docs/about/relations-as-linear-maps.md +++ b/docs/about/relations-as-linear-maps.md @@ -131,8 +131,8 @@ $`x`$ over generators and $`y`$ over buses, \langle M x, y \rangle = \sum_{b} y_b \sum_{g} \mathbf{1}_R(g, b)\, x_g = \sum_{g} x_g \sum_{b} \mathbf{1}_R(g, b)\, y_b = \langle x, M^{\mathsf{T}} y \rangle, ``` -which is why the program's `Lookup` is the same `Join` as its `GroupSum`, read -without the group-by. A bare relation has the same matrix without the +which is why the program lowers `at` to a `Join` node and `sum(by=)` to the +same `Join` under a `Sum`. A bare relation has the same matrix without the functional claim. A column of $`M`$ may hold several ones, so the sum fans out and no group is one row. That is why `at` through a bare relation is refused. diff --git a/docs/contributing.md b/docs/contributing.md index 6e9d5fe3..8a9c4fd9 100644 --- a/docs/contributing.md +++ b/docs/contributing.md @@ -107,13 +107,13 @@ suffix says which layer: 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. -| File verb | Node | What the node names | -| ------------------ | ----------- | ------------------------------ | -| `sum(over=)` | `Sum` | dims removed from the result | -| `sum(by=)` | `GroupSum` | a join and group-by | -| `at(by=)` | `Lookup` | a join with no group-by | -| `shift(along=)` | `Translate` | a re-index along one dimension | -| `sum_back(along=)` | `WindowSum` | a sum over a trailing window | +| File verb | Node | What the node names | +| ------------------ | ------------------- | ------------------------------------------ | +| `sum(over=)` | `Sum` | dims removed from the result | +| `sum(by=)` | `Sum` over a `Join` | a join, and the sum over the dims it drops | +| `at(by=)` | `Join` | a join with no sum over it | +| `shift(along=)` | `Translate` | a re-index along one dimension | +| `sum_back(along=)` | `WindowSum` | a sum over a trailing window | Nothing is abbreviated. diff --git a/src/math_spec/_expression_parser.py b/src/math_spec/_expression_parser.py index 715fa978..753011f2 100644 --- a/src/math_spec/_expression_parser.py +++ b/src/math_spec/_expression_parser.py @@ -23,7 +23,7 @@ if TYPE_CHECKING: from collections.abc import Callable, Iterator, Mapping - from math_spec.program import Join, Partition, Predicate + from math_spec.program import JoinColumns, 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. @@ -133,12 +133,12 @@ def __str__(self) -> str: @dataclass(frozen=True) class JoinNode: - """A resolved ``by=`` on ``sum`` or ``at``: the relation, as the :class:`Join` the call names.""" + """A resolved ``by=`` on ``sum`` or ``at``: the relation, as the :class:`JoinColumns` the call names.""" - join: Join + columns: JoinColumns def __str__(self) -> str: - return self.join.name + return self.columns.name @dataclass(frozen=True) diff --git a/src/math_spec/advice.py b/src/math_spec/advice.py index aef61352..85ba75b7 100644 --- a/src/math_spec/advice.py +++ b/src/math_spec/advice.py @@ -14,7 +14,7 @@ from math_spec.boundedness import unbounded_notes from math_spec.errors import Advice from math_spec.lowering import to_program -from math_spec.program import GroupSum, Lookup, walk +from math_spec.program import Join, walk if TYPE_CHECKING: from collections.abc import Mapping @@ -74,6 +74,6 @@ def _grouped_axes(program: Program) -> set[str]: """ axes: set[str] = set() for node in walk(*program.roots): - if isinstance(node, GroupSum | Lookup): - axes.update(node.join.added_dims) + if isinstance(node, Join): + axes.update(node.columns.added_dims) return axes diff --git a/src/math_spec/boundedness.py b/src/math_spec/boundedness.py index 994d2646..89d8cb03 100644 --- a/src/math_spec/boundedness.py +++ b/src/math_spec/boundedness.py @@ -24,8 +24,7 @@ Divide, Dual, Expression, - GroupSum, - Lookup, + Join, Multiply, Negate, Parameter, @@ -162,7 +161,7 @@ def _record_signs(node: Expression, sign: Sign, signs: dict[str, Sign]) -> None: _record_signs(node.base, None, signs) _record_signs(node.exponent, None, signs) return - if isinstance(node, Sum | GroupSum | Lookup | Translate | WindowSum | Cases): + if isinstance(node, Sum | Join | Translate | WindowSum | Cases): for child in children(node): _record_signs(child, sign, signs) return diff --git a/src/math_spec/dimensions.py b/src/math_spec/dimensions.py index 67cfaad1..9fdd817e 100644 --- a/src/math_spec/dimensions.py +++ b/src/math_spec/dimensions.py @@ -46,7 +46,7 @@ DimensionComparison, DimensionPosition, ExpressionComparison, - Join, + JoinColumns, Mask, ParameterComparison, ParameterDefined, @@ -161,7 +161,7 @@ def _sum_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spec, conte return inner - {summed.name} assert isinstance(by, JoinNode), 'resolution reads sum(by=) as a join' - join = by.join + join = by.columns if missing := sorted(set(join.dropped_dims) - inner): raise DimensionError( _not_carried( @@ -178,7 +178,7 @@ def _at_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spec, contex """``at`` is the join of ``sum(by=)`` with no group-by: it joins on the dims a sum groups by and groups by the ones a sum joins on.""" by = node.kwargs['by'] assert isinstance(by, JoinNode), 'resolution reads at(by=) as a join' - join = by.join + join = by.columns if absent := sorted(set(join.dropped_dims) - inner): raise DimensionError( f'{context}: at(by={join.name}) joins on ' @@ -212,7 +212,7 @@ def _translation_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spe return inner -def _join_dims(call: str, join: Join, inner: frozenset[str], context: str) -> frozenset[str]: +def _join_dims(call: str, join: JoinColumns, inner: frozenset[str], context: str) -> frozenset[str]: """The dims after *join*: the operand's, less the dims joined on, plus the dims grouped by. A column grouped by and not joined on brings its dim, so the operand does @@ -234,7 +234,7 @@ def _join_dims(call: str, join: Join, inner: frozenset[str], context: str) -> fr return (inner - set(join.joined_dims)) | set(join.grouped_dims) -def _check_joined(call: str, use: Join | Partition, inner: frozenset[str], context: str) -> None: +def _check_joined(call: str, use: JoinColumns | Partition, inner: frozenset[str], context: str) -> None: """The columns a call joins on are matched at their dimensions, so the operand carries every one, each once. Two joined columns over one dimension would match the operand's one diff --git a/src/math_spec/lowering.py b/src/math_spec/lowering.py index eb76992a..305c7d77 100644 --- a/src/math_spec/lowering.py +++ b/src/math_spec/lowering.py @@ -304,9 +304,9 @@ def _predicate(self, node: program.Predicate) -> program.Predicate: 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. + One node either way: a ``sum(by=)`` is a :class:`~math_spec.program.Sum` + over a :class:`~math_spec.program.Join`, its ``over`` the dims the join + drops, so ``by=`` decides only what the sum stands over. """ by_node = node.kwargs.get('by') operand = self.expr(node.args[0]) @@ -317,13 +317,14 @@ def sum(self, node: FunctionCallNode) -> program.Expression: assert isinstance(summed, DimensionNode), 'resolution refuses a over= that is not a dimension' return program.Sum(operand, (summed.name,)) assert isinstance(by_node, JoinNode), 'resolution reads sum(by=) as a join' - return program.GroupSum(operand, join=by_node.join) + columns = by_node.columns + return program.Sum(program.Join(operand, columns), columns.dropped_dims) def at(self, node: FunctionCallNode) -> program.Expression: - """``at(x, by=relation)`` — the join of :meth:`sum`'s ``by=`` form with no group-by.""" + """``at(x, by=relation)`` — the :class:`~math_spec.program.Join` of :meth:`sum`'s ``by=`` form, with no sum over it.""" by_node = node.kwargs['by'] assert isinstance(by_node, JoinNode), 'resolution reads at(by=) as a join' - return program.Lookup(self.expr(node.args[0]), join=by_node.join) + return program.Join(self.expr(node.args[0]), by_node.columns) def sum_back(self, node: FunctionCallNode) -> program.Expression: """``sum_back(x, along=d, window=w)`` — a trailing window along one dimension. diff --git a/src/math_spec/program.py b/src/math_spec/program.py index 7eb29beb..130c6260 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -64,11 +64,10 @@ 'FanIn', 'FirstOf', 'Footprint', - 'GroupSum', 'Holds', 'Join', + 'JoinColumns', 'LastOf', - 'Lookup', 'Mask', 'MaskOf', 'Multiply', @@ -241,38 +240,32 @@ class Divide: @dataclass(frozen=True) class Sum: - """Sum ``operand`` over the named dims, removing them from the result.""" + """Sum ``operand`` over the named dims, removing them from the result. - operand: Expression - over: tuple[str, ...] - - -@dataclass(frozen=True) -class GroupSum: - """Sum ``operand`` through a relation: a join on ``join.joined``, then a group-by on ``join.grouped`` with a sum. - - The operand carries every dim joined on. The result drops the dims joined - on and not grouped by, keeps the ones both joined on and grouped by, and - gains the ones grouped by and not joined on. + ``sum(x, by=relation, over=a, into=b)`` lowers to a ``Sum`` over a + :class:`Join`, ``over`` naming the dims the join drops: the group-by is + this node, and the join is its operand. """ operand: Expression - join: Join + over: tuple[str, ...] @dataclass(frozen=True) -class Lookup: - """Read ``operand`` through a relation: the join of :class:`GroupSum` with no group-by. - - The grouped columns hold the relation's whole key, which the loader - checks, so each row of the result meets one row of the relation and - reads one value. The join fans out where several key tuples share the - values joined on, at each coordinate of the columns both joined on and - grouped by. +class Join: + """Join ``operand`` to a relation on the columns ``columns`` joins on, and carry the columns it groups by. + + The operand carries every dim joined on. The result keeps every dim the + operand carries and gains the dims grouped by and not joined on; the dims + joined on and not grouped by leave only under a :class:`Sum` over them. + ``at(x, by=relation, over=a, into=b)`` lowers to a bare ``Join``: the + grouped columns hold the relation's whole key, which the loader checks, so + each row of the result meets one row of the relation and reads one value. + The join fans out where several key tuples share the values joined on. """ operand: Expression - join: Join + columns: JoinColumns @dataclass(frozen=True) @@ -375,8 +368,7 @@ class Cases: | Power | Divide | Sum - | GroupSum - | Lookup + | Join | Translate | WindowSum | Cases @@ -389,13 +381,13 @@ def fan_in(expression: Expression) -> FanIn: For the absence rules, both classes other than ``'one-to-one'`` sum several input slots into an output row. """ - if isinstance(expression, (Sum, GroupSum)): + if isinstance(expression, Sum): return 'many-to-one' if isinstance(expression, WindowSum): return 'one-to-many' if isinstance( expression, - (Constant, Parameter, Variable, Dual, Negate, Add, Multiply, Power, Divide, Lookup, Translate, Cases), + (Constant, Parameter, Variable, Dual, Negate, Add, Multiply, Power, Divide, Join, Translate, Cases), ): return 'one-to-one' assert_never(expression) @@ -411,7 +403,7 @@ def children(expression: Expression) -> tuple[Expression, ...]: return (expression.numerator, expression.divisor) if isinstance(expression, Power): return (expression.base, expression.exponent) - if isinstance(expression, (Sum, GroupSum, Lookup, Translate, WindowSum)): + if isinstance(expression, (Sum, Join, Translate, WindowSum)): return (expression.operand,) if isinstance(expression, Cases): return tuple(region.value for region in expression.regions) @@ -465,7 +457,7 @@ def dim(self, role: str) -> str: @dataclass(frozen=True) -class Join: +class JoinColumns: """One relation as one call joins it: the columns joined on, and the columns grouped by. The declaration fixes no direction; the call does, and this is the one it @@ -1440,8 +1432,8 @@ def _names_under(*expressions: Expression) -> frozenset[str]: for node in walk(*expressions): if isinstance(node, Cases): names.update(*(region.when.names_read for region in node.regions)) - elif isinstance(node, (GroupSum, Lookup)): - names.add(node.join.name) + elif isinstance(node, Join): + names.add(node.columns.name) elif isinstance(node, (Translate, WindowSum)): if node.partition is not None: names.add(node.partition.name) diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index 1c265d01..501baf53 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -74,7 +74,7 @@ CountComparison, DimensionComparison, DimensionPosition, - Join, + JoinColumns, Mask, Not, Or, @@ -227,7 +227,7 @@ class Resolved: objective: The objective's expression, ``None`` where the file declares none. relations: Each relation's columns and key, as declared — the one - copy, which every :class:`~math_spec.program.Join` and + copy, which every :class:`~math_spec.program.JoinColumns` and :class:`~math_spec.program.Partition` in the trees holds. assumptions: Each ``assumptions:`` entry's predicate and the mask it is checked under. @@ -594,8 +594,8 @@ def _relation_ref( return value if partition is None else PartitionNode(partition) if not ({'over', 'into'} <= set(named)): return value # the call shape refused it already, with the wording that names the rewrite - join = self._join(name, operator, named['over'], named['into']) - return value if join is None else JoinNode(join) + columns = self._join(name, operator, named['over'], named['into']) + return value if columns is None else JoinNode(columns) 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.""" @@ -614,7 +614,7 @@ def _join( operator: str, from_roles: tuple[str, ...], into_roles: tuple[str, ...], - ) -> Join | None: + ) -> JoinColumns | None: """How ``sum`` or ``at`` joins relation *name*, between the columns the call named. Both ends arrive written: the call shape refuses a call that leaves @@ -657,7 +657,7 @@ def _join( ) return None kept = tuple(r for r in shape.key if r not in from_roles and r not in into_roles) - join = Join(name, shape, (*from_roles, *kept), (*into_roles, *kept)) + join = JoinColumns(name, shape, (*from_roles, *kept), (*into_roles, *kept)) one_row_per_group = set(shape.key) <= set(join.grouped) if not forward and not one_row_per_group: self.errors.append( diff --git a/src/math_spec/separability.py b/src/math_spec/separability.py index c61fca43..e17f027a 100644 --- a/src/math_spec/separability.py +++ b/src/math_spec/separability.py @@ -11,8 +11,7 @@ from math_spec.program import ( Cases, DimensionPosition, - GroupSum, - Lookup, + Join, Mask, Reach, Separability, @@ -94,9 +93,20 @@ def waits_on(dimension: str, label: str, name: str, kind: Literal['offset', 'par if row is not None: rows[label] = row masks: list[Mask | None] = [mask] + grouped = { + id(node.operand) for node in walk(*nodes) if isinstance(node, Sum) and isinstance(node.operand, Join) + } for node in walk(*nodes): if isinstance(node, Cases): masks.extend(region.when for region in node.regions) + elif isinstance(node, Sum) and isinstance(node.operand, Join): + for dimension in node.over: + report( + 'coupled', + dimension, + label, + f'groups {dimension} into {", ".join(node.operand.columns.added_dims)} — window that dimension instead, or cut only at the group edges', + ) elif isinstance(node, Sum): if reductions_couple: for dimension in node.over: @@ -106,17 +116,11 @@ def waits_on(dimension: str, label: str, name: str, kind: Literal['offset', 'par label, f'sums over {dimension} — a rolling sum_back(window=n) windows, a total over the horizon does not', ) - elif isinstance(node, GroupSum): - for dimension in node.join.dropped_dims: - report( - 'coupled', - dimension, - label, - f'groups {dimension} into {", ".join(node.join.added_dims)} — window that dimension instead, or cut only at the group edges', - ) - elif isinstance(node, Lookup): - for dimension in node.join.dropped_dims: - waits_on(dimension, label, node.join.name, 'coordinate') + elif isinstance(node, Join): + if id(node) in grouped: + continue + for dimension in node.columns.dropped_dims: + waits_on(dimension, label, node.columns.name, 'coordinate') elif isinstance(node, (Translate, WindowSum)): dimension = node.along if node.wrap: diff --git a/src/math_spec/typesetting/walk.py b/src/math_spec/typesetting/walk.py index bee957ed..1f502fbd 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -43,7 +43,7 @@ DimensionComparison, DimensionPosition, ExpressionComparison, - Join, + JoinColumns, Mask, Not, Or, @@ -465,7 +465,7 @@ def _call(self, node: FunctionCallNode, ctx: _Context) -> tuple[str, int]: by = node.kwargs['by'] assert isinstance(by, JoinNode) outer = ctx - join = by.join + join = by.columns at = {r: outer.subscript(join.dim(r)) for r in join.grouped} for read in join.dropped: ctx = ctx.looked_up(join.dim(read), self._relation_read(join.name, at, read)) @@ -473,7 +473,7 @@ def _call(self, node: FunctionCallNode, ctx: _Context) -> tuple[str, int]: if (by := node.kwargs.get('by')) is not None: assert isinstance(by, JoinNode) - join = by.join + join = by.columns dummies: dict[str, str] = {} inner = ctx for d in join.dropped_dims: @@ -496,7 +496,7 @@ def _call(self, node: FunctionCallNode, ctx: _Context) -> tuple[str, int]: domain = self.format.joined(memberships, '') return self.format.summation(domain, self._reduction_body(node.args[0], inner)), _PRECEDENCE['+'] - def _grouping(self, join: Join, dummies: Mapping[str, str], ctx: _Context) -> list[str]: + def _grouping(self, join: JoinColumns, dummies: Mapping[str, str], ctx: _Context) -> list[str]: """The conditions a grouped sum's domain carries for one join: what it fixes of the row it joins on. A join fixes its relation's key either way, so each value column it diff --git a/tests/test_lowering.py b/tests/test_lowering.py index 6fa63bd9..411fba92 100644 --- a/tests/test_lowering.py +++ b/tests/test_lowering.py @@ -36,10 +36,9 @@ Expression, ExpressionComparison, Footprint, - GroupSum, Holds, Join, - Lookup, + JoinColumns, Mask, Multiply, Negate, @@ -91,7 +90,7 @@ #: `lk` as `sum` joins it: joined on the key, grouped by the value, no key column left unnamed. LK = RelationDeclaration((('g', 'g'), ('h', 'h')), ('g',)) LK2 = RelationDeclaration((('g', 'g'), ('z', 'z')), ('g',)) -LK_JOIN = Join('lk', LK, ('g',), ('h',)) +LK_JOIN = JoinColumns('lk', LK, ('g',), ('h',)) AT_BUS = RelationDeclaration((('g', 'g'), ('bus', 'bus')), ('g',)) #: `fixtures.SMALL_MODEL` plus a second relation and a per-entity @@ -427,7 +426,7 @@ def test_a_comparison_of_expressions_lowers_to_program_expressions_on_both_sides ) mask = program.constraints['w'].where assert mask is not None and isinstance(mask.root, ExpressionComparison) - assert isinstance(mask.root.right, Add) and isinstance(mask.root.right.left, Lookup) + assert isinstance(mask.root.right, Add) and isinstance(mask.root.right.left, Join) assert mask.names_read == frozenset({'c', 'zc', 'lk2'}), ( 'the relation a pullback and a partition read through is data the consumer binds too' ) @@ -560,13 +559,13 @@ def test_a_power_lowers_to_a_node_of_its_own(dispatch_schema): pytest.param('sum(q, over=h)', Sum(Variable('q'), ('h',)), id='an-over-sums-away-the-dim-it-names'), pytest.param( 'sum(p, by=lk, over=g, into=h)', - GroupSum(Variable('p'), join=LK_JOIN), - id='a-grouped-sum-names-the-dim-it-joins-on-and-the-one-it-groups-by', + Sum(Join(Variable('p'), LK_JOIN), ('g',)), + id='a-grouped-sum-is-a-sum-over-a-join-of-the-dim-the-join-drops', ), pytest.param( 'at(r, by=lk, over=h, into=g)', - Lookup(Variable('r'), join=Join('lk', LK, ('h',), ('g',))), - id='a-lookup-joins-the-same-table-the-other-way', + Join(Variable('r'), JoinColumns('lk', LK, ('h',), ('g',))), + id='an-at-is-the-same-join-the-other-way-with-no-sum-over-it', ), pytest.param( "shift(p, along=g, offset=1, edge='wrap')", @@ -694,29 +693,30 @@ def test_a_relation_lowers_with_the_join_each_call_names(): declared = RelationDeclaration(columns, ('generator', 'snapshot')) assert program.relations == {'zone_of': declared}, 'the relation sits once in the program, under its name' zonal = program.constraints['zonal'].lhs - assert zonal == GroupSum( - Variable('p'), join=Join('zone_of', declared, ('generator', 'snapshot'), ('zone', 'snapshot')) - ), ( - 'a grouped sum joins on the over= column and the unnamed key column, and groups by the into= column and that key column' + columns = JoinColumns('zone_of', declared, ('generator', 'snapshot'), ('zone', 'snapshot')) + assert zonal == Sum(Join(Variable('p'), columns), ('generator',)), ( + 'a grouped sum is a sum over a join: the join names the over= column and the unnamed key column as joined on, ' + 'the into= column and that key column as grouped by, and the sum stands over the dim the join drops' ) - assert isinstance(zonal, GroupSum) - assert (zonal.join.dropped_dims, zonal.join.added_dims, zonal.join.kept) == ( + assert isinstance(zonal, Sum) and isinstance(zonal.operand, Join) + assert (zonal.operand.columns.dropped_dims, zonal.operand.columns.added_dims, zonal.operand.columns.kept) == ( ('generator',), ('zone',), ('snapshot',), ), 'the dims a consumer reads are read off the join: dropped, added, and the key columns kept' - assert zonal.join.relation is program.relations['zone_of'], ( + assert zonal.operand.columns.relation is program.relations['zone_of'], ( 'the join holds the one declaration the program holds, not an equal copy built again' ) - assert program.constraints['history'].lhs == GroupSum( - Variable('p'), join=Join('zone_of', declared, ('snapshot', 'generator'), ('zone', 'generator')) + assert program.constraints['history'].lhs == Sum( + Join(Variable('p'), JoinColumns('zone_of', declared, ('snapshot', 'generator'), ('zone', 'generator'))), + ('snapshot',), ), 'the same table joined on its other key column' priced = program.constraints['priced'].rhs - assert priced == Lookup( - Parameter('price'), join=Join('zone_of', declared, ('zone', 'snapshot'), ('generator', 'snapshot')) - ), 'and the lookup joins on the value column and groups by the key column' - assert isinstance(priced, Lookup) - assert (priced.join.dropped_dims, priced.join.added_dims, priced.join.kept) == ( + assert priced == Join( + Parameter('price'), JoinColumns('zone_of', declared, ('zone', 'snapshot'), ('generator', 'snapshot')) + ), 'and an at is the bare join, on the value column, grouped by the key column' + assert isinstance(priced, Join) + assert (priced.columns.dropped_dims, priced.columns.added_dims, priced.columns.kept) == ( ('zone',), ('generator',), ('snapshot',), @@ -739,13 +739,13 @@ def test_a_binary_variable_lowers_to_a_binary_domain(): assert program.variables['dispatch'].domain == 'binary' -def test_a_divisor_under_a_lookup_is_still_named(): +def test_a_divisor_under_a_join_is_still_named(): """`children` has to descend through every node, or a refusal loses its name.""" quotient = Divide(Variable('x'), Parameter('rate')) component_of = RelationDeclaration((('flow', 'flow'), ('component', 'component')), ('flow',)) - looked_up = Lookup(quotient, join=Join('component_of', component_of, ('component',), ('flow',))) + looked_up = Join(quotient, JoinColumns('component_of', component_of, ('component',), ('flow',))) - assert divisor_parameters(looked_up) == frozenset({'rate'}), 'the walk descends through `Lookup`' + assert divisor_parameters(looked_up) == frozenset({'rate'}), 'the walk descends through `Join`' assert divisor_parameters(Sum(looked_up, ('flow',))) == frozenset({'rate'}), 'and through a `Sum` over it' @@ -818,8 +818,8 @@ def test_walk_is_the_node_column_of_walk_regions(): Power(Parameter('c'), Constant(2.0)): 'one-to-one', Divide(Variable('p'), Parameter('c')): 'one-to-one', Sum(Variable('p'), ('g',)): 'many-to-one', - GroupSum(Variable('p'), join=Join('at_bus', AT_BUS, ('g',), ('bus',))): 'many-to-one', - Lookup(Variable('p'), join=Join('at_bus', AT_BUS, ('bus',), ('g',))): 'one-to-one', + Sum(Join(Variable('p'), JoinColumns('at_bus', AT_BUS, ('g',), ('bus',))), ('g',)): 'many-to-one', + Join(Variable('p'), JoinColumns('at_bus', AT_BUS, ('bus',), ('g',))): 'one-to-one', Translate(Variable('p'), 't', offset=1, wrap=False, fill=0.0): 'one-to-one', WindowSum(Variable('p'), 't', width=2, wrap=False): 'one-to-many', Cases((Region(Mask(ParameterDefined('c', ('g',))), Variable('p')),)): 'one-to-one', diff --git a/tests/test_parser.py b/tests/test_parser.py index 1db9619a..11089ced 100644 --- a/tests/test_parser.py +++ b/tests/test_parser.py @@ -47,7 +47,7 @@ from math_spec.program import ( And, BooleanLiteral, - Join, + JoinColumns, Not, Or, Partition, @@ -554,7 +554,7 @@ def test_a_node_prints_as_the_file_writes_it(text, printed): pytest.param(DimensionNode('t'), 't', id='a-dimension'), pytest.param(DualNode('budget'), 'dual(budget)', id='a-dual'), pytest.param( - JoinNode(Join('zone_of', _ZONE_OF, ('u',), ('zone',))), + JoinNode(JoinColumns('zone_of', _ZONE_OF, ('u',), ('zone',))), 'zone_of', id='a-relation-as-a-call-joins-it', ), diff --git a/tests/typesetting/test_golden.py b/tests/typesetting/test_golden.py index 21744688..03a7c7f1 100644 --- a/tests/typesetting/test_golden.py +++ b/tests/typesetting/test_golden.py @@ -160,12 +160,12 @@ def _rendered_trees() -> Iterator[object]: } #: A dataclass the walk steps *through* rather than renders: an arm has no -#: branch of its own — its ``when`` and ``value`` do — a join 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', 'Join', 'Mask', 'Partition', 'RelationDeclaration'} +#: branch of its own — its ``when`` and ``value`` do — the columns a join +#: names 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', 'JoinColumns', 'Mask', 'Partition', 'RelationDeclaration'} def test_the_golden_model_carries_every_node_kind_the_walk_renders(): From 5cbb7acc24948af2508ec5cf8affe848aaee629c Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 22 Sep 2026 19:58:04 +0000 Subject: [PATCH 06/18] chore: a rule the package wrote twice is written once The duplicate-keyword refusal, the name-or-list read, a translation's step count, a raw-model section and the predicate union each have one home. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_017XwKgY5wZXv1bKgkfCW2q1 --- src/math_spec/_expression_parser.py | 28 +++++++++++--------- src/math_spec/_where_parser.py | 11 ++------ src/math_spec/lowering.py | 28 +++++++++----------- src/math_spec/piecewise.py | 18 +++++-------- src/math_spec/program.py | 41 +++++++++-------------------- src/math_spec/resolution.py | 8 +++--- src/math_spec/sos.py | 10 +++---- src/math_spec/typesetting/README.md | 2 +- 8 files changed, 57 insertions(+), 89 deletions(-) diff --git a/src/math_spec/_expression_parser.py b/src/math_spec/_expression_parser.py index 5fe55bc9..511468a8 100644 --- a/src/math_spec/_expression_parser.py +++ b/src/math_spec/_expression_parser.py @@ -21,7 +21,7 @@ from math_spec.operators import EDGE_WRAP if TYPE_CHECKING: - from collections.abc import Callable, Iterator, Mapping + from collections.abc import Callable, Iterable, Iterator, Mapping from math_spec.program import Direction, Partition, Predicate @@ -462,18 +462,22 @@ def _build_grammar() -> tuple[pp.ParserElement, pp.ParserElement]: def _make_func_call(tokens: pp.ParseResults) -> FunctionCallNode: """The callee is cast: a ParseResults element is untyped, and the grammar guarantees an identifier in position 0.""" name = cast('str', tokens[0]) - args = [] - kwargs = {} + args: list[ArithmeticNode] = [] + pairs: list[tuple[str, ArithmeticNode]] = [] for item in tokens[1:]: - if isinstance(item, tuple) and len(item) == 2: - k, v = item - if k in kwargs: - msg = f'{name}({k}=) is given twice. A keyword names one value; drop one of them.' - raise SchemaError(msg) - kwargs[k] = v - else: - args.append(item) - return FunctionCallNode(name=name, args=tuple(args), kwargs=kwargs) + (pairs if isinstance(item, tuple) and len(item) == 2 else args).append(item) + return FunctionCallNode(name=name, args=tuple(args), kwargs=keywords(name, pairs)) + + +def keywords[V](name: str, pairs: Iterable[tuple[str, V]]) -> dict[str, V]: + """A call's keywords, in the order written; a keyword given twice is refused, in both grammars.""" + kwargs: dict[str, V] = {} + for key, value in pairs: + if key in kwargs: + msg = f'{name}({key}=) is given twice. A keyword names one value; drop one of them.' + raise SchemaError(msg) + kwargs[key] = value + return kwargs def _make_left_assoc(tokens: pp.ParseResults) -> ArithmeticNode: diff --git a/src/math_spec/_where_parser.py b/src/math_spec/_where_parser.py index 6c99399e..d19e7c64 100644 --- a/src/math_spec/_where_parser.py +++ b/src/math_spec/_where_parser.py @@ -20,9 +20,8 @@ import pyparsing as pp -from math_spec._expression_parser import ARITHMETIC, NAME, ArithmeticNode, children, parse_text +from math_spec._expression_parser import ARITHMETIC, NAME, ArithmeticNode, children, keywords, parse_text from math_spec._sealed import Sealed -from math_spec.errors import SchemaError from math_spec.program import ( And, BooleanLiteral, @@ -140,13 +139,7 @@ class UnresolvedComparisonNode: def _predicate_call(tokens: pp.ParseResults) -> UnresolvedPredicateCallNode: """The call node, with a keyword given twice refused as the arithmetic grammar refuses it.""" name, operand, *pairs = tokens - kwargs: dict[str, ArithmeticNode] = {} - for key, value in pairs: - if key in kwargs: - msg = f'{name}({key}=) is given twice. A keyword names one value; drop one of them.' - raise SchemaError(msg) - kwargs[key] = value - return UnresolvedPredicateCallNode(name, operand, kwargs) + return UnresolvedPredicateCallNode(name, operand, keywords(name, pairs)) def _build_where_grammar() -> pp.ParserElement: diff --git a/src/math_spec/lowering.py b/src/math_spec/lowering.py index b82f2b04..12feaf83 100644 --- a/src/math_spec/lowering.py +++ b/src/math_spec/lowering.py @@ -347,16 +347,11 @@ def sum_back(self, node: FunctionCallNode) -> program.Expression: """ over_node = node.kwargs['along'] assert isinstance(over_node, DimensionNode), 'resolution refuses an along= that is not a dimension' - window_node = node.kwargs['window'] operand = self.expr(node.args[0]) wrap = isinstance(node.kwargs.get('edge'), EdgeNode) - width: int | str - if isinstance(window_node, ParameterNode): - width = window_node.name - else: - assert isinstance(window_node, NumberNode), 'a window= that is neither is refused at load' - width = int(window_node.value) - return program.WindowSum(operand, over_node.name, width=width, wrap=wrap, partition=_partition_of(node)) + 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. @@ -366,19 +361,12 @@ def shift(self, node: FunctionCallNode) -> program.Expression: """ over_node = node.kwargs['along'] assert isinstance(over_node, DimensionNode), 'resolution refuses an along= that is not a dimension' - by_node = node.kwargs['offset'] operand = self.expr(node.args[0]) edge = node.kwargs.get('edge') - by: int | str - if isinstance(by_node, ParameterNode): - by = by_node.name - else: - assert isinstance(by_node, NumberNode), 'an offset= that is neither is refused at load' - by = int(by_node.value) return program.Translate( operand, over_node.name, - offset=by, + offset=_amount(node.kwargs['offset']), wrap=isinstance(edge, EdgeNode), fill=edge.value if isinstance(edge, NumberNode) else None, partition=_partition_of(node), @@ -394,6 +382,14 @@ def shift(self, node: FunctionCallNode) -> program.Expression: } +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. diff --git a/src/math_spec/piecewise.py b/src/math_spec/piecewise.py index 9788d57a..4ed53f4a 100644 --- a/src/math_spec/piecewise.py +++ b/src/math_spec/piecewise.py @@ -23,7 +23,7 @@ from math_spec.model import Curvature, PiecewiseBlock, PiecewiseMethod, Spec, undeclared_dimension from math_spec.program import PiecewiseDeclaration from math_spec.resolution import Namespace, resolve_expression -from math_spec.sos import Emitted, emit +from math_spec.sos import Emitted, emit, section if TYPE_CHECKING: from collections.abc import Iterable, Iterator @@ -265,31 +265,25 @@ 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. """ - section = self._section('assumptions') + assumptions = 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 - section[name] = entry + assumptions[name] = entry # -- emitters ---------------------------------------------------------- - def _section(self, name: str) -> dict[str, object]: - """The *name* section of the raw model, created empty where the file declares none.""" - section = self.raw.setdefault(name, {}) - assert isinstance(section, dict), f'{name}: is a mapping in a validated model' - return section - def _weight(self, name: str, **fields: object) -> None: """A variable over the frame and the breakpoint dim, masked as the block is.""" - self._section('variables')[name] = { + 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: - self._section('constraints')[name] = { + section(self.raw, 'constraints')[name] = { 'dims': dims, **({'where': where} if where else {}), 'expression': expression, @@ -313,7 +307,7 @@ def _weights(self) -> None: f'({link.expression}) {link.sign} sum({self.lam} * {link.values}, over={d})', ) if self.pw.method in ('sos2', 'adjacency'): - self._section('sos')[self.name] = {'variable': self.lam, 'over': d, 'type': 2} + section(self.raw, 'sos')[self.name] = {'variable': self.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)``. diff --git a/src/math_spec/program.py b/src/math_spec/program.py index 21cb1220..0fc9af79 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -124,7 +124,7 @@ #: How a shape operator's output rows relate to its input slots, answered by #: :func:`fan_in` for every node. FanIn = Literal['one-to-one', 'many-to-one', 'one-to-many'] -ObjectiveSense = Literal['minimize', 'maximize'] +ObjectiveSense = _model.ObjectiveSense #: Where a degree-2 product may stand in the math a solver sees. An objective #: and a constraint take ``variable * variable``; a bound and a ``piecewise:`` @@ -1222,32 +1222,6 @@ class Or: right: Predicate -#: Every resolved predicate node. A lowered mask's ``root`` holds every member -#: but :class:`ArithmeticComparison`, which lowering rewrites into an -#: :class:`ExpressionComparison`, so a consumer walking a program never meets -#: one. The parser's ``Unresolved*`` nodes are not members: they live with the -#: grammar in :mod:`math_spec._where_parser`, and resolution rewrites them away -#: before anything here is asked. -Predicate = ( - BooleanLiteral - | DimensionPosition - | ParameterDefined - | VariableDefined - | ParameterComparison - | ExpressionComparison - | ArithmeticComparison - | DimensionComparison - | RelationComparison - | RelationPairComparison - | RelationDefined - | CountComparison - | TranslatedPredicate - | PulledBackPredicate - | Not - | And - | Or -) - #: Every predicate resolution has typed: it names a declaration and the kind is #: settled. Resolution passes these straight through, having nothing left to #: decide about them. @@ -1273,6 +1247,14 @@ class Or: #: tree shares them — the transient impurity resolution normalizes away. Connective = Not | And | Or +#: Every resolved predicate node. A lowered mask's ``root`` holds every member +#: but :class:`ArithmeticComparison`, which lowering rewrites into an +#: :class:`ExpressionComparison`, so a consumer walking a program never meets +#: one. The parser's ``Unresolved*`` nodes are not members: they live with the +#: grammar in :mod:`math_spec._where_parser`, and resolution rewrites them away +#: before anything here is asked. +Predicate = BooleanLiteral | TypedPredicate | Connective + def where_children(where: Predicate) -> tuple[Predicate, ...]: """The predicates under *where* — a connective's operands, and nothing under a leaf. @@ -1327,6 +1309,9 @@ def _atom_dims(atom: TypedPredicate) -> frozenset[str]: | VariableDefined() | CountComparison() | TranslatedPredicate() + | RelationComparison() + | RelationPairComparison() + | RelationDefined() | PulledBackPredicate() ): return frozenset(atom.dims) @@ -1334,8 +1319,6 @@ def _atom_dims(atom: TypedPredicate) -> frozenset[str]: return frozenset({atom.name}) case DimensionPosition(): return frozenset({atom.name, *(atom.partition.joined_dims if atom.partition is not None else ())}) - case RelationComparison() | RelationPairComparison() | RelationDefined(): - return frozenset(atom.dims) case _: assert_never(atom) diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index 02d5326b..1f82ddf4 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -599,10 +599,8 @@ def _relation_ref( 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.""" - if isinstance(value, NameNode): - return (value.name,) - if isinstance(value, NameListNode): - return value.names + if names := names_in(value): + return names self.errors.append( f'{self.context}: {operator}({key}=...) names columns of the relation — a bare name, or a list of them.' ) @@ -1311,7 +1309,7 @@ def _position_shape(call: FunctionCallNode) -> tuple[str, str | None, tuple[str, return None if within is not None and not isinstance(within, NameNode | NameListNode): return None - into = within.names if isinstance(within, NameListNode) else (within.name,) if within is not None else None + into = names_in(within) if within is not None else None return call.args[0].name, by.name if by is not None else None, into diff --git a/src/math_spec/sos.py b/src/math_spec/sos.py index 6e2df8a3..2ac7f635 100644 --- a/src/math_spec/sos.py +++ b/src/math_spec/sos.py @@ -105,23 +105,23 @@ def emit(raw: dict[str, object], name: str) -> None: it runs over. name: Which set to lower. """ - sets = _section(raw, 'sos') + sets = section(raw, 'sos') block = sets.pop(name) assert isinstance(block, dict), 'a validated model carries each set as a mapping' variable, over, order = block['variable'], block['over'], block['type'] - member = _section(raw, 'variables')[variable] + member = section(raw, 'variables')[variable] assert isinstance(member, dict), 'a validated model carries each variable as a mapping' dims = list(member['dims']) emitted = Emitted.of(name, order) - _section(raw, 'variables')[emitted.seg] = { + section(raw, 'variables')[emitted.seg] = { 'dims': dims, **({'where': member['where']} if member.get('where') else {}), 'domain': 'binary', 'description': _SEGMENTS[order], } picked = emitted.seg if order == 1 else f'{emitted.seg} + shift({emitted.seg}, along={over}, offset=1, edge=0)' - constraints = _section(raw, 'constraints') + constraints = section(raw, 'constraints') constraints[emitted.pick] = { 'dims': [d for d in dims if d != over], 'expression': f'sum({emitted.seg}, over={over}) <= 1', @@ -157,7 +157,7 @@ def _coefficients(member: dict[str, object]) -> tuple[float | str, float | str]: return below, above -def _section(raw: dict[str, object], name: str) -> dict[str, object]: +def section(raw: dict[str, object], name: str) -> dict[str, object]: """The *name* section of the raw model, created empty where the file declares none.""" section = raw.setdefault(name, {}) assert isinstance(section, dict), f'{name}: is a mapping in a validated model' diff --git a/src/math_spec/typesetting/README.md b/src/math_spec/typesetting/README.md index 173b9362..372b718b 100644 --- a/src/math_spec/typesetting/README.md +++ b/src/math_spec/typesetting/README.md @@ -3,7 +3,7 @@ SPDX-FileCopyrightText: math-spec Contributors SPDX-License-Identifier: MIT --> -# `typeset/` — the model, printed +# `typesetting/` — the model, printed This package is a consumer of the resolved core syntax tree. It builds no model and binds no data. It walks the typed tree that `to_spec` validates, and prints From 638a2bfc4cf493abe0fd3158169c55882ca92608 Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 22 Sep 2026 20:00:48 +0000 Subject: [PATCH 07/18] chore: a where string's bare name and quoted label are the expression grammar's nodes UnresolvedNameNode was NameNode and QuotedNode was KeywordNode under other names. A predicate call's kwargs are now typed as what they hold. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_017XwKgY5wZXv1bKgkfCW2q1 --- src/math_spec/_where_parser.py | 47 +++++++++++++------------------- src/math_spec/resolution.py | 16 +++++------ tests/test_parser.py | 21 ++++++-------- tests/typesetting/test_golden.py | 8 ++---- 4 files changed, 38 insertions(+), 54 deletions(-) diff --git a/src/math_spec/_where_parser.py b/src/math_spec/_where_parser.py index d19e7c64..ef8d1038 100644 --- a/src/math_spec/_where_parser.py +++ b/src/math_spec/_where_parser.py @@ -20,7 +20,16 @@ import pyparsing as pp -from math_spec._expression_parser import ARITHMETIC, NAME, ArithmeticNode, children, keywords, parse_text +from math_spec._expression_parser import ( + ARITHMETIC, + NAME, + ArithmeticNode, + KeywordNode, + NameNode, + children, + keywords, + parse_text, +) from math_spec._sealed import Sealed from math_spec.program import ( And, @@ -41,13 +50,6 @@ # --------------------------------------------------------------------------- -@dataclass(frozen=True) -class UnresolvedNameNode: - """A bare name — unresolved. ``resolution.py`` types it.""" - - name: str - - @dataclass(frozen=True) class ColumnNode: """``relation.column`` on a side of a comparison — the one place the language names a column.""" @@ -61,18 +63,6 @@ def shown(self) -> str: return f'{self.relation}.{self.column}' -@dataclass(frozen=True) -class QuotedNode: - """A right-hand side that arrived in quotes. - - A bare word is ambiguous — it may name a declaration — and resolution - refuses it for that reason; a quoted one is unambiguously a label, which - is the only way to write ``combined-cycle`` or a date. - """ - - value: str - - @dataclass(frozen=True) class UnresolvedPredicateCallNode: """``([, …])`` — an operator reading a predicate rather than arithmetic. @@ -113,22 +103,23 @@ class UnresolvedComparisonNode: A side is the expression grammar's arithmetic, so a name, a number and a ``position(...)`` call all arrive as the nodes an expression would carry - them in; a relation column and a quoted label have nodes of their own. + them in, a quoted label as the :class:`KeywordNode` a quoted kwarg is, and a + relation column in a node of its own. """ left: ArithmeticNode | ColumnNode op: PredicateOperator - right: ArithmeticNode | ColumnNode | QuotedNode + right: ArithmeticNode | ColumnNode | KeywordNode #: What resolution rewrites away on the where side — the nodes whose leaves #: are still names the schema has not been asked about. -UnresolvedWhereNode = UnresolvedNameNode | UnresolvedComparisonNode | UnresolvedPredicateCallNode | UnresolvedCountNode +UnresolvedWhereNode = NameNode | UnresolvedComparisonNode | UnresolvedPredicateCallNode | UnresolvedCountNode #: Every node a parsed where string is built of: the connectives and literals, #: the unresolved leaves, and the arithmetic and the two side nodes under a #: comparison. What the depth measurement walks. -_ParsedWhere = Predicate | UnresolvedWhereNode | ArithmeticNode | ColumnNode | QuotedNode +_ParsedWhere = Predicate | UnresolvedWhereNode | ArithmeticNode | ColumnNode | KeywordNode # --------------------------------------------------------------------------- @@ -159,7 +150,7 @@ def _build_where_grammar() -> pp.ParserElement: column = pp.Regex(rf'({NAME})\.({NAME})').set_parse_action(lambda t: ColumnNode(*t[0].split('.'))) quoted = (pp.QuotedString("'", esc_char='\\') | pp.QuotedString('"', esc_char='\\')).set_parse_action( # pyrefly: ignore[implicit-any-lambda] - lambda t: QuotedNode(t[0]) + lambda t: KeywordNode(t[0]) ) comparator = pp.one_of(list(get_args(PredicateOperator))) @@ -190,7 +181,7 @@ def _call(head: pp.ParserElement) -> pp.ParserElement: # pyrefly: ignore[implicit-any-lambda] ).set_parse_action(lambda t: UnresolvedComparisonNode(t[0], t[1], t[2])) # pyrefly: ignore[implicit-any-lambda] - existence = name.copy().set_parse_action(lambda t: UnresolvedNameNode(t[0])) + existence = name.copy().set_parse_action(lambda t: NameNode(t[0])) atom = ( true_lit @@ -281,8 +272,8 @@ def parse_where(text: str) -> Predicate | UnresolvedWhereNode: """Parse a where string into an AST, its leaves still unresolved. The connectives and literals are the resolved vocabulary's own; the leaves - naming declarations are ``Unresolved*`` nodes, which only - :func:`~math_spec.resolution.resolve_where` takes. + naming declarations are a bare :class:`NameNode` or an ``Unresolved*`` + node, which only :func:`~math_spec.resolution.resolve_where` takes. Raises: SchemaError: If *text* is not a where string of the language. A diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index 1f82ddf4..7f0d5c6f 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -47,10 +47,8 @@ ) from math_spec._where_parser import ( ColumnNode, - QuotedNode, UnresolvedComparisonNode, UnresolvedCountNode, - UnresolvedNameNode, UnresolvedPredicateCallNode, UnresolvedWhereNode, parse_where, @@ -758,7 +756,7 @@ def where(self, node: Predicate | UnresolvedWhereNode) -> Predicate | Unresolved """One predicate node typed, or returned unresolved with its refusal appended.""" if isinstance(node, BooleanLiteral | TypedPredicate): return node - if isinstance(node, UnresolvedNameNode): + if isinstance(node, NameNode): return self._where_name(node) if isinstance(node, UnresolvedComparisonNode): return self._comparison(node) @@ -778,7 +776,7 @@ def _child(self, node: Predicate | UnresolvedWhereNode) -> Predicate: """A connective's child, typed as resolved: an unresolved one survives only with its refusal appended.""" return cast('Predicate', self.where(node)) - def _where_name(self, node: UnresolvedNameNode) -> Predicate | UnresolvedWhereNode: + def _where_name(self, node: NameNode) -> Predicate | UnresolvedWhereNode: """A bare name: a parameter's or relation's definedness, or a variable's existence.""" ns, context = self.ns, self.context kind = ns.kind(node.name) @@ -965,8 +963,8 @@ def _plain(self, node: UnresolvedComparisonNode) -> _Plain | None: ns = self.ns name, right = _side_name(node.left), node.right value: float | str | None - quoted = isinstance(right, QuotedNode) - if isinstance(right, QuotedNode): + quoted = isinstance(right, KeywordNode) + if isinstance(right, KeywordNode): value = right.value elif isinstance(right, ColumnNode): value = right.shown @@ -991,7 +989,7 @@ def _expression_comparison(self, node: UnresolvedComparisonNode) -> ArithmeticCo found = len(self.errors) sides = [] for side in (node.left, node.right): - if isinstance(side, ColumnNode | QuotedNode): + if isinstance(side, ColumnNode | KeywordNode): self.errors.append(_not_arithmetic(context, side)) continue if any(isinstance(n, FunctionCallNode) and n.name == 'count' for n in nodes(side)): @@ -1051,7 +1049,7 @@ def _position( ) return node dimension, by, into = shape - index = None if isinstance(node.right, ColumnNode | QuotedNode) else _literal(node.right) + index = None if isinstance(node.right, ColumnNode | KeywordNode) else _literal(node.right) if index is None or not index.value.is_integer(): self.errors.append( f'{context}: position({dimension}) is compared against an integer index, where 0 is first and a ' @@ -1376,7 +1374,7 @@ def _literal(value: ArithmeticNode) -> NumberNode | None: return None -def _not_arithmetic(context: str, side: ColumnNode | QuotedNode) -> str: +def _not_arithmetic(context: str, side: ColumnNode | KeywordNode) -> str: """Why a relation column or a quoted label may not stand on a side of a comparison of expressions.""" if isinstance(side, ColumnNode): return ( diff --git a/tests/test_parser.py b/tests/test_parser.py index d0d73806..c7593b32 100644 --- a/tests/test_parser.py +++ b/tests/test_parser.py @@ -27,6 +27,7 @@ DualNode, EdgeNode, FunctionCallNode, + KeywordNode, NameListNode, NameNode, NumberNode, @@ -38,9 +39,7 @@ ) from math_spec._where_parser import ( ColumnNode, - QuotedNode, UnresolvedComparisonNode, - UnresolvedNameNode, parse_where, ) from math_spec.errors import SchemaError @@ -263,7 +262,7 @@ def test_a_name_may_begin_with_inf(name): ('text', 'node_type', 'attrs'), [ pytest.param('True', BooleanLiteral, {'value': True}, id='a-literal'), - pytest.param('p_max', UnresolvedNameNode, {'name': 'p_max'}, id='a-bare-name'), + pytest.param('p_max', NameNode, {'name': 'p_max'}, id='a-bare-name'), pytest.param('p_max > 0', UnresolvedComparisonNode, {'op': '>', 'right': NumberNode(0)}, id='a-comparison'), pytest.param('a AND b', And, {}, id='and'), pytest.param('a OR b', Or, {}, id='or'), @@ -280,9 +279,7 @@ def test_a_where_string_parses_to_its_node(text, node_type, attrs): def test_and_binds_tighter_than_or(): - assert parse_where('a OR b AND c') == Or( - UnresolvedNameNode('a'), And(UnresolvedNameNode('b'), UnresolvedNameNode('c')) - ) + assert parse_where('a OR b AND c') == Or(NameNode('a'), And(NameNode('b'), NameNode('c'))) @pytest.mark.parametrize( @@ -315,12 +312,12 @@ def test_conjuncts_does_not_split_or_or_not(text): @pytest.mark.parametrize( ('text', 'right'), [ - ("g == 'wind'", QuotedNode('wind')), - ('g == "wind"', QuotedNode('wind')), - ("g == 'combined-cycle'", QuotedNode('combined-cycle')), - ("g == 'CCGT 400MW'", QuotedNode('CCGT 400MW')), - ("t > '2030-01-01'", QuotedNode('2030-01-01')), - ("g == 'it\\'s'", QuotedNode("it's")), + ("g == 'wind'", KeywordNode('wind')), + ('g == "wind"', KeywordNode('wind')), + ("g == 'combined-cycle'", KeywordNode('combined-cycle')), + ("g == 'CCGT 400MW'", KeywordNode('CCGT 400MW')), + ("t > '2030-01-01'", KeywordNode('2030-01-01')), + ("g == 'it\\'s'", KeywordNode("it's")), ('g == wind', NameNode('wind')), ], ids=['single', 'double', 'hyphen', 'space', 'date', 'escaped quote', 'bare'], diff --git a/tests/typesetting/test_golden.py b/tests/typesetting/test_golden.py index 5d86b95c..253bb457 100644 --- a/tests/typesetting/test_golden.py +++ b/tests/typesetting/test_golden.py @@ -142,17 +142,15 @@ def _rendered_trees() -> Iterator[object]: yield from links -#: What resolution never hands the walk: the four nodes a where carries before -#: its sides are read, the three an expression only carries before names are -#: resolved, and the lowered form of a comparison of expressions, which only a +#: What resolution never hands the walk: the two nodes a where carries before +#: its sides are read, the three an expression and a where carry before names +#: are resolved, and the lowered form of a comparison of expressions, which only a #: program carries. 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 = { - 'UnresolvedNameNode', 'UnresolvedComparisonNode', 'ColumnNode', - 'QuotedNode', 'NameNode', 'NameListNode', 'KeywordNode', From 3d7584dad3636fa4dc5be7598cde84ebc4845f5f Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 22 Sep 2026 20:04:07 +0000 Subject: [PATCH 08/18] fix(language): a macro template nothing calls is held to every rule a call site is The template check was a second walker that looked at names only. It is now the resolver, with the template's formals left bare, so a label parameter used as a value or an unknown relation column is refused at load. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_017XwKgY5wZXv1bKgkfCW2q1 --- src/math_spec/resolution.py | 29 +++++++++--- src/math_spec/validation.py | 89 +------------------------------------ tests/test_expansion.py | 35 +++++++++++++-- 3 files changed, 58 insertions(+), 95 deletions(-) diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index 7f0d5c6f..cfddcb15 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -288,16 +288,21 @@ def resolve_expression( 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. + 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. """ before = len(errors) - resolved = _Resolver(ns, context, errors).expression(node) + resolved = _Resolver(ns, context, errors, formals=formals).expression(node) return None if len(errors) > before else resolved @@ -350,13 +355,19 @@ class _Resolver: 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. + 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. """ ns: Namespace context: str errors: list[str] self_variable: str | None = None + formals: frozenset[str] = frozenset() + + def _formal(self, value: ArithmeticNode) -> bool: + """Whether *value* is a formal, left for the call site to bind.""" + return isinstance(value, NameNode) and value.name in self.formals # -- expressions ------------------------------------------------------- @@ -374,7 +385,7 @@ def _arith(self, node: ArithmeticNode, *, amount: bool = False) -> ArithmeticNod numeric check here stands aside for it. 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): + if isinstance(node, NumberNode | VariableNode | ParameterNode | DualNode | KwargNode) or self._formal(node): return node if isinstance(node, NameNode): return self._name(node, amount=amount) @@ -429,7 +440,7 @@ def _name(self, node: NameNode, *, amount: bool) -> ArithmeticNode: ) return node case _: - self.errors.append(self.ns.unknown(node.name, self.context, allow_dims=False)) + self.errors.append(self.ns.unknown(node.name, self.context, allow_dims=False, formals=self.formals)) return node def _call(self, node: FunctionCallNode) -> ArithmeticNode: @@ -495,6 +506,8 @@ def _amount(self, value: ArithmeticNode, operator: str, key: str) -> ArithmeticN def _edge(self, value: ArithmeticNode, operator: str) -> ArithmeticNode: """``edge=``: the closed keyword ``wrap``, or a number to contribute; a name here is a typo.""" + if self._formal(value): + return value if isinstance(value, KeywordNode): if value.value == EDGE_WRAP: return EdgeNode() @@ -519,6 +532,8 @@ def _edge(self, value: ArithmeticNode, operator: str) -> ArithmeticNode: def _dim_ref(self, value: ArithmeticNode, operator: str, key: str) -> ArithmeticNode: """An operator kwarg whose *value* must name a declared dimension.""" + if self._formal(value): + return value if not isinstance(value, NameNode): self.errors.append(f'{self.context}: {operator}({key}=...) must name a dimension.') return value @@ -536,6 +551,8 @@ def _dual(self, node: FunctionCallNode) -> ArithmeticNode: (:mod:`math_spec.validation`); this pass only types the name. """ (value,) = node.args + if self._formal(value): + return node if not isinstance(value, NameNode): self.errors.append( f'{self.context}: dual() takes the name of a declared constraint, written bare — ' @@ -543,7 +560,7 @@ def _dual(self, node: FunctionCallNode) -> ArithmeticNode: ) return node if value.name not in self.ns.constraints: - self.errors.append(self.ns.unknown_constraint(value.name, self.context)) + self.errors.append(self.ns.unknown_constraint(value.name, self.context, formals=self.formals)) return node return DualNode(value.name) @@ -576,6 +593,8 @@ def _relation_ref( ) return value name = names[0] + if name in self.formals or any(n in self.formals for v in roles.values() for n in names_in(v)): + return value if (problem := self._not_a_relation(name, operator, key)) is not None: self.errors.append(problem) diff --git a/src/math_spec/validation.py b/src/math_spec/validation.py index e976ebce..f62df8d4 100644 --- a/src/math_spec/validation.py +++ b/src/math_spec/validation.py @@ -7,29 +7,17 @@ from __future__ import annotations from collections.abc import Mapping -from typing import TYPE_CHECKING, Literal, assert_never, overload +from typing import TYPE_CHECKING, Literal, overload import math_spec.degree as degree from math_spec._expression_parser import ( ArithmeticNode, - BinaryOperatorNode, CaseArm, CasesNode, ComparisonNode, DefinitionNode, - DualNode, - FunctionCallNode, - KeywordNode, - KwargNode, - NameListNode, - NameNode, - NumberNode, - ParameterNode, ParsedNode, - UnaryOperatorNode, - VariableNode, case_context, - children, ) from math_spec._yaml import read_model from math_spec.dimensions import check_schema @@ -37,7 +25,6 @@ from math_spec.exclusivity import overlapping from math_spec.expansion import expand, parse_and_expand, parse_template from math_spec.model import AssumptionBlock, Spec -from math_spec.operators import BUILTINS, call_shape_error, unknown_operator_message from math_spec.piecewise import assumptions_of from math_spec.program import BooleanLiteral, Mask, VariableDefined from math_spec.resolution import ( @@ -46,7 +33,6 @@ ResolvedAssumption, ResolvedConstraint, mask_of, - names_in, resolve_expression, resolve_where_text, ) @@ -141,7 +127,7 @@ def validate_expressions(schema: Spec) -> Resolved: f'ambiguous with the dimension itself.' for f in sorted(formals & ns.dimensions) ) - _check_template_names(body_ast, context, ns, formals, errors) + resolve_expression(body_ast, ns, context, errors, formals=formals) expressions: dict[str, CasesNode | DefinitionNode] = {} for ename, block in schema.expressions.items(): @@ -373,74 +359,3 @@ def _check_expression( ) return None return resolved - - -def _check_template_names( - node: ArithmeticNode, - context: str, - ns: Namespace, - formals: frozenset[str], - errors: list[str], -) -> None: - """Check a macro body's names and call shapes, treating formals as bound — not resolution, since a formal has no kind until a call site binds it. - - An operator call is refused by its signature here, as at a call site, so a - keyword the operator does not declare is caught in a template nothing calls. - A case arm's value only: its ``when`` is the declaration's, checked there. - """ - if isinstance(node, NumberNode | VariableNode | ParameterNode | DualNode | KwargNode | KeywordNode | NameListNode): - return - - if isinstance(node, NameNode): - if node.name not in formals and ns.kind(node.name) is None: - errors.append(ns.unknown(node.name, context, allow_dims=False, formals=formals)) - return - - if isinstance(node, UnaryOperatorNode | BinaryOperatorNode | CasesNode | DefinitionNode): - for child in children(node): - _check_template_names(child, context, ns, formals, errors) - return - - if isinstance(node, FunctionCallNode): - builtin = BUILTINS.get(node.name) - if builtin is None: - errors.append(f'{context}: {unknown_operator_message(node.name)}') - else: - shape_error = call_shape_error(node.name, len(node.args), node.kwargs) - if shape_error is not None: - errors.append(f'{context}: {shape_error}') - if node.name == 'dual': - errors.extend( - ns.unknown_constraint(arg.name, context, formals=formals) - for arg in node.args - if isinstance(arg, NameNode) and arg.name not in formals and arg.name not in ns.constraints - ) - return - for arg in node.args: - _check_template_names(arg, context, ns, formals, errors) - for kwarg, value in node.kwargs.items(): - with_relation = builtin is not None and any(k in node.kwargs for k in builtin.relation_kwargs) - match builtin.kind_of(kwarg, with_relation=with_relation) if builtin else 'value': - case 'dimension': - if isinstance(value, NameNode) and value.name not in ns.dimensions | formals: - errors.append( - f'{context}: {node.name}({kwarg}={value.name}) does not name a ' - f'declared dimension or a formal of this macro.' - ) - case 'relation': - errors.extend( - f'{context}: {node.name}({kwarg}={one}) does not name a relation or a formal of this macro.' - for one in names_in(value) - if one not in formals and ns.kind(one) != 'relation' - ) - case 'value': - _check_template_names(value, context, ns, formals, errors) - case 'role': - pass - case 'edge': - pass # a keyword or a number: nothing in it to name - case None: - pass # a keyword the operator does not declare; the shape error above named it - return - - assert_never(node) diff --git a/tests/test_expansion.py b/tests/test_expansion.py index b79bc6e6..f1eb067f 100644 --- a/tests/test_expansion.py +++ b/tests/test_expansion.py @@ -13,7 +13,7 @@ from math_spec._expression_parser import ComparisonNode, DefinitionNode, parse_expression, with_children from math_spec.errors import LanguageError from math_spec.expansion import parse_and_expand -from tests.fixtures import DISPATCH_MODEL, schema_of +from tests.fixtures import DISPATCH_MODEL, SMALL_MODEL, schema_of WEIGHTED_SUM = { 'args': ['array', 'weights'], @@ -223,8 +223,8 @@ def test_macro_collisions_rejected(patch, match): ), pytest.param( {'grouped': {'args': ['x'], 'template': 'sum(x, by=[nope, also])'}}, - r"Macro 'grouped'.*sum\(by=nope\) does not name a relation", - id='a-typo-in-a-relation-list', + r"Macro 'grouped'.*sum\(by=\[nope, also\]\) names 2 relations", + id='a-list-of-relations', ), ], ) @@ -234,8 +234,37 @@ def test_macro_templates_validated_even_when_unused(macros, match): schema(macros=macros) +@pytest.mark.parametrize( + ('template', 'match'), + [ + pytest.param('x * tag', "Macro 'm': 'tag' is declared dtype: str", id='a-label-parameter-as-a-value'), + pytest.param('sum(x, by=lk, over=nope, into=h)', "over=nope names no column of 'lk'", id='a-typo-in-a-column'), + ], +) +def test_a_template_is_held_to_the_rules_a_call_site_is(template, match): + """A template nothing calls was checked for names only: a label parameter or an unknown column passed load.""" + with pytest.raises(LanguageError, match=match): + schema_of(SMALL_MODEL, macros={'m': {'args': ['x'], 'template': template}}) + + @pytest.mark.parametrize('fragment', ['my_python_helper', 'macros:', 'escape']) def test_an_unknown_operator_is_refused_at_load_with_the_rewrite(fragment): with pytest.raises(LanguageError) as exc: schema(constraints={'c': {'dims': ['snapshot'], 'expression': 'my_python_helper(p) <= load'}}) assert fragment in str(exc.value) + + +@pytest.mark.parametrize( + ('formals', 'template'), + [ + pytest.param(['x', 'e'], 'shift(x, along=g, offset=1, edge=e)', id='an-edge'), + pytest.param(['row'], 'dual(row)', id='a-constraint'), + pytest.param(['x', 'rel', 'a', 'b'], 'sum(x, by=rel, over=a, into=b)', id='a-relation-and-its-columns'), + pytest.param(['x', 'a', 'b'], 'sum(x, by=lk, over=a, into=b)', id='the-columns-of-a-declared-relation'), + ], +) +def test_a_formal_stands_where_a_call_site_will_bind_it(formals, template): + """A formal has no kind until a call binds it, so the template check leaves it bare in every slot.""" + assert ( + schema_of(SMALL_MODEL, macros={'m': {'args': formals, 'template': template}}).macros['m'].template == template + ) From a57176de0ec420a40e9a404073d26d09e0d3bd6a Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 22 Sep 2026 20:06:09 +0000 Subject: [PATCH 09/18] chore(typeset): a legend section is titled the way an equation section is A format's glossary() returns the rows and section() sets the title, so the Glossary class and three copies of each heading go. cases_row leaves the Format protocol, since only a format's own cases() read it. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_017XwKgY5wZXv1bKgkfCW2q1 --- src/math_spec/typesetting/__init__.py | 4 ++-- src/math_spec/typesetting/format.py | 14 +++----------- src/math_spec/typesetting/latex.py | 7 +++---- src/math_spec/typesetting/markdown.py | 4 ++-- src/math_spec/typesetting/typst.py | 8 +++----- src/math_spec/typesetting/walk.py | 13 ++++--------- 6 files changed, 17 insertions(+), 33 deletions(-) diff --git a/src/math_spec/typesetting/__init__.py b/src/math_spec/typesetting/__init__.py index ab8c6e5b..38d1e585 100644 --- a/src/math_spec/typesetting/__init__.py +++ b/src/math_spec/typesetting/__init__.py @@ -151,7 +151,7 @@ def typeset( blocks = [format_.note(format_.escape(schema.description))] if schema.description else [] if legend: - blocks += [format_.glossary(group.title, group.entries) for group in walk.glossaries(noticed)] + blocks += [format_.section(title, format_.glossary(entries)) for title, entries in walk.glossaries(noticed)] blocks += [format_.note(text) for text in walk.convention_notes()] blocks += [format_.note(text) for text in walk.translation_notes(noticed)] blocks += [format_.note(text) for text in walk.position_notes(noticed)] @@ -194,7 +194,7 @@ def typeset_declaration( Raises: ValueError: *fmt* names no format. LanguageError: A model that does not compile; it does not print. - SchemaError: *name* is declared as none of the four, or as two — a + SchemaError: *name* is declared as none of the five, or as two — a constraint may share a variable's name; or a symbol table entry names nothing in the model. """ diff --git a/src/math_spec/typesetting/format.py b/src/math_spec/typesetting/format.py index 4b024614..b6cf4613 100644 --- a/src/math_spec/typesetting/format.py +++ b/src/math_spec/typesetting/format.py @@ -137,14 +137,6 @@ class Entry: meaning: str -@dataclass(frozen=True) -class Glossary: - """One legend section: its title, and the entries under it.""" - - title: str - entries: list[Entry] - - #: The one notation author prose carries: a name in backticks, set in monospace. _CODE_SPAN = re.compile(r'`([^`]+)`') @@ -169,8 +161,6 @@ class Format(Protocol): operators: ClassVar[Mapping[OperatorName, str]] #: The em dash in prose: TeX and Typst read ``---`` as one, Markdown does not. dash: ClassVar[str] - #: Between the rows of a ``cases`` block. - cases_row: ClassVar[str] # -- atoms ------------------------------------------------------------- @@ -255,7 +245,9 @@ def equation(self, line: Line) -> str: def equations(self, lines: list[Line], *, numbered: bool) -> str: ... - def glossary(self, title: str, entries: list[Entry]) -> str: ... + def glossary(self, entries: list[Entry]) -> str: + """A legend section's rows; :meth:`section` sets its title, as it does for the equations.""" + ... def section(self, title: str, body: str) -> str: ... diff --git a/src/math_spec/typesetting/latex.py b/src/math_spec/typesetting/latex.py index 3e3296f5..38c5c634 100644 --- a/src/math_spec/typesetting/latex.py +++ b/src/math_spec/typesetting/latex.py @@ -47,7 +47,6 @@ class LatexFormat: notation: ClassVar[Notation] = 'latex' #: TeX's own em-dash ligature. dash: ClassVar[str] = '---' - cases_row: ClassVar[str] = r' \\ ' operators: ClassVar[Mapping[OperatorName, str]] = {name: latex for name, (latex, _) in OPERATOR_SPELLINGS.items()} @@ -102,7 +101,7 @@ def set_of(self, members: str, condition: str) -> str: return rf'\{{ {members} {self.operators["such_that"]} {condition} \}}' def cases(self, arms: list[tuple[str, str]]) -> str: - rows = self.cases_row.join(f'{value} & {condition}' for value, condition in arms) + rows = r' \\ '.join(f'{value} & {condition}' for value, condition in arms) return rf'\begin{{cases}} {rows} \end{{cases}}' def summation(self, domain: str, body: str) -> str: @@ -125,9 +124,9 @@ def equations(self, lines: list[Line], *, numbered: bool) -> str: body = ' \\\\\n'.join(aligned_rows(lines, self, gap=' && ')) return f'\\begin{{{environment}}}\n{body}\n\\end{{{environment}}}' - def glossary(self, title: str, entries: list[Entry]) -> str: + def glossary(self, entries: list[Entry]) -> str: rows = '\n'.join(rf'\item[{{{self.math(e.symbol)}}}] {e.meaning}' for e in entries) - return f'\\paragraph{{{title}}}\n\\begin{{description}}\n{rows}\n\\end{{description}}' + return f'\\begin{{description}}\n{rows}\n\\end{{description}}' def section(self, title: str, body: str) -> str: return f'\\paragraph{{{title}}}\n{body}' diff --git a/src/math_spec/typesetting/markdown.py b/src/math_spec/typesetting/markdown.py index 7f80a280..9fb6b17c 100644 --- a/src/math_spec/typesetting/markdown.py +++ b/src/math_spec/typesetting/markdown.py @@ -86,9 +86,9 @@ def equations(self, lines: list[Line], *, numbered: bool) -> str: return '\n\n'.join(blocks) @override - def glossary(self, title: str, entries: list[Entry]) -> str: + def glossary(self, entries: list[Entry]) -> str: rows = '\n'.join(f'| {_cell(self.math(e.symbol))} | {_cell(e.meaning)} |' for e in entries) - return f'#### {title}\n\n| Symbol | Meaning |\n|---|---|\n{rows}' + return f'| Symbol | Meaning |\n|---|---|\n{rows}' @override def section(self, title: str, body: str) -> str: diff --git a/src/math_spec/typesetting/typst.py b/src/math_spec/typesetting/typst.py index d83a1aeb..b3c7cce0 100644 --- a/src/math_spec/typesetting/typst.py +++ b/src/math_spec/typesetting/typst.py @@ -55,7 +55,6 @@ class TypstFormat: notation: ClassVar[Notation] = 'typst' #: Typst applies the same substitution TeX does. dash: ClassVar[str] = '---' - cases_row: ClassVar[str] = ', ' operators: ClassVar[Mapping[OperatorName, str]] = {name: typst for name, (_, typst) in OPERATOR_SPELLINGS.items()} @@ -110,7 +109,7 @@ def set_of(self, members: str, condition: str) -> str: return f'{{{members} {self.operators["such_that"]} {condition}}}' def cases(self, arms: list[tuple[str, str]]) -> str: - return 'cases({})'.format(self.cases_row.join(f'{value} & {condition}' for value, condition in arms)) + return 'cases({})'.format(', '.join(f'{value} & {condition}' for value, condition in arms)) def summation(self, domain: str, body: str) -> str: return f'sum_({domain}) {body}' @@ -133,9 +132,8 @@ def equations(self, lines: list[Line], *, numbered: bool) -> str: numbering = '#set math.equation(numbering: "(1)")\n' if numbered else '' return f'{numbering}$ {body} $' - def glossary(self, title: str, entries: list[Entry]) -> str: - rows = '\n'.join(f'/ {self.math(e.symbol)}: {e.meaning}' for e in entries) - return f'== {title}\n{rows}' + def glossary(self, entries: list[Entry]) -> str: + return '\n'.join(f'/ {self.math(e.symbol)}: {e.meaning}' for e in entries) def section(self, title: str, body: str) -> str: return f'== {title}\n{body}' diff --git a/src/math_spec/typesetting/walk.py b/src/math_spec/typesetting/walk.py index a65c4846..d82e2266 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -58,7 +58,7 @@ TranslatedPredicate, VariableDefined, ) -from math_spec.typesetting.format import Entry, Glossary, Line, OperatorName +from math_spec.typesetting.format import Entry, Line, OperatorName if TYPE_CHECKING: import datetime @@ -1018,7 +1018,7 @@ def _sorted(self, dims: frozenset[str]) -> list[str]: # -- legend ------------------------------------------------------------ - def glossaries(self, noticed: Noticed) -> list[Glossary]: + def glossaries(self, noticed: Noticed) -> list[tuple[str, list[Entry]]]: fmt = self.format sets = [ self._entry( @@ -1041,13 +1041,8 @@ def glossaries(self, noticed: Noticed) -> list[Glossary]: for e, block in self.schema.expressions.items() if e in self._defined() ] - groups = ( - Glossary('Sets', sets), - Glossary('Parameters', parameters), - Glossary('Variables', variables), - Glossary('Definitions', definitions), - ) - return [group for group in groups if group.entries] + groups = (('Sets', sets), ('Parameters', parameters), ('Variables', variables), ('Definitions', definitions)) + return [(title, entries) for title, entries in groups if entries] def _entry(self, symbol: str, what: str, description: str | None) -> Entry: meaning = f'{what} {self.format.dash} {self.format.escape(description)}' if description else what From b279760c2d90250083d925019d7c7aa6355ff4e8 Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 22 Sep 2026 20:11:00 +0000 Subject: [PATCH 10/18] refactor(program): a comparison of expressions is one node before and after lowering ArithmeticComparison was ExpressionComparison with its sides still in the core syntax tree. ExpressionComparison now takes its side type as a parameter, and program.ArithmeticComparison is gone. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_017XwKgY5wZXv1bKgkfCW2q1 --- docs/reference/reading.md | 4 +--- src/math_spec/dimensions.py | 3 +-- src/math_spec/exclusivity.py | 29 ++++++++++++----------- src/math_spec/lowering.py | 2 +- src/math_spec/program.py | 38 +++++++------------------------ src/math_spec/resolution.py | 8 ++++--- src/math_spec/typesetting/walk.py | 10 +++----- tests/typesetting/test_golden.py | 6 ++--- 8 files changed, 35 insertions(+), 65 deletions(-) diff --git a/docs/reference/reading.md b/docs/reference/reading.md index 9ea57a2a..d5237914 100644 --- a/docs/reference/reading.md +++ b/docs/reference/reading.md @@ -138,9 +138,7 @@ 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. -Every `where` arrives as a `Mask`. Its `.root` is the resolved predicate. One -member of the `Predicate` union never reaches you. Lowering rewrites every -`ArithmeticComparison` into an `ExpressionComparison`. The +Every `where` arrives as a `Mask`. Its `.root` is the resolved predicate. The mask also answers four questions: - `.conjuncts` flattens the `AND` spine, and stops at an `OR` or a `NOT`. diff --git a/src/math_spec/dimensions.py b/src/math_spec/dimensions.py index 6ae0b740..f7faaf6b 100644 --- a/src/math_spec/dimensions.py +++ b/src/math_spec/dimensions.py @@ -41,7 +41,6 @@ from math_spec.errors import DimensionError from math_spec.operators import BUILTINS from math_spec.program import ( - ArithmeticComparison, CountComparison, DimensionComparison, DimensionPosition, @@ -561,7 +560,7 @@ def _check_where_dims( leaf = f"where-dimension '{atom.name}'" case RelationComparison() | RelationPairComparison() | RelationDefined(): leaf = f"where-relation '{atom.name}'" - case ArithmeticComparison() | ExpressionComparison(): + case ExpressionComparison(): leaf = 'a where-comparison of expressions' case CountComparison(): leaf = f"a where-count over '{atom.over}'" diff --git a/src/math_spec/exclusivity.py b/src/math_spec/exclusivity.py index 4c316867..930453b7 100644 --- a/src/math_spec/exclusivity.py +++ b/src/math_spec/exclusivity.py @@ -24,7 +24,6 @@ from math_spec._expression_parser import NumberNode, ParameterNode, UnaryOperatorNode from math_spec.program import ( And, - ArithmeticComparison, BooleanLiteral, CountComparison, DimensionComparison, @@ -47,6 +46,7 @@ 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 @@ -223,12 +223,12 @@ def _undecided(mask: Predicate) -> str | None: comparison of expressions falls before the numbers arrive. """ for atom in Mask(mask).atoms: - if isinstance(atom, ArithmeticComparison | ExpressionComparison): + if isinstance(atom, ExpressionComparison): return _expression_rewrite(atom) return None -def _expression_rewrite(node: ArithmeticComparison | ExpressionComparison) -> str: +def _expression_rewrite(node: ExpressionComparison[ArithmeticNode]) -> 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 @@ -237,16 +237,15 @@ def _expression_rewrite(node: ArithmeticComparison | ExpressionComparison) -> st other resolves to a :class:`~math_spec.program.ParameterComparison` and never reaches here, and a quoted label cannot stand on the left at all. """ - if isinstance(node, ArithmeticComparison): - left, right = node.left, node.right - number = isinstance(left, NumberNode) or ( - isinstance(left, UnaryOperatorNode) and isinstance(left.operand, NumberNode) + 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): + 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}' ) - if number and isinstance(right, ParameterNode): - 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}' - ) return ( 'it compares expressions, whose values only the data decides — compare one parameter against a ' 'literal, or precompute the test as a boolean parameter and test that' @@ -261,7 +260,7 @@ def _observe( ``position()`` converts the dimension to an integer, so an ordering over a rank is an ordering of integers and every comparator is admitted there. """ - if isinstance(node, ArithmeticComparison | ExpressionComparison): + if isinstance(node, ExpressionComparison): raise Undecidable(_expression_rewrite(node)) if isinstance(node, CountComparison): msg = ( @@ -318,7 +317,7 @@ def _subject_of(node: TypedPredicate) -> Subject: return Subject('relation', name) case RelationPairComparison(name=name, other=other): return Subject('relation_pair', name, other) - case ArithmeticComparison() | ExpressionComparison(): + case ExpressionComparison(): return Subject('expression', 'a comparison of expressions') case CountComparison(): return Subject('expression', 'a count of the coordinates a predicate admits') @@ -517,7 +516,7 @@ def _atom(node: TypedPredicate, cell: dict[Subject, Cell], grid: _Grid) -> bool: return bool(value) case RelationPairComparison(op=op): return bool(value) if op == '==' else not value - case ArithmeticComparison() | ExpressionComparison(): + case ExpressionComparison(): msg = 'a comparison of expressions is refused as undecidable before any cell is read' raise AssertionError(msg) case CountComparison() | TranslatedPredicate() | PulledBackPredicate(): diff --git a/src/math_spec/lowering.py b/src/math_spec/lowering.py index 12feaf83..396f8213 100644 --- a/src/math_spec/lowering.py +++ b/src/math_spec/lowering.py @@ -296,7 +296,7 @@ def _mask(self, mask: program.Mask) -> program.Mask: return program.Mask(self._predicate(mask.root)) def _predicate(self, node: program.Predicate) -> program.Predicate: - if isinstance(node, program.ArithmeticComparison): + 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)) diff --git a/src/math_spec/program.py b/src/math_spec/program.py index 0fc9af79..a1087608 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -34,15 +34,12 @@ import datetime from collections.abc import Iterator - from math_spec._expression_parser import ArithmeticNode - #: What ``math_spec.program`` promises a consumer, sorted. __all__ = [ 'QUADRATIC_POSITIONS', 'Add', 'And', - 'ArithmeticComparison', 'Assumption', 'BooleanLiteral', 'Cases', @@ -1056,34 +1053,22 @@ class ParameterComparison: @dataclass(frozen=True) -class ExpressionComparison: +class ExpressionComparison[Side]: """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. - """ - - left: Expression - op: PredicateOperator - right: Expression - dims: tuple[str, ...] - - -@dataclass(frozen=True) -class ArithmeticComparison: - """The same comparison as resolution types it, its sides in the core syntax tree. - What the spec-side readers walk — the typesetter, the dim rules, the - exclusivity check. :func:`~math_spec.lowering.lower_program` rebuilds - every mask with an :class:`ExpressionComparison` in its place, so a - program never carries one. + 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: ArithmeticNode + left: Side op: PredicateOperator - right: ArithmeticNode + right: Side dims: tuple[str, ...] @@ -1227,8 +1212,8 @@ 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 - | ArithmeticComparison | ParameterDefined | VariableDefined | DimensionComparison @@ -1247,10 +1232,7 @@ class Or: #: tree shares them — the transient impurity resolution normalizes away. Connective = Not | And | Or -#: Every resolved predicate node. A lowered mask's ``root`` holds every member -#: but :class:`ArithmeticComparison`, which lowering rewrites into an -#: :class:`ExpressionComparison`, so a consumer walking a program never meets -#: one. The parser's ``Unresolved*`` nodes are not members: they live with the +#: Every resolved predicate node. The parser's ``Unresolved*`` nodes are not members: they live with the #: grammar in :mod:`math_spec._where_parser`, and resolution rewrites them away #: before anything here is asked. Predicate = BooleanLiteral | TypedPredicate | Connective @@ -1304,7 +1286,6 @@ def _atom_dims(atom: TypedPredicate) -> frozenset[str]: case ( ParameterComparison() | ExpressionComparison() - | ArithmeticComparison() | ParameterDefined() | VariableDefined() | CountComparison() @@ -1341,9 +1322,6 @@ def _atom_names(atom: TypedPredicate) -> frozenset[str]: return frozenset({atom.name, atom.other}) case ExpressionComparison(): return _names_under(atom.left, atom.right) - case ArithmeticComparison(): - msg = 'a resolved mask is asked what it reads; lowering rebuilds it first, and the program mask answers.' - raise AssertionError(msg) case CountComparison(): return atom.predicate.names_read case TranslatedPredicate(): diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index cfddcb15..06bf9a94 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -67,12 +67,12 @@ ) from math_spec.program import ( And, - ArithmeticComparison, BooleanLiteral, CountComparison, DimensionComparison, DimensionPosition, Direction, + ExpressionComparison, Mask, Not, Or, @@ -997,7 +997,9 @@ def _plain(self, node: UnresolvedComparisonNode) -> _Plain | None: return None return _Plain(name, node.op, value, quoted) - def _expression_comparison(self, node: UnresolvedComparisonNode) -> ArithmeticComparison | UnresolvedComparisonNode: + def _expression_comparison( + self, node: UnresolvedComparisonNode + ) -> ExpressionComparison[ArithmeticNode] | 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 @@ -1053,7 +1055,7 @@ def _expression_comparison(self, node: UnresolvedComparisonNode) -> ArithmeticCo f'the comparison.' ) return node - return ArithmeticComparison(left, node.op, right, tuple(d for d in ns.schema.dimensions if d in dims)) + return ExpressionComparison(left, node.op, right, tuple(d for d in ns.schema.dimensions if d in dims)) def _position( self, call: FunctionCallNode, node: UnresolvedComparisonNode diff --git a/src/math_spec/typesetting/walk.py b/src/math_spec/typesetting/walk.py index d82e2266..b4037963 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -37,7 +37,6 @@ from math_spec.dimensions import dims_of from math_spec.program import ( And, - ArithmeticComparison, BooleanLiteral, CountComparison, DimensionComparison, @@ -82,7 +81,8 @@ #: align on the way it aligns a constraint. AlignedComparison = ( ParameterComparison - | ArithmeticComparison + # pyrefly: ignore[implicit-any-type-argument] # the union is an isinstance target, which takes no parameterized class + | ExpressionComparison | CountComparison | DimensionComparison | DimensionPosition @@ -602,10 +602,6 @@ def _where(self, node: Predicate, ctx: _Context) -> tuple[str, int]: left, right = self.sides(node, ctx) return f'{left} {right}', comparison - if isinstance(node, ExpressionComparison): - msg = 'a lowered comparison reached the typesetter; it prints the resolved tree, which lowering rebuilds.' - raise AssertionError(msg) - if isinstance(node, TranslatedPredicate): moved = ctx.translated(node.along, _Step(node.offset, 'plain')) return self._where(node.operand.root, moved) @@ -642,7 +638,7 @@ def sides(self, node: AlignedComparison, ctx: _Context) -> tuple[str, str]: """ if isinstance(node, ParameterComparison): left, right = ctx.indexed(self.symbols.name[node.name], list(node.dims)), self._literal(node.value) - elif isinstance(node, ArithmeticComparison): + elif isinstance(node, ExpressionComparison): left, right = self._expression(node.left, ctx), self._expression(node.right, ctx) elif isinstance(node, DimensionComparison): if isinstance(node.value, int | float): diff --git a/tests/typesetting/test_golden.py b/tests/typesetting/test_golden.py index 253bb457..c1a01a32 100644 --- a/tests/typesetting/test_golden.py +++ b/tests/typesetting/test_golden.py @@ -143,9 +143,8 @@ def _rendered_trees() -> Iterator[object]: #: What resolution never hands the walk: the two nodes a where carries before -#: its sides are read, the three an expression and a where carry before names -#: are resolved, and the lowered form of a comparison of expressions, which only a -#: program carries. The walk raises on each rather than rendering it, so a +#: 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 = { @@ -154,7 +153,6 @@ def _rendered_trees() -> Iterator[object]: 'NameNode', 'NameListNode', 'KeywordNode', - 'ExpressionComparison', } #: A dataclass the walk steps *through* rather than renders: an arm has no From 968f1f95bf73aea691a02cd96b7a49baad3eaf7f Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 22 Sep 2026 20:17:48 +0000 Subject: [PATCH 11/18] refactor(language): each named expression is resolved once, and every use reads that node Namespace.named resolves an expressions: entry the first time anything reads it. Expansion inlines that node, and the resolver passes it through. Expansion no longer parses named expressions, the resolver no longer re-reads a cased expression's arms, and validation drops its deduplication of repeated arm errors. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_017XwKgY5wZXv1bKgkfCW2q1 --- src/math_spec/expansion.py | 104 +++++++++---------------- src/math_spec/piecewise.py | 2 +- src/math_spec/resolution.py | 147 +++++++++++++++++++++++++++++++----- src/math_spec/validation.py | 81 +++----------------- tests/fixtures.py | 2 +- tests/test_expansion.py | 46 +++++++++-- tests/test_validation.py | 2 +- 7 files changed, 211 insertions(+), 173 deletions(-) diff --git a/src/math_spec/expansion.py b/src/math_spec/expansion.py index f3952b15..0af39a52 100644 --- a/src/math_spec/expansion.py +++ b/src/math_spec/expansion.py @@ -10,48 +10,47 @@ from math_spec._expression_parser import ( ArithmeticNode, - CaseArm, - CasesNode, ComparisonNode, - DefinitionNode, FunctionCallNode, NameNode, ParsedNode, parse_expression, with_children, ) -from math_spec._where_parser import parse_where from math_spec.errors import SchemaError if TYPE_CHECKING: - from math_spec.model import ExpressionBlock, MacroBlock, Spec + from math_spec.model import MacroBlock + from math_spec.resolution import Namespace -def parse_and_expand(text: str, schema: Spec, context: str) -> ParsedNode: +def parse_and_expand(text: str, ns: Namespace, context: str) -> ParsedNode: """Parse *text* and expand named sub-expressions and macros to core AST. Args: text: The expression as the file wrote it. - schema: Where names and macros are declared. + ns: Where names and macros are declared, and where a named expression is resolved. context: What an error names. """ - return expand(parse_expression(text), schema, context) + return expand(parse_expression(text), ns, context) @overload -def expand(node: ArithmeticNode, schema: Spec, context: str, *, shadow: frozenset[str] = ...) -> ArithmeticNode: ... +def expand(node: ArithmeticNode, ns: Namespace, context: str, *, shadow: frozenset[str] = ...) -> ArithmeticNode: ... @overload -def expand(node: ComparisonNode, schema: Spec, context: str, *, shadow: frozenset[str] = ...) -> ComparisonNode: ... +def expand(node: ComparisonNode, ns: Namespace, context: str, *, shadow: frozenset[str] = ...) -> ComparisonNode: ... -def expand(node: ParsedNode, schema: Spec, context: str, *, shadow: frozenset[str] = frozenset()) -> ParsedNode: +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 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. Args: node: The parsed expression. - schema: Where names and macros are declared. + ns: Where names and 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. @@ -59,10 +58,10 @@ def expand(node: ParsedNode, schema: Spec, context: str, *, shadow: frozenset[st if isinstance(node, ComparisonNode): return ComparisonNode( node.op, - _expand(node.left, schema, context, (), shadow), - _expand(node.right, schema, context, (), shadow), + _expand(node.left, ns, context, (), shadow), + _expand(node.right, ns, context, (), shadow), ) - return _expand(node, schema, context, (), shadow) + return _expand(node, ns, context, (), shadow) def macro_signature(name: str, macro: MacroBlock) -> str: @@ -73,73 +72,41 @@ def macro_signature(name: str, macro: MacroBlock) -> str: def parse_template(name: str, macro: MacroBlock, context: str) -> ArithmeticNode: """Parse a macro template, rejecting comparisons.""" - return _parse_body(macro.template, f"macro '{name}' template", context) + body = parse_expression(macro.template) + if isinstance(body, ComparisonNode): + msg = f"{context}: macro '{name}' template must not contain a comparison operator. Got: {macro.template!r}" + raise SchemaError(msg) + return body def _expand( node: ArithmeticNode, - schema: Spec, + ns: Namespace, context: str, stack: tuple[str, ...], shadow: frozenset[str], ) -> ArithmeticNode: - def _cycle(name: str, kind: str) -> None: - if name in stack: - chain = ' -> '.join([*stack, name]) - msg = f'{context}: circular {kind} reference: {chain}' - raise SchemaError(msg) - - if isinstance(node, NameNode) and node.name in schema.expressions and node.name not in shadow: - _cycle(node.name, 'expression') - body = _expand(_parse_named(node.name, schema, context), schema, context, (*stack, node.name), shadow) - return body if isinstance(body, CasesNode) else DefinitionNode(node.name, body) - - if isinstance(node, FunctionCallNode) and node.name in schema.macros: - _cycle(node.name, 'macro') - return _expand_macro(node, schema, context, stack, shadow) - - return with_children(node, lambda child: _expand(child, schema, context, stack, shadow)) - - -def _parse_named(name: str, schema: Spec, context: str) -> ArithmeticNode: - block = schema.expressions[name] - if block.cases: - return _parse_cased(name, block, context) - assert block.expression is not None - return _parse_body(block.expression, f"named expression '{name}'", context) - - -def _parse_cased(name: str, block: ExpressionBlock, context: str) -> CasesNode: - """A cased expression as the node that stands where its name was: the arms in file order, ``otherwise:`` last.""" - arms = [] - for label, case in block.cases.items(): - value = _parse_body(case.expression, f"named expression '{name}', case '{label}'", context) - # pyrefly: ignore[bad-argument-type] # the field is typed as resolution leaves it - arms.append(CaseArm(label, parse_where(case.when), value)) - assert block.otherwise is not None - fallback = _parse_body(block.otherwise, f"named expression '{name}', otherwise", context) - arms.append(CaseArm('otherwise', None, fallback)) - return CasesNode(name, tuple(arms)) + if isinstance(node, NameNode) and node.name in ns.schema.expressions and node.name not in shadow: + return ns.named(node.name, context) + 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) -def _parse_body(text: str, subject: str, context: str) -> ArithmeticNode: - """Parse one expression string that stands for a value, not a relation.""" - body = parse_expression(text) - if isinstance(body, ComparisonNode): - msg = f'{context}: {subject} must not contain a comparison operator. Got: {text!r}' - raise SchemaError(msg) - return body + return with_children(node, lambda child: _expand(child, ns, context, stack, shadow)) def _expand_macro( call: FunctionCallNode, - schema: Spec, + ns: Namespace, context: str, stack: tuple[str, ...], shadow: frozenset[str], ) -> ArithmeticNode: """Call-by-value: arguments are expanded before substitution, and the substituted body is expanded again.""" - macro = schema.macros[call.name] + macro = ns.schema.macros[call.name] signature = macro_signature(call.name, macro) if len(call.args) != len(macro.args): msg = ( @@ -156,15 +123,12 @@ def _expand_macro( raise SchemaError(msg) bindings = { - **{ - formal: _expand(arg, schema, context, stack, shadow) - for formal, arg in zip(macro.args, call.args, strict=True) - }, - **{formal: _expand(call.kwargs[formal], schema, context, stack, shadow) for formal in macro.kwargs}, + **{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}, } body = parse_template(call.name, macro, context) substituted = _substitute(body, bindings) - return _expand(substituted, schema, context, (*stack, call.name), shadow) + return _expand(substituted, ns, context, (*stack, call.name), shadow) def _substitute(node: ArithmeticNode, bindings: dict[str, ArithmeticNode]) -> ArithmeticNode: diff --git a/src/math_spec/piecewise.py b/src/math_spec/piecewise.py index 4ed53f4a..f5378075 100644 --- a/src/math_spec/piecewise.py +++ b/src/math_spec/piecewise.py @@ -503,7 +503,7 @@ def _nothing_collides(self) -> None: 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.schema, ctx) + 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] = [] diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index 06bf9a94..3c39878d 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -13,7 +13,7 @@ import datetime import re -from dataclasses import dataclass, replace +from dataclasses import dataclass from functools import cached_property from typing import TYPE_CHECKING, Literal, NamedTuple, assert_never, cast @@ -54,8 +54,9 @@ parse_where, ) from math_spec.dimensions import dims_of, pulled_back_dims -from math_spec.errors import DimensionError, LanguageError, did_you_mean, prefixed -from math_spec.expansion import expand +from math_spec.errors import DimensionError, LanguageError, SchemaError, 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 ( BUILTINS, @@ -94,7 +95,7 @@ if TYPE_CHECKING: from collections.abc import Iterable, Mapping - from math_spec.model import DeclaredDtype, Spec + from math_spec.model import DeclaredDtype, ExpressionBlock, Spec #: What a name a file may write turns out to be. Answered by @@ -109,7 +110,18 @@ class Namespace: A name has one kind: model.py refuses one declared under two sections. """ - __slots__ = ('constraints', 'dimensions', 'dtypes', 'leaf_dims', 'parameters', 'relations', 'schema', 'variables') + __slots__ = ( + '_loading', + '_named', + 'constraints', + 'dimensions', + 'dtypes', + 'leaf_dims', + 'parameters', + 'relations', + 'schema', + 'variables', + ) def __init__(self, schema: Spec) -> None: #: The schema the names come from — what an expression is expanded and @@ -140,6 +152,42 @@ def __init__(self, schema: Spec) -> None: **{p: tuple(pd.dims) for p, pd in schema.parameters.items()}, **{v: tuple(vd.dims) for v, vd in schema.variables.items()}, } + #: 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, ...]]] = {} + #: The named expressions being resolved, outermost first — a cycle's chain. + self._loading: list[str] = [] + + def named(self, name: str, context: str) -> CasesNode | DefinitionNode: + """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 + for, and read from then on, so a fault in it is reported once. + + Raises: + SchemaError: The entry reads itself, or does not load. + """ + if name in self._loading: + chain = ' -> '.join([*self._loading[self._loading.index(name) :], name]) + msg = f'{context}: circular expression reference: {chain}' + raise SchemaError(msg) + node, _ = self.named_entry(name) + if node is None: + msg = f"{context}: named expression '{name}' does not load. Its refusal is listed with it." + raise SchemaError(msg) + return node + + def named_entry(self, name: str) -> tuple[CasesNode | DefinitionNode | None, tuple[str, ...]]: + """The ``expressions:`` entry *name* resolved, or ``None``, with every refusal it earned.""" + if name not in self._named: + errors: list[str] = [] + self._loading.append(name) + try: + node = _named(name, self.schema.expressions[name], self, errors) + finally: + self._loading.pop() + self._named[name] = (node, tuple(errors)) + return self._named[name] def kind(self, name: str) -> DeclarationKind | None: """What *name* was declared as, or ``None`` where the file declares it nowhere.""" @@ -347,6 +395,71 @@ 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. + + 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. + """ + 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) + + found = len(errors) + arms: list[CaseArm] = [] + masks: dict[str, Predicate] = {} + for case_name, case in block.cases.items(): + arm_context = case_context(name, case_name) + when = resolve_where_text(case.when, ns, arm_context, errors) + if isinstance(when, BooleanLiteral): + 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) + if when is not None and value is not None: + arms.append(CaseArm(case_name, when, value)) + assert block.otherwise is not None + fallback = _value(block.otherwise, ns, case_context(name, None), errors) + 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}') + return None + resolved = resolve_expression(ast, ns, context, errors) + assert not isinstance(resolved, ComparisonNode), 'an arithmetic tree resolves to arithmetic' + return resolved + + +def _constant_arm(context: str, *, value: bool) -> str: + """The refusal for a case arm whose mask the connectives already decided. + + Cases are proved apart rather than ranked, so an always-true arm is not + one that shadows the arms under it — it is one no other arm can be proved + apart from, and the ``otherwise`` it leaves is empty. An always-false arm + is the plainer half: nothing to apply to. + """ + if value: + return ( + f'{context}: the mask admits every row, so no other arm can hold anywhere ' + f'and `otherwise:` covers nothing. Write the expression without `cases:`, ' + f'or narrow the `when`.' + ) + return f'{context}: the mask admits no row, so this arm never applies. Delete the arm, or widen the `when`.' + + @dataclass(frozen=True) class _Resolver: """One resolution walk, and the three things every step of it reads. @@ -383,13 +496,18 @@ def _arith(self, node: ArithmeticNode, *, amount: bool = False) -> ArithmeticNod *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. + arithmetic arrives through a macro formal bound to one. A named + expression arrives resolved, from :meth:`Namespace.named`, and passes. """ - if isinstance(node, NumberNode | VariableNode | ParameterNode | DualNode | KwargNode) or self._formal(node): + if isinstance( + node, NumberNode | VariableNode | ParameterNode | DualNode | KwargNode | CasesNode | DefinitionNode + ): + return node + if self._formal(node): return node if isinstance(node, NameNode): return self._name(node, amount=amount) - if isinstance(node, UnaryOperatorNode | BinaryOperatorNode | DefinitionNode): + if isinstance(node, UnaryOperatorNode | BinaryOperatorNode): return with_children(node, self._arith) if isinstance(node, FunctionCallNode): return self._call(node) @@ -407,8 +525,6 @@ def _arith(self, node: ArithmeticNode, *, amount: bool = False) -> ArithmeticNod f'terms out and add them.' ) return node - if isinstance(node, CasesNode): - return self._cases(node) assert_never(node) def _name(self, node: NameNode, *, amount: bool) -> ArithmeticNode: @@ -479,15 +595,6 @@ def _call(self, node: FunctionCallNode) -> ArithmeticNode: pass # a keyword the operator does not declare; the shape error already named it return FunctionCallNode(node.name, args, kwargs) - def _cases(self, node: CasesNode) -> CasesNode: - """Each arm's value and ``when`` typed under the arm's own context.""" - arms = [] - for arm in node.arms: - arm_context = case_context(node.name, None if arm.when is None else arm.label) - when = None if arm.when is None else resolve_where(arm.when, self.ns, arm_context, self.errors) - arms.append(CaseArm(arm.label, when, replace(self, context=arm_context)._arith(arm.value))) - return CasesNode(node.name, tuple(arms)) - def _amount(self, value: ArithmeticNode, operator: str, key: str) -> ArithmeticNode: """``offset=`` or ``window=``: a number or a parameter name, never an expression. @@ -1020,7 +1127,7 @@ def _expression_comparison( ) continue try: - expanded = expand(side, ns.schema, context) + expanded = expand(side, ns, context) except ValueError as e: self.errors.append(prefixed(context, e)) continue diff --git a/src/math_spec/validation.py b/src/math_spec/validation.py index f62df8d4..b876837d 100644 --- a/src/math_spec/validation.py +++ b/src/math_spec/validation.py @@ -12,17 +12,14 @@ import math_spec.degree as degree from math_spec._expression_parser import ( ArithmeticNode, - CaseArm, CasesNode, ComparisonNode, DefinitionNode, ParsedNode, - case_context, ) 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.exclusivity import overlapping 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 @@ -40,9 +37,6 @@ if TYPE_CHECKING: from pathlib import Path - from math_spec.model import ExpressionBlock - from math_spec.program import Predicate - def to_spec(model: str | Path | Mapping[str, object] | Spec) -> Spec: """Load and validate a model definition — the language's front door. @@ -72,17 +66,6 @@ def to_spec(model: str | Path | Mapping[str, object] | Spec) -> Spec: return Spec.model_validate(model if isinstance(model, Mapping) else read_model(model)) -def _once(errors: list[str]) -> str: - """The errors as one message, an identical sentence kept only the first time. - - A cased expression is expanded at every use, so a fault in one of its arms - is found again at each constraint naming it. Every error carries the - context it was found in, so two that differ at all are two faults and an - exact repeat is one seen twice. - """ - return '\n'.join(dict.fromkeys(errors)) - - def validate_expressions(schema: Spec) -> Resolved: """Validate and resolve every expression and where string in *schema*, once for every reader. @@ -117,7 +100,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), schema, context, shadow=formals) + body_ast = expand(parse_template(mname, macro, context), ns, context, shadow=formals) except ValueError as e: errors.append(prefixed(context, e)) continue @@ -130,9 +113,13 @@ def validate_expressions(schema: Spec) -> Resolved: resolve_expression(body_ast, ns, context, errors, formals=formals) expressions: dict[str, CasesNode | DefinitionNode] = {} - for ename, block in schema.expressions.items(): - if (node := _named(ename, block, ns, errors)) is not None: + for ename in schema.expressions: + node, refusals = ns.named_entry(ename) + errors.extend(refusals) + if node is not None: expressions[ename] = node + if errors: + raise SchemaError('\n'.join(errors)) variables = { vname: mask_of(resolve_where_text(vdef.where, ns, f"Variable '{vname}'", errors, self_variable=vname)) @@ -174,46 +161,13 @@ def validate_expressions(schema: Spec) -> Resolved: piecewise[pname] = tuple(link for link in links if link is not None) if errors: - raise SchemaError(_once(errors)) + raise SchemaError('\n'.join(errors)) resolved = Resolved(expressions, variables, constraints, objective, ns.relations, assumptions, piecewise) check_schema(schema, resolved) return resolved -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. - - 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. - """ - context = f"Named expression '{name}'" - if not block.cases: - assert block.expression is not None - body = _check_expression(block.expression, ns, context, errors, comparison=False, ceiling=None) - return None if body is None else DefinitionNode(name, body) - - found = len(errors) - arms: list[CaseArm] = [] - masks: dict[str, Predicate] = {} - for case_name, case in block.cases.items(): - arm_context = case_context(name, case_name) - when = resolve_where_text(case.when, ns, arm_context, errors) - if isinstance(when, BooleanLiteral): - errors.append(_constant_arm(arm_context, value=when.value)) - elif when is not None: - masks[case_name] = when - value = _check_expression(case.expression, ns, arm_context, errors, comparison=False, ceiling=None) - if when is not None and value is not None: - arms.append(CaseArm(case_name, when, value)) - assert block.otherwise is not None - fallback = _check_expression(block.otherwise, ns, case_context(name, None), errors, comparison=False, 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 _assumption(name: str, block: AssumptionBlock, ns: Namespace, errors: list[str]) -> ResolvedAssumption | None: """One ``assumptions:`` entry typed, or ``None`` once anything in it failed. @@ -270,23 +224,6 @@ def _decided_where(context: str, text: str, *, value: bool) -> str: ) -def _constant_arm(context: str, *, value: bool) -> str: - """The refusal for a case arm whose mask the connectives already decided. - - Cases are proved apart rather than ranked, so an always-true arm is not - one that shadows the arms under it — it is one no other arm can be proved - apart from, and the ``otherwise`` it leaves is empty. An always-false arm - is the plainer half: nothing to apply to. - """ - if value: - return ( - f'{context}: the mask admits every row, so no other arm can hold anywhere ' - f'and `otherwise:` covers nothing. Write the expression without `cases:`, ' - f'or narrow the `when`.' - ) - return f'{context}: the mask admits no row, so this arm never applies. Delete the arm, or widen the `when`.' - - @overload def _check_expression( expression: str, @@ -329,7 +266,7 @@ def _check_expression( declared. """ try: - ast = parse_and_expand(expression, ns.schema, context) + ast = parse_and_expand(expression, ns, context) except ValueError as e: errors.append(prefixed(context, e)) return None diff --git a/tests/fixtures.py b/tests/fixtures.py index 09794c11..05be8513 100644 --- a/tests/fixtures.py +++ b/tests/fixtures.py @@ -100,7 +100,7 @@ def raw_of(source: str | Path | dict[str, Any]) -> dict[str, Any]: def expression_of(text: str, ns: Namespace, context: str) -> ParsedNode: """Parse, expand and resolve one expression, raising every problem at once rather than collecting.""" errors: list[str] = [] - resolved = resolve_expression(parse_and_expand(text, ns.schema, context), ns, context, errors) + resolved = resolve_expression(parse_and_expand(text, ns, context), ns, context, errors) if errors: raise LanguageError('\n'.join(errors)) assert resolved is not None diff --git a/tests/test_expansion.py b/tests/test_expansion.py index f1eb067f..f3410df2 100644 --- a/tests/test_expansion.py +++ b/tests/test_expansion.py @@ -10,10 +10,11 @@ import pytest -from math_spec._expression_parser import ComparisonNode, DefinitionNode, parse_expression, with_children +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 tests.fixtures import DISPATCH_MODEL, SMALL_MODEL, schema_of +from math_spec.resolution import Namespace +from tests.fixtures import DISPATCH_MODEL, SMALL_MODEL, expression_of, schema_of WEIGHTED_SUM = { 'args': ['array', 'weights'], @@ -107,15 +108,17 @@ 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.""" - expanded = parse_and_expand(call, schema(expressions=expressions, macros=macros), 'expression') - assert _bodies(expanded) == parse_expression(want) + ns = Namespace(schema(expressions=expressions, macros=macros)) + assert _bodies(expression_of(call, ns, 'expression')) == expression_of(want, ns, 'expression') def test_a_named_expression_arrives_under_the_node_carrying_its_name(): - expanded = parse_and_expand('sum(gen_cost, over=generator)', schema(expressions={'gen_cost': 'p * cost'}), 'e') - assert expanded.args[0] == DefinitionNode('gen_cost', parse_expression('p * cost')), ( - 'the body is inlined and the name kept, for the typesetter to define it once' + 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')), ( + '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' @pytest.mark.parametrize( @@ -151,7 +154,7 @@ def test_a_refusal_names_its_context_once(): ) def test_macro_arity_errors(call, match): with pytest.raises(LanguageError, match=match): - parse_and_expand(call, schema(macros={'ws': WEIGHTED_SUM}), 'expression') + parse_and_expand(call, Namespace(schema(macros={'ws': WEIGHTED_SUM})), 'expression') @pytest.mark.parametrize( @@ -268,3 +271,30 @@ def test_a_formal_stands_where_a_call_site_will_bind_it(formals, template): assert ( schema_of(SMALL_MODEL, macros={'m': {'args': formals, 'template': template}}).macros['m'].template == template ) + + +def test_a_named_expression_is_resolved_once_however_many_uses(monkeypatch): + """Every use parsed, expanded and resolved the entry again, and a cased one's arms with it.""" + from math_spec import resolution + + resolved: list[str] = [] + named = resolution._named + + def counted(name, *args): + resolved.append(name) + return named(name, *args) + + monkeypatch.setattr(resolution, '_named', counted) + schema( + expressions={'gen_cost': 'p * cost', 'total': 'sum(gen_cost, over=generator) + sum(gen_cost, over=generator)'}, + constraints={'balance': {'dims': ['snapshot'], 'expression': 'sum(gen_cost, over=generator) >= load'}}, + ) + assert sorted(resolved) == ['gen_cost', 'total'], 'three uses of gen_cost, one resolution' + + +def test_a_use_of_a_refused_named_expression_names_it_rather_than_repeating_its_fault(): + with pytest.raises(LanguageError) as exc: + schema(expressions={'a': 'b + 1', 'b': 'nope'}) + message = str(exc.value) + assert message.count("'nope' not found") == 1, 'the fault is the entry that holds it' + assert "Named expression 'a': named expression 'b' does not load" in message diff --git a/tests/test_validation.py b/tests/test_validation.py index c1a76d71..8e51e6a7 100644 --- a/tests/test_validation.py +++ b/tests/test_validation.py @@ -1822,7 +1822,7 @@ def test_a_constraint_naming_it_carries_the_declared_frame(self): to_spec(model) def test_a_fault_in_an_arm_names_the_declaration_and_is_reported_once(self): - """The block is expanded at every use, and the fault is in one place. + """The block is resolved once, and the fault is in one place. Naming the use site would report a case on a constraint that has none, and one sentence per constraint reading the expression is the same From a70ba12fdbad4e1d9b9d2cc7a6dcbf3e8f7b8e61 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 23 Sep 2026 05:18:09 +0000 Subject: [PATCH 12/18] chore(resolution): one function reads an expression string into its typed tree `resolve_expression_text` in resolution.py is what validation.py's `_check_expression` was, and resolution's `_value` was the same function with `comparison=False, ceiling=None`. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01DwZZXWXoJfSMabvXTUndtn --- src/math_spec/resolution.py | 85 +++++++++++++++++++++++------- src/math_spec/validation.py | 101 +++++------------------------------- 2 files changed, 78 insertions(+), 108 deletions(-) diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index 3c39878d..8a059361 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -15,7 +15,7 @@ import re from dataclasses import dataclass from functools import cached_property -from typing import TYPE_CHECKING, Literal, NamedTuple, assert_never, cast +from typing import TYPE_CHECKING, Literal, NamedTuple, assert_never, cast, overload import math_spec.degree as degree from math_spec._expression_parser import ( @@ -395,6 +395,66 @@ def resolve_where_text( return resolve_where(node, ns, context, errors, self_variable) +@overload +def resolve_expression_text( + text: str, ns: Namespace, context: str, errors: list[str], *, comparison: Literal[True], ceiling: int | None +) -> ComparisonNode | None: ... +@overload +def resolve_expression_text( + text: str, ns: Namespace, context: str, errors: list[str], *, comparison: Literal[False], ceiling: int | None +) -> ArithmeticNode | None: ... + + +def resolve_expression_text( + text: str, ns: Namespace, context: str, errors: list[str], *, comparison: bool, ceiling: int | None +) -> ParsedNode | None: + """Parse, expand, resolve and degree-check one expression string, as :func:`resolve_where_text` reads a where string. + + *comparison* says what the position holds: a constraint carries exactly one + comparison, with a variable on a side (#1171), and every other position + carries none. *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 and piecewise link, and not where an entry is declared. + + Returns: + The typed tree, or ``None`` once anything failed, the problem appended + to *errors*. + """ + try: + ast = parse_and_expand(text, 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: {text!r}' + ) + return None + if not comparison and 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 + 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: {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 resolved + + 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. @@ -404,7 +464,7 @@ def _named(name: str, block: ExpressionBlock, ns: Namespace, errors: list[str]) context = f"Named expression '{name}'" if not block.cases: assert block.expression is not None - body = _value(block.expression, ns, context, errors) + body = resolve_expression_text(block.expression, ns, context, errors, comparison=False, ceiling=None) return None if body is None else DefinitionNode(name, body) found = len(errors) @@ -417,32 +477,19 @@ 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, comparison=False, ceiling=None) if when is not None and value is not None: arms.append(CaseArm(case_name, 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, comparison=False, 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}') - return None - resolved = resolve_expression(ast, ns, context, errors) - assert not isinstance(resolved, ComparisonNode), 'an arithmetic tree resolves to arithmetic' - return resolved - - def _constant_arm(context: str, *, value: bool) -> str: """The refusal for a case arm whose mask the connectives already decided. diff --git a/src/math_spec/validation.py b/src/math_spec/validation.py index b876837d..ecaa228b 100644 --- a/src/math_spec/validation.py +++ b/src/math_spec/validation.py @@ -7,20 +7,12 @@ 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.errors import SchemaError, prefixed +from math_spec.expansion import 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 @@ -31,12 +23,15 @@ ResolvedConstraint, mask_of, resolve_expression, + resolve_expression_text, resolve_where_text, ) if TYPE_CHECKING: from pathlib import Path + from math_spec._expression_parser import CasesNode, DefinitionNode + def to_spec(model: str | Path | Mapping[str, object] | Spec) -> Spec: """Load and validate a model definition — the language's front door. @@ -130,13 +125,13 @@ def validate_expressions(schema: Spec) -> Resolved: 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) + expression = resolve_expression_text(cdef.expression, ns, context, errors, comparison=True, ceiling=2) if expression is not None: constraints[cname] = ResolvedConstraint(expression, mask_of(where)) objective = None if schema.objective is not None: - objective = _check_expression( + objective = resolve_expression_text( schema.objective.expression, ns, 'The objective', errors, comparison=False, ceiling=2 ) @@ -154,7 +149,9 @@ def validate_expressions(schema: Spec) -> Resolved: 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, comparison=False, ceiling=1 + ) for i, link in enumerate(pdef.links) ] if all(link is not None for link in links): @@ -222,77 +219,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 From c509d8d5eaf459179d93660ccf12e4058967eb7c Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 23 Sep 2026 05:21:04 +0000 Subject: [PATCH 13/18] refactor(language): a piecewise block is checked once at load, and the expansion only writes rows `Spec` computes `resolved` before it expands, so the expansion reads each link's typed tree instead of parsing the text again. The names a block references and the names it emits are checked with the other `Spec` reference rules, and its frame in `curve_frame`, which the typesetter reads too. `PiecewiseExpansionError` is gone: a block is refused as a `SchemaError` or a `DimensionError` like every other declaration. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01DwZZXWXoJfSMabvXTUndtn --- docs/reference/language/errors.md | 13 +- src/math_spec/__init__.py | 2 - src/math_spec/errors.py | 4 - src/math_spec/model.py | 82 +++++++- src/math_spec/piecewise.py | 338 ++++++++++++------------------ src/math_spec/typesetting/walk.py | 17 +- src/math_spec/validation.py | 9 +- tests/test_piecewise.py | 18 +- tests/test_public_surface.py | 2 +- 9 files changed, 225 insertions(+), 260 deletions(-) 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/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/errors.py b/src/math_spec/errors.py index 2c7ad7ad..ce60957e 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) diff --git a/src/math_spec/model.py b/src/math_spec/model.py index bfc28102..dc6f608a 100644 --- a/src/math_spec/model.py +++ b/src/math_spec/model.py @@ -612,6 +612,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 +948,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 +1123,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/piecewise.py b/src/math_spec/piecewise.py index f5378075..10920b10 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 dataclasses import dataclass from typing import TYPE_CHECKING, Literal, NamedTuple -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 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._expression_parser import ArithmeticNode #: The suffix on the second gate row, where the gate variable does not exist. @@ -218,19 +213,113 @@ def _bends(block: str, pw: PiecewiseBlock, x: str, y: str, curvature: Curvature) return Assumed(bend.format('<=' if curvature == 'convex' else '>='), interior, 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. + + 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[ArithmeticNode]) -> tuple[str, ...]: + """The dimensions block *name* builds one curve per coordinate of: every one its links and its gate carry. - 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. + 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 +327,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,7 +346,7 @@ 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: @@ -276,14 +357,14 @@ def _assumptions(self) -> None: 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 +374,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 +431,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 +451,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 +463,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/typesetting/walk.py b/src/math_spec/typesetting/walk.py index b4037963..818b1aec 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -35,6 +35,7 @@ VariableNode, ) from math_spec.dimensions import dims_of +from math_spec.piecewise import curve_frame from math_spec.program import ( And, BooleanLiteral, @@ -920,7 +921,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 +990,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 ecaa228b..b9326db8 100644 --- a/src/math_spec/validation.py +++ b/src/math_spec/validation.py @@ -14,7 +14,7 @@ from math_spec.errors import SchemaError, prefixed from math_spec.expansion import expand, parse_template from math_spec.model import AssumptionBlock, Spec -from math_spec.piecewise import assumptions_of +from math_spec.piecewise import assumptions_of, curve_frame from math_spec.program import BooleanLiteral, Mask, VariableDefined from math_spec.resolution import ( Namespace, @@ -76,8 +76,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 @@ -162,6 +163,8 @@ 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 diff --git a/tests/test_piecewise.py b/tests/test_piecewise.py index 76b82367..76811785 100644 --- a/tests/test_piecewise.py +++ b/tests/test_piecewise.py @@ -14,7 +14,7 @@ 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 @@ -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})) @@ -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 From 4145de77935a145787a06f0e9d55a2340984fc81 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 23 Sep 2026 05:22:25 +0000 Subject: [PATCH 14/18] chore: a shape the package already has is not declared a second time `piecewise.assumptions_of` returns `AssumptionBlock`s, which is what it made from its `Assumed` tuples; `Resolved.assumptions` holds `program.Holds`, which `ResolvedAssumption` had the fields of; the `Assumption` alias of `Holds` and the `LinkSign` alias of `ComparisonOperator` are gone. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01DwZZXWXoJfSMabvXTUndtn --- src/math_spec/lowering.py | 10 ++-- src/math_spec/model.py | 5 +- src/math_spec/piecewise.py | 76 ++++++++++++------------------- src/math_spec/program.py | 15 ++---- src/math_spec/resolution.py | 19 ++------ src/math_spec/typesetting/walk.py | 3 +- src/math_spec/validation.py | 15 +++--- tests/typesetting/test_golden.py | 8 ++-- 8 files changed, 56 insertions(+), 95 deletions(-) diff --git a/src/math_spec/lowering.py b/src/math_spec/lowering.py index 396f8213..1c330aff 100644 --- a/src/math_spec/lowering.py +++ b/src/math_spec/lowering.py @@ -185,7 +185,7 @@ def lower_program(expanded: Spec) -> program.Program: ) -def _assumptions(expanded: Spec) -> dict[str, program.Assumption]: +def _assumptions(expanded: Spec) -> dict[str, program.Holds]: """Everything the data has to satisfy, in the order the model states it. One mapping rather than two, because a consumer binding data checks them @@ -193,12 +193,12 @@ def _assumptions(expanded: Spec) -> dict[str, program.Assumption]: already here: the expansion writes them into ``assumptions:``, and a load derives the same text for a block the file still declares. """ - assumptions: dict[str, program.Assumption] = {} - for name, (holds, where, description) in expanded.resolved.assumptions.items(): + assumptions: dict[str, program.Holds] = {} + for name, holds in expanded.resolved.assumptions.items(): lowering = _Lowering(expanded, f"assumption '{name}'") - predicate = lowering.mask(holds) + predicate = lowering.mask(holds.predicate) 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) + assumptions[name] = program.Holds(predicate, lowering.mask(holds.where), holds.description) return assumptions diff --git a/src/math_spec/model.py b/src/math_spec/model.py index dc6f608a..4b4ee38f 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] @@ -550,7 +547,7 @@ class PiecewiseLink(_StrictBlock): expression: str values: str - sign: LinkSign = '==' + sign: ComparisonOperator = '==' @model_validator(mode='before') @classmethod diff --git a/src/math_spec/piecewise.py b/src/math_spec/piecewise.py index 10920b10..7e92cc57 100644 --- a/src/math_spec/piecewise.py +++ b/src/math_spec/piecewise.py @@ -15,12 +15,12 @@ from __future__ import annotations from dataclasses import dataclass -from typing import TYPE_CHECKING, Literal, NamedTuple +from typing import TYPE_CHECKING, Literal import math_spec.sos as sos from math_spec.dimensions import dims_of from math_spec.errors import DimensionError -from math_spec.model import Curvature, PiecewiseBlock, PiecewiseMethod, Spec +from math_spec.model import AssumptionBlock, Curvature, PiecewiseBlock, PiecewiseMethod, Spec from math_spec.program import PiecewiseDeclaration if TYPE_CHECKING: @@ -64,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 @@ -89,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. ' + ( @@ -110,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 @@ -183,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, @@ -205,12 +191,13 @@ 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 + ) @dataclass(frozen=True) @@ -348,10 +335,7 @@ def _assumptions(self) -> None: """ 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 ---------------------------------------------------------- diff --git a/src/math_spec/program.py b/src/math_spec/program.py index a1087608..51a62aa7 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -40,7 +40,6 @@ 'QUADRATIC_POSITIONS', 'Add', 'And', - 'Assumption', 'BooleanLiteral', 'Cases', 'Connective', @@ -543,7 +542,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:`Holds` like any other, under :attr:`Program.assumptions`; what is left here is the curve. Attributes: @@ -576,15 +575,7 @@ 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: +def assumption_message(name: str, assumption: Holds) -> str: """The sentence a consumer raises when the data bound to *assumption*, called *name*, fails it. The language's own wording, so every consumer refuses in the same words; @@ -840,7 +831,7 @@ class Program: #: what each ``piecewise:`` block's method assumes of its breakpoints. The #: language decides none of it, so the consumer binding the data checks #: each and refuses with :func:`assumption_message`. - assumptions: Mapping[str, Assumption] = Sealed({}) + assumptions: Mapping[str, Holds] = Sealed({}) #: Declared ``expressions:``, lowered, each saying whether the math reads #: it. None builds a row of its own — one the math reads is inlined where #: it is read — but all are lowered with the program, so a file whose diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index 8a059361..6681a317 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -74,6 +74,7 @@ DimensionPosition, Direction, ExpressionComparison, + Holds, Mask, Not, Or, @@ -239,19 +240,6 @@ class ResolvedConstraint(NamedTuple): 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. @@ -277,7 +265,8 @@ class Resolved: copy, which every :class:`~math_spec.program.Direction` and :class:`~math_spec.program.Partition` in the trees holds. assumptions: Each ``assumptions:`` entry's predicate and the mask it - is checked under. + is checked under, as a program carries it, its comparisons of + expressions still in the core syntax tree. piecewise: Each ``piecewise:`` block's link expressions, in link order. """ @@ -286,7 +275,7 @@ class Resolved: constraints: dict[str, ResolvedConstraint] objective: ArithmeticNode | None relations: dict[str, RelationDeclaration] - assumptions: dict[str, ResolvedAssumption] + assumptions: dict[str, Holds] piecewise: dict[str, tuple[ArithmeticNode, ...]] @cached_property diff --git a/src/math_spec/typesetting/walk.py b/src/math_spec/typesetting/walk.py index 818b1aec..94c7c34f 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -901,7 +901,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): diff --git a/src/math_spec/validation.py b/src/math_spec/validation.py index b9326db8..50482254 100644 --- a/src/math_spec/validation.py +++ b/src/math_spec/validation.py @@ -13,13 +13,12 @@ from math_spec.dimensions import check_schema from math_spec.errors import SchemaError, prefixed from math_spec.expansion import expand, parse_template -from math_spec.model import AssumptionBlock, Spec +from math_spec.model import Spec from math_spec.piecewise import assumptions_of, curve_frame -from math_spec.program import BooleanLiteral, Mask, VariableDefined +from math_spec.program import BooleanLiteral, Holds, Mask, VariableDefined from math_spec.resolution import ( Namespace, Resolved, - ResolvedAssumption, ResolvedConstraint, mask_of, resolve_expression, @@ -31,6 +30,7 @@ from pathlib import Path from math_spec._expression_parser import CasesNode, DefinitionNode + from math_spec.model import AssumptionBlock def to_spec(model: str | Path | Mapping[str, object] | Spec) -> Spec: @@ -136,15 +136,14 @@ def validate_expressions(schema: Spec) -> Resolved: schema.objective.expression, ns, 'The objective', errors, comparison=False, ceiling=2 ) - assumptions: dict[str, ResolvedAssumption] = {} + assumptions: dict[str, Holds] = {} 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 = {} @@ -168,7 +167,7 @@ def validate_expressions(schema: Spec) -> Resolved: 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]) -> Holds | 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 +197,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 Holds(Mask(holds), mask_of(where), block.description) def _decided_assumption(context: str, text: str, *, value: bool) -> str: diff --git a/tests/typesetting/test_golden.py b/tests/typesetting/test_golden.py index c1a01a32..6d2f3c1b 100644 --- a/tests/typesetting/test_golden.py +++ b/tests/typesetting/test_golden.py @@ -133,10 +133,10 @@ def _rendered_trees() -> Iterator[object]: 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 From ec9089b0c2c1c6ae962ba69ef46c368d73ec9a76 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 23 Sep 2026 05:56:04 +0000 Subject: [PATCH 15/18] refactor(language): an expression resolves straight into the program's own nodes Resolution builds `program.Expression` nodes from the syntax tree, so the dim rules, the degree rules, the exclusivity check, the typesetter and lowering read one vocabulary. The parser keeps its eight syntax nodes; the ten typed ones, the node groups and the assertion arms that policed them are gone, and lowering packages declarations with each `Named` use of an `expressions:` entry inlined. The form of an `offset=`, `window=` and `edge=` is decided where it is read, in resolution. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01DwZZXWXoJfSMabvXTUndtn --- docs/contributing.md | 23 +- docs/reference/reading.md | 4 + src/math_spec/_expression_parser.py | 218 +-------- src/math_spec/degree.py | 133 +++--- src/math_spec/dimensions.py | 414 +++++------------ src/math_spec/errors.py | 11 + src/math_spec/exclusivity.py | 26 +- src/math_spec/expansion.py | 64 +-- src/math_spec/lowering.py | 352 +++----------- src/math_spec/model.py | 4 +- src/math_spec/operators.py | 41 +- src/math_spec/piecewise.py | 4 +- src/math_spec/program.py | 30 +- src/math_spec/resolution.py | 669 ++++++++++++++++++--------- src/math_spec/typesetting/symbols.py | 7 +- src/math_spec/typesetting/walk.py | 360 +++++++------- src/math_spec/validation.py | 28 +- tests/fixtures.py | 24 +- tests/test_degree.py | 19 +- tests/test_dimensions.py | 6 +- tests/test_expansion.py | 39 +- tests/test_lowering.py | 33 +- tests/test_operators.py | 42 +- tests/test_parser.py | 46 -- tests/test_validation.py | 7 +- tests/typesetting/test_golden.py | 64 +-- 26 files changed, 1120 insertions(+), 1548 deletions(-) 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/reading.md b/docs/reference/reading.md index d5237914..065bf622 100644 --- a/docs/reference/reading.md +++ b/docs/reference/reading.md @@ -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/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 ce60957e..97220d2b 100644 --- a/src/math_spec/errors.py +++ b/src/math_spec/errors.py @@ -93,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 1c330aff..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.Holds]: - """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.Holds] = {} - for name, holds in expanded.resolved.assumptions.items(): - lowering = _Lowering(expanded, f"assumption '{name}'") - predicate = lowering.mask(holds.predicate) - assert predicate is not None, 'a predicate that admits every row was refused as deciding nothing' - assumptions[name] = program.Holds(predicate, lowering.mask(holds.where), holds.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 4b4ee38f..c756590f 100644 --- a/src/math_spec/model.py +++ b/src/math_spec/model.py @@ -340,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' 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 7e92cc57..c00a0224 100644 --- a/src/math_spec/piecewise.py +++ b/src/math_spec/piecewise.py @@ -26,7 +26,7 @@ if TYPE_CHECKING: from collections.abc import Iterable - from math_spec._expression_parser import ArithmeticNode + from math_spec.program import Expression #: The suffix on the second gate row, where the gate variable does not exist. @@ -257,7 +257,7 @@ def by_kind(self) -> tuple[tuple[str, tuple[str, ...]], ...]: ) -def curve_frame(schema: Spec, name: str, pw: PiecewiseBlock, links: Iterable[ArithmeticNode]) -> tuple[str, ...]: +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 diff --git a/src/math_spec/program.py b/src/math_spec/program.py index 51a62aa7..cf0a5687 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -63,6 +63,7 @@ 'Holds', 'Mask', 'Multiply', + 'Named', 'Negate', 'Not', 'ObjectiveDeclaration', @@ -346,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 @@ -395,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)): @@ -1044,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, ...] @@ -1203,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 6681a317..7dd44593 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 @@ -15,35 +16,21 @@ import re from dataclasses import dataclass from functools import cached_property -from typing import TYPE_CHECKING, Literal, NamedTuple, assert_never, cast, overload +from typing import TYPE_CHECKING, Literal, NamedTuple, assert_never, cast import math_spec.degree as degree 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,35 +55,58 @@ unknown_operator_message, ) from math_spec.program import ( + Add, And, BooleanLiteral, + Cases, + Constant, + ConstraintDeclaration, CountComparison, DimensionComparison, DimensionPosition, Direction, + Divide, + Dual, + Expression, ExpressionComparison, + GroupSum, Holds, 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 @@ -104,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. @@ -155,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 @@ -178,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] = [] @@ -233,50 +248,42 @@ 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 - - @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. assumptions: Each ``assumptions:`` entry's predicate and the mask it - is checked under, as a program carries it, its comparisons of - expressions still in the core syntax tree. + is checked under. 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, Holds] - piecewise: dict[str, tuple[ArithmeticNode, ...]] + piecewise: dict[str, tuple[Expression, ...]] @cached_property def read_by_the_math(self) -> frozenset[str]: @@ -289,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)) # --------------------------------------------------------------------------- @@ -315,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 @@ -384,55 +404,59 @@ def resolve_where_text( return resolve_where(node, ns, context, errors, self_variable) -@overload -def resolve_expression_text( - text: str, ns: Namespace, context: str, errors: list[str], *, comparison: Literal[True], ceiling: int | None -) -> ComparisonNode | None: ... -@overload def resolve_expression_text( - text: str, ns: Namespace, context: str, errors: list[str], *, comparison: Literal[False], ceiling: int | None -) -> ArithmeticNode | None: ... + 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. - -def resolve_expression_text( - text: str, ns: Namespace, context: str, errors: list[str], *, comparison: bool, ceiling: int | None -) -> ParsedNode | None: - """Parse, expand, resolve and degree-check one expression string, as :func:`resolve_where_text` reads a where string. - - *comparison* says what the position holds: a constraint carries exactly one - comparison, with a variable on a side (#1171), and every other position - carries none. *ceiling* is the degree the position honours, and ``None`` - for an ``expressions:`` entry's body: what the math admits + *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 and piecewise link, and not where an entry is declared. + 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*. """ - try: - ast = parse_and_expand(text, 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: {text!r}' - ) + ast = _parsed(text, ns, context, errors) + if ast is None: return None - if not comparison and isinstance(ast, ComparisonNode): + 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 - try: - degree.check_expression(resolved, context, ceiling=ceiling) - except LanguageError as e: - errors.append(str(e)) + 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 isinstance(resolved, ComparisonNode) and not degree.carries_variable(resolved): + 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' @@ -441,23 +465,45 @@ def resolve_expression_text( f'bound, or state the fact under `assumptions:`, where the consumer binding the data checks it.' ) return None - return resolved + return left, ast.op, right -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 _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 = resolve_expression_text(block.expression, ns, context, errors, comparison=False, ceiling=None) - 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) @@ -466,17 +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 = resolve_expression_text(case.expression, ns, arm_context, errors, comparison=False, ceiling=None) + 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 = resolve_expression_text( - block.otherwise, ns, case_context(name, None), errors, comparison=False, ceiling=None - ) + 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))) + if len(errors) > found: + return None + 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: @@ -500,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 @@ -520,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, *, amount: bool = False) -> ArithmeticNode: - """One arithmetic node typed. + def arith(self, node: ArithmeticNode) -> Expression | None: + """The program node *node* stands for, or ``None``. - *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): @@ -553,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 " @@ -582,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 " @@ -590,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) -> 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. - def _call(self, node: FunctionCallNode) -> ArithmeticNode: - """An operator call: its shape checked, and each kwarg typed 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: @@ -615,97 +691,213 @@ 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) + + 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))) - Closed so that :func:`math_spec.dimensions._check_named_amount` sees every - parameter an amount carries. + 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)) - 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, @@ -713,49 +905,47 @@ 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 or any(n in self.formals for v in roles.values() for n in names_in(v)): - return value + return None if (problem := self._not_a_relation(name, operator, key)) is not None: self.errors.append(problem) - 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) 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.""" @@ -1046,14 +1236,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. @@ -1140,9 +1330,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 @@ -1151,7 +1339,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)) @@ -1167,12 +1355,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.' @@ -1519,18 +1709,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): @@ -1602,3 +1787,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 94c7c34f..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,55 +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 @@ -82,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 @@ -117,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. @@ -222,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 @@ -272,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] @@ -347,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, Cases): + return self.format.cases(self._arms(node, ctx)), _ATOM - if isinstance(node, UnresolvedNode | KwargNode): - msg = f'{type(node).__name__} reached the typesetter; resolve the expression first.' - raise AssertionError(msg) + 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)) - ``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. + 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. + + ``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.""" @@ -518,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: @@ -540,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 ----------------------------------------------------- @@ -732,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. @@ -750,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), ) @@ -785,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, @@ -819,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. diff --git a/src/math_spec/validation.py b/src/math_spec/validation.py index 50482254..d360e820 100644 --- a/src/math_spec/validation.py +++ b/src/math_spec/validation.py @@ -15,12 +15,12 @@ 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 BooleanLiteral, Holds, Mask, VariableDefined +from math_spec.program import BooleanLiteral, ConstraintDeclaration, Holds, Mask, ObjectiveDeclaration, VariableDefined from math_spec.resolution import ( Namespace, Resolved, - ResolvedConstraint, mask_of, + resolve_constraint_text, resolve_expression, resolve_expression_text, resolve_where_text, @@ -29,8 +29,8 @@ if TYPE_CHECKING: from pathlib import Path - from math_spec._expression_parser import CasesNode, DefinitionNode from math_spec.model import AssumptionBlock + from math_spec.program import Named def to_spec(model: str | Path | Mapping[str, object] | Spec) -> Spec: @@ -96,7 +96,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 @@ -108,7 +108,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) @@ -122,19 +122,19 @@ 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 = resolve_expression_text(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 = resolve_expression_text( - 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, Holds] = {} for aname, adef in schema.assumptions.items(): @@ -149,9 +149,7 @@ def validate_expressions(schema: Spec) -> Resolved: piecewise = {} for pname, pdef in schema.piecewise.items(): links = [ - resolve_expression_text( - 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): 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 f3410df2..11108d86 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..e980ba6f 100644 --- a/tests/test_lowering.py +++ b/tests/test_lowering.py @@ -11,15 +11,14 @@ 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, @@ -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( @@ -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_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 6d2f3c1b..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,11 +124,13 @@ 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 @@ -142,26 +143,13 @@ def _rendered_trees() -> Iterator[object]: 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:', From fb58db984ba5f47df02adb951ab6a26b7db93177 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 23 Sep 2026 05:56:44 +0000 Subject: [PATCH 16/18] chore(schema): the published schema carries the macro block's current sentence Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01DwZZXWXoJfSMabvXTUndtn --- schema/math-spec.schema.json | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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": [], From 3270e4809bd1067748832dd14b6c70f94d794ebc Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 23 Sep 2026 07:00:35 +0000 Subject: [PATCH 17/18] refactor(program): an assumption is the class named Assumption `program.Holds` is renamed `Assumption`: the file's section is `assumptions:`, every other declaration class is a noun, and the alias that carried the noun stood in for kinds that never came. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01DwZZXWXoJfSMabvXTUndtn --- docs/reference/reading.md | 6 +++--- src/math_spec/program.py | 10 +++++----- src/math_spec/resolution.py | 4 ++-- src/math_spec/validation.py | 15 +++++++++++---- tests/test_lowering.py | 6 +++--- tests/test_piecewise.py | 4 ++-- 6 files changed, 26 insertions(+), 19 deletions(-) diff --git a/docs/reference/reading.md b/docs/reference/reading.md index 065bf622..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 diff --git a/src/math_spec/program.py b/src/math_spec/program.py index cf0a5687..4521e973 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -40,6 +40,7 @@ 'QUADRATIC_POSITIONS', 'Add', 'And', + 'Assumption', 'BooleanLiteral', 'Cases', 'Connective', @@ -60,7 +61,6 @@ 'FanIn', 'Footprint', 'GroupSum', - 'Holds', 'Mask', 'Multiply', 'Named', @@ -561,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 - :class:`Holds` 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: @@ -576,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 @@ -594,7 +594,7 @@ class Holds: description: str | None = None -def assumption_message(name: str, assumption: Holds) -> str: +def assumption_message(name: str, assumption: Assumption) -> str: """The sentence a consumer raises when the data bound to *assumption*, called *name*, fails it. The language's own wording, so every consumer refuses in the same words; @@ -850,7 +850,7 @@ class Program: #: what each ``piecewise:`` block's method assumes of its breakpoints. The #: language decides none of it, so the consumer binding the data checks #: each and refuses with :func:`assumption_message`. - assumptions: Mapping[str, Holds] = Sealed({}) + assumptions: Mapping[str, Assumption] = Sealed({}) #: Declared ``expressions:``, lowered, each saying whether the math reads #: it. None builds a row of its own — one the math reads is inlined where #: it is read — but all are lowered with the program, so a file whose diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index 7dd44593..e83881d9 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -57,6 +57,7 @@ from math_spec.program import ( Add, And, + Assumption, BooleanLiteral, Cases, Constant, @@ -70,7 +71,6 @@ Expression, ExpressionComparison, GroupSum, - Holds, Mask, Multiply, Named, @@ -282,7 +282,7 @@ class Resolved: constraints: dict[str, ConstraintDeclaration] objective: ObjectiveDeclaration | None relations: dict[str, RelationDeclaration] - assumptions: dict[str, Holds] + assumptions: dict[str, Assumption] piecewise: dict[str, tuple[Expression, ...]] @cached_property diff --git a/src/math_spec/validation.py b/src/math_spec/validation.py index d360e820..1e47c205 100644 --- a/src/math_spec/validation.py +++ b/src/math_spec/validation.py @@ -15,7 +15,14 @@ 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 BooleanLiteral, ConstraintDeclaration, Holds, Mask, ObjectiveDeclaration, VariableDefined +from math_spec.program import ( + Assumption, + BooleanLiteral, + ConstraintDeclaration, + Mask, + ObjectiveDeclaration, + VariableDefined, +) from math_spec.resolution import ( Namespace, Resolved, @@ -136,7 +143,7 @@ def validate_expressions(schema: Spec) -> Resolved: if expression is not None: objective = ObjectiveDeclaration(schema.objective.sense, expression) - assumptions: dict[str, Holds] = {} + assumptions: dict[str, Assumption] = {} for aname, adef in schema.assumptions.items(): if (assumption := _assumption(aname, adef, ns, errors)) is not None: assumptions[aname] = assumption @@ -165,7 +172,7 @@ def validate_expressions(schema: Spec) -> Resolved: return resolved -def _assumption(name: str, block: AssumptionBlock, ns: Namespace, errors: list[str]) -> Holds | 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 @@ -195,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 Holds(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: diff --git a/tests/test_lowering.py b/tests/test_lowering.py index e980ba6f..ff7586c9 100644 --- a/tests/test_lowering.py +++ b/tests/test_lowering.py @@ -24,6 +24,7 @@ QUADRATIC_POSITIONS, Add, And, + Assumption, BooleanLiteral, Cases, Constant, @@ -37,7 +38,6 @@ ExpressionComparison, Footprint, GroupSum, - Holds, Mask, Multiply, Negate, @@ -498,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 == [ @@ -514,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' diff --git a/tests/test_piecewise.py b/tests/test_piecewise.py index 76811785..46601e11 100644 --- a/tests/test_piecewise.py +++ b/tests/test_piecewise.py @@ -17,7 +17,7 @@ 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 @@ -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'}), ( From 0455e935d8d2dc437845acb5a2c230d81f559332 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 23 Sep 2026 08:45:33 +0000 Subject: [PATCH 18/18] feat(program): a sum through a relation is a sum over the axes its join opens A Join keeps each column it joins on and does not group by as an axis of its own, named relation.column, so a column dropped and a column added over one dimension stay two axes. sum(by=) lowers to a Sum over those axes, and at() to the bare Join, which opens none because its groups are one row. Each node's frame is its own again, and a reduction other than a sum can follow the same join. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_013M6RG2cnsrFCEbaUuU8KyJ --- docs/about/relations-as-linear-maps.md | 12 ++++--- docs/contributing.md | 19 +++++------ src/math_spec/degree.py | 3 +- src/math_spec/dimensions.py | 10 +++--- src/math_spec/program.py | 47 ++++++++++++++++++-------- src/math_spec/resolution.py | 2 +- src/math_spec/separability.py | 27 ++++++++------- src/math_spec/typesetting/walk.py | 23 ++++++------- tests/test_dimensions.py | 19 ++++++++++- tests/test_lowering.py | 29 ++++++++-------- tests/test_separability.py | 4 +++ tests/typesetting/test_golden.py | 2 +- 12 files changed, 119 insertions(+), 78 deletions(-) diff --git a/docs/about/relations-as-linear-maps.md b/docs/about/relations-as-linear-maps.md index 327747b5..e94118d6 100644 --- a/docs/about/relations-as-linear-maps.md +++ b/docs/about/relations-as-linear-maps.md @@ -131,11 +131,13 @@ $`x`$ over generators and $`y`$ over buses, \langle M x, y \rangle = \sum_{b} y_b \sum_{g} \mathbf{1}_R(g, b)\, x_g = \sum_{g} x_g \sum_{b} \mathbf{1}_R(g, b)\, y_b = \langle x, M^{\mathsf{T}} y \rangle, ``` -which is why the program lowers `at` and `sum(by=)` to one `Join` node. The -node is the contraction, and whether each group is one row tells which call it -is. A bare relation has the same matrix without the -functional claim. A column of $`M`$ may hold several ones, so the sum fans out -and no group is one row. That is why `at` through a bare relation is refused. +which is why the program lowers `at` to a `Join` node and `sum(by=)` to the +same `Join` under a `Sum`. The join keeps the column it sums away as an axis +named for the relation's column, `gen_zone.generator`, and the `Sum` stands +over that axis. So a map into its own dimension, which drops and adds one +dimension, still has two axes between the join and the sum. A bare relation has +the same matrix without the functional claim. A column of $`M`$ may hold several +ones, so the sum fans out and no group is one row. That is why `at` through a bare relation is refused. ## The join and the aggregate diff --git a/docs/contributing.md b/docs/contributing.md index 786a5e36..56f33171 100644 --- a/docs/contributing.md +++ b/docs/contributing.md @@ -105,16 +105,15 @@ suffix says which layer: | Program (`math_spec.program`) | none / `Declaration` | `Variable`, `VariableDeclaration` | A node names the operation, not the verb a file writes. One verb can resolve -to two nodes and two verbs to one, so the file's spelling cannot decide the -name. - -| File verb | Node | What the node names | -| ------------------ | ----------- | ---------------------------------------------------- | -| `sum(over=)` | `Sum` | dims removed from the result | -| `sum(by=)` | `Join` | a join, and the sum over each group it groups by | -| `at(by=)` | `Join` | the same join, where each group is one row: a lookup | -| `shift(along=)` | `Translate` | a re-index along one dimension | -| `sum_back(along=)` | `WindowSum` | a sum over a trailing window | +to two nodes, so the file's spelling cannot decide the name. + +| File verb | Node | What the node names | +| ------------------ | ------------------- | ---------------------------------------------------- | +| `sum(over=)` | `Sum` | dims removed from the result | +| `sum(by=)` | `Sum` over a `Join` | a join, and the sum over the axes it opens | +| `at(by=)` | `Join` | a join whose groups are one row, with no sum over it | +| `shift(along=)` | `Translate` | a re-index along one dimension | +| `sum_back(along=)` | `WindowSum` | a sum over a trailing window | Nothing is abbreviated. diff --git a/src/math_spec/degree.py b/src/math_spec/degree.py index f2863258..ac840f8c 100644 --- a/src/math_spec/degree.py +++ b/src/math_spec/degree.py @@ -29,7 +29,6 @@ Divide, Dual, Expression, - Join, Multiply, Power, Sum, @@ -178,7 +177,7 @@ def _joins_terms(node: Expression) -> bool: A lookup and a translation re-index and are not reductions: they move a term, leaving one term where there was one. """ - if isinstance(node, Sum | WindowSum) or (isinstance(node, Join) and not node.columns.one_row_per_group): + if isinstance(node, Sum | WindowSum): return carries_variable(node.operand) return isinstance(node, Add) and carries_variable(node.left) and carries_variable(node.right) diff --git a/src/math_spec/dimensions.py b/src/math_spec/dimensions.py index f2fe10e6..6521fffa 100644 --- a/src/math_spec/dimensions.py +++ b/src/math_spec/dimensions.py @@ -128,10 +128,12 @@ def _sum_dims(node: Sum, inner: frozenset[str], context: str) -> frozenset[str]: def join_dims(columns: JoinColumns, inner: frozenset[str], context: str, operand: str) -> frozenset[str]: - """The dims *inner* has once *columns* joins it and sums each group, an expression's or a predicate's alike. + """The dims *inner* has once *columns* joins it, an expression's or a predicate's alike. - The dims joined on go and the dims grouped by arrive, so a column both - joined on and grouped by keeps its dim. The call a refusal quotes is + The dims joined on go, the dims grouped by arrive, and each column joined + on and not grouped by opens its own axis (:attr:`JoinColumns.axes`), which + the :class:`~math_spec.program.Sum` over the join takes away. A column + both joined on and grouped by keeps its dim. The call a refusal quotes is ``at`` where each group is one row, and ``sum`` otherwise. Raises: @@ -159,7 +161,7 @@ def join_dims(columns: JoinColumns, inner: frozenset[str], context: str, operand f'or group by a column over another dimension.' ) _check_joined(call, columns, inner, context) - return (inner - set(columns.joined_dims)) | set(columns.grouped_dims) + return (inner - set(columns.joined_dims)) | set(columns.grouped_dims) | set(columns.axes) #: The verb a file writes each translation with, which its refusals quote. diff --git a/src/math_spec/program.py b/src/math_spec/program.py index 9f9addcc..b78f3da0 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -234,7 +234,14 @@ class Divide: @dataclass(frozen=True) class Sum: - """Sum ``operand`` over the named dims, removing them from the result.""" + """Sum ``operand`` over the named dims, removing them from the result. + + A name in ``over`` is a dimension, or one of the axes a :class:`Join` + under it opens (:attr:`JoinColumns.axes`): ``sum(x, by=relation, over=a, + into=b)`` lowers to a ``Sum`` over a ``Join``, over the axis of each column + the join does not group by. The join is the join, and this node is the + group-by that follows it. + """ operand: Expression over: tuple[str, ...] @@ -242,19 +249,19 @@ class Sum: @dataclass(frozen=True) class Join: - """Join ``operand`` to a relation on the columns ``columns`` joins on, and sum each group of the columns it groups by. + """Join ``operand`` to a relation on the columns ``columns`` joins on, one row per matching row of the relation. The operand carries every dim joined on. The result has the operand's - dims, less the dims joined on, plus the dims grouped by: the join and the - sum over each group are one contraction, and this node is both. - ``sum(x, by=relation, over=a, into=b)`` and ``at(x, by=relation, over=a, - into=b)`` both lower to a ``Join``. Where the grouped columns hold the - relation's whole key (:attr:`JoinColumns.one_row_per_group`), each group - is one row and the node is ``at``'s lookup: one value per row of the - result, fanned out where several key tuples share the values joined on. - Elsewhere each group sums several rows, which is ``sum``'s. The loader - refuses a ``sum`` of the first shape and an ``at`` of the second, so - :func:`fan_in` reads which one a node is off its columns. + dims, less the dims joined on, plus the dims grouped by, plus one axis + per column joined on and not grouped by (:attr:`JoinColumns.axes`), named + for the relation's column and not for its dimension. So a column dropped + and a column added over one dimension stay two axes. A :class:`Sum` over + those axes is the group-by that follows the join, which is how + ``sum(x, by=relation, over=a, into=b)`` lowers. Where the grouped columns + hold the relation's whole key, they determine every other column, the + join opens no axis, and the bare ``Join`` is ``at(x, by=relation, + over=a, into=b)``: one value per row of the result, fanned out where + several key tuples share the values joined on. """ operand: Expression @@ -392,13 +399,11 @@ def fan_in(expression: Expression) -> FanIn: """ if isinstance(expression, Sum): return 'many-to-one' - if isinstance(expression, Join): - return 'one-to-one' if expression.columns.one_row_per_group else 'many-to-one' if isinstance(expression, WindowSum): return 'one-to-many' if isinstance( expression, - (Constant, Parameter, Variable, Dual, Negate, Add, Multiply, Power, Divide, Translate, Cases), + (Constant, Parameter, Variable, Dual, Negate, Add, Multiply, Power, Divide, Join, Translate, Cases), ): return 'one-to-one' assert_never(expression) @@ -538,6 +543,18 @@ def one_row_per_group(self) -> bool: """Whether the grouped columns hold the relation's whole key, so each group is one row: a lookup, not a sum.""" return set(self.relation.key) <= set(self.grouped) + @property + def axes(self) -> tuple[str, ...]: + """The axis the join opens for each column it drops, ``relation.column``, empty where each group is one row. + + A dimension's name holds no dot, so an axis never meets one. A column + the grouped columns determine opens none: in a lookup they hold the + key, and the key determines every column. + """ + if self.one_row_per_group: + return () + return tuple(f'{self.name}.{role}' for role in self.dropped) + @dataclass(frozen=True) class Partition: diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index 12f9577f..651ca2aa 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -732,7 +732,7 @@ def _built( if operator == 'sum': if read is not None: assert isinstance(read, JoinColumns), 'a sum joins its relation' - return Join(operand, read) + return Sum(Join(operand, read), read.axes) if (over := dims.get('over')) is not None: return Sum(operand, (over,)) return self._bare_sum(operand) diff --git a/src/math_spec/separability.py b/src/math_spec/separability.py index 8c227b1e..4534241e 100644 --- a/src/math_spec/separability.py +++ b/src/math_spec/separability.py @@ -97,25 +97,26 @@ def waits_on(dimension: str, label: str, name: str, kind: Literal['offset', 'par if isinstance(node, Cases): masks.extend(region.when for region in node.regions) elif isinstance(node, Sum): - if reductions_couple: - for dimension in node.over: + join = node.operand.columns if isinstance(node.operand, Join) and node.operand.columns.axes else None + grouped = dict(zip(join.axes, join.dropped_dims, strict=True)) if join is not None else {} + for name in node.over: + if join is not None and name in grouped: report( 'coupled', - dimension, + grouped[name], label, - f'sums over {dimension} — a rolling sum_back(window=n) windows, a total over the horizon does not', + f'groups {grouped[name]} into {", ".join(join.added_dims)} — window that dimension instead, or cut only at the group edges', ) - elif isinstance(node, Join) and node.columns.one_row_per_group: + elif reductions_couple: + report( + 'coupled', + name, + label, + f'sums over {name} — a rolling sum_back(window=n) windows, a total over the horizon does not', + ) + elif isinstance(node, Join) and not node.columns.axes: for dimension in node.columns.dropped_dims: waits_on(dimension, label, node.columns.name, 'coordinate') - elif isinstance(node, Join): - for dimension in node.columns.dropped_dims: - report( - 'coupled', - dimension, - label, - f'groups {dimension} into {", ".join(node.columns.added_dims)} — window that dimension instead, or cut only at the group edges', - ) elif isinstance(node, (Translate, WindowSum)): dimension = node.along if node.wrap: diff --git a/src/math_spec/typesetting/walk.py b/src/math_spec/typesetting/walk.py index 36c24164..6bfac647 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -424,7 +424,9 @@ def _binary(self, node: Add | Multiply | Divide | Power, ctx: _Context) -> tuple return self.format.joined([left, right], self._op(names[op])), precedence def _sum(self, node: Sum, ctx: _Context) -> tuple[str, int]: - """A reduction over named dims: one dummy index per dim, in declaration order.""" + """A reduction over named dims: one dummy index per dim, in declaration order, or the group-by over a join.""" + if isinstance(node.operand, Join) and node.operand.columns.axes: + return self._grouped_sum(node.operand, ctx) memberships = [] inner = ctx for d in self._sorted(frozenset(node.over)): @@ -433,16 +435,9 @@ def _sum(self, node: Sum, ctx: _Context) -> tuple[str, int]: domain = self.format.joined(memberships, '') return self.format.summation(domain, self._reduction_body(node.operand, inner)), _PRECEDENCE['+'] - def _join(self, node: Join, ctx: _Context) -> tuple[str, int]: - """A join through a relation: a lookup re-indexes the operand, and a grouped sum sums the rows each group joins. - - A lookup emits no operator of its own, so the read shows at the leaves. - A grouped sum takes a dummy per dim it drops, and the row it joins on - as the domain's condition. - """ - columns = node.columns - if columns.one_row_per_group: - return self._arithmetic(node.operand, self._looked_up(columns, ctx)) + def _grouped_sum(self, join: Join, ctx: _Context) -> tuple[str, int]: + """A sum over the axes a join opens: a dummy per dim it drops, and the row it joins on as the domain's condition.""" + columns = join.columns dummies: dict[str, str] = {} inner = ctx for d in columns.dropped_dims: @@ -452,7 +447,11 @@ def _join(self, node: Join, ctx: _Context) -> tuple[str, int]: f'{self.format.joined([self._membership(d, dummies[d]) for d in columns.dropped_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['+'] + return self.format.summation(domain, self._reduction_body(join.operand, inner)), _PRECEDENCE['+'] + + def _join(self, node: Join, ctx: _Context) -> tuple[str, int]: + """``at`` emits no operator of its own: it re-indexes the operand, so the lookup shows at the leaves.""" + return self._arithmetic(node.operand, self._looked_up(node.columns, 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. diff --git a/tests/test_dimensions.py b/tests/test_dimensions.py index bb53ea9b..c061b1e6 100644 --- a/tests/test_dimensions.py +++ b/tests/test_dimensions.py @@ -12,7 +12,7 @@ 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.program import Join, Mask, RelationPairComparison, Sum from math_spec.resolution import Namespace from math_spec.validation import to_spec from tests.fixtures import expression_of, override, schema_of, where_of @@ -228,6 +228,23 @@ def test_a_lookup_carries_the_whole_key_and_what_the_operand_brings_beside_it(): assert _dims('at(zone_load, by=gen_bz, over=zone, into=generator)') == {'generator', 'snapshot'} +def test_a_join_opens_an_axis_for_the_column_it_drops_and_the_sum_over_it_closes_it(): + """A map into its own dimension drops and adds one dimension, so the join names the dropped column for the relation. + + Named for its dimension, the column the join drops and the column it + groups by are one name, and the sum over the join takes away the dim the + row keeps. + """ + s = _schema() + node = expression_of('sum(p, by=rep_of, over=snapshot, into=rep)', Namespace(s), 't') + assert isinstance(node, Sum) and isinstance(node.operand, Join) + assert node.over == ('rep_of.snapshot',), 'the sum stands over the axis the join opens' + assert dims_of(node.operand, s, 't') == {'generator', 'snapshot', 'rep_of.snapshot'}, ( + 'the join keeps the dropped column beside the dimension it groups by' + ) + assert dims_of(node, s, 't') == {'generator', 'snapshot'}, 'and the sum over it leaves the frame the row keeps' + + def test_a_sum_joins_on_a_key_column_and_a_value_column_together(): """The columns a sum joins on and sums away are not one kind: it needs one key column, and may name a value column beside it.""" assert _dims('sum(p * load, by=gen_bz, over=[generator, bus], into=zone)') == {'snapshot', 'zone'} diff --git a/tests/test_lowering.py b/tests/test_lowering.py index e38aff5b..88c3e673 100644 --- a/tests/test_lowering.py +++ b/tests/test_lowering.py @@ -578,13 +578,13 @@ def test_a_power_resolves_to_a_node_of_its_own(dispatch_schema): pytest.param('sum(q, over=h)', Sum(Variable('q'), ('h',)), id='an-over-sums-away-the-dim-it-names'), pytest.param( 'sum(p, by=lk, over=g, into=h)', - Join(Variable('p'), LK_JOIN), - id='a-grouped-sum-is-a-join-grouped-by-the-into-column', + Sum(Join(Variable('p'), LK_JOIN), ('lk.g',)), + id='a-grouped-sum-is-a-sum-over-the-axis-its-join-opens', ), pytest.param( 'at(r, by=lk, over=h, into=g)', Join(Variable('r'), JoinColumns('lk', LK, ('h',), ('g',))), - id='an-at-is-the-same-join-the-other-way-grouped-by-the-key', + id='an-at-is-the-same-join-the-other-way-with-no-sum-over-it', ), pytest.param( "shift(p, along=g, offset=1, edge='wrap')", @@ -712,27 +712,28 @@ def test_a_relation_lowers_with_the_join_each_call_names(): assert program.relations == {'zone_of': declared}, 'the relation sits once in the program, under its name' zonal = program.constraints['zonal'].lhs columns = JoinColumns('zone_of', declared, ('generator', 'snapshot'), ('zone', 'snapshot')) - assert zonal == Join(Variable('p'), columns), ( - 'a grouped sum is a join: it names the over= column and the unnamed key column as joined on, ' - 'and the into= column and that key column as grouped by' + assert zonal == Sum(Join(Variable('p'), columns), ('zone_of.generator',)), ( + 'a grouped sum is a sum over a join: the join names the over= column and the unnamed key column as joined ' + 'on, the into= column and that key column as grouped by, and the sum stands over the axis the join opens ' + 'for the column it drops' ) - assert isinstance(zonal, Join) - assert (zonal.columns.dropped_dims, zonal.columns.added_dims, zonal.columns.kept) == ( + assert isinstance(zonal, Sum) and isinstance(zonal.operand, Join) + assert (zonal.operand.columns.dropped_dims, zonal.operand.columns.added_dims, zonal.operand.columns.kept) == ( ('generator',), ('zone',), ('snapshot',), ), 'the dims a consumer reads are read off the join: dropped, added, and the key columns kept' - assert not zonal.columns.one_row_per_group, 'grouping by zone and snapshot leaves several generators in a group' - assert zonal.columns.relation is program.relations['zone_of'], ( + assert zonal.operand.columns.relation is program.relations['zone_of'], ( 'the join holds the one declaration the program holds, not an equal copy built again' ) - assert program.constraints['history'].lhs == Join( - Variable('p'), JoinColumns('zone_of', declared, ('snapshot', 'generator'), ('zone', 'generator')) + assert program.constraints['history'].lhs == Sum( + Join(Variable('p'), JoinColumns('zone_of', declared, ('snapshot', 'generator'), ('zone', 'generator'))), + ('zone_of.snapshot',), ), 'the same table joined on its other key column' priced = program.constraints['priced'].rhs assert priced == Join( Parameter('price'), JoinColumns('zone_of', declared, ('zone', 'snapshot'), ('generator', 'snapshot')) - ), 'and an at is the same node, joined on the value column and grouped by the key columns' + ), 'and an at is the bare join, on the value column, grouped by the key columns' assert isinstance(priced, Join) assert (priced.columns.dropped_dims, priced.columns.added_dims, priced.columns.kept) == ( ('zone',), @@ -836,7 +837,7 @@ def test_walk_is_the_node_column_of_walk_regions(): Power(Parameter('c'), Constant(2.0)): 'one-to-one', Divide(Variable('p'), Parameter('c')): 'one-to-one', Sum(Variable('p'), ('g',)): 'many-to-one', - Join(Variable('p'), JoinColumns('at_bus', AT_BUS, ('g',), ('bus',))): 'many-to-one', + Sum(Join(Variable('p'), JoinColumns('at_bus', AT_BUS, ('g',), ('bus',))), ('at_bus.g',)): 'many-to-one', Join(Variable('p'), JoinColumns('at_bus', AT_BUS, ('bus',), ('g',))): 'one-to-one', Translate(Variable('p'), 't', offset=1, wrap=False, fill=0.0): 'one-to-one', WindowSum(Variable('p'), 't', width=2, wrap=False): 'one-to-many', diff --git a/tests/test_separability.py b/tests/test_separability.py index c275553c..f3a649e2 100644 --- a/tests/test_separability.py +++ b/tests/test_separability.py @@ -239,6 +239,10 @@ def test_a_grouping_that_sums_the_axis_away_couples_it(): ) verdict = program.separability['u'] assert not verdict.windowable, 'the grouping sums u away, so a window of u is a different sum' + assert verdict.coupled == { + "constraint 'z'": 'groups u into zone — window that dimension instead, or cut only at the group edges' + }, 'the sum over the join is the grouping, reported once' + assert not verdict.undecided, 'the join under the sum is not also a lookup waiting on the relation' def test_every_declared_axis_has_a_verdict_and_nothing_else_does(): diff --git a/tests/typesetting/test_golden.py b/tests/typesetting/test_golden.py index 4d8b7606..3e2a74c9 100644 --- a/tests/typesetting/test_golden.py +++ b/tests/typesetting/test_golden.py @@ -178,7 +178,7 @@ def test_the_golden_model_calls_every_operator_in_the_language(): verbs = {Sum: 'sum', Translate: 'shift', WindowSum: 'sum_back', Dual: 'dual'} nodes = [node for tree in _rendered_trees() for node in _nodes(tree)] calls = {verb for node in nodes for kind, verb in verbs.items() if isinstance(node, kind)} - calls |= {'at' if node.columns.one_row_per_group else 'sum' for node in nodes if isinstance(node, Join)} + calls |= {'at' for node in nodes if isinstance(node, Join) and not node.columns.axes} 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.'