diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index b7ce4c3d..25de5aee 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -24,7 +24,7 @@ jobs: fetch-depth: 0 persist-credentials: false - name: Set up pixi - uses: prefix-dev/setup-pixi@a09b6247153796b190642a2b53fac4241043cf6f # v0.10.0 + uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 with: locked: true # Off, deliberately: this job publishes the artifact a tag ships, and diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index f6a1830f..bb093c04 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -50,7 +50,7 @@ jobs: persist-credentials: false - name: Set up pixi - uses: prefix-dev/setup-pixi@a09b6247153796b190642a2b53fac4241043cf6f # v0.10.0 + uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 with: # `--locked`: pixi.lock is the environment, and a run that quietly # re-solved would not be testing what a contributor has. @@ -64,7 +64,7 @@ jobs: # is not a conda package. Keyed on the lockfile so a tectonic upgrade # gets a fresh bundle rather than a stale one. - name: Cache the tectonic bundle - uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0 + uses: actions/cache@55cc8345863c7cc4c66a329aec7e433d2d1c52a9 # v6.1.0 with: path: ~/.cache/Tectonic key: tectonic-${{ runner.os }}-${{ hashFiles('pixi.lock') }} diff --git a/.github/workflows/pypsa-references.yml b/.github/workflows/pypsa-references.yml index 83185330..5ea7d1d2 100644 --- a/.github/workflows/pypsa-references.yml +++ b/.github/workflows/pypsa-references.yml @@ -31,7 +31,7 @@ jobs: - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false - - uses: prefix-dev/setup-pixi@a09b6247153796b190642a2b53fac4241043cf6f # v0.10.0 + - uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 with: run-install: false - name: every rung solves to its record diff --git a/.github/workflows/scorecard.yml b/.github/workflows/scorecard.yml index da504b01..ccd249a5 100644 --- a/.github/workflows/scorecard.yml +++ b/.github/workflows/scorecard.yml @@ -78,6 +78,6 @@ jobs: # Upload the results to GitHub's code scanning dashboard (optional). # Commenting out will disable upload of results to your repo's Code Scanning dashboard - name: "Upload to code-scanning" - uses: github/codeql-action/upload-sarif@5595ccaf912efad79be6eef63a5619ff05969be3 # v4.37.6 + uses: github/codeql-action/upload-sarif@db488ddef3bf6cb639b32c2e9a7c0a7ea8271d28 # v4.37.8 with: sarif_file: results.sarif diff --git a/.github/workflows/update-lockfiles.yml b/.github/workflows/update-lockfiles.yml index 9989ebfd..34d4f0fb 100644 --- a/.github/workflows/update-lockfiles.yml +++ b/.github/workflows/update-lockfiles.yml @@ -18,7 +18,7 @@ jobs: steps: - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Set up pixi - uses: prefix-dev/setup-pixi@a09b6247153796b190642a2b53fac4241043cf6f # v0.10.0 + uses: prefix-dev/setup-pixi@f00437f565399d418b0acc85936d12c1fb668347 # v0.10.1 with: run-install: false - name: Update lockfiles diff --git a/.release-please-manifest.json b/.release-please-manifest.json index fcee1931..4cb21f18 100644 --- a/.release-please-manifest.json +++ b/.release-please-manifest.json @@ -1,3 +1,3 @@ { - ".": "0.0.0-alpha.53" + ".": "0.0.0-alpha.67" } diff --git a/CHANGELOG.md b/CHANGELOG.md index b2127a3f..726ca2b4 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -17,6 +17,110 @@ contained a literal `## [X.Y.Z]` heading, release-please inserts above the first `##` it finds, and so the entire release landed inside the comment and rendered nowhere. +## [0.0.0-alpha.67](https://github.com/energy-models/math-spec/compare/v0.0.0-alpha.66...v0.0.0-alpha.67) (2026-09-02) + + +### Refactoring + +* a built-in's one positional argument is stated once, and a degree question names the tree it walks ([#364](https://github.com/energy-models/math-spec/issues/364)) ([556af30](https://github.com/energy-models/math-spec/commit/556af3053e0bb32ca8bf60e656b9455b98068981)) +* every rule two passes shared has one home, and a docstring says what a caller needs rather than why ([#363](https://github.com/energy-models/math-spec/issues/363)) ([f24ddd3](https://github.com/energy-models/math-spec/commit/f24ddd3089da12aa75715560992822a5a3ee818b)) +* **model:** each cross-declaration rule is one method, so a refusal names the rule that raised it ([#369](https://github.com/energy-models/math-spec/issues/369)) ([b25a602](https://github.com/energy-models/math-spec/commit/b25a602d21a6bb91d597f663f8bcf0b97f95aed8)) +* **piecewise:** one block expands itself, holding its names, frame and mask once ([#367](https://github.com/energy-models/math-spec/issues/367)) ([537eadf](https://github.com/energy-models/math-spec/commit/537eadf76f52615c83235e0dfbbdd2af5fa70588)) +* resolution is one method per node kind, and each operator's dim rule is one function ([#368](https://github.com/energy-models/math-spec/issues/368)) ([c7c2833](https://github.com/energy-models/math-spec/commit/c7c283347869fbd45e0ff851d0112bfcbc09f98c)) +* **typesetting:** the legend reads what the equations returned rather than state left on the walk ([#366](https://github.com/energy-models/math-spec/issues/366)) ([db96a54](https://github.com/energy-models/math-spec/commit/db96a54d55996a9e4cc89a3f837303149475211b)) + +## [0.0.0-alpha.66](https://github.com/energy-models/math-spec/compare/v0.0.0-alpha.65...v0.0.0-alpha.66) (2026-09-01) + + +### Performance + +* a model loads, lowers and typesets three to eight times faster ([#357](https://github.com/energy-models/math-spec/issues/357)) ([ea6fe79](https://github.com/energy-models/math-spec/commit/ea6fe798c85118750294b642463b16aa0935065f)) + +## [0.0.0-alpha.65](https://github.com/energy-models/math-spec/compare/v0.0.0-alpha.64...v0.0.0-alpha.65) (2026-09-01) + + +### Refactoring + +* **language:** a translation policy and a bound's side name the values they can be, rather than being a string ([#354](https://github.com/energy-models/math-spec/issues/354)) ([d212157](https://github.com/energy-models/math-spec/commit/d2121572496bf7e624cfa5e38c348cfbed4b71ad)) +* **typesetting:** a format spells the operators the language names, rather than any string a walk happens to ask for ([#352](https://github.com/energy-models/math-spec/issues/352)) ([12b4041](https://github.com/energy-models/math-spec/commit/12b404126dbe00e28ad16b4736dbe90dfda833d0)) + +## [0.0.0-alpha.64](https://github.com/energy-models/math-spec/compare/v0.0.0-alpha.63...v0.0.0-alpha.64) (2026-09-01) + + +### Features + +* **parser:** a refused where string names the rewrite for pandas and C connective habits ([#346](https://github.com/energy-models/math-spec/issues/346)) ([3dbd9b2](https://github.com/energy-models/math-spec/commit/3dbd9b261dd82fc5cd52924ecdd03b18ddd88c14)) + +## [0.0.0-alpha.63](https://github.com/energy-models/math-spec/compare/v0.0.0-alpha.62...v0.0.0-alpha.63) (2026-09-01) + + +### Refactoring + +* **language:** an operator, a declaration kind and a notation name the values they can be, rather than being a string ([#345](https://github.com/energy-models/math-spec/issues/345)) ([3f3f858](https://github.com/energy-models/math-spec/commit/3f3f8581baabc51882ad4e1d97fb774ae1944dc3)) + +## [0.0.0-alpha.62](https://github.com/energy-models/math-spec/compare/v0.0.0-alpha.61...v0.0.0-alpha.62) (2026-09-01) + + +### Bug Fixes + +* **language:** a declaration named what no expression could write is refused, rather than loading unreferenceable ([#340](https://github.com/energy-models/math-spec/issues/340)) ([b865bc1](https://github.com/energy-models/math-spec/commit/b865bc15fde7e5a7714cdf809f5f0b9e6e6f44e5)) + +## [0.0.0-alpha.61](https://github.com/energy-models/math-spec/compare/v0.0.0-alpha.60...v0.0.0-alpha.61) (2026-09-01) + + +### Performance + +* **program:** a mask walks its leaves once and every question reads that walk ([#338](https://github.com/energy-models/math-spec/issues/338)) ([337d169](https://github.com/energy-models/math-spec/commit/337d1697724980356c8d911bf21242a19a6f517a)) + +## [0.0.0-alpha.60](https://github.com/energy-models/math-spec/compare/v0.0.0-alpha.59...v0.0.0-alpha.60) (2026-09-01) + + +### Bug Fixes + +* **language:** the language reference states which case arms are refused, and the refusal says what actually breaks ([#336](https://github.com/energy-models/math-spec/issues/336)) ([01920d1](https://github.com/energy-models/math-spec/commit/01920d1be4f6dbd6965f6f0e7e683543384cc744)) + +## [0.0.0-alpha.59](https://github.com/energy-models/math-spec/compare/v0.0.0-alpha.58...v0.0.0-alpha.59) (2026-09-01) + + +### Bug Fixes + +* **parser:** a parsed expression cannot be rewritten under another pass ([#329](https://github.com/energy-models/math-spec/issues/329)) ([fcfb7b8](https://github.com/energy-models/math-spec/commit/fcfb7b8a3e2c66cd03da316994154b9c2dd493d0)) + +## [0.0.0-alpha.58](https://github.com/energy-models/math-spec/compare/v0.0.0-alpha.57...v0.0.0-alpha.58) (2026-09-01) + + +### Features + +* **program:** a resolved where is a first-class Mask whose leaves carry their dims, and the where grammar is package-private ([#327](https://github.com/energy-models/math-spec/issues/327)) ([53cc352](https://github.com/energy-models/math-spec/commit/53cc3522e917a5849ce3150585c2c9e05a8ea162)) + +## [0.0.0-alpha.57](https://github.com/energy-models/math-spec/compare/v0.0.0-alpha.56...v0.0.0-alpha.57) (2026-09-01) + + +### Bug Fixes + +* **program:** a program describes a mathematical program rather than being one, and claims neither linearity nor a storage format ([#315](https://github.com/energy-models/math-spec/issues/315)) ([5dfa6be](https://github.com/energy-models/math-spec/commit/5dfa6be817fb49999571b2fdc8b8b820371f80d0)) + +## [0.0.0-alpha.56](https://github.com/energy-models/math-spec/compare/v0.0.0-alpha.55...v0.0.0-alpha.56) (2026-09-01) + + +### Features + +* **program:** the conjuncts of a where mask are the program's to give, not each consumer's to re-derive ([#313](https://github.com/energy-models/math-spec/issues/313)) ([db63d3c](https://github.com/energy-models/math-spec/commit/db63d3ca90079e4031a9339ea7749c0567b31be9)) + +## [0.0.0-alpha.55](https://github.com/energy-models/math-spec/compare/v0.0.0-alpha.54...v0.0.0-alpha.55) (2026-08-31) + + +### Documentation + +* **ceiling:** load-time unit checking has its own refusal, where the data-prep row used to answer for it ([#272](https://github.com/energy-models/math-spec/issues/272)) ([8dfd17d](https://github.com/energy-models/math-spec/commit/8dfd17df17993f1209fe479734c3dee756fcc384)) + +## [0.0.0-alpha.54](https://github.com/energy-models/math-spec/compare/v0.0.0-alpha.53...v0.0.0-alpha.54) (2026-08-31) + + +### Documentation + +* drop the effects row from the PyPSA-1.3.0 parity table, a feature that release does not have ([#305](https://github.com/energy-models/math-spec/issues/305)) ([92e5ed8](https://github.com/energy-models/math-spec/commit/92e5ed84c14c59e8068a7814722950f936e2925a)) + ## [0.0.0-alpha.53](https://github.com/energy-models/math-spec/compare/v0.0.0-alpha.52...v0.0.0-alpha.53) (2026-08-31) diff --git a/README.md b/README.md index 9f7cab45..0d037f6f 100644 --- a/README.md +++ b/README.md @@ -128,14 +128,15 @@ builds reads the second. That seam is [one page](docs/reference/language/reading.md), and it is the whole of it. -And the same file says, in print: +And that same `spec` says, in print — read and checked once, then printed +three ways: ```python symbols = 'dispatch.symbols.yaml' # optional: a dict, a path, or a SymbolTable -ms.to_latex('dispatch.yaml', symbols=symbols) # amsmath align -ms.to_typst('dispatch.yaml') # compiles without a TeX toolchain -ms.to_markdown('dispatch.yaml') # renders as-is on GitHub +ms.to_latex(spec, symbols=symbols) # amsmath align +ms.to_typst(spec) # compiles without a TeX toolchain +ms.to_markdown(spec) # renders as-is on GitHub ``` Drop the symbol table and the same model prints as $\mathit{load}_t$, diff --git a/docs/about/ceiling.md b/docs/about/ceiling.md index 7ab24ec2..4e00c7e9 100644 --- a/docs/about/ceiling.md +++ b/docs/about/ceiling.md @@ -235,16 +235,17 @@ for and refused, with the reason and the rewrite — so a request that has alrea been answered is answered once rather than re-argued. Parity with another tool is not by itself a reason to add anything. -| Request | Why | Instead | -| ------------------------------------------------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | ---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | -| Data prep — resampling, clustering, IO, units | not math | preprocess; pass a parameter | -| Arbitrary array ops (`merge`, `reindex`) | unbounded; xarray with extra steps | data prep | -| Domain helpers (`reduce_carrier_dim`) | encodes one domain into the language | component libraries over generic primitives | -| A tracked-metric vocabulary — `impacts:`, `effects:`, a `costs` dimension | the three fates are already reference-it-or-don't | an `impact` dim and one named expression: cap it with a constraint whose dual is the shadow price, weight it in the objective, read it with `result.expression` ([#124](https://github.com/fluxopt/lpspec/issues/124)) | -| `**` with a **variable** base or exponent | the exponent would decide the degree, and no data is read at load — `p ** n` is affine at 1, quadratic at 2 and over the ceiling at 3, and the file says which only once the numbers arrive | `x * x` for a square; above degree 2 there is no rewrite. Over variable-free operands `**` **is** in the language ([#1175](https://github.com/fluxopt/lpspec/issues/1175)) | -| Normalisation (`x / sum(x)`) | a _variable divisor_ is rational, not polynomial — no sink takes it at any degree | state the ratio as a constraint, or fix the denominator | -| Conditionals, iteration, data-dependent structure **inside one plan** | destroys the closed AST | `where` masks + `foreach` dims. A _process_ may loop over plans | -| A Python API for constructing models | hard rule 5 — the model is the file you review and diff | YAML. Whether Python may _emit_ declarations is [#381](https://github.com/fluxopt/lpspec/issues/381) | +| Request | Why | Instead | +| ------------------------------------------------------------------------- | ---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | ---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| Data prep — resampling, clustering, IO, unit conversion | not math | preprocess; pass a parameter | +| Unit _checking_ at load, with conversion staying refused | a `unit:` is a claim about data that nothing checks against the data, so a clean pass says "no two annotated operands disagreed" and never "the columns are in these units" — and an unannotated declaration propagating permissively makes even that advice. It would also open a unit grammar and a base-unit vocabulary the language then has to close ([#125](https://github.com/fluxopt/lpspec/issues/125)) | preprocess to one unit system, and name it in the declaration's `description:` | +| Arbitrary array ops (`merge`, `reindex`) | unbounded; xarray with extra steps | data prep | +| Domain helpers (`reduce_carrier_dim`) | encodes one domain into the language | component libraries over generic primitives | +| A tracked-metric vocabulary — `impacts:`, `effects:`, a `costs` dimension | the three fates are already reference-it-or-don't | an `impact` dim and one named expression: cap it with a constraint whose dual is the shadow price, weight it in the objective, read it with `result.expression` ([#124](https://github.com/fluxopt/lpspec/issues/124)) | +| `**` with a **variable** base or exponent | the exponent would decide the degree, and no data is read at load — `p ** n` is affine at 1, quadratic at 2 and over the ceiling at 3, and the file says which only once the numbers arrive | `x * x` for a square; above degree 2 there is no rewrite. Over variable-free operands `**` **is** in the language ([#1175](https://github.com/fluxopt/lpspec/issues/1175)) | +| Normalisation (`x / sum(x)`) | a _variable divisor_ is rational, not polynomial — no sink takes it at any degree | state the ratio as a constraint, or fix the denominator | +| Conditionals, iteration, data-dependent structure **inside one plan** | destroys the closed AST | `where` masks + `foreach` dims. A _process_ may loop over plans | +| A Python API for constructing models | hard rule 5 — the model is the file you review and diff | YAML. Whether Python may _emit_ declarations is [#381](https://github.com/fluxopt/lpspec/issues/381) | Genuinely unsayable math goes to a declared `escape:` island ([#38](https://github.com/fluxopt/lpspec/issues/38)) — named in the file, diff --git a/docs/examples/pypsa.md b/docs/examples/pypsa.md index e7cf5865..931b1bc3 100644 --- a/docs/examples/pypsa.md +++ b/docs/examples/pypsa.md @@ -425,7 +425,6 @@ each type is three blocks by sense. | [`tech_capacity_expansion_limit`](#tech_capacity_expansion_limit) | split | a block per sense | | `Bus-nom_min/max_{carrier}` | out | deprecated in PyPSA | | [`Carrier-growth_limit`](pypsa_multi_period.md) | done | rung 15, a file of its own | -| `effect_limit`, priced effects | open | `effects.py` not inventoried | > ✔ `pypsa 1.3.0` solves this rung's network at objective `10282.833333333332`, 102 rows. diff --git a/docs/index.md b/docs/index.md index d5213794..5a879a27 100644 --- a/docs/index.md +++ b/docs/index.md @@ -195,9 +195,11 @@ Only the notation is a choice, and **How** shows the one that was made here. }, } - ms.to_latex('dispatch.yaml', symbols=symbols) # amsmath align - ms.to_typst('dispatch.yaml') # compiles without a TeX toolchain - ms.to_markdown('dispatch.yaml') # renders as-is on GitHub + spec = ms.to_spec('dispatch.yaml') # read and checked once, then printed three ways + + ms.to_latex(spec, symbols=symbols) # amsmath align + ms.to_typst(spec) # compiles without a TeX toolchain + ms.to_markdown(spec) # renders as-is on GitHub ``` `symbols` is optional — drop it and the same model prints as diff --git a/docs/reference/language/expressions.md b/docs/reference/language/expressions.md index bf03c72b..ea39c6ed 100644 --- a/docs/reference/language/expressions.md +++ b/docs/reference/language/expressions.md @@ -82,6 +82,10 @@ keep the LP, its duals and its warm start. ## Name resolution +**A name is a letter or an underscore, then letters, digits or underscores** — +the spelling an expression uses to refer to one. A declaration keyed by +anything else is a load error, because nothing in the file could ever write it. + **One flat namespace** covers dimensions, parameters, variables, named expressions, macros and the built-in operators. A collision is a load error naming both declarations — there is no shadowing, because under it declaring a @@ -169,20 +173,20 @@ POSITION ::= "position" "(" NAME [ "," "by" "=" NAME ] ")" QUOTED ::= "'" chars "'" | '"' chars '"' ``` -| Surface | Names a… | Meaning | -| -------------------------------- | -------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | -| `name` (bare) | parameter | what defined means is the **declaration's** to say: a `bool` is its own answer, a `str` is defined wherever the table has a row, and a number has to be finite as well — `0.0` counts, `inf` does not, though it is a value everywhere else | -| `name` (bare) | variable | the variable exists at this coordinate — the counterpart of the parameter row, and how you say which coordinates the row-dropping rule applies to | -| `name` (bare) | dimension | load error: it is true everywhere, so it reads as a condition and is not one. Compare it instead | -| `name OP value` | parameter | element-wise; a null compares false. The right-hand side is a literal number, or a bare name read as a string coordinate | -| `name OP value` | dimension | a filter on the frame's own coordinate column | -| `name` (bare) | lookup | defined: the label maps somewhere. A lookup may be [partial](dimensions.md#lookups), and this is how a declaration asks for the labels that do map | -| `name OP value` | lookup | a filter on the lookup's column of its `over` dimension's index — which therefore has to be in the frame. A null value is **false**, whatever the comparator | -| `name OP name` | two lookups | the one comparison whose both sides are structure. Legal only where both map out of the **same** dimension _and_ into the **same** one — `from != to` excludes a self-loop | -| `position(name) OP i` | one dimension | where the row sits along that dimension's own order, as an integer — `0` is first, negative counts from the end. Both sides are integers, so every comparator reads the one way | -| `position(name, by=lookup) OP i` | a dimension and a lookup over it | the same, counted **within each group** the lookup makes — every period's first snapshot, whatever each period's length | -| `AND` `OR` `NOT` | — | case-insensitive; `NOT` binds tighter than `AND`, which binds tighter than `OR` | -| `True` / `False` | — | literals, decided at load wherever they stand: `True` is the same as no `where`, `False` is a declaration with no rows, and one under an `AND` or an `OR` settles that side — `x AND False` is the declaration with no rows too | +| Surface | Names a… | Meaning | +| -------------------------------- | -------------------------------- | --------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `name` (bare) | parameter | what defined means is the **declaration's** to say: a `bool` is its own answer, a `str` is defined wherever the table has a row, and a number has to be finite as well — `0.0` counts, `inf` does not, though it is a value everywhere else | +| `name` (bare) | variable | the variable exists at this coordinate — the counterpart of the parameter row, and how you say which coordinates the row-dropping rule applies to | +| `name` (bare) | dimension | load error: it is true everywhere, so it reads as a condition and is not one. Compare it instead | +| `name OP value` | parameter | element-wise; a null compares false. The right-hand side is a literal number, or a bare name read as a string coordinate | +| `name OP value` | dimension | a filter on the frame's own coordinate column | +| `name` (bare) | lookup | defined: the label maps somewhere. A lookup may be [partial](dimensions.md#lookups), and this is how a declaration asks for the labels that do map | +| `name OP value` | lookup | a filter on the lookup's column of its `over` dimension's index — which therefore has to be in the frame. A null value is **false**, whatever the comparator | +| `name OP name` | two lookups | the one comparison whose both sides are structure. Legal only where both map out of the **same** dimension _and_ into the **same** one — `from != to` excludes a self-loop | +| `position(name) OP i` | one dimension | where the row sits along that dimension's own order, as an integer — `0` is first, negative counts from the end. Both sides are integers, so every comparator reads the one way | +| `position(name, by=lookup) OP i` | a dimension and a lookup over it | the same, counted **within each group** the lookup makes — every period's first snapshot, whatever each period's length | +| `AND` `OR` `NOT` | — | case-insensitive; `NOT` binds tighter than `AND`, which binds tighter than `OR` | +| `True` / `False` | — | literals, decided at load wherever they stand: `True` is the same as no `where`, `False` is a declaration with no rows, and one under an `AND` or an `OR` settles that side — `x AND False` is the declaration with no rows too. A double negation goes the same way, `NOT NOT x` being `x`, so what a page prints is what the mask decides rather than how it was spelled. A case [`when:`](#the-rules) is the one place a mask that folds to a literal is refused instead | The mask's dims must not exceed the frame it sits in ([dim algebra](#dim-algebra)), and an undeclared bare name is a @@ -392,6 +396,21 @@ names the pair, a coordinate they both claim, and the rewrite: That is why `boundary` above says `committable and`. +**A `when:` the data cannot decide is not a case.** A mask the connectives +settle on their own — `True`, `False`, or anything that folds to one, like +`committable OR True` — states no condition for the data to answer, so it names +no region. Both halves are refused at load, and the refusal names the rewrite: + +> `Named expression 'previous_status', case 'always_on'`: the mask admits every +> row, so no other arm can hold anywhere and `otherwise:` covers nothing. Write +> the expression without `cases:`, or narrow the `when`. + +An always-false arm is the other half — it never applies, so delete it or widen +it. A **declaration's** `where:` is not held to this rule and cannot be: there +`False` is how a file says the declaration has no rows, and `True` is the same +as writing no mask at all. It is the `when:` on an arm that has to be a +question, because the arms are kept apart by proof. + **A pair the check cannot decide is refused too**, and that refusal names its rewrite as well. The one that comes up is `position(snapshot) == 0` against `position(snapshot) == -1`. On an axis with a single member those two pick the diff --git a/docs/reference/language/reading.md b/docs/reference/language/reading.md index 6bcca9d1..d6915464 100644 --- a/docs/reference/language/reading.md +++ b/docs/reference/language/reading.md @@ -101,9 +101,24 @@ means the caller binds it. **Nothing here is built by hand.** The program's nodes are exported to be dispatched on with `isinstance` and read, which is why what ships beside them -is the walk (`children()`) and not builders. A mask is the language's own -resolved `where` node rather than a second set spelling the same predicates — -one home, so the two cannot come to disagree about what a comparison is. +is the walk (`children()`) and not builders. A mask is `Mask`: the language's +own resolved `where` as its `.root` — the node an engine still dispatches on +with `isinstance` — and every question derived from it, the way a dimension +carries `.maps`. `.conjuncts` flattens the `AND` spine and stops at an `OR` or +a `NOT`; `.names_read` gives the declarations the mask names; `.atoms` its +leaves, connectives removed; and `.dims` the dimensions it is read at — read +off the leaves, which resolution stamped with their declarations' dims the way +a lookup leaf carries the dimension it maps out of. So a predicate a consumer +builds from resolved pieces answers exactly as a declaration's own does: wrap +it in `Mask`, or build it there with `~`, `&` and `|`. Construction +folds — a double negation cancels, a literal flips or is absorbed rather than +buried — so a boolean literal stands at a mask's root or nowhere, derived or +carried alike, and a tree with unresolved leaves is refused at the door. A +consumer asks the mask rather than re-deriving any of these from `.root`, so +two cannot come to disagree about what a conjunct, a name or a comparison is. +A `Region`'s `when` arrives in the same carrier, and the node classes a +`.root` is built of live in `math_spec.program` beside every other node a +consumer dispatches on. ## Asking what a program uses @@ -139,3 +154,68 @@ The footprint stops at the kind. A sink that takes a window but not a wrapped one reads `Window in footprint.shapes` and then walks: `wrap`, `partition` and a named width are refinements without end, and each is one line once the set has said where to look. + +## Asking whether an axis can be cut + +A driver that solves a horizon in windows — a rolling horizon, a myopic +pathway — needs one thing from the model before it starts: **is every row it +builds complete inside some window?** Storage carried over a snapshot is, once +the windows overlap by a row. An annual budget never is, and the windows still +solve, so nothing else would say so. + +```python +program.separability['bp'].windowable # False +tied = program.separability['generator'].coupled["constraint 'target'"] +tied.partition(' — ')[0] # 'sums over generator' +'sum_back(within=n)' in tied # True +``` + +Neither axis of the model above may be cut, and the report says which +declaration ties each one — including the three the `piecewise:` block emitted, +so a coupling introduced by an expansion is named under the name the expansion +gave it rather than under the block a reader wrote. + +It is the locality [the ceiling](../../about/ceiling.md) already argues in — +pointwise, bounded halo and global — asked about a dimension rather than about +an operator. Every declared axis has an entry, walked once and held like +[`footprint`](#asking-what-a-program-uses): answering for every axis costs what +answering for one did, since every construct that ties an axis names the axis it +ties. + +`behind` and `ahead` are how many coordinates a window must see before its +first row and after its last — `0` where every row is pointwise, `1` behind for +a `shift` of one, `n - 1` behind for a `sum_back` of `n`, `2` ahead for a shift +of `-2`. They are two numbers because a driver supplies them differently: a +lookahead is rows it solves and does not keep, a history is rows it carries +from the window before. + +What would break comes in three kinds, so a driver can act on each. `coupled` +names each declaration that ties the axis together — a sum over it in a +constraint, a grouping that consumes it, a wrapped shift, a set — and, after +the dash, the one modelling change that would lift it: a horizon total becomes +a rolling `sum_back`, a wrap becomes an opening-state seed, a grouping is +windowed along the dimension it groups into. No window satisfies a coupling and +no rewrite keeps the model's meaning, so the remedy is named and not applied. +`undecided` names each declaration whose reach only the data can say, to the +parameter or lookup that says it: an offset or width taken from a parameter, a +shift inside the groups a lookup makes, a read through a lookup with `at()`; a +driver holding the data computes the reach from it. `restarts` names each +declaration counting a `position()` along the axis, which a window restarts at +its first row. `windowable` is false while anything is coupled or undecided; a +restart does not count against it. + +The same walk answers a second driver. `independent` is whether each +coordinate builds on its own — windowable, reading nothing behind or ahead, +counting no position — which is what a scenario sweep asks before solving one +coordinate per slice, and what licenses solving the slices in any order or at +once. A restart does count against it: with one coordinate per slice a +`position()` holds everywhere, which changes what the mask means. + +**A reduction means opposite things by position**, which is the whole of the +care: in a constraint a sum over the axis ties every window to every other, and +in the objective it is additively separable, an objective being a sum already. +What is not decided here is whether the windowed answer is the whole-horizon +one — a store carried over one row windows cleanly and a rolling solve of it is +still a different answer — nor whether the modeller _wanted_ a restart: a +`position(t) == 0` seed fires once over a horizon and once per window, and both +are models somebody means. diff --git a/docs/reference/typeset.md b/docs/reference/typeset.md index de3c9b7d..a899f32c 100644 --- a/docs/reference/typeset.md +++ b/docs/reference/typeset.md @@ -17,11 +17,18 @@ question is whether the notation is right, rather than how to print it. ```python import math_spec as ms -print(ms.to_latex('model.yaml')) # amsmath align -print(ms.to_typst('model.yaml')) # compiles without a TeX toolchain -print(ms.to_markdown('model.yaml')) # renders as-is on GitHub +spec = ms.to_spec('model.yaml') # read and checked once, then printed three ways + +print(ms.to_latex(spec)) # amsmath align +print(ms.to_typst(spec)) # compiles without a TeX toolchain +print(ms.to_markdown(spec)) # renders as-is on GitHub ``` +Each of the three takes what `to_spec` takes — a path, the YAML, a mapping — +and reads it. Hand it the `Spec` instead and the file is read and checked once +rather than once per format, which is also how a `Spec` you already hold gets +printed without a second trip through the loader. + Or from a shell, where this belongs in a Makefile next to `pdflatex`: ```bash diff --git a/pixi.lock b/pixi.lock index 8c2bf4ab..692183a2 100644 --- a/pixi.lock +++ b/pixi.lock @@ -219,6 +219,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/tomlkit-0.15.1-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/trove-classifiers-2026.6.1.19-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/twine-7.0.0-pyhcf101f3_0.conda + - conda: https://conda.anaconda.org/conda-forge/noarch/types-pyyaml-6.0.12.20260815-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-inspection-0.4.4-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda @@ -335,6 +336,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/tomlkit-0.15.1-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/trove-classifiers-2026.6.1.19-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/twine-7.0.0-pyhcf101f3_0.conda + - conda: https://conda.anaconda.org/conda-forge/noarch/types-pyyaml-6.0.12.20260815-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-inspection-0.4.4-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda @@ -521,6 +523,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/tomlkit-0.15.1-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/trove-classifiers-2026.6.1.19-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/twine-7.0.0-pyhcf101f3_0.conda + - conda: https://conda.anaconda.org/conda-forge/noarch/types-pyyaml-6.0.12.20260815-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-inspection-0.4.4-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda @@ -709,6 +712,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/tomlkit-0.15.1-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/trove-classifiers-2026.6.1.19-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/twine-7.0.0-pyhcf101f3_0.conda + - conda: https://conda.anaconda.org/conda-forge/noarch/types-pyyaml-6.0.12.20260815-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-inspection-0.4.4-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda @@ -838,6 +842,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/pytest-xdist-3.8.0-pyhd8ed1ab_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/python_abi-3.12-8_cp312.conda - conda: https://conda.anaconda.org/conda-forge/noarch/tomli-2.4.1-pyhcf101f3_0.conda + - conda: https://conda.anaconda.org/conda-forge/noarch/types-pyyaml-6.0.12.20260815-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-inspection-0.4.4-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda @@ -861,6 +866,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/pytest-xdist-3.8.0-pyhd8ed1ab_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/python_abi-3.12-8_cp312.conda - conda: https://conda.anaconda.org/conda-forge/noarch/tomli-2.4.1-pyhcf101f3_0.conda + - conda: https://conda.anaconda.org/conda-forge/noarch/types-pyyaml-6.0.12.20260815-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-inspection-0.4.4-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda @@ -900,6 +906,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/pytest-xdist-3.8.0-pyhd8ed1ab_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/python_abi-3.12-8_cp312.conda - conda: https://conda.anaconda.org/conda-forge/noarch/tomli-2.4.1-pyhcf101f3_0.conda + - conda: https://conda.anaconda.org/conda-forge/noarch/types-pyyaml-6.0.12.20260815-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-inspection-0.4.4-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda @@ -940,6 +947,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/pytest-xdist-3.8.0-pyhd8ed1ab_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/python_abi-3.12-8_cp312.conda - conda: https://conda.anaconda.org/conda-forge/noarch/tomli-2.4.1-pyhcf101f3_0.conda + - conda: https://conda.anaconda.org/conda-forge/noarch/types-pyyaml-6.0.12.20260815-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-inspection-0.4.4-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda @@ -1012,6 +1020,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/pytest-xdist-3.8.0-pyhd8ed1ab_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/python_abi-3.13-8_cp313.conda - conda: https://conda.anaconda.org/conda-forge/noarch/tomli-2.4.1-pyhcf101f3_0.conda + - conda: https://conda.anaconda.org/conda-forge/noarch/types-pyyaml-6.0.12.20260815-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-inspection-0.4.4-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda @@ -1035,6 +1044,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/pytest-xdist-3.8.0-pyhd8ed1ab_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/python_abi-3.13-8_cp313.conda - conda: https://conda.anaconda.org/conda-forge/noarch/tomli-2.4.1-pyhcf101f3_0.conda + - conda: https://conda.anaconda.org/conda-forge/noarch/types-pyyaml-6.0.12.20260815-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-inspection-0.4.4-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda @@ -1075,6 +1085,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/pytest-xdist-3.8.0-pyhd8ed1ab_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/python_abi-3.13-8_cp313.conda - conda: https://conda.anaconda.org/conda-forge/noarch/tomli-2.4.1-pyhcf101f3_0.conda + - conda: https://conda.anaconda.org/conda-forge/noarch/types-pyyaml-6.0.12.20260815-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-inspection-0.4.4-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda @@ -1116,6 +1127,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/pytest-xdist-3.8.0-pyhd8ed1ab_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/python_abi-3.13-8_cp313.conda - conda: https://conda.anaconda.org/conda-forge/noarch/tomli-2.4.1-pyhcf101f3_0.conda + - conda: https://conda.anaconda.org/conda-forge/noarch/types-pyyaml-6.0.12.20260815-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-inspection-0.4.4-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda @@ -1189,6 +1201,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/pytest-xdist-3.8.0-pyhd8ed1ab_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/python_abi-3.14-8_cp314.conda - conda: https://conda.anaconda.org/conda-forge/noarch/tomli-2.4.1-pyhcf101f3_0.conda + - conda: https://conda.anaconda.org/conda-forge/noarch/types-pyyaml-6.0.12.20260815-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-inspection-0.4.4-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda @@ -1212,6 +1225,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/pytest-xdist-3.8.0-pyhd8ed1ab_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/python_abi-3.14-8_cp314.conda - conda: https://conda.anaconda.org/conda-forge/noarch/tomli-2.4.1-pyhcf101f3_0.conda + - conda: https://conda.anaconda.org/conda-forge/noarch/types-pyyaml-6.0.12.20260815-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-inspection-0.4.4-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda @@ -1253,6 +1267,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/pytest-xdist-3.8.0-pyhd8ed1ab_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/python_abi-3.14-8_cp314.conda - conda: https://conda.anaconda.org/conda-forge/noarch/tomli-2.4.1-pyhcf101f3_0.conda + - conda: https://conda.anaconda.org/conda-forge/noarch/types-pyyaml-6.0.12.20260815-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-inspection-0.4.4-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda @@ -1295,6 +1310,7 @@ environments: - conda: https://conda.anaconda.org/conda-forge/noarch/pytest-xdist-3.8.0-pyhd8ed1ab_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/python_abi-3.14-8_cp314.conda - conda: https://conda.anaconda.org/conda-forge/noarch/tomli-2.4.1-pyhcf101f3_0.conda + - conda: https://conda.anaconda.org/conda-forge/noarch/types-pyyaml-6.0.12.20260815-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing-inspection-0.4.4-pyhcf101f3_0.conda - conda: https://conda.anaconda.org/conda-forge/noarch/typing_extensions-4.16.0-pyhcf101f3_0.conda @@ -4862,6 +4878,18 @@ packages: run_exports: {} size: 42936 timestamp: 1785252141368 +- conda: https://conda.anaconda.org/conda-forge/noarch/types-pyyaml-6.0.12.20260815-pyhcf101f3_0.conda + sha256: 9790e34c023c33e3cebf9c9eeed411d468803eb6edbab3bc61116522fdaa32aa + md5: cd5c2725760ce602617c80a2b4c26d6d + depends: + - python >=3.10 + - python + license: Apache-2.0 AND MIT + purls: + - pkg:pypi/types-pyyaml?source=compressed-mapping + run_exports: {} + size: 27436 + timestamp: 1786859675692 - conda: https://conda.anaconda.org/conda-forge/noarch/typing-extensions-4.16.0-h69aa097_0.conda sha256: b141933ece3518f6d7b75dfb59451e2f26b405a44c18e2518a83e9a02e09315c md5: c680b5747e8c4c8f23dca0bb7042a8fc diff --git a/pixi.toml b/pixi.toml index 27ab6485..666f7207 100644 --- a/pixi.toml +++ b/pixi.toml @@ -60,6 +60,10 @@ pytest-xdist = ">=3.8" # that reports one more thing turns a green branch red without a commit. Held # to the same version the lint feature pins ruff to. pyrefly = "==1.2.0" +# yaml ships no annotations of its own, and `untyped-import` is an error: the +# stubs are what make the two `import yaml` lines in src/ answerable rather +# than suppressed. +types-pyyaml = ">=6.0" [feature.test.pypi-dependencies] # The PyPI bindings, which compile in-process — not conda-forge's `typst`, # which is the CLI binary and cannot be imported. Three compile tests in diff --git a/pyproject.toml b/pyproject.toml index 44fced4a..97044dfa 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -46,7 +46,7 @@ dependencies = [ # every model error to SchemaError. "pydantic>=2.1", # The two grammars — `expression:` and `where:` are strings to the schema and - # are parsed here (expression_parser.py, where_parser.py). + # are parsed here (expression_parser.py, _where_parser.py). "pyparsing>=3.1", "pyyaml>=6.0", ] @@ -115,7 +115,7 @@ extend-fixable = ["B", "SIM", "RUF", "C4", "UP"] "T20", "D", ] # likewise a generator, and only reachable as a command -"src/math_spec/where_parser.py" = [ +"src/math_spec/_where_parser.py" = [ "N806", ] # grammar tokens (NOT/AND/OR) are named after keywords "docs/static/hooks.py" = ["PERF401", "RUF015", "D"] @@ -142,6 +142,10 @@ runtime-evaluated-base-classes = ["pydantic.BaseModel"] python-version = "3.12" project-includes = ["src/math_spec"] preset = "strict" +# A callable taking `*args: Any, **kwargs: Any` otherwise matches every +# signature, which is how a parse action or a validator with the wrong shape +# gets passed without a word. Nothing in src/ relies on the loose reading. +strict-callable-subtyping = true # Warn-by-default rules promoted to error, so a regression fails the gate # instead of scrolling past. All of these are already clean; each one is a @@ -157,20 +161,23 @@ preset = "strict" # than switched off so a lambda outside those grammars still has to say what # it takes. [tool.pyrefly.errors] -implicit-import = true -missing-import = true -non-exhaustive-match = true -not-required-key-access = true -redundant-cast = true -redundant-condition = true -unknown-name = true -unnecessary-comparison = true -unnecessary-type-conversion = true -unreachable = true -unresolvable-dunder-all = true -untyped-import = true -unused-ignore = true -variance-mismatch = true +implicit-import = "error" +missing-import = "error" +no-any-return-explicit = "error" +no-any-return-implicit = "error" +non-exhaustive-match = "error" +not-required-key-access = "error" +redundant-cast = "error" +redundant-condition = "error" +unknown-name = "error" +unknown-variable-type = "error" +unnecessary-comparison = "error" +unnecessary-type-conversion = "error" +unreachable = "error" +unresolvable-dunder-all = "error" +untyped-import = "error" +unused-ignore = "error" +variance-mismatch = "error" [tool.pytest.ini_options] testpaths = ["tests"] diff --git a/schema/math-spec.schema.json b/schema/math-spec.schema.json index 2666a975..5a6870ed 100644 --- a/schema/math-spec.schema.json +++ b/schema/math-spec.schema.json @@ -2,7 +2,7 @@ "$defs": { "BoundsBlock": { "additionalProperties": false, - "description": "Variable bounds \u2014 each side is a number or parameter name.\n\nlinopy's defaults (``add_variables(lower=-inf, upper=inf)``): omitting a\nbound leaves the variable unbounded on that side, not implicitly\nnon-negative. Non-negativity is a real constraint, so the file says it.", + "description": "Variable bounds \u2014 each side is a number or parameter name.\n\nAn omitted bound leaves the variable unbounded on that side, not\nimplicitly non-negative.", "properties": { "lower": { "anyOf": [ @@ -112,7 +112,7 @@ "anyOf": [ { "additionalProperties": false, - "description": "A named quantity: one arithmetic expression, readable after a solve.\n\nWritten in YAML as a bare string, or as a mapping once it carries a\n``description:`` \u2014 and serialised back to whichever form it was written in,\nso a round trip through :meth:`Spec.to_yaml` reproduces the file::\n\n expressions:\n total_generation: sum(p, over=generator)\n emissions:\n expression: sum(p * rate, over=generator)\n description: CO2 released, the quantity the cap bounds\n\nA quantity whose value varies by **region** is written as ``cases:``\ninstead \u2014 one case per region over a declared ``foreach:``, no two of them\nclaiming one coordinate, and an ``otherwise:`` for the rest::\n\n previous_status:\n foreach: [snapshot, generator]\n cases:\n always_on: { when: \"not committable\", expression: 1 }\n boundary: { when: \"committable and position(snapshot) == 0\", expression: status_initial }\n otherwise: shift(status, over=snapshot, offset=1)\n\nSo the constraint that needs it names it, rather than being forked into one\ncopy per regime.", + "description": "A named quantity: one arithmetic expression, readable after a solve.\n\nWritten in YAML as a bare string, or as a mapping once it carries a\n``description:`` \u2014 and serialised back to whichever form it was written in,\nso a round trip through :meth:`Spec.to_yaml` reproduces the file::\n\n expressions:\n total_generation: sum(p, over=generator)\n emissions:\n expression: sum(p * rate, over=generator)\n description: CO2 released, the quantity the cap bounds\n\nA quantity whose value varies by region is written as ``cases:`` over a\ndeclared ``foreach:``, with an ``otherwise:`` for the rest \u2014 see the\nlanguage reference.", "properties": { "cases": { "additionalProperties": { @@ -395,7 +395,7 @@ }, "PiecewiseBlock": { "additionalProperties": false, - "description": "N expressions jointly pinned to a breakpoint-indexed piecewise curve.\n\nMirrors ``linopy.Spec.add_piecewise_formulation``. Each link is\n``[expression, values_parameter]`` or ``[expression, values_parameter,\nsign]``: *expression* is any affine expression string, *values_parameter*\nnames a parameter carrying the ``over`` dim, and *sign* bounds the link by\nthe curve instead of pinning it (at most one non-``\"==\"``, and only with\nexactly two links).\n\n``over`` names the breakpoint dimension; ``method`` is which of\n:data:`PIECEWISE_METHODS` restricts the weights; ``activity`` names what the weights sum\nto \u2014 1 where the block is unconditional, and a binary where a curve applies\nonly when something runs, which pins the formulation to 0 when it is 0; ``points`` names a\nboolean parameter saying how far each curve runs, for a model whose curves\nare not all the same length. Expanded before building into plain variables\nand constraints \u2014 see ``math_spec.piecewise``.", + "description": "N expressions jointly pinned to a breakpoint-indexed piecewise curve.\n\nMirrors ``linopy.Spec.add_piecewise_formulation``. Each link is\n``[expression, values_parameter]`` or ``[expression, values_parameter,\nsign]``: *expression* is any affine expression string, *values_parameter*\nnames a parameter carrying the ``over`` dim, and *sign* bounds the link by\nthe curve instead of pinning it (at most one non-``\"==\"``, and only with\nexactly two links).", "properties": { "activity": { "anyOf": [ @@ -476,9 +476,9 @@ "sign": { "default": "==", "enum": [ - "==", "<=", - ">=" + ">=", + "==" ], "title": "Sign", "type": "string" @@ -626,7 +626,7 @@ }, "$schema": "https://json-schema.org/draft/2020-12/schema", "additionalProperties": false, - "description": "The declared math \u2014 one YAML file, or one dict, validated. Nothing here has seen data.\n\nThe API is the ten declaration sections plus ``version`` and\n``description``, and two ways back out: :meth:`to_dict` for the model as\ndata, :meth:`to_yaml` for the file a reviewer reads. In goes through\n``to_spec``, which raises\n:class:`~math_spec.errors.LanguageError` on a model the language refuses.\n\nEverything else on this class is pydantic's, not a contract this package\nkeeps \u2014 ``model_json_schema()`` describes the shape pydantic validates\nrather than the language (checked in for editors as\n``schema/math_spec.schema.json``), and ``model_construct()`` skips validation\nentirely, so a ``Spec`` is valid when it was built the normal way.", + "description": "The declared math \u2014 one YAML file, or one dict, validated. Nothing here has seen data.\n\nA ``Spec`` that exists has passed the whole language: constructing one by\nany route \u2014 ``to_spec``, :meth:`model_validate`, the constructor \u2014 runs\nevery load-time check, expansion and expression pass included, and raises\n:class:`~math_spec.errors.LanguageError` on a model the language refuses.\nHolding one is the proof, so nothing downstream checks it again.\n\nThe API is the ten declaration sections plus ``version`` and\n``description``, and two ways back out: :meth:`to_dict` for the model as\ndata, :meth:`to_yaml` for the file a reviewer reads. Everything else on\nthis class is pydantic's, not a contract this package keeps.", "properties": { "constraints": { "additionalProperties": { diff --git a/src/math_spec/__init__.py b/src/math_spec/__init__.py index d1c7eeaa..5d8de90c 100644 --- a/src/math_spec/__init__.py +++ b/src/math_spec/__init__.py @@ -8,8 +8,7 @@ a :class:`~math_spec.program.Program` is what it *means* — and a conversion to each. The AST between them is this package's own: it is reachable by module path for a renderer that needs it, and out of ``__all__`` because a consumer -reads a program instead. ``__all__`` is the public surface, pinned by -``tests/test_public_surface.py``. +reads a program instead. """ from math_spec import program @@ -42,9 +41,6 @@ edge_error, unknown_operator_message, ) - -# Last: `math_spec.typesetting` reaches back for the two conversions, so those -# must be bound before it is imported. from math_spec.typesetting import ( FORMATS, SymbolTable, diff --git a/src/math_spec/__main__.py b/src/math_spec/__main__.py index 4fd183f6..c55c9e1f 100644 --- a/src/math_spec/__main__.py +++ b/src/math_spec/__main__.py @@ -5,9 +5,7 @@ """``python -m math_spec model.yaml`` — the shell front. ``check`` loads the file and prints the language's advice; one further verb -per typeset format, read off :data:`math_spec.typesetting.FORMATS`. No entry -point: ``python -m`` says which environment it ran in, which a bare name on -``PATH`` does not. +per typeset format, read off :data:`math_spec.typesetting.FORMATS`. """ from __future__ import annotations @@ -22,7 +20,7 @@ def parser() -> argparse.ArgumentParser: - """The verbs, built from ``FORMATS``; separate from :func:`main` so a test can read them off it.""" + """The verbs, built from ``FORMATS``.""" front = argparse.ArgumentParser(prog='python -m math_spec') verbs = front.add_subparsers(dest='verb', required=True) @@ -43,8 +41,7 @@ def parser() -> argparse.ArgumentParser: def main(argv: list[str] | None = None) -> int: """Run one verb; a refused file is its message on stderr and exit status 1. - Advice is not a refusal: ``check`` prints it and exits 0, since a note is - what a half-written model looks like too. + Advice is not a refusal: ``check`` prints it and exits 0. """ args = parser().parse_args(argv) if args.verb == 'check': @@ -64,7 +61,7 @@ def main(argv: list[str] | None = None) -> int: numbered=not args.no_numbers, ) if args.out: - Path(args.out).write_text(text) + Path(args.out).write_text(text, encoding='utf-8') else: sys.stdout.write(text) return 0 diff --git a/src/math_spec/_where_parser.py b/src/math_spec/_where_parser.py new file mode 100644 index 00000000..638bc63a --- /dev/null +++ b/src/math_spec/_where_parser.py @@ -0,0 +1,200 @@ +# SPDX-FileCopyrightText: math-spec Contributors +# +# SPDX-License-Identifier: MIT + +"""The where-string grammar and the ``Unresolved*`` nodes it emits, package-private. + +The resolved vocabulary lives in :mod:`math_spec.program`. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from functools import lru_cache +from typing import TYPE_CHECKING, Any, cast, get_args + +import pyparsing as pp + +from math_spec.expression_parser import NAME, REAL, parse_text +from math_spec.program import AndNode, BooleanLiteralNode, NotNode, OrNode, PredicateOperator + +if TYPE_CHECKING: + from collections.abc import Callable + + from math_spec.program import WhereNode + +# --------------------------------------------------------------------------- +# AST nodes +# --------------------------------------------------------------------------- + + +@dataclass(frozen=True) +class UnresolvedNameNode: + """A bare name — unresolved. ``resolution.py`` types it.""" + + name: str + + +@dataclass(frozen=True) +class UnresolvedComparisonNode: + """A comparison against an unresolved name. ``resolution.py`` types it.""" + + name: str + op: PredicateOperator + value: float | str + #: Whether the right-hand side 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. Consumed by resolution, never lowered. + quoted: bool = False + + +@dataclass(frozen=True) +class UnresolvedPositionNode: + """``position(dim[, by=lookup]) i`` before the names are checked; ``resolution.py`` types it.""" + + dimension: str + op: PredicateOperator + position: int + by: str | None = None + + +#: What resolution rewrites away on the where side — the three nodes whose +#: left-hand side is still a name the schema has not been asked about. +UnresolvedWhereNode = UnresolvedNameNode | UnresolvedComparisonNode | UnresolvedPositionNode + + +# --------------------------------------------------------------------------- +# Grammar +# --------------------------------------------------------------------------- + + +class _Quoted(str): + """A right-hand side that arrived in quotes; :func:`_comparison` turns it back into a flag.""" + + __slots__ = () + + +def _position_comparison(tokens: pp.ParseResults) -> UnresolvedPositionNode: + """``position(dim[, by=lookup]) i`` off the tokens the grammar captured.""" + *call, op, at = tokens + dimension, by = call[0], call[1] if len(call) > 1 else None + return UnresolvedPositionNode(str(dimension), op, at, None if by is None else str(by)) + + +def _comparison(tokens: pp.ParseResults) -> UnresolvedComparisonNode: + """``name literal`` off the tokens the grammar captured, the quoted marker turned into a flag.""" + name, op, value = tokens + quoted = isinstance(value, _Quoted) + return UnresolvedComparisonNode(str(name), op, str(value) if quoted else value, quoted) + + +def _build_where_grammar() -> pp.ParserElement: + """Build the pyparsing grammar for where strings. + + Both quote characters are accepted because YAML already owns one of them. + ``NOT`` binds tightest, then ``AND``, then ``OR``. ``position(...)`` leads + the alternation, since ``position`` would otherwise be read as a bare name. + """ + where_expr = pp.Forward() + + true_lit = pp.CaselessKeyword('True').set_parse_action(lambda: BooleanLiteralNode(True)) + false_lit = pp.CaselessKeyword('False').set_parse_action(lambda: BooleanLiteralNode(False)) + + # pyrefly: ignore[implicit-any-lambda] + number = pp.Regex(rf'-?({REAL}|\d+)').set_parse_action(lambda t: float(t[0])) + # pyrefly: ignore[implicit-any-lambda] + position = pp.Regex(r'-?\d+').set_parse_action(lambda t: int(t[0])) + + name = pp.Regex(NAME) + + quoted = (pp.QuotedString("'", esc_char='\\') | pp.QuotedString('"', esc_char='\\')).set_parse_action( + lambda t: _Quoted(t[0]) + ) + + grouped_by = pp.Suppress(',') + pp.Suppress(pp.Keyword('by')) + pp.Suppress('=') + name + comparator = pp.one_of(list(get_args(PredicateOperator))) + + position_call = ( + pp.Suppress(pp.Keyword('position')) + pp.Suppress('(') + name + pp.Optional(grouped_by) + pp.Suppress(')') + ) + position_comparison = (position_call + comparator + position).set_parse_action(_position_comparison) + + comparison = (name + comparator + (number | quoted | name)).set_parse_action(_comparison) + # pyrefly: ignore[implicit-any-lambda] + existence = name.copy().set_parse_action(lambda t: UnresolvedNameNode(t[0])) + + atom = ( + true_lit + | false_lit + | position_comparison + | comparison + | existence + | (pp.Suppress('(') + where_expr + pp.Suppress(')')) + ) + + NOT = pp.CaselessKeyword('NOT').suppress() + # pyrefly: ignore[implicit-any-lambda] + not_expr = (NOT + atom).set_parse_action(lambda t: NotNode(t[0])) | atom + + AND = pp.CaselessKeyword('AND').suppress() + and_expr = not_expr + pp.ZeroOrMore(AND + not_expr) + and_expr.set_parse_action(_folder(AndNode)) + + OR = pp.CaselessKeyword('OR').suppress() + or_expr = and_expr + pp.ZeroOrMore(OR + and_expr) + or_expr.set_parse_action(_folder(OrNode)) + + where_expr <<= or_expr + return where_expr + + +def _folder(node_type: type[AndNode] | type[OrNode]) -> Callable[[pp.ParseResults], Any]: + """A parse action left-folding a flat operator chain into *node_type*.""" + + def fold(tokens: pp.ParseResults) -> Any: + items = list(tokens) + result: WhereNode | UnresolvedWhereNode = items[0] + for item in items[1:]: + result = node_type(cast('WhereNode', result), item) + return result + + return fold + + +_WHERE_GRAMMAR = _build_where_grammar() + + +def _named_rewrite(text: str, loc: int) -> str | None: + """The rewrite for a connective habit of pandas or C at the token where the grammar gave up, or ``None``. + + ``!=``, ``<`` and ``>`` are legal here, so only the tokens no predicate + admits are diagnosed. + """ + rest = text[loc:].lstrip() + if rest.startswith('&'): + return "'&' is not the conjunction — both predicates at once is written AND." + if rest.startswith('|'): + return "'|' is not the disjunction — either predicate is written OR." + if rest.startswith(('~', '!')) and not rest.startswith('!='): + return f"'{rest[0]}' is not the negation — it is written NOT, before the predicate." + if rest.startswith('=') and not rest.startswith('=='): + return "'=' compares nothing — equality is written ==." + return None + + +@lru_cache(maxsize=4096) +def parse_where(text: str) -> WhereNode | 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. + + Raises: + SchemaError: If *text* is not a where string of the language. A + predictable mistake — ``&``/``|``/``~``/``!`` for a connective, a + lone ``=`` — is named with its rewrite beside the grammar's own + complaint. + """ + return cast('WhereNode | UnresolvedWhereNode', parse_text(_WHERE_GRAMMAR, text, 'where string', _named_rewrite)) diff --git a/src/math_spec/_yaml.py b/src/math_spec/_yaml.py index 197db7a6..375c65e8 100644 --- a/src/math_spec/_yaml.py +++ b/src/math_spec/_yaml.py @@ -4,27 +4,16 @@ """How this project reads a YAML file. -`yaml.safe_load` implements YAML 1.1, and two of its rules are actively wrong -for a language whose scalars are user data. The loader is the only layer that -can see them, so both are fixed here: +`yaml.safe_load` implements YAML 1.1, and two of its rules are wrong for a +language whose scalars are user data; both are fixed here: - **1.2 booleans.** ``on``/``off``/``yes``/``no``/``y``/``n`` are ordinary names in this language — a country code as a dimension, a mode as a lookup. - YAML 1.1 resolves them to ``True``/``False``, so the declaration the file - writes is not the one that reaches the schema. Only ``true``/``false`` are + YAML 1.1 resolves them to ``True``/``False``; only ``true``/``false`` are booleans here, which is the YAML 1.2 core schema. - **Duplicate keys.** 1.1 lets the last one win silently, discarding a declaration the file plainly contains. -Two further 1.1 coercions survive on purpose — the implicit timestamp -(``2024-01-01`` → ``date``) and sexagesimal ints (``12:30`` → ``750``). Neither -reaches a coordinate, which is data and never written here; a literal on the -other side of a ``where`` comparison is where one would be read as a label, and -there it is checked against the declared ``dtype`` (``resolution.py``). -``dtype: datetime`` is implemented — a label needs only an order and equality, -and nothing does arithmetic on a coordinate — so the timestamp coercion is the -*useful* reading there, not a hazard to route around. - The output is plain ``dict``/``str``: no loader wrapper reaches the schema or the AST. """ @@ -33,7 +22,7 @@ import re from pathlib import Path -from typing import Any +from typing import TYPE_CHECKING, Any import yaml @@ -43,15 +32,22 @@ _BOOL_1_2 = re.compile(r'^(?:true|True|TRUE|false|False|FALSE)$') -class _StrictLoader(yaml.SafeLoader): +if TYPE_CHECKING: + # Typed as SafeLoader: typeshed declares CSafeLoader unconditionally, and a PyYAML without libyaml lacks it. + _BaseLoader = yaml.SafeLoader +else: + # Same document either way: both drive the Python Resolver and SafeConstructor. + _BaseLoader = getattr(yaml, 'CSafeLoader', yaml.SafeLoader) + + +class _StrictLoader(_BaseLoader): """SafeLoader with 1.2 booleans. Duplicate keys are checked on the nodes.""" -#: The resolver table is rebuilt, not edited in place: it is inherited from -#: ``SafeLoader``, and mutating it would reconfigure PyYAML for the whole process. +# Rebuilt rather than edited: the table is inherited, and mutating it reconfigures PyYAML process-wide. _StrictLoader.yaml_implicit_resolvers = { ch: [(tag, rx) for tag, rx in pairs if tag != 'tag:yaml.org,2002:bool'] - for ch, pairs in yaml.SafeLoader.yaml_implicit_resolvers.items() + for ch, pairs in _BaseLoader.yaml_implicit_resolvers.items() } _StrictLoader.add_implicit_resolver('tag:yaml.org,2002:bool', _BOOL_1_2, list('tTfF')) @@ -69,7 +65,8 @@ def _check_duplicate_keys(node: yaml.Node, origin: str) -> None: """ if isinstance(node, yaml.MappingNode): seen: dict[Any, int] = {} - for key_node, value_node in node.value: + pairs: list[tuple[yaml.Node, yaml.Node]] = node.value + for key_node, value_node in pairs: line = key_node.start_mark.line + 1 if not isinstance(key_node, yaml.ScalarNode): msg = f'{origin}:{line}: a key must be a scalar — a name, not a list or a mapping.' @@ -94,17 +91,12 @@ def _check_duplicate_keys(node: yaml.Node, origin: str) -> None: def read_yaml(path: Path | str) -> dict[str, Any]: """Read *path* off disk and parse it, in YAML 1.2's reading of scalars.""" - return parse_yaml(Path(path).read_text(), str(path)) + return parse_yaml(Path(path).read_text(encoding='utf-8'), str(path)) def parse_yaml(text: str, origin: str = '') -> dict[str, Any]: """Parse YAML *text* as a mapping of sections. - The half of :func:`read_yaml` that does not touch the filesystem, so a - caller holding the text rather than the path — a test fixture, a doc block - — resolves scalars the same way. ``yaml.safe_load`` is 1.1 and would read - ``no`` as a boolean, which is the divergence this module exists to remove. - Args: text: The YAML source. origin: What a load error calls this source — a file's path, or the default for text that never was one. diff --git a/src/math_spec/advice.py b/src/math_spec/advice.py index 9cfc146a..7e6a082a 100644 --- a/src/math_spec/advice.py +++ b/src/math_spec/advice.py @@ -4,37 +4,29 @@ """Advice — what is decidable without data and is a note rather than a refusal. -One door, :func:`advice`, over every pass of that kind. A consumer prints what -it returns; the sentences are the language's, so two consumers cannot come to -disagree about them. +One door, :func:`advice`, over every pass of that kind. """ from __future__ import annotations from typing import TYPE_CHECKING -from math_spec import program as program_ 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 At, GroupSum, walk if TYPE_CHECKING: from pathlib import Path from typing import Any from math_spec.model import Spec - from math_spec.program import ExpressionNode, Program + from math_spec.program import Program def advice(model: str | Path | dict[str, Any] | Spec | Program) -> tuple[Advice, ...]: """Everything the language advises about *model* — never an error, decidable without data. - A dimension that is never an axis is a label space wearing the wrong - declaration, or unused. A variable the objective drives toward an open - bound with no constraint naming it makes the model unbounded for any data - there is. Both are what a half-written model looks like too, which is why - they are advice and ``to_spec`` stays open to them. - Args: model: A YAML path, a mapping, a loaded :class:`Spec`, or a :class:`Program`. Both passes read the program, so the four @@ -53,8 +45,7 @@ def _never_an_axis(program: Program) -> list[Advice]: axes: set[str] = set() for declaration in (*program.parameters.values(), *program.variables.values(), *program.constraints.values()): axes.update(declaration.dims) - for e in program.expressions: - axes |= _produced_axes(e) + axes |= _produced_axes(program) targeted = {lk.target: (dimension, lk.name) for dimension, lk in program.lookups if lk.target is not None} notes: list[Advice] = [] @@ -81,18 +72,15 @@ def _never_an_axis(program: Program) -> list[Advice]: return notes -def _produced_axes(e: ExpressionNode) -> set[str]: - """The axes an expression *creates*, beyond what its declarations index. +def _produced_axes(program: Program) -> set[str]: + """The axes the expressions create beyond what any declaration indexes. - ``sum(by=)`` lands terms on its target and ``at()`` spreads onto its fine - dimension, so both are axes even when no declaration is indexed by them — - an objective may group into a dimension and then implicitly sum it away. + ``sum(by=)`` lands on its target and ``at()`` spreads onto its fine dimension. """ - out: set[str] = set() - if isinstance(e, program_.GroupSum): - out |= set(e.into) - if isinstance(e, program_.At): - out.add(e.over) - for child in program_.children(e): - out |= _produced_axes(child) - return out + axes: set[str] = set() + for node in walk(*program.expressions): + if isinstance(node, GroupSum): + axes.update(node.into) + elif isinstance(node, At): + axes.add(node.over) + return axes diff --git a/src/math_spec/boundedness.py b/src/math_spec/boundedness.py index 8fe1629e..daf03c7d 100644 --- a/src/math_spec/boundedness.py +++ b/src/math_spec/boundedness.py @@ -4,15 +4,11 @@ """Provably unbounded models, named before a solver says a bare ``unbounded``. -A variable that is unbounded on the side its objective term improves toward -**and** appears in no constraint runs to infinity for any data at all. Advice -rather than a refusal, because the same shape is what a half-written model -looks like. - -Which side improves is read off the *sign* the variable enters the objective -with: under ``minimize`` a ``+v`` term runs down toward ``lower``. Where that -sign is not decidable without data — a parameter coefficient, or occurrences -of both signs — nothing is claimed. +A variable unbounded on the side its objective term improves toward, and named +by no constraint, runs to infinity for any data. Which side is read off the +sign the variable enters the objective with: under ``minimize`` a ``+v`` term +runs down toward ``lower``. Where that sign is not decidable without data — a +parameter coefficient, or occurrences of both signs — nothing is claimed. """ from __future__ import annotations @@ -37,6 +33,7 @@ Translate, Variable, Window, + children, variables_of, ) @@ -48,19 +45,20 @@ #: appearing with both signs (which may cancel), a divisor carrying one. Sign = Literal['+', '-'] | None +#: Which of a variable's two bounds a term drives it toward. +BoundSide = Literal['lower', 'upper'] + #: The bound value that leaves each side open. A ``lower`` of ``+inf`` is not #: this — that model is empty, not unbounded — so the match is by value. -_OPEN = {'lower': -math.inf, 'upper': math.inf} +_OPEN: dict[BoundSide, float] = {'lower': -math.inf, 'upper': math.inf} def unbounded_notes(program: Program) -> list[Advice]: """Name every variable the objective can drive to infinity unopposed. - Asked of the program rather than the file: every fact the rule reads is a - declaration — the objective's sense and its terms, the variables each - constraint names, the two bounds — and by the time a program exists a - ``piecewise:`` block has already become the constraints it expands into, - which is where the variables it names are held. + Args: + program: The lowered program, in which ``piecewise:`` has already + become the constraints it expands into. Returns: One note per variable that is unbounded on the side its objective term @@ -74,14 +72,14 @@ def unbounded_notes(program: Program) -> list[Advice]: constrained |= variables_of(constraint.lhs, constraint.rhs) signs: dict[str, Sign] = {} - _walk(program.objective.expression, '+', signs) + _record_signs(program.objective.expression, '+', signs) minimize = program.objective.sense == 'minimize' notes: list[Advice] = [] for vname, sign in signs.items(): if sign is None or vname in constrained: continue - side = 'lower' if minimize == (sign == '+') else 'upper' + side: BoundSide = 'lower' if minimize == (sign == '+') else 'upper' if _is_open(program.variables[vname], side): notes.append( Advice( @@ -97,13 +95,10 @@ def unbounded_notes(program: Program) -> list[Advice]: return notes -def _is_open(vdef: VariableDeclaration, side: str) -> bool: - """Whether *vdef* declares nothing at all on *side*. +def _is_open(vdef: VariableDeclaration, side: BoundSide) -> bool: + """Whether *vdef*'s bound on *side* is the open value itself. - A bound naming a parameter is finite or not by data this pass does not - have, which is why the match is against the open value rather than for a - missing bound. A ``binary`` variable needs no case of its own: it reaches - the program with the 0/1 bounds its domain fixes. + A bound naming a parameter is finite or not by data, so it does not count. """ bound = vdef.lower if side == 'lower' else vdef.upper return bound == Constant(_OPEN[side]) @@ -132,15 +127,13 @@ def _coefficient_sign(node: ExpressionNode) -> Sign: return None -def _walk(node: ExpressionNode, sign: Sign, signs: dict[str, Sign]) -> None: +def _record_signs(node: ExpressionNode, sign: Sign, signs: dict[str, Sign]) -> None: """Record the sign each variable under *node* carries into the objective. - *signs* accumulates, and a variable reached twice with different signs — or - once with an undecidable one — lands on ``None``, which claims nothing. - Every shape node sums its operand's terms with coefficient 1, being a - reduction, a re-index or a window, so each hands *sign* on unchanged. A - power carries no sign in either half, which is what makes a degree-2 term - claim nothing. + A variable reached twice with different signs, or once with an undecidable + one, lands on ``None``, which claims nothing. A reduction, a re-index, a + window and a cases selection pass *sign* to their children unchanged; a + power carries no sign in either half. """ if isinstance(node, Variable): signs[node.name] = sign if signs.setdefault(node.name, sign) == sign else None @@ -148,30 +141,26 @@ def _walk(node: ExpressionNode, sign: Sign, signs: dict[str, Sign]) -> None: if isinstance(node, Constant | Parameter): return if isinstance(node, Negate): - _walk(node.operand, _flip(sign), signs) + _record_signs(node.operand, _flip(sign), signs) return if isinstance(node, Add): - _walk(node.left, sign, signs) - _walk(node.right, sign, signs) + _record_signs(node.left, sign, signs) + _record_signs(node.right, sign, signs) return if isinstance(node, Multiply): - _walk(node.left, _times(sign, _coefficient_sign(node.right)), signs) - _walk(node.right, _times(sign, _coefficient_sign(node.left)), signs) + _record_signs(node.left, _times(sign, _coefficient_sign(node.right)), signs) + _record_signs(node.right, _times(sign, _coefficient_sign(node.left)), signs) return if isinstance(node, Divide): - _walk(node.numerator, _times(sign, _coefficient_sign(node.divisor)), signs) - _walk(node.divisor, None, signs) + _record_signs(node.numerator, _times(sign, _coefficient_sign(node.divisor)), signs) + _record_signs(node.divisor, None, signs) return if isinstance(node, Power): - _walk(node.base, None, signs) - _walk(node.exponent, None, signs) - return - if isinstance(node, Sum | GroupSum | At | Translate | Window): - _walk(node.operand, sign, signs) + _record_signs(node.base, None, signs) + _record_signs(node.exponent, None, signs) return - if isinstance(node, Cases): - # a selection, not a sum: whichever region applies stands where the whole value does - for region in node.regions: - _walk(region.value, sign, signs) + if isinstance(node, Sum | GroupSum | At | Translate | Window | Cases): + for child in children(node): + _record_signs(child, sign, signs) return assert_never(node) diff --git a/src/math_spec/degree.py b/src/math_spec/degree.py index b478b8f5..995450dc 100644 --- a/src/math_spec/degree.py +++ b/src/math_spec/degree.py @@ -38,9 +38,10 @@ def carries_variable(node: ExpressionNode) -> bool: - """Whether *node* contains a decision variable. + """Whether *node* contains a decision variable, over the core AST. - An unresolved node reaching here is a resolution bug, so it is refused + :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. """ if isinstance(node, VariableNode): @@ -70,23 +71,7 @@ def _adds(node: ExpressionNode) -> bool: return False -def is_quadratic(node: ExpressionNode) -> bool: - """Whether *node* multiplies two variable-carrying operands. - - What :func:`check_binary` refuses at ``ceiling=1``, asked of a whole - expression rather than of one node. - """ - if ( - isinstance(node, BinaryOperatorNode) - and node.op == '*' - and carries_variable(node.left) - and carries_variable(node.right) - ): - return True - return any(is_quadratic(child) for child in children(node)) - - -def check_binary(node: BinaryOperatorNode, context: str | None = None, *, ceiling: int = 1) -> None: +def check_binary(node: BinaryOperatorNode, context: str, *, ceiling: int) -> None: """Check that *node* stays inside the degree its position allows. Args: @@ -159,7 +144,6 @@ def _above_the_ceiling_message(where: str, degree: int) -> str: def _a_variable_under_a_power_message(where: str) -> str: - """A variable base is a degree question; a variable exponent has no degree until the data arrives.""" return ( f'{where}`**` is not in the language over variables: it takes a base and an exponent that ' f'carry none.\n' @@ -170,7 +154,6 @@ def _a_variable_under_a_power_message(where: str) -> str: def _degree_two_here_message(where: str) -> str: - """Names the position, not the math — the same product is admissible one declaration away.""" return ( f'{where}both factors of a product contain variables, which is degree 2. ' f'The **objective and constraints** take that; a bound, a named expression ' @@ -202,14 +185,15 @@ def _multi_term(node: ExpressionNode) -> bool: 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. """ - if isinstance(node, FunctionCallNode): - if node.name in _REDUCTIONS and any(carries_variable(a) for a in node.args): - return True - return any(_multi_term(c) for c in children(node)) - if isinstance(node, BinaryOperatorNode): - if node.op in ('+', '-') and carries_variable(node.left) and carries_variable(node.right): - return True - return _multi_term(node.left) or _multi_term(node.right) + if isinstance(node, FunctionCallNode) and node.name in _REDUCTIONS and any(carries_variable(a) for a in node.args): + return True + if ( + isinstance(node, BinaryOperatorNode) + and node.op in ('+', '-') + and carries_variable(node.left) + and carries_variable(node.right) + ): + return True return any(_multi_term(c) for c in children(node)) diff --git a/src/math_spec/dimensions.py b/src/math_spec/dimensions.py index 64d1804b..5b76e554 100644 --- a/src/math_spec/dimensions.py +++ b/src/math_spec/dimensions.py @@ -4,29 +4,18 @@ """Static dim-set checking — a type system whose type is a set of dim names. -Parameter ``dims`` are declared, variable ``foreach`` is declared, and operator -dimension arguments are name-checked, so **every node's dim set is computable -before any data is bound**. That is the whole basis of this pass: it runs at -load time, on the resolved core AST, so every consumer gets the same answer by -construction. The per-node rules are the "Dim algebra" table in -``docs/reference/language/expressions.md``; at the declaration level:: - - constraint -> the dims of both sides together must *equal* foreach - where -> the predicate's dims must not exceed the frame - bounds -> the bound parameter's dims must not exceed foreach - -The direction that matters most is the *stray* dim: one the frame does not -declare broadcasts silently at build time, so the same YAML quietly builds a -bigger model than it reads as. The missing direction is checked too, a foreach -dim the equation never uses just repeating one row across it — nearly always a -typo. +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 +``docs/reference/language/expressions.md``; a constraint's two sides together +must equal its ``foreach``, and a where or a bound may not exceed the frame. """ from __future__ import annotations -from typing import TYPE_CHECKING, assert_never +import math +from typing import TYPE_CHECKING, NamedTuple, assert_never -from math_spec.degree import carries_variable +import math_spec.degree as degree from math_spec.errors import DimensionError from math_spec.expression_parser import ( ArithmeticNode, @@ -46,20 +35,23 @@ VariableNode, case_context, ) -from math_spec.operators import BUILTINS, edge_error -from math_spec.resolution import Namespace, expression_of, where_of -from math_spec.where_parser import ( +from math_spec.operators import BUILTINS +from math_spec.program import ( DimensionComparisonNode, DimensionPositionNode, + LookupComparisonNode, + LookupDefinedNode, + LookupPairComparisonNode, + Mask, ParameterComparisonNode, ParameterDefinedNode, VariableDefinedNode, - WhereNode, - _atom_dims, - atoms, ) +from math_spec.resolution import Namespace, expression_of, where_of if TYPE_CHECKING: + from collections.abc import Callable + from math_spec.model import Spec @@ -83,7 +75,7 @@ def _dims( schema: Spec, context: str, ) -> frozenset[str]: - """The recursive worker under :func:`dims_of`; a binary operator takes the union of its sides, unchecked, since :func:`check_schema` compares the result with the declared frame.""" + """The recursive worker under :func:`dims_of`.""" if isinstance(node, NumberNode): return frozenset() @@ -107,167 +99,194 @@ def _dims( return _dims_call(node, schema, context) if isinstance(node, CasesNode): - # the declared frame, not the union of the cases: one narrower than it - # broadcasts, as a parameter with fewer dims does - return frozenset(schema.expressions[node.name].foreach or ()) + return _cases_dims(node, schema) assert_never(node) -def _dims_call( - node: FunctionCallNode, - schema: Spec, - context: str, -) -> frozenset[str]: - """The dim rule of one operator call. - - ``sum`` with neither ``over=`` nor ``by=`` takes every dim the operand - carries, so its result is scalar. ``by=`` reduces the lookup's own dim - *into* its target rather than away. - ``at`` is the adjoint of ``sum``, one mapping table walked either way: - ``sum`` consumes the dim the lookup is *over*, ``at`` the dim it maps - *into*, and each produces the other. +def _cases_dims(node: CasesNode, schema: Spec) -> frozenset[str]: + """The declared frame rather than the union of the arms. + + A narrower arm broadcasts, as a parameter with fewer dims does. """ - if node.name == 'sum': - inner = _dims(node.args[0], schema, context) - 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: - over = node.kwargs['over'] - assert isinstance(over, DimensionNode) - if over.name not in inner: - raise DimensionError( - f'{context}: sum(over={over.name}) but the expression has dims ' - f'{sorted(inner)}. Summing over a dim the operand does not carry ' - f'is a no-op that builds and solves wrong — drop the sum, or fix ' - f'the dim.' - ) - return inner - {over.name} - - assert isinstance(by, LookupNode) - if by.dimension not in inner: + return frozenset(schema.expressions[node.name].foreach or ()) + + +def _not_carried(context: str, call: str, inner: frozenset[str], rewrite: str) -> str: + """The refusal for an operator walking a dim its operand does not carry; *rewrite* is the operator's own.""" + return ( + f'{context}: {call} but the expression has dims {sorted(inner)}. An operator over a dim the ' + f'operand does not carry is a no-op that builds and solves wrong — {rewrite}.' + ) + + +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: FunctionCallNode, inner: frozenset[str], schema: Spec, context: str) -> frozenset[str]: + """``sum`` reduces a dim away, or through a lookup into the dim it maps *into*.""" + by = node.kwargs.get('by') + if by is None and 'over' not in node.kwargs: + if not inner: raise DimensionError( - f"{context}: sum(by={by.shown}) consumes '{by.dimension}', the dim it maps " - f'out of, but the expression has dims {sorted(inner)}. Summing ' - f'over a dim the operand does not carry is a no-op that builds and ' - f'solves wrong — drop the sum, or fix the dim.' + 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.' ) - collides = sorted(set(by.into) & (inner - {by.dimension})) - if collides: - raise DimensionError( - f'{context}: sum(by={by.shown}) targets {collides}, ' - f'which the expression already carries ({sorted(inner)}). The result would ' - f"need {collides} twice — once as the operand's own dim and once as the " - f'group it is placed into. Sum over one of the two first, ' - f'or group into a dimension the operand does not have.' + return frozenset() + if by is None: + over = node.kwargs['over'] + assert isinstance(over, DimensionNode) + if over.name not in inner: + raise DimensionError(_not_carried(context, f'sum(over={over.name})', inner, 'drop the sum, or fix the dim')) + return inner - {over.name} + + assert isinstance(by, LookupNode) + if by.dimension not in inner: + raise DimensionError( + _not_carried( + context, + f"sum(by={by.shown}) consumes '{by.dimension}', the dim it maps out of,", + inner, + 'drop the sum, or fix the dim', ) - return (inner - {by.dimension}) | set(by.into) - - if node.name == 'at': - inner = _dims(node.args[0], schema, context) - by = node.kwargs['by'] - assert isinstance(by, LookupNode) - absent = sorted(set(by.into) - inner) - if absent: - raise DimensionError( - f'{context}: at(by={by.shown}) reads through ' - 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.' + ) + collides = sorted(set(by.into) & (inner - {by.dimension})) + if collides: + raise DimensionError( + f'{context}: sum(by={by.shown}) targets {collides}, ' + f'which the expression already carries ({sorted(inner)}). The result would ' + f"need {collides} twice — once as the operand's own dim and once as the " + f'group it is placed into. Sum over one of the two first, ' + f'or group into a dimension the operand does not have.' + ) + return (inner - {by.dimension}) | set(by.into) + + +def _at_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spec, context: str) -> frozenset[str]: + """``at`` is the adjoint of ``sum(by=)``: it consumes the dim a lookup maps *into* and produces the one it is over.""" + by = node.kwargs['by'] + assert isinstance(by, LookupNode) + absent = sorted(set(by.into) - inner) + if absent: + raise DimensionError( + f'{context}: at(by={by.shown}) reads through ' + 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.' + ) + if by.dimension in inner - set(by.into): + raise DimensionError( + f'{context}: at(by={by.shown}) places terms onto ' + f"'{by.dimension}', which the expression already carries ({sorted(inner)}). " + f"The result would need '{by.dimension}' twice — once as the operand's own " + f'dim and once as the dim it is spread onto. Sum over one of the two first.' + ) + return (inner - set(by.into)) | {by.dimension} + + +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['over'] + assert isinstance(over, DimensionNode) + if over.name not in inner: + raise DimensionError( + _not_carried( + context, + f'{node.name}(over={over.name})', + inner, + f'walk a dim the operand carries, or drop the {node.name}', ) - if by.dimension in inner - set(by.into): + ) + _check_named_amount(node, over.name, inner, schema, context) + _check_amount_form(node, context) + _check_edge(node, context) + partition = node.kwargs.get('by') + if partition is not None: + assert isinstance(partition, LookupNode) + if len(partition.names) > 1: raise DimensionError( - f'{context}: at(by={by.shown}) places terms onto ' - f"'{by.dimension}', which the expression already carries ({sorted(inner)}). " - f"The result would need '{by.dimension}' twice — once as the operand's own " - f'dim and once as the dim it is spread onto. Sum over one of the two first.' + f'{context}: {node.name}(over={over.name}, by={partition.shown}) partitions by ' + f'several lookups at once. A partition says which rows are neighbours rather than ' + f'which group a term lands in, so it names one lookup — partition by a lookup whose ' + f'values already distinguish them.' ) - return (inner - set(by.into)) | {by.dimension} - - if node.name in ('shift', 'sum_back'): - inner = _dims(node.args[0], schema, context) - over = node.kwargs['over'] - assert isinstance(over, DimensionNode) - if over.name not in inner: + if partition.dimension != over.name: raise DimensionError( - f'{context}: {node.name}(over={over.name}) but the expression has dims {sorted(inner)}.' + f'{context}: {node.name}(over={over.name}, by={partition.shown}) walks ' + f"'{over.name}' but groups by a lookup over '{partition.dimension}'. No row of " + f"'{over.name}' carries it, so no coordinate has a neighbour inside a group — " + f"partition by a lookup over '{over.name}'." ) - _check_named_amount(node, over.name, inner, schema, context) - _check_amount_form(node, context) - _check_edge(node, context) - partition = node.kwargs.get('by') - if partition is not None: - assert isinstance(partition, LookupNode) - if len(partition.names) > 1: - raise DimensionError( - f'{context}: {node.name}(over={over.name}, by={partition.shown}) partitions by ' - f'several lookups at once. A partition says which rows are neighbours rather than ' - f'which group a term lands in, so it names one lookup — partition by a lookup whose ' - f'values already distinguish them.' - ) - if partition.dimension != over.name: - raise DimensionError( - f'{context}: {node.name}(over={over.name}, by={partition.shown}) walks ' - f"'{over.name}' but groups by a lookup over '{partition.dimension}'. No row of " - f"'{over.name}' carries it, so no coordinate has a neighbour inside a group — " - f"partition by a lookup over '{over.name}'." - ) - return inner - - msg = f"operator '{node.name}' reached the dim checker without a rule; resolution admits only BUILTINS." - raise AssertionError(msg) - - -#: Per axis-walking operator: the word its errors call the amount it takes, -#: why negating a named one at the call site is not what the caller means, and -#: what a named one that varies along the axis it walks becomes. -_AMOUNT_WORDING = { - 'shift': ( + return inner + + +#: 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 an axis-walking operator's errors 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 along the axis it walks 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': ( + '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 _check_amount_form(node: FunctionCallNode, context: str) -> None: - """What an ``offset=`` or ``within=`` may be written as, before it is read. - - A whole number or a parameter name, and nothing else — decidable from the - file, so refused here rather than by whoever tries to build it. - """ +def _amount_of(node: FunctionCallNode) -> tuple[str, ArithmeticNode]: + """The kwarg an axis-walking operator takes its amount through, and the value written there.""" (kwarg,) = BUILTINS[node.name].required_value_kwargs - amount = node.kwargs[kwarg] - if isinstance(amount, ParameterNode): - return - if node.name == 'shift': - if not (isinstance(amount, NumberNode) and int(amount.value) == amount.value): - raise DimensionError( - f'{context}: shift(offset=...) must be a whole number, or the name of an integer ' - f'parameter when the offset differs per entity — a lead time, a transit time, a ' - f'minimum up time.' - ) + 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 ``within=`` 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 - if not (isinstance(amount, NumberNode) and int(amount.value) == amount.value and amount.value >= 1): - raise DimensionError( - f'{context}: sum_back(within=...) needs a whole number of positions of at least 1, or ' - f'the name of an integer parameter when the window differs per entity. A width of 1 is ' - f'the operand itself.' - ) + raise DimensionError(f'{context}: {node.name}({kwarg}=...) {_AMOUNTS[node.name].form}') def _check_edge(node: FunctionCallNode, context: str) -> None: @@ -290,7 +309,7 @@ def _check_edge(node: FunctionCallNode, context: str) -> None: if isinstance(edge, EdgeNode): return fill = _edge_fill(edge, context) - has_var = carries_variable(node.args[0]) + 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 ' @@ -320,9 +339,10 @@ 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 - if not isinstance(edge, NumberNode): - raise DimensionError(f'{context}: {edge_error("shift", "...")}') - return float(edge.value) + 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: @@ -342,14 +362,7 @@ def _named_offset_edge_message(name: str) -> str: def _shift_over_data_message(context: str) -> str: - """The three ways out, one of which is two things at once. - - A ``where`` is a *companion* to ``edge=``, not an alternative: the refusal - is decided on the expression alone so a mask does not lift it, and - ``edge=0`` alone leaves a row at the vacated coordinate whose bound is that - zero — the silent pinning this refusal exists to prevent. Either one alone - is wrong, so the message says so rather than listing them as alternatives. - """ + """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' @@ -362,22 +375,12 @@ def _shift_over_data_message(context: str) -> str: def _check_named_amount(node: FunctionCallNode, over: str, inner: frozenset[str], schema: Spec, context: str) -> None: - """The rules that hold of an ``offset=`` or ``within=`` naming a parameter. - - They are about the *amount* rather than about a dim set, but they live here - because here is where the schema is in hand — a parameter's ``dtype`` and - its ``dims`` are read off the same declaration, and splitting a documented - set across two passes would give one rule of it several voices. - - A literal breaks none of them: it parses as a number, and a number has - neither a dtype to declare nor dims to vary over. - """ - (kwarg,) = BUILTINS[node.name].required_value_kwargs - noun, negated, varies = _AMOUNT_WORDING[node.name] - amount = node.kwargs[kwarg] + """The rules that hold of an ``offset=`` or ``within=`` 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 {noun}. {negated}' + f'{context}: {node.name}({kwarg}={amount.op}{amount.operand.name}) negates a named {words.noun}. {words.negated}' ) if not isinstance(amount, ParameterNode): return @@ -387,24 +390,24 @@ def _check_named_amount(node: FunctionCallNode, over: str, inner: frozenset[str] f'{context}: {node.name}({kwarg}={amount.name}) counts positions along ' f"'{over}', but '{amount.name}' 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 {noun} has nowhere to arrive from.' + f'column, so a fractional {words.noun} has nowhere to arrive from.' ) if over in declared.dims: raise DimensionError( f'{context}: {node.name}({kwarg}={amount.name}) walks ' f"'{over}', but '{amount.name}' is declared over {sorted(declared.dims)}, which " - f'carries it. A named {noun} that varies along the axis it walks is {varies} ' + f'carries it. A named {words.noun} that varies along the axis it walks is {words.varies} ' f"— declare '{amount.name}' over dims '{over}' is not one of." ) partition = node.kwargs.get('by') groups = frozenset(partition.into) if isinstance(partition, LookupNode) else frozenset() if stray := sorted(frozenset(declared.dims) - inner - groups): raise DimensionError( - f'{context}: {node.name}({kwarg}={amount.name}) reads its {noun} at the coordinate it ' + f'{context}: {node.name}({kwarg}={amount.name}) reads its {words.noun} at the coordinate it ' f"walks, but '{amount.name}' 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 lookup into " - f'one of {stray}, so that each group is reached by its own {noun}.' + f'one of {stray}, so that each group is reached by its own {words.noun}.' ) @@ -424,7 +427,7 @@ def check_schema(schema: Spec) -> None: for vname, vdef in schema.variables.items(): frame = frozenset(vdef.foreach) context = f"Variable '{vname}'" - _check_where_dims(where_of(vdef.where, ns, context), schema, frame, context) + _check_where_dims(where_of(vdef.where, ns, context), frame, context) for side in ('lower', 'upper'): bound = getattr(vdef.bounds, side) if isinstance(bound, str): @@ -442,7 +445,7 @@ def check_schema(schema: Spec) -> None: frame = frozenset(block.foreach or []) for case_name, case in block.cases.items(): context = case_context(ename, case_name) - _check_where_dims(where_of(case.when, ns, context), schema, frame, context) + _check_where_dims(where_of(case.when, ns, context), frame, context) _check_value_dims(case.expression, schema, ns, frame, context) assert block.otherwise is not None _check_value_dims(block.otherwise, schema, ns, frame, case_context(ename, None)) @@ -450,7 +453,7 @@ def check_schema(schema: Spec) -> None: for cname, cdef in schema.constraints.items(): frame = frozenset(cdef.foreach) context = f"Constraint '{cname}'" - _check_where_dims(where_of(cdef.where, ns, context), schema, frame, context) + _check_where_dims(where_of(cdef.where, ns, context), frame, context) got = dims_of(expression_of(cdef.expression, schema, ns, context), schema, context) if got != frame: stray, missing = sorted(got - frame), sorted(frame - got) @@ -497,58 +500,34 @@ def _check_value_dims( def _check_where_dims( - node: WhereNode | None, - schema: Spec, + mask: Mask | None, frame: frozenset[str], context: str, ) -> None: - """A predicate may only test dims the frame carries. + """A predicate may only test dims the frame carries; reducing an outside dim to fit would fail open. - Reducing an outside dim to fit — with ``any()``, say — is a mask that fails - *open*, silently including everything. It is rejected here, at load time. - - **Which dims a leaf reads is not decided here.** That is - :func:`~math_spec.where_parser.dims_read`'s rule, and this walks the same - leaves by the same reading, so a mask cannot be checked against one answer - and built against another. What is decided here is the wording, and the - wording is per leaf: a reader told only that the predicate leaves the frame - would still have to work out which half of it did. + The refusal names the leaf that left the frame, reading its dims as + :attr:`~math_spec.program.Mask.dims` does. """ - if node is None: + if mask is None: return - name_dims = _name_dims(schema) - for atom in atoms(node): - if not (outside := sorted(_atom_dims(atom, name_dims) - frame)): + for atom in mask.atoms: + if not (outside := sorted(Mask(atom).dims - frame)): continue - if isinstance(atom, (ParameterDefinedNode, ParameterComparisonNode)): - raise DimensionError( - f"{context}: where-parameter '{atom.name}' has dims " - f'{outside} outside the frame {sorted(frame)}. Reducing ' - f'a mask over an unlisted dim would silently widen it.' - ) - if isinstance(atom, VariableDefinedNode): - raise DimensionError( - f"{context}: where-variable '{atom.name}' has dims " - f'{outside} outside the frame {sorted(frame)}. A mask ' - f'reducing over an unlisted dim would silently widen it — say which ' - f'reduction you mean.' - ) - if isinstance(atom, (DimensionComparisonNode, DimensionPositionNode)): - raise DimensionError( - f"{context}: where-comparison on dimension '{atom.name}', which is not in the frame {sorted(frame)}." - ) + match atom: + case ParameterDefinedNode() | ParameterComparisonNode(): + noun = 'parameter' + case VariableDefinedNode(): + noun = 'variable' + case DimensionComparisonNode() | DimensionPositionNode(): + noun = 'dimension' + case LookupComparisonNode() | LookupPairComparisonNode() | LookupDefinedNode(): + noun = 'lookup' + case _: + assert_never(atom) raise DimensionError( - f"{context}: where-comparison on lookup '{atom.name}', which is over " - f"dimension '{atom.over}' — not in the frame {sorted(frame)}. A lookup is " - f'read on the dim it maps out of, so that dim has to be one the ' - f'declaration ranges over.' + f"{context}: where-{noun} '{atom.name}' reads dims {outside} outside the frame {sorted(frame)}. " + f'Reducing a mask over an unlisted dim would silently widen it — add the dim to foreach, ' + f'or test a name the frame carries.' ) - - -def _name_dims(schema: Spec) -> dict[str, tuple[str, ...]]: - """Every declared name to the dims it is read through — what :func:`~math_spec.where_parser.dims_read` takes.""" - return { - **{name: tuple(pdef.dims) for name, pdef in schema.parameters.items()}, - **{name: tuple(vdef.foreach) for name, vdef in schema.variables.items()}, - } diff --git a/src/math_spec/errors.py b/src/math_spec/errors.py index b19f529e..1c467ae9 100644 --- a/src/math_spec/errors.py +++ b/src/math_spec/errors.py @@ -2,24 +2,19 @@ # # SPDX-License-Identifier: MIT -"""What the language says back about a file: the errors it raises, and the advice it gives. - -:class:`LanguageError` is the file saying something the language does not -accept — decidable at load time, with no data bound. :class:`MathSpecError` -is the root a consumer's own errors may derive from, so one ``except`` covers -the package. :class:`Advice` is the other kind of sentence: about a file the -language accepts, decidable without data all the same. -""" +"""What the language says back about a file: the errors it raises, and the advice it gives.""" from __future__ import annotations import difflib from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, Literal, get_args +from typing import TYPE_CHECKING, Literal, get_args if TYPE_CHECKING: from collections.abc import Iterable + from pydantic import ValidationError + #: Which pass an :class:`Advice` comes from. Closed, like the operator set: a #: consumer filtering on it can enumerate every value. @@ -58,13 +53,7 @@ class LanguageError(MathSpecError): class SchemaError(LanguageError): - """**The declarations themselves are wrong**, before any expression is read. - - An unknown key, a bad ``dtype``, a duplicate YAML key, a - version this reader does not know — as against a bare - :class:`LanguageError`, which is sound declarations saying something the - language rejects (an undeclared name, a dim rule, degree 2). - """ + """What a load refuses: an unknown key, a bad ``dtype``, a duplicate YAML key, an unparseable or unresolvable expression.""" class DimensionError(LanguageError): @@ -76,11 +65,7 @@ class PiecewiseExpansionError(LanguageError): 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. - - Only the clause is shared — an unknown declaration, an unknown YAML key and - an unknown symbol-table entry each frame it with a sentence of their own. - """ + """The repair clause for an unrecognised name: the near miss, or the set.""" candidates = sorted(known) near = difflib.get_close_matches(name, candidates, n=1, cutoff=0.6) if near: @@ -88,18 +73,16 @@ def did_you_mean(name: str, known: Iterable[str], *, label: str = 'Declared') -> return f'{label}: {", ".join(candidates) or "nothing"}.' -def schema_error(exc: Any) -> LanguageError: - """A pydantic ``ValidationError`` as one of ours, keeping the class. +def schema_error(exc: ValidationError) -> LanguageError: + """A pydantic ``ValidationError`` as one of ours. - Pydantic wraps whatever a validator raises, so our own class cannot reach - the caller from inside the model — but the original survives under - ``ctx['error']``, so a :class:`DimensionError` comes back one. Anything - else, including several errors at once, is a :class:`SchemaError`. + Returns the original :class:`LanguageError` subclass where exactly one + error carries one, and a :class:`SchemaError` otherwise. """ errors = exc.errors() lines = [] for error in errors: - message = str(error.get('msg', '')).removeprefix('Value error, ') + message = error.get('msg', '').removeprefix('Value error, ') where = '.'.join(str(part) for part in error.get('loc', ())) lines.append(f'{where}: {message}' if where else message) text = '\n'.join(lines) or str(exc) diff --git a/src/math_spec/exclusivity.py b/src/math_spec/exclusivity.py index bbe1e5a1..ee0e1bde 100644 --- a/src/math_spec/exclusivity.py +++ b/src/math_spec/exclusivity.py @@ -4,39 +4,12 @@ """Can two of a named expression's cases claim one coordinate? Decided without data. -A named expression with ``cases:`` is one quantity whose value varies by -region — the regime a unit is in, which end of the horizon a row sits at. It is -*one* quantity only if no coordinate is claimed twice, and nothing about the -data decides that — so it is decided at load, and a file leaving two cases free -to collide does not load. - -The other half of being a quantity — a value *everywhere* — is the block's -shape rather than anything proved: the ``otherwise:`` beside the cases takes -whatever they leave. Only the ``when`` strings are checked, and only -against each other, pair by pair: ``when_i AND when_j`` unsatisfiable. - -Every atom in the where-grammar talks about exactly one **subject** — a -parameter, a dimension's coordinates, a dimension's *rank*, a lookup, a pair of -lookups. Atoms with different subjects are independent; atoms sharing one are -not, and that is where a propositional reading goes wrong: on ``kind == -'battery'`` and ``kind == 'h2'`` it invents a world where both hold and reports -an overlap that no data can produce. - -So each subject is split into **cells** — finitely many regions its value can -sit in, chosen so that every atom over that subject is constant on each cell. -The cells of the pair's subjects are multiplied out and both masks evaluated on -each. A cell where both are true is a witness. Because the cells cover every -value a subject can take, "no witness" is a proof and not a sample. - -Independence between subjects is an **over**-approximation: the product of -cells contains worlds the data may never produce, so a spurious world can only -manufacture a witness, never hide one. Every outcome here is therefore -conservative — this refuses case sets that would have been fine, and admits -none that would not. - -A pair the procedure will not reason about is refused exactly as an overlapping -one is, and the refusal names the rewrite: a checker that guesses where it -cannot decide buys nothing over no checker. +Each ``when`` is read over cells: regions of one subject's value on which every +atom is constant. The cells of the pair's subjects are multiplied out, both +masks are evaluated on each, and a cell where both hold is a witness. +Independence between subjects over-approximates, so it can manufacture a +witness but never hide one. The rule itself is stated in +``docs/reference/language/expressions.md``. """ from __future__ import annotations @@ -46,10 +19,9 @@ import math from dataclasses import dataclass from enum import Enum -from typing import TYPE_CHECKING, Any, Literal, cast +from typing import TYPE_CHECKING, Any, Literal, assert_never, cast -from math_spec.resolution import Namespace -from math_spec.where_parser import ( +from math_spec.program import ( AndNode, BooleanLiteralNode, DimensionComparisonNode, @@ -57,6 +29,7 @@ LookupComparisonNode, LookupDefinedNode, LookupPairComparisonNode, + Mask, NotNode, OrNode, ParameterComparisonNode, @@ -67,39 +40,35 @@ if TYPE_CHECKING: from collections.abc import Iterable, Iterator, Mapping - from math_spec.model import Spec - from math_spec.where_parser import PredicateOperator, WhereNode + from math_spec.model import DeclaredDtype + from math_spec.program import PredicateOperator, TypedPredicateNode, WhereNode -#: The product of the pair's subjects' cells is enumerated, so the bound is on -#: the product rather than on any one subject. Two real masks carry two to four -#: atoms between them; a pair that blows this is telling you it is several -#: expressions. +#: The most cells one pair may multiply out to; a pair past it is several expressions. CELL_BUDGET = 8192 #: The dtypes an ordering is decided against. Everything else compares only #: with == and !=, which need no order on the values. -_ORDERED_DTYPES = ('float', 'int', 'datetime') +_ORDERED_DTYPES: tuple[DeclaredDtype, ...] = ('float', 'int', 'datetime') class Undecidable(Exception): # noqa: N818 """A pair this procedure will not reason about. Carries the rewrite.""" -def overlapping(cases: Mapping[str, WhereNode], schema: Spec) -> Iterator[str]: +def overlapping(cases: Mapping[str, WhereNode], dtypes: Mapping[str, DeclaredDtype]) -> Iterator[str]: """One refusal per pair of cases that could both claim a coordinate. Args: cases: The ``when`` of every case, keyed by the case's name. The block's ``otherwise`` is not among them: it claims what the rest leave, so it overlaps nothing by construction. - schema: Read for the dtype of every name a mask compares against. + dtypes: The declared dtype of every name a mask compares against. Yields: A sentence per pair, naming both cases and either a coordinate they both claim or what stopped the pair being decided. Empty where every pair is proved apart. """ - dtypes = Namespace.of(schema).dtypes for (first, left), (second, right) in itertools.combinations(cases.items(), 2): try: witness = _witness(left, right, dtypes) @@ -119,9 +88,10 @@ def overlapping(cases: Mapping[str, WhereNode], schema: Spec) -> Iterator[str]: ) -def _witness(first: WhereNode, second: WhereNode, dtypes: Mapping[str, str]) -> str | None: +def _witness(first: WhereNode, second: WhereNode, dtypes: Mapping[str, DeclaredDtype]) -> str | None: """A coordinate both masks claim, rendered — ``None`` where no cell holds both.""" - frame = _Frame.of([first, second], dtypes) + masks = (Mask(first), Mask(second)) + frame = _Frame.of(masks, dtypes) if frame.size > CELL_BUDGET: msg = ( f'{frame.size} regions to check exceeds the budget of {CELL_BUDGET} — ' @@ -129,7 +99,7 @@ def _witness(first: WhereNode, second: WhereNode, dtypes: Mapping[str, str]) -> ) raise Undecidable(msg) for cell in frame.cells(): - if _evaluate(first, cell, frame) and _evaluate(second, cell, frame): + if all(_evaluate(mask.root, cell, frame) for mask in masks): return frame.witness(cell) return None @@ -142,15 +112,13 @@ def _witness(first: WhereNode, second: WhereNode, dtypes: Mapping[str, str]) -> class Special(Enum): """Values a cell can hold that are not values of the subject's own type.""" - #: No row in the table. A null compares false whatever the comparator, and - #: is not `defined`. + #: No row in the table: compares false under every comparator, and is not `defined`. NULL = 'null' - #: A magnitude, and the one that is a *value* everywhere else but is not - #: `defined` — see the bare-name row of the where-string table. + #: A magnitude every comparison reads normally, and the one that is not `defined`. POS_INF = '+inf' + #: Its negative twin. NEG_INF = '-inf' - #: A label none of the masks names. Stands for every such label at once, - #: which they cannot tell apart. + #: A label none of the masks names — every such label at once. OTHER = 'other' @@ -184,23 +152,20 @@ def __str__(self) -> str: class _Frame: """The cells to check, and what reading an atom on one of them needs. - ``subjects`` is keyed by ``id(node)`` because the where-AST nodes are - ``@dataclass`` with ``eq=True`` and so unhashable. It is a memo of a pure - function: without it every atom re-derives and re-allocates its subject once - per cell, which is the hot path here. + ``subjects`` is keyed by ``id(node)``: the where nodes are ``@dataclass`` + with ``eq=True`` and so unhashable. The memo is valid while the masks are alive. """ domains: dict[Subject, list[Cell]] subjects: dict[int, Subject] @classmethod - def of(cls, masks: Iterable[WhereNode], dtypes: Mapping[str, str]) -> _Frame: + def of(cls, masks: Iterable[Mask], dtypes: Mapping[str, DeclaredDtype]) -> _Frame: values: dict[Subject, set[Any]] = {} subjects: dict[int, Subject] = {} for mask in masks: - for node in _walk(mask): - if (subject := _subject_of(node)) is None: - continue + for node in mask.atoms: + subject = _subject_of(node) subjects[id(node)] = subject _observe(node, subject, values.setdefault(subject, set()), dtypes) return cls({s: _cells_for(s, seen, dtypes) for s, seen in values.items()}, subjects) @@ -217,22 +182,13 @@ def witness(self, cell: dict[Subject, Cell]) -> str: return ', '.join(f'{subject} is {_shown(subject, value)}' for subject, value in cell.items()) -def _walk(node: WhereNode) -> Iterator[WhereNode]: - """Every atom in *node*; the connectives are stepped through.""" - if isinstance(node, NotNode): - yield from _walk(node.operand) - elif isinstance(node, AndNode | OrNode): - yield from _walk(node.left) - yield from _walk(node.right) - else: - yield node +def _observe(node: TypedPredicateNode, subject: Subject, values: set[Any], dtypes: Mapping[str, DeclaredDtype]) -> None: + """Record what *node* says about its subject: a position, or a literal. - -def _observe(node: WhereNode, subject: Subject, values: set[Any], dtypes: Mapping[str, str]) -> None: - """Record what *node* says about its subject: a position, or a literal.""" + ``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, DimensionPositionNode): - # Every comparator reads here: `position()` converts the dimension to - # an integer, so an ordering is an ordering of integers (#32). values.add(node.position) elif isinstance(node, LookupPairComparisonNode): if node.op not in ('==', '!='): @@ -253,10 +209,8 @@ def _observe(node: WhereNode, subject: Subject, values: set[Any], dtypes: Mappin values.add(node.value) -def _subject_of(node: WhereNode) -> Subject | None: +def _subject_of(node: TypedPredicateNode) -> Subject: match node: - case BooleanLiteralNode(): - return None case ParameterDefinedNode(name=name) | ParameterComparisonNode(name=name): return Subject('param', name) case VariableDefinedNode(name=name): @@ -270,13 +224,10 @@ def _subject_of(node: WhereNode) -> Subject | None: case LookupPairComparisonNode(name=name, other=other): return Subject('lookup_pair', name, other) case _: - # As in `dimensions.py` and the typesetter: an unresolved node here - # is a caller that skipped `resolve_where`, not a model to refuse. - msg = f'{type(node).__name__} reached the exclusivity check unresolved.' - raise AssertionError(msg) + assert_never(node) -def _cells_for(subject: Subject, values: set[Any], dtypes: Mapping[str, str]) -> list[Cell]: +def _cells_for(subject: Subject, values: set[Any], dtypes: Mapping[str, DeclaredDtype]) -> list[Cell]: """Every region *subject*'s value can sit in — ordinary values first. The order is the order :func:`_witness` searches, so a refusal names an @@ -300,19 +251,23 @@ def _cells_for(subject: Subject, values: set[Any], dtypes: Mapping[str, str]) -> cells: list[Cell] = list( _ordered_cells(values, discrete=dated or dtype == 'int') if numeric or dated else _label_cells(values) ) - # A dimension's coordinates are its own index, so there is no null among - # them; everything else may be absent, and absence is a region of its own - # because a null compares false and is not `defined`. - if subject.kind != 'dim': - cells.append(Special.NULL) - if numeric: - # `defined` excludes an infinity, so it needs a region where every - # comparison still reads normally but the bare name is false. - cells.extend([Special.NEG_INF, Special.POS_INF]) + cells.extend(_absence_cells(subject, numeric=numeric)) return cells -def _numeric(dtype: str | None, literals: set[Any]) -> bool: +def _absence_cells(subject: Subject, *, numeric: bool) -> list[Cell]: + """The regions where *subject* has no ordinary value. + + A dimension's coordinates have no null. ``defined`` excludes an infinity, + so a magnitude needs a region where every comparison still reads normally + and the bare name is false. + """ + if subject.kind == 'dim': + return [] + return [Special.NULL, Special.NEG_INF, Special.POS_INF] if numeric else [Special.NULL] + + +def _numeric(dtype: DeclaredDtype | None, literals: set[Any]) -> bool: """Is this subject a magnitude? The declaration says so where it is known.""" if dtype is not None: return dtype in ('float', 'int') @@ -351,18 +306,14 @@ def _step(value: Any) -> Any: def _between(value: Any, following: Any, step: Any, *, discrete: bool) -> Any | None: - """A value strictly between two literals, where the type admits one. - - A continuous magnitude always admits one — the midpoint. A **discrete** - subject, an ``int`` or a date, need not: between 0 and 1 there is no - integer and between two adjacent days no date, so the gap has to be wider - than one unit before there is anything in it to stand for. A midpoint - invented there is a coordinate the subject cannot take, and the only thing - it can do is manufacture a witness — refusing ``n < 1`` against ``n > 0``, - which no integer claims twice, at a coordinate named ``0.5``. + """A value strictly between two literals, or ``None`` where the type admits none. + + A discrete subject — an ``int`` or a date — has one only where the gap is wider than one unit. """ if discrete: return value + step if following - value > step else None + # pyrefly: ignore[no-any-return-implicit] -- declaring `Any` would silence this and stop + # saying that the discrete branch has nothing to return. return (value + following) / 2.0 @@ -439,13 +390,11 @@ def _evaluate(node: WhereNode, cell: dict[Subject, Cell], frame: _Frame) -> bool return _atom(node, cell, frame) -def _atom(node: WhereNode, cell: dict[Subject, Cell], frame: _Frame) -> bool: +def _atom(node: TypedPredicateNode, cell: dict[Subject, Cell], frame: _Frame) -> bool: subject = frame.subjects[id(node)] value = cell[subject] match node: case ParameterDefinedNode() | LookupDefinedNode(): - # What `defined` means is the declaration's to say: a bool is its - # own answer, and a number has to be finite as well. if isinstance(value, bool): return value return value not in (Special.NULL, Special.POS_INF, Special.NEG_INF) @@ -456,41 +405,32 @@ def _atom(node: WhereNode, cell: dict[Subject, Cell], frame: _Frame) -> bool: case DimensionPositionNode(op=op, position=position): return _compare(value, op, position) case ParameterComparisonNode(op=op, value=literal) | LookupComparisonNode(op=op, value=literal): - # A null compares false, whatever the comparator. if value is Special.NULL: return False return _compare(value, op, literal) case DimensionComparisonNode(op=op, value=literal): return _compare(value, op, literal) case _: - msg = f'{type(node).__name__} reached the exclusivity check unresolved.' - raise AssertionError(msg) + assert_never(node) def _compare(value: Cell, op: PredicateOperator, literal: Any) -> bool: """One atom's truth in one cell. Both sides are already this cell's frame.""" if isinstance(value, Special): if value is Special.OTHER: - # A label none of the masks names sorts nowhere; `_observe` has - # already refused the ordering that would reach the second arm. + # a label none of the masks names sorts nowhere if op in ('==', '!='): return op == '!=' msg = f'a label neither case names is ordered with {op!r} — compare labels with == or != instead' raise Undecidable(msg) magnitude = math.inf if value is Special.POS_INF else -math.inf - return _numeric_compare(magnitude, op, float(literal)) + return _ordered(magnitude, op, float(literal)) if isinstance(value, int | float) and isinstance(literal, int | float) and not isinstance(value, bool): - return _numeric_compare(float(value), op, float(literal)) - if type(value) is not type(literal) and not isinstance(value, type(literal)): - # Unreachable while one declared dtype types every literal of a subject. - if op in ('==', '!='): - return op == '!=' - msg = f'{value!r} is ordered against {literal!r}, and the two carry no order — compare them with == or !=' - raise Undecidable(msg) - return _numeric_compare(value, op, literal) + return _ordered(float(value), op, float(literal)) + return _ordered(value, op, literal) -def _numeric_compare(left: Any, op: PredicateOperator, right: Any) -> bool: +def _ordered(left: Any, op: PredicateOperator, right: Any) -> bool: match op: case '==': return bool(left == right) diff --git a/src/math_spec/expansion.py b/src/math_spec/expansion.py index dca9f2e7..8d08a10a 100644 --- a/src/math_spec/expansion.py +++ b/src/math_spec/expansion.py @@ -2,62 +2,45 @@ # # SPDX-License-Identifier: MIT -"""Named sub-expressions and macros, both declared in the YAML and expanded into the AST before anything reads the expression. - -There is no Python operator registry: the built-in set is closed, macros cover -composition, and math the language cannot say goes in a declared ``escape:`` -island (#38). -""" +"""Named sub-expressions and macros, expanded into the core AST before anything reads the expression.""" from __future__ import annotations -from typing import TYPE_CHECKING, assert_never, overload +from typing import TYPE_CHECKING, overload +from math_spec._where_parser import parse_where from math_spec.errors import SchemaError from math_spec.expression_parser import ( ArithmeticNode, - BinaryOperatorNode, CaseArm, CasesNode, ComparisonNode, ExpressionNode, FunctionCallNode, - LeafNode, NameNode, - UnaryOperatorNode, parse_expression, + with_children, ) -from math_spec.where_parser import parse_where if TYPE_CHECKING: - from collections.abc import Callable - from math_spec.model import ExpressionBlock, MacroBlock, Spec -def parse_and_expand(text: str, schema: Spec, context: str = 'expression') -> ExpressionNode: +def parse_and_expand(text: str, schema: Spec, context: str) -> ExpressionNode: """Parse *text* and expand named sub-expressions and macros to core AST.""" return expand(parse_expression(text), schema, context) @overload -def expand( - node: ArithmeticNode, schema: Spec, context: str = ..., *, shadow: frozenset[str] = ... -) -> ArithmeticNode: ... +def expand(node: ArithmeticNode, schema: Spec, context: str, *, shadow: frozenset[str] = ...) -> ArithmeticNode: ... @overload -def expand( - node: ComparisonNode, schema: Spec, context: str = ..., *, shadow: frozenset[str] = ... -) -> ComparisonNode: ... +def expand(node: ComparisonNode, schema: Spec, context: str, *, shadow: frozenset[str] = ...) -> ComparisonNode: ... -def expand( - node: ExpressionNode, schema: Spec, context: str = 'expression', *, shadow: frozenset[str] = frozenset() -) -> ExpressionNode: +def expand(node: ExpressionNode, schema: Spec, context: str, *, shadow: frozenset[str] = frozenset()) -> ExpressionNode: """Expand all named sub-expressions and macro calls under *node*. - Expansion never changes the shape of the root: a comparison stays a - comparison, an arithmetic node stays arithmetic. The overloads say so, so - callers holding an ``ArithmeticNode`` keep it across the call. + A comparison stays a comparison and an arithmetic node stays arithmetic. Args: node: The parsed expression. @@ -83,36 +66,7 @@ def macro_signature(name: str, macro: MacroBlock) -> str: def parse_template(name: str, macro: MacroBlock, context: str) -> ArithmeticNode: """Parse a macro template, rejecting comparisons.""" - 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 _descend(node: ArithmeticNode, recurse: Callable[[ArithmeticNode], ArithmeticNode]) -> ArithmeticNode: - """Rebuild *node* with *recurse* applied to each child. - - The structural half of a tree walk, shared by the two walks below: they - differ only in what they do at NameNode and FunctionCallNode, and duplicating - the other four cases is how the two drift apart. - """ - if isinstance(node, LeafNode): - return node - if isinstance(node, UnaryOperatorNode): - return UnaryOperatorNode(node.op, recurse(node.operand)) - if isinstance(node, BinaryOperatorNode): - return BinaryOperatorNode(node.op, recurse(node.left), recurse(node.right)) - if isinstance(node, FunctionCallNode): - return FunctionCallNode( - node.name, - [recurse(a) for a in node.args], - {k: recurse(v) for k, v in node.kwargs.items()}, - ) - if isinstance(node, CasesNode): - # the values only: a `when` is a mask over the frame, checked where the cases are declared - return CasesNode(node.name, tuple(CaseArm(a.label, a.when, recurse(a.value)) for a in node.arms)) - assert_never(node) + return _parse_body(macro.template, f"macro '{name}' template", context) def _expand( @@ -137,7 +91,7 @@ def _cycle(name: str, kind: str) -> None: _cycle(node.name, 'macro') return _expand_macro(node, schema, context, stack, shadow) - return _descend(node, lambda child: _expand(child, 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: @@ -149,16 +103,11 @@ def _parse_named(name: str, schema: Spec, context: str) -> ArithmeticNode: def _parse_cased(name: str, block: ExpressionBlock, context: str) -> CasesNode: - """A cased expression, as the node that stands where its name was. - - The arms come out in file order, which is the order they print in, and - carry unresolved ``when`` masks — expansion runs before resolution, so - :mod:`math_spec.resolution` types those along with everything else. The - block's ``otherwise:`` becomes the last arm, the only one with no mask. - """ + """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) @@ -215,4 +164,4 @@ def _substitute(node: ArithmeticNode, bindings: dict[str, ArithmeticNode]) -> Ar """Replace formal-name NameNodes in *node* with their bound subtrees.""" if isinstance(node, NameNode) and node.name in bindings: return bindings[node.name] - return _descend(node, lambda child: _substitute(child, bindings)) + return with_children(node, lambda child: _substitute(child, bindings)) diff --git a/src/math_spec/expression_parser.py b/src/math_spec/expression_parser.py index 174246d3..f9a0b754 100644 --- a/src/math_spec/expression_parser.py +++ b/src/math_spec/expression_parser.py @@ -2,69 +2,79 @@ # # SPDX-License-Identifier: MIT -"""pyparsing-based expression parser for math expressions. +"""The core AST every pass reads, and the pyparsing grammar that builds it. -Parses strings like ``sum(p * cost, over=generator) == load`` into an AST -that can be evaluated against a namespace of linopy variables and xarray -parameters. - -``ArithmeticNode`` is the arithmetic-only union: every nested expression -position (operands, args, kwargs) accepts it and nothing else, and -``ComparisonNode`` appears only at the top of a parsed expression. +Arithmetic nests anywhere; a comparison appears only at the top of a parsed +expression. """ from __future__ import annotations from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any, Literal, cast +from functools import lru_cache +from types import MappingProxyType +from typing import TYPE_CHECKING, Any, Literal, assert_never, cast, get_args import pyparsing as pp from math_spec.errors import SchemaError if TYPE_CHECKING: - from math_spec.where_parser import WhereNode + from collections.abc import Callable, Mapping + + from math_spec.program import WhereNode +#: 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['<=', '>=', '=='] +#: The sign a unary operator applies to its operand. +UnaryOperator = Literal['+', '-'] + +#: The arithmetic a binary operator may spell. Closed by the grammar, and the +#: vocabulary a renderer dispatching on :attr:`BinaryOperatorNode.op` +#: switches over — it keeps no list of its own. +BinaryOperator = Literal['+', '-', '*', '/', '**'] + +#: What an expression writes to refer to a declaration, and so what a +#: declaration may be named. +NAME = r'[a-zA-Z_][a-zA-Z0-9_]*' + +#: A float — a fractional part or an exponent. A sign is the unary operator's. +REAL = r'\d+\.\d*([eE][+-]?\d+)?|\d+[eE][+-]?\d+' + # --------------------------------------------------------------------------- # AST nodes # --------------------------------------------------------------------------- -@dataclass +@dataclass(frozen=True) class NumberNode: value: float -@dataclass +@dataclass(frozen=True) class NameNode: - """An unresolved token — a name whose *kind* is not yet known. - - The parser cannot know whether ``p`` is a variable, a parameter or a - dimension; only the schema knows. ``resolution.py`` rewrites every one of - these into one of the typed nodes below, so a NameNode never reaches a - backend. If you find one there, resolution was skipped. - """ + """A bare name whose kind only the schema knows; resolution rewrites every one into a typed node.""" name: str -@dataclass +@dataclass(frozen=True) class VariableNode: """A resolved reference to a declared decision variable.""" name: str -@dataclass +@dataclass(frozen=True) class ParameterNode: """A resolved reference to a declared parameter.""" name: str -@dataclass +@dataclass(frozen=True) class DimensionNode: """A resolved reference to a declared dimension. @@ -75,7 +85,7 @@ class DimensionNode: name: str -@dataclass +@dataclass(frozen=True) class NameListNode: """A bracketed list of names in a kwarg value — ``sum(x, by=[a, b])``. @@ -90,7 +100,7 @@ def shown(self) -> str: return shown(self.names) -@dataclass +@dataclass(frozen=True) class LookupNode: """A resolved reference to one or more declared lookups, legal only in a kwarg value. @@ -109,7 +119,7 @@ def shown(self) -> str: return shown(self.names) -@dataclass +@dataclass(frozen=True) class KeywordNode: """A quoted closed keyword in a kwarg value — ``shift(..., edge='wrap')``. @@ -119,38 +129,41 @@ class KeywordNode: value: str -@dataclass +@dataclass(frozen=True) class EdgeNode: - """A resolved edge policy, legal only as an ``edge=`` value. - - A *number* in the same position stays a :class:`NumberNode`: the value the - vacated positions contribute. - """ + """The resolved ``edge='wrap'``; a number in the same position stays a :class:`NumberNode`.""" - policy: str - -@dataclass +@dataclass(frozen=True) class UnaryOperatorNode: - op: str + op: UnaryOperator operand: ArithmeticNode -@dataclass +@dataclass(frozen=True) class BinaryOperatorNode: - op: str + op: BinaryOperator left: ArithmeticNode right: ArithmeticNode -@dataclass +@dataclass(frozen=True) class FunctionCallNode: + """An operator or macro call. + + ``kwargs`` is held behind a read-only view and excluded from the hash; + equal nodes still hash equal on ``name`` and ``args``. + """ + name: str - args: list[ArithmeticNode] = field(default_factory=list) - kwargs: dict[str, ArithmeticNode] = field(default_factory=dict) + args: tuple[ArithmeticNode, ...] = () + kwargs: Mapping[str, ArithmeticNode] = field(default_factory=dict, hash=False) + + def __post_init__(self) -> None: + object.__setattr__(self, 'kwargs', MappingProxyType(dict(self.kwargs))) -@dataclass +@dataclass(frozen=True) class CaseArm: """One region of a :class:`CasesNode`: where it applies, and the value there. @@ -166,16 +179,11 @@ class CaseArm: def case_context(name: str, label: str | None) -> str: - """Where an error inside one arm of a cased expression is reported: the declaration, not the use site. - - A cased expression is expanded where its name stood, so the context in hand - at that point is the constraint's — and naming it would report a case on a - constraint that has none. + """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:``, - which is not a case and is not named as one. + label: The case's name, or ``None`` for the block's ``otherwise:``. Returns: The context prefix an error message carries. @@ -184,20 +192,14 @@ def case_context(name: str, label: str | None) -> str: return f"Named expression '{name}', {where}" -@dataclass +@dataclass(frozen=True) class CasesNode: - """A value defined by region — a named expression's ``cases:``, inlined. - - Built by :mod:`math_spec.expansion` where a reference to a cased expression - stood; there is no grammar for it, since a file writes the cases on the - declaration rather than at the use site. - - Exactly one arm applies at every coordinate: no two ``when`` masks can hold - at once, which :mod:`math_spec.exclusivity` proves at load, and the last arm - — the block's ``otherwise:`` — carries no ``when`` and so takes whatever the - rest leave. So the arms may be read in any order; the file's is kept because - it is the order they print in. The frame is not carried here: it is on the - declaration, which every consumer needing it already holds. + """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 @@ -221,7 +223,7 @@ class CasesNode: ) -@dataclass +@dataclass(frozen=True) class ComparisonNode: op: ComparisonOperator left: ArithmeticNode @@ -261,12 +263,10 @@ def children(node: ExpressionNode) -> 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. A pass whose - *answer* differs per node type dispatches itself and keeps its - ``assert_never``; this is for the ones that only need to get everywhere. - - An operator's kwargs are children too — a dimension or coordinate is an + 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 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,) @@ -275,32 +275,47 @@ def children(node: ExpressionNode) -> tuple[ArithmeticNode, ...]: if isinstance(node, FunctionCallNode): return (*node.args, *node.kwargs.values()) if isinstance(node, CasesNode): - # the values only: a `when` is a mask over the frame, not a value in it return tuple(arm.value for arm in node.arms) return () +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): + return node + if isinstance(node, UnaryOperatorNode): + return UnaryOperatorNode(node.op, recurse(node.operand)) + if isinstance(node, BinaryOperatorNode): + return BinaryOperatorNode(node.op, recurse(node.left), recurse(node.right)) + if isinstance(node, FunctionCallNode): + return FunctionCallNode( + node.name, + 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)) + assert_never(node) + + # --------------------------------------------------------------------------- # Grammar # --------------------------------------------------------------------------- def _build_grammar() -> pp.ParserElement: - """Build the pyparsing grammar for math expressions. - - ``inf`` is a ``pp.Keyword``, not a ``pp.Literal``: a ``Literal`` matches a - prefix, so it would eat the first three characters of ``inflow`` and leave - the parser meeting ``low`` where it expects the end of the expression. A - quoted value or a bracketed list of names is admitted only in a kwarg - value; a comparison appears at most once, and only at the top. - """ + """``inf`` is a ``pp.Keyword`` rather than a ``pp.Literal``, which would match the prefix of ``inflow``.""" arith = pp.Forward() inf_literal = (pp.Keyword('.inf') | pp.Keyword('inf')).set_parse_action(lambda: NumberNode(float('inf'))) # pyrefly: ignore[implicit-any-lambda] number = inf_literal | pp.Regex(rf'{REAL}|\d+').set_parse_action(lambda t: NumberNode(float(t[0]))) - name = pp.Regex(r'[a-zA-Z_][a-zA-Z0-9_]*') + name = pp.Regex(NAME) quoted = (pp.QuotedString("'") | pp.QuotedString('"')).set_parse_action(lambda t: KeywordNode(str(t[0]))) name_list = (pp.Suppress('[') + pp.DelimitedList(name) + pp.Suppress(']')).set_parse_action( @@ -328,18 +343,14 @@ def _build_grammar() -> pp.ParserElement: arith <<= add_sub - comparator = pp.one_of('<= >= ==') + comparator = pp.one_of(list(get_args(ComparisonOperator))) return (arith + pp.Optional(comparator + arith)).set_parse_action( lambda t: ComparisonNode(t[1], t[0], t[2]) if len(t) == 3 else t[0] ) def _make_func_call(tokens: pp.ParseResults) -> FunctionCallNode: - """Build a FunctionCallNode from parsed tokens. - - A ParseResults element is untyped, so the callee is cast; the grammar - guarantees an identifier in position 0. - """ + """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 = {} @@ -352,11 +363,10 @@ def _make_func_call(tokens: pp.ParseResults) -> FunctionCallNode: kwargs[k] = v else: args.append(item) - return FunctionCallNode(name=name, args=args, kwargs=kwargs) + return FunctionCallNode(name=name, args=tuple(args), kwargs=kwargs) def _make_left_assoc(tokens: pp.ParseResults) -> Any: - """Fold tokens into a left-associative BinaryOperatorNode chain.""" result, *rest = tokens for op, right in zip(rest[::2], rest[1::2], strict=True): result = BinaryOperatorNode(op, result, right) @@ -369,21 +379,60 @@ def _make_power(tokens: pp.ParseResults) -> Any: return items[0] if len(items) == 1 else BinaryOperatorNode('**', items[0], items[2]) -#: A float — a fractional part or an exponent. A sign is the unary operator's. -REAL = r'\d+\.\d*([eE][+-]?\d+)?|\d+[eE][+-]?\d+' - _GRAMMAR = _build_grammar() -def parse_expression(text: str) -> ExpressionNode: - """Parse a math expression string into an AST. +def parse_text(grammar: pp.ParserElement, text: str, what: str, rewrite: Callable[[str, int], str | None]) -> Any: + """Parse the whole of *text* with *grammar*, or raise :class:`SchemaError` naming *what* failed to parse. - Raises: - SchemaError: If *text* is not an expression of the language. + *rewrite* is asked for the predictable mistake at the failure position; its + sentence, if any, precedes the grammar's own complaint. """ try: - result = _GRAMMAR.parse_string(text, parse_all=True) + result = grammar.parse_string(text, parse_all=True) except pp.ParseException as e: - msg = f'Failed to parse expression: {text!r}\n{e}' + hint = rewrite(text, e.loc) + msg = f'Failed to parse {what}: {text!r}\n{f"{hint}\n" if hint is not None else ""}{e}' raise SchemaError(msg) from e - return cast('ExpressionNode', result[0]) + return result[0] + + +def _named_rewrite(text: str, loc: int) -> str | None: + """The rewrite for a predictable mistake at the token where the grammar gave up, or ``None``. + + A two-character token is tested before its one-character prefix. + """ + rest = text[loc:].lstrip() + if rest.startswith(get_args(ComparisonOperator)): + return ( + f"'{rest[:2]}' follows a complete comparison, and an expression carries " + f'one comparison, at the top. Split the chain into two constraints.' + ) + if rest.startswith('!='): + return ( + "'!=' is not a constraint sense — the senses are <=, >= and ==. " + 'Holding rows apart is a where matter: write the test in where:, where != is legal.' + ) + if rest.startswith(('<', '>')): + return f"'{rest[0]}' is not a constraint sense — the senses are <=, >= and ==. Write the bound inclusive." + if rest.startswith('='): + return ( + "'=' on its own is how a kwarg is written inside a call, like sum(x, over=d). " + 'Equality between two sides is written ==.' + ) + if rest.startswith('^'): + return "power is written '**', not '^'." + return None + + +@lru_cache(maxsize=4096) +def parse_expression(text: str) -> ExpressionNode: + """Parse a math expression string into an AST. + + Raises: + SchemaError: If *text* is not an expression of the language. A + predictable mistake — a strict or chained comparison, ``!=``, a + lone ``=``, ``^`` for power — is named with its rewrite before the + grammar's own complaint. + """ + return cast('ExpressionNode', parse_text(_GRAMMAR, text, 'expression', _named_rewrite)) diff --git a/src/math_spec/lowering.py b/src/math_spec/lowering.py index f51b3097..180c8efb 100644 --- a/src/math_spec/lowering.py +++ b/src/math_spec/lowering.py @@ -2,38 +2,21 @@ # # SPDX-License-Identifier: MIT -"""Lower a parsed YAML schema (typed AST) to a :class:`~math_spec.program.Program`. - -One lowering, whatever builds the result: it reads the typed AST and emits -declarations with names resolved and shapes fixed. It lives on the language -side, so no consumer needs YAML knowledge and this module reaches no consumer -— which is what makes two consumers agreeing about a file structural rather -than careful. - -Constructs with no lowering raise :class:`~math_spec.errors.LanguageError` naming -the construct and its rewrite, never a pointer at some other implementation: -a rejection here is a language gap (docs/about/roadmap.md) rather than a -routing decision. - -The rules a lowered program then carries: - -- a reduction over a dim the operand does not carry is an error, not a silent - identity — ``math_spec.dimensions`` owns that rule and this module asks it; -- a constraint is **one rule** carrying its own name, so a row is read back by - the name the file writes, with no positional suffix to guess; -- a file declares one objective, likewise one expression; -- an objective is scalar, so every reduction in it is one the file wrote and - nothing sums on its own behalf. +"""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. """ from __future__ import annotations from dataclasses import dataclass -from typing import TYPE_CHECKING, Literal, assert_never, cast +from typing import TYPE_CHECKING, assert_never import math_spec.program as program from math_spec.dimensions import dims_of -from math_spec.errors import LanguageError from math_spec.expression_parser import ( ArithmeticNode, BinaryOperatorNode, @@ -53,7 +36,6 @@ from math_spec.piecewise import declaration_of, derivations_of, expand_piecewise from math_spec.resolution import Namespace, expression_of, where_of from math_spec.validation import to_spec -from math_spec.where_parser import AndNode, NotNode, WhereNode if TYPE_CHECKING: from collections.abc import Callable @@ -62,32 +44,20 @@ from math_spec.model import Spec, _ExpandedSpec -_SENSES = {'==', '<=', '>='} - -def _none_of(masks: list[WhereNode]) -> WhereNode: +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 = _negated(masks[0]) + remainder = ~masks[0] for mask in masks[1:]: - remainder = AndNode(remainder, _negated(mask)) + remainder = remainder & ~mask return remainder -def _negated(mask: WhereNode) -> WhereNode: - """*mask* negated, cancelling a negation rather than stacking one. - - ``not (not committable)`` is a term every consumer would evaluate twice to - reach the answer it started from. The regions are built here, so this is - the one place that can spell them without it. - """ - return mask.operand if isinstance(mask, NotNode) else NotNode(mask) - - def to_program(spec: str | Path | dict[str, Any] | Spec | program.Program) -> program.Program: """*spec* as a :class:`~math_spec.program.Program` — the public door. @@ -95,10 +65,7 @@ def to_program(spec: str | Path | dict[str, Any] | Spec | program.Program) -> pr model, or a program already. Idempotent, so a caller that does not know which it holds can call this and be sure. - Not memoised. :func:`~math_spec.piecewise.expand_piecewise` is, because - validators reach for the expansion as well as consumers and the same model - is expanded more than once on one pass; nothing has that shape here. Add it - the day something lowers one model twice. + Not memoised; :func:`~math_spec.piecewise.expand_piecewise` is. Args: spec: What to read the declarations from. @@ -117,23 +84,15 @@ def to_program(spec: str | Path | dict[str, Any] | Spec | program.Program) -> pr return lower_program(expand_piecewise(to_spec(spec))) -def lower_program(schema: _ExpandedSpec) -> program.Program: - """Compile a :class:`_ExpandedSpec` into a :class:`Program`. - - Takes the expanded model rather than expanding one: a program is built from - declarations, and `_ExpandedSpec` is the type that guarantees they are all - there. Every caller already held one — the expansion is memoised on the - model — so this moves no work, it only stops the guarantee being a - convention four consumers happened to observe. +def lower_program(expanded: _ExpandedSpec) -> program.Program: + """Compile an expanded model into a :class:`~math_spec.program.Program`. - A ``domain: binary`` variable lowers with fixed 0/1 bounds, so the domain - needs no separate carrier. + A ``domain: binary`` variable lowers with fixed 0/1 bounds. Raises: - LanguageError: A construct outside the streaming language, named with - its rewrite. + LanguageError: A construct outside the language, named with its + rewrite. """ - expanded = schema ns = Namespace.of(expanded) derivations = { name: how @@ -147,7 +106,7 @@ def lower_program(schema: _ExpandedSpec) -> program.Program: variables = {} for vname, vdef in expanded.variables.items(): - variable_type = cast('program.VariableType', vdef.domain) + variable_type = vdef.domain if variable_type == 'binary': lower, upper = program.Constant(0.0), program.Constant(1.0) else: @@ -158,20 +117,14 @@ def lower_program(schema: _ExpandedSpec) -> program.Program: lower=lower, upper=upper, variable_type=variable_type, - absence=cast('program.VariableAbsence', vdef.absence), + absence=vdef.absence, ) constraints = {} for cname, cdef in expanded.constraints.items(): where = where_of(cdef.where, ns, f"constraint '{cname}'") ast = expression_of(cdef.expression, expanded, ns, f"constraint '{cname}'") - if not isinstance(ast, ComparisonNode): - raise LanguageError( - f"constraint '{cname}': expression must contain exactly one " - f'comparison operator (<=, >=, ==). Got: {cdef.expression!r}' - ) - if ast.op not in _SENSES: - raise LanguageError(f"constraint '{cname}': unsupported sense '{ast.op}'") + assert isinstance(ast, ComparisonNode), 'load-time validation refuses a constraint without a comparison' lowering = _Lowering(expanded, f"constraint '{cname}'") constraints[cname] = program.ConstraintDeclaration( tuple(cdef.foreach), @@ -184,8 +137,7 @@ def lower_program(schema: _ExpandedSpec) -> program.Program: objective = None if (odef := expanded.objective) is not None: ast = expression_of(odef.expression, expanded, ns, 'the objective') - if isinstance(ast, ComparisonNode): - raise LanguageError('the objective: expression must not contain a comparison operator') + assert not isinstance(ast, ComparisonNode), 'load-time validation refuses a comparison in the objective' objective = program.ObjectiveDeclaration( odef.sense, _Lowering(expanded, 'the objective').expr(ast), @@ -206,7 +158,7 @@ def lower_program(schema: _ExpandedSpec) -> program.Program: sname: program.SosDeclaration( sdef.variable, sdef.over, - sos_type=cast('Literal[1, 2]', sdef.type), + sos_type=sdef.type, big_m=sdef.big_m, ) for sname, sdef in expanded.sos.items() @@ -228,7 +180,7 @@ def _lower_expression(schema: _ExpandedSpec, ns: Namespace, name: str) -> progra """Compile the named expression *name* into a program expression. Raises: - LanguageError: A construct outside the streaming language. + LanguageError: A construct outside the language, named with its rewrite. """ context = f"named expression '{name}'" ast = expression_of(name, schema, ns, context) @@ -243,26 +195,13 @@ def _lower_expression(schema: _ExpandedSpec, ns: Namespace, name: str) -> progra @dataclass(frozen=True) class _Lowering: - """One expression walk, and the two things every step of it reads. - - ``schema`` and ``context`` are fixed for a whole walk, so they are its - state rather than two arguments every recursion repeats. Extend this - rather than adding a parameter to :meth:`expr` and every operator seam. - """ + """One expression walk, and the two things every step of it reads.""" schema: _ExpandedSpec context: str def expr(self, node: ArithmeticNode) -> program.ExpressionNode: - """Rewrite one resolved core-AST expression as a program expression. - - Nothing is judged here: the expression passed every language rule at - load, and this walk only decides which node a call becomes and the - shapes a node cannot represent — a ``GroupSum`` groups by a declared - lookup, a ``Translate`` distance is an integer literal. ``Sum`` and - ``GroupSum`` stay two nodes under one surface verb, reducing a dim away - and reducing it into another being different relational shapes. - """ + """Rewrite one resolved core-AST expression as a program expression.""" if isinstance(node, NumberNode): return program.Constant(node.value) @@ -298,11 +237,7 @@ def expr(self, node: ArithmeticNode) -> program.ExpressionNode: raise AssertionError(f'{self.context}: operator {node.op!r} reached lowering') if isinstance(node, FunctionCallNode): - try: - lower_call = _CALLS[node.name] - except KeyError: - raise LanguageError(f"{self.context}: built-in '{node.name}' declares no lowering case") from None - return lower_call(self, node) + return _CALLS[node.name](self, node) if isinstance(node, CasesNode): return self._cases(node) @@ -317,11 +252,14 @@ def _cases(self, node: CasesNode) -> program.Cases: 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 = [arm.when for arm in node.arms if arm.when is not None] + stated = [program.Mask(arm.when) for arm in node.arms if arm.when is not None] regions = [] for arm in node.arms: - when = arm.when if arm.when is not None else _none_of(stated) + when = program.Mask(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)) @@ -338,18 +276,15 @@ def sum(self, node: FunctionCallNode) -> program.ExpressionNode: return program.Sum(operand, tuple(sorted(dims_of(node.args[0], self.schema, self.context)))) if by_node is None: over_node = node.kwargs['over'] - if not isinstance(over_node, DimensionNode): - raise LanguageError(f'{self.context}: sum(over=...) must name a dimension') + assert isinstance(over_node, DimensionNode), 'resolution refuses an over= that is not a dimension' return program.Sum(operand, (over_node.name,)) - if not isinstance(by_node, LookupNode): - raise LanguageError(f'{self.context}: sum(by=...) must name a lookup') + assert isinstance(by_node, LookupNode), 'resolution refuses a by= that is not a lookup' return program.GroupSum(operand, over=by_node.dimension, coordinate=by_node.names, into=by_node.into) def at(self, node: FunctionCallNode) -> program.ExpressionNode: """``at(x, by=lookup)`` — the adjoint of :meth:`sum`'s ``by=`` form.""" by_node = node.kwargs['by'] - if not isinstance(by_node, LookupNode): - raise LanguageError(f'{self.context}: at(by=...) must name a lookup') + assert isinstance(by_node, LookupNode), 'resolution refuses a by= that is not a lookup' return program.At( self.expr(node.args[0]), over=by_node.dimension, @@ -369,8 +304,7 @@ def sum_back(self, node: FunctionCallNode) -> program.ExpressionNode: held it to one lookup over the walked dimension. """ over_node = node.kwargs['over'] - if not isinstance(over_node, DimensionNode): - raise LanguageError(f'{self.context}: sum_back(over=...) must name a dimension') + assert isinstance(over_node, DimensionNode), 'resolution refuses an over= that is not a dimension' within_node = node.kwargs['within'] operand = self.expr(node.args[0]) wrap = isinstance(node.kwargs.get('edge'), EdgeNode) @@ -389,8 +323,7 @@ def shift(self, node: FunctionCallNode) -> program.ExpressionNode: language has already held it to the keyword or a number. """ over_node = node.kwargs['over'] - if not isinstance(over_node, DimensionNode): - raise LanguageError(f'{self.context}: shift(over=...) must name a dimension') + assert isinstance(over_node, DimensionNode), 'resolution refuses an over= that is not a dimension' by_node = node.kwargs['offset'] operand = self.expr(node.args[0]) edge = node.kwargs.get('edge') @@ -410,11 +343,7 @@ def shift(self, node: FunctionCallNode) -> program.ExpressionNode: ) -#: One lowering per name in the language's ``BUILTIN_NAMES``. A table rather -#: than a chain of ``if``s because the set is *closed* — nothing registers into -#: it, and a name the language declares with no entry here is refused by name -#: rather than by ``KeyError``. Each method is named for the operator it -#: lowers, so the table reads as the identity it nearly is. +#: One lowering per name in the language's ``BUILTIN_NAMES``. _CALLS: dict[str, Callable[[_Lowering, FunctionCallNode], program.ExpressionNode]] = { 'sum': _Lowering.sum, 'at': _Lowering.at, diff --git a/src/math_spec/model.py b/src/math_spec/model.py index 5f10d978..624c1393 100644 --- a/src/math_spec/model.py +++ b/src/math_spec/model.py @@ -4,17 +4,15 @@ """The YAML surface's types — every block a file may contain, rooted at :class:`Spec`. -A block per declaration kind, and one strict base: an unrecognised key is an -error naming the near miss rather than a shrug, because a dropped ``bounds:`` -leaves a variable unbounded and says nothing. - Nothing here has seen data. """ from __future__ import annotations import math -from typing import TYPE_CHECKING, Annotated, Any, ClassVar, Literal, Self, get_args, override +import re +from collections import Counter +from typing import TYPE_CHECKING, Annotated, Any, ClassVar, Literal, Self, cast, get_args, override from pydantic import ( BaseModel, @@ -31,10 +29,11 @@ ) from math_spec.errors import did_you_mean, schema_error +from math_spec.expression_parser import NAME, ComparisonOperator from math_spec.operators import BUILTIN_NAMES if TYPE_CHECKING: - from collections.abc import Iterable + from collections.abc import Iterable, Iterator from pydantic import GetJsonSchemaHandler from pydantic.json_schema import JsonSchemaValue @@ -52,7 +51,7 @@ class _StrictBlock(BaseModel): model_config = ConfigDict(extra='forbid') #: What this model is called in a YAML file, for the error message. - _label: ClassVar[str] = '' + _label: ClassVar[str] @model_validator(mode='before') @classmethod @@ -67,10 +66,9 @@ def _reject_unknown_keys(cls, data: Any) -> Any: known = set(cls.model_fields) unknown = [k for k in data if isinstance(k, str) and k not in known] if unknown: - label = cls._label or cls.__name__ raise ValueError( '\n'.join( - f"unknown key '{k}' in {label}. {did_you_mean(k, known, label='Valid keys')}" for k in unknown + f"unknown key '{k}' in {cls._label}. {did_you_mean(k, known, label='Valid keys')}" for k in unknown ) ) return data @@ -88,6 +86,11 @@ def _reject_unknown_keys(cls, data: Any) -> Any: #: indexes by. ParameterDtype = Literal['float', 'int', 'bool', 'str'] +#: What a *name* a where comparison tests may be — a parameter's dtype or a +#: dimension's, since a lookup's is its target's. The union rather than either +#: half, because a mask names all three kinds and reads the dtype the same way. +DeclaredDtype = ParameterDtype | DimensionDtype + #: The domain a variable may declare. VariableDomain = Literal['continuous', 'integer', 'binary'] @@ -101,7 +104,7 @@ def _reject_unknown_keys(cls, data: Any) -> Any: ObjectiveSense = Literal['minimize', 'maximize'] #: The relation a link may pin its expression to the curve with. -LinkSign = Literal['==', '<=', '>='] +LinkSign = ComparisonOperator #: The order of special ordered set. SosType = Literal[1, 2] @@ -122,7 +125,7 @@ def _reject_unknown_keys(cls, data: Any) -> Any: PARAMETER_DTYPES = frozenset(get_args(ParameterDtype)) #: The parameter dtypes that stand where a number belongs — a coefficient, a #: term, a divisor, a bound. A label selects and a flag masks; neither is one. -NUMERIC_DTYPES = frozenset({'float', 'int'}) +NUMERIC_DTYPES: frozenset[ParameterDtype] = frozenset({'float', 'int'}) VARIABLE_DOMAINS = frozenset(get_args(VariableDomain)) VARIABLE_ABSENCE = frozenset(get_args(VariableAbsence)) CURVATURES = frozenset(get_args(Curvature)) @@ -208,9 +211,8 @@ class ParameterBlock(_StrictBlock): class BoundsBlock(_StrictBlock): """Variable bounds — each side is a number or parameter name. - linopy's defaults (``add_variables(lower=-inf, upper=inf)``): omitting a - bound leaves the variable unbounded on that side, not implicitly - non-negative. Non-negativity is a real constraint, so the file says it. + An omitted bound leaves the variable unbounded on that side, not + implicitly non-negative. """ _label: ClassVar[str] = 'a bounds block' @@ -255,13 +257,7 @@ class VariableBlock(_StrictBlock): @model_validator(mode='after') def _absence_needs_a_mask(self) -> VariableBlock: - """``absence:`` says what a *missing* coordinate means, so one must be missable. - - A variable's only source of absence is its own ``where:`` — ``foreach`` - is a product of declared dimensions and has every coordinate. Without a - mask the key selects between two readings of a case that cannot arise, - which is a setting the reader has to interpret and nothing can reach. - """ + """``absence:`` says what a *missing* coordinate means, so one must be missable.""" if self.absence != 'undefined' and self.where is None: msg = ( f'absence: {self.absence} needs a `where:` — a variable with no mask exists at every ' @@ -361,34 +357,19 @@ class ExpressionBlock(_StrictBlock): expression: sum(p * rate, over=generator) description: CO2 released, the quantity the cap bounds - A quantity whose value varies by **region** is written as ``cases:`` - instead — one case per region over a declared ``foreach:``, no two of them - claiming one coordinate, and an ``otherwise:`` for the rest:: - - previous_status: - foreach: [snapshot, generator] - cases: - always_on: { when: "not committable", expression: 1 } - boundary: { when: "committable and position(snapshot) == 0", expression: status_initial } - otherwise: shift(status, over=snapshot, offset=1) - - So the constraint that needs it names it, rather than being forked into one - copy per regime. + A quantity whose value varies by region is written as ``cases:`` over a + declared ``foreach:``, with an ``otherwise:`` for the rest — see the + language reference. """ _label: ClassVar[str] = 'a named expression' expression: Expression | None = None - #: The frame the cases are read over — required with them and refused - #: without, since no one case's body gives a cased expression its shape. + #: The frame the cases are read over — required with them, refused without. foreach: list[str] | None = None - #: The regions this quantity is defined by, keyed by the name labelling the - #: row it prints. Each ``when`` is proved apart from every other, so the - #: order is the page's rather than the meaning's. + #: The regions, keyed by the name labelling the row each prints; every ``when`` is proved apart from the others. cases: Annotated[dict[str, ExpressionCase], Field(min_length=1)] = {} - #: The value wherever no case's ``when`` holds, which is what makes the - #: quantity whole. Written as the bare value — it has nothing else to - #: carry — and printed as the last row, the one that reads "otherwise". + #: The value wherever no case's ``when`` holds, printed as the last row. otherwise: Expression | None = None description: str | None = None @@ -399,13 +380,7 @@ def _from_string(cls, data: Any) -> Any: @model_validator(mode='after') def _one_form_or_the_other(self) -> Self: - """One ``expression:``, or ``cases:`` with the ``otherwise:`` and ``foreach:`` they need. - - Each near-miss gets its own sentence, being a different mistake: both - forms is not knowing which wins, neither is an empty declaration, a - ``foreach:`` alone is a second answer to what the body already answers, - and ``cases:`` without ``otherwise:`` is a quantity with a hole in it. - """ + """One ``expression:``, or ``cases:`` with the ``otherwise:`` and ``foreach:`` they need.""" if bool(self.cases) == (self.expression is not None): got = 'both' if self.cases else 'neither' msg = ( @@ -503,11 +478,6 @@ def _as_list(self) -> list[str]: #: its words too, and mean the same things. ``adjacency`` and ``convex`` are #: ours, linopy having no name for the first and reaching the second only as a #: fallback. -#: -#: ``lp`` is the one that emits no weights: it states the curve as its segment -#: lines directly, which is exact for a convex or concave curve bounded on the -#: matching side, and needs the domain rows because a line does not stop at a -#: breakpoint. PIECEWISE_METHODS = { 'adjacency': 'a binary per segment, and a row making the two nonzero weights neighbours', 'sos2': 'the same weights, restricted by a set the solver branches on (the sos rules)', @@ -525,22 +495,18 @@ class PiecewiseBlock(_StrictBlock): names a parameter carrying the ``over`` dim, and *sign* bounds the link by the curve instead of pinning it (at most one non-``"=="``, and only with exactly two links). - - ``over`` names the breakpoint dimension; ``method`` is which of - :data:`PIECEWISE_METHODS` restricts the weights; ``activity`` names what the weights sum - to — 1 where the block is unconditional, and a binary where a curve applies - only when something runs, which pins the formulation to 0 when it is 0; ``points`` names a - boolean parameter saying how far each curve runs, for a model whose curves - are not all the same length. Expanded before building into plain variables - and constraints — see ``math_spec.piecewise``. """ _label: ClassVar[str] = 'a piecewise declaration' + #: The breakpoint dimension. over: str links: list[PiecewiseLink] + #: Which of :data:`PIECEWISE_METHODS` restricts the weights. method: PiecewiseMethod = 'adjacency' + #: What the weights sum to — 1 where absent, or a binary that pins the formulation to 0 when it is 0. activity: str | None = None + #: A boolean parameter saying how far each curve runs, for curves of unequal length. points: str | None = None description: str | None = None @@ -548,9 +514,7 @@ class PiecewiseBlock(_StrictBlock): def curve(self) -> tuple[PiecewiseLink, PiecewiseLink]: """The two links as ``(x, y)``, the bounded one last. - Only a two-link block has a curve to speak of, and only a bounded link - can be the wrong way round in ``links:`` — so this is what reads the - pair anywhere the ``y`` side is the one being stated. + Two-link blocks only. """ x, y = self.links return (y, x) if x.sign != '==' else (x, y) @@ -559,7 +523,7 @@ def curve(self) -> tuple[PiecewiseLink, PiecewiseLink]: @classmethod def _check_method(cls, v: Any, handler: ValidatorFunctionWrapHandler) -> PiecewiseMethod: try: - return handler(v) + return cast('PiecewiseMethod', handler(v)) except ValidationError: options = '\n'.join(f' {name}: {what}' for name, what in PIECEWISE_METHODS.items()) msg = f'unknown piecewise method {v!r}. The formulations are:\n{options}' @@ -627,14 +591,17 @@ class SosBlock(_StrictBlock): big_m: float | None = None description: str | None = None - @field_validator('type', mode='before') + @field_validator('type', mode='wrap') @classmethod - def _check_type(cls, v: Any) -> Any: - if type(v) is not int or v not in SOS_TYPES: - orders = ' or '.join(str(t) for t in sorted(SOS_TYPES)) - msg = f'sos type must be {orders}, got {v!r}. A set of any other order is not a construct solvers carry.' + def _check_type(cls, v: Any, handler: ValidatorFunctionWrapHandler) -> SosType: + orders = ' or '.join(str(t) for t in sorted(SOS_TYPES)) + msg = f'sos type must be {orders}, got {v!r}. A set of any other order is not a construct solvers carry.' + if type(v) is not int: # True == 1 == 1.0, and a set of order True is nothing raise ValueError(msg) - return v + try: + return cast('SosType', handler(v)) + except ValidationError: + raise ValueError(msg) from None @field_validator('big_m') @classmethod @@ -658,17 +625,6 @@ def undeclared_dimension(kind: str, name: str, dimension: str) -> str: return f"{kind} '{name}' references undeclared dimension '{dimension}'. Declare it under 'dimensions:'." -def _repeated(items: Iterable[Any]) -> list[Any]: - """Each value that appears more than once, once, in first-repeat order — for labels that may be unhashable.""" - seen: list[Any] = [] - repeats: list[Any] = [] - for item in items: - if item in seen and item not in repeats: - repeats.append(item) - seen.append(item) - return repeats - - def _without_absence(value: Any) -> Any: """*value* with every absent entry stripped, recursively — see :meth:`Spec._drop_absence`.""" if not isinstance(value, dict): @@ -691,17 +647,16 @@ def _is_absent(value: Any) -> bool: class Spec(_StrictBlock): """The declared math — one YAML file, or one dict, validated. Nothing here has seen data. - The API is the ten declaration sections plus ``version`` and - ``description``, and two ways back out: :meth:`to_dict` for the model as - data, :meth:`to_yaml` for the file a reviewer reads. In goes through - ``to_spec``, which raises + A ``Spec`` that exists has passed the whole language: constructing one by + any route — ``to_spec``, :meth:`model_validate`, the constructor — runs + every load-time check, expansion and expression pass included, and raises :class:`~math_spec.errors.LanguageError` on a model the language refuses. + Holding one is the proof, so nothing downstream checks it again. - Everything else on this class is pydantic's, not a contract this package - keeps — ``model_json_schema()`` describes the shape pydantic validates - rather than the language (checked in for editors as - ``schema/math_spec.schema.json``), and ``model_construct()`` skips validation - entirely, so a ``Spec`` is valid when it was built the normal way. + The API is the ten declaration sections plus ``version`` and + ``description``, and two ways back out: :meth:`to_dict` for the model as + data, :meth:`to_yaml` for the file a reviewer reads. Everything else on + this class is pydantic's, not a contract this package keeps. """ _label: ClassVar[str] = 'the top level of the file' @@ -709,7 +664,7 @@ class Spec(_StrictBlock): #: The :class:`_ExpandedSpec` built from this model. Owned entirely — written #: and read — by :func:`~math_spec.piecewise.expand_piecewise`; only #: the slot lives here. - _expansion: Any = PrivateAttr(default=None) + _expansion: _ExpandedSpec | None = PrivateAttr(default=None) #: Which language surface this file is written against. Absent means 0, so #: the field is additive. **0 means unstable** — the surface may change in @@ -774,38 +729,56 @@ def _drop_absence(self, handler: Any) -> dict[str, Any]: []`` is a scalar). On the serializer so that ``model_dump``, :meth:`to_dict` and :meth:`to_yaml` agree. """ - return _without_absence(handler(self)) + return cast('dict[str, Any]', _without_absence(handler(self))) def to_dict(self) -> dict[str, Any]: """The model as plain data. ``to_spec(m.to_dict())`` reproduces it.""" return self.model_dump() def to_yaml(self) -> str: - """The file a reviewer reads — including for a model that never had one. - - Generated rather than authored, so length costs a reader nothing and - being unambiguous saves them knowing this package's defaults at all. - """ + """The file a reviewer reads — including for a model that never had one.""" import yaml return yaml.safe_dump(self.to_dict(), sort_keys=False, allow_unicode=True) @model_validator(mode='after') - def _validate_references(self) -> Spec: - """Every cross-declaration rule the schema can decide without data. - - Names share one flat namespace, shadowing being how a new declaration - would silently change what an existing expression means. Every lookup - joins it — a lookup named after a dimension, its own target included, - is a collision, so each map carries a name of its own. + def _names_are_names(self) -> Spec: + """Every declaration is keyed by something an expression could write. - A lookup's target must be a declared dimension other than the one it - is over — grouping a dim into itself is a no-op that reads as a - reduction. Bounds look like the expression language but are not it, so - their error says what they actually accept. + Read off the model's own mappings rather than a list of sections, so a + section added later cannot be forgotten here — every mapping a Spec + carries is keyed by a declaration name. """ - errors = [] + errors = [ + f'{section}: {name!r} is not a name. A declaration is named the way an expression ' + f'writes it — a letter or an underscore, then letters, digits or underscores — so ' + f'nothing can refer to this one. Rename it.' + for section, value in self + if isinstance(value, dict) + for name in value + if not re.fullmatch(NAME, name) + ] + if errors: + raise ValueError('\n'.join(errors)) + return self + + @model_validator(mode='after') + def _validate_references(self) -> Spec: + """Every cross-declaration rule the schema can decide without data, collected rather than raised on the first.""" + errors = [ + *self._name_collisions(), + *self._frame_dimensions(), + *self._lookup_targets(), + *self._bound_names(), + *self._sos_shapes(), + ] + if errors: + raise ValueError('\n'.join(errors)) + return self + + def _name_collisions(self) -> Iterator[str]: + """A name is declared once, and never as a built-in operator.""" kinds: list[tuple[str, Iterable[str]]] = [ ('dimension', self.dimensions), ('lookup', self.lookups), @@ -818,19 +791,21 @@ def _validate_references(self) -> Spec: for kind, group in kinds: for name in group: if name in BUILTIN_NAMES: - errors.append( + yield ( f"{kind.capitalize()} '{name}' collides with the built-in operator " f"'{name}'. The operator set is closed and its names are reserved; " f'rename the {kind}.' ) if name in seen: - errors.append( + yield ( f"{kind.capitalize()} '{name}' collides with the {seen[name]} of " f'the same name. Names share one flat namespace — rename one of them.' ) else: seen[name] = kind + def _frame_dimensions(self) -> Iterator[str]: + """Every frame is a product of distinct, declared dimensions.""" frames = [ *(('Parameter', name, p.dims) for name, p in self.parameters.items()), *(('Variable', name, v.foreach) for name, v in self.variables.items()), @@ -838,27 +813,30 @@ def _validate_references(self) -> Spec: *(('Named expression', name, e.foreach or []) for name, e in self.expressions.items()), ] for kind, name, dims in frames: - errors.extend(undeclared_dimension(kind, name, d) for d in dims if d not in self.dimensions) - errors.extend( + yield from (undeclared_dimension(kind, name, d) for d in dims if d not in self.dimensions) + yield from ( f"{kind} '{name}' names dimension '{d}' twice. A frame is a product of distinct dimensions." - for d in _repeated(dims) + for d, count in Counter(dims).items() + if count > 1 ) + def _lookup_targets(self) -> Iterator[str]: + """A lookup is over a declared dimension and maps into a different declared one.""" for lname, lk in self.lookups.items(): if lk.over not in self.dimensions: - errors.append(undeclared_dimension('Lookup', lname, lk.over)) + yield (undeclared_dimension('Lookup', lname, lk.over)) if lk.into is not None: if lk.into not in self.dimensions: - errors.append( + yield ( f"Lookup '{lname}' targets undeclared dimension '{lk.into}'. " f"Declare it under 'dimensions:' — the target is what the " f'lookup values are checked against.' ) elif lk.into == lk.over: - errors.append( - f"Lookup '{lname}' maps '{lk.over}' into itself. A lookup maps into a different dimension." - ) + yield (f"Lookup '{lname}' maps '{lk.over}' into itself. A lookup maps into a different dimension.") + def _bound_names(self) -> Iterator[str]: + """A named bound is a numeric parameter.""" for vname, vdef in self.variables.items(): for side in ('lower', 'upper'): val = getattr(vdef.bounds, side) @@ -867,7 +845,7 @@ def _validate_references(self) -> Spec: if val in self.parameters: dtype = self.parameters[val].dtype if dtype not in NUMERIC_DTYPES: - errors.append( + yield ( f"Variable '{vname}' bounds.{side}: '{val}' is a {dtype} parameter, and a bound " f'is a number. Declare it dtype: float or int, or bound the variable by another.' ) @@ -878,27 +856,29 @@ def _validate_references(self) -> Spec: else f'bounds accept a parameter name or a number, not an expression (got {val!r}). ' f'Precompute it as a parameter' ) - errors.append(f"Variable '{vname}' bounds.{side}: {detail}.") + yield (f"Variable '{vname}' bounds.{side}: {detail}.") + def _sos_shapes(self) -> Iterator[str]: + """A set runs along one dim of one declared variable, and a variable carries one set.""" claimed: dict[str, str] = {} for sname, block in self.sos.items(): context = f"Sos '{sname}'" if block.over not in self.dimensions: - errors.append(undeclared_dimension('Sos', sname, block.over)) + yield (undeclared_dimension('Sos', sname, block.over)) elif block.variable not in self.variables: - errors.append( + yield ( f"{context}: '{block.variable}' is not a declared variable.\n" f' Variables: {sorted(self.variables)}\n' f'A set is over one variable, so a parameter or an expression cannot carry one.' ) elif block.over not in self.variables[block.variable].foreach: - errors.append( + yield ( f"{context}: over '{block.over}' is not a dim of variable " f"'{block.variable}' (foreach {self.variables[block.variable].foreach}). The set runs " f"along one of the variable's own dims — one set per coordinate of the rest." ) elif block.variable in claimed: - errors.append( + yield ( f"{context}: variable '{block.variable}' already carries the set declared by " f"'{claimed[block.variable]}'. A variable holds one set — declare a second " f'variable, or state the other restriction as a constraint.' @@ -906,19 +886,11 @@ def _validate_references(self) -> Spec: else: claimed[block.variable] = sname - if errors: - raise ValueError('\n'.join(errors)) - - return self - @model_validator(mode='after') def _validate_expressions(self) -> Spec: """Every expression and where string — after expansion, whose emitted declarations are language too. - The :class:`_ExpandedSpec` expansion builds runs every validator of its own, - so a file with ``piecewise:`` has its references checked twice and its - expressions once, there. The checkers import this module, so the - imports are local. + The checkers import this module, so the imports are local. """ from math_spec.piecewise import expand_piecewise from math_spec.validation import validate_expressions @@ -938,6 +910,8 @@ class ExpandedPiecewise(_StrictBlock): edge flags an ``lp`` block under a mask sits its domain rows on. """ + _label: ClassVar[str] = 'an expanded piecewise block' + block: PiecewiseBlock points: str | None = None starts: str | None = None diff --git a/src/math_spec/operators.py b/src/math_spec/operators.py index 84bd73ae..cb17b910 100644 --- a/src/math_spec/operators.py +++ b/src/math_spec/operators.py @@ -4,21 +4,14 @@ """The closed set of built-in operators and their call shapes. -Closed: there is no Python registry, so every consumer accepts exactly the -same language. Compositions belong in ``macros:``; math the language cannot say belongs in a -declared ``escape:`` island (#38), not in an operator that reads like a built-in. - -The *language* side of an operator — its name and signature, nothing else. The -signature lives here because more than one pass needs it (resolution types -the dimension arguments, validation name-checks macro bodies, a consumer builds -the call), and an arity spelled out once per pass is one the passes can -disagree about. Dependency-free on purpose: counts and keyword names, no AST. +One home for each signature: a composition is a macro, and math the language +cannot say is a declared ``escape:``. """ from __future__ import annotations from dataclasses import dataclass -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Literal if TYPE_CHECKING: from collections.abc import Iterable @@ -36,12 +29,12 @@ class Builtin: ``required_value_kwargs`` are ordinary values that must be present — a number, never a name to resolve (``shift(..., offset=1)``). - Every dimension or lookup an operator names arrives in a kwarg *value*, - which is what lets a macro pass one as a formal. ``usage`` is the wording - every refusal quotes back. + Every operator takes one positional argument, the expression; every + dimension or lookup it names arrives in a kwarg *value*, which is what + lets a macro pass one as a formal. ``usage`` is the wording every refusal + quotes back. """ - positional: int usage: str dimension_kwargs: tuple[str, ...] = () lookup_kwargs: tuple[str, ...] = () @@ -58,43 +51,41 @@ class Builtin: optional_kwargs: tuple[str, ...] = () @property - def keywords(self) -> frozenset[str]: - """Every keyword the call must carry, when they are named at all.""" + def required(self) -> frozenset[str]: + """Every keyword the call must carry.""" return ( (frozenset(self.dimension_kwargs) | frozenset(self.lookup_kwargs) | frozenset(self.required_value_kwargs)) - frozenset(self.at_most_one_of) - frozenset(self.optional_kwargs) ) - @property - def optional(self) -> frozenset[str]: - """Every keyword the call may carry but need not.""" - return frozenset(self.edge_kwargs) | frozenset(self.at_most_one_of) | frozenset(self.optional_kwargs) + def kind_of(self, kwarg: str) -> Literal['dimension', 'lookup', 'edge', 'value']: + """What resolution turns the value of *kwarg* into: a dimension, a lookup, an edge policy, or a plain value.""" + if kwarg in self.dimension_kwargs: + return 'dimension' + if kwarg in self.lookup_kwargs: + return 'lookup' + if kwarg in self.edge_kwargs: + return 'edge' + return 'value' #: The closed operator set. ``by=`` is the one keyword that addresses a lookup, -#: and a lookup carries its own dimensions, so the sibling kwargs that used to -#: restate them (``sum``'s ``over=`` beside ``group_by=``, ``at``'s ``onto=``) -#: are gone — what the two-keyword spelling once said, the name's *kind* now -#: says, checked at load. ``by=`` on ``shift`` and -#: ``sum_back`` partitions the axis the operator walks, which is the same -#: lookup in a different position: it says which rows are neighbours, not which -#: group a term lands in. +#: and a lookup carries its own dimensions, so no sibling kwarg restates them. +#: On ``shift`` and ``sum_back`` it partitions the axis the operator walks: it +#: says which rows are neighbours, not which group a term lands in. BUILTINS: dict[str, Builtin] = { 'sum': Builtin( - 1, 'sum(), sum(, over=) or sum(, by=)', dimension_kwargs=('over',), lookup_kwargs=('by',), at_most_one_of=('over', 'by'), ), 'at': Builtin( - 1, 'at(, by=)', lookup_kwargs=('by',), ), 'sum_back': Builtin( - 1, "sum_back(, over=, within=[, edge='wrap'][, by=])", dimension_kwargs=('over',), lookup_kwargs=('by',), @@ -103,7 +94,6 @@ def optional(self) -> frozenset[str]: optional_kwargs=('by',), ), 'shift': Builtin( - 1, "shift(, over=, offset=[, edge='wrap'|][, by=])", dimension_kwargs=('over',), lookup_kwargs=('by',), @@ -141,7 +131,8 @@ def call_shape_error(name: str, positional: int, kwargs: Iterable[str]) -> str | f'its own dimensions, so by= leaves over= nothing to add.\n' f'Write: {builtin.usage}' ) - fits = positional == builtin.positional and keys - builtin.optional == builtin.keywords + optional = {*builtin.edge_kwargs, *builtin.at_most_one_of, *builtin.optional_kwargs} + fits = positional == 1 and keys - optional == builtin.required return None if fits else f'{name}() expects {builtin.usage}' diff --git a/src/math_spec/piecewise.py b/src/math_spec/piecewise.py index e46ac16a..5bc9d53f 100644 --- a/src/math_spec/piecewise.py +++ b/src/math_spec/piecewise.py @@ -4,38 +4,16 @@ """Expand ``piecewise:`` blocks into plain variables and constraints. -A ``piecewise:`` block becomes ordinary affine declarations before anything -reads the model. The λ convex-combination method needs only the breakpoint -parameters themselves, no derived data. For a block - - piecewise: - curve: - over: bp - links: - - [power, power_bp] - - [fuel * eff, fuel_bp, "<="] - -with F = the union of the links' dims, it emits: - - variables: - curve_lam(F, bp) in [0, 1] - curve_seg(F, bp) binary (method: adjacency) - constraints: - curve_convexity(F): sum(curve_lam, over=bp) == 1 - curve_pick(F): sum(curve_seg, over=bp) == 1 (method: adjacency) - curve_adjacency(F, bp): curve_lam <= curve_seg + shift(curve_seg, over=bp, offset=1, edge=0) - curve_link0(F): (power) == sum(curve_lam * power_bp, over=bp) - curve_link1(F): (fuel * eff) <= sum(curve_lam * fuel_bp, over=bp) - -Only the restriction on λ varies (:data:`~math_spec.model.PIECEWISE_METHODS`); -``lp`` emits no weights at all. A link expression is judged before expansion, -so ``p * p`` is refused against the link the user wrote rather than -``curve_link0``. +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. """ from __future__ import annotations -from typing import TYPE_CHECKING, Any +from typing import Any from math_spec.degree import check_expression from math_spec.dimensions import dims_of @@ -57,21 +35,14 @@ ) from math_spec.resolution import Namespace, resolve_expression -if TYPE_CHECKING: - from collections.abc import Sequence +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 -def _mask_of(block: str, pw: PiecewiseBlock) -> str | None: - """The parameter a block masks its weights with, or ``None`` for a whole curve. - ``points:`` may name the mask itself, or one of the block's own values - parameters — "the curve runs as far as this does" — in which case the - mask is derived from that parameter when data binds, under the name this - returns. - """ - if pw.points is None: - return None - return f'{block}_points' if pw.points in {link.values for link in pw.links} else pw.points +#: The suffix on the second gate row, where the gate variable does not exist. +_UNGATED = '_ungated' def _curvature_required(pw: PiecewiseBlock) -> Curvature | None: @@ -106,7 +77,7 @@ def declaration_of(expanded: ExpandedPiecewise) -> PiecewiseDeclaration: if pw.method == 'lp': checks.append(AtLeastTwo(pw.over, expanded.points)) if expanded.points is not None: - checks.append(Contiguous(expanded.points, _nominated(expanded))) + checks.append(Contiguous(expanded.points, _nominated(pw))) return PiecewiseDeclaration( over=pw.over, method=pw.method, @@ -124,7 +95,7 @@ def derivations_of(block: str, expanded: ExpandedPiecewise) -> dict[str, Derivat if (mask := expanded.points) is None: return {} derivations: dict[str, Derivation] = {} - if (values := _nominated(expanded)) is not None: + if (values := _nominated(expanded.block)) is not None: derivations[mask] = MaskOf(block, values) if expanded.starts is not None: derivations[expanded.starts] = FirstOf(block, mask) @@ -133,43 +104,281 @@ def derivations_of(block: str, expanded: ExpandedPiecewise) -> dict[str, Derivat return derivations -def _nominated(expanded: ExpandedPiecewise) -> str | None: - """The breakpoint parameter the mask was derived from, or ``None`` where the file supplied the mask itself.""" - return expanded.block.points if expanded.block.points != expanded.points else None +class _Block: + """One ``piecewise:`` block being expanded into the raw model it writes. + + Every name the expansion may write is spelled once here, so the emitters + and the collision check read the same table. ``points`` is the derived + mask, written only where ``nominated`` names the values parameter it is + derived from; ``mask`` is whichever parameter masks the weights, or + ``None`` for a whole curve. + + Raises: + PiecewiseExpansionError: A block naming something that does not exist, + or emitting a name the file already declares. + """ + def __init__(self, schema: Spec, raw: dict[str, Any], name: str, pw: PiecewiseBlock) -> None: + self.schema = schema + self.raw = raw + self.name = name + self.pw = pw + self.nominated = _nominated(pw) + self.lam = f'{name}_lam' + self.seg = f'{name}_seg' + self.starts = f'{name}_starts' + self.ends = f'{name}_ends' + self.points = f'{name}_points' + self.convexity = f'{name}_convexity' + self.pick = f'{name}_pick' + self.adjacency = f'{name}_adjacency' + 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.mask = self.points if self.nominated is not None else pw.points + self.frame = self._validated_frame() + self.record: dict[str, Any] = {'block': raw['piecewise'][name], 'points': self.mask} + + def expand(self) -> dict[str, Any]: + """Write the block's declarations, and return the record ``expanded_piecewise`` keeps for it.""" + if self.nominated is not None: + self._parameter( + self.points, + list(self.schema.parameters[self.nominated].dims), + f"where '{self.nominated}' has a row, and so where the curve runs", + ) + if self.pw.method == 'lp': + self._segment_lines() + else: + self._weights() + return self.record + + # -- emitters ---------------------------------------------------------- + + def _parameter(self, name: str, dims: list[str], description: str) -> None: + """A ``bool`` parameter the expansion derives.""" + self.raw.setdefault('parameters', {})[name] = {'dims': dims, 'dtype': 'bool', 'description': description} + + def _weight(self, name: str, **fields: Any) -> None: + """A variable over the frame and the breakpoint dim, masked as the block is.""" + self.raw['variables'][name] = { + 'foreach': [*self.frame, self.pw.over], + **({'where': self.mask} if self.mask else {}), + **fields, + } -def _gate_rows(schema: Spec, pw: PiecewiseBlock) -> tuple[tuple[str, str | None, str], ...]: - """What the weights sum to, as ``(name suffix, where, right-hand side)``. + def _constraint(self, name: str, foreach: list[str], expression: str, where: str | None = None) -> None: + self.raw['constraints'][name] = { + 'foreach': foreach, + **({'where': where} if where else {}), + 'expression': expression, + } - One row where the gate exists at every coordinate the block builds a curve - for, and **two** where it does not. A gate is a variable, so a masked one - has coordinates where it does not exist — and there the block is ungated, - which is the ``1`` a block with no ``activity:`` gets. Written as a single - row it would instead be *no row*: absence does not spread out of a - reduction, so the right-hand side would take the row with it and leave the - weights without the convexity that makes them a curve at all (#1158). + 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, + 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( + cname, + list(self.frame), + f'({link.expression}) {link.sign} sum({self.lam} * {link.values}, over={d})', + ) + if self.pw.method == 'sos2': + self.raw.setdefault('sos', {})[self.name] = {'variable': self.lam, 'over': d, 'type': 2} + elif self.pw.method == 'adjacency': + self._weight(self.seg, domain='binary', bounds={}) + for suffix, where, rhs in gated: + self._constraint(self.pick + suffix, list(self.frame), f'sum({self.seg}, over={d}) == {rhs}', where) + self._constraint( + self.adjacency, + [*self.frame, d], + f'{self.lam} <= {self.seg} + shift({self.seg}, over={d}, offset=1, edge=0)', + ) - ``absence: zero`` is the other reading and stays one row — the gate is 0 - where it does not exist, so the curve is pinned off there. - """ - if pw.activity is None: - return (('', None, '1'),) - gate = schema.variables[pw.activity] - if gate.where is None or gate.absence == 'zero': - return (('', None, f'({pw.activity})'),) - return (('', pw.activity, f'({pw.activity})'), ('_ungated', f'NOT {pw.activity}', '1')) + def _gate_rows(self) -> tuple[tuple[str, str | None, str], ...]: + """What the weights sum to, as ``(name suffix, where, right-hand side)``. + + One row where the gate exists at every coordinate the block builds a curve + for, and **two** where it does not. A gate is a variable, so a masked one + has coordinates where it does not exist — and there the block is ungated, + which is the ``1`` a block with no ``activity:`` gets. Written as a single + row it would instead be *no row*: absence does not spread out of a + reduction, so the right-hand side would take the row with it and leave the + weights without the convexity that makes them a curve at all (#1158). + + ``absence: zero`` is the other reading and stays one row — the gate is 0 + where it does not exist, so the curve is pinned off there. + """ + activity = self.pw.activity + if activity is None: + return (('', None, '1'),) + gate = self.schema.variables[activity] + if gate.where is None or gate.absence == 'zero': + return (('', None, f'({activity})'),) + return (('', activity, f'({activity})'), (_UNGATED, f'NOT {activity}', '1')) + + def _segment_lines(self) -> None: + """The segment-line form: a row per segment, and the two domain rows. + + The chord sits at the later breakpoint, so the first has none and its + ``where:`` and ``edge=0`` travel together — without the exclusion the + vacated position is a spurious line through the origin; under a mask the + first breakpoint is the curve's own, which is what ``_starts`` names. The + row is multiplied through by the run rather than dividing, which keeps its + sense only because the breakpoints are strictly monotone. The domain rows + are ``linopy``'s ``_add_lp`` rows under its names; under ``points:`` they + sit on the derived ``_starts``/``_ends`` flags, which is why the mask has to + be a prefix. + """ + x_link, y_link = self.pw.curve + d = self.pw.over + mask = self.mask + run = f'({x_link.values} - shift({x_link.values}, over={d}, offset=1, edge=0))' + rise = f'({y_link.values} - shift({y_link.values}, over={d}, offset=1, edge=0))' + interior = f'{mask} AND NOT {self.starts}' if mask else f'position({d}) != 0' + self._constraint( + self.chord, + [*self.frame, d], + f'({y_link.expression}) * {run} {y_link.sign} ' + f'{rise} * (({x_link.expression}) - {x_link.values}) + {y_link.values} * {run}', + interior, + ) + edges = ((self.domain_lo, '>=', self.starts), (self.domain_hi, '<=', self.ends)) + axis = ((self.domain_lo, '>=', f'position({d}) == 0'), (self.domain_hi, '<=', f'position({d}) == -1')) + for cname, sense, at in edges if mask else axis: + if mask: + self.record['starts' if sense == '>=' else 'ends'] = at + self._parameter( + at, + self.raw['parameters'][mask]['dims'], + f'the {"first" if sense == ">=" else "last"} breakpoint of each curve', + ) + self._constraint(cname, [*self.frame, d], f'({x_link.expression}) {sense} {x_link.values}', at) + + # -- 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.""" + return ( + ('variable', (self.lam, self.seg)), + ('parameter', (self.starts, self.ends, *((self.points,) if self.nominated is not None else ()))), + ( + 'constraint', + ( + self.convexity, + self.convexity + _UNGATED, + self.pick, + self.pick + _UNGATED, + self.adjacency, + self.chord, + self.domain_lo, + self.domain_hi, + *self.links, + ), + ), + ('sos', (self.name,)), + ) + + def _validated_frame(self) -> tuple[str, ...]: + """Check references and infer the frame (union of the links' 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. + """ + schema, pw = self.schema, self.pw + ctx = f"piecewise '{self.name}'" + if pw.over not in schema.dimensions: + raise PiecewiseExpansionError(undeclared_dimension('piecewise', self.name, pw.over)) + + frame: list[str] = [] + for i, link in enumerate(pw.links): + values = link.values + if values not in schema.parameters: + raise PiecewiseExpansionError(f"{ctx}: link {i} values references undeclared parameter '{values}'") + if pw.over not in schema.parameters[values].dims: + raise PiecewiseExpansionError( + f"{ctx}: link {i} values parameter '{values}' must carry dim " + f"'{pw.over}' (has {schema.parameters[values].dims})" + ) + for d in _declared_order(schema, _expr_dims(schema, link.expression, f'{ctx} link {i}')): + if d == pw.over: + raise PiecewiseExpansionError( + f"{ctx}: link {i} expression already carries the breakpoint dim '{pw.over}'" + ) + if d not in frame: + frame.append(d) + + if pw.activity is not None: + if pw.activity not in schema.variables: + raise PiecewiseExpansionError( + f"{ctx}: activity '{pw.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 schema.variables[pw.activity].domain != 'binary': + raise PiecewiseExpansionError(f"{ctx}: activity variable '{pw.activity}' must be binary") + for d in _declared_order(schema, _expr_dims(schema, pw.activity, f'{ctx} activity')): + if d == pw.over: + raise PiecewiseExpansionError(f"{ctx}: activity must not carry 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 PiecewiseExpansionError( + f"{ctx}: 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 self.nominated is None: + if pw.points not in schema.parameters: + raise PiecewiseExpansionError(f"{ctx}: points references undeclared parameter '{pw.points}'") + if (dtype := 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 = 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" + ) + + declared = { + 'variable': schema.variables, + 'parameter': schema.parameters, + 'constraint': schema.constraints, + 'sos': schema.sos, + } + for kind, names in self._emitted_by_kind(): + for one in names: + if one in declared[kind]: + raise PiecewiseExpansionError(f"{ctx}: emitted {kind} '{one}' collides with a declared {kind}") + return tuple(frame) def expand_piecewise(schema: Spec) -> _ExpandedSpec: """Return *schema* as a :class:`_ExpandedSpec` — every ``piecewise:`` block expanded away. - The adjacency row shifts with ``edge=0``: a bare ``shift`` would drop the - first breakpoint's row and leave its weight unconstrained, a wrong MILP - with no error (#289). ``points:`` masks the weights and the segment - binaries and no constraint — every emitted row reduces over the breakpoint - axis or carries a masked weight. The result is memoised on *schema*, and a - :class:`_ExpandedSpec` comes straight back; a model with no ``piecewise:`` is - retyped with ``model_construct``, its validation already done on the way in. + Memoised on *schema*. Raises: PiecewiseExpansionError: A block naming something that does not exist, @@ -186,223 +395,13 @@ def expand_piecewise(schema: Spec) -> _ExpandedSpec: raw = schema.model_dump() raw.setdefault('variables', {}) raw.setdefault('constraints', {}) - records: dict[str, dict[str, Any]] = {} - raw['expanded_piecewise'] = records - for name, pw in schema.piecewise.items(): - frame = _validate_block(schema, name, pw) - mask, nominated = _mask_of(name, pw), pw.points - record = records[name] = {'block': raw['piecewise'][name], 'points': mask} - if mask is not None and nominated is not None and mask != nominated: - _emit_parameter( - raw, - mask, - list(schema.parameters[nominated].dims), - f"where '{nominated}' has a row, and so where the curve runs", - ) - if pw.method == 'lp': - _expand_lp(raw, record, name, pw, frame, mask, schema.parameters[pw.points].dims if pw.points else ()) - continue - lam = f'{name}_lam' - - raw['variables'][lam] = { - 'foreach': [*frame, pw.over], - **({'where': mask} if mask else {}), - 'bounds': {'lower': 0.0, 'upper': 1.0}, - 'description': 'convex-combination weight on a breakpoint', - } - gated = _gate_rows(schema, pw) - for suffix, where, rhs in gated: - raw['constraints'][f'{name}_convexity{suffix}'] = { - 'foreach': list(frame), - **({'where': where} if where else {}), - 'expression': f'sum({lam}, over={pw.over}) == {rhs}', - } - for i, link in enumerate(pw.links): - raw['constraints'][f'{name}_link{i}'] = { - 'foreach': list(frame), - 'expression': (f'({link.expression}) {link.sign} sum({lam} * {link.values}, over={pw.over})'), - } - if pw.method == 'sos2': - raw.setdefault('sos', {})[name] = {'variable': lam, 'over': pw.over, 'type': 2} - elif pw.method == 'adjacency': - seg = f'{name}_seg' - raw['variables'][seg] = { - 'foreach': [*frame, pw.over], - **({'where': mask} if mask else {}), - 'domain': 'binary', - 'bounds': {}, - } - for suffix, where, rhs in gated: - raw['constraints'][f'{name}_pick{suffix}'] = { - 'foreach': list(frame), - **({'where': where} if where else {}), - 'expression': f'sum({seg}, over={pw.over}) == {rhs}', - } - raw['constraints'][f'{name}_adjacency'] = { - 'foreach': [*frame, pw.over], - 'expression': f'{lam} <= {seg} + shift({seg}, over={pw.over}, offset=1, edge=0)', - } - + raw['expanded_piecewise'] = {name: _Block(schema, raw, name, pw).expand() for name, pw in schema.piecewise.items()} raw['piecewise'].clear() expanded = _ExpandedSpec.model_validate(raw) schema._expansion = expanded return expanded -def _expand_lp( - raw: dict[str, Any], - record: dict[str, Any], - name: str, - pw: PiecewiseBlock, - frame: tuple[str, ...], - mask: str | None, - schema_dims: Sequence[str], -) -> None: - """Emit the segment-line form: a row per segment, and the two domain rows. - - The chord sits at the later breakpoint, so the first has none and its - ``where:`` and ``edge=0`` travel together — without the exclusion the - vacated position is a spurious line through the origin; under a mask the - first breakpoint is the curve's own, which is what ``_starts`` names. The - row is multiplied through by the run rather than dividing, which keeps its - sense only because the breakpoints are strictly monotone. The domain rows - are ``linopy``'s ``_add_lp`` rows under its names; under ``points:`` they - sit on the derived ``_starts``/``_ends`` flags, which is why the mask has to - be a prefix. - """ - x_link, y_link = pw.curve - d = pw.over - run = f'({x_link.values} - shift({x_link.values}, over={d}, offset=1, edge=0))' - rise = f'({y_link.values} - shift({y_link.values}, over={d}, offset=1, edge=0))' - interior = f'{mask} AND NOT {name}_starts' if mask else f'position({d}) != 0' - raw['constraints'][f'{name}_chord'] = { - 'foreach': [*frame, d], - 'where': interior, - 'expression': ( - f'({y_link.expression}) * {run} {y_link.sign} ' - f'{rise} * (({x_link.expression}) - {x_link.values}) + {y_link.values} * {run}' - ), - } - edges = (('domain_lo', '>=', f'{name}_starts'), ('domain_hi', '<=', f'{name}_ends')) - axis = (('domain_lo', '>=', f'position({d}) == 0'), ('domain_hi', '<=', f'position({d}) == -1')) - for suffix, sense, at in edges if mask else axis: - if mask: - record['starts' if sense == '>=' else 'ends'] = at - _emit_parameter( - raw, - at, - list(schema_dims), - f'the {"first" if sense == ">=" else "last"} breakpoint of each curve', - ) - raw['constraints'][f'{name}_{suffix}'] = { - 'foreach': [*frame, d], - 'where': at, - 'expression': f'({x_link.expression}) {sense} {x_link.values}', - } - - -def _validate_block(schema: Spec, name: str, pw: PiecewiseBlock) -> tuple[str, ...]: - """Check references and infer the frame (union of the links' 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. - """ - ctx = f"piecewise '{name}'" - if pw.over not in schema.dimensions: - raise PiecewiseExpansionError(undeclared_dimension('piecewise', name, pw.over)) - - frame: list[str] = [] - for i, link in enumerate(pw.links): - values = link.values - if values not in schema.parameters: - raise PiecewiseExpansionError(f"{ctx}: link {i} values references undeclared parameter '{values}'") - if pw.over not in schema.parameters[values].dims: - raise PiecewiseExpansionError( - f"{ctx}: link {i} values parameter '{values}' must carry dim " - f"'{pw.over}' (has {schema.parameters[values].dims})" - ) - for d in _declared_order(schema, _expr_dims(schema, link.expression, f'{ctx} link {i}')): - if d == pw.over: - raise PiecewiseExpansionError( - f"{ctx}: link {i} expression already carries the breakpoint dim '{pw.over}'" - ) - if d not in frame: - frame.append(d) - - if pw.activity is not None: - if pw.activity not in schema.variables: - raise PiecewiseExpansionError( - f"{ctx}: activity '{pw.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 schema.variables[pw.activity].domain != 'binary': - raise PiecewiseExpansionError(f"{ctx}: activity variable '{pw.activity}' must be binary") - for d in _declared_order(schema, _expr_dims(schema, pw.activity, f'{ctx} activity')): - if d == pw.over: - raise PiecewiseExpansionError(f"{ctx}: activity must not carry 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 PiecewiseExpansionError( - f"{ctx}: 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 _mask_of(name, pw) == pw.points: - if pw.points not in schema.parameters: - raise PiecewiseExpansionError(f"{ctx}: points references undeclared parameter '{pw.points}'") - if (dtype := 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 = 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" - ) - - emitted_constraints = ( - f'{name}_convexity', - f'{name}_convexity_ungated', - f'{name}_pick', - f'{name}_pick_ungated', - f'{name}_adjacency', - f'{name}_chord', - f'{name}_domain_lo', - f'{name}_domain_hi', - *(f'{name}_link{i}' for i in range(len(pw.links))), - ) - derived = (f'{name}_points',) if _mask_of(name, pw) != pw.points else () - for kind, emitted, declared in ( - ('variable', (f'{name}_lam', f'{name}_seg'), schema.variables), - ('parameter', (f'{name}_starts', f'{name}_ends', *derived), schema.parameters), - ('constraint', emitted_constraints, schema.constraints), - ('sos', (name,), schema.sos), - ): - for one in emitted: - if one in declared: - raise PiecewiseExpansionError(f"{ctx}: emitted {kind} '{one}' collides with a declared {kind}") - return tuple(frame) - - -def _emit_parameter(raw: dict[str, Any], name: str, dims: list[str], description: str) -> None: - """Write a ``bool`` parameter the expansion derives.""" - raw.setdefault('parameters', {})[name] = {'dims': dims, 'dtype': 'bool', 'description': description} - - def _declared_order(schema: Spec, dims: frozenset[str]) -> list[str]: """*dims* in declaration order — iterating the set varies the emitted ``foreach``, and every column index behind it, per process.""" return [d for d in schema.dimensions if d in dims] diff --git a/src/math_spec/program.py b/src/math_spec/program.py index 2399d5ef..e217e42a 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -4,54 +4,17 @@ """The program: what a file declares, with names resolved and shapes fixed. -A :class:`Program` is a complete declarative description of a linear program -over named tidy tables — every declaration a file makes, and no data in it at -all. Data is bound against these declarations by whatever builds the model; -:func:`~math_spec.lowering.to_program` is what produces one from a spec. - -**It is the second public state, and the one a consumer reads.** A -:class:`~math_spec.model.Spec` is what the file *says*; a program is what it -*means*, with macros expanded, names typed, operators resolved to nodes and -every dim rule already checked. Consumers dispatch on these nodes and read -them; nothing here is built by hand, so what ships beside the nodes is the -walk (:func:`children`), not builders. A program is trusted by construction: -:func:`~math_spec.lowering.to_program` is the only thing that builds one, and -nothing checks one assembled by hand. The language's refusals happen at load, -where the file and its author are, and a program put together some other way -is outside that guarantee rather than inside a pass restating it. - -What a consumer needs from this module falls in three, and only the middle -one has to be *called* to be got right: - -- **Types to match on** — every node and declaration class, the - :data:`ExpressionNode` union, and the ``Literal`` vocabularies. A backend - dispatches on these and calls none of them. -- **Rules to call** — :func:`children` and :func:`fan_in`. Neither is visible - in a node's own structure, so a consumer deriving them derives them wrongly - the day a node is added. -- **Questions over the walk** — :func:`walk` and the filters beside it. Each is - a line a consumer could write; they are here so two consumers cannot write it - differently. - -A mask is the language's own resolved ``where`` node -(:mod:`math_spec.where_parser`) rather than a second set spelling the same -predicates — one home, so the two cannot come to disagree about what a -comparison is. Its literals are already decided: a mask admitting every row -arrives as ``None`` and one admitting none as ``BooleanLiteralNode(False)``, so -that node stands at the root of a mask or nowhere in it, and no consumer needs -a constant folder of its own to agree with the others about which rows exist. - -The declaration vocabularies are the language's own for the same reason -(:mod:`math_spec.model`): a ``dtype``, a domain and an absence reading cross -into a program by a cast, and a member added to one spelling alone would -arrive as a string no consumer's branch recognises. - -Frozen dataclasses only — no execution logic, and nothing imported from a -consumer. - -Expressions support operator sugar so programs read naturally in Python: - -balance = GroupSum(Variable("p"), over="generator", coordinate=("bus",), into=("bus",)) - Parameter("load") +The second public state, and the one a consumer reads. A :class:`Program` is +every declaration a file makes and no data at all; +:func:`~math_spec.lowering.to_program` is the only thing that builds one, so +nothing here re-checks a hand-built one. + +Node and declaration classes are matched with ``isinstance``. The rules a +node's structure does not show are :func:`children` and :func:`fan_in`; the +questions over the walk are :func:`walk` and the filters beside it. A +resolved ``where`` arrives as a :class:`Mask`. Frozen dataclasses only — no +execution logic, and nothing imported from a consumer. How a consumer reads +one: ``docs/reference/language/reading.md``. """ from __future__ import annotations @@ -64,34 +27,34 @@ import math_spec.model as _model from math_spec.errors import did_you_mean +from math_spec.expression_parser import ComparisonOperator if TYPE_CHECKING: + import datetime from collections.abc import Iterator - from math_spec.where_parser import WhereNode - -#: What ``math_spec.program`` promises. The package's ``__all__`` exports this -#: *module* rather than its names, so this is where a consumer's imports are -#: pinned. Sorted, like the package's own: the grouping a reader wants is in -#: ``tests/test_public_surface.py``, which derives this set from the module -#: rather than restating it, so the two cannot come apart. +#: What ``math_spec.program`` promises a consumer, sorted. __all__ = [ 'QUADRATIC_POSITIONS', 'Add', + 'AndNode', 'At', 'AtLeastTwo', + 'BooleanLiteralNode', 'Cases', 'Check', - 'ComparisonOperator', + 'ConnectiveWhereNode', 'Constant', 'ConstraintDeclaration', 'ConstraintSense', 'Contiguous', 'Curved', 'Derivation', + 'DimensionComparisonNode', 'DimensionDeclaration', 'DimensionDtype', + 'DimensionPositionNode', 'Divide', 'Expression', 'ExpressionNode', @@ -101,27 +64,40 @@ 'GroupSum', 'Increasing', 'LastOf', + 'LookupComparisonNode', 'LookupDeclaration', + 'LookupDefinedNode', + 'LookupPairComparisonNode', + 'Mask', 'MaskOf', 'Multiply', 'Negate', + 'NotNode', 'ObjectiveDeclaration', 'ObjectiveSense', + 'OrNode', 'Parameter', + 'ParameterComparisonNode', 'ParameterDeclaration', + 'ParameterDefinedNode', 'ParameterDtype', 'PiecewiseDeclaration', 'Power', + 'PredicateOperator', 'Program', 'QuadraticPosition', 'Region', + 'Separability', 'SosDeclaration', 'Sum', 'Translate', + 'TypedPredicateNode', 'Variable', 'VariableAbsence', 'VariableDeclaration', + 'VariableDefinedNode', 'VariableType', + 'WhereNode', 'Window', 'carries_variable', 'check_message', @@ -136,16 +112,10 @@ ] -ConstraintSense = Literal['==', '<=', '>='] +ConstraintSense = ComparisonOperator -#: How a shape operator's output rows relate to its input slots — the absence -#: rules' own distinction, as a field. ``sum``, ``sum(by=)`` and a window put -#: several input slots into one output row, so an absent slot there is one -#: summand fewer and the row stands; a pullback and a translation are one slot -#: for one, so an absent input *is* the output and takes the row with it. -#: Answered by :func:`fan_in` for every node, because a consumer keeping its -#: own list of which is which would be deciding a rule the language has -#: already decided. +#: 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'] @@ -158,7 +128,6 @@ #: ``QUADRATIC_POSITIONS <= handled`` is how one says it covers every position #: and hears about it when the language admits another. QUADRATIC_POSITIONS = frozenset(get_args(QuadraticPosition)) -ComparisonOperator = Literal['==', '!=', '<=', '>=', '<', '>'] #: What a dimension's labels are — the language's own vocabulary #: (:data:`~math_spec.model.DimensionDtype`), under the name a consumer reads @@ -186,27 +155,16 @@ class Expression: """Base class for expressions over variables and parameters. - Affine everywhere but the objective, where a :class:`Multiply` of two - variable-carrying operands is degree 2; which position allows what is - ``math_spec.degree``'s to say and no node here records. - - The four operators exist for the tests that compose plans by hand; - constructing Programs in Python is not supported API, so there is no - scalar coercion and no reflected form. + Affine everywhere but where ``math_spec.degree`` admits a :class:`Multiply` + of two variable-carrying operands; no node records which position that is. """ def __add__(self: ExpressionNode, other: ExpressionNode) -> ExpressionNode: return Add(self, other) - def __sub__(self: ExpressionNode, other: ExpressionNode) -> ExpressionNode: - return Add(self, Negate(other)) - def __mul__(self: ExpressionNode, other: ExpressionNode) -> ExpressionNode: return Multiply(self, other) - def __neg__(self: ExpressionNode) -> ExpressionNode: - return Negate(self) - @dataclass(frozen=True) class Constant(Expression): @@ -244,10 +202,8 @@ class Add(Expression): class Multiply(Expression): """Product of two operands. - Affine where at least one factor is variable-free. **Degree 2 where neither - is**, which the language allows in the objective alone - (``math_spec.degree``) — so a consumer that cannot represent a quadratic - term is told which position it is compiling rather than assuming it. + Affine where at least one factor is variable-free; degree 2 where neither + is, which ``math_spec.degree`` admits in a :data:`QuadraticPosition` alone. """ left: ExpressionNode @@ -288,17 +244,11 @@ class Sum(Expression): class GroupSum(Expression): """Sum ``operand`` through coordinates declared on dim ``over``. - ``coordinate`` names coordinates carried by dim ``over`` whose values are + ``coordinate`` names lookups carried by dim ``over`` whose values are labels of the matching dim in ``into``; the result replaces ``over`` with - all of them. - - ``into`` restates each coordinate's declared target, because a node is - read on its own — a consumer places terms from one without consulting the - program — and lowering is the only thing that writes it. - - Several coordinates are one grouping into a product of targets, not a - composition of groupings — they are consumed in a single join, so the pair - of tuples is always the same length and their order pairs them up. + all of them. The two tuples are the same length and their order pairs + them: several coordinates are one grouping into a product of targets, + consumed in a single join. """ operand: ExpressionNode @@ -313,12 +263,7 @@ class At(Expression): Same mapping table, walked the other way: ``GroupSum`` consumes ``over`` and produces ``into``, this consumes ``into`` and produces ``over``. The - fields are named for the *table* rather than the direction, so the pair - reads as one relation; the surface says which end you stand on - (``sum(by=)`` consumes it, ``at(by=)`` produces it, the lookup names the map). - - The join fans out, many ``over`` labels sharing one ``into`` tuple — the - fan-out ``GroupSum`` pays in reverse, so the locality class is unchanged. + join fans out, many ``over`` labels sharing one ``into`` tuple. """ operand: ExpressionNode @@ -329,34 +274,20 @@ class At(Expression): @dataclass(frozen=True) class Translate(Expression): - """Re-index along one dimension: the result at *t* is ``operand`` at *t - by*. - - One node for the whole of ``shift``, whose ``edge=`` decides ``wrap``: - ``edge='wrap'`` is periodic, absent or numeric is not. + """Re-index along one dimension: the result at *t* is ``operand`` at *t - offset*. - ``wrap`` carries no default, on this node or on :class:`Window`. Whether an - axis closes onto itself is the difference between a battery that must end - as it started and one that need not, and there is no reading of a - translation that leaves it unsaid — a node that guessed would be answering - for the file. + ``wrap`` is ``edge='wrap'`` in the file: periodic, and stated on every + node. ``fill`` is what an acyclic shift leaves behind: ``None`` leaves the + vacated positions absent, so the row drops; a number makes them present + and contribute it. Always ``None`` under ``wrap``. - ``fill`` decides what an acyclic shift leaves behind. ``None``, what bare - ``shift`` lowers to, leaves the vacated positions **absent**: they carry no - value, the absence rules propagate that, and the row drops. A number makes - them present and contribute it, which is the only way a file can say - "before the axis starts, read zero" without inventing coordinates. Always - ``None`` under ``wrap``, a cyclic map vacating nothing. + ``offset`` is an integer, or the name of an integer parameter that does + not depend on ``dimension`` and carries its sign in the values. - ``offset`` is how far back to reach: an integer, or the name of an integer - parameter when it differs per entity — a construction lead time, a transit - time, a minimum up time. A named offset may not depend on the dimension - being translated, and carries its sign in the values. - - ``partition`` names a lookup over ``dimension``, and then the translation - happens **inside each group** it makes: the neighbour of a coordinate is the - one before it *in its own group*, the edge is that group's edge, and a wrap - closes each group onto itself. A coordinate the lookup sends nowhere is in - no group and reaches nothing. + ``partition`` names a lookup over ``dimension``, and the translation then + happens inside each group it makes: the neighbour is the one before in + the same group, the edge is the group's, and a wrap closes each group onto + itself. A coordinate the lookup sends nowhere reaches nothing. """ operand: ExpressionNode @@ -381,19 +312,11 @@ class Window(Expression): horizon. A named width may not depend on the dimension being summed over. ``wrap`` says whether the window reaches around the start of the axis - instead of stopping short at it, and is stated at every construction for - the reason :class:`Translate` gives. + instead of stopping short at it, and is stated on every node. ``partition`` names a lookup over that dimension, and the window then stops - at each group's edge: a representative day, a season, a scenario's own run - of hours. Positions are counted inside the group rather than along the - axis, so a coordinate the lookup places nowhere reaches nothing at all — - not even itself. - - One node rather than a sum of ``Translate``s, because the number of terms - would then be read from data and the program's *shape* is fixed before any - data is bound. What data supplies is the mask's cardinality, exactly as it - supplies how many snapshots there are. + at each group's edge. Positions are counted inside the group, so a + coordinate the lookup places nowhere reaches nothing — not even itself. """ operand: ExpressionNode @@ -407,14 +330,11 @@ class Window(Expression): class Region: """One region of a :class:`Cases`: where it applies, and the value there. - ``when`` is stated on every region, the one the file wrote as - ``otherwise:`` included — its mask is the negation of the others, resolved - once here rather than by each consumer in turn. A consumer builds a region - without holding the rest in mind, and both facts it needs are on the region - it is reading. + ``when`` is stated on every region; the one the file wrote as + ``otherwise:`` carries the negation of the others. """ - when: WhereNode + when: Mask value: ExpressionNode @@ -422,16 +342,9 @@ class Region: class Cases(Expression): """A value defined by region — exactly one region applies at each coordinate. - The language proves the regions apart before any data binds, and the - file's ``otherwise:`` covers whatever the rest leave, so they are disjoint - and total by construction: a consumer adds the regions rather than ranking - them, and needs neither an order nor a tie-break. - - Not a shape operator — every region spans the dims the expression does, and - this neither reduces nor replicates. What it adds is the one thing no other - node here carries: **a mask in a value position**. A consumer that can - restrict rows but cannot weigh a term by a predicate builds each region - against its own mask and adds the results. + The regions are disjoint and total, so a consumer adds them rather than + ranking them. Not a shape operator: every region spans the dims the + expression does. """ regions: tuple[Region, ...] @@ -464,18 +377,8 @@ class Cases(Expression): def fan_in(expression: ExpressionNode) -> FanIn: """How *expression*'s output rows relate to its input slots. - Total over the node set, so a consumer asks any node rather than keeping a - list of which kinds carry the answer. Arithmetic and the leaves reshape - nothing, which is one slot for one row — the same class a pullback and a - translation are in, reached for a different reason. - - Exhaustive rather than defaulted: a node added without a case here is a - type error at this function, where the absence rule it needs is decided, - instead of silently inheriting the class that reshapes nothing. - - :class:`Cases` is in that class too, for a reason of its own: its regions - are disjoint, so an output row reads exactly one of them — the several - values it holds are alternatives rather than slots summed together. + For the absence rules, both classes other than ``'one-to-one'`` sum + several input slots into an output row. """ if isinstance(expression, (Sum, GroupSum)): return 'many-to-one' @@ -490,12 +393,7 @@ def fan_in(expression: ExpressionNode) -> FanIn: def children(expression: ExpressionNode) -> tuple[ExpressionNode, ...]: - """The sub-expressions of *expression* — the structural half of any walk. - - Every walk over a program's expressions recurses through here and differs only in - what it does at the leaves. Enumerating the children once is how a node - added later reaches all of them rather than one. - """ + """The sub-expressions of *expression* — what every walk recurses through.""" if isinstance(expression, Negate): return (expression.operand,) if isinstance(expression, (Add, Multiply)): @@ -725,7 +623,7 @@ class ParameterDeclaration: @dataclass(frozen=True) class VariableDeclaration: dims: tuple[str, ...] - where: WhereNode | None = None + where: Mask | None = None lower: ExpressionNode = field(default_factory=lambda: Constant(float('-inf'))) upper: ExpressionNode = field(default_factory=lambda: Constant(float('inf'))) variable_type: VariableType = 'continuous' @@ -745,7 +643,7 @@ class ConstraintDeclaration: lhs: ExpressionNode sense: ConstraintSense rhs: ExpressionNode - where: WhereNode | None = None + where: Mask | None = None @dataclass(frozen=True) @@ -779,41 +677,17 @@ class ObjectiveDeclaration: @dataclass(frozen=True) class Footprint: - """Which of the language's constructs one program actually reaches for. - - A *subset*, never the whole: the language admits more than any one file - uses, and an empty field says this program does not use that construct — - not that the construct does not exist. Every field is a set, so - ``if footprint.x`` asks whether it appears at all and ``y in footprint.x`` - asks about one kind, and a construct admitted later widens a set rather - than needing a field a consumer does not yet read. + """Which of the language's constructs one program uses. - Facts only. What a sink can ingest is a separate axis - (``docs/about/ceiling.md``, "Capability is not the ceiling"), where a - capability is neither a flat set nor one verdict per construct — so there - is deliberately no verdict here to read instead of giving one. - - Nothing below the kind, either: a sink that takes a window but not a - wrapped one reads ``Window in shapes`` and then walks, because ``wrap``, - ``partition`` and a named width are refinements without end and each is one - line once the set has said where to look. + A subset, never the whole: an empty field says this program does not use + the construct. Attributes: quadratic: Each position a product of two variable-carrying operands - stands in. Empty is affine throughout. Convexity is not here: it is - a property of the whole Hessian rather than of any term, and the - coefficients deciding it arrive with the data — so, as with a - curve's shape (:class:`Curved`), - this names where the products are and the caller holding the - numbers does the checking. - variable_types: Every domain declared, ``{'continuous'}`` alone being - the pure-LP case. - sos_types: The order of each special-ordered set declared. Empty where - the file declares none. - shapes: Every expression node kind that appears, complete rather than - curated — picking the interesting ones would be the judgement this - leaves to the consumer, and a node added later is reported without - anyone remembering a filter. + stands in; empty is affine throughout. + variable_types: Every domain declared. + sos_types: The order of each special-ordered set declared. + shapes: Every expression node kind that appears. """ quadratic: frozenset[QuadraticPosition] @@ -830,9 +704,76 @@ def _declared[Declaration](items: Mapping[str, Declaration], name: str, kind: st raise KeyError(f"unknown {kind} '{name}'. " + did_you_mean(name, list(items))) from None +@dataclass(frozen=True) +class Separability: + """What building one dimension a window at a time asks of a driver, and what it would break. + + A rolling-horizon or myopic driver cuts an axis into windows and builds + each on its own. What the program can say is whether every row it builds + is then complete inside some window: how far a row reads along the axis, + and which declarations tie the axis together so that no window holds them. + It cannot say whether the windowed answer is the one a whole-horizon solve + would give — a store carried over one row windows cleanly, and a rolling + solve of it is still a different answer — which is the driver's design and + not the model's. + + Attributes: + dimension: The axis asked about. + behind: Coordinates a window must see before its first row for every + row it builds to be complete — what a trailing window or a + positive ``shift`` reads. ``0`` is pointwise; a ``shift`` of one is + ``1``; a ``sum_back`` of ``n`` is ``n - 1``. + ahead: The same after its last row — what a negative ``shift`` reads. + 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 + 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. + undecided: Each declaration whose reach along the axis only data can + say, to the parameter or lookup that says it — a named offset or + width, a partition whose groups a window may cut, a read through a + lookup at a coordinate the data chooses. With data bound, the reach + is the driver's to compute. + restarts: Each declaration counting a position along the axis, which a + window restarts at its first row. Whether that is wanted — a seed + once per window, or once per horizon — is the modeller's, so it is + reported rather than refused. + """ + + dimension: str + behind: int + ahead: int + coupled: Mapping[str, str] + undecided: Mapping[str, str] + restarts: Mapping[str, str] + + @property + def windowable(self) -> bool: + """Whether every row builds complete inside a window overlapping by :attr:`behind` and :attr:`ahead`. + + ``False`` while a reach is :attr:`undecided`, which a driver holding + the data may resolve; :attr:`restarts` do not count against it. + """ + return not self.coupled and not self.undecided + + @property + def independent(self) -> bool: + """Whether each coordinate builds on its own: no row reads another, nothing ties them, nothing counts them. + + What a driver solving one coordinate per slice — a scenario sweep — + asks, and what licenses solving the slices in any order or at once. + A :attr:`restarts` entry counts against it, unlike for + :attr:`windowable`: with one coordinate per slice a ``position()`` + holds everywhere, which changes what the mask means rather than + where it restarts. + """ + return self.windowable and not self.behind and not self.ahead and not self.restarts + + @dataclass(frozen=True, kw_only=True) class Program: - """A complete linear program over named tidy tables. + """A complete declarative description of a mathematical program, with no data in it. Every group of declarations is keyed by the name the file wrote, in the order it wrote them, and is read-only: the mappings are wrapped at @@ -859,45 +800,31 @@ class Program: named_expressions: Mapping[str, ExpressionNode] = MappingProxyType({}) def __post_init__(self) -> None: - """Seal every group, so a program handed out cannot be written to. - - ``frozen=True`` stops a field being rebound and says nothing about the - mapping behind it. Wrapping here rather than trusting the caller is - what makes the guarantee hold for every construction path. - """ + """Seal every group, so a program handed out cannot be written to.""" for f in fields(self): group = getattr(self, f.name) if isinstance(group, Mapping): object.__setattr__(self, f.name, MappingProxyType(dict(group))) + def _by_position(self) -> Iterator[tuple[QuadraticPosition, tuple[ExpressionNode, ...]]]: + """The row-building expressions, grouped by the position they stand in.""" + yield 'objective', (self.objective.expression,) if self.objective is not None else () + yield 'constraint', tuple(side for c in self.constraints.values() for side in (c.lhs, c.rhs)) + @property def expressions(self) -> tuple[ExpressionNode, ...]: """Every expression a row is built from — the objective and both sides of each constraint. - What a walk over the program *a solver sees* takes. A declared - :attr:`named_expressions` entry is not among them: it builds no row, so - a question asked about what will be solved would answer wrongly if it - counted one. + A :attr:`named_expressions` entry builds no row and is not among them. """ - return ( - *((self.objective.expression,) if self.objective is not None else ()), - *(side for c in self.constraints.values() for side in (c.lhs, c.rhs)), - ) + return tuple(e for _, group in self._by_position() for e in group) @cached_property def footprint(self) -> Footprint: - """Which constructs this program uses — walked once, then held. - - Safe to hold: a program cannot change after construction, its groups - being sealed and every node under them frozen. - """ - objective = (self.objective.expression,) if self.objective is not None else () - sides = tuple(side for c in self.constraints.values() for side in (c.lhs, c.rhs)) + """Which constructs this program uses — walked once, then held.""" return Footprint( quadratic=frozenset( - position - for position, group in (('objective', objective), ('constraint', sides)) - if any(is_quadratic(e) for e in group) + position for position, group in self._by_position() if any(is_quadratic(e) for e in group) ), variable_types=frozenset(v.variable_type for v in self.variables.values()), sos_types=frozenset(s.sos_type for s in self.sos.values()), @@ -909,12 +836,7 @@ def dimension(self, name: str) -> DimensionDeclaration: @property def lookups(self) -> tuple[tuple[str, LookupDeclaration], ...]: - """Every targeted map in the program, with the dimension it is over. - - One walk for the several shapes consumers want it in — name to target, - target to origin, the set of targets — because the nested comprehension - that produces any of them is the same walk written again. - """ + """Every lookup in the program, targeted and label-space alike, with the dimension it is over.""" return tuple((dimension, lk) for dimension, d in self.dimensions.items() for lk in d.lookups) def parameter(self, name: str) -> ParameterDeclaration: @@ -923,10 +845,41 @@ def parameter(self, name: str) -> ParameterDeclaration: def variable(self, name: str) -> VariableDeclaration: return _declared(self.variables, name, 'variable') + @cached_property + def separability(self) -> Mapping[str, Separability]: + """Every axis, to what building it a window at a time asks and what it would break. + + The locality :doc:`the ceiling ` argues in — pointwise, + bounded halo, global — asked about the axes rather than about the + operators, so a driver may know before it cuts a horizon whether every + row it builds is complete inside some window. + + **A reduction means opposite things by position**, which is the whole of + the care: in a constraint a sum over the axis ties every window to every + other, and in the objective it is additively separable, an objective + being a sum already. + + Every declared dimension has an entry, an axis nothing mentions being + trivially windowable. Walked once and held, like :attr:`footprint` and + for the same reason — a program cannot change after construction — and + answering for every axis costs what answering for one did, every + construct that ties an axis naming the axis it ties (#248). + """ + return MappingProxyType(_separabilities(self)) -# -------------------------------------------------------------------------- -# Walks, and the questions asked through them -# -------------------------------------------------------------------------- + def _built_blocks(self) -> Iterator[tuple[str, tuple[ExpressionNode, ...], Mask | None, bool]]: + """Every block that builds rows, labelled as the lowering's own messages label it. + + A named expression is not one: it is inlined where it is referenced, so + walking the constraint sides reaches it, and walking it again would + report one coupling twice. + """ + for name, block in self.constraints.items(): + yield f"constraint '{name}'", (block.lhs, block.rhs), block.where, True + for name, variable in self.variables.items(): + yield f"variable '{name}'", (variable.lower, variable.upper), variable.where, True + if self.objective is not None: + yield 'the objective', (self.objective.expression,), None, False def walk(*expressions: ExpressionNode) -> Iterator[ExpressionNode]: @@ -997,3 +950,505 @@ def divisor_parameters(*expressions: ExpressionNode) -> frozenset[str]: rows a declaration builds. """ return frozenset().union(*(parameters_of(q.divisor) for q in quotients(*expressions))) + + +# -------------------------------------------------------------------------- +# Walks, and the questions asked through them +# -------------------------------------------------------------------------- + + +def walk(*expressions: ExpressionNode) -> Iterator[ExpressionNode]: + """Every node under *expressions*, each expression itself included, parents first. + + The traversal every *question* about a program is a filter of — which names + it mentions, whether a variable stands under it, which divisions it + contains. One generator rather than that five-line recursion once per + question: how a program is traversed is one fact, so a node kind + :func:`children` learns to descend into reaches every caller at once + rather than the callers that remembered. + """ + for expression in expressions: + yield expression + yield from walk(*children(expression)) + + +def is_quadratic(expression: ExpressionNode) -> bool: + """Whether *expression* contains a product of two variable-carrying operands. + + A structural question over the program, and unrelated consumers ask it — + what a solver must support, which declarations to build last, whether this + form can be represented at all — so it is answered once here beside the + other walks rather than once per consumer in its own terms. + + Whether a degree *may be written* is the language's verdict, and this is + not a second opinion on it: by the time a program exists the question is + which shape the expression has, and the program is what is in hand to + answer it. + """ + return any( + isinstance(node, Multiply) and all(carries_variable(side) for side in (node.left, node.right)) + for node in walk(expression) + ) + + +def carries_variable(expression: ExpressionNode) -> bool: + """Whether a variable appears anywhere under *expression*.""" + return any(isinstance(node, Variable) for node in walk(expression)) + + +def parameters_of(*expressions: ExpressionNode) -> frozenset[str]: + """Every parameter named anywhere under *expressions*.""" + return frozenset(node.name for node in walk(*expressions) if isinstance(node, Parameter)) + + +def variables_of(*expressions: ExpressionNode) -> frozenset[str]: + """Every variable named anywhere under *expressions*.""" + return frozenset(node.name for node in walk(*expressions) if isinstance(node, Variable)) + + +def quotients(*expressions: ExpressionNode) -> tuple[Divide, ...]: + """Every division under *expressions*, each kept whole. + + The divisor and the numerator answer different questions and one consumer + needs them paired: a divisor is judged against the rows the declaration + builds *narrowed by the variables in its own numerator*, which the flat + :func:`divisor_parameters` cannot say. + """ + return tuple(node for node in walk(*expressions) if isinstance(node, Divide)) + + +def divisor_parameters(*expressions: ExpressionNode) -> frozenset[str]: + """Every parameter named anywhere in a divisor under *expressions*.""" + return frozenset().union(*(parameters_of(q.divisor) for q in quotients(*expressions))) + + +# --------------------------------------------------------------------------- +# the resolved where vocabulary +# --------------------------------------------------------------------------- + + +PredicateOperator = Literal['<=', '>=', '==', '!=', '<', '>'] + + +@dataclass(frozen=True) +class BooleanLiteralNode: + value: bool + + +@dataclass(frozen=True) +class ParameterDefinedNode: + """True wherever the named parameter is non-null and finite. + + ``dims`` is the parameter's own, copied off the declaration during + resolution; every leaf below that names a declaration carries its dims + (or ``over``) the same way. + """ + + name: str + dims: tuple[str, ...] + + +@dataclass(frozen=True) +class VariableDefinedNode: + """True at the coordinates where the named variable exists.""" + + name: str + dims: tuple[str, ...] + + +@dataclass(frozen=True) +class ParameterComparisonNode: + """Compare a parameter against a literal, element-wise.""" + + name: str + op: PredicateOperator + value: float | str + dims: tuple[str, ...] + + +@dataclass(frozen=True) +class DimensionComparisonNode: + """Compare a dimension's own coordinates against a literal.""" + + name: str + op: PredicateOperator + value: float | str | datetime.date + + +@dataclass(frozen=True) +class DimensionPositionNode: + """Compare where a row sits along a dimension against a position — ``position(snapshot) == 0``. + + Both sides are integers, negative counting from the end. With ``by`` the + position is counted within each group the lookup makes. + """ + + name: str + op: PredicateOperator + position: int + by: str | None = None + + +@dataclass(frozen=True) +class LookupComparisonNode: + """Compare a lookup's values against a literal — ``period_of == 2030``. + + ``over`` is the dimension the lookup maps out of. + """ + + name: str + over: str + op: PredicateOperator + value: float | str | datetime.date + + +@dataclass(frozen=True) +class LookupPairComparisonNode: + """Compare two lookups over one dimension — ``from != to``, row by row on that dimension's table.""" + + name: str + other: str + over: str + op: PredicateOperator + + +@dataclass(frozen=True) +class LookupDefinedNode: + """True where the named lookup has a value — a null says the label belongs to no group.""" + + name: str + over: str + + +@dataclass(frozen=True) +class NotNode: + operand: WhereNode + + +@dataclass(frozen=True) +class AndNode: + left: WhereNode + right: WhereNode + + +@dataclass(frozen=True) +class OrNode: + left: WhereNode + right: WhereNode + + +#: Every resolved predicate node — what a lowered mask's ``root`` is built of. +#: 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. +WhereNode = ( + BooleanLiteralNode + | DimensionPositionNode + | ParameterDefinedNode + | VariableDefinedNode + | ParameterComparisonNode + | DimensionComparisonNode + | LookupComparisonNode + | LookupPairComparisonNode + | LookupDefinedNode + | NotNode + | AndNode + | OrNode +) + +#: Every predicate resolution has typed: it names a declaration and the kind is +#: settled. Resolution passes these straight through, having nothing left to +#: decide about them. +TypedPredicateNode = ( + ParameterComparisonNode + | ParameterDefinedNode + | VariableDefinedNode + | DimensionComparisonNode + | DimensionPositionNode + | LookupComparisonNode + | LookupPairComparisonNode + | LookupDefinedNode +) + +#: The boolean connectives — the only where nodes carrying other where nodes, +#: and so the only place a walk over a predicate recurses. The grammar builds +#: these classes directly, over leaves still unresolved, so a pre-resolution +#: tree shares them — the transient impurity resolution normalizes away. +ConnectiveWhereNode = NotNode | AndNode | OrNode + + +def _atoms(where: WhereNode) -> Iterator[TypedPredicateNode]: + """Every node in *where* that reads a declaration, connectives removed. + + A boolean literal yields nothing. + + Raises: + AssertionError: An unresolved node reached the walk. + """ + if isinstance(where, NotNode): + yield from _atoms(where.operand) + elif isinstance(where, (AndNode, OrNode)): + yield from _atoms(where.left) + yield from _atoms(where.right) + elif isinstance(where, BooleanLiteralNode): + return + elif isinstance(where, TypedPredicateNode): + yield where + else: + msg = f'{type(where).__name__} reached a predicate walk unresolved.' + raise AssertionError(msg) + + +def _atom_dims(atom: TypedPredicateNode) -> frozenset[str]: + """One leaf's dims — the rule :attr:`Mask.dims` is the union of. + + A parameter or variable leaf carries its own dims off the declaration; a + comparison on a dimension is read through that dimension, and a lookup + through the dimension it maps out of — the dim it leaves, not the one it + lands in. Separate from the union because the load-time frame check + reports per leaf. Closed by ``assert_never``: a predicate node added + without a reading is a type error here, at the one place that has to grow + a branch, rather than a wrong dim set at the first model to use it. + """ + match atom: + case ParameterComparisonNode() | ParameterDefinedNode() | VariableDefinedNode(): + return frozenset(atom.dims) + case DimensionComparisonNode() | DimensionPositionNode(): + return frozenset({atom.name}) + case LookupComparisonNode() | LookupPairComparisonNode() | LookupDefinedNode(): + return frozenset({atom.over}) + case _: + assert_never(atom) + + +def _atom_names(atom: TypedPredicateNode) -> frozenset[str]: + """One leaf's declarations, its dimension apart — the rule :attr:`Mask.names_read` is the union of. + + A comparison on a dimension names no declaration — a coordinate is not + data to feed — and a lookup pair names both maps it compares. + ``assert_never``-closed for the reason :func:`_atom_dims` is: a predicate + node added without a reading is a type error at this one branch rather + than a name silently dropped at the first model to use it. + """ + match atom: + case ( + ParameterComparisonNode() + | ParameterDefinedNode() + | VariableDefinedNode() + | LookupComparisonNode() + | LookupDefinedNode() + ): + return frozenset({atom.name}) + case LookupPairComparisonNode(): + return frozenset({atom.name, atom.other}) + case DimensionComparisonNode() | DimensionPositionNode(): + return frozenset() + case _: + assert_never(atom) + + +def _conjuncts(where: WhereNode) -> tuple[WhereNode, ...]: + """The flatten rule behind :attr:`Mask.conjuncts` — the one home of the split. + + ``a AND b AND c`` gives three, and a predicate that is not an ``AND`` gives + itself. The walk stops at the first node that is not an ``AND``: the + conjuncts of ``a AND (b OR c)`` are ``a`` and ``b OR c``, and of + ``NOT (a AND b)`` the single ``NOT`` — neither an ``OR`` nor a ``NOT`` is a + claim the predicate makes on its own, so neither is split. + """ + if isinstance(where, AndNode): + return _conjuncts(where.left) + _conjuncts(where.right) + return (where,) + + +def _fold(node: WhereNode) -> WhereNode: + """*node* with every connective a literal or a double negation decides evaluated away. + + ``X AND True`` is ``X``, ``X OR True`` is every row, ``X AND False`` is + none, ``NOT True`` is ``False`` and ``NOT NOT X`` is ``X``. What survives + is a predicate over data, or the one literal the whole mask reduces to — + the invariant :class:`Mask` applies at construction, so it holds wherever + a mask is built. + """ + if isinstance(node, NotNode): + operand = _fold(node.operand) + if isinstance(operand, BooleanLiteralNode): + return BooleanLiteralNode(not operand.value) + if isinstance(operand, NotNode): + return operand.operand + return NotNode(operand) + if isinstance(node, AndNode): + left, right = _fold(node.left), _fold(node.right) + if isinstance(left, BooleanLiteralNode): + return right if left.value else left + if isinstance(right, BooleanLiteralNode): + return left if right.value else right + return AndNode(left, right) + if isinstance(node, OrNode): + left, right = _fold(node.left), _fold(node.right) + if isinstance(left, BooleanLiteralNode): + return left if left.value else right + if isinstance(right, BooleanLiteralNode): + return right if right.value else left + return OrNode(left, right) + return node + + +@dataclass(frozen=True) +class Mask: + """A resolved ``where`` and the questions the language answers about it. + + ``root`` is the predicate a consumer dispatches on with ``isinstance``; + every question below is derived from it. Construction folds, so a boolean + literal stands at the root or nowhere in it, and refuses an unresolved + tree. + + Attributes: + root: The resolved predicate the mask restricts rows by, folded. + """ + + root: WhereNode + + def __post_init__(self) -> None: + object.__setattr__(self, 'root', _fold(self.root)) + _ = self.atoms # the walk is the refusal, and runs after the fold + + @cached_property + def atoms(self) -> tuple[TypedPredicateNode, ...]: + """The mask's leaves, connectives removed — the one walk the other questions read. + + Held rather than re-walked: construction takes this walk anyway, to + refuse an unresolved tree, and a mask cannot change afterwards. + """ + return tuple(_atoms(self.root)) + + @property + def conjuncts(self) -> tuple[WhereNode, ...]: + """The predicates the mask joins with ``AND`` — its ``AND`` spine flattened, stopping at an ``OR`` or a ``NOT``.""" + return _conjuncts(self.root) + + @property + def names_read(self) -> frozenset[str]: + """The parameters, lookups and variables the mask names.""" + return frozenset(name for atom in self.atoms for name in _atom_names(atom)) + + @property + def dims(self) -> frozenset[str]: + """The dims the mask is read at — the union of what each leaf carries. + + Empty for a mask over nothing but literals. Read off the leaves, which + resolution stamped with their declarations' dims, so a predicate built + from resolved pieces answers exactly as a declaration's own does. + """ + return frozenset(dim for atom in self.atoms for dim in _atom_dims(atom)) + + def __invert__(self) -> Mask: + """The mask admitting exactly the rows this one refuses — construction folds a double negation or a literal flip.""" + return Mask(NotNode(self.root)) + + def __and__(self, other: Mask) -> Mask: + """Both masks at once — construction absorbs a literal side rather than burying it.""" + return Mask(AndNode(self.root, other.root)) + + def __or__(self, other: Mask) -> Mask: + """Either mask — construction absorbs a literal side rather than burying it.""" + return Mask(OrNode(self.root, other.root)) + + +def _separabilities(program: Program) -> dict[str, Separability]: + """Every axis's verdict, in one walk. + + One traversal rather than one per axis, because every construct that ties an + axis together names the axis it ties: asking each node *which* dimension it + is about answers for all of them at what answering for one cost. + + ``reductions_couple`` is the position a block stands in rather than anything + about the block — a sum over the axis couples a constraint row to the whole + horizon and leaves an objective additively separable. A translation reads + behind for a positive offset and ahead for a negative one, and a trailing + window behind by its width less one. Each coupling carries the one + modelling change that would lift it, after the dash. + """ + behind = dict.fromkeys(program.dimensions, 0) + ahead = dict.fromkeys(program.dimensions, 0) + reasons: dict[str, dict[str, dict[str, list[str]]]] = { + kind: {dimension: {} for dimension in program.dimensions} for kind in ('coupled', 'undecided', 'restarts') + } + + def report(kind: str, dimension: str, label: str, reason: str) -> None: + reasons[kind][dimension].setdefault(label, []).append(reason) + + for label, nodes, mask, reductions_couple in program._built_blocks(): + masks: list[Mask | None] = [mask] + for node in walk(*nodes): + 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: + report( + 'coupled', + dimension, + label, + f'sums over {dimension} — a rolling sum_back(within=n) windows, a total over the horizon does not', + ) + elif isinstance(node, GroupSum): + report( + 'coupled', + node.over, + label, + f'groups {node.over} into {", ".join(node.into)} — window that dimension instead, or cut only at the group edges', + ) + elif isinstance(node, At): + for dimension in node.into: + for lookup in node.coordinate: + report('undecided', dimension, label, lookup) + elif isinstance(node, (Translate, Window)): + dimension = node.dimension + if node.wrap: + report( + 'coupled', + dimension, + label, + f'wraps around {dimension}, so its first row reads its last — an opening-state seed at ' + f'position({dimension}) == 0 is what a rolling horizon replaces the wrap with', + ) + continue + if node.partition is not None: + report('undecided', dimension, label, node.partition) + reach = node.offset if isinstance(node, Translate) else node.width + if isinstance(reach, str): + report('undecided', dimension, label, reach) + elif isinstance(node, Window): + behind[dimension] = max(behind[dimension], reach - 1) + elif reach > 0: + behind[dimension] = max(behind[dimension], reach) + else: + ahead[dimension] = max(ahead[dimension], -reach) + for candidate in masks: + for atom in candidate.atoms if candidate is not None else (): + if isinstance(atom, DimensionPositionNode): + report('restarts', atom.name, label, f'counts a position along {atom.name}') + + for name, block in program.sos.items(): + report( + 'coupled', + block.over, + f"set '{name}'", + f'is a set over {block.over}, which a window would cut — only a window holding every whole set keeps it', + ) + + def joined(kind: str, dimension: str) -> dict[str, str]: + return {label: ', '.join(dict.fromkeys(found)) for label, found in reasons[kind][dimension].items()} + + return { + dimension: Separability( + dimension=dimension, + behind=behind[dimension], + ahead=ahead[dimension], + coupled=joined('coupled', dimension), + undecided=joined('undecided', dimension), + restarts=joined('restarts', dimension), + ) + for dimension in program.dimensions + } diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index e6d883a9..aee11d85 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -5,18 +5,25 @@ """Name resolution — the pass that makes the core AST fully typed. Parsers emit unresolved names; this module rewrites each into the typed node -its kind asks for, so the AST reaching a consumer holds none. Done once here, -every consumer scopes identically by construction. The rules live in the -language reference; the namespace is flat, and macro formals are the one scope. +its kind asks for, so the AST reaching a consumer holds none. The rules live in +the language reference. """ from __future__ import annotations import datetime import re -from typing import TYPE_CHECKING, assert_never +from dataclasses import dataclass, replace +from typing import TYPE_CHECKING, Literal, assert_never, cast -from math_spec.errors import LanguageError +from math_spec._where_parser import ( + UnresolvedComparisonNode, + UnresolvedNameNode, + UnresolvedPositionNode, + UnresolvedWhereNode, + parse_where, +) +from math_spec.errors import LanguageError, did_you_mean from math_spec.expansion import parse_and_expand from math_spec.expression_parser import ( ArithmeticNode, @@ -48,7 +55,7 @@ edge_error, unknown_operator_message, ) -from math_spec.where_parser import ( +from math_spec.program import ( AndNode, BooleanLiteralNode, DimensionComparisonNode, @@ -56,33 +63,35 @@ LookupComparisonNode, LookupDefinedNode, LookupPairComparisonNode, + Mask, NotNode, OrNode, ParameterComparisonNode, ParameterDefinedNode, TypedPredicateNode, - UnresolvedComparisonNode, - UnresolvedNameNode, - UnresolvedPositionNode, VariableDefinedNode, WhereNode, - parse_where, ) if TYPE_CHECKING: from collections.abc import Iterable, Mapping - from math_spec.model import Spec + from math_spec.model import DeclaredDtype, Spec + + +#: What a name a file may write turns out to be. Answered by +#: :meth:`Namespace.kind`, so a pass reading a name switches over this rather +#: than over the stores it would otherwise have to try in order. +DeclarationKind = Literal['variable', 'parameter', 'dimension', 'lookup'] class Namespace: """The declared names of one schema, by kind. - Flat by construction: :meth:`kind` is a single lookup, not an ordered - walk through several stores. + A name has one kind: model.py refuses one declared under two sections. """ - __slots__ = ('dimensions', 'dtypes', 'lookups', 'parameters', 'variables') + __slots__ = ('dimensions', 'dtypes', 'leaf_dims', 'lookups', 'parameters', 'variables') def __init__( self, @@ -90,17 +99,22 @@ def __init__( parameters: Iterable[str], dimensions: Iterable[str], lookups: Mapping[str, tuple[str, str | None]], - dtypes: Mapping[str, str], + dtypes: Mapping[str, DeclaredDtype], + leaf_dims: Mapping[str, tuple[str, ...]], ) -> None: self.variables = frozenset(variables) self.parameters = frozenset(parameters) self.dimensions = frozenset(dimensions) #: name -> declared dtype, for dimensions, parameters and lookups alike; #: what a where comparison checks its literal against. - self.dtypes: dict[str, str] = dict(dtypes) + self.dtypes: dict[str, DeclaredDtype] = dict(dtypes) #: lookup name -> ``(over, into)``; ``into`` is ``None`` for a label #: space, which owns its values. self.lookups: dict[str, tuple[str, str | None]] = dict(lookups) + #: parameter or variable name -> the dims it is read through — + #: parameters by their ``dims``, variables by their frame. Stamped onto + #: each leaf a where names, the way a lookup leaf carries ``over``. + self.leaf_dims: dict[str, tuple[str, ...]] = dict(leaf_dims) def groupable(self) -> dict[str, str]: """The lookups a ``by=`` may name: name -> the dimension it maps into. @@ -128,10 +142,14 @@ def of(cls, schema: Spec) -> Namespace: **{n: schema.dimensions[lk.into].dtype for n, lk in schema.lookups.items() if lk.into is not None}, **{n: lk.dtype for n, lk in schema.lookups.items() if lk.dtype is not None}, }, + { + **{p: tuple(pd.dims) for p, pd in schema.parameters.items()}, + **{v: tuple(vd.foreach) for v, vd in schema.variables.items()}, + }, ) - def kind(self, name: str) -> str | None: - """``'variable'`` | ``'parameter'`` | ``'dimension'`` | ``'lookup'`` | ``None``.""" + def kind(self, name: str) -> DeclarationKind | None: + """What *name* was declared as, or ``None`` where the file declares it nowhere.""" if name in self.variables: return 'variable' if name in self.parameters: @@ -150,8 +168,17 @@ def into_of(self, lookup: str) -> str | None: """The dimension *lookup*'s values are labels of, ``None`` for a label space.""" return self.lookups[lookup][1] - def _unknown(self, name: str, context: str, *, allow_dims: bool) -> str: - shown = ( + def unknown(self, name: str, context: str, *, allow_dims: bool, formals: Iterable[str] = ()) -> str: + """The refusal for a *name* declared nowhere, listing what it could have been. + + Args: + name: The name the file wrote. + context: The declaration it was found in. + allow_dims: Whether a dimension would have been accepted there. + formals: A macro's formals, listed first when there are any. + """ + shown: list[tuple[str, Iterable[str]]] = [('Formals', formals)] if formals else [] + shown += ( [('Parameters', self.parameters), ('Dimensions', self.dimensions)] if allow_dims else [('Variables', self.variables), ('Parameters', self.parameters)] @@ -168,10 +195,6 @@ def _unknown(self, name: str, context: str, *, allow_dims: bool) -> str: def expression_of(text: str, schema: Spec, ns: Namespace, context: str) -> ExpressionNode: """Parse, expand and resolve *text* — the only way a consumer gets an AST. - ``validation.py`` runs the same path at load time, so a consumer calling - this gets a *typed* tree off a result already known to be clean, without - duplicating the pass. - Raises: LanguageError: Listing every problem the text has. """ @@ -183,58 +206,23 @@ def expression_of(text: str, schema: Spec, ns: Namespace, context: str) -> Expre return resolved -def where_of(text: str | None, ns: Namespace, context: str, self_variable: str | None = None) -> WhereNode | None: - """Parse, resolve and fold a where string — ``None`` for no mask, however the file spelled it. +def where_of(text: str | None, ns: Namespace, context: str, self_variable: str | None = None) -> Mask | None: + """Parse and resolve a where string into the :class:`~math_spec.program.Mask` a declaration carries. - Every constant the connectives decide is folded away, so a mask that - admits every row arrives as ``None`` and one that admits none as - ``BooleanLiteralNode(False)``; that node stands at the root of a mask or - nowhere in it, and every reader — a program, a typeset page — gets the - same predicate. + ``None`` for no mask, however the file spelled it: a mask that admits every + row is dropped, and one that admits none arrives as a mask over + ``BooleanLiteralNode(False)``. Raises: LanguageError: Listing every problem the predicate has. """ - if text is None: - return None errors: list[str] = [] - resolved = resolve_where(parse_where(text), ns, context, errors, self_variable) + resolved = resolve_where_text(text, ns, context, errors, self_variable) if errors: raise LanguageError('\n'.join(errors)) - assert resolved is not None - folded = _fold(resolved) - if isinstance(folded, BooleanLiteralNode) and folded.value: + if resolved is None or (isinstance(resolved, BooleanLiteralNode) and resolved.value): return None - return folded - - -def _fold(node: WhereNode) -> WhereNode: - """*node* with every literal a connective decides evaluated away. - - ``X AND True`` is ``X``, ``X OR True`` is every row, ``X AND False`` is - none, and ``NOT True`` is ``False``. What survives is a predicate over - data, or the one literal the whole mask reduces to. - """ - if isinstance(node, NotNode): - operand = _fold(node.operand) - if isinstance(operand, BooleanLiteralNode): - return BooleanLiteralNode(not operand.value) - return NotNode(operand) - if isinstance(node, AndNode): - left, right = _fold(node.left), _fold(node.right) - if isinstance(left, BooleanLiteralNode): - return right if left.value else left - if isinstance(right, BooleanLiteralNode): - return left if right.value else right - return AndNode(left, right) - if isinstance(node, OrNode): - left, right = _fold(node.left), _fold(node.right) - if isinstance(left, BooleanLiteralNode): - return left if left.value else right - if isinstance(right, BooleanLiteralNode): - return right if right.value else left - return OrNode(left, right) - return node + return Mask(resolved) # --------------------------------------------------------------------------- @@ -248,11 +236,7 @@ def resolve_expression( context: str, errors: list[str], ) -> ExpressionNode | None: - """Rewrite every ``NameNode`` under *node* to a typed node. - - Operator *call shapes* are checked here too (``operators.call_shape_error``). - Arity is a language rule, and this is the pass every consumer goes through, - so no consumer has to state a signature a second time. + """Rewrite every ``NameNode`` under *node* to a typed node, checking operator call shapes on the way. Returns: The typed tree, or ``None`` once anything failed — appending to @@ -260,51 +244,125 @@ def resolve_expression( whole schema reports them together. """ before = len(errors) - if isinstance(node, ComparisonNode): - resolved: ExpressionNode = ComparisonNode( - node.op, - _resolve_arith(node.left, ns, context, errors), - _resolve_arith(node.right, ns, context, errors), - ) - else: - resolved = _resolve_arith(node, ns, context, errors) + resolved = _Resolver(ns, context, errors).expression(node) return None if len(errors) > before else resolved -def _resolve_arith( - node: ArithmeticNode, +def resolve_where( + node: WhereNode | UnresolvedWhereNode, ns: Namespace, context: str, errors: list[str], - *, - amount: bool = False, -) -> ArithmeticNode: - """The recursive worker under :func:`resolve_expression`. - - *amount* marks an ``offset=``/``within=`` 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. + self_variable: str | None = None, +) -> WhereNode | None: + """Rewrite a parsed where AST into typed predicates, folded as :class:`~math_spec.program.Mask` folds. + + Returns: + The typed tree — a mask admitting every row or none comes back as the + one ``BooleanLiteralNode`` — or ``None`` once anything failed, with the + problems appended to *errors*. """ - if isinstance(node, NumberNode): - return node + before = len(errors) + resolved = _Resolver(ns, context, errors, self_variable).where(node) + return None if len(errors) > before else Mask(cast('WhereNode', resolved)).root - if isinstance(node, VariableNode | ParameterNode | KwargNode): - return node - if isinstance(node, NameNode): - match ns.kind(node.name): +def resolve_where_text( + text: str | None, + ns: Namespace, + context: str, + errors: list[str], + self_variable: str | None = None, +) -> WhereNode | None: + """Parse and resolve one where string as :func:`resolve_where` does, a parse failure appended to *errors*. + + Returns: + ``None`` where there is no mask to read, and where reading it failed. + """ + if text is None: + return None + try: + node = parse_where(text) + except ValueError as e: + errors.append(f'{context}: {e}') + return None + return resolve_where(node, ns, context, errors, self_variable) + + +@dataclass(frozen=True) +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. + """ + + ns: Namespace + context: str + errors: list[str] + self_variable: str | None = None + + # -- expressions ------------------------------------------------------- + + def expression(self, node: ExpressionNode) -> ExpressionNode: + """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. + + *amount* marks an ``offset=``/``within=`` 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. + """ + if isinstance(node, NumberNode | VariableNode | ParameterNode | KwargNode): + return node + if isinstance(node, NameNode): + return self._name(node, amount=amount) + if isinstance(node, UnaryOperatorNode): + return UnaryOperatorNode(node.op, self._arith(node.operand)) + if isinstance(node, BinaryOperatorNode): + return BinaryOperatorNode(node.op, self._arith(node.left), self._arith(node.right)) + if isinstance(node, FunctionCallNode): + return self._call(node) + if isinstance(node, KeywordNode): + self.errors.append( + f'{self.context}: {node.value!r} is a quoted keyword, which is only legal as a ' + f"operator kwarg value such as shift(..., edge='wrap'). In an expression, quote " + f'nothing — names resolve and numbers are written bare.' + ) + return node + if isinstance(node, NameListNode): + self.errors.append( + f'{self.context}: {node.shown} 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 + if isinstance(node, CasesNode): + return self._cases(node) + assert_never(node) + + def _name(self, node: NameNode, *, amount: bool) -> ArithmeticNode: + """A bare name as the variable or parameter it declares; a dimension or lookup is not a value.""" + match self.ns.kind(node.name): case 'variable': return VariableNode(node.name) case 'parameter': - dtype = ns.dtypes.get(node.name) + dtype = self.ns.dtypes.get(node.name) if not amount and dtype is not None and dtype not in NUMERIC_DTYPES: - errors.append(_not_a_number(node.name, dtype, context)) + self.errors.append(_not_a_number(node.name, dtype, self.context)) return node return ParameterNode(node.name) case 'dimension': - errors.append( - f"{context}: '{node.name}' is a dimension, and a dimension is " + self.errors.append( + f"{self.context}: '{node.name}' is a dimension, and a dimension is " f'not a value in an expression. Dimensions appear in ' f"'foreach:', in operator arguments (sum(x, over={node.name})), " f'and in where-comparisons — to use its coordinates as data, ' @@ -312,73 +370,353 @@ def _resolve_arith( ) return node case 'lookup': - errors.append( - f"{context}: '{node.name}' is a lookup, and a lookup is structure " + self.errors.append( + f"{self.context}: '{node.name}' is a lookup, and a lookup is structure " f'rather than data, so it is not a value in an expression. A lookup ' 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 case _: - errors.append(ns._unknown(node.name, context, allow_dims=False)) + self.errors.append(self.ns.unknown(node.name, self.context, allow_dims=False)) return node - if isinstance(node, UnaryOperatorNode): - return UnaryOperatorNode(node.op, _resolve_arith(node.operand, ns, context, errors)) - - if isinstance(node, BinaryOperatorNode): - return BinaryOperatorNode( - node.op, - _resolve_arith(node.left, ns, context, errors), - _resolve_arith(node.right, ns, context, errors), - ) - - if isinstance(node, FunctionCallNode): + def _call(self, node: FunctionCallNode) -> ArithmeticNode: + """An operator call: its shape checked, and each kwarg typed by the kind the operator declares for it.""" if node.name not in BUILTINS: - errors.append(f'{context}: {unknown_operator_message(node.name)}') + self.errors.append(f'{self.context}: {unknown_operator_message(node.name)}') return node builtin = BUILTINS[node.name] shape_error = call_shape_error(node.name, len(node.args), node.kwargs) if shape_error is not None: - errors.append(f'{context}: {shape_error}') - args = [_resolve_arith(a, ns, context, errors) for a in node.args] + self.errors.append(f'{self.context}: {shape_error}') + args = tuple(self._arith(a) for a in node.args) kwargs: dict[str, ArithmeticNode] = {} for key, value in node.kwargs.items(): - if key in builtin.edge_kwargs: - kwargs[key] = _resolve_edge(value, context, node.name, errors) - elif key in builtin.dimension_kwargs: - kwargs[key] = _resolve_dim_ref(value, ns, context, node.name, key, errors) - elif key in builtin.lookup_kwargs: - kwargs[key] = _resolve_lookup_ref(value, ns, context, node.name, key, errors) - else: - kwargs[key] = _resolve_amount(value, ns, context, node.name, key, errors) + match builtin.kind_of(key): + case 'edge': + kwargs[key] = self._edge(value, node.name) + case 'dimension': + kwargs[key] = self._dim_ref(value, node.name, key) + case 'lookup': + kwargs[key] = self._lookup_ref(value, node.name, key) + case 'value': + kwargs[key] = self._amount(value, node.name, key) return FunctionCallNode(node.name, args, kwargs) - if isinstance(node, KeywordNode): - errors.append( - f'{context}: {node.value!r} is a quoted keyword, which is only legal as a ' - f"operator kwarg value such as shift(..., edge='wrap'). In an expression, quote " - f'nothing — names resolve and numbers are written bare.' + 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 ``within=``: a number or a parameter name, never an expression. + + Closed so that :func:`math_spec.dimensions._check_named_amount` sees every + parameter an amount carries. + """ + if (literal := _literal(value)) is not None: + return literal + if not isinstance(_without_sign(value), 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) + + 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 isinstance(value, KeywordNode): + if value.value == EDGE_WRAP: + return EdgeNode() + self.errors.append(f'{self.context}: {edge_error(operator, repr(value.value))}') + return value + 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 + self.errors.append(f'{self.context}: {edge_error(operator, value.name)}') + return value + 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 + + def _dim_ref(self, value: ArithmeticNode, operator: str, key: str) -> ArithmeticNode: + """An operator kwarg whose *value* must name a declared dimension.""" + if not isinstance(value, NameNode): + self.errors.append(f'{self.context}: {operator}({key}=...) must name a dimension.') + return value + 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) + + def _lookup_ref(self, value: ArithmeticNode, operator: str, key: str) -> ArithmeticNode: + """An operator kwarg whose *value* must name groupable lookups. + + A lookup carries its own dimensions, so nothing else in the call is + consulted: the names alone decide both the dim the operator consumes and + the ones it produces. A bracketed list is one grouping through several + maps at once rather than a composition of groupings, so its members must + share the dim they are over and must not target the same dim twice. + """ + if isinstance(value, NameListNode): + names = value.names + elif isinstance(value, NameNode): + names = (value.name,) + else: + self.errors.append(f'{self.context}: {operator}({key}=...) must name a lookup.') + return value + + ns = self.ns + groupable = ns.groupable() + named = [self._ungroupable(name, groupable, operator, key) for name in names] + if any(problem is not None for problem in named): + self.errors.extend(problem for problem in named if problem is not None) + return value + + over = {ns.over_of(name) for name in names} + if len(over) > 1: + self.errors.append( + f'{self.context}: {operator}({key}={shown(names)}) groups through lookups over ' + f'different dimensions ({", ".join(f"{n} over {ns.over_of(n)}" for n in names)}). ' + f'One grouping consumes one dimension, so every lookup in the list must be ' + f'over the same one — group through them in turn instead, one call each.' + ) + return value + + targets = tuple(groupable[name] for name in names) + repeated = sorted({t for t in targets if targets.count(t) > 1}) + if repeated: + self.errors.append( + f'{self.context}: {operator}({key}={shown(names)}) targets {repeated} more than once. ' + f'Each lookup in the list produces its own dimension, so two that land on the ' + f'same one would need it twice — drop one, or group into a dimension of its own.' + ) + return value + + return LookupNode(names, dimension=next(iter(over)), into=targets) + + def _ungroupable(self, name: str, groupable: Mapping[str, str], operator: str, key: str) -> str | None: + """Why *name* is not a groupable lookup; ``None`` where it is one.""" + ns, context = self.ns, self.context + if name in ns.lookups and name not in groupable: + over = ns.over_of(name) + return ( + f'{context}: {operator}({key}={name}): ' + f"'{name}' is a label space over '{over}', not a groupable lookup — " + f'it targets no dimension for the terms to land on. To group into it, ' + f'declare the axis and target it under a name of its own:\n' + f' dimensions:\n' + f' {name}: {{...}}\n' + f' lookups:\n' + f' {name}_of: {{over: {over}, into: {name}}}' + ) + if name in groupable: + return None + if name in ns.dimensions: + into_here = sorted(n for n, into in groupable.items() if into == name) + hint = f" Lookups into '{name}': {into_here}" if into_here else f" No lookup maps into '{name}'." + return ( + f"{context}: {operator}({key}={name}): '{name}' is a dimension, and " + f'{key}= takes a lookup — the named map out of a dimension.\n{hint}' + ) + return ( + f'{context}: {operator}({key}={name}) does not name a lookup. ' + f'{did_you_mean(name, groupable, label="Lookups")}\n' + f"Declare it under 'lookups:' — {name}: {{over: , into: }}.' ) + + # -- where strings ----------------------------------------------------- + + def where(self, node: WhereNode | UnresolvedWhereNode) -> WhereNode | UnresolvedWhereNode: + """One predicate node typed, or returned unresolved with its refusal appended.""" + if isinstance(node, BooleanLiteralNode | TypedPredicateNode): + return node + if isinstance(node, UnresolvedNameNode): + return self._where_name(node) + if isinstance(node, UnresolvedPositionNode): + return self._position(node) + if isinstance(node, UnresolvedComparisonNode): + return self._comparison(node) + if isinstance(node, NotNode): + return NotNode(self._child(node.operand)) + if isinstance(node, AndNode): + return AndNode(self._child(node.left), self._child(node.right)) + if isinstance(node, OrNode): + return OrNode(self._child(node.left), self._child(node.right)) + assert_never(node) + + def _child(self, node: WhereNode | UnresolvedWhereNode) -> WhereNode: + """A connective's child, typed as resolved: an unresolved one survives only with its refusal appended.""" + return cast('WhereNode', self.where(node)) + + def _where_name(self, node: UnresolvedNameNode) -> WhereNode | UnresolvedWhereNode: + """A bare name: a parameter's or lookup's definedness, or a variable's existence.""" + ns, context = self.ns, self.context + kind = ns.kind(node.name) + if kind is None: + self.errors.append(ns.unknown(node.name, context, allow_dims=True)) + return node + match kind: + case 'parameter': + return ParameterDefinedNode(node.name, ns.leaf_dims[node.name]) + case 'dimension': + self.errors.append( + f"{context}: '{node.name}' is a dimension, and a bare dimension " + f'name is true at every coordinate — the mask has no effect. ' + f'Remove it, or compare it: where: "{node.name} > 0".' + ) + case 'lookup': + return LookupDefinedNode(node.name, ns.over_of(node.name)) + case 'variable': + if node.name == self.self_variable: + self.errors.append( + f"{context}: variable '{node.name}' asks whether it exists in its own " + f'where, which nothing can answer — the mask is what decides where it ' + f'exists. Test a parameter, or another variable declared before it.' + ) + else: + return VariableDefinedNode(node.name, ns.leaf_dims[node.name]) return node - if isinstance(node, NameListNode): - errors.append( - f'{context}: {node.shown} 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.' - ) + def _position(self, node: UnresolvedPositionNode) -> DimensionPositionNode | UnresolvedPositionNode: + """``position(dim[, by=lookup]) i``: the name a dimension, ``by=`` a lookup over it.""" + ns, context = self.ns, self.context + if node.dimension not in ns.dimensions: + self.errors.append( + f"{context}: position() counts along a dimension's coordinates, and " + f"'{node.dimension}' is {_declared_as(ns, node.dimension)}. " + f'{did_you_mean(node.dimension, ns.dimensions, label="Dimensions")}' + ) + return node + if node.by is None: + return DimensionPositionNode(node.dimension, node.op, node.position, node.by) + call = f'position({node.dimension}, by={node.by})' + if ns.kind(node.by) != 'lookup': + self.errors.append( + f"{context}: '{call}' groups by '{node.by}', which is {_declared_as(ns, node.by)}. " + f'``by=`` takes a lookup over that dimension — either kind, since counting ' + f'inside a group lands no terms, unlike sum(by=) and at(by=). ' + f'{did_you_mean(node.by, ns.lookups, label="Lookups")}' + ) + return node + over = ns.over_of(node.by) + if over != node.dimension: + self.errors.append( + f"{context}: '{call}' counts positions along '{node.dimension}' but groups by a " + f"lookup over '{over}'. No row of '{node.dimension}' carries it, so there is no " + f"position within a group to name — group by a lookup over '{node.dimension}'." + ) + return node + return DimensionPositionNode(node.dimension, node.op, node.position, node.by) + + def _comparison(self, node: UnresolvedComparisonNode) -> WhereNode | UnresolvedWhereNode: + """``name literal``, or the one structural form ``lookup lookup``.""" + ns, context = self.ns, self.context + value = node.value + if not node.quoted and isinstance(value, str) and (rhs_kind := ns.kind(value)) is not None: + if rhs_kind == 'lookup' and ns.kind(node.name) == 'lookup': + if (refusal := _lookup_pair_error(context, node, value, ns)) is not None: + self.errors.append(refusal) + return node + return LookupPairComparisonNode(node.name, value, ns.over_of(node.name), node.op) + self.errors.append(_declared_rhs_error(context, node, value, rhs_kind)) + return node + + kind = ns.kind(node.name) + if kind is None: + self.errors.append(ns.unknown(node.name, context, allow_dims=True)) + return node + if kind in ('parameter', 'dimension', 'lookup'): + typed = self._typed_literal(node, ns.dtypes[node.name]) + if typed is None: + return node + value = typed + + match kind: + case 'parameter': + assert not isinstance(value, datetime.date) + return ParameterComparisonNode(node.name, node.op, value, ns.leaf_dims[node.name]) + case 'dimension': + return DimensionComparisonNode(node.name, node.op, value) + case 'lookup': + return LookupComparisonNode(node.name, ns.over_of(node.name), node.op, value) + case 'variable': + self.errors.append( + f"{context}: where references variable '{node.name}'. A where " + f'mask is built before variables exist — it may test parameters ' + f'and dimension coordinates only.' + ) return node - if isinstance(node, CasesNode): - 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, ns, arm_context, errors) - arms.append(CaseArm(arm.label, when, _resolve_arith(arm.value, ns, arm_context, errors))) - return CasesNode(node.name, tuple(arms)) + def _typed_literal( + self, node: UnresolvedComparisonNode, dtype: DeclaredDtype + ) -> float | str | datetime.date | None: + """The comparison's literal, checked against the declared dtype. - assert_never(node) + Getting it wrong is silent: polars reads a datetime column against an + integer as an epoch offset, so ``snapshot > 0`` drops every coordinate + before 1970 without a word (#460). Returns ``None`` once it has recorded + an error, so the caller leaves the node unresolved. + """ + context = self.context + value = node.value + text = isinstance(value, str) + + if dtype == 'datetime': + if not text: + self.errors.append( + f"{context}: '{node.name}' is a datetime dimension, so comparing it to " + f'{value!r} compares against the epoch — {node.name} > 0 means "after ' + f'1970-01-01", not what it looks like. Quote an ISO date instead: ' + f"{node.name} {node.op} '2030-01-01'." + ) + return None + try: + return ( + datetime.datetime.fromisoformat(value) + if _HAS_TIME.search(value) + else datetime.date.fromisoformat(value) + ) + except ValueError: + self.errors.append( + f"{context}: '{node.name}' is a datetime dimension and {value!r} is not an " + f"ISO date. Write '2030-01-01' or '2030-01-01T06:00'." + ) + return None + + if dtype == 'str' and not text: + self.errors.append( + f"{context}: '{node.name}' has dtype 'str', so comparing it to the number " + f'{value!r} matches no label. Quote it if it is one: {node.name} {node.op} ' + f"'{value:g}'." + ) + return None + if dtype in ('int', 'float', 'bool') and text: + self.errors.append( + f"{context}: '{node.name}' has dtype '{dtype}', so comparing it to the string " + f'{value!r} matches nothing. Drop the quotes if it is a number.' + ) + return None + return value + + +#: An ISO literal carrying a time-of-day, which decides date vs datetime. +_HAS_TIME = re.compile(r'[T ]\d') def _not_a_number(name: str, dtype: str, context: str) -> str: @@ -401,16 +739,21 @@ def _not_a_number(name: str, dtype: str, context: str) -> str: ) -def _undeclared_dim(context: str, operator: str, shown: str, name: str, ns: Namespace) -> str: +def _undeclared_dim(context: str, operator: str, call: str, name: str, ns: Namespace) -> str: return ( - f'{context}: {operator}({shown}) does not name a declared dimension.\n' - f' Dimensions: {sorted(ns.dimensions)}\n' + f'{context}: {operator}({call}) does not name a declared dimension. ' + f'{did_you_mean(name, ns.dimensions, label="Dimensions")}\n' f"Declare '{name}' under 'dimensions:', or fix the typo — an unknown " f'dimension makes {operator}() a silent no-op rather than an error.' ) -def _unsigned(value: ArithmeticNode) -> ArithmeticNode: +def _declared_as(ns: Namespace, name: str) -> str: + kind = ns.kind(name) + return f'a {kind}' if kind else 'not declared' + + +def _without_sign(value: ArithmeticNode) -> ArithmeticNode: """*value* under its sign, if it carries one.""" return value.operand if isinstance(value, UnaryOperatorNode) else value @@ -429,278 +772,30 @@ def _literal(value: ArithmeticNode) -> NumberNode | None: return None -def _resolve_amount( - value: ArithmeticNode, ns: Namespace, context: str, operator: str, key: str, errors: list[str] -) -> ArithmeticNode: - """Resolve ``offset=`` or ``within=``: a number or a parameter name, never an expression. - - Closed so that :func:`math_spec.dimensions._check_named_amount` sees every - parameter an amount carries. - """ - if (literal := _literal(value)) is not None: - return literal - if not isinstance(_unsigned(value), NameNode): - errors.append( - f'{context}: {operator}({key}=) takes a number or the name of an integer parameter. ' - f'Precompute it as a parameter.' - ) - return value - return _resolve_arith(value, ns, context, errors, amount=True) - - -def _resolve_edge( - value: ArithmeticNode, - context: str, - operator: str, - errors: list[str], -) -> ArithmeticNode: - """Resolve ``edge=``: the closed keyword ``wrap``, or a number to contribute. - - Takes no namespace: the keyword set is closed, so a name here is a typo - rather than a lookup. - """ - if isinstance(value, EdgeNode): - return value - if isinstance(value, KeywordNode): - if value.value == EDGE_WRAP: - return EdgeNode(EDGE_WRAP) - errors.append(f'{context}: {edge_error(operator, repr(value.value))}') - return value - if isinstance(value, NameNode): - if value.name == EDGE_WRAP: - errors.append( - f'{context}: {operator}(edge={EDGE_WRAP}) is a bare name where a keyword belongs. ' - f"Write edge='{EDGE_WRAP}', quoted." - ) - return value - errors.append(f'{context}: {edge_error(operator, value.name)}') - return value - if (literal := _literal(value)) is None: - errors.append( - f"{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 - - -def _resolve_dim_ref( - value: ArithmeticNode, - ns: Namespace, - context: str, - operator: str, - key: str, - errors: list[str], -) -> ArithmeticNode: - """Resolve an operator kwarg whose *value* must name a declared dimension.""" - if isinstance(value, DimensionNode): - return value - if not isinstance(value, NameNode): - errors.append(f'{context}: {operator}({key}=...) must name a dimension.') - return value - if value.name not in ns.dimensions: - errors.append(_undeclared_dim(context, operator, f'{key}={value.name}', value.name, ns)) - return value - return DimensionNode(value.name) - - -def _resolve_lookup_ref( - value: ArithmeticNode, - ns: Namespace, - context: str, - operator: str, - key: str, - errors: list[str], -) -> ArithmeticNode: - """Resolve an operator kwarg whose *value* must name groupable lookups. - - A lookup carries its own dimensions, so nothing else in the call is - consulted: the names alone decide both the dim the operator consumes and - the ones it produces. A bracketed list is one grouping through several - maps at once rather than a composition of groupings, so its members must - share the dim they are over and must not target the same dim twice — - both checked here, where the declarations are still in hand. - """ - if isinstance(value, LookupNode): - return value - if isinstance(value, NameListNode): - names = value.names - elif isinstance(value, NameNode): - names = (value.name,) - else: - errors.append(f'{context}: {operator}({key}=...) must name a lookup.') - return value - - groupable = ns.groupable() - named = [_ungroupable(name, ns, groupable, context, operator, key) for name in names] - if any(problem is not None for problem in named): - errors.extend(problem for problem in named if problem is not None) - return value - - over = {ns.over_of(name) for name in names} - if len(over) > 1: - errors.append( - f'{context}: {operator}({key}={shown(names)}) groups through lookups over ' - f'different dimensions ({", ".join(f"{n} over {ns.over_of(n)}" for n in names)}). ' - f'One grouping consumes one dimension, so every lookup in the list must be ' - f'over the same one — group through them in turn instead, one call each.' - ) - return value - - targets = tuple(groupable[name] for name in names) - repeated = sorted({t for t in targets if targets.count(t) > 1}) - if repeated: - errors.append( - f'{context}: {operator}({key}={shown(names)}) targets {repeated} more than once. ' - f'Each lookup in the list produces its own dimension, so two that land on the ' - f'same one would need it twice — drop one, or group into a dimension of its own.' - ) - return value - - return LookupNode(names, dimension=next(iter(over)), into=targets) - - -def _ungroupable( - name: str, - ns: Namespace, - groupable: Mapping[str, str], - context: str, - operator: str, - key: str, -) -> str | None: - """Why *name* is not a groupable lookup; ``None`` where it is one.""" - if name in ns.lookups and name not in groupable: - over = ns.over_of(name) - return ( - f'{context}: {operator}({key}={name}): ' - f"'{name}' is a label space over '{over}', not a groupable lookup — " - f'it targets no dimension for the terms to land on. To group into it, ' - f'declare the axis and target it under a name of its own:\n' - f' dimensions:\n' - f' {name}: {{...}}\n' - f' lookups:\n' - f' {name}_of: {{over: {over}, into: {name}}}' - ) - if name in groupable: - return None - if name in ns.dimensions: - into_here = sorted(n for n, into in groupable.items() if into == name) - hint = f" Lookups into '{name}': {into_here}" if into_here else f" No lookup maps into '{name}'." - return ( - f"{context}: {operator}({key}={name}): '{name}' is a dimension, and " - f'{key}= takes a lookup — the named map out of a dimension.\n{hint}' - ) - listing = f' Lookups: {sorted(groupable)}' if groupable else ' No lookups are declared.' - return ( - f'{context}: {operator}({key}={name}) does not name a lookup.\n{listing}\n' - f"Declare it under 'lookups:' — {name}: {{over: , into: }}.' - ) - - -# --------------------------------------------------------------------------- -# where strings -# --------------------------------------------------------------------------- - - -def resolve_where( - node: WhereNode, - ns: Namespace, - context: str, - errors: list[str], - self_variable: str | None = None, -) -> WhereNode | None: - """Rewrite a parsed where AST into typed predicates. - - Both parameters and dimensions are legal here — a where-string is a - predicate over the frame, and the frame carries its own coordinates. What - is *not* legal is an unknown name: read as "scalar False" it would mask - every row out and produce an empty model in silence. - """ - before = len(errors) - resolved = _resolve_where(node, ns, context, errors, self_variable) - return None if len(errors) > before else resolved - - -#: An ISO literal carrying a time-of-day, which decides date vs datetime. -_HAS_TIME = re.compile(r'[T ]\d') - - -def _typed_literal( - node: UnresolvedComparisonNode, - dtype: str, - context: str, - errors: list[str], -) -> float | str | datetime.date | None: - """The comparison's literal, checked against the declared dtype. - - Getting it wrong is silent: polars reads a datetime column against an - integer as an epoch offset, so ``snapshot > 0`` drops every coordinate - before 1970 without a word (#460). Returns ``None`` once it has recorded - an error, so the caller leaves the node unresolved. - """ - value = node.value - text = isinstance(value, str) - - if dtype == 'datetime': - if not text: - errors.append( - f"{context}: '{node.name}' is a datetime dimension, so comparing it to " - f'{value!r} compares against the epoch — {node.name} > 0 means "after ' - f'1970-01-01", not what it looks like. Quote an ISO date instead: ' - f"{node.name} {node.op} '2030-01-01'." - ) - return None - try: - return ( - datetime.datetime.fromisoformat(str(value)) - if _HAS_TIME.search(str(value)) - else datetime.date.fromisoformat(str(value)) - ) - except ValueError: - errors.append( - f"{context}: '{node.name}' is a datetime dimension and {value!r} is not an " - f"ISO date. Write '2030-01-01' or '2030-01-01T06:00'." - ) - return None - - if dtype == 'str' and not text: - errors.append( - f"{context}: '{node.name}' has dtype 'str', so comparing it to the number " - f'{value!r} matches no label. Quote it if it is one: {node.name} {node.op} ' - f"'{value:g}'." - ) - return None - if dtype in ('int', 'float', 'bool') and text: - errors.append( - f"{context}: '{node.name}' has dtype '{dtype}', so comparing it to the string " - f'{value!r} matches nothing. Drop the quotes if it is a number.' - ) - return None - return value - - def _declared_rhs_error(context: str, node: UnresolvedComparisonNode, value: str, kind: str) -> str: """Why the right-hand side of a where-comparison may not name a declaration.""" - shown = f"'{node.name} {node.op} {value}'" + comparison = f"'{node.name} {node.op} {value}'" if kind == 'parameter': return ( - f'{context}: {shown} compares two parameters, which is not in the ' + f'{context}: {comparison} compares two parameters, which is not in the ' f'language — a where-comparison tests one parameter or dimension against ' f'a literal. Precompute the comparison as a boolean parameter in data ' f'prep and test that.' ) if kind == 'variable': - return f'{context}: {shown} compares against variable {value!r}. A where mask is built before variables exist.' + return ( + f'{context}: {comparison} compares against variable {value!r}. ' + f'A where mask is built before variables exist.' + ) if kind == 'lookup': return ( - f'{context}: {shown} compares {node.name!r} against lookup {value!r}, and a ' + f'{context}: {comparison} compares {node.name!r} against lookup {value!r}, and a ' f'lookup is structure rather than data — every other comparison tests a name ' f'against a literal. A lookup on the right-hand side is the one exception, and ' f'only where the left-hand side is a lookup sharing its dimension and its target.' ) return ( - f'{context}: {shown} compares against dimension {value!r}, which the RHS reads ' + f'{context}: {comparison} compares against dimension {value!r}, which the RHS reads ' f'as the literal coordinate {value!r} and so masks everything out. Comparing two ' f'dimensions is not in the language; if {value!r} is a coordinate rather than the ' f'dimension, rename one of the two.' @@ -720,11 +815,11 @@ def _lookup_pair_error(context: str, node: UnresolvedComparisonNode, other: str, the same one, or no value of one is ever a value of the other. Both wrong answers are silent, and a build's data library decides which one. """ - shown = f"'{node.name} {node.op} {other}'" + comparison = f"'{node.name} {node.op} {other}'" left_over, right_over = ns.over_of(node.name), ns.over_of(other) if left_over != right_over: return ( - f'{context}: {shown} compares lookups over different dimensions ' + f'{context}: {comparison} compares lookups over different dimensions ' f"('{left_over}' and '{right_over}') — there is no row carrying both, so the " f'comparison has nothing to test. Two lookups may be compared only where they ' f'map out of the same dimension.' @@ -732,138 +827,9 @@ def _lookup_pair_error(context: str, node: UnresolvedComparisonNode, other: str, left, right = ns.into_of(node.name), ns.into_of(other) if left is None or right is None or left != right: return ( - f'{context}: {shown} compares {_label_set_of(ns, node.name)} with ' + f'{context}: {comparison} compares {_label_set_of(ns, node.name)} with ' f'{_label_set_of(ns, other)}. No value of one is ever a value of the other, so ' f'the predicate can only mask everything out. Two lookups may be compared only ' f'where they map into the same dimension.' ) return None - - -def _resolve_position(node: UnresolvedPositionNode, ns: Namespace, context: str, errors: list[str]) -> WhereNode: - """Type ``position(dim[, by=lookup]) i``. - - The name has to be a dimension, the one thing with an order to count - along; ``by=`` has to be a lookup over that dimension, or no row of it - carries a group for a position to be counted in. - """ - if node.dimension not in ns.dimensions: - kind = ns.kind(node.dimension) - was = f'a {kind}' if kind else 'not declared' - errors.append( - f"{context}: position() counts along a dimension's coordinates, and " - f"'{node.dimension}' is {was}.\n Dimensions: {sorted(ns.dimensions)}" - ) - return node - if node.by is None: - return DimensionPositionNode(node.dimension, node.op, node.position, node.by) - call = f'position({node.dimension}, by={node.by})' - if (kind := ns.kind(node.by)) != 'lookup': - was = f'a {kind}' if kind else 'not declared' - errors.append( - f"{context}: '{call}' groups by '{node.by}', which is {was}. " - f'``by=`` takes a lookup over that dimension — either kind, since counting ' - f'inside a group lands no terms, unlike sum(by=) and at(by=).\n' - f' Lookups: {sorted(ns.lookups)}' - ) - return node - over = ns.over_of(node.by) - if over != node.dimension: - errors.append( - f"{context}: '{call}' counts positions along '{node.dimension}' but groups by a " - f"lookup over '{over}'. No row of '{node.dimension}' carries it, so there is no " - f"position within a group to name — group by a lookup over '{node.dimension}'." - ) - return node - return DimensionPositionNode(node.dimension, node.op, node.position, node.by) - - -def _resolve_where( - node: WhereNode, ns: Namespace, context: str, errors: list[str], self_variable: str | None = None -) -> WhereNode: - if isinstance(node, BooleanLiteralNode): - return node - - if isinstance(node, TypedPredicateNode): - return node - - if isinstance(node, UnresolvedNameNode): - match ns.kind(node.name): - case 'parameter': - return ParameterDefinedNode(node.name) - case 'dimension': - errors.append( - f"{context}: '{node.name}' is a dimension, and a bare dimension " - f'name is true at every coordinate — the mask has no effect. ' - f'Remove it, or compare it: where: "{node.name} > 0".' - ) - return node - case 'lookup': - return LookupDefinedNode(node.name, ns.over_of(node.name)) - case 'variable': - if node.name == self_variable: - errors.append( - f"{context}: variable '{node.name}' asks whether it exists in its own " - f'where, which nothing can answer — the mask is what decides where it ' - f'exists. Test a parameter, or another variable declared before it.' - ) - return node - return VariableDefinedNode(node.name) - case _: - errors.append(ns._unknown(node.name, context, allow_dims=True)) - return node - - if isinstance(node, UnresolvedPositionNode): - return _resolve_position(node, ns, context, errors) - - if isinstance(node, UnresolvedComparisonNode): - value = node.value - if not node.quoted and isinstance(value, str) and (rhs_kind := ns.kind(value)) is not None: - if rhs_kind == 'lookup' and ns.kind(node.name) == 'lookup': - if (refusal := _lookup_pair_error(context, node, value, ns)) is not None: - errors.append(refusal) - return node - return LookupPairComparisonNode(node.name, value, ns.over_of(node.name), node.op) - errors.append(_declared_rhs_error(context, node, value, rhs_kind)) - return node - - kind = ns.kind(node.name) - if kind in ('parameter', 'dimension', 'lookup'): - typed = _typed_literal(node, ns.dtypes[node.name], context, errors) - if typed is None: - return node - value = typed - - match kind: - case 'parameter': - assert not isinstance(value, datetime.date) - return ParameterComparisonNode(node.name, node.op, value) - case 'dimension': - return DimensionComparisonNode(node.name, node.op, value) - case 'lookup': - return LookupComparisonNode(node.name, ns.over_of(node.name), node.op, value) - case 'variable': - errors.append( - f"{context}: where references variable '{node.name}'. A where " - f'mask is built before variables exist — it may test parameters ' - f'and dimension coordinates only.' - ) - return node - case _: - errors.append(ns._unknown(node.name, context, allow_dims=True)) - return node - - if isinstance(node, NotNode): - return NotNode(_resolve_where(node.operand, ns, context, errors, self_variable)) - if isinstance(node, AndNode): - return AndNode( - _resolve_where(node.left, ns, context, errors, self_variable), - _resolve_where(node.right, ns, context, errors, self_variable), - ) - if isinstance(node, OrNode): - return OrNode( - _resolve_where(node.left, ns, context, errors, self_variable), - _resolve_where(node.right, ns, context, errors, self_variable), - ) - - assert_never(node) diff --git a/src/math_spec/typesetting/__init__.py b/src/math_spec/typesetting/__init__.py index 9659dff7..f6ceaca3 100644 --- a/src/math_spec/typesetting/__init__.py +++ b/src/math_spec/typesetting/__init__.py @@ -4,18 +4,6 @@ """Typeset a validated model — a *reading* of the math. -A consumer of the resolved core AST that produces no model and binds no data. -It exists because a declared thing can be printed the way a paper prints it, -which is the cheapest review tool available for "does this YAML say what I -meant". - -It reads the model as :func:`~math_spec.to_spec` validates it: expand -``piecewise:``, resolve names, walk. Expansion runs first, so a ``piecewise:`` -block prints as the λ-formulation it *is* rather than the sugar it was written -as. - -**One walk, many formats** — the split is :mod:`math_spec.typesetting.format`'s. - Symbols are **derived** by default, aiming at unambiguous rather than beautiful, so it prints with no setup; a :class:`~math_spec.typesetting.symbols.SymbolTable` (``--symbols``) makes it @@ -78,7 +66,10 @@ def typeset( """Render *model*'s math in *fmt*. Args: - model: Anything :func:`math_spec.to_spec` accepts. + model: Anything :func:`math_spec.to_spec` accepts. A + :class:`~math_spec.model.Spec` is rendered as it stands, so + printing one model in several formats reads and checks the file + once rather than once per format. fmt: What spells the math — one of :data:`FORMATS`. symbols: How names print, as a :class:`SymbolTable`, a path or a mapping. Names it does not carry are derived, and it must be @@ -98,25 +89,21 @@ def typeset( table written in a notation *fmt* does not read. """ schema = expand_piecewise(to_spec(model)) + namespace = Namespace.of(schema) if symbols is None: symbols = SymbolTable(fmt.notation) table = symbols if isinstance(symbols, SymbolTable) else SymbolTable.load(symbols) - walk = Walk(schema, Namespace.of(schema), Symbols(schema, fmt, table.checked_against(schema)), fmt) - - sections = [ - ('Objective', walk.objective()), - ('Subject to', walk.constraints()), - ('Definitions', walk.definitions()), - ('Variable domains', walk.variables()), - ] + walk = Walk(schema, namespace, Symbols(schema, namespace, fmt, table.checked_against(schema)), fmt) + + sections, noticed = walk.equations() rendered = [fmt.section(title, fmt.equations(lines, numbered=numbered)) for title, lines in sections if lines] blocks = [fmt.note(fmt.escape(schema.description))] if schema.description else [] if legend: - blocks += [fmt.glossary(group.title, group.entries) for group in walk.glossaries()] + blocks += [fmt.glossary(group.title, group.entries) for group in walk.glossaries(noticed)] blocks += [fmt.note(text) for text in walk.convention_notes()] - blocks += [fmt.note(text) for text in walk.translation_notes()] - blocks += [fmt.note(text) for text in walk.position_notes()] + blocks += [fmt.note(text) for text in walk.translation_notes(noticed)] + blocks += [fmt.note(text) for text in walk.position_notes(noticed)] return fmt.document([*blocks, *rendered], standalone=standalone) diff --git a/src/math_spec/typesetting/format.py b/src/math_spec/typesetting/format.py index 04e344d9..f69a6e05 100644 --- a/src/math_spec/typesetting/format.py +++ b/src/math_spec/typesetting/format.py @@ -2,30 +2,63 @@ # # SPDX-License-Identifier: MIT -r"""The seam between *what* a model says and *how* a format spells it. +"""The seam between *what* a model says and *how* a format spells it. -One walk, many formats. :mod:`math_spec.typesetting.walk` decides where a bracket is needed, -which dimension a reduction binds and where a mask belongs; a :class:`Format` -decides only that a sum is ``\sum_{…}`` or ``sum_(…)``. - -Two rules make the split hold: - -- **Everything a walk emits is *bare math*.** No ``$``, no environment; a - format wraps it with :meth:`Format.math` to embed it in prose, so the walk - never knows which mode it is in. -- **A format spells; it never decides.** No method takes an AST node or a - schema. If a format had to look at the model, the question belongs in the - walk. +The split, and each module's role in it, are in ``README.md`` beside this file. """ from __future__ import annotations from dataclasses import dataclass -from typing import TYPE_CHECKING, ClassVar, Protocol +from typing import TYPE_CHECKING, ClassVar, Literal, Protocol, get_args if TYPE_CHECKING: from collections.abc import Mapping +#: The language a symbol table's entries are written in, and the one a format +#: reads them as. Markdown is absent because its math is MathJax's, so it reads +#: ``latex``; nothing translates between the two. +Notation = Literal['latex', 'typst'] + +#: The set form, for the sidecar that has to check a string against it. +NOTATIONS = frozenset(get_args(Notation)) + +#: Every operator a walk can name. A walk asks for one by name and a format +#: spells it, so neither keeps a list of its own; the spellings are below, one +#: row per name, and a name with no row is a type error at that row's table. +OperatorName = Literal[ + 'cdot', + 'plus', + 'minus', + 'equal', + 'le', + 'ge', + 'lt', + 'gt', + 'ne', + 'in', + 'and', + 'or', + 'not', + 'false', + 'forall', + 'such_that', + 'infinity', + 'cyclic_minus', + 'cyclic_plus', + 'edge_minus', + 'edge_plus', + 'times', + 'maps_to', + 'reals', + 'integers', + 'binary_set', + 'sos_set', + 'position', + 'minimize', + 'maximize', +] + #: Every operator a walk can emit, by the name the walk uses for it, with its #: LaTeX spelling first and its Typst spelling second — one row per operator, #: so no format can be missing one. ``such_that`` is the colon in @@ -33,7 +66,7 @@ #: ``maps_to`` is the → in a coordinate map, and the three translations are #: three models: plain leaves the vacated position absent, ``cyclic_*`` wraps, #: ``edge_*`` fills it with the value it carries as a subscript. -OPERATOR_SPELLINGS: dict[str, tuple[str, str]] = { +OPERATOR_SPELLINGS: dict[OperatorName, tuple[str, str]] = { 'cdot': (r'\cdot', 'dot'), 'plus': ('+', '+'), 'minus': ('-', '-'), @@ -66,8 +99,8 @@ 'maximize': (r'\max', 'max'), } -#: The operator vocabulary itself. -OPERATOR_NAMES = frozenset(OPERATOR_SPELLINGS) +#: The set form, for the test pinning each format's table against the vocabulary. +OPERATOR_NAMES = frozenset(get_args(OperatorName)) @dataclass(frozen=True) @@ -104,15 +137,14 @@ class Glossary: class Format(Protocol): """How one output format spells what a walk emits.""" - #: File suffix, for the CLI's default output name. - suffix: ClassVar[str] - #: The notation a symbol table must be written in — ``latex`` or ``typst``; - #: markdown reads ``latex``, its math being MathJax's. - notation: ClassVar[str] + #: The notation a symbol table must be written in. + notation: ClassVar[Notation] #: Spelling for each of :data:`OPERATOR_NAMES`. - operators: ClassVar[Mapping[str, str]] + 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 ------------------------------------------------------------- @@ -161,7 +193,7 @@ def superscript(self, base: str, tail: str) -> str: ... def parenthesise(self, inner: str) -> str: ... def cardinality(self, inner: str) -> str: - """How many members a set has: ``|T|``. A fence, so not an infix entry in :data:`OPERATOR_NAMES`.""" + """An absolute-value fence: ``|x|``.""" ... def fraction(self, numerator: str, denominator: str) -> str: ... @@ -169,15 +201,11 @@ def fraction(self, numerator: str, denominator: str) -> str: ... def summation(self, domain: str, body: str) -> str: ... def cases(self, arms: list[tuple[str, str]]) -> str: - """A value defined by region: ``(value, condition)`` per arm, in order. - - Both halves arrive rendered — which arm is the fallback is the walk's - to decide, and this only stacks the rows. - """ + """A value defined by region: ``(value, condition)`` per arm, in order.""" ... def apply(self, function: str, argument: str) -> str: - """A coordinate map applied to an index: ``bus(g)``.""" + """A function applied to an argument: ``f(x)``.""" ... def joined(self, parts: list[str], operator: str) -> str: @@ -197,3 +225,18 @@ def note(self, text: str) -> str: ... def document(self, blocks: list[str], *, standalone: bool) -> str: ... + + +def aligned_rows(lines: list[Line], fmt: Format, *, gap: str) -> list[str]: + """One alignment row per line — label, left, right, condition — *gap* around the relation, trailing empty cells stripped.""" + return [ + f'{fmt.prose(line.label) if line.label else ""}{gap}{line.left} & {line.right}{gap}{line.condition}'.rstrip( + ' &' + ) + for line in lines + ] + + +def paragraphs(blocks: list[str]) -> str: + """*blocks* separated by blank lines, ending in a newline.""" + return '\n\n'.join(blocks) + '\n' diff --git a/src/math_spec/typesetting/latex.py b/src/math_spec/typesetting/latex.py index f28f5e40..dd0563bb 100644 --- a/src/math_spec/typesetting/latex.py +++ b/src/math_spec/typesetting/latex.py @@ -2,23 +2,18 @@ # # SPDX-License-Identifier: MIT -"""LaTeX (amsmath). The format that lands in a journal. - -Verbose source and a toolchain to compile it, in exchange for being the one -target a paper actually accepts. CI compiles every example with a two-package -TeX, which is also a check that the preamble stays installable from a small one. -""" +"""LaTeX (amsmath). The preamble must stay installable from a two-package TeX, which is what CI compiles with.""" from __future__ import annotations from typing import TYPE_CHECKING, ClassVar -from math_spec.typesetting.format import OPERATOR_SPELLINGS +from math_spec.typesetting.format import OPERATOR_SPELLINGS, aligned_rows, paragraphs if TYPE_CHECKING: from collections.abc import Mapping - from math_spec.typesetting.format import Entry, Line + from math_spec.typesetting.format import Entry, Line, Notation, OperatorName _ESCAPES = { '\\': r'\textbackslash{}', @@ -49,14 +44,12 @@ def _escape(text: str) -> str: class LatexFormat: """See :class:`math_spec.typesetting.format.Format`.""" - suffix: ClassVar[str] = '.tex' - notation: ClassVar[str] = 'latex' + notation: ClassVar[Notation] = 'latex' #: TeX's own em-dash ligature. dash: ClassVar[str] = '---' - #: Between the rows of a ``cases`` block. cases_row: ClassVar[str] = r' \\ ' - operators: ClassVar[Mapping[str, str]] = {name: latex for name, (latex, _) in OPERATOR_SPELLINGS.items()} + operators: ClassVar[Mapping[OperatorName, str]] = {name: latex for name, (latex, _) in OPERATOR_SPELLINGS.items()} # -- atoms ------------------------------------------------------------- @@ -76,13 +69,7 @@ def prose(self, text: str) -> str: return rf'\text{{{_escape(text)}}}' def quoted(self, label: str) -> str: - r"""The quotes in text mode, the label itself upright in math mode. - - Not one ``\text{'label'}``: MathJax renders ``\_`` inside ``\text`` - as a literal backslash, and this format is what `to_markdown` inherits - — in math mode the escape works everywhere, which is how every name - already prints. - """ + r"""The quotes in text mode and the label upright in math mode, since MathJax renders ``\_`` inside ``\text`` literally.""" return rf"\text{{'}}{self.upright(label)}\text{{'}}" def mono(self, text: str) -> str: @@ -128,13 +115,7 @@ def joined(self, parts: list[str], operator: str) -> str: def equations(self, lines: list[Line], *, numbered: bool) -> str: environment = 'align' if numbered else 'align*' - rows = [ - f'{self.prose(line.label) if line.label else ""} && {line.left} & {line.right} && {line.condition}'.rstrip( - ' &' - ) - for line in lines - ] - body = ' \\\\\n'.join(rows) + 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: @@ -148,5 +129,5 @@ def note(self, text: str) -> str: return f'\\noindent {text}' def document(self, blocks: list[str], *, standalone: bool) -> str: - body = '\n\n'.join(blocks) + '\n' + body = paragraphs(blocks) return f'{_PREAMBLE}\n{body}\n\\end{{document}}\n' if standalone else body diff --git a/src/math_spec/typesetting/markdown.py b/src/math_spec/typesetting/markdown.py index 452154c5..139da288 100644 --- a/src/math_spec/typesetting/markdown.py +++ b/src/math_spec/typesetting/markdown.py @@ -2,24 +2,19 @@ # # SPDX-License-Identifier: MIT -r"""GitHub-flavoured Markdown. The format that renders where the docs already live. - -Markdown has no math of its own — GitHub delegates to MathJax, which reads -LaTeX — so the math is :class:`LatexFormat`'s and only the document layer -differs. It exists so `docs/examples/` does not write its math by hand with -nothing checking it against the model — see `test_the_gallery_math_is_current`. -""" +"""GitHub-flavoured Markdown. GitHub renders math with MathJax, so the math is :class:`LatexFormat`'s and only the document layer differs.""" from __future__ import annotations from typing import TYPE_CHECKING, ClassVar, override +from math_spec.typesetting.format import paragraphs from math_spec.typesetting.latex import LatexFormat if TYPE_CHECKING: from collections.abc import Mapping - from math_spec.typesetting.format import Entry, Line + from math_spec.typesetting.format import Entry, Line, OperatorName def _cell(text: str) -> str: @@ -30,7 +25,6 @@ def _cell(text: str) -> str: class MarkdownFormat(LatexFormat): """See :class:`math_spec.typesetting.format.Format`. Math is LaTeX's; prose is not.""" - suffix: ClassVar[str] = '.md' #: The character, not TeX's ligature: no Markdown renderer this output #: is aimed at substitutes one, so `---` reaches the reader as three #: hyphens in the middle of a legend row. @@ -40,12 +34,9 @@ class MarkdownFormat(LatexFormat): #: those two backslashes, so MathJax would never break the row. cases_row: ClassVar[str] = r' \cr ' - #: LaTeX's, except where the spelling uses a backslash before punctuation. - #: GitHub runs Markdown's escape processing *inside* `$$`, so `\,` arrives - #: as a literal comma and `\;` as a semicolon — `\forall\, s` renders as - #: "∀, s". Letter-named macros (`\thinspace`, `\quad`) pass through - #: untouched, and MathJax treats them identically. - operators: ClassVar[Mapping[str, str]] = { + #: LaTeX's, with letter-named spacing macros: GitHub's escape pass runs inside + #: ``$$`` and turns ``\,`` into a bare comma, while ``\thinspace`` passes through. + operators: ClassVar[Mapping[OperatorName, str]] = { **LatexFormat.operators, 'forall': r'\forall\thinspace', 'such_that': r'\thinspace:\thinspace', @@ -102,5 +93,5 @@ def note(self, text: str) -> str: @override def document(self, blocks: list[str], *, standalone: bool) -> str: """No preamble: ``standalone`` adds the heading a fragment is pasted under.""" - body = '\n\n'.join(blocks) + '\n' + body = paragraphs(blocks) return f'## The math\n\n{body}' if standalone else body diff --git a/src/math_spec/typesetting/symbols.py b/src/math_spec/typesetting/symbols.py index ff3c1191..38972127 100644 --- a/src/math_spec/typesetting/symbols.py +++ b/src/math_spec/typesetting/symbols.py @@ -2,13 +2,9 @@ # # SPDX-License-Identifier: MIT -"""Which symbol each declared name prints as — and the sidecar that overrides it. +"""Which symbol each declared name prints as, and the sidecar that overrides it. -Derivation aims at *unambiguous*, not beautiful, so a model prints with no -setup; :class:`SymbolTable` is where a reader makes it conventional, in a file -of its own, since presentation is not language. What a declaration *is* stays -``description:`` on the declaration. This module decides *which* symbol a name -gets; a :class:`~math_spec.typesetting.format.Format` decides how it is written. +This module decides *which* symbol a name gets; a :class:`~math_spec.typesetting.format.Format` decides how it is written. """ from __future__ import annotations @@ -17,18 +13,20 @@ from collections.abc import Mapping from dataclasses import dataclass, field from pathlib import Path -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, cast +import math_spec.degree as degree from math_spec._yaml import read_yaml -from math_spec.degree import carries_variable from math_spec.errors import SchemaError, did_you_mean -from math_spec.resolution import Namespace, expression_of +from math_spec.resolution import expression_of +from math_spec.typesetting.format import NOTATIONS if TYPE_CHECKING: from collections.abc import Iterator from math_spec.model import ExpressionBlock, _ExpandedSpec - from math_spec.typesetting.format import Format + from math_spec.resolution import Namespace + from math_spec.typesetting.format import Format, Notation __all__ = ['SymbolTable', 'Symbols'] @@ -89,7 +87,7 @@ def printed_expressions(schema: _ExpandedSpec) -> tuple[str, ...]: return tuple(name for name, block in schema.expressions.items() if block.cases) -def chosen_expressions(schema: _ExpandedSpec) -> frozenset[str]: +def chosen_expressions(schema: _ExpandedSpec, namespace: Namespace) -> frozenset[str]: """The cased expressions the solver decides, rather than is handed. A ``when`` does not move one: a variable there asks whether the variable @@ -97,12 +95,11 @@ def chosen_expressions(schema: _ExpandedSpec) -> frozenset[str]: variable does — through a second cased expression's arms too, since :func:`~math_spec.expression_of` expands those where the name stood. """ - namespace = Namespace.of(schema) return frozenset( name for name in printed_expressions(schema) if any( - carries_variable(expression_of(text, schema, namespace, f"expression '{name}', {where}")) + degree.carries_variable(expression_of(text, schema, namespace, f"expression '{name}', {where}")) for text, where in _values_of(schema.expressions[name]) ) ) @@ -133,7 +130,7 @@ class Symbols: SchemaError: If *table* is written in a notation *fmt* does not read. """ - def __init__(self, schema: _ExpandedSpec, fmt: Format, table: SymbolTable) -> None: + def __init__(self, schema: _ExpandedSpec, namespace: Namespace, fmt: Format, table: SymbolTable) -> None: if table.notation != fmt.notation: msg = ( f'symbol table: written in {table.notation}, but this is a {fmt.notation} render ' @@ -141,13 +138,11 @@ def __init__(self, schema: _ExpandedSpec, fmt: Format, table: SymbolTable) -> No ) raise SchemaError(msg) printed = printed_expressions(schema) - chosen = frozenset(schema.variables) | chosen_expressions(schema) + chosen = frozenset(schema.variables) | chosen_expressions(schema, namespace) names = (*schema.parameters, *schema.variables, *printed) declared = frozenset(names) - #: Names whose symbol came from the table rather than the derivation; - #: the convention note quotes only the others, a table being free to - #: map a parameter to an italic symbol. + #: Names the table spelled; the convention note quotes only derived symbols. self.overridden = frozenset(table.names) & declared self.name: dict[str, str] = { name: table.names[name] @@ -208,11 +203,11 @@ class SymbolTable: An entry naming nothing in the model is an error naming the near miss. Attributes: - notation: The language the entries are written in, ``latex`` or - ``typst``; :meth:`load` lower-cases it. + notation: The language the entries are written in; :meth:`load` + lower-cases it. """ - notation: str + notation: Notation indices: dict[str, str] = field(default_factory=dict) sets: dict[str, str] = field(default_factory=dict) names: dict[str, str] = field(default_factory=dict) @@ -234,7 +229,7 @@ def load(cls, source: str | Path | Mapping[str, Any]) -> SymbolTable: msg = "symbol table: 'notation:' is required — latex or typst, the language the entries are written in." raise SchemaError(msg) notation = str(raw['notation']).lower() - if notation not in ('latex', 'typst'): + if notation not in NOTATIONS: msg = f'symbol table: unknown notation {raw["notation"]!r}. Valid notations: latex, typst.' raise SchemaError(msg) @@ -254,7 +249,7 @@ def load(cls, source: str | Path | Mapping[str, Any]) -> SymbolTable: sets[dim] = str(spec['set']) return cls( - notation=notation, + notation=cast('Notation', notation), indices=indices, sets=sets, names={k: str(v) for k, v in (raw.get('names') or {}).items()}, diff --git a/src/math_spec/typesetting/typst.py b/src/math_spec/typesetting/typst.py index df212421..85ef4fab 100644 --- a/src/math_spec/typesetting/typst.py +++ b/src/math_spec/typesetting/typst.py @@ -2,24 +2,19 @@ # # SPDX-License-Identifier: MIT -"""Typst. The format that compiles without a toolchain. - -The compiler is one self-contained binary (a pip wheel, so the suite compiles -every example without apt), and multi-letter identifiers in math are upright -by default, which is why names go through ``italic("…")``. -""" +"""Typst. Multi-letter identifiers in math are upright by default, which is why names go through ``italic("…")``.""" from __future__ import annotations import re from typing import TYPE_CHECKING, ClassVar -from math_spec.typesetting.format import OPERATOR_SPELLINGS +from math_spec.typesetting.format import OPERATOR_SPELLINGS, aligned_rows, paragraphs if TYPE_CHECKING: from collections.abc import Mapping - from math_spec.typesetting.format import Entry, Line + from math_spec.typesetting.format import Entry, Line, Notation, OperatorName _PREAMBLE = """#set page(margin: 2.5cm) #set text(size: 11pt) @@ -57,12 +52,12 @@ class TypstFormat: ``minus.circle`` does not compile. """ - suffix: ClassVar[str] = '.typ' - notation: ClassVar[str] = 'typst' + notation: ClassVar[Notation] = 'typst' #: Typst applies the same substitution TeX does. dash: ClassVar[str] = '---' + cases_row: ClassVar[str] = ', ' - operators: ClassVar[Mapping[str, str]] = {name: typst for name, (_, typst) in OPERATOR_SPELLINGS.items()} + operators: ClassVar[Mapping[OperatorName, str]] = {name: typst for name, (_, typst) in OPERATOR_SPELLINGS.items()} # -- atoms ------------------------------------------------------------- @@ -112,7 +107,7 @@ def fraction(self, numerator: str, denominator: str) -> str: return f'frac({numerator}, {denominator})' def cases(self, arms: list[tuple[str, str]]) -> str: - return 'cases({})'.format(', '.join(f'{value} & {condition}' for value, condition in arms)) + return 'cases({})'.format(self.cases_row.join(f'{value} & {condition}' for value, condition in arms)) def summation(self, domain: str, body: str) -> str: return f'sum_({domain}) {body}' @@ -127,13 +122,7 @@ def joined(self, parts: list[str], operator: str) -> str: def equations(self, lines: list[Line], *, numbered: bool) -> str: """A block equation, aligned on ``&`` as amsmath does.""" - rows = [ - f'{self.prose(line.label) if line.label else ""} & {line.left} & {line.right} & {line.condition}'.rstrip( - ' &' - ) - for line in lines - ] - body = ' \\\n '.join(rows) + body = ' \\\n '.join(aligned_rows(lines, self, gap=' & ')) numbering = '#set math.equation(numbering: "(1)")\n' if numbered else '' return f'{numbering}$ {body} $' @@ -149,5 +138,5 @@ def note(self, text: str) -> str: def document(self, blocks: list[str], *, standalone: bool) -> str: """No preamble/body split: ``standalone`` only adds the page setup.""" - body = '\n\n'.join(blocks) + '\n' + body = paragraphs(blocks) return f'{_PREAMBLE}\n{body}' if standalone else body diff --git a/src/math_spec/typesetting/walk.py b/src/math_spec/typesetting/walk.py index b53fd85a..a11589f3 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -13,11 +13,12 @@ from __future__ import annotations from dataclasses import dataclass, field, replace -from typing import TYPE_CHECKING, assert_never +from typing import TYPE_CHECKING, Literal, assert_never from math_spec.dimensions import dims_of from math_spec.expression_parser import ( ArithmeticNode, + BinaryOperator, BinaryOperatorNode, CasesNode, ComparisonNode, @@ -32,13 +33,7 @@ UnresolvedNode, VariableNode, ) -from math_spec.resolution import ( - expression_of, - where_of, -) -from math_spec.typesetting.format import Entry, Glossary, Line -from math_spec.typesetting.symbols import printed_expressions -from math_spec.where_parser import ( +from math_spec.program import ( AndNode, BooleanLiteralNode, DimensionComparisonNode, @@ -46,14 +41,21 @@ LookupComparisonNode, LookupDefinedNode, LookupPairComparisonNode, + Mask, NotNode, OrNode, ParameterComparisonNode, ParameterDefinedNode, - UnresolvedWhereNode, + PredicateOperator, VariableDefinedNode, WhereNode, ) +from math_spec.resolution import ( + expression_of, + where_of, +) +from math_spec.typesetting.format import Entry, Glossary, Line, OperatorName +from math_spec.typesetting.symbols import printed_expressions if TYPE_CHECKING: import datetime @@ -67,25 +69,35 @@ #: Operator precedence, for deciding brackets. A reduction sits at the bottom #: with ``+``: an unbracketed sum reads as capturing whatever follows it, so as #: a factor it has to be bracketed. -_PRECEDENCE = {'+': 1, '-': 1, '*': 2, '/': 2, '**': 3} +_PRECEDENCE: dict[BinaryOperator, int] = {'+': 1, '-': 1, '*': 2, '/': 2, '**': 3} _ATOM = 5 -_PREDICATES = {'==': 'equal', '!=': 'ne', '<=': 'le', '>=': 'ge', '<': 'lt', '>': 'gt'} +_PREDICATES: dict[PredicateOperator, OperatorName] = { + '==': 'equal', + '!=': 'ne', + '<=': 'le', + '>=': 'ge', + '<': 'lt', + '>': 'gt', +} + +#: What a translation does with the row the shift vacates. Three policies get +#: three spellings because they are three different equations at the boundary. +TranslationPolicy = Literal['plain', 'wrap', 'edge'] -#: Edge policy -> the operator pair that renders it, backward then forward. -#: Three policies get three spellings because they are three different -#: equations at the boundary — the vacated row dropped, wrapped, or filled. -_TRANSLATIONS = { +#: The positional forms an equation can print, each of which the legend explains once. +PositionForm = Literal['plain', 'grouped', 'from_end'] + +#: Edge policy -> the operator pair that renders it, backward then forward — +#: the vacated row dropped, wrapped, or filled. +_TRANSLATIONS: dict[TranslationPolicy, tuple[OperatorName, OperatorName]] = { 'plain': ('minus', 'plus'), 'wrap': ('cyclic_minus', 'cyclic_plus'), 'edge': ('edge_minus', 'edge_plus'), } -PRIME = "'" - - def _amount(node: ArithmeticNode) -> int | str: """``shift``'s ``offset=``: a signed number, or the name of a parameter. @@ -110,11 +122,9 @@ class _Step: """ by: int | str - policy: str + policy: TranslationPolicy fill: str = '' - #: The rendered group a partitioned translation walks inside, if it has - #: one. It rides the operator rather than the index: what changes is where - #: the axis ends, not which coordinate is being written. + #: The rendered group a partitioned translation walks inside, or empty. within: str = '' def merged(self, other: _Step) -> _Step | None: @@ -143,10 +153,7 @@ class _Context: walk: Walk offsets: dict[str, tuple[_Step, ...]] = field(default_factory=dict) - #: dim -> the subscript that replaces its own index. ``at`` re-indexes its - #: operand exactly as ``shift`` does, so it shows up at the *leaves* too — - #: but through a coordinate rather than an offset, so it renders as an - #: application, ``period(t)``, and not as arithmetic on the index. + #: dim -> the rendered subscript that replaces its index, as ``at`` re-indexes a leaf. pullbacks: 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. @@ -168,7 +175,8 @@ def reducing(self, dim: str) -> tuple[str, _Context]: once per enclosing use of the same dimension, so ``sum(q, by=bus_of)`` under ``∀ g`` sums over ``g'`` and its condition can still name ``g``. """ - dummy = f'{self.walk.symbols.index[dim]}{PRIME * self.bound.count(dim)}' + 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)) return dummy, body @@ -187,7 +195,7 @@ def subscript(self, dim: str) -> str: continue base = self.walk.format.parenthesise(text) if translated else text amount = self.walk.symbols.name[step.by] if isinstance(step.by, str) else str(abs(step.by)) - text = f'{base} {self.walk.translation(step)} {amount}' + text = f'{base} {self.walk._translation(step)} {amount}' translated = True return text @@ -204,14 +212,22 @@ def _unsigned(node: ArithmeticNode) -> ArithmeticNode | None: return None +@dataclass +class Noticed: + """What the equations printed that the legend has to explain.""" + + policies: set[TranslationPolicy] = field(default_factory=set) + grouped: bool = False + positions: set[PositionForm] = field(default_factory=set) + numeric_coordinates: set[str] = field(default_factory=set) + + class Walk: """Walks a validated schema, emitting :class:`Line`s in one format. - Stateful only in what it has *noticed* — which edge policies appeared, - whether a translation was counted inside a group, which positional forms - printed, and which dimensions were compared against a coordinate that is a - number; every one of them something the legend has to explain once the - equations print it. + :meth:`equations` prints every section and returns what it :class:`Noticed`; + the legend methods take that record, so they can only describe symbols the + equations printed. """ def __init__(self, schema: _ExpandedSpec, namespace: Namespace, symbols: Symbols, fmt: Format) -> None: @@ -219,52 +235,48 @@ def __init__(self, schema: _ExpandedSpec, namespace: Namespace, symbols: Symbols self.namespace = namespace self.symbols = symbols self.format = fmt - self.policies: set[str] = set() - self.grouped = False - self.positions: set[str] = set() - self.numeric_coordinates: set[str] = set() + self.noticed = Noticed() - def op(self, name: str) -> str: + def _op(self, name: OperatorName) -> str: return self.format.operators[name] - def translation(self, step: _Step) -> str: + def _translation(self, step: _Step) -> str: """The operator for one translation, its fill below and its group above. - Two subscripts is a TeX error rather than a rendering (#1165), and - comma-joined in one, ``0,season_of(t)`` said nothing about which was - the fill and which the group. A named offset is always backward, since - ``offset=-p`` is refused at load. + Two slots, because one subscript holding both says nothing about which + is the fill and which the group. A named offset is always backward, + since ``offset=-p`` is refused at load. """ backward, forward = _TRANSLATIONS[step.policy] - operator = self.op(backward if isinstance(step.by, str) or step.by > 0 else forward) + operator = self._op(backward if isinstance(step.by, str) or step.by > 0 else forward) if step.fill: operator = self.format.subscript(operator, [step.fill]) if not step.within: return operator - self.grouped = True + self.noticed.grouped = True return self.format.superscript(operator, step.within) - def lookup(self, name: str, index: str) -> str: + def _lookup(self, name: str, index: str) -> str: """A coordinate map applied to an index: ``bus(g)``.""" return self.format.apply(self.format.upright(name), index) - def context(self, frame: Iterable[str] = ()) -> _Context: + def _context(self, frame: Iterable[str] = ()) -> _Context: return _Context(self, bound=tuple(frame)) - def number(self, value: float) -> str: + def _number(self, value: float) -> str: if value == float('inf'): - return self.op('infinity') + return self._op('infinity') if value == int(value): return str(int(value)) mantissa, _, exponent = repr(value).partition('e') if not exponent: return mantissa power = self.format.superscript('10', str(int(exponent))) - return power if mantissa == '1' else f'{mantissa} {self.op("times")} {power}' + return power if mantissa == '1' else f'{mantissa} {self._op("times")} {power}' # -- arithmetic -------------------------------------------------------- - def arithmetic(self, node: ArithmeticNode, ctx: _Context, *, need: int = 0) -> str: + def _expression(self, node: ArithmeticNode, ctx: _Context, *, need: int = 0) -> str: text, precedence = self._arithmetic(node, ctx) return self.format.parenthesise(text) if precedence < need else text @@ -277,7 +289,7 @@ def _arithmetic(self, node: ArithmeticNode, ctx: _Context) -> tuple[str, int]: rendering decision. """ if isinstance(node, NumberNode): - return self.number(node.value), _ATOM if node.value >= 0 else 1 + return self._number(node.value), _ATOM if node.value >= 0 else 1 if isinstance(node, ParameterNode): return ctx.indexed(self.symbols.name[node.name], list(self.schema.parameters[node.name].dims)), _ATOM @@ -290,7 +302,7 @@ def _arithmetic(self, node: ArithmeticNode, ctx: _Context) -> tuple[str, int]: return self._arithmetic(node.operand, ctx) text, precedence = self._arithmetic(node.operand, ctx) operand = self.format.parenthesise(text) if precedence < 2 else text - return f'{self.op("minus")}{operand}', 2 + return f'{self._op("minus")}{operand}', 2 if isinstance(node, BinaryOperatorNode): return self._binary(node, ctx) @@ -299,7 +311,6 @@ def _arithmetic(self, node: ArithmeticNode, ctx: _Context) -> tuple[str, int]: return self._call(node, ctx) if isinstance(node, CasesNode): - # the symbol, not the block: :meth:`definitions` prints that, once return ctx.indexed(self.symbols.name[node.name], self._frame(node.name)), _ATOM if isinstance(node, UnresolvedNode | KwargNode): @@ -316,26 +327,26 @@ def _binary(self, node: BinaryOperatorNode, ctx: _Context) -> tuple[str, int]: 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, ``x^{a}^{b}`` being a LaTeX - error. + atomic to everything but another power, a stacked superscript being + ambiguous. """ if node.op == '/': - top = self.arithmetic(node.left, ctx) - bottom = self.arithmetic(node.right, ctx) + top = self._expression(node.left, ctx) + bottom = self._expression(node.right, ctx) return self.format.fraction(top, bottom), _ATOM if node.op == '**': - base = self.arithmetic(node.left, ctx, need=_PRECEDENCE['**'] + 1) - return self.format.superscript(base, self.arithmetic(node.right, ctx)), _PRECEDENCE['**'] + base = self._expression(node.left, ctx, need=_PRECEDENCE['**'] + 1) + return self.format.superscript(base, self._expression(node.right, ctx)), _PRECEDENCE['**'] precedence = _PRECEDENCE[node.op] - left = self.arithmetic(node.left, ctx, need=precedence) + 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 == '-' need = _ATOM if negated_factor else _PRECEDENCE[op] + (1 if op == '-' else 0) - right = self.arithmetic(operand, ctx, need=need) - names = {'*': 'cdot', '+': 'plus', '-': 'minus'} - return self.format.joined([left, right], self.op(names[op])), precedence + 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. @@ -349,7 +360,7 @@ def _call(self, node: FunctionCallNode, ctx: _Context) -> tuple[str, int]: dim = node.kwargs['over'] assert isinstance(dim, DimensionNode) step = self._step(_amount(node.kwargs['offset']), node.kwargs.get('edge')) - self.policies.add(step.policy) + 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)) @@ -358,46 +369,46 @@ def _call(self, node: FunctionCallNode, ctx: _Context) -> tuple[str, int]: 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.policies.add(step.policy) + self.noticed.policies.add(step.policy) source, inner = ctx.reducing(over.name) - lag = f'{ctx.subscript(over.name)} {self.translation(step)} {source}' + 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["within"])}' + 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["within"])}' ) - body = self.reduction_body(node.args[0], inner) + 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, LookupNode) for name, into in zip(by.names, by.into, strict=True): - ctx = ctx.pulled_back(into, self.lookup(name, ctx.subscript(by.dimension))) + ctx = ctx.pulled_back(into, self._lookup(name, ctx.subscript(by.dimension))) return self._arithmetic(node.args[0], ctx) if (by := node.kwargs.get('by')) is not None: assert isinstance(by, LookupNode) dummy, inner = ctx.reducing(by.dimension) conditions = [ - f'{self.lookup(name, dummy)} {self.op("equal")} {ctx.subscript(into)}' + f'{self._lookup(name, dummy)} {self._op("equal")} {ctx.subscript(into)}' for name, into in zip(by.names, by.into, strict=True) ] domain = ( - f'{self.membership(by.dimension, dummy)} {self.op("such_that")} ' - f'{self.format.joined(conditions, self.op("and"))}' + f'{self._membership(by.dimension, dummy)} {self._op("such_that")} ' + f'{self.format.joined(conditions, self._op("and"))}' ) elif (over := node.kwargs.get('over')) is not None: assert isinstance(over, DimensionNode) dummy, inner = ctx.reducing(over.name) - domain = self.membership(over.name, dummy) + domain = self._membership(over.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)) + memberships.append(self._membership(d, dummy)) domain = self.format.joined(memberships, '') - return self.format.summation(domain, self.reduction_body(node.args[0], inner)), _PRECEDENCE['+'] + return self.format.summation(domain, self._reduction_body(node.args[0], inner)), _PRECEDENCE['+'] def _group(self, by: ArithmeticNode | None, dim: str) -> str: """A ``by=`` as the superscript its translation operator carries. @@ -409,7 +420,7 @@ def _group(self, by: ArithmeticNode | None, dim: str) -> str: if by is None: return '' assert isinstance(by, LookupNode) - return self.lookup(by.names[0], self.symbols.index[dim]) + return self._lookup(by.names[0], self.symbols.index[dim]) def _width(self, node: ArithmeticNode) -> str: """``sum_back``'s ``within=``: a number, or a parameter's own symbol. @@ -422,7 +433,7 @@ def _width(self, node: ArithmeticNode) -> str: if isinstance(node, ParameterNode): return self.symbols.name[node.name] assert isinstance(node, NumberNode) - return self.number(node.value) + 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. @@ -437,139 +448,138 @@ def _step(self, by: int | str, edge: ArithmeticNode | None) -> _Step: if edge is None: return _Step(by, 'plain') assert isinstance(edge, NumberNode) - return _Step(by, 'edge', self.number(edge.value)) + return _Step(by, 'edge', self._number(edge.value)) - 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 _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: ArithmeticNode, 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. The precedence rule would - bracket that too, and a renderer that brackets everything is one nobody - trusts to bracket the thing that matters. + a nested reduction, which is unambiguous. """ additive = isinstance(node, UnaryOperatorNode) or ( isinstance(node, BinaryOperatorNode) and node.op in ('+', '-') ) - return self.arithmetic(node, ctx, need=2 if additive else 0) + return self._expression(node, ctx, need=2 if additive else 0) # -- where strings ----------------------------------------------------- - def where(self, node: WhereNode, ctx: _Context, *, need: int = 0) -> str: + def _predicate(self, node: WhereNode, ctx: _Context, *, need: int = 0) -> str: text, precedence = self._where(node, ctx) return self.format.parenthesise(text) if precedence < need else text def _where(self, node: WhereNode, ctx: _Context) -> tuple[str, int]: if isinstance(node, BooleanLiteralNode): - assert not node.value, 'resolution folds a True literal away before anything prints it' - return self.op('false'), _ATOM + assert not node.value, 'an always-true mask is folded away or refused before anything prints it' + return self._op('false'), _ATOM if isinstance(node, ParameterDefinedNode): - block = self.schema.parameters[node.name] - indexed = ctx.indexed(self.symbols.name[node.name], list(block.dims)) - if block.dtype == 'bool': + indexed = ctx.indexed(self.symbols.name[node.name], list(node.dims)) + if self.schema.parameters[node.name].dtype == 'bool': return indexed, _ATOM return f'{indexed} {self.format.prose(" is defined")}', 2 if isinstance(node, VariableDefinedNode): - dims = list(self.schema.variables[node.name].foreach) - return f'{ctx.indexed(self.symbols.name[node.name], dims)} {self.format.prose(" exists")}', 2 + return f'{ctx.indexed(self.symbols.name[node.name], list(node.dims))} {self.format.prose(" exists")}', 2 if isinstance(node, ParameterComparisonNode): - dims = list(self.schema.parameters[node.name].dims) - left = ctx.indexed(self.symbols.name[node.name], dims) - return f'{left} {self.op(_PREDICATES[node.op])} {self.literal(node.value)}', 2 + left = ctx.indexed(self.symbols.name[node.name], list(node.dims)) + return f'{left} {self._op(_PREDICATES[node.op])} {self._literal(node.value)}', 2 if isinstance(node, DimensionComparisonNode): - if isinstance(node.value, (int, float)): - self.numeric_coordinates.add(node.name) - return f'{ctx.subscript(node.name)} {self.op(_PREDICATES[node.op])} {self.literal(node.value)}', 2 + if isinstance(node.value, int | float): + self.noticed.numeric_coordinates.add(node.name) + return f'{ctx.subscript(node.name)} {self._op(_PREDICATES[node.op])} {self._literal(node.value)}', 2 if isinstance(node, DimensionPositionNode): - grouping = None if node.by is None else self.lookup(node.by, ctx.subscript(node.name)) - place = self.position(ctx.subscript(node.name), grouping) - ordinal = self.ordinal(node.name, node.position, grouping) - return f'{place} {self.op(_PREDICATES[node.op])} {ordinal}', 2 + grouping = None if node.by is None else self._lookup(node.by, ctx.subscript(node.name)) + place = self._position(ctx.subscript(node.name), grouping) + ordinal = self._ordinal(node.name, node.position, grouping) + return f'{place} {self._op(_PREDICATES[node.op])} {ordinal}', 2 if isinstance(node, LookupComparisonNode): - applied = self.lookup(node.name, ctx.subscript(node.over)) - return f'{applied} {self.op(_PREDICATES[node.op])} {self.literal(node.value)}', 2 + applied = self._lookup(node.name, ctx.subscript(node.over)) + return f'{applied} {self._op(_PREDICATES[node.op])} {self._literal(node.value)}', 2 if isinstance(node, LookupPairComparisonNode): index = ctx.subscript(node.over) - left = self.lookup(node.name, index) - right = self.lookup(node.other, index) - return f'{left} {self.op(_PREDICATES[node.op])} {right}', 2 + left = self._lookup(node.name, index) + right = self._lookup(node.other, index) + return f'{left} {self._op(_PREDICATES[node.op])} {right}', 2 if isinstance(node, LookupDefinedNode): - applied = self.lookup(node.name, ctx.subscript(node.over)) + applied = self._lookup(node.name, ctx.subscript(node.over)) return f'{applied} {self.format.prose(" is defined")}', 2 if isinstance(node, NotNode): - return f'{self.op("not")} {self.where(node.operand, ctx, need=3)}', 3 + return f'{self._op("not")} {self._predicate(node.operand, ctx, need=3)}', 3 if isinstance(node, AndNode): - sides = [self.where(node.left, ctx, need=1), self.where(node.right, ctx, need=1)] - return self.format.joined(sides, self.op('and')), 1 + sides = [self._predicate(node.left, ctx, need=1), self._predicate(node.right, ctx, need=1)] + return self.format.joined(sides, self._op('and')), 1 if isinstance(node, OrNode): - sides = [self.where(node.left, ctx, need=0), self.where(node.right, ctx, need=0)] - return self.format.joined(sides, self.op('or')), 0 - - if isinstance(node, UnresolvedWhereNode): - msg = f'{type(node).__name__} reached the typesetter; resolve the where string first.' - raise AssertionError(msg) + sides = [self._predicate(node.left, ctx, need=0), self._predicate(node.right, ctx, need=0)] + return self.format.joined(sides, self._op('or')), 0 assert_never(node) - def literal(self, value: float | str | datetime.date) -> str: - return self.number(value) if isinstance(value, (int, float)) else self.format.quoted(str(value)) + def _literal(self, value: float | str | datetime.date) -> str: + return self._number(value) if isinstance(value, int | float) else self.format.quoted(str(value)) - def position(self, index: str, grouping: str | None) -> str: + def _position(self, index: str, grouping: str | None) -> str: """``position(dim)`` applied to the row, *grouping* as a subscript — as an argument it read as a second position.""" - self.positions.add('grouped' if grouping is not None else 'plain') - symbol = self.op('position') + self.noticed.positions.add('grouped' if grouping is not None else 'plain') + symbol = self._op('position') if grouping is not None: symbol = self.format.subscript(symbol, [grouping]) return self.format.apply(symbol, index) - def ordinal(self, dimension: str, at: int, grouping: str | None) -> str: + def _ordinal(self, dimension: str, at: int, grouping: str | None) -> str: """The position compared against; a negative one counts back from the size of the set it is a position in — the group's where grouped.""" if at >= 0: - return self.number(at) - self.positions.add('from_end') + return self._number(at) + self.noticed.positions.add('from_end') size = self.symbols.set[dimension] if grouping is not None: size = self.format.subscript(size, [grouping]) - return f'{self.format.cardinality(size)} {self.op("minus")} {self.number(-at)}' + return f'{self.format.cardinality(size)} {self._op("minus")} {self._number(-at)}' - def conjoined(self, ctx: _Context, *nodes: WhereNode | None) -> str: - """The mask on a quantifier, as one condition. + def _condition(self, ctx: _Context, mask: Mask | None) -> str: + """The mask on a quantifier, printed. A mask every row passes arrives as ``None`` — resolution folds it, so this prints what a program carries — and a quantifier with no condition prints none. """ - kept = [n for n in nodes if n is not None] - parts = [self.where(n, ctx, need=1 if len(kept) > 1 else 0) for n in kept] - return self.format.joined(parts, self.op('and')) if parts else '' + return '' if mask is None else self._predicate(mask.root, ctx) - def quantifier(self, dims: list[str], condition: str) -> str: + def _quantifier(self, dims: list[str], condition: str) -> str: if not dims and not condition: return '' - over = self.format.joined([self.membership(d) for d in dims], '') + over = self.format.joined([self._membership(d) for d in dims], '') if not condition: - return f'{self.op("forall")} {over}' + return f'{self._op("forall")} {over}' if not over: return f'{self.format.prose("where ")} {condition}' - return f'{self.op("forall")} {over} {self.op("such_that")} {condition}' + return f'{self._op("forall")} {over} {self._op("such_that")} {condition}' # -- declarations ------------------------------------------------------ - def objective(self) -> list[Line]: + def equations(self) -> tuple[list[tuple[str, list[Line]]], Noticed]: + """Every titled section of equations, and what printing them noticed for the legend.""" + sections = [ + ('Objective', self._objective()), + ('Subject to', self._constraints()), + ('Definitions', self._definitions()), + ('Variable domains', self._variables()), + ] + return sections, self.noticed + + def _objective(self) -> list[Line]: """The objective's line. The expression is scalar — every reduction in it is one the file wrote @@ -579,12 +589,12 @@ def objective(self) -> list[Line]: block = self.schema.objective if block is None: return [] - sense = self.op('minimize' if block.sense == 'minimize' else 'maximize') + sense = self._op('minimize' if block.sense == 'minimize' else 'maximize') node = expression_of(block.expression, self.schema, self.namespace, 'the objective') assert not isinstance(node, ComparisonNode) - return [Line(label='', left=sense, right=self.arithmetic(node, self.context()))] + return [Line(label='', left=sense, right=self._expression(node, self._context()))] - def constraints(self) -> list[Line]: + def _constraints(self) -> list[Line]: lines = [] for name, block in self.schema.constraints.items(): context = f"constraint '{name}'" @@ -592,42 +602,36 @@ def constraints(self) -> list[Line]: if not isinstance(node, ComparisonNode): msg = f'{context}: expected a comparison, got {type(node).__name__}' raise AssertionError(msg) - ctx = self.context(frame=block.foreach) - condition = self.conjoined(ctx, where_of(block.where, self.namespace, context)) + ctx = self._context(frame=block.foreach) + condition = self._condition(ctx, where_of(block.where, self.namespace, context)) lines.append( Line( label=name, - left=self.arithmetic(node.left, ctx), - right=f'{self.op(_PREDICATES[node.op])} {self.arithmetic(node.right, ctx)}', - condition=self.quantifier(list(block.foreach), condition), + left=self._expression(node.left, ctx), + right=f'{self._op(_PREDICATES[node.op])} {self._expression(node.right, ctx)}', + condition=self._quantifier(list(block.foreach), condition), ) ) return lines - def definitions(self) -> list[Line]: + def _definitions(self) -> list[Line]: """One line per cased expression, in declaration order, defining it. - Inlining the block where its name stood is what the AST does and the - wrong thing to print: three arms are three rows tall, so whatever - follows sits beside the middle one. So a use prints the symbol and the - block prints here, as a paper states a quantity defined by region. - - Every declared one prints, used or not — the rule a variable's domain - follows, and what keeps this section independent of the others having - run. + A use prints the symbol and the block prints here, as a paper states a + quantity defined by region. Every declared one prints, used or not. """ lines = [] for name in printed_expressions(self.schema): node = expression_of(name, self.schema, self.namespace, f"expression '{name}'") assert isinstance(node, CasesNode) frame = self._frame(name) - ctx = self.context(frame) + ctx = self._context(frame) lines.append( Line( label=name, left=ctx.indexed(self.symbols.name[name], frame), - right=f'{self.op("equal")} {self.format.cases(self._arms(node, ctx))}', - condition=self.quantifier(frame, ''), + right=f'{self._op("equal")} {self.format.cases(self._arms(node, ctx))}', + condition=self._quantifier(frame, ''), ) ) return lines @@ -647,12 +651,12 @@ def _arms(self, node: CasesNode, ctx: _Context) -> list[tuple[str, str]]: when = ( self.format.prose('otherwise') if arm.when is None - else f'{self.format.prose("if ")} {self.where(arm.when, ctx, need=1)}' + else f'{self.format.prose("if ")} {self._predicate(arm.when, ctx, need=1)}' ) - arms.append((self.arithmetic(arm.value, ctx), when)) + arms.append((self._expression(arm.value, ctx), when)) return arms - def variables(self) -> list[Line]: + def _variables(self) -> list[Line]: """One line per variable, and one more for a set the variable carries. A ``sos:`` block restricts the *domain* — which members of a family may @@ -663,53 +667,48 @@ def variables(self) -> list[Line]: sets = {block.variable: block for block in self.schema.sos.values()} lines = [] for name, block in self.schema.variables.items(): - ctx = self.context(frame=block.foreach) + ctx = self._context(frame=block.foreach) symbol = ctx.indexed(self.symbols.name[name], list(block.foreach)) where = where_of(block.where, self.namespace, f"variable '{name}'", self_variable=name) - condition = self.quantifier(list(block.foreach), self.conjoined(ctx, where)) + condition = self._quantifier(list(block.foreach), self._condition(ctx, where)) lower, upper = block.bounds.lower, block.bounds.upper if block.domain == 'binary': - left, right = symbol, f'{self.op("in")} {self.op("binary_set")}' + left, right = symbol, f'{self._op("in")} {self._op("binary_set")}' else: below, above = lower == float('-inf'), upper == float('inf') if below and above: - domain = self.op('integers' if block.domain == 'integer' else 'reals') - left, right = symbol, f'{self.op("in")} {domain}' + domain = self._op('integers' if block.domain == 'integer' else 'reals') + left, right = symbol, f'{self._op("in")} {domain}' elif below: - left, right = symbol, f'{self.op("le")} {self._bound(ctx, upper)}' + left, right = symbol, f'{self._op("le")} {self._bound(ctx, upper)}' elif above: - left, right = symbol, f'{self.op("ge")} {self._bound(ctx, lower)}' + left, right = symbol, f'{self._op("ge")} {self._bound(ctx, lower)}' else: - left = f'{self._bound(ctx, lower)} {self.op("le")} {symbol}' - right = f'{self.op("le")} {self._bound(ctx, upper)}' + left = f'{self._bound(ctx, lower)} {self._op("le")} {symbol}' + right = f'{self._op("le")} {self._bound(ctx, upper)}' if block.domain == 'integer' and not (below and above): - right = f'{right}, {symbol} {self.op("in")} {self.op("integers")}' + right = f'{right}, {symbol} {self._op("in")} {self._op("integers")}' lines.append(Line(label=name, left=left, right=right, condition=condition)) if name in sets: lines.append(self._sos(name, sets[name], ctx)) return lines def _sos(self, name: str, block: SosBlock, ctx: _Context) -> Line: - """``(x_{s,o})_{o ∈ O} ∈ SOS2 ∀ s ∈ S`` — the family, and its order. - - The set runs along one dim and there is one of it per coordinate of the - rest, which is exactly the split between the subscript on the family - and the quantifier beside it. - """ + """The variable's family along the set's dim, as one member of the SOS set, quantified over the other dims.""" foreach = self.schema.variables[name].foreach family = self.format.parenthesise(ctx.indexed(self.symbols.name[name], list(foreach))) return Line( label=f'{name} sos', - left=self.format.subscript(family, [self.membership(block.over)]), - right=f'{self.op("in")} {self.op("sos_set")}{block.type}', - condition=self.quantifier([d for d in foreach if d != block.over], ''), + left=self.format.subscript(family, [self._membership(block.over)]), + right=f'{self._op("in")} {self._op("sos_set")}{block.type}', + condition=self._quantifier([d for d in foreach if d != block.over], ''), ) 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)) - return self.number(value) + return self._number(value) def _sorted(self, dims: frozenset[str]) -> list[str]: order = list(self.schema.dimensions) @@ -717,12 +716,12 @@ def _sorted(self, dims: frozenset[str]) -> list[str]: # -- legend ------------------------------------------------------------ - def glossaries(self) -> list[Glossary]: + def glossaries(self, noticed: Noticed) -> list[Glossary]: fmt = self.format sets = [ self._entry( self.symbols.set[d], - f'index {fmt.math(self.symbols.index[d])} {fmt.dash} {fmt.mono(d)}{self._coords(d)}', + f'index {fmt.math(self.symbols.index[d])} {fmt.dash} {fmt.mono(d)}{self._coords(d, noticed)}', block.description, ) for d, block in self.schema.dimensions.items() @@ -745,10 +744,10 @@ def _entry(self, symbol: str, what: str, description: str | None) -> Entry: def _over(self, dims: list[str]) -> str: if not dims: return ' (scalar)' - product = self.format.joined([self.symbols.set[d] for d in dims], self.op('times')) + product = self.format.joined([self.symbols.set[d] for d in dims], self._op('times')) return f' over {self.format.math(product)}' - def _coords(self, dim: str) -> str: + def _coords(self, dim: str, noticed: Noticed) -> str: """The dimension's carried structure, groupable maps before plain labels. A targeted lookup renders as the map it is (``bus_of: G ↦ B``); a @@ -760,12 +759,12 @@ def _coords(self, dim: str) -> str: targeted = self.schema.targeted_of(dim) labels = self.schema.labels_of(dim) clauses = [] - if dim in self.numeric_coordinates: + if dim in noticed.numeric_coordinates: clauses.append(f' ({self.format.mono(self.schema.dimensions[dim].dtype)} coordinates)') if targeted: maps = self.format.joined( [ - f'{self.format.upright(c)}: {self.symbols.set[dim]} {self.op("maps_to")} {self.symbols.set[target]}' + f'{self.format.upright(c)}: {self.symbols.set[dim]} {self._op("maps_to")} {self.symbols.set[target]}' for c, target in targeted.items() ], '', @@ -797,38 +796,33 @@ def convention_notes(self) -> list[str]: f'An index is italic too, being what a quantifier chooses, and a set is script.' ] - def translation_notes(self) -> list[str]: - """A sentence for each translation symbol the model actually printed. - - Only those: a legend explaining a symbol that is nowhere on the page is - a dead end, and plain ``t-k`` needs no note until something else stands - beside it. - """ + def translation_notes(self, noticed: Noticed) -> list[str]: + """A sentence for each translation symbol the model printed; plain ``t-k`` needs none.""" notes = [] - if 'wrap' in self.policies: - cyclic = self.format.math(f't {self.op("cyclic_minus")} k') + if 'wrap' in noticed.policies: + cyclic = self.format.math(f't {self._op("cyclic_minus")} k') notes.append( f'{cyclic} denotes cyclic translation: index {self.format.math("t-k")} taken modulo the size of ' f'the dimension ({self.format.mono("roll")}). Plain {self.format.math("t-k")} ' f'({self.format.mono("shift")}) has no wraparound {self.format.dash} terms translated past ' f'the edge are simply absent.' ) - if 'edge' in self.policies: - filled = self.format.math(f't {self.format.subscript(self.op("edge_minus"), ["v"])} k') + if 'edge' in noticed.policies: + filled = self.format.math(f't {self.format.subscript(self._op("edge_minus"), ["v"])} k') notes.append( f'{filled} denotes translation with {self.format.math("v")} standing where index ' f'{self.format.math("t-k")} leaves the dimension ({self.format.mono("shift(edge=v)")}), so the row ' f'at that boundary is built and carries {self.format.math("v")} rather than being dropped.' ) - if self.grouped: - applied = self.lookup('lookup', 't') - counted = self.format.math(f't {self.format.superscript(self.op("cyclic_minus"), applied)} k') + if noticed.grouped: + applied = self._lookup('lookup', 't') + counted = self.format.math(f't {self.format.superscript(self._op("cyclic_minus"), applied)} k') note = ( f'{counted} denotes a translation counted inside the group a lookup puts {self.format.math("t")} ' f'in ({self.format.mono("shift(by=lookup)")}), so a term never crosses out of its own group.' ) - if 'edge' in self.policies: - both = self.format.superscript(self.format.subscript(self.op('edge_minus'), ['v']), applied) + if 'edge' in noticed.policies: + both = self.format.superscript(self.format.subscript(self._op('edge_minus'), ['v']), applied) note += ( f' The two modifiers take different slots {self.format.dash} the group above, the fill ' f'below {self.format.dash} so {self.format.math(f"t {both} k")} is both at once.' @@ -836,17 +830,12 @@ def translation_notes(self) -> list[str]: notes.append(note) return notes - def position_notes(self) -> list[str]: - """A sentence for each positional symbol the model actually printed. - - The first is what the page cannot go without: a reader arrives from - papers where the index *is* the ordinal, so a page printing both - ``pos(t) = 0`` and ``t >= 3`` has to say once which is the position. - """ + def position_notes(self, noticed: Noticed) -> list[str]: + """A sentence for each positional symbol the model printed; the first says which of ``pos(t)`` and ``t`` is the position.""" notes = [] - if self.positions: + if noticed.positions: index = self.format.math('t') - place = self.format.math(self.format.apply(self.op('position'), 't')) + place = self.format.math(self.format.apply(self._op('position'), 't')) dash = self.format.dash notes.append( f"{place} denotes where index {index} sits along its dimension's own order {dash} the order " @@ -854,17 +843,17 @@ def position_notes(self) -> list[str]: f'{self.format.math("0")}. The index itself stays the coordinate, so {index} compares against ' f'labels and {place} against positions.' ) - if 'grouped' in self.positions: - applied = self.lookup('lookup', 't') - grouped = self.format.math(self.format.apply(self.format.subscript(self.op('position'), [applied]), 't')) + if 'grouped' in noticed.positions: + applied = self._lookup('lookup', 't') + grouped = self.format.math(self.format.apply(self.format.subscript(self._op('position'), [applied]), 't')) group = self.format.math(self.format.subscript(self.format.script('T'), [applied])) notes.append( f'{grouped} counts within the group a lookup puts {self.format.math("t")} in: the subscript names ' f'the map, {group} is the group it lands in, and that group has a first position of its own.' ) - if 'from_end' in self.positions: + if 'from_end' in noticed.positions: size = self.format.cardinality(self.format.script('T')) - last = self.format.math(f'{size} {self.op("minus")} {self.number(1)}') + last = self.format.math(f'{size} {self._op("minus")} {self._number(1)}') notes.append( f'{self.format.math(size)} denotes the size of the set being counted along, and a position ' f'counted from the end prints against it {self.format.dash} {last} is the last position, one ' diff --git a/src/math_spec/validation.py b/src/math_spec/validation.py index 5e4496ac..7391648e 100644 --- a/src/math_spec/validation.py +++ b/src/math_spec/validation.py @@ -2,15 +2,15 @@ # # SPDX-License-Identifier: MIT -"""Load-time validation: every expression and where string is parsed, expanded and resolved through the same pass the backends use, collecting every problem rather than raising on the first.""" +"""Load-time validation: the front door, and the pass that decides every expression.""" from __future__ import annotations from pathlib import Path -from typing import Any, assert_never +from typing import TYPE_CHECKING, Any, assert_never +import math_spec.degree as degree from math_spec._yaml import read_yaml -from math_spec.degree import carries_variable, check_expression from math_spec.dimensions import check_schema from math_spec.errors import LanguageError, SchemaError from math_spec.exclusivity import overlapping @@ -33,8 +33,11 @@ ) from math_spec.model import Spec from math_spec.operators import BUILTINS, unknown_operator_message -from math_spec.resolution import Namespace, resolve_expression, resolve_where -from math_spec.where_parser import WhereNode, parse_where +from math_spec.program import BooleanLiteralNode +from math_spec.resolution import Namespace, resolve_expression, resolve_where_text + +if TYPE_CHECKING: + from math_spec.program import WhereNode def to_spec(model: str | Path | dict[str, Any] | Spec) -> Spec: @@ -54,8 +57,8 @@ def to_spec(model: str | Path | dict[str, Any] | Spec) -> Spec: LanguageError: Anything the language does not accept. """ if isinstance(model, (list, tuple)): - msg = 'a model is one file, one dict or one Spec, never a list of them; merge the declarations into one dict (#30).' - raise TypeError(msg) + msg = 'a model is one file, one dict or one Spec, never a list of them; merge the declarations into one dict.' + raise SchemaError(msg) if isinstance(model, Spec): return model return Spec.model_validate(model if isinstance(model, dict) else read_yaml(Path(model))) @@ -107,7 +110,7 @@ def validate_expressions(schema: Spec) -> None: f'ambiguous with the dimension itself.' for f in sorted(formals & ns.dimensions) ) - _check_template_names(body_ast, macro.template, context, ns, formals, errors) + _check_template_names(body_ast, context, ns, formals, errors) for ename, block in schema.expressions.items(): context = f"Named expression '{ename}'" @@ -119,20 +122,23 @@ def validate_expressions(schema: Spec) -> None: masks: dict[str, WhereNode] = {} for case_name, case in block.cases.items(): arm_context = case_context(ename, case_name) - if (mask := _check_where(case.when, ns, arm_context, errors)) is not None: - masks[case_name] = mask + if (mask := resolve_where_text(case.when, ns, arm_context, errors)) is not None: + if isinstance(mask, BooleanLiteralNode): + errors.append(_constant_arm(arm_context, value=mask.value)) + else: + masks[case_name] = mask _check_expression(case.expression, schema, ns, arm_context, errors, comparison=False, ceiling=1) assert block.otherwise is not None _check_expression(block.otherwise, schema, ns, case_context(ename, None), errors, comparison=False, ceiling=1) if len(errors) == found: - errors.extend(f'{context}: {problem}' for problem in overlapping(masks, schema)) + errors.extend(f'{context}: {problem}' for problem in overlapping(masks, ns.dtypes)) for vname, vdef in schema.variables.items(): - _check_where(vdef.where, ns, f"Variable '{vname}'", errors, self_variable=vname) + resolve_where_text(vdef.where, ns, f"Variable '{vname}'", errors, self_variable=vname) for cname, cdef in schema.constraints.items(): context = f"Constraint '{cname}'" - _check_where(cdef.where, ns, context, errors) + resolve_where_text(cdef.where, ns, context, errors) _check_expression(cdef.expression, schema, ns, context, errors, comparison=True, ceiling=2) if schema.objective is not None: @@ -149,6 +155,23 @@ def _prefixed(context: str, e: ValueError) -> str: return str(e) if str(e).startswith(context) else f'{context}: {e}' +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`.' + + def _check_expression( expression: str, schema: Spec, @@ -176,7 +199,7 @@ def _check_expression( resolved = resolve_expression(ast, ns, context, errors) if resolved is None: return - if isinstance(resolved, ComparisonNode) and not carries_variable(resolved): + 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' @@ -185,31 +208,11 @@ def _check_expression( f'bound, or drop the declaration and check the fact where the data is prepared.' ) try: - check_expression(resolved, context, ceiling=ceiling) + degree.check_expression(resolved, context, ceiling=ceiling) except LanguageError as e: errors.append(str(e)) -def _check_where( - text: str | None, - ns: Namespace, - context: str, - errors: list[str], - self_variable: str | None = None, -) -> WhereNode | None: - """Parse and resolve one mask, returning it — ``None`` where there is none to read, and where reading it failed.""" - if text is None: - return None - try: - node = parse_where(text) - except ValueError as e: - errors.append(f'{context}: {e}') - return None - found = len(errors) - resolved = resolve_where(node, ns, context, errors, self_variable) - return resolved if len(errors) == found else None - - def _names_in(value: ArithmeticNode) -> tuple[str, ...]: """The names a lookup kwarg carries: one bare, several bracketed, none otherwise.""" if isinstance(value, NameNode): @@ -219,34 +222,30 @@ def _names_in(value: ArithmeticNode) -> tuple[str, ...]: def _check_template_names( node: ArithmeticNode, - template: str, context: str, ns: Namespace, formals: frozenset[str], errors: list[str], ) -> None: - """Name-check a macro body treating formals as bound — not resolution, since a formal has no kind until a call site binds it.""" + """Name-check a macro body treating formals as bound — not resolution, since a formal has no kind until a call site binds it. + + A case arm's value only: its ``when`` is the declaration's, checked there. + """ if isinstance(node, NumberNode | VariableNode | ParameterNode | KwargNode | KeywordNode | NameListNode): return if isinstance(node, NameNode): if node.name not in formals and ns.kind(node.name) is None: - errors.append( - f"{context}: '{node.name}' not found in template {template!r}.\n" - f' Formals: {sorted(formals)}\n' - f' Variables: {sorted(ns.variables)}\n' - f' Parameters: {sorted(ns.parameters)}\n' - f"Check for typos, or ensure '{node.name}' is declared." - ) + errors.append(ns.unknown(node.name, context, allow_dims=False, formals=formals)) return if isinstance(node, UnaryOperatorNode): - _check_template_names(node.operand, template, context, ns, formals, errors) + _check_template_names(node.operand, context, ns, formals, errors) return if isinstance(node, BinaryOperatorNode): - _check_template_names(node.left, template, context, ns, formals, errors) - _check_template_names(node.right, template, context, ns, formals, errors) + _check_template_names(node.left, context, ns, formals, errors) + _check_template_names(node.right, context, ns, formals, errors) return if isinstance(node, FunctionCallNode): @@ -254,31 +253,30 @@ def _check_template_names( if builtin is None: errors.append(f'{context}: {unknown_operator_message(node.name)}') for arg in node.args: - _check_template_names(arg, template, context, ns, formals, errors) - dimension_kwargs, lookup_kwargs, edge_kwargs = ( - (builtin.dimension_kwargs, builtin.lookup_kwargs, builtin.edge_kwargs) if builtin else ((), (), ()) - ) + _check_template_names(arg, context, ns, formals, errors) for kwarg, value in node.kwargs.items(): - if kwarg in dimension_kwargs: - 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.' + match builtin.kind_of(kwarg) 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 'lookup': + errors.extend( + f'{context}: {node.name}({kwarg}={one}) does not name a lookup or a formal of this macro.' + for one in _names_in(value) + if one not in formals and ns.kind(one) != 'lookup' ) - elif kwarg in lookup_kwargs: - errors.extend( - f'{context}: {node.name}({kwarg}={one}) does not name a lookup or a formal of this macro.' - for one in _names_in(value) - if one not in formals and ns.kind(one) != 'lookup' - ) - elif kwarg not in edge_kwargs: - _check_template_names(value, template, context, ns, formals, errors) + case 'value': + _check_template_names(value, context, ns, formals, errors) + case 'edge': + pass # a keyword or a number: nothing in it to name return if isinstance(node, CasesNode): - # the values only: a `when` is the declaration's, checked there for arm in node.arms: - _check_template_names(arm.value, template, context, ns, formals, errors) + _check_template_names(arm.value, context, ns, formals, errors) return assert_never(node) diff --git a/src/math_spec/where_parser.py b/src/math_spec/where_parser.py deleted file mode 100644 index 66daa8eb..00000000 --- a/src/math_spec/where_parser.py +++ /dev/null @@ -1,443 +0,0 @@ -# SPDX-FileCopyrightText: math-spec Contributors -# -# SPDX-License-Identifier: MIT - -"""pyparsing-based parser for where strings — grammar and AST only. - -Parses strings like ``"p_max > 0 AND NOT is_must_run"`` into an AST. What a -mask *means* is the consumer's business: it evaluates the AST against the data -it holds. -""" - -from __future__ import annotations - -import re -from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, Literal, assert_never, cast - -import pyparsing as pp - -from math_spec.errors import SchemaError -from math_spec.expression_parser import REAL - -if TYPE_CHECKING: - import datetime - from collections.abc import Callable, Iterator, Mapping, Sequence - -PredicateOperator = Literal['<=', '>=', '==', '!=', '<', '>'] - -# --------------------------------------------------------------------------- -# AST nodes -# --------------------------------------------------------------------------- - - -@dataclass(frozen=True) -class BooleanLiteralNode: - value: bool - - -@dataclass(frozen=True) -class UnresolvedNameNode: - """A bare name — unresolved. ``resolution.py`` types it.""" - - name: str - - -@dataclass(frozen=True) -class UnresolvedComparisonNode: - """A comparison against an unresolved name. ``resolution.py`` types it.""" - - name: str - op: PredicateOperator - value: float | str - #: Whether the right-hand side 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. Consumed by resolution, never lowered. - quoted: bool = False - - -@dataclass(frozen=True) -class UnresolvedPositionNode: - """``position(dim) i`` before the name is checked. - - Kept apart from :class:`UnresolvedComparisonNode` because its left-hand - side is not a name but an *application* to one, which no bare name can - carry. ``resolution.py`` types it into :class:`DimensionPositionNode`. - """ - - dimension: str - op: PredicateOperator - position: int - by: str | None = None - - -@dataclass(frozen=True) -class ParameterDefinedNode: - """True wherever the named parameter is non-null and finite.""" - - name: str - - -@dataclass(frozen=True) -class VariableDefinedNode: - """True at the coordinates where the named variable exists. - - The variable counterpart of :class:`ParameterDefinedNode`, and spelled the - same way — a bare name. A parameter's bare name asks whether it has a value - here; a variable's asks whether it exists here. - """ - - name: str - - -@dataclass(frozen=True) -class ParameterComparisonNode: - """Compare a parameter against a literal, element-wise.""" - - name: str - op: PredicateOperator - value: float | str - - -@dataclass(frozen=True) -class DimensionComparisonNode: - """Compare a dimension's own coordinates against a literal.""" - - name: str - op: PredicateOperator - value: float | str | datetime.date - - -@dataclass(frozen=True) -class DimensionPositionNode: - """Compare where a row sits along a dimension against a position — ``position(snapshot) == 0``. - - Both sides are integers, negative counting from the end; comparing - coordinates against the label *at* a position would read differently on an - axis whose coordinates do not arrive sorted (#32). With ``by`` the position - is counted within each group the lookup makes. - """ - - name: str - op: PredicateOperator - position: int - by: str | None = None - - -@dataclass(frozen=True) -class LookupComparisonNode: - """Compare a lookup's values against a literal — ``period_of == 2030``. - - ``over`` is the dimension the lookup maps out of, copied off the - declaration during resolution so the frame check and every consumer read it - here rather than looking the lookup up again. - """ - - name: str - over: str - op: PredicateOperator - value: float | str | datetime.date - - -@dataclass(frozen=True) -class LookupPairComparisonNode: - """Compare two lookups over one dimension — ``from != to``. - - The one comparison whose both sides are structure: two maps out of the - same dimension, tested row by row on that dimension's own table. Over - different dims there is no row to compare them on, which resolution - refuses. - """ - - name: str - other: str - over: str - op: PredicateOperator - - -@dataclass(frozen=True) -class LookupDefinedNode: - """True where the named lookup has a value — the partial-lookup case. - - A lookup may be partial: a null says the label belongs to no group (a - generator on no bus, a line with one open end). This is how a declaration - asks for the labels that *do* map, spelled as a bare name exactly as a - parameter's definedness is. - """ - - name: str - over: str - - -@dataclass(frozen=True) -class NotNode: - operand: WhereNode - - -@dataclass(frozen=True) -class AndNode: - left: WhereNode - right: WhereNode - - -@dataclass(frozen=True) -class OrNode: - left: WhereNode - right: WhereNode - - -WhereNode = ( - BooleanLiteralNode - | UnresolvedNameNode - | UnresolvedComparisonNode - | UnresolvedPositionNode - | DimensionPositionNode - | ParameterDefinedNode - | VariableDefinedNode - | ParameterComparisonNode - | DimensionComparisonNode - | LookupComparisonNode - | LookupPairComparisonNode - | LookupDefinedNode - | NotNode - | AndNode - | OrNode -) - -#: What resolution rewrites away on the where side — the three nodes whose -#: left-hand side is still a name the schema has not been asked about. The -#: expression side has :data:`~math_spec.expression_parser.UnresolvedNode` for -#: the same reason, and a pass meeting either ran before resolution. -UnresolvedWhereNode = UnresolvedNameNode | UnresolvedComparisonNode | UnresolvedPositionNode - -#: Every predicate resolution has typed: it names a declaration and the kind is -#: settled. Resolution passes these straight through, having nothing left to -#: decide about them. -TypedPredicateNode = ( - ParameterComparisonNode - | ParameterDefinedNode - | VariableDefinedNode - | DimensionComparisonNode - | DimensionPositionNode - | LookupComparisonNode - | LookupPairComparisonNode - | LookupDefinedNode -) - -#: The boolean connectives — the only where nodes carrying other where nodes, -#: and so the only place a walk over a predicate recurses. -ConnectiveWhereNode = NotNode | AndNode | OrNode - - -# --------------------------------------------------------------------------- -# What a predicate reads -# --------------------------------------------------------------------------- - - -def atoms(where: WhereNode) -> Iterator[TypedPredicateNode]: - """Every node in *where* that reads a declaration, connectives removed. - - A predicate is a tree of :data:`ConnectiveWhereNode` over leaves that each - name one declaration, so every question about what a mask *reads* is asked - of the leaves and answered by taking them together. A boolean literal reads - nothing and yields nothing. - - Raises: - AssertionError: An unresolved node, which is a pass running before - resolution rather than a predicate with a property to read. - """ - if isinstance(where, NotNode): - yield from atoms(where.operand) - elif isinstance(where, (AndNode, OrNode)): - yield from atoms(where.left) - yield from atoms(where.right) - elif isinstance(where, UnresolvedWhereNode): - msg = f'{type(where).__name__} reached a predicate walk unresolved.' - raise AssertionError(msg) - elif not isinstance(where, BooleanLiteralNode): - yield where - - -def dims_read(where: WhereNode, name_dims: Mapping[str, Sequence[str]]) -> frozenset[str]: - """Which dims *where* reads, given what each declared name is read through. - - The dim rule for the predicate side, stated once: **a mask is read at the - coordinates its leaves are read at**. A parameter is read through its own - dims, a variable through the frame it is declared over, a comparison on a - dimension through that dimension, and a lookup through the dimension it - maps out of — a lookup being read on the dim it leaves, not the one it - lands in. - - A consumer masking rows needs this to know which coordinates a mask can - restrict, and answering it separately is the mistake - ``what-counts-as-language.md`` forbids: two consumers deciding differently - would mask the same model differently, with no error anywhere. - - Args: - where: A resolved predicate. - name_dims: Every declared name to the dims it is read through — - parameters by their ``dims`` and variables by their ``foreach``, - one flat mapping because the language has one flat namespace. - - Returns: - The dims read, which is empty for a predicate over nothing but - literals. - """ - return frozenset(dim for atom in atoms(where) for dim in _atom_dims(atom, name_dims)) - - -def _atom_dims(atom: TypedPredicateNode, name_dims: Mapping[str, Sequence[str]]) -> frozenset[str]: - """One leaf's dims — the rule :func:`dims_read` is the union of. - - Private because a caller wanting one leaf's answer wants the whole - predicate's, and separate because the load-time frame check reports per - leaf and so cannot take the union. Closed by ``assert_never``: a predicate - node added without a reading is a type error here, at the one place that - has to grow a branch, rather than a wrong dim set at the first model to use - it. - """ - match atom: - case ParameterComparisonNode() | ParameterDefinedNode() | VariableDefinedNode(): - return frozenset(name_dims.get(atom.name, ())) - case DimensionComparisonNode() | DimensionPositionNode(): - return frozenset({atom.name}) - case LookupComparisonNode() | LookupPairComparisonNode() | LookupDefinedNode(): - return frozenset({atom.over}) - case _: - assert_never(atom) - - -# --------------------------------------------------------------------------- -# Grammar -# --------------------------------------------------------------------------- - - -class _Quoted(str): - """A right-hand side that arrived in quotes; :func:`_comparison` turns it back into a flag.""" - - __slots__ = () - - -def _position_comparison(tokens: pp.ParseResults) -> UnresolvedPositionNode: - """``position(dim[, by=lookup]) i`` off the tokens the grammar captured.""" - *call, op, at = tokens - dimension, by = call[0], call[1] if len(call) > 1 else None - return UnresolvedPositionNode( - str(dimension), cast('PredicateOperator', op), cast('int', at), None if by is None else str(by) - ) - - -def _comparison(tokens: pp.ParseResults) -> UnresolvedComparisonNode: - """``name literal`` off the tokens the grammar captured, the quoted marker turned into a flag.""" - name, op, value = tokens - quoted = isinstance(value, _Quoted) - return UnresolvedComparisonNode(str(name), cast('PredicateOperator', op), str(value) if quoted else value, quoted) - - -def _build_where_grammar() -> pp.ParserElement: - """Build the pyparsing grammar for where strings. - - Both quote characters are accepted because YAML already owns one of them. - ``NOT`` binds tightest, then ``AND``, then ``OR``. - """ - where_expr = pp.Forward() - - true_lit = pp.CaselessKeyword('True').set_parse_action(lambda: BooleanLiteralNode(True)) - false_lit = pp.CaselessKeyword('False').set_parse_action(lambda: BooleanLiteralNode(False)) - - # pyrefly: ignore[implicit-any-lambda] - number = pp.Regex(rf'-?({REAL}|\d+)').set_parse_action(lambda t: float(t[0])) - # pyrefly: ignore[implicit-any-lambda] - position = pp.Regex(r'-?\d+').set_parse_action(lambda t: int(t[0])) - - name = pp.Regex(r'[a-zA-Z_][a-zA-Z0-9_]*') - - quoted = (pp.QuotedString("'", esc_char='\\') | pp.QuotedString('"', esc_char='\\')).set_parse_action( - lambda t: _Quoted(t[0]) - ) - - grouped_by = pp.Suppress(',') + pp.Suppress(pp.Keyword('by')) + pp.Suppress('=') + name - comparator = pp.one_of('<= >= == != < >') - - position_call = ( - pp.Suppress(pp.Keyword('position')) + pp.Suppress('(') + name + pp.Optional(grouped_by) + pp.Suppress(')') - ) - position_comparison = (position_call + comparator + position).set_parse_action(_position_comparison) - - comparison = (name + comparator + (number | quoted | name)).set_parse_action(_comparison) - # pyrefly: ignore[implicit-any-lambda] - existence = name.copy().set_parse_action(lambda t: UnresolvedNameNode(t[0])) - - # `position_comparison` leads: it starts with a keyword that `existence` - # would otherwise take for a bare name, and `comparison` for a parameter. - # See `DimensionPositionNode` for why it converts on the left (#32). - atom = ( - true_lit - | false_lit - | position_comparison - | comparison - | existence - | (pp.Suppress('(') + where_expr + pp.Suppress(')')) - ) - - NOT = pp.CaselessKeyword('NOT').suppress() - # pyrefly: ignore[implicit-any-lambda] - not_expr = (NOT + atom).set_parse_action(lambda t: NotNode(t[0])) | atom - - AND = pp.CaselessKeyword('AND').suppress() - and_expr = not_expr + pp.ZeroOrMore(AND + not_expr) - and_expr.set_parse_action(_folder(AndNode)) - - OR = pp.CaselessKeyword('OR').suppress() - or_expr = and_expr + pp.ZeroOrMore(OR + and_expr) - or_expr.set_parse_action(_folder(OrNode)) - - where_expr <<= or_expr - return where_expr - - -def _folder(node_type: type[AndNode] | type[OrNode]) -> Callable[[pp.ParseResults], Any]: - """A parse action left-folding a flat operator chain into *node_type*. - - ``AND`` and ``OR`` differ only in the node they build; the fold is the - grammar's associativity, which is one rule. - """ - - def fold(tokens: pp.ParseResults) -> Any: - items = list(tokens) - result = items[0] - for item in items[1:]: - result = node_type(result, item) - return result - - return fold - - -_WHERE_GRAMMAR = _build_where_grammar() - -#: The spelling this grammar dropped, and its rewrite (#32). A retired syntax -#: speaks before the generic mismatch, the same way a retired kwarg does in -#: `operators.call_shape_error`: "Expected end of text, found '('" is what every -#: model written against the old spelling would otherwise get. -_INDEX_CALL = re.compile(r'\bindex\s*\(') -_INDEX_REWRITE = ( - "\n\n index() is now position(), and converts on the left: write 'position(dim) == i' " - "for 'dim == index(dim, i)', and 'position(dim, by=lookup) == i' for the grouped form." -) - - -def parse_where(text: str) -> WhereNode: - """Parse a where string into an AST. - - Raises: - SchemaError: If *text* is not a where string of the language. - """ - try: - result = _WHERE_GRAMMAR.parse_string(text, parse_all=True) - except pp.ParseException as e: - msg = f'Failed to parse where string: {text!r}\n{e}' - if _INDEX_CALL.search(text): - msg += _INDEX_REWRITE - raise SchemaError(msg) from e - return cast('WhereNode', result[0]) diff --git a/tests/fixtures.py b/tests/fixtures.py index b9cea1be..73f1008b 100644 --- a/tests/fixtures.py +++ b/tests/fixtures.py @@ -16,12 +16,14 @@ if TYPE_CHECKING: from math_spec import Spec +EXAMPLES = Path(__file__).resolve().parent.parent / 'examples' + #: One construct per file; `tools/spec_math.py` renders the operator reference #: from the same directory, so a probe added for the page is swept here too. -OPERATOR_PROBES = sorted((Path(__file__).resolve().parent.parent / 'examples' / 'operators').glob('*.yaml')) +OPERATOR_PROBES = sorted((EXAMPLES / 'operators').glob('*.yaml')) -#: The same math as ``examples/dispatch.yaml``, as a dict a test can vary with -#: :func:`override`. +#: ``examples/dispatch.yaml`` without its ``where:`` and with the constraint +#: named ``balance``, as a dict a test can vary with :func:`override`. DISPATCH_MODEL: dict[str, Any] = { 'dimensions': {'snapshot': {'dtype': 'int'}, 'generator': {'dtype': 'str'}}, 'parameters': { diff --git a/tests/test_advice.py b/tests/test_advice.py index bbcab004..a68136f0 100644 --- a/tests/test_advice.py +++ b/tests/test_advice.py @@ -31,7 +31,7 @@ @pytest.mark.parametrize( - ('patch', 'expected'), + ('patch', 'fragments'), [ pytest.param( {}, @@ -39,21 +39,27 @@ id='a-target-nothing-reaches-is-a-label-space', ), pytest.param({'lookups': {}}, ["dimension 'h' is never used"], id='a-dimension-nothing-reaches-is-unused'), + ], +) +def test_a_dimension_that_is_never_an_axis_is_named(patch, fragments): + (note,) = advice(override(LABEL_SPACE, **patch)) + assert (note.kind, note.subject) == ('never-an-axis', 'h') + for fragment in fragments: + assert fragment in str(note) + + +@pytest.mark.parametrize( + 'patch', + [ pytest.param( {'constraints': {'cap': {'foreach': ['h'], 'expression': 'sum(p, by=lk) <= k'}}}, - [], - id='grouping-into-it-makes-it-an-axis', + id='grouping-into-it', ), - pytest.param({'variables.r': {'foreach': ['h']}}, [], id='indexing-by-it-makes-it-an-axis'), + pytest.param({'variables.r': {'foreach': ['h']}}, id='indexing-by-it'), ], ) -def test_a_dimension_that_is_never_an_axis_is_named(patch, expected): - notes = advice(override(LABEL_SPACE, **patch)) - assert len(notes) == (1 if expected else 0), 'one dimension is never an axis, so one piece of advice or none' - for fragment in expected: - assert fragment in str(notes[0]) - if expected: - assert (notes[0].kind, notes[0].subject) == ('never-an-axis', 'h') +def test_a_dimension_something_reaches_is_an_axis(patch): + assert not advice(override(LABEL_SPACE, **patch)), 'a dimension a declaration indexes or groups into is an axis' #: A model with one note of each kind: `h` is a label space, and `p` is driven @@ -86,13 +92,7 @@ def _written(model: dict, tmp_path: Path) -> Path: ], ) def test_the_answer_does_not_turn_on_which_state_it_is_asked_of(form, tmp_path): - """A `Program` was advised of one kind and every other input of two (#210). - - The unboundedness pass read the file, and the arm that had already lowered - skipped it — so a consumer that lowers first, which is every consumer, - since lowering is what it wanted the program for, got the shorter answer - and no signal that a rule had been skipped. - """ + """A `Program` was advised of one kind and every other input of two (#210), with no signal that a rule had been skipped.""" assert [(n.kind, n.subject) for n in advice(form(BOTH_KINDS, tmp_path))] == [ ('never-an-axis', 'h'), ('unbounded', 'p'), diff --git a/tests/test_boundedness.py b/tests/test_boundedness.py index 00b71c80..7a136bc9 100644 --- a/tests/test_boundedness.py +++ b/tests/test_boundedness.py @@ -6,10 +6,6 @@ A note is a proof that no data can bound the variable, so most rows here are reasons it must stay silent: a false note is worse than none. - -Asked of the program, because that is the state the rule reads: every case -here lowers first, and `test_advice.py` is where the four ways to hand the -same model in are held to one answer. """ from __future__ import annotations @@ -28,8 +24,12 @@ ) +def _advice(**patch): + return unbounded_notes(to_program(schema_of(BASE, **patch))) + + def _notes(**patch) -> list[str]: - return [str(a) for a in unbounded_notes(to_program(schema_of(BASE, **patch)))] + return [str(a) for a in _advice(**patch)] @pytest.mark.parametrize( @@ -101,7 +101,7 @@ def test_nothing_is_claimed_where_the_file_does_not_decide_it(patch): @pytest.mark.parametrize('builtin', sorted(BUILTIN_NAMES)) def test_every_operator_hands_its_sign_to_its_operand(builtin): - """`_walk` gives all four built-ins one arm, on a claim each of them has to keep. + """`_record_signs` gives all four built-ins one arm, on a claim each of them has to keep. The claim is that every operator sums its argument's terms with coefficient 1 — being a reduction, a re-index or a window — so the sign passes through @@ -111,7 +111,7 @@ def test_every_operator_hands_its_sign_to_its_operand(builtin): """ assert builtin in THROUGH_EACH_OPERATOR, ( f"the built-in '{builtin}' has no case here. Add the objective that drives a free variable " - f'through it — or, if it does not hand its sign to its operand, split the shape-node arm of `_walk`.' + f'through it — or, if it does not hand its sign to its operand, split the shape-node arm of `_record_signs`.' ) notes = _notes(**THROUGH_EACH_OPERATOR[builtin]) assert len(notes) == 1, 'one variable is driven and unopposed, so one note' @@ -119,7 +119,7 @@ def test_every_operator_hands_its_sign_to_its_operand(builtin): def test_every_unopposed_variable_is_named(): - advice = unbounded_notes(to_program(schema_of(BASE, **{'objective.expression': 'sum(v + w, over=g)'}))) + advice = _advice(**{'objective.expression': 'sum(v + w, over=g)'}) assert [(a.kind, a.subject) for a in advice] == [('unbounded', 'v'), ('unbounded', 'w')], ( 'one piece of advice per variable, in objective order' ) @@ -132,13 +132,10 @@ def test_the_note_names_the_rewrite(): def test_a_curve_holds_its_variables_through_the_rows_it_emits(): """A piecewise block names no constraint in the file; its expansion does.""" - model = override( - BASE, - **{ - 'dimensions.bp': {'dtype': 'int'}, - 'parameters.bx': {'dims': ['bp']}, - 'parameters.by': {'dims': ['bp']}, - 'piecewise': {'curve': {'over': 'bp', 'links': [['v', 'bx'], ['w', 'by']]}}, - }, - ) - assert unbounded_notes(to_program(schema_of(model))) == [] + curve = { + 'dimensions.bp': {'dtype': 'int'}, + 'parameters.bx': {'dims': ['bp']}, + 'parameters.by': {'dims': ['bp']}, + 'piecewise': {'curve': {'over': 'bp', 'links': [['v', 'bx'], ['w', 'by']]}}, + } + assert _notes(**curve) == [], 'the emitted link rows pin v and w, so neither is unopposed' diff --git a/tests/test_degree.py b/tests/test_degree.py index a11084e8..83535eea 100644 --- a/tests/test_degree.py +++ b/tests/test_degree.py @@ -13,7 +13,7 @@ import pytest from math_spec import LanguageError -from math_spec.degree import carries_variable, check_binary, check_expression, is_quadratic +from math_spec.degree import carries_variable, check_binary, check_expression from math_spec.expression_parser import NameNode from math_spec.resolution import Namespace, expression_of from tests.fixtures import SMALL_MODEL, schema_of @@ -94,26 +94,16 @@ def test_degree_two_is_one_term_against_one_term_and_no_higher(text, fragment): assert fragment in str(exc.value) -def test_the_context_is_optional_and_prefixes_the_sentence(): - with pytest.raises(LanguageError, match=r"^Constraint 'k': both factors"): - check_binary(_ast('p * q'), "Constraint 'k'") - with pytest.raises(LanguageError, match=r'^both factors'): - check_binary(_ast('p * q')) - - @pytest.mark.parametrize( - ('text', 'quadratic'), + ('context', 'opening'), [ - pytest.param('p * q', True, id='a-product'), - pytest.param('sum(p * c * q, over=g)', True, id='under-a-reduction'), - pytest.param('c + shift(p * q, over=g, offset=1)', True, id='under-an-operator'), - pytest.param('p * c', False, id='affine'), - pytest.param('p + q', False, id='a-sum-is-not-a-product'), - pytest.param('c ** 2', False, id='no-variable'), + pytest.param("Constraint 'k'", r"^Constraint 'k': both factors", id='a-context-prefixes-the-sentence'), + pytest.param('', r'^both factors', id='an-empty-one-leaves-it-bare'), ], ) -def test_is_quadratic_finds_a_product_of_variables_anywhere(text, quadratic): - assert is_quadratic(_ast(text)) is quadratic +def test_the_context_prefixes_the_sentence_and_an_empty_one_leaves_it_bare(context, opening): + with pytest.raises(LanguageError, match=opening): + check_binary(_ast('p * q'), context, ceiling=1) def test_carries_variable_refuses_an_unresolved_name(): diff --git a/tests/test_dimensions.py b/tests/test_dimensions.py index e7063d43..eb3c7420 100644 --- a/tests/test_dimensions.py +++ b/tests/test_dimensions.py @@ -2,11 +2,7 @@ # # SPDX-License-Identifier: MIT -"""Dim sets are a type system, checked before any data is bound. - -Every case here used to build a model and solve it — wrongly, or larger than -the file reads as. None of them needs data to be caught. -""" +"""Dim sets are a type system, checked before any data is bound.""" from __future__ import annotations @@ -14,11 +10,11 @@ import pytest -from math_spec.dimensions import DimensionError, _check_where_dims, _name_dims, check_schema, dims_of +from math_spec.dimensions import DimensionError, _check_where_dims, dims_of +from math_spec.program import LookupPairComparisonNode, Mask from math_spec.resolution import Namespace, expression_of, where_of from math_spec.validation import to_spec -from math_spec.where_parser import dims_read -from tests.fixtures import OPERATOR_PROBES, override, schema_of +from tests.fixtures import override, schema_of if TYPE_CHECKING: from math_spec.model import Spec @@ -67,6 +63,11 @@ def _dims(expr: str) -> frozenset[str]: return dims_of(expression_of(expr, s, Namespace.of(s), 't'), s, 't') +@pytest.fixture +def namespace() -> Namespace: + return Namespace.of(_schema()) + + # --------------------------------------------------------------------------- # the rules # --------------------------------------------------------------------------- @@ -88,9 +89,16 @@ def _dims(expr: str) -> frozenset[str]: ("shift(p, over=snapshot, offset=1, edge='wrap')", {'snapshot', 'generator'}), ("shift(p, over=snapshot, offset=spinup, edge='wrap')", {'snapshot', 'generator'}), ('sum_back(p, over=snapshot, within=spinup)', {'snapshot', 'generator'}), - # the same offset a `by=` makes readable: one lag per group it maps into - ("shift(p, over=snapshot, offset=bus_lead, edge='wrap', by=snap_bus)", {'snapshot', 'generator'}), - ('sum_back(p, over=snapshot, within=bus_lead, by=snap_bus)', {'snapshot', 'generator'}), + pytest.param( + "shift(p, over=snapshot, offset=bus_lead, edge='wrap', by=snap_bus)", + {'snapshot', 'generator'}, + id='a-by-makes-an-offset-over-another-dim-readable-one-lag-per-group', + ), + pytest.param( + 'sum_back(p, over=snapshot, within=bus_lead, by=snap_bus)', + {'snapshot', 'generator'}, + id='a-by-makes-a-width-over-another-dim-readable-one-window-per-group', + ), pytest.param('p + 1', {'snapshot', 'generator'}, id='a-scalar-broadcasts'), ], ) @@ -173,7 +181,9 @@ def test_an_outer_product_is_legal_and_carries_both_dim_sets(): piecewise epigraph, which multiplies a per-segment slope by a per-snapshot variable on purpose. The guard is the constraint rule below: the *frame* has to declare the result.""" - assert _dims('cost + load') == {'generator', 'snapshot', 'bus'} + assert _dims('cost + load') == {'generator', 'snapshot', 'bus'}, ( + 'a binary operator unions its two sides rather than requiring one to contain the other' + ) # --------------------------------------------------------------------------- @@ -196,12 +206,12 @@ def test_an_outer_product_is_legal_and_carries_both_dim_sets(): ), pytest.param( {'variables.cap': {'foreach': ['generator'], 'where': 'load > 0'}}, - r"where-parameter 'load' has dims \['bus', 'snapshot'\]", + r"where-parameter 'load' reads dims \['bus', 'snapshot'\]", id='where-dim-outside-the-frame', ), pytest.param( {'variables.cap': {'foreach': ['generator'], 'where': 'snapshot > 0'}}, - "where-comparison on dimension 'snapshot'", + "where-dimension 'snapshot'", id='where-comparison-on-a-dim-outside-the-frame', ), pytest.param( @@ -216,19 +226,11 @@ def test_an_ill_dimensioned_declaration_is_rejected(patch, match): _schema(**patch) -@pytest.mark.parametrize('path', OPERATOR_PROBES, ids=lambda p: p.name) -def test_every_operator_probe_typechecks(path): - check_schema(schema_of(path)) - - class TestTheEdgeRulesAreDecidedAtLoad: - """`to_spec` refuses what `to_program` used to, so the two cannot disagree. + """A file accepted by one door and refused by the next is the bug these close (#193). - 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. A file - accepted by one door and refused by the next is the bug these close (#193): - a repository of models compiled in CI would pass, and only a consumer that - built them would find out. + 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. """ BASE: ClassVar[dict[str, Any]] = { @@ -244,20 +246,38 @@ def _refused(self, expression: str) -> str: to_spec(raw) return str(caught.value) - def test_a_shift_over_data_with_no_edge_is_refused_by_to_spec(self): - assert 'leaves vacated positions with no value' in self._refused('p <= shift(cap, over=g, offset=1)') - - def test_a_named_offset_with_no_edge_is_refused_by_to_spec(self): - assert 'per-entity offset cannot say yet' in self._refused('p <= shift(p, over=t, offset=lead)') - - def test_a_nonzero_edge_over_a_variable_is_refused_by_to_spec(self): - assert 'only fill=0 is representable' in self._refused('p <= shift(p, over=t, offset=1, edge=2)') - - def test_a_numeric_edge_on_a_window_is_refused_by_to_spec(self): - assert "takes 'wrap' or nothing" in self._refused('p <= sum_back(p, over=t, within=2, edge=0)') - - def test_a_fractional_amount_is_refused_by_to_spec(self): - assert 'must be a whole number' in self._refused("p <= shift(p, over=t, offset=1.5, edge='wrap')") + @pytest.mark.parametrize( + ('expression', 'fragment'), + [ + pytest.param( + 'p <= shift(cap, over=g, offset=1)', + 'leaves vacated positions with no value', + id='a-shift-over-data-with-no-edge', + ), + pytest.param( + 'p <= shift(p, over=t, offset=lead)', + 'per-entity offset cannot say yet', + id='a-named-offset-with-no-edge', + ), + pytest.param( + 'p <= shift(p, over=t, offset=1, edge=2)', + 'only fill=0 is representable', + id='a-nonzero-edge-over-a-variable', + ), + pytest.param( + 'p <= sum_back(p, over=t, within=2, edge=0)', + "takes 'wrap' or nothing", + id='a-numeric-edge-on-a-window', + ), + pytest.param( + "p <= shift(p, over=t, offset=1.5, edge='wrap')", + 'must be a whole number', + id='a-fractional-amount', + ), + ], + ) + def test_an_edge_rule_is_refused_by_to_spec(self, expression, fragment): + assert fragment in self._refused(expression) @pytest.mark.parametrize('width', ['0', '1.5', '-2'], ids=['zero', 'fractional', 'negative']) def test_a_literal_width_below_one_is_refused_by_to_spec(self, width): @@ -275,8 +295,7 @@ def test_a_zero_step_vacates_nothing_and_needs_no_edge(self): literal zero vacates none, so there is nothing for an `edge=` to answer for. A *named* offset may be zero in the data and is not known here. """ - raw = override(self.BASE, **{'constraints.k.expression': 'p <= shift(cap, over=g, offset=0)'}) - assert to_spec(raw) is not None, 'a zero step is none at all, so no edge policy is owed' + to_spec(override(self.BASE, **{'constraints.k.expression': 'p <= shift(cap, over=g, offset=0)'})) # --------------------------------------------------------------------------- @@ -296,44 +315,65 @@ def test_a_zero_step_vacates_nothing_and_needs_no_edge(self): pytest.param('False', set(), id='a-literal-reads-nothing'), ], ) -def test_a_predicate_is_read_at_the_coordinates_its_leaves_are_read_at(predicate, expected): - """The dim rule for the predicate side, which a consumer masking rows needs. - - Stated here rather than in whichever engine builds the mask: two consumers - answering it differently would restrict the same model differently, with no - error anywhere to say which was meant. - """ - schema = to_spec(BASE) - ns = Namespace.of(schema) - name_dims = _name_dims(schema) - - where = where_of(predicate, ns, 'test') +def test_a_predicate_is_read_at_the_coordinates_its_leaves_are_read_at(namespace, predicate, expected): + """The dim rule for the predicate side.""" + where = where_of(predicate, namespace, 'test') assert where is not None, 'a predicate the connectives cannot settle survives the fold' - assert dims_read(where, name_dims) == expected, f'{predicate!r} reads {expected}' + assert where.dims == expected -def test_a_predicate_that_admits_every_row_has_no_leaves_left_to_read(): +def test_a_predicate_that_admits_every_row_has_no_leaves_left_to_read(namespace): """`where_of` folds an always-true mask to `None`, so there is no node to ask.""" - schema = to_spec(BASE) + assert where_of('True', namespace, 'test') is None, 'folded away, not a predicate over nothing' - assert where_of('True', Namespace.of(schema), 'test') is None, 'folded away, not a predicate over nothing' - -def test_the_frame_check_and_the_reading_walk_the_same_leaves(): +def test_the_frame_check_and_the_reading_walk_the_same_leaves(namespace): """One rule, two readers — the check reports per leaf and so cannot take the union. A predicate outside the frame is refused by the *name* of the leaf that - left it, and that leaf is one `dims_read` counted: a check passing a mask + left it, and that leaf is one `Mask.dims` counted: a check passing a mask the builder then reads wider would be the divergence this shares a walk to prevent. """ - schema = to_spec(BASE) - ns = Namespace.of(schema) - - where = where_of('p_max > 0', ns, 'test') + where = where_of('p_max > 0', namespace, 'test') assert where is not None - assert dims_read(where, _name_dims(schema)) == {'generator'}, 'read at the generator axis' - with pytest.raises(DimensionError, match=r"where-parameter 'p_max' has dims \['generator'\]"): - _check_where_dims(where, schema, frozenset({'snapshot'}), 'test') + assert where.dims == {'generator'}, 'read at the generator axis' + with pytest.raises(DimensionError, match=r"where-parameter 'p_max' reads dims \['generator'\]"): + _check_where_dims(where, frozenset({'snapshot'}), 'test') + + +@pytest.mark.parametrize( + ('predicate', 'expected'), + [ + pytest.param('p_max > 0', {'p_max'}, id='a-parameter-comparison-names-the-parameter'), + pytest.param('spinup', {'spinup'}, id='a-parameter-bare-names-the-parameter'), + pytest.param('p', {'p'}, id='a-variable-bare-names-the-variable'), + pytest.param('snap_bus == "b1"', {'snap_bus'}, id='a-lookup-comparison-names-the-lookup'), + pytest.param('gen_bus', {'gen_bus'}, id='a-lookup-bare-names-the-lookup'), + pytest.param('snapshot == 0', set(), id='a-dimension-names-nothing-it-is-a-coordinate'), + pytest.param('position(snapshot) == 0', set(), id='a-position-names-nothing'), + pytest.param('p_max > 0 AND snapshot == 0', {'p_max'}, id='a-conjunction-drops-the-dimension-side'), + pytest.param('p_max > 0 AND snap_bus == "b1"', {'p_max', 'snap_bus'}, id='a-conjunction-unions-both-names'), + pytest.param('NOT p_max > 0', {'p_max'}, id='a-negation-names-what-it-negates'), + pytest.param('False', set(), id='a-literal-names-nothing'), + ], +) +def test_a_predicate_names_the_declarations_its_leaves_test(namespace, predicate, expected): + """A dimension names no declaration — it is a coordinate — so `names_read` drops it where `dims` keeps it.""" + where = where_of(predicate, namespace, 'test') + + assert where is not None, 'a predicate the connectives cannot settle survives the fold' + assert where.names_read == expected + + +def test_names_read_takes_both_sides_of_a_lookup_pair(): + """The one leaf that names two declarations — two maps compared on the dimension they share. + + BASE has one lookup per dimension, so the pair is built directly rather than + resolved from a predicate string. + """ + where = LookupPairComparisonNode('from_bus', 'to_bus', 'line', '!=') + + assert Mask(where).names_read == {'from_bus', 'to_bus'}, 'a lookup pair names both maps it compares' diff --git a/tests/test_docs.py b/tests/test_docs.py index 02ecda87..f1a643e3 100644 --- a/tests/test_docs.py +++ b/tests/test_docs.py @@ -2,14 +2,7 @@ # # SPDX-License-Identifier: MIT -"""The generated half of the documentation, held to what generates it. - -The site shows models beside the math the typesetter prints from them. Written -by hand, that math would be a claim nothing checks — on a site whose subject is -the math a file means, and in a repository that owns the renderer which would -have caught it. So the blocks are generated, and this is what makes -"generated" true of the committed files rather than of scripts nobody runs. -""" +"""The committed pages a generator writes, held to their generator.""" from __future__ import annotations @@ -63,13 +56,10 @@ def test_every_generator_is_asked(): def test_every_piecewise_method_has_a_model_on_the_notation_page(): - """What the page's `_curves()` claims: one row per `method:`, all of them. - - `PIECEWISE_METHODS` is the closed set, so a method added to the language - lands here as a missing key rather than as a section quietly showing three - of four formulations. - """ - assert set(notation.PIECEWISE) == set(PIECEWISE_METHODS) + """What the page's `_curves()` claims: one row per `method:`, all of them.""" + assert set(notation.PIECEWISE) == set(PIECEWISE_METHODS), ( + 'a method added to the language lands here as a missing key rather than as a formulation the page omits' + ) def _card_bodies(page: Path) -> list[tuple[int, str]]: diff --git a/tests/test_exclusivity.py b/tests/test_exclusivity.py index 8fdf8881..6302dc84 100644 --- a/tests/test_exclusivity.py +++ b/tests/test_exclusivity.py @@ -17,14 +17,15 @@ import pytest +from math_spec._where_parser import parse_where from math_spec.exclusivity import CELL_BUDGET, Special, Subject, _evaluate, _Frame, overlapping +from math_spec.program import AndNode, Mask, NotNode, OrNode from math_spec.resolution import Namespace, resolve_where from math_spec.validation import to_spec -from math_spec.where_parser import AndNode, NotNode, OrNode, parse_where if TYPE_CHECKING: from math_spec.model import Spec - from math_spec.where_parser import WhereNode + from math_spec.program import WhereNode #: A storage model carrying one atom of every kind a `when` can be built from. #: Every axis takes its coordinates from data, so nothing here sizes one. @@ -56,7 +57,7 @@ def schema() -> Spec: def refusals(schema: Spec, cases: dict[str, str]) -> list[str]: """Resolve each case's `when` against *schema*, then decide every pair.""" namespace = Namespace.of(schema) - return list(overlapping({name: _mask(when, namespace, name) for name, when in cases.items()}, schema)) + return list(overlapping({name: _mask(when, namespace, name) for name, when in cases.items()}, namespace.dtypes)) def _mask(text: str, namespace: Namespace, name: str) -> WhereNode: @@ -69,26 +70,50 @@ def _mask(text: str, namespace: Namespace, name: str) -> WhereNode: class TestProvesApart: - def test_the_storage_split_from_the_issue(self, schema: Spec): - """Two atoms, four regions, and the two masks are exact complements. - - Spelled as the issue spells it, less `cyclic_state_of_charge == True`, - which is a load error: a bool's bare name *is* its value. - """ - cases = { - 'first_ts': 'not cyclic and position(snapshot) == 0', - 'all_other_ts': '(not cyclic and position(snapshot) != 0) or cyclic', - } - assert refusals(schema, cases) == [], 'the two regimes are written as complements' - - def test_a_category_split(self, schema: Spec): - """Equality against distinct labels is exclusive — the theory step. - - A purely propositional reading invents a storage that is both a battery - and hydrogen, and refuses the split the feature exists for. - """ - cases = {'battery': "kind == 'battery'", 'hydrogen': "kind == 'h2'"} - assert refusals(schema, cases) == [], 'a storage carries one kind, so the two labels are apart' + @pytest.mark.parametrize( + ('cases', 'claim'), + [ + pytest.param( + { + 'first_ts': 'not cyclic and position(snapshot) == 0', + 'all_other_ts': '(not cyclic and position(snapshot) != 0) or cyclic', + }, + 'the two regimes are exact complements, spelled as the issue spells them less ' + "`cyclic_state_of_charge == True`, which is a load error: a bool's bare name is its value", + id='the-storage-split-from-the-issue', + ), + pytest.param( + {'battery': "kind == 'battery'", 'hydrogen': "kind == 'h2'"}, + 'a storage carries one kind, so two labels are apart where a propositional reading invents a ' + 'storage that is both', + id='a-category-split', + ), + pytest.param( + {'first': 'position(snapshot) == 0', 'rest': 'position(snapshot) > 0'}, + 'no row sits at rank 0 and past it at once, whatever order the coordinates arrive in (#32)', + id='an-ordering-on-a-position', + ), + pytest.param( + {'final_two': 'position(snapshot) >= -2', 'earlier': 'position(snapshot) < -2'}, + 'a rank is inside the last two or before them: ranks run away from -1, and nothing follows it', + id='a-band-counted-from-the-back', + ), + pytest.param( + {'small': 'capacity and capacity <= 10', 'large': 'capacity and capacity > 10'}, + 'a capacity is at most 10 or above it, with the absent capacity a region of its own', + id='numeric-bands', + ), + pytest.param( + {'new': 'age < 1', 'old': 'age > 0'}, + '`age` is declared int, so the bands are complements over the integers; a midpoint invented in ' + 'the gap is a coordinate the subject cannot take, and the refusal it manufactures names 0.5 — ' + 'a value no data produces, with a rewrite the file has already followed', + id='an-integer-admits-no-value-between-its-bands', + ), + ], + ) + def test_a_pair_that_shares_no_coordinate_is_proved_apart(self, schema: Spec, cases: dict[str, str], claim: str): + assert refusals(schema, cases) == [], claim @pytest.mark.parametrize( 'cases', @@ -102,40 +127,9 @@ def test_a_category_split(self, schema: Spec): ], ) def test_the_ramp_regimes_from_the_issue(self, schema: Spec, cases: dict[str, str]): - """The three quantities #2 factors a PyPSA ramp limit into, less each one's `otherwise`. - - The point of putting cases on an expression rather than on the - constraint: three independent axes multiply into eight constraint cases - and add into seven expression cases, and the inequality is written once - instead of eight times. - """ + """The three quantities #2 factors a PyPSA ramp limit into, less each one's `otherwise`.""" assert refusals(schema, cases) == [], f'{cases} claims no coordinate twice' - def test_an_ordering_on_a_position(self, schema: Spec): - """Every rank is either 0 or greater, whatever order the coordinates arrive in (#32).""" - cases = {'first': 'position(snapshot) == 0', 'rest': 'position(snapshot) > 0'} - assert refusals(schema, cases) == [], 'no row sits at rank 0 and past it at once' - - def test_a_band_counted_from_the_back(self, schema: Spec): - """The mirrored frame: ranks run away from -1, and nothing follows it.""" - cases = {'final_two': 'position(snapshot) >= -2', 'earlier': 'position(snapshot) < -2'} - assert refusals(schema, cases) == [], 'a rank is inside the last two or before them' - - def test_numeric_bands(self, schema: Spec): - """And the same for a magnitude, with the absent capacity a region of its own.""" - cases = {'small': 'capacity and capacity <= 10', 'large': 'capacity and capacity > 10'} - assert refusals(schema, cases) == [], 'a capacity is at most 10 or above it' - - def test_an_integer_admits_no_value_between_its_bands(self, schema: Spec): - """`age` is declared `int`, and the two bands are complements over the integers. - - A midpoint invented in the gap is a coordinate the subject cannot take, - and the refusal it manufactures names `0.5` — a value no data produces, - with a rewrite the file has already followed. - """ - cases = {'new': 'age < 1', 'old': 'age > 0'} - assert refusals(schema, cases) == [], 'an integer is below 1 or above 0, never between' - def test_a_magnitude_still_admits_one(self, schema: Spec): """The mirror: `capacity` is a float, so 0.5 is a coordinate it can take.""" cases = {'small': 'capacity < 1', 'large': 'capacity > 0'} @@ -151,11 +145,24 @@ def test_a_when_of_true_is_not_a_fallback(self, schema: Spec): assert refusals(schema, cases), 'a case whose `when` is True claims every coordinate' +@pytest.fixture(scope='module') +def overlap(schema: Spec) -> str: + """The one refusal a bool against a label draws, read once for every fragment asserted on it.""" + [refusal] = refusals(schema, {'cyclic': 'cyclic', 'battery': "kind == 'battery'"}) + return refusal + + class TestRefuses: - def test_an_overlap_names_both_cases_and_a_witness(self, schema: Spec): - [refusal] = refusals(schema, {'cyclic': 'cyclic', 'battery': "kind == 'battery'"}) - assert "cases 'cyclic' and 'battery' both claim the value where" in refusal - assert 'cyclic is true' in refusal + @pytest.mark.parametrize( + 'fragment', + [ + pytest.param("cases 'cyclic' and 'battery' both claim the value where", id='both-cases'), + pytest.param('cyclic is true', id='a-witness'), + pytest.param('narrow one of the two `when:` strings by the negation of the other', id='the-rewrite'), + ], + ) + def test_an_overlap_names_both_cases_a_witness_and_the_rewrite(self, overlap: str, fragment: str): + assert fragment in overlap def test_every_overlapping_pair_is_named(self, schema: Spec): """Not the first: a set with three problems has three sentences.""" @@ -167,10 +174,6 @@ def test_defined_is_not_non_zero(self, schema: Spec): [refusal] = refusals(schema, {'has_initial': 'soc_initial', 'zero': 'soc_initial == 0'}) assert 'soc_initial is 0.0' in refusal - def test_the_refusal_names_the_rewrite(self, schema: Spec): - [refusal] = refusals(schema, {'cyclic': 'cyclic', 'battery': "kind == 'battery'"}) - assert 'narrow one of the two `when:` strings by the negation of the other' in refusal - class TestWillNotDecide: def test_both_ends_of_one_axis(self, schema: Spec): @@ -217,18 +220,9 @@ class TestSoundness: subject can take, so "no witness among the cells" means "no witness". This walks a concrete grid — several points inside single cells, both infinities, an absent value, labels the masks never name — and asserts that - nothing it proved apart has a point claimed by both. - - **The two masks are drawn independently**, and only the pairs the check - proves apart are walked. A pair built as a complement — `m` against - `not m` — makes the assertion `X and not X`, false at every point under - every implementation, so a fuzz over those shapes cannot fail and certifies - nothing. - - What it does not test is the reading of an individual atom: ground truth - here evaluates through the same `_evaluate` the checker uses, so a misread - atom would agree with itself. That is what `TestProvesApart` and - `TestRefuses` pin, one atom at a time. + nothing it proved apart has a point claimed by both. Only pairs the check + proves apart are walked; a complement pair asserts X and not X and cannot + fail. """ ATOMS: ClassVar[tuple[str, ...]] = ( @@ -255,12 +249,12 @@ class TestSoundness: 'storage': [0, 1, 2], } - def _mask(self, rng: random.Random, atoms: list[Any], depth: int = 0) -> Any: + def _random_mask(self, rng: random.Random, atoms: list[Any], depth: int = 0) -> Any: if depth >= 2 or rng.random() < 0.45: atom = rng.choice(atoms) return NotNode(atom) if rng.random() < 0.25 else atom - left = self._mask(rng, atoms, depth + 1) - right = self._mask(rng, atoms, depth + 1) + left = self._random_mask(rng, atoms, depth + 1) + right = self._random_mask(rng, atoms, depth + 1) node = AndNode(left, right) if rng.random() < 0.5 else OrNode(left, right) return NotNode(node) if rng.random() < 0.15 else node @@ -283,13 +277,11 @@ def test_a_pair_proved_apart_stays_apart_on_a_finer_grid(self, schema: Spec, see dtypes = namespace.dtypes proved = 0 for _ in range(2000): - first, second = self._mask(rng, atoms), self._mask(rng, atoms) - if list(overlapping({'a': first, 'b': second}, schema)): + first, second = self._random_mask(rng, atoms), self._random_mask(rng, atoms) + if list(overlapping({'a': first, 'b': second}, dtypes)): continue proved += 1 - # The same frame the check built, so ground truth reads each atom - # the way it did — what differs is the grid, which is finer. - frame = _Frame.of([first, second], dtypes) + frame = _Frame.of([Mask(first), Mask(second)], dtypes) for point in grid: both = _evaluate(first, point, frame) and _evaluate(second, point, frame) assert not both, f'both cases claim {point} — the cells hid a witness' diff --git a/tests/test_expansion.py b/tests/test_expansion.py index 15898af4..cbd245c9 100644 --- a/tests/test_expansion.py +++ b/tests/test_expansion.py @@ -96,7 +96,9 @@ ], ) def test_a_call_expands_to_core_ast(expressions, macros, call, want): - assert parse_and_expand(call, schema(expressions=expressions, macros=macros)) == parse_expression(want) + assert parse_and_expand(call, schema(expressions=expressions, macros=macros), 'expression') == parse_expression( + want + ) @pytest.mark.parametrize( @@ -120,7 +122,7 @@ def test_a_bad_named_expression_is_refused_at_load(expressions, match): def test_a_refusal_names_its_context_once(): with pytest.raises(LanguageError) as exc: schema(expressions={'a': 'a + 1'}) - assert str(exc.value).count("Named expression 'a'") == 1, str(exc.value) + assert str(exc.value).count("Named expression 'a'") == 1, 'the context is prefixed once, not once per pass' @pytest.mark.parametrize( @@ -132,7 +134,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})) + parse_and_expand(call, schema(macros={'ws': WEIGHTED_SUM}), 'expression') @pytest.mark.parametrize( diff --git a/tests/test_lowering.py b/tests/test_lowering.py index e9fda9ac..d3705ffd 100644 --- a/tests/test_lowering.py +++ b/tests/test_lowering.py @@ -4,27 +4,19 @@ """The lowering pass: a resolved model in, a logical plan out. -The plan is read back node by node rather than through the answer it produces: -it is the contract both lanes are written against, so its *shape* is the thing -under test here and what either lane then builds from it is not. - -Nothing in this module binds data, builds a model or names a lane. That is the -point of it — the pass has one input and one output, both of them values, and a -test that needed a solver to reach it would be testing the assembly instead. -Lowering's verdict reaching a caller is ``test_language_boundary.py``; the two -lanes agreeing about it is ``test_degree_parity.py`` and its siblings. +The plan is read back node by node rather than through the answer it produces — +it is the contract consumers are written against. """ from __future__ import annotations from dataclasses import FrozenInstanceError -from pathlib import Path from typing import TYPE_CHECKING, get_args import pytest -import math_spec.program as program_module -from math_spec import LanguageError, Spec +from math_spec import LanguageError, Spec, to_program +from math_spec._where_parser import parse_where from math_spec.exclusivity import overlapping from math_spec.expression_parser import FunctionCallNode, NumberNode from math_spec.lowering import _Lowering, lower_program @@ -32,18 +24,26 @@ from math_spec.program import ( QUADRATIC_POSITIONS, Add, + AndNode, At, + BooleanLiteralNode, Cases, Constant, + DimensionComparisonNode, DimensionDeclaration, Divide, ExpressionNode, Footprint, GroupSum, LookupDeclaration, + Mask, Multiply, Negate, + NotNode, + OrNode, Parameter, + ParameterComparisonNode, + ParameterDefinedNode, Power, Program, Region, @@ -58,19 +58,38 @@ walk, ) from math_spec.resolution import Namespace, expression_of, where_of -from math_spec.where_parser import ( - AndNode, - BooleanLiteralNode, - DimensionComparisonNode, - ParameterComparisonNode, - ParameterDefinedNode, -) -from tests.fixtures import DISPATCH_MODEL, SMALL_MODEL, override, schema_of +from tests.fixtures import DISPATCH_MODEL, EXAMPLES, SMALL_MODEL, override, schema_of if TYPE_CHECKING: from math_spec.expression_parser import ArithmeticNode -EXAMPLES_DIR = Path(__file__).resolve().parents[1] / 'examples' +DISPATCH_YAML = EXAMPLES / 'dispatch.yaml' + +#: The mask `examples/dispatch.yaml` puts on `p`, as the plan carries it. +P_MAX_POSITIVE = ParameterComparisonNode('p_max', '>', 0.0, ('generator',)) + +#: One dimension, one parameter, one bounded variable and a scalar constraint: +#: the smallest model that loads, for a claim about the plan's record rather +#: than about the math in it. A test adds what it judges with :func:`override`. +TINY = { + 'dimensions': {'g': {}}, + 'parameters': {'cost': {'dims': ['g']}}, + 'variables': {'p': {'foreach': ['g'], 'bounds': {'lower': 0, 'upper': 1}}}, + 'constraints': {'c': {'foreach': [], 'expression': 'sum(p, over=g) >= 1'}}, +} + +#: `fixtures.SMALL_MODEL` plus a second groupable lookup 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 +#: and two lookups over one of them. +SHAPES_MODEL = override( + SMALL_MODEL, + **{ + 'dimensions.z': {'dtype': 'str'}, + 'lookups.lk2': {'over': 'g', 'into': 'z'}, + 'parameters.lead': {'dims': ['g'], 'dtype': 'int'}, + }, +) def resolved(text: str, schema: Spec) -> ArithmeticNode: @@ -83,95 +102,61 @@ def resolved(text: str, schema: Spec) -> ArithmeticNode: return expression_of(text, schema, Namespace.of(schema), 't') -DISPATCH_YAML = EXAMPLES_DIR / 'dispatch.yaml' - - @pytest.fixture def dispatch_schema() -> Spec: return schema_of(DISPATCH_YAML) +@pytest.fixture +def dispatch_program(dispatch_schema) -> Program: + return lower_program(expand_piecewise(dispatch_schema)) + + +@pytest.fixture +def shapes_schema() -> Spec: + return schema_of(SHAPES_MODEL) + + # --------------------------------------------------------------------------- # the plan the language lowers to # --------------------------------------------------------------------------- -def test_lower_program_structure(dispatch_schema): - program = lower_program(expand_piecewise(dispatch_schema)) - - assert list(program.parameters) == ['p_max', 'load', 'cost'], 'keyed by name, in declaration order' - ((vname, v),) = program.variables.items() +def test_lower_program_structure(dispatch_program): + assert list(dispatch_program.parameters) == ['p_max', 'load', 'cost'], 'keyed by name, in declaration order' + ((vname, v),) = dispatch_program.variables.items() assert vname == 'p' - assert v.dims == ('snapshot', 'generator') - assert v.where == ParameterComparisonNode('p_max', '>', 0.0) + assert v.dims == ('snapshot', 'generator'), 'the frame is the foreach, in the order the file wrote it' + assert v.where == Mask(P_MAX_POSITIVE) assert v.upper == Parameter('p_max') - ((cname, c),) = program.constraints.items() + ((cname, c),) = dispatch_program.constraints.items() assert cname == 'power_balance' - assert c.dims == ('snapshot',) + assert c.dims == ('snapshot',), 'the frame is the foreach, in the order the file wrote it' assert c.lhs == Sum(Variable('p'), ('generator',)) - assert c.sense == '==' + assert c.sense == '==', "the comparison crosses as the file's own operator, untranslated" assert c.rhs == Parameter('load') - assert program.objective.sense == 'minimize', "the program carries the language's spelling, untranslated" - assert program.objective.expression == Sum(Variable('p') * Parameter('cost'), ('generator', 'snapshot')), ( + assert dispatch_program.objective.sense == 'minimize', "the program carries the language's spelling, untranslated" + assert dispatch_program.objective.expression == Sum(Variable('p') * Parameter('cost'), ('generator', 'snapshot')), ( 'the objective carries the sum the file wrote, over the dims it named none of' ) @pytest.mark.parametrize('sense', [pytest.param('minimize', id='minimize'), pytest.param('maximize', id='maximize')]) def test_the_objective_sense_crosses_untranslated(sense: str): - """One spelling from the file to the program, and each sink translates at its own edge. - - A second spelling here would be two names for one axis inside one package - once the program is declared beside the language, and a stale one is not a - type error: the sense is a ``Literal``, so a program built by hand with a - retired spelling reaches a sink whose comparison quietly fails and flips - the model rather than refusing it. - - Both directions, because a translation reintroduced for one of them is what - a single case would miss. - """ - model = { - 'dimensions': {'g': {'dtype': 'str'}}, - 'parameters': {'cost': {'dims': ['g']}}, - 'variables': {'p': {'foreach': ['g'], 'bounds': {'lower': 0, 'upper': 1}}}, - 'constraints': {'c': {'foreach': [], 'expression': 'sum(p, over=g) >= 1'}}, - 'objective': {'sense': sense, 'expression': 'sum(p * cost, over=g)'}, - } - program = lower_program(expand_piecewise(Spec.model_validate(model))) + """One spelling from the file to the program, in both directions — each sink translates at its own edge.""" + program = to_program(override(TINY, objective={'sense': sense, 'expression': 'sum(p * cost, over=g)'})) assert program.objective is not None assert program.objective.sense == sense, "the file's own word for the direction, unchanged" def test_a_file_with_no_objective_lowers_to_no_sense(): """A feasibility problem has no direction, and nothing downstream invents one.""" - model = { - 'dimensions': {'g': {'dtype': 'str'}}, - 'variables': {'p': {'foreach': ['g'], 'bounds': {'lower': 0, 'upper': 1}}}, - 'constraints': {'c': {'foreach': [], 'expression': 'sum(p, over=g) >= 1'}}, - } - program = lower_program(expand_piecewise(Spec.model_validate(model))) + program = to_program(TINY) assert program.objective is None, 'no objective declared is no objective, not a minimisation of nothing' -@pytest.mark.parametrize( - ('where', 'expected'), - [ - pytest.param(None, None, id='no-where-at-all'), - pytest.param('True', None, id='True-is-no-mask'), - pytest.param('p_max', ParameterDefinedNode('p_max'), id='a-bare-parameter-name'), - pytest.param( - 'snapshot > 5', - DimensionComparisonNode('snapshot', '>', 5), - id='a-dimension-coordinate-compares-like-a-parameter', - ), - ], -) -def test_where_lowering(dispatch_schema, where, expected): - assert where_of(where, Namespace.of(dispatch_schema), 't') == expected - - def test_a_literal_amount_resolves_to_one_signed_number(dispatch_schema): """`offset=-1` parses as a unary minus over `1`; after resolution it is `-1`, for every reader alike.""" ns = Namespace.of(dispatch_schema) @@ -180,49 +165,54 @@ def test_a_literal_amount_resolves_to_one_signed_number(dispatch_schema): assert (node.kwargs['offset'], node.kwargs['edge']) == (NumberNode(-1.0), NumberNode(2.0)) -def test_a_compound_where_lowers_to_something(dispatch_schema): - assert where_of('p_max > 0 AND NOT load == 0', Namespace.of(dispatch_schema), 't') is not None - - @pytest.mark.parametrize( ('where', 'expected'), [ + pytest.param(None, None, id='no-where-at-all'), + pytest.param('True', None, id='True-is-no-mask'), + pytest.param('p_max', ParameterDefinedNode('p_max', ('generator',)), id='a-bare-parameter-name'), + pytest.param( + 'snapshot > 5', + DimensionComparisonNode('snapshot', '>', 5), + id='a-dimension-coordinate-compares-like-a-parameter', + ), + pytest.param( + 'p_max > 0 AND NOT load == 0', + AndNode(P_MAX_POSITIVE, NotNode(ParameterComparisonNode('load', '==', 0.0, ('snapshot',)))), + id='a-compound-where-keeps-its-connectives', + ), pytest.param('False', BooleanLiteralNode(False), id='the-empty-declaration-keeps-its-own-spelling'), - pytest.param('p_max > 0 AND True', ParameterComparisonNode('p_max', '>', 0.0), id='and-true-is-the-other-side'), - pytest.param('p_max > 0 OR False', ParameterComparisonNode('p_max', '>', 0.0), id='or-false-is-the-other-side'), + pytest.param('p_max > 0 AND True', P_MAX_POSITIVE, id='and-true-is-the-other-side'), + pytest.param('p_max > 0 OR False', P_MAX_POSITIVE, id='or-false-is-the-other-side'), pytest.param('p_max > 0 OR True', None, id='or-true-is-no-mask-at-all'), pytest.param('p_max > 0 AND False', BooleanLiteralNode(False), id='and-false-is-the-empty-declaration'), pytest.param('NOT True', BooleanLiteralNode(False), id='not-true-is-false'), pytest.param('NOT False', None, id='not-false-is-no-mask'), pytest.param('NOT (p_max > 0 AND False)', None, id='a-branch-folded-away-folds-the-one-above-it'), + pytest.param( + 'NOT (NOT p_max)', + ParameterDefinedNode('p_max', ('generator',)), + id='a-double-negation-cancels-on-the-load-path', + ), pytest.param( '(p_max > 0 OR True) AND load', - ParameterDefinedNode('load'), + ParameterDefinedNode('load', ('snapshot',)), id='an-absorbed-side-takes-its-own-branch-with-it', ), ], ) -def test_a_literal_is_folded_wherever_it_stands(dispatch_schema, where, expected): +def test_a_where_is_one_resolved_predicate_with_every_literal_folded(dispatch_schema, where, expected): """One mask had two lowerings: `True` was dropped at the root and kept under a connective. - So a consumer that met `where: "True"` first — no mask at all — had no - reason to expect a `BooleanLiteralNode` under an `AND`, and `p_max > 0 AND - False` reached it as a tree that only says "no rows" once someone - evaluates it. Everything decidable without data is decided at load, and - which rows a mask admits is decidable wherever a literal meets a - connective. - - The fold then lived in lowering alone, and the typesetter — reading the - same `where_of` — printed `True AND x` as written while the program said - `x`. It is resolution's now, so every reader of a mask gets one predicate. - - What the table asserts between the rows: a `BooleanLiteralNode` is a node - a consumer meets at the root or nowhere. + A `BooleanLiteralNode` is a node a consumer meets at the root or nowhere. """ - assert where_of(where, Namespace.of(dispatch_schema), 't') == expected + mask = where_of(where, Namespace.of(dispatch_schema), 't') + assert (mask.root if mask is not None else None) == expected, ( + 'the Mask carries exactly the resolved predicate, folded at resolution however the file spelled it' + ) -def test_a_folded_mask_reaches_the_declaration_the_shorter_spelling_would_have(dispatch_schema): +def test_a_folded_mask_reaches_the_declaration_the_shorter_spelling_would_have(): """The fold is the program's, not a helper's: two files, one declaration.""" written_out = lower_program( expand_piecewise(schema_of(DISPATCH_MODEL, **{'variables.p.where': 'p_max > 0 AND True'})) @@ -238,48 +228,145 @@ def test_an_unknown_where_name_is_an_error_at_lowering_too(dispatch_schema): where_of('no_such_param', Namespace.of(dispatch_schema), 't') -def test_a_lowered_mask_cannot_be_rewritten_in_place(dispatch_schema): +def test_a_lowered_mask_cannot_be_rewritten_in_place(dispatch_program): """A consumer handed a program could invert the mask another one reads. The where nodes were plain dataclasses while every declaration embedding - them was frozen, so `variable.where.op = '!='` rewrote `p_max > 0` into + them was frozen, so `variable.where.root.op = '!='` rewrote `p_max > 0` into `p_max != 0` on the shared object — two consumers disagreeing about one file, which is the failure a program exists to prevent. It also left hashability depending on the file: an unmasked declaration hashed and a masked one raised TypeError. """ - program = lower_program(expand_piecewise(dispatch_schema)) - (v,) = program.variables.values() - assert v.where == ParameterComparisonNode('p_max', '>', 0.0) + (v,) = dispatch_program.variables.values() + assert v.where == Mask(P_MAX_POSITIVE) with pytest.raises(FrozenInstanceError): - v.where.op = '!=' - assert v.where == ParameterComparisonNode('p_max', '>', 0.0), 'the mask the file wrote, unchanged' + v.where.root.op = '!=' + assert v.where == Mask(P_MAX_POSITIVE), 'the mask the file wrote, unchanged' assert isinstance(hash(v), int), 'a masked declaration hashes like an unmasked one' -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_lowered_where_is_a_mask_that_answers_from_its_root(dispatch_program): + """The `where` a lowering carries is a `Mask`, and its questions are its root's. + A consumer asks the mask — `where.names_read`, `where.conjuncts` — the way it + asks a dimension `dimension.maps`, rather than reaching for a free function + with the raw node. + """ + (v,) = dispatch_program.variables.values() -#: `fixtures.SMALL_MODEL` plus a second groupable lookup 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 -#: and two lookups over one of them. -SHAPES_MODEL = override( - SMALL_MODEL, - **{ - 'dimensions.z': {'dtype': 'str'}, - 'lookups.lk2': {'over': 'g', 'into': 'z'}, - 'parameters.lead': {'dims': ['g'], 'dtype': 'int'}, - }, + assert v.where == Mask(P_MAX_POSITIVE) + assert v.where.names_read == {'p_max'}, 'the declarations the mask names' + assert v.where.conjuncts == (P_MAX_POSITIVE,), 'a mask that is not an AND is its own only conjunct' + assert v.where.atoms == (P_MAX_POSITIVE,), 'a single leaf, connectives removed' + + +@pytest.mark.parametrize( + ('variable', 'where', 'dims', 'conjuncts', 'atoms'), + [ + pytest.param( + 'q', + "lk == 'east' and position(h) == 0", + {'g', 'h'}, + 2, + 2, + id='a-lookup-is-read-at-the-dim-it-maps-out-of-and-a-position-at-its-own', + ), + pytest.param('p', 'k > 0', set(), 1, 1, id='a-scalar-parameter-is-read-at-no-coordinate'), + pytest.param( + 'p', + 'flag and (c > 0 or k > 0)', + {'g'}, + 2, + 3, + id='atoms-cross-the-or-that-conjuncts-stop-at', + ), + ], ) +def test_a_lowered_mask_answers_its_dims_conjuncts_and_atoms(variable, where, dims, conjuncts, atoms): + """`Mask.dims` is read off the leaves, which carry their declarations' dims; + `atoms` crosses the `OR` that `conjuncts` stops at.""" + mask = to_program(override(SMALL_MODEL, **{f'variables.{variable}.where': where})).variables[variable].where + assert mask.dims == frozenset(dims) + assert len(mask.conjuncts) == conjuncts, 'an OR is one conjunct, a leaf is one conjunct' + assert len(mask.atoms) == atoms, 'the leaves of every arm, connectives removed' + + +def test_a_synthetic_predicate_answers_its_own_dims(): + """A tree built from resolved pieces answers like a declaration's own mask. + + A consumer builds region complements and conjunctions — `NotNode(root)`, + `AndNode(a, b)` — with no declaration behind them. Because the leaves carry + their dims, wrapping any such tree in `Mask` answers without a name-to-dims + mapping, which is what let the mapping die everywhere. + """ + b = ParameterDefinedNode('load', ('snapshot',)) + + assert Mask(NotNode(P_MAX_POSITIVE)).dims == {'generator'}, 'negation keeps the dims it negates' + assert (Mask(P_MAX_POSITIVE) & Mask(b)).dims == {'generator', 'snapshot'}, 'conjunction unions both sides' + assert (Mask(P_MAX_POSITIVE) & Mask(b)).root == AndNode(P_MAX_POSITIVE, b), ( + 'the conjunction joins the roots under one AND' + ) -@pytest.fixture -def shapes_schema() -> Spec: - return schema_of(SHAPES_MODEL) + +def test_mask_construction_folds_so_a_literal_stands_at_the_root_or_nowhere(): + """The fold lives in the constructor, so the invariant holds however a mask is built.""" + x = ParameterDefinedNode('committable', ('g',)) + empty, every = Mask(BooleanLiteralNode(False)), Mask(BooleanLiteralNode(True)) + + assert Mask(OrNode(BooleanLiteralNode(True), x)) == every, 'a True side absorbs the OR at the door' + assert Mask(AndNode(BooleanLiteralNode(False), x)) == empty, 'a False side dominates the AND at the door' + assert Mask(NotNode(BooleanLiteralNode(True))) == empty, 'NOT over a literal flips at the door' + assert Mask(NotNode(NotNode(x))) == Mask(x), 'a double negation cancels at the door' + + assert ~Mask(x) == Mask(NotNode(x)), 'a plain predicate negated gains one NOT' + assert ~Mask(NotNode(x)) == Mask(x), '`not (not x)` cancels rather than stacking, so no consumer evaluates it twice' + assert ~empty == every, 'the empty mask negated admits every row, with no NOT stacked' + assert ~every == empty, 'and back again' + assert empty & Mask(x) == empty, 'a False root dominates the conjunction' + assert Mask(x) & empty == empty, 'from either side' + assert every & Mask(x) == Mask(x), 'a True root is the other side' + assert Mask(x) | empty == Mask(x), 'a False root is the other side of an OR' + assert Mask(x) | every == every, 'a True root dominates the OR' + + +def test_a_held_leaf_walk_is_taken_after_the_fold_absorbed_a_branch(): + """`atoms` is held from construction, and construction folds first — so the fold's losses are not in it.""" + absorbed = Mask(AndNode(BooleanLiteralNode(False), ParameterDefinedNode('committable', ('g',)))) + + assert absorbed.root == BooleanLiteralNode(False) + assert absorbed.atoms == (), 'the absorbed leaf is not among them' + assert absorbed.names_read == frozenset(), 'nor named' + assert absorbed.dims == frozenset(), 'nor read at any dim' + + +def test_a_mask_over_an_unresolved_tree_is_refused_at_construction(): + """A tree whose leaves are unresolved is refused where it is wrapped, not where it is read.""" + with pytest.raises(AssertionError, match='reached a predicate walk unresolved'): + Mask(parse_where('a AND b')) + + +def test_an_unwritten_where_lowers_to_none_not_an_empty_mask(): + lowered = to_program(DISPATCH_MODEL) + (v,) = lowered.variables.values() + (c,) = lowered.constraints.values() + + assert v.where is None, 'no `where:` in the file means no mask, not a mask over nothing' + assert c.where is None, 'the constraint arm makes the same fold' + + +def test_a_constraint_where_is_a_mask_like_a_variable_s(): + lowered = to_program(override(DISPATCH_MODEL, **{'constraints.balance.where': 'load > 0'})) + (c,) = lowered.constraints.values() + + assert c.where == Mask(ParameterComparisonNode('load', '>', 0.0, ('snapshot',))) + + +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' @pytest.mark.parametrize( @@ -345,51 +432,27 @@ def shapes_schema() -> Spec: ], ) def test_a_construct_lowers_to_its_node(shapes_schema, expression, expected): - """Which node each surface construct becomes, and every field it arrives with. - - The nodes are frozen dataclasses, so one `==` asserts the kind and all of - `over`, `into`, `wrap`, `fill`, `partition` and `width` at once — the - fields a partial assertion skips, which is where a lowering goes astray - while still producing a node of the right kind. - """ + """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 node, so no field is asserted by omission' + assert lowered == expected, 'the whole frozen node, so no field is asserted by omission' def test_a_binary_variable_lowers_to_a_vtype(): - program = lower_program( - expand_piecewise(schema_of(DISPATCH_YAML, **{'variables.p.domain': 'binary', 'variables.p.bounds': {}})) - ) + program = to_program(schema_of(DISPATCH_YAML, **{'variables.p.domain': 'binary', 'variables.p.bounds': {}})) assert program.variable('p').variable_type == 'binary' def test_a_divisor_under_a_pullback_is_still_named(): - """`children` has to descend through every node, or a refusal loses its name. - - `divisor_parameters` is what turns "a coefficient came out null" into a - message naming the parameter the caller has to fix, and it finds those names - by walking `children`. `At` was missing from that walk, so a quotient inside - `at(...)` reported an uncovered divisor with an empty list where the name - belongs — the refusal still fired, and stopped saying what to do about it. - - Asked of the walk directly rather than through a build: the walk is static, - and a test that needed data to reach it would be testing the assembly. - """ + """`children` has to descend through every node, or a refusal loses its name.""" quotient = Divide(Variable('x'), Parameter('rate')) pulled = At(quotient, over='flow', coordinate=('component',), into=('component',)) - assert divisor_parameters(pulled) == frozenset({'rate'}) - assert divisor_parameters(Sum(pulled, ('flow',))) == frozenset({'rate'}) + assert divisor_parameters(pulled) == frozenset({'rate'}), 'the walk descends through `At`' + assert divisor_parameters(Sum(pulled, ('flow',))) == frozenset({'rate'}), 'and through a `Sum` over it' def test_a_quotient_is_found_whole_so_its_two_halves_stay_paired(): - """`divisor_parameters` flattens, and one caller cannot use the flat answer. - - A divisor has to have values wherever the row is built *and the numerator - exists*, so the eager lane narrows the mask by the variables in that - quotient's own numerator — which needs the pair, not the union. Two - quotients in one expression is the case a flattened set gets wrong. - """ + """`divisor_parameters` flattens, and one caller cannot use the flat answer.""" left = Divide(Variable('x'), Parameter('rate')) right = Divide(Variable('y'), Parameter('loss')) @@ -417,34 +480,25 @@ def test_a_quotient_is_found_whole_so_its_two_halves_stay_paired(): At(Variable('p'), over='g', coordinate=('at_bus',), into=('bus',)): 'one-to-one', Translate(Variable('p'), 't', offset=1, wrap=False, fill=0.0): 'one-to-one', Window(Variable('p'), 't', width=2, wrap=False): 'one-to-many', - Cases((Region(ParameterDefinedNode('c'), Variable('p')),)): 'one-to-one', + Cases((Region(Mask(ParameterDefinedNode('c', ('g',))), Variable('p')),)): 'one-to-one', } -def test_every_expression_node_answers_fan_in(): - """`fan_in` was a ClassVar on five nodes, so `Add(...).fan_in` was an AttributeError. - - A consumer had to keep its own list of which node kinds carry the answer, - which is the list the declaration existed to spare it. The table below is - checked for completeness against `ExpressionNode` so a node added later - fails here rather than reaching a consumer unclassified. - """ +def test_every_expression_node_is_classified_by_fan_in(): + """`fan_in` was a ClassVar on five nodes, so `Add(...).fan_in` was an AttributeError.""" covered = {type(node) for node in FAN_IN} assert covered == set(get_args(ExpressionNode)), ( 'every node in the ExpressionNode union is classified, and nothing retired lingers' ) - assert {type(node).__name__: fan_in(node) for node in FAN_IN} == { - type(node).__name__: expected for node, expected in FAN_IN.items() - } -def test_a_lookup_names_the_dimension_its_values_label(): - """Five sites asked this and each walked for it; the plan answers it once. +@pytest.mark.parametrize(('node', 'expected'), FAN_IN.items(), ids=[type(node).__name__ for node in FAN_IN]) +def test_a_node_answers_its_fan_in(node, expected): + assert fan_in(node) == expected - Both shapes are here because both had callers: one dimension's maps, for an - operator that partitions along it, and every map in the program, for - binding, which reads them all before it knows which are used. - """ + +def test_a_lookup_names_the_dimension_its_values_label(): + """Five sites asked this and each walked for it; the plan answers it once.""" program = Program( parameters={}, variables={}, @@ -467,26 +521,17 @@ def test_a_lookup_names_the_dimension_its_values_label(): def test_a_label_space_keeps_its_dtype_and_has_no_target(): - """The file's claim about a label-space column used to be dropped at lowering. - - ``period: {over: snapshot, dtype: int}`` became a bare name, so a consumer - binding the column had nothing to check it against — the one claim - ``dtype`` makes for a dimension and a parameter, missing for this column. - """ - program = lower_program( - expand_piecewise( - schema_of( - { - 'dimensions': {'snapshot': {'dtype': 'int'}, 'season': {}}, - 'lookups': { - 'season_of': {'over': 'snapshot', 'into': 'season'}, - 'period': {'over': 'snapshot', 'dtype': 'int'}, - }, - 'variables': {'p': {'foreach': ['snapshot'], 'where': 'period == 1'}}, - 'constraints': {'k': {'foreach': ['season'], 'expression': 'sum(p, by=season_of) >= 1'}}, - } - ) - ) + """The file's claim about a label-space column used to be dropped at lowering.""" + program = to_program( + { + 'dimensions': {'snapshot': {'dtype': 'int'}, 'season': {}}, + 'lookups': { + 'season_of': {'over': 'snapshot', 'into': 'season'}, + 'period': {'over': 'snapshot', 'dtype': 'int'}, + }, + 'variables': {'p': {'foreach': ['snapshot'], 'where': 'period == 1'}}, + 'constraints': {'k': {'foreach': ['season'], 'expression': 'sum(p, by=season_of) >= 1'}}, + } ) assert program.dimension('snapshot').lookups == ( @@ -498,28 +543,8 @@ def test_a_label_space_keeps_its_dtype_and_has_no_target(): def test_an_unknown_dimension_is_a_near_miss_rather_than_an_empty_declaration(): - """A typo used to return an empty declaration, silently dropping every join. - - `dimension()` answered an unknown name with `DimensionDeclaration(name)` — - no lookups, no label spaces — while its siblings `parameter()` and - `variable()` raised. So a consumer misspelling a dimension read a - declaration that mapped nowhere and placed no terms, rather than being - told. Every dimension a valid model can name is declared: `Spec`'s - reference check refuses an undeclared one in `dims:`, `foreach:`, a - lookup's `over:`/`into:` and an sos's `over:`, so the fallback was - reachable only by a mistake. - """ - program = lower_program( - expand_piecewise( - schema_of( - { - 'dimensions': {'snapshot': {}, 'generator': {}}, - 'variables': {'p': {'foreach': ['generator'], 'bounds': {'lower': 0, 'upper': 1}}}, - 'constraints': {'k': {'foreach': [], 'expression': 'sum(p, over=generator) >= 1'}}, - } - ) - ) - ) + """A typo used to return an empty declaration, silently dropping every join.""" + program = to_program(override(TINY, **{'dimensions.snapshot': {}})) assert program.dimension('snapshot').dtype == 'str', 'a declared dimension still comes back' with pytest.raises(KeyError, match='snapshto') as excinfo: @@ -528,61 +553,33 @@ def test_an_unknown_dimension_is_a_near_miss_rather_than_an_empty_declaration(): def test_a_program_is_built_by_keyword_so_a_field_added_later_cannot_reorder_an_old_call(): - """Positional construction made every field's *position* part of the contract. - - `Program(parameters, variables, constraints, objective, dimensions, sos, - expressions)` is seven positional slots on a record consumers read; a field - inserted anywhere but the end silently rebound the ones after it, with no - type error where the arguments happen to share a shape. - """ + """Positional construction made every field's *position* part of the contract.""" with pytest.raises(TypeError, match='positional'): Program({}, {}, {}, None) # pyrefly: ignore[bad-argument-count] the point of the test -def test_a_program_seals_its_declaration_groups(dispatch_schema): - """`frozen=True` sealed the fields and said nothing about what was behind them. - - `Program.expressions` was a plain dict, so a consumer could add or replace - a declaration on the program another consumer was reading — the same - two-consumers-disagree failure a mutable where node allowed. Every group is - keyed now, so the seal has to hold for all of them rather than the one. - """ - program = lower_program(expand_piecewise(dispatch_schema)) - - for group in (program.parameters, program.variables, program.constraints, program.dimensions, program.sos): - with pytest.raises(TypeError): - group['sneak'] = None # pyrefly: ignore[unsupported-operation] the point of the test - assert list(program.parameters) == ['p_max', 'load', 'cost'], "and the file's own order survives the seal" +@pytest.mark.parametrize('group', ['parameters', 'variables', 'constraints', 'dimensions', 'sos']) +def test_a_program_seals_its_declaration_groups(dispatch_program, group): + """`frozen=True` sealed the fields and said nothing about what was behind them.""" + with pytest.raises(TypeError): + getattr(dispatch_program, group)['sneak'] = None # pyrefly: ignore[unsupported-operation] the point of the test def test_expressions_are_the_ones_a_row_is_built_from(): - """`expressions` named the *declared* ones, which build no row at all. - - Two readers walk the objective and both constraint sides — `advice` and - anything asking what a solver must support — and each wrote that list out. - A named expression counted among them would answer wrongly about what is - solved, which is why the declared ones are `named_expressions` now. - """ - program = lower_program( - expand_piecewise( - schema_of( - { - 'dimensions': {'g': {}}, - 'parameters': {'cost': {'dims': ['g']}}, - 'variables': {'p': {'foreach': ['g'], 'bounds': {'lower': 0, 'upper': 1}}}, - 'constraints': {'k': {'foreach': [], 'expression': 'sum(p, over=g) >= 1'}}, - 'expressions': {'spend': 'sum(cost, over=g)'}, - 'objective': {'sense': 'minimize', 'expression': 'sum(p * cost, over=g)'}, - } - ) + """`expressions` named the *declared* ones, which build no row at all.""" + program = to_program( + override( + TINY, + expressions={'spend': 'sum(cost, over=g)'}, + objective={'sense': 'minimize', 'expression': 'sum(p * cost, over=g)'}, ) ) assert list(program.named_expressions) == ['spend'], 'the declared ones keep their own name' assert program.expressions == ( program.objective.expression, - program.constraints['k'].lhs, - program.constraints['k'].rhs, + program.constraints['c'].lhs, + program.constraints['c'].rhs, ), 'the objective first, then both sides of each constraint, in declaration order' assert program.named_expressions['spend'] not in program.expressions, ( 'a named expression builds no row, so it is not one of the expressions a row is built from' @@ -591,16 +588,11 @@ def test_expressions_are_the_ones_a_row_is_built_from(): def _footprint_of(constraint: str, objective: str) -> Footprint: - return lower_program( - expand_piecewise( - schema_of( - { - 'dimensions': {'g': {}}, - 'variables': {'p': {'foreach': ['g'], 'bounds': {'lower': 0, 'upper': 1}}}, - 'constraints': {'k': {'foreach': ['g'], 'expression': constraint}}, - 'objective': {'sense': 'minimize', 'expression': objective}, - } - ) + return to_program( + override( + TINY, + constraints={'k': {'foreach': ['g'], 'expression': constraint}}, + objective={'sense': 'minimize', 'expression': objective}, ) ).footprint @@ -612,9 +604,11 @@ def test_the_footprint_says_which_position_a_quadratic_stands_in(): actually make — quadratic is bounded "by convexity and again by what it stands beside" — and leave the sink walking the program to recover it. """ - assert _footprint_of('p <= 1', 'sum(p * p, over=g)').quadratic == {'objective'} - assert _footprint_of('p * p <= 1', 'sum(p, over=g)').quadratic == {'constraint'} - assert _footprint_of('p * p <= 1', 'sum(p * p, over=g)').quadratic == {'objective', 'constraint'} + assert _footprint_of('p <= 1', 'sum(p * p, over=g)').quadratic == {'objective'}, 'a quadratic objective alone' + assert _footprint_of('p * p <= 1', 'sum(p, over=g)').quadratic == {'constraint'}, 'a quadratic constraint alone' + assert _footprint_of('p * p <= 1', 'sum(p * p, over=g)').quadratic == {'objective', 'constraint'}, ( + 'both positions, each named' + ) assert _footprint_of('p <= 1', 'sum(p, over=g)').quadratic == frozenset(), 'affine throughout is the empty set' @@ -627,7 +621,7 @@ def test_a_construct_the_file_does_not_use_is_an_empty_set_rather_than_none(): footprint = _footprint_of('p <= 1', 'sum(p, over=g)') assert footprint.sos_types == frozenset(), 'a file declaring no sos' - assert footprint.quadratic == frozenset() + assert footprint.quadratic == frozenset(), 'a file with no quadratic anywhere' assert footprint.variable_types == {'continuous'}, 'never empty — a program has variables' assert {type(f) for f in (footprint.sos_types, footprint.quadratic, footprint.shapes)} == {frozenset}, ( 'every field is a set, so one rule reads all of them' @@ -638,27 +632,14 @@ def test_a_construct_the_file_does_not_use_is_an_empty_set_rather_than_none(): ) -def test_the_footprint_is_walked_once_and_held(dispatch_schema): +def test_the_footprint_is_walked_once_and_held(dispatch_program): """Safe to hold only because the program cannot change under it.""" - program = lower_program(expand_piecewise(dispatch_schema)) - assert program.footprint is program.footprint + assert dispatch_program.footprint is dispatch_program.footprint def test_a_named_expression_is_not_in_the_footprint(): """It builds no row, so counting it would answer wrongly about what is solved.""" - program = lower_program( - expand_piecewise( - schema_of( - { - 'dimensions': {'g': {}}, - 'parameters': {'cost': {'dims': ['g']}}, - 'variables': {'p': {'foreach': ['g'], 'bounds': {'lower': 0, 'upper': 1}}}, - 'constraints': {'k': {'foreach': [], 'expression': 'sum(p, over=g) >= 1'}}, - 'expressions': {'spend': 'sum(p * cost, over=g)'}, - } - ) - ) - ) + program = to_program(override(TINY, expressions={'spend': 'sum(p * cost, over=g)'})) assert Parameter not in program.footprint.shapes, "the named expression's parameter reaches no row" assert Parameter in {type(n) for n in walk(program.named_expressions['spend'])}, 'though it is in the expression' @@ -668,18 +649,9 @@ def test_a_dimension_carries_the_dtype_its_labels_are_checked_against(): """The declared type travels with the dimension, as a parameter's does. A dimension is read from whatever table carries it, so nothing downstream - can infer what the column should have been — ``sources.py`` checks the - labels against this and has no other way to know. + can infer what the column should have been. """ - schema = schema_of( - { - 'dimensions': {'t': {'dtype': 'int'}, 'g': {}}, - 'parameters': {'c': {'dims': ['g']}}, - 'variables': {'p': {'foreach': ['g'], 'bounds': {'lower': 0, 'upper': 1}}}, - 'constraints': {'k': {'foreach': [], 'expression': 'sum(p, over=g) >= 1'}}, - } - ) - program = lower_program(expand_piecewise(schema)) + program = to_program(override(TINY, **{'dimensions.t': {'dtype': 'int'}})) assert program.dimension('t').dtype == 'int', 'a declared dtype reaches the plan' assert program.dimension('g').dtype == 'str', "and the schema's default does too, rather than nothing" @@ -703,18 +675,17 @@ def test_a_dimension_carries_the_dtype_its_labels_are_checked_against(): } -def _cases_in(program: Program) -> program_module.Cases: +def _cases_in(program: Program) -> Cases: """The one cased node the fixture's constraint carries.""" sides = [side for c in program.constraints.values() for side in (c.lhs, c.rhs)] - found = [n for n in walk(*sides) if isinstance(n, program_module.Cases)] + found = [n for n in walk(*sides) if isinstance(n, Cases)] assert len(found) == 1, 'the fixture has exactly one cased expression, inlined where it is named' return found[0] def test_a_cased_expression_lowers_to_one_region_per_case(): """The regions come out in file order, values lowered like any other expression.""" - program = lower_program(expand_piecewise(schema_of(CASED))) - cases = _cases_in(program) + cases = _cases_in(to_program(CASED)) assert len(cases.regions) == 3, 'one region per case, the `otherwise` among them' assert [type(r.value).__name__ for r in cases.regions] == ['Constant', 'Parameter', 'Translate'], ( @@ -728,15 +699,33 @@ def test_the_fallback_region_carries_the_mask_the_file_left_unwritten(): A consumer adds regions rather than working out which one is left over, so the remainder is resolved once here instead of once per consumer. """ - program = lower_program(expand_piecewise(schema_of(CASED))) - remainder = _cases_in(program).regions[-1] + remainder = _cases_in(to_program(CASED)).regions[-1] - assert isinstance(remainder.when, AndNode), 'two stated cases, so the remainder is a conjunction of two negations' - assert remainder.when.left == ParameterDefinedNode('committable'), ( + assert isinstance(remainder.when.root, AndNode), ( + 'two stated cases, so the remainder is a conjunction of two negations' + ) + assert remainder.when.root.left == ParameterDefinedNode('committable', ('g',)), ( 'the negation of `not committable` is the term itself, not a second `not` around it' ) +def test_a_region_s_when_is_a_mask_with_its_own_dims(): + """`Region.when` arrives in the same carrier as a declaration's `where`. + + It was the one mask left as a bare node, so a helper written over `Mask` + branched on where a mask came from — the divergence the carrier exists to + prevent. The synthesized remainder gets its dims like any stated case. + """ + always_on, boundary, remainder = _cases_in(to_program(CASED)).regions + + assert all(isinstance(r.when, Mask) for r in (always_on, boundary, remainder)), ( + 'every region, the synthesized remainder included, carries its predicate as a Mask' + ) + assert always_on.when.dims == frozenset({'g'}), "`not committable` reads the parameter's dims" + assert boundary.when.dims == frozenset({'g', 't'}), 'the position comparison adds its dimension' + assert remainder.when.dims == frozenset({'g', 't'}), 'the remainder reads every dim the stated cases do' + + def test_the_lowered_regions_are_still_proved_apart(): """The remainder does not collide with the cases it was built from. @@ -745,16 +734,16 @@ def test_the_lowered_regions_are_still_proved_apart(): same prover, against each stated case, and must overlap none of them. """ spec = schema_of(CASED) - regions = _cases_in(lower_program(expand_piecewise(spec))).regions - named = {f'region{i}': r.when for i, r in enumerate(regions)} + regions = _cases_in(to_program(spec)).regions + named = {f'region{i}': r.when.root for i, r in enumerate(regions)} - assert list(overlapping(named, spec)) == [], 'no two lowered regions can claim one coordinate' + assert list(overlapping(named, Namespace.of(spec).dtypes)) == [], 'no two lowered regions can claim one coordinate' def test_a_cased_expression_is_readable_by_the_name_the_file_wrote(): """`Program.expressions` carries it, so a consumer reads it back whole.""" - program = lower_program(expand_piecewise(schema_of(CASED))) + program = to_program(CASED) - assert isinstance(program.named_expressions['previous'], program_module.Cases), ( + assert isinstance(program.named_expressions['previous'], Cases), ( 'a cased expression reaches the program as the node, not as its fallback arm alone' ) diff --git a/tests/test_parser.py b/tests/test_parser.py index b7425071..1c58462c 100644 --- a/tests/test_parser.py +++ b/tests/test_parser.py @@ -8,8 +8,18 @@ ``NameNode``/``Unresolved*`` nodes. """ +import operator +from dataclasses import FrozenInstanceError + import pytest +import math_spec.program as program_module +from math_spec._where_parser import ( + UnresolvedComparisonNode, + UnresolvedNameNode, + UnresolvedPositionNode, + parse_where, +) from math_spec.errors import SchemaError from math_spec.expression_parser import ( BinaryOperatorNode, @@ -21,29 +31,34 @@ UnaryOperatorNode, parse_expression, ) -from math_spec.where_parser import ( - AndNode, - BooleanLiteralNode, - NotNode, - OrNode, - UnresolvedComparisonNode, - UnresolvedNameNode, - UnresolvedPositionNode, - parse_where, -) +from math_spec.program import AndNode, BooleanLiteralNode, NotNode, OrNode, _conjuncts + + +def test_the_grammar_builds_the_program_s_own_node_classes(): + """The connectives and literals in a parse are `math_spec.program`'s classes. + + The parser constructs the resolved vocabulary's connectives directly, so a + consumer's `isinstance` against the program's classes holds on any tree — + two homes for `AndNode` would make it hold on neither. + """ + tree = parse_where('a AND NOT b OR True') + assert type(tree) is program_module.OrNode + assert type(tree.left) is program_module.AndNode + assert type(tree.left.right) is program_module.NotNode + assert type(tree.right) is program_module.BooleanLiteralNode @pytest.mark.parametrize( ('text', 'node_type', 'attrs'), [ - ('42', NumberNode, {'value': 42}), - ('3.14', NumberNode, {'value': pytest.approx(3.14)}), - ('p_max', NameNode, {'name': 'p_max'}), - ('a + b', BinaryOperatorNode, {'op': '+'}), - ('-x', UnaryOperatorNode, {'op': '-'}), - ('p <= p_max', ComparisonNode, {'op': '<='}), - ('sum(p, over=g) == load', ComparisonNode, {'op': '=='}), - ('sum(p, over=generator)', FunctionCallNode, {'name': 'sum'}), + pytest.param('42', NumberNode, {'value': 42}, id='an-integer'), + pytest.param('3.14', NumberNode, {'value': pytest.approx(3.14)}, id='a-decimal'), + pytest.param('p_max', NameNode, {'name': 'p_max'}, id='a-name'), + pytest.param('a + b', BinaryOperatorNode, {'op': '+'}, id='a-binary-operator'), + pytest.param('-x', UnaryOperatorNode, {'op': '-'}, id='a-unary-operator'), + pytest.param('p <= p_max', ComparisonNode, {'op': '<='}, id='a-comparison'), + pytest.param('sum(p, over=g) == load', ComparisonNode, {'op': '=='}, id='a-comparison-over-a-call'), + pytest.param('sum(p, over=generator)', FunctionCallNode, {'name': 'sum'}, id='a-call'), ], ) def test_an_expression_parses_to_its_node(text, node_type, attrs): @@ -84,13 +99,69 @@ def test_precedence(text, tree): def test_a_call_carries_its_positional_and_keyword_arguments(): node = parse_expression('sum(p * cost, over=generator)') - assert len(node.args) == 1 + assert len(node.args) == 1, 'one positional argument; the keyword is not among them' assert isinstance(node.args[0], BinaryOperatorNode), 'the argument is an expression, not just a name' assert 'over' in node.kwargs -def test_an_unparseable_expression_is_an_error(): - with pytest.raises(SchemaError, match='Failed to parse'): +@pytest.mark.parametrize( + ('rewrite', 'error', 'match'), + [ + pytest.param( + lambda node: setattr(node, 'op', '>='), FrozenInstanceError, 'cannot assign', id='a-comparison-sense' + ), + pytest.param( + lambda node: operator.setitem(node.left.kwargs, 'over', NameNode('snapshot')), + TypeError, + 'does not support item assignment', + id='a-reduction-axis', + ), + ], +) +def test_a_parsed_expression_cannot_be_rewritten_under_another_pass(rewrite, error, match): + """A pass handed a parsed tree could rewrite the operand another one reads. + + The expression nodes were plain dataclasses while every where and program + node was frozen (#197): `node.op = '<='` flipped a shared comparison and + `node.kwargs['over'] = ...` re-aimed a reduction, with no error anywhere. + """ + node = parse_expression('sum(p * cost, over=generator) == load') + with pytest.raises(error, match=match): + rewrite(node) + + +def test_a_call_copies_the_kwargs_it_is_handed(): + """A caller's own dict is copied on the way in, so holding it is not a back door either.""" + passed = {'over': NameNode('generator')} + built = FunctionCallNode('sum', (NameNode('p'),), passed) + passed['over'] = NameNode('snapshot') + assert built.kwargs == {'over': NameNode('generator')}, 'the dict handed in was copied, not aliased' + assert isinstance(hash(built), int), 'kwargs sits outside the hash, so a call hashes like every other node' + + +@pytest.mark.parametrize( + ('text', 'rewrite'), + [ + pytest.param('p < p_max', r'the senses are <=, >= and ==\. Write the bound inclusive', id='strict-less'), + pytest.param('p > 0', r'the senses are <=, >= and ==\. Write the bound inclusive', id='strict-greater'), + pytest.param('status != 0', r'write the test in where:, where != is legal', id='not-equals-is-a-where-matter'), + pytest.param('p = p_max', r'Equality between two sides is written ==', id='a-lone-equals'), + pytest.param('p ^ 2', r"power is written '\*\*', not '\^'", id='caret-for-power'), + pytest.param('0 <= p <= p_max', r'Split the chain into two constraints', id='a-chained-comparison'), + ], +) +def test_a_parse_failure_names_the_rewrite(text, rewrite): + """The predictable mistakes are refused with their rewrite, not the grammar's complaint alone.""" + with pytest.raises(SchemaError, match=rewrite): + parse_expression(text) + + +@pytest.mark.parametrize( + 'fragment', + [pytest.param('Failed to parse', id='the-refusal'), pytest.param('Expected', id='the-grammar-s-complaint')], +) +def test_a_failure_with_no_diagnosis_still_shows_the_grammar_s_complaint(fragment): + with pytest.raises(SchemaError, match=fragment): parse_expression('a +') @@ -123,20 +194,19 @@ def test_a_list_of_names_is_a_kwarg_value(): ], ) def test_a_list_the_grammar_cannot_read_is_refused_at_load(text): - """A list is a kwarg value and nothing else, and the last three say so. - - Which is a claim about the *grammar*: a list admitted as a term would be - a second thing `[a, b]` could mean, and one read past a missing comma - would be a grouping the file does not write. Neither is decidable later — - a parse is what every consumer starts from. - """ + """A list is a kwarg value and nothing else, and the last three say so.""" with pytest.raises(SchemaError, match='Failed to parse expression'): parse_expression(text) @pytest.mark.parametrize( ('text', 'value'), - [('1e5', 1e5), ('2.5E-3', 2.5e-3), ('1e+3', 1e3), ('7.e2', 700.0)], + [ + pytest.param('1e5', 1e5, id='a-bare-exponent'), + pytest.param('2.5E-3', 2.5e-3, id='an-upper-case-negative-exponent'), + pytest.param('1e+3', 1e3, id='a-signed-exponent'), + pytest.param('7.e2', 700.0, id='a-trailing-point-mantissa'), + ], ) def test_scientific_notation_is_a_number(text, value): assert parse_expression(text) == NumberNode(value) @@ -158,12 +228,12 @@ def test_a_name_may_begin_with_inf(name): @pytest.mark.parametrize( ('text', 'node_type', 'attrs'), [ - ('True', BooleanLiteralNode, {'value': True}), - ('p_max', UnresolvedNameNode, {'name': 'p_max'}), - ('p_max > 0', UnresolvedComparisonNode, {'op': '>', 'value': 0}), - ('a AND b', AndNode, {}), - ('a OR b', OrNode, {}), - ('NOT a', NotNode, {}), + pytest.param('True', BooleanLiteralNode, {'value': True}, id='a-literal'), + pytest.param('p_max', UnresolvedNameNode, {'name': 'p_max'}, id='a-bare-name'), + pytest.param('p_max > 0', UnresolvedComparisonNode, {'op': '>', 'value': 0}, id='a-comparison'), + pytest.param('a AND b', AndNode, {}, id='and'), + pytest.param('a OR b', OrNode, {}, id='or'), + pytest.param('NOT a', NotNode, {}, id='not'), ], ) def test_a_where_string_parses_to_its_node(text, node_type, attrs): @@ -181,6 +251,33 @@ def test_and_binds_tighter_than_or(): ) +@pytest.mark.parametrize( + ('text', 'expected'), + [ + ('a', ['a']), + ('a AND b', ['a', 'b']), + ('a AND b AND c', ['a', 'b', 'c']), + ], + ids=['single', 'pair', 'chain'], +) +def test_conjuncts_flattens_the_and_spine(text, expected): + """A chain the grammar left-folds into nested `AndNode`s comes back flat (#312). + + `_conjuncts` is the one home of the flatten rule; `Mask.conjuncts` is the + door a consumer asks it through.""" + assert [n.name for n in _conjuncts(parse_where(text))] == expected + + +@pytest.mark.parametrize( + 'text', + ['a OR b', 'NOT a', 'a AND b OR c', 'NOT (a AND b)'], + ids=['or', 'not', 'or-of-and', 'not-of-and'], +) +def test_conjuncts_does_not_split_or_or_not(text): + result = _conjuncts(parse_where(text)) + assert result == (parse_where(text),), 'a non-AND top node is its own only conjunct' + + @pytest.mark.parametrize( ('text', 'value', 'quoted'), [ @@ -230,15 +327,45 @@ def test_a_position_is_not_confused_with_a_name(): assert isinstance(parse_where('position(t) == 0 AND p_max > 0'), AndNode) -def test_the_old_index_spelling_names_its_rewrite(): - """`index(dim, i)` is what every model wrote before #32.""" - with pytest.raises(SchemaError) as excinfo: - parse_where('snapshot == index(snapshot, 0)') - assert 'index() is now position()' in str(excinfo.value) - assert "write 'position(dim) == i'" in str(excinfo.value) +@pytest.mark.parametrize( + ('text', 'rewrite'), + [ + pytest.param('p_max > 0 & committable', r'written AND', id='ampersand'), + pytest.param('p_max > 0 && committable', r'written AND', id='doubled-ampersand'), + pytest.param('p_max > 0 | committable', r'written OR', id='pipe'), + pytest.param('p_max > 0 || committable', r'written OR', id='doubled-pipe'), + pytest.param('~committable', r'written NOT, before the predicate', id='tilde'), + pytest.param('!committable', r'written NOT, before the predicate', id='bang'), + pytest.param('status = 0', r'equality is written ==', id='a-lone-equals'), + ], +) +def test_a_where_parse_failure_names_the_rewrite(text, rewrite): + """The connective habits of pandas and C are refused with their rewrite, not the grammar's complaint alone.""" + with pytest.raises(SchemaError, match=rewrite): + parse_where(text) + + +def test_a_legal_where_operator_is_never_diagnosed(): + """`!=`, `<` and `>` are predicates here, unlike on the expression side, so no diagnosis may fire on them.""" + assert parse_where('status != 0') == UnresolvedComparisonNode('status', '!=', 0.0) + assert parse_where('p_max < 5') == UnresolvedComparisonNode('p_max', '<', 5.0) def test_an_unrelated_parse_failure_says_nothing_about_positions(): with pytest.raises(SchemaError) as excinfo: parse_where('p_max >') assert 'position()' not in str(excinfo.value) + + +def test_a_string_parses_to_one_shared_tree(): + """Drop the memo and this passes on `==` alone — `is` is the claim.""" + text = 'sum(p * cost, over=generator) == load' + assert parse_expression(text) is parse_expression(text), 'the same expression string parses to one tree' + assert parse_where('p_max > 0') is parse_where('p_max > 0'), 'and so does the same where string' + + +def test_a_parse_failure_is_raised_every_time_it_is_asked_for(): + """A memo that cached the failure would hand the second caller a traceback from the first.""" + for _ in range(2): + with pytest.raises(SchemaError, match=r"power is written '\*\*'"): + parse_expression('p ^ 2') diff --git a/tests/test_piecewise.py b/tests/test_piecewise.py index 50be3b79..323a1ed4 100644 --- a/tests/test_piecewise.py +++ b/tests/test_piecewise.py @@ -82,6 +82,8 @@ 'variables.running': {'foreach': ['snapshot'], 'domain': 'binary'}, }, ) +#: The ``lp`` curve masked by one of its own values-parameters, so every check a block can carry is on it. +LP_MASKED = override(LP, **{'piecewise.cost_curve.points': 'bp_x'}) #: Two dims in the frame, so the emitted ``foreach`` has an order to get wrong. TWO_DIM = override( raw_of(NONCONVEX_YAML), @@ -100,7 +102,7 @@ def test_expansion_emits_the_lambda_declarations(): expanded = expand_piecewise(schema_of(NONCONVEX_YAML)) - assert not expanded.piecewise + assert not expanded.piecewise, 'the block is spent once its declarations are emitted' assert 'cost_curve_lam' in expanded.variables assert expanded.variables['cost_curve_seg'].domain == 'binary' assert set(expanded.constraints) >= { @@ -110,7 +112,7 @@ def test_expansion_emits_the_lambda_declarations(): 'cost_curve_link0', 'cost_curve_link1', 'balance', - } + }, "the adjacency formulation's five rows, one link each, beside the constraint the file wrote" def test_an_emitted_set_may_not_collide_with_a_declared_one(): @@ -134,10 +136,6 @@ def test_the_file_is_not_an_expansion_and_the_expansion_is(): assert isinstance(expand_piecewise(schema), _ExpandedSpec) -def test_a_curve_stated_as_lines_expands_to_an_expanded_spec(): - assert isinstance(expand_piecewise(schema_of(LP)), _ExpandedSpec) - - def test_expansion_is_memoised_and_idempotent(): """One object from every call: validation already built the expansion, and an `_ExpandedSpec` is its own.""" schema = schema_of(NONCONVEX_YAML) @@ -158,7 +156,13 @@ def test_an_expansion_will_not_be_built_around_a_curve(): _ExpandedSpec.model_validate(raw_of(NONCONVEX_YAML)) -@pytest.mark.parametrize('order', [['snapshot', 'generator', 'bp'], ['generator', 'snapshot', 'bp']]) +@pytest.mark.parametrize( + 'order', + [ + pytest.param(['snapshot', 'generator', 'bp'], id='snapshot-first'), + pytest.param(['generator', 'snapshot', 'bp'], id='generator-first'), + ], +) def test_the_emitted_foreach_follows_declaration_order(order): """The frame is a set until something orders it, and a set iterates the same way for the same names within one process — so a run that reads the @@ -261,8 +265,8 @@ def test_a_malformed_block_is_refused(model, patch, match): @pytest.mark.parametrize( ('link_expression', 'message'), [ - ('p ** 2', 'over variables'), - ('p * p', 'both factors of a product contain variables'), + pytest.param('p ** 2', 'over variables', id='a-power-of-a-variable'), + pytest.param('p * p', 'both factors of a product contain variables', id='a-product-of-variables'), ], ) def test_a_link_outside_the_language_is_named_where_the_user_wrote_it(link_expression, message): @@ -308,9 +312,7 @@ def test_a_gate_that_is_not_a_variable_is_refused(activity, match): @pytest.mark.parametrize(('raw', 'expected'), _CURVATURE_CASES) def test_a_method_names_the_curvature_it_is_exact_for(raw, expected): - """The consumer holding the breakpoints checks the shape; this says what to - check for. It is the block's own semantics, so it is answered here rather - than re-derived by every repository that binds data to a curve.""" + """The consumer holding the breakpoints checks the shape; this says what to check for.""" answer = next((c.curvature for c in to_program(raw).piecewise['cost_curve'].checks if isinstance(c, Curved)), None) assert answer == expected assert answer is None or answer in CURVATURES, ( @@ -337,7 +339,7 @@ def test_an_emitted_parameter_says_how_it_is_filled(): derivation that fills it, and every parameter the file declared carries none. """ - program = lower_program(expand_piecewise(schema_of(LP, **{'piecewise.cost_curve.points': 'bp_x'}))) + program = lower_program(expand_piecewise(schema_of(LP_MASKED))) assert {n: p.derivation for n, p in program.parameters.items() if p.derivation is not None} == { 'cost_curve_points': MaskOf('cost_curve', 'bp_x'), @@ -362,12 +364,8 @@ def test_a_file_supplied_mask_derives_nothing(): def test_a_block_is_kept_as_the_checks_a_consumer_binding_it_runs(): - """What a block assumes of its numbers used to be readable only off the file. - - A consumer asserted the data conditions in its own words, against names - it re-spelled; each condition now arrives carrying its own subjects. - """ - curve = to_program(override(LP, **{'piecewise.cost_curve.points': 'bp_x'})).piecewise['cost_curve'] + """Every condition a curve puts on its data arrives carrying its own subjects.""" + curve = to_program(LP_MASKED).piecewise['cost_curve'] assert curve.breakpoints == ('bp_x', 'bp_y'), 'the values parameters, in link order' assert set(curve.checks) == { @@ -383,7 +381,7 @@ def test_a_block_is_kept_as_the_checks_a_consumer_binding_it_runs(): @pytest.mark.parametrize('kind', get_args(Check), ids=lambda k: k.__name__) def test_every_check_has_a_sentence(kind): - curve = to_program(override(LP, **{'piecewise.cost_curve.points': 'bp_x'})).piecewise['cost_curve'] + curve = to_program(LP_MASKED).piecewise['cost_curve'] check = next((c for c in curve.checks if isinstance(c, kind)), None) assert check is not None, 'the fixture is the block that assumes everything' assert check_message('cost_curve', curve, check).startswith("piecewise 'cost_curve':") diff --git a/tests/test_program_nodes.py b/tests/test_program_nodes.py index f1db565e..04c348e2 100644 --- a/tests/test_program_nodes.py +++ b/tests/test_program_nodes.py @@ -4,11 +4,8 @@ """Every node a `Program` can carry is one some file actually lowers to. -The sibling of `test_the_golden_model_carries_every_node_kind_the_walk_renders`, -one state along. That one holds the *renderer* to the AST; this holds the -*lowering* to the program, and the two cannot share a fixture: rendering -accepts more than lowering does, so the golden model carries a shift over a -variable-free expression that lowering refuses outright. +The lowering-side sibling of `test_the_golden_model_carries_every_node_kind_the_walk_renders`, +on a fixture of its own because rendering accepts what lowering refuses. Without this, a node can join `ExpressionNode` with nothing producing it and the suite stays green — `assert_never` fires only where some test happens to @@ -21,6 +18,8 @@ from pathlib import Path from typing import get_args +import pytest + import math_spec as ms from math_spec.program import ExpressionNode, Program, walk @@ -41,22 +40,25 @@ def _expressions(program: Program) -> list[ExpressionNode]: return trees -def test_every_program_node_is_one_some_file_lowers_to(): - """A node nothing produces is a node no consumer has been asked to build.""" +@pytest.fixture(scope='module') +def kinds() -> tuple[set[str], set[str]]: + """The node classes the fixture lowers to, and the ones `ExpressionNode` declares.""" program = ms.to_program(FIXTURE) reached = {type(node).__name__ for node in walk(*_expressions(program))} declared = {node.__name__ for node in get_args(ExpressionNode)} + return reached, declared + +def test_every_program_node_is_one_some_file_lowers_to(kinds): + """A node nothing produces is a node no consumer has been asked to build.""" + reached, declared = kinds assert declared <= reached, ( f'{FIXTURE.name} lowers to none of {sorted(declared - reached)}. A node no file reaches is ' f'one whose lowering nobody has run — add a declaration using the construct it stands for.' ) -def test_the_fixture_carries_nothing_the_program_has_no_node_for(): +def test_the_fixture_carries_nothing_the_program_has_no_node_for(kinds): """The other direction, so the fixture cannot drift into asserting nothing.""" - program = ms.to_program(FIXTURE) - reached = {type(node).__name__ for node in walk(*_expressions(program))} - declared = {node.__name__ for node in get_args(ExpressionNode)} - + reached, declared = kinds assert reached <= declared, f'{FIXTURE.name} lowers to {sorted(reached - declared)}, which is not a program node' diff --git a/tests/test_public_surface.py b/tests/test_public_surface.py index 039ca765..8761dc45 100644 --- a/tests/test_public_surface.py +++ b/tests/test_public_surface.py @@ -13,6 +13,8 @@ import ast from pathlib import Path +import pytest + import math_spec from math_spec import program, typesetting @@ -39,6 +41,13 @@ } ) # fmt: skip +#: The modules whose `__all__` a consumer imports from. +MODULES = [ + pytest.param(math_spec, id='math_spec'), + pytest.param(typesetting, id='typesetting'), + pytest.param(program, id='program'), +] + def test_all_matches_the_pinned_surface(): """Both directions, because either alone rots.""" @@ -48,20 +57,17 @@ def test_all_matches_the_pinned_surface(): ) -def test_every_exported_name_is_bound(): - """`__all__` naming something the package does not bind is a broken import.""" - missing = sorted(n for n in math_spec.__all__ if not hasattr(math_spec, n)) - assert not missing, f'__all__ names unbound attributes: {missing}' - +@pytest.mark.parametrize('module', MODULES) +def test_every_exported_name_is_bound(module): + """`__all__` naming something the module does not bind is a broken import.""" + missing = sorted(n for n in module.__all__ if not hasattr(module, n)) + assert not missing, f'{module.__name__}.__all__ names unbound attributes: {missing}' -def test_all_names_nothing_twice(): - names = list(math_spec.__all__) - assert len(names) == len(set(names)), 'duplicate name in __all__' - -def test_the_typeset_subpackage_binds_what_it_exports(): - missing = sorted(n for n in typesetting.__all__ if not hasattr(typesetting, n)) - assert not missing, f'math_spec.typesetting.__all__ names unbound attributes: {missing}' +@pytest.mark.parametrize('module', MODULES) +def test_all_names_nothing_twice(module): + names = list(module.__all__) + assert len(names) == len(set(names)), f'duplicate name in {module.__name__}.__all__' def _defined_by(module: object) -> set[str]: @@ -84,25 +90,10 @@ def _defined_by(module: object) -> set[str]: def test_the_program_module_exports_everything_it_defines(): - """`math_spec.__all__` exports the *module*, so this is the consumers' surface. - - Without an `__all__` the module's namespace was the surface, which made - every import it happens to make — `dataclass`, `Mapping`, `Literal` — - part of what a consumer could reach. Both directions, so a public name - added without a decision fails here rather than shipping unnoticed. - """ + """`math_spec.__all__` exports the module, so this is the consumers' surface — both + directions, so a public name added without a decision fails here.""" declared = set(program.__all__) defined = _defined_by(program) assert declared == defined, ( f'only in __all__: {sorted(declared - defined)}; defined but unexported: {sorted(defined - declared)}' ) - - -def test_the_program_module_binds_what_it_exports(): - missing = sorted(n for n in program.__all__ if not hasattr(program, n)) - assert not missing, f'math_spec.program.__all__ names unbound attributes: {missing}' - - -def test_the_program_module_names_nothing_twice(): - names = list(program.__all__) - assert len(names) == len(set(names)), 'duplicate name in math_spec.program.__all__' diff --git a/tests/test_pypsa_references.py b/tests/test_pypsa_references.py index e7491f96..fbe86c64 100644 --- a/tests/test_pypsa_references.py +++ b/tests/test_pypsa_references.py @@ -4,16 +4,8 @@ """What the PyPSA references pin the model files to, without any engine. -This repository holds the corpus — the model files, the reference networks -as PyPSA scripts with the data inline, and `references.json`: what PyPSA -solved each of them to, checked by the `PyPSA references` workflow. Everything here asserts over those -committed files alone: names both directions between what PyPSA built and -what the files declare, a record per rung from the pinned pypsa, generic -spine weightings. What an engine makes of the rungs — one objective across -the fence, coverage, a model-for-model verdict — is that engine's own record -and its own tests (lpspec keeps both under `differential/pypsa/`). The page -blocks the records feed are held current by ``tests/test_docs.py`` through -``tools.gallery``. +Everything here asserts over the committed model files, reference scripts and +`references.json` alone. """ from __future__ import annotations @@ -33,7 +25,9 @@ SCRIPT = REFERENCES / 'reference.py' PAGE_TEXTS = [(gallery.PAGES / page).read_text() for page in DECLARED] -MODELS = [to_spec(path) for path in DECLARED.values()] +SPECS = {page: to_spec(path) for page, path in DECLARED.items()} +MODELS = list(SPECS.values()) +BASE = SPECS['pypsa.md'] ROWS_DECLARED = {_stands_for(name, block.description) for m in MODELS for name, block in m.constraints.items()} COLUMNS_DECLARED = {_stands_for(name, block.description) for m in MODELS for name, block in m.variables.items()} #: The five GlobalConstraint formulas open with their *type* — PyPSA names @@ -115,13 +109,10 @@ def _stated(name: str, row: str) -> bool: return re.fullmatch(re.sub(r'\\\{[a-z]\\\}', '.+', re.escape(name)), row) is not None -BASE = to_spec(DECLARED['pypsa.md']) - - @pytest.mark.parametrize('page', [page for page in DECLARED if page != 'pypsa.md']) def test_a_file_of_its_own_shares_its_declarations_with_the_base(page: str): """A keyword file restates the base surface; a shared name keeps its PyPSA name and its dtype, or it has drifted.""" - own = to_spec(DECLARED[page]) + own = SPECS[page] drifted = [] for section in ('parameters', 'lookups', 'variables', 'constraints'): theirs, ours = getattr(BASE, section), getattr(own, section) diff --git a/tests/test_reading_page.py b/tests/test_reading_page.py index f426db44..0117f308 100644 --- a/tests/test_reading_page.py +++ b/tests/test_reading_page.py @@ -51,10 +51,10 @@ def test_the_page_shows_the_declarations_the_expansion_emits(tmp_path, monkeypat namespace: dict[str, object] = {} claims: list[tuple[str, object]] = [] for code in _blocks('python'): - # the page is the input, so running it is the point rather than a smell + # the page is the input, so running it is the point exec(compile(code, str(PAGE), 'exec'), namespace) claims.extend(_claims(code)) - assert len(claims) == 7, 'every `expression # value` line on the page is checked; one without one is not' + assert len(claims) == 10, 'every `expression # value` line on the page is checked; one without one is not' for expression, claimed in claims: assert eval(expression, namespace) == claimed, f'reading.md says `{expression}` is {claimed}' diff --git a/tests/test_schema.py b/tests/test_schema.py index ec0d5e6a..d4abbf06 100644 --- a/tests/test_schema.py +++ b/tests/test_schema.py @@ -7,8 +7,7 @@ `schema/math-spec.schema.json` is a generated artefact that ships in the repository so an editor can offer completion without importing the package. Nothing regenerates it on the way to a release, so the only thing keeping it -equal to the models is this file. What a validation rule *means* is -`test_validation.py`'s business; this one is about the artefact. +equal to the models is this file. """ import json @@ -26,18 +25,20 @@ def test_the_checked_in_json_schema_has_not_drifted(): ) -def test_the_json_schema_admits_what_the_loader_admits(): - """The two shorthands live in before-validators, which pydantic's generated - schema cannot see — each needs its own schema hook in model.py, and losing a - hook loses the shorthand from every editor silently.""" - doc = json.loads(schema.PATH.read_text()) - link = doc['$defs']['PiecewiseLink'] - assert any(form.get('type') == 'array' for form in link.get('anyOf', [])), ( - 'the schema lost the `[expression, values, sign?]` link shorthand the loader accepts' - ) - expression = doc['$defs']['ExpressionBlock'] - assert {'type': 'string'} in expression.get('anyOf', []), ( - 'the schema lost the bare-string form a named expression is written in' +@pytest.mark.parametrize( + ('definition', 'shorthand', 'spelling'), + [ + pytest.param('PiecewiseLink', 'array', '`[expression, values, sign?]`', id='link-shorthand'), + pytest.param('ExpressionBlock', 'string', 'bare-string', id='bare-string-expression'), + ], +) +def test_the_json_schema_admits_the_shorthand_the_loader_admits(definition, shorthand, spelling): + """A shorthand lives in a before-validator, which pydantic's generated schema + cannot see — each needs its own schema hook in model.py, and losing the hook + loses the shorthand from every editor silently.""" + forms = json.loads(schema.PATH.read_text())['$defs'][definition].get('anyOf', []) + assert any(form.get('type') == shorthand for form in forms), ( + f'the schema lost the {spelling} form of {definition} the loader accepts' ) @@ -47,7 +48,9 @@ def test_no_definition_refers_only_to_itself(): Rendered rather than read from the file, so it fails on whichever pydantic is installed.""" doc = json.loads(schema.rendered()) for name, entry in doc['$defs'].items(): - assert {'$ref': f'#/$defs/{name}'} not in entry.get('anyOf', []), name + assert {'$ref': f'#/$defs/{name}'} not in entry.get('anyOf', []), ( + f'{name} refers to itself, so its mapping form is unreachable from the schema' + ) @pytest.mark.parametrize( @@ -74,4 +77,6 @@ def test_a_closed_vocabulary_is_published_as_an_enum(block, field, alias): def test_the_piecewise_method_vocabulary_has_one_home(): """`PiecewiseMethod` types the field and `PIECEWISE_METHODS` says what each emits.""" - assert set(get_args(model.PiecewiseMethod)) == set(model.PIECEWISE_METHODS) + assert set(get_args(model.PiecewiseMethod)) == set(model.PIECEWISE_METHODS), ( + 'the typed methods and the emitting ones disagree, so a method is accepted that emits nothing or the reverse' + ) diff --git a/tests/test_separability.py b/tests/test_separability.py new file mode 100644 index 00000000..288eac4c --- /dev/null +++ b/tests/test_separability.py @@ -0,0 +1,208 @@ +# SPDX-FileCopyrightText: math-spec Contributors +# +# SPDX-License-Identifier: MIT + +"""Whether a horizon may be built in windows, asked before any data binds. + +The verdict is what a rolling-horizon or myopic driver needs and cannot +currently get: a model with an annual budget windows into feasible pieces whose +rows are incomplete, and nothing says so. Every case below is one model shape +and the verdict it earns, because the value of the pass is entirely in getting +the boundary between the categories right. +""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any + +import pytest + +import math_spec as ms + +FIXTURE = Path(__file__).resolve().parent / 'fixtures' / 'every_program_node.yaml' + +BASE: dict[str, Any] = { + 'dimensions': {'h': {'dtype': 'int'}, 'u': {'dtype': 'str'}, 'zone': {'dtype': 'str'}, 'day': {'dtype': 'int'}}, + 'lookups': {'zone_of': {'over': 'u', 'into': 'zone'}, 'day_of': {'over': 'h', 'into': 'day'}}, + 'parameters': { + 'cost': {'dims': ['u']}, + 'budget': {'dims': []}, + 'width': {'dims': ['u'], 'dtype': 'int'}, + 'cap': {'dims': ['zone']}, + }, + 'variables': {'p': {'foreach': ['h', 'u'], 'bounds': {'lower': 0}}}, + 'objective': {'sense': 'minimize', 'expression': 'sum(p * cost)'}, +} + + +def _verdict(dimension: str = 'h', **patch: Any): + return ms.to_program({**BASE, **patch}).separability[dimension] + + +def _rows(expression: str, *, foreach: list[str] | None = None, **block: Any) -> dict[str, Any]: + return {'constraints': {'k': {'foreach': foreach or ['h', 'u'], 'expression': expression, **block}}} + + +@pytest.mark.parametrize( + ('patch', 'behind', 'ahead'), + [ + pytest.param(_rows('p >= 0'), 0, 0, id='pointwise-needs-no-overlap'), + pytest.param(_rows('p >= shift(p, over=h, offset=1, edge=0)'), 1, 0, id='a-shift-of-one-reads-one-row-behind'), + pytest.param(_rows('p >= shift(p, over=h, offset=3, edge=0)'), 3, 0, id='a-shift-of-three-reads-three-behind'), + pytest.param(_rows('p >= shift(p, over=h, offset=-2, edge=0)'), 0, 2, id='a-negative-shift-reads-ahead'), + pytest.param(_rows('sum_back(p, over=h, within=4) >= 0'), 3, 0, id='a-window-of-four-reads-three-behind'), + pytest.param( + _rows('p >= shift(p, over=u, offset=1, edge=0)'), 0, 0, id='a-shift-along-another-axis-is-nothing' + ), + ], +) +def test_a_separable_model_reports_the_overlap_a_window_needs_on_each_side(patch, behind, ahead): + verdict = _verdict(**patch) + assert verdict.windowable, 'nothing here ties the axis together' + assert (verdict.behind, verdict.ahead) == (behind, ahead), ( + 'a window must see what the widest translation reads before its first row and after its last' + ) + + +@pytest.mark.parametrize( + ('patch', 'fragment'), + [ + pytest.param(_rows('sum(p, over=h) <= budget', foreach=['u']), 'sums over h', id='a-budget-over-the-horizon'), + pytest.param(_rows("p >= shift(p, over=h, offset=1, edge='wrap')"), 'wraps around h', id='a-cyclic-shift'), + ], +) +def test_a_model_the_axis_ties_together_names_what_ties_it(patch, fragment): + verdict = _verdict(**patch) + assert not verdict.windowable, 'this shape does not survive being cut into windows' + assert fragment in verdict.coupled["constraint 'k'"], 'the report names the construct, not just the declaration' + assert not verdict.undecided and not verdict.restarts, 'a coupling is not something data or a driver resolves' + + +@pytest.mark.parametrize( + ('patch', 'named'), + [ + pytest.param(_rows('p >= shift(p, over=h, offset=1, by=day_of, edge=0)'), 'day_of', id='a-shift-inside-groups'), + pytest.param(_rows('p >= shift(p, over=h, offset=width, edge=0)'), 'width', id='an-offset-from-data'), + pytest.param(_rows('sum_back(p, over=h, within=width) >= 0'), 'width', id='a-width-from-data'), + ], +) +def test_a_reach_only_data_can_say_names_what_says_it(patch, named): + """A driver holding the data can compute this reach itself — the max of an + offset parameter, the runs a partition makes — so the verdict names the + parameter or lookup rather than refusing the model.""" + verdict = _verdict(**patch) + assert not verdict.windowable, 'undecided until data binds' + assert verdict.undecided["constraint 'k'"] == named, 'the report names what the driver has to read' + assert not verdict.coupled, 'and nothing structural ties the axis' + + +def test_a_read_through_a_lookup_is_undecided_on_the_axis_it_reads(): + """`at(cap, by=zone_of)` reads `zone` at whatever coordinate the lookup + chooses, so how far that reaches along `zone` is the lookup's data to say.""" + verdict = _verdict('zone', **_rows('p - at(cap, by=zone_of) <= 0')) + assert not verdict.windowable and not verdict.coupled, 'undecided until the lookup binds' + assert verdict.undecided["constraint 'k'"] == 'zone_of', 'the report names the lookup a driver has to read' + + +@pytest.mark.parametrize( + ('patch', 'independent'), + [ + pytest.param(_rows('p >= 0'), True, id='pointwise-slices-in-any-order'), + pytest.param(_rows('p >= shift(p, over=h, offset=1, edge=0)'), False, id='a-row-reading-another-is-not'), + pytest.param(_rows('p >= 0', where='position(h) == 0'), False, id='a-position-means-something-else-per-slice'), + pytest.param(_rows('p >= shift(p, over=h, offset=width, edge=0)'), False, id='an-undecided-reach-is-not'), + ], +) +def test_independence_is_windowability_with_nothing_read_across_and_nothing_counted(patch, independent): + verdict = _verdict(**patch) + assert verdict.independent is independent, 'one coordinate per slice needs no row to read or count another' + + +def test_a_coupling_names_the_change_that_would_lift_it(): + coupled = _verdict(**_rows('sum(p, over=h) <= budget', foreach=['u'])).coupled["constraint 'k'"] + assert 'sum_back(within=n)' in coupled, 'a horizon total becomes a rolling one' + wrapped = _verdict(**_rows("p >= shift(p, over=h, offset=1, edge='wrap')")).coupled["constraint 'k'"] + assert 'position(h) == 0' in wrapped, 'a wrap becomes an opening-state seed' + + +def test_a_mask_counting_a_position_is_reported_and_not_refused(): + """`position(h) == 0` fires once over a horizon and once per window, and a + rolling horizon seeding its opening state means the second. The verdict + stays windowable and says where a window would restart the count.""" + verdict = _verdict(**_rows('p >= 0', where='position(h) == 0')) + assert verdict.windowable, 'a seed is a modelling intent, not a coupling' + assert verdict.restarts == {"constraint 'k'": 'counts a position along h'}, 'the report names the declaration' + + +def test_a_sum_over_the_axis_couples_a_constraint_and_leaves_the_objective_alone(): + """The crux. An objective *is* a sum, so summing the windows' objectives is + summing the model's; a constraint row summing the axis ties every window to + every other. A verdict treating the two alike would refuse every windowable + model there is — and `BASE`'s objective sums over `h` in every case above.""" + assert _verdict(**_rows('p >= 0')).windowable, 'the objective sums over h and that is not a coupling' + coupled = _verdict(**_rows('sum(p, over=h) <= budget', foreach=['u'])) + assert not coupled.windowable, 'the same sum in a constraint is one' + + +def test_a_position_inside_a_cased_region_is_found(): + """`children` descends into a region's value and not its `when`, so a mask + written inside `cases:` is reachable by no expression walk — and seeding a + quantity at the start of the axis is exactly what a rolling horizon does.""" + verdict = _verdict( + expressions={ + 'prev': { + 'foreach': ['h', 'u'], + 'cases': {'opening': {'when': 'position(h) == 0', 'expression': 0}}, + 'otherwise': 'shift(p, over=h, offset=1, edge=0)', + } + }, + **_rows('p - prev <= 1'), + ) + assert "constraint 'k'" in verdict.restarts, 'the seed fires once over a horizon and once per window' + + +def test_the_overlap_is_the_widest_reach_of_any_block(): + verdict = _verdict( + constraints={ + 'near': {'foreach': ['h', 'u'], 'expression': 'p >= shift(p, over=h, offset=1, edge=0)'}, + 'far': {'foreach': ['h', 'u'], 'expression': 'p >= shift(p, over=h, offset=5, edge=0)'}, + } + ) + assert verdict.behind == 5, 'one window must see behind its first row as far as any block reads' + + +def test_a_grouping_that_consumes_the_axis_couples_it(): + program = ms.to_program( + {**BASE, 'constraints': {'z': {'foreach': ['h', 'zone'], 'expression': 'sum(p, by=zone_of) <= cap'}}} + ) + verdict = program.separability['u'] + assert not verdict.windowable, 'the grouping consumes u, so a window of u is a different sum' + + +def test_every_declared_axis_has_a_verdict_and_nothing_else_does(): + """The mapping is complete over the program's dimensions, so an axis nothing + mentions is trivially windowable rather than missing, and a name that is not + an axis is a `KeyError` rather than a verdict nobody should trust.""" + program = ms.to_program({**BASE, **_rows('p >= 0')}) + assert sorted(program.separability) == sorted(program.dimensions), 'every declared axis is answered for' + assert program.separability['zone'].windowable, 'an axis no construct mentions is trivially windowable' + with pytest.raises(KeyError): + program.separability['hh'] + + +@pytest.mark.parametrize('dimension', ['t', 'g', 'zone']) +def test_every_node_a_program_can_carry_is_judged_without_raising(dimension): + """The fixture the node fence maintains carries every construct, so this is + the pass meeting each of them at least once.""" + verdict = ms.to_program(FIXTURE).separability[dimension] + assert isinstance(verdict.behind, int), 'a verdict comes back for every axis of the widest model there is' + + +def test_a_reduction_over_several_axes_couples_every_one_of_them(): + """`sum(p)` with no `over=` collapses every dimension its operand carries, + 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': {'foreach': [], 'expression': 'sum(p) <= budget'}}}) + assert not program.separability['h'].windowable, 'the reduction consumes h' + assert not program.separability['u'].windowable, 'and u, in the same node' diff --git a/tests/test_validation.py b/tests/test_validation.py index 130ea7af..b74d449a 100644 --- a/tests/test_validation.py +++ b/tests/test_validation.py @@ -14,9 +14,9 @@ from math_spec._yaml import parse_yaml from math_spec.errors import DimensionError, LanguageError, SchemaError +from math_spec.program import DimensionPositionNode from math_spec.resolution import Namespace, where_of from math_spec.validation import to_spec -from math_spec.where_parser import DimensionPositionNode from tests.fixtures import DISPATCH_MODEL, OPERATOR_PROBES, SMALL_MODEL, override if TYPE_CHECKING: @@ -27,63 +27,14 @@ def _schema(**patch) -> Spec: return to_spec(override(SMALL_MODEL, **patch)) -class TestValidateExpressions: - @pytest.mark.parametrize( - ('patch', 'fragments'), - [ - pytest.param( - {'constraints': {'cap': {'foreach': ['g'], 'expression': 'nope <= c'}}}, - ("'nope' not found", "Constraint 'cap'", 'c'), - id='an-unknown-name-in-a-constraint', - ), - pytest.param( - {'constraints': {'cap': {'foreach': ['g'], 'expression': 'p + c'}}}, - ('exactly one comparison',), - id='a-constraint-without-a-comparison', - ), - pytest.param( - {'objective': {'expression': 'sum(p, over=g) <= 5'}}, - ('must not contain a comparison',), - id='an-objective-with-a-comparison', - ), - pytest.param( - {'constraints': {'cap': {'foreach': ['g'], 'expression': 'c <= 1'}}}, - ('decides nothing', "Constraint 'cap'", "'c <= 1'"), - id='a-comparison-with-no-variable-in-it', - ), - pytest.param( - {'constraints': {'cap': {'foreach': ['g'], 'expression': 'p * p * p <= c'}}}, - ("Constraint 'cap'", 'this product is degree 3'), - id='a-cubic-constraint', - ), - pytest.param( - {'objective': {'expression': 'sum(p ** 2, over=g)'}}, - ('The objective', '`**` is not in the language over variables'), - id='a-variable-under-a-power', - ), - pytest.param( - {'expressions': {'sq': 'p * p'}}, - ("Named expression 'sq'", 'which is degree 2'), - id='a-quadratic-named-expression', - ), - pytest.param( - {'constraints': {'cap': {'foreach': ['g'], 'where': 'c >', 'expression': 'p <= c'}}}, - ('Failed to parse where string',), - id='a-malformed-where-string', - ), - pytest.param( - {'constraints': {'cap': {'foreach': ['g'], 'where': 'not_a_param > 0', 'expression': 'p <= c'}}}, - ("'not_a_param' not found",), - id='an-unknown-name-in-a-where-used-to-evaluate-to-false', - ), - ], - ) - def test_a_bad_declaration_is_refused_at_load(self, patch, fragments): - with pytest.raises(LanguageError) as exc: - _schema(**patch) - for fragment in fragments: - assert fragment in str(exc.value) +def _refusal(model: dict[str, Any] = SMALL_MODEL, **patch: Any) -> str: + """The message `to_spec` refuses *model* patched with — and it has to refuse.""" + with pytest.raises(LanguageError) as caught: + to_spec(override(model, **patch)) + return str(caught.value) + +class TestValidateExpressions: def test_the_objective_and_a_constraint_take_degree_two(self): _schema( constraints={'floor': {'foreach': ['g'], 'expression': 'p * p >= 1'}}, @@ -91,46 +42,43 @@ def test_the_objective_and_a_constraint_take_degree_two(self): ) def test_multiple_errors_collected(self): - with pytest.raises(LanguageError) as exc_info: - _schema( - constraints={ - 'a': {'foreach': ['g'], 'expression': 'nope <= 1'}, - 'b': {'foreach': ['g'], 'expression': 'p + 1'}, - }, - ) - msg = str(exc_info.value) - assert "'nope' not found" in msg - assert 'exactly one comparison' in msg + message = _refusal( + constraints={ + 'a': {'foreach': ['g'], 'expression': 'nope <= 1'}, + 'b': {'foreach': ['g'], 'expression': 'p + 1'}, + }, + ) + assert "'nope' not found" in message + assert 'exactly one comparison' in message, 'the second fault is reported beside the first, not behind it' -class TestDimensionKwargs: - """A dim kwarg that names nothing is a silent no-op, not an error — `sum(p, over=snapshto)` used to load.""" +def _kwarg_model(expression: str, foreach: list[str] | None = None) -> dict[str, Any]: + """A model over (snapshot, generator), with `zone` a lookup into `bus`. - @staticmethod - def _schema(expression: str, foreach: list[str] | None = None) -> Spec: - """A model over (snapshot, generator), with `zone` a lookup into `bus`. + `zone` deliberately targets a dim `p` does *not* carry: grouping into + one it already has needs that dim twice, which is its own error. + `season` is a label space over the same dim, for the refusals below. + An explicit ``foreach=[]`` is a scalar constraint; ``None`` is the + default frame over `snapshot`. + """ + return { + 'dimensions': { + 'snapshot': {'dtype': 'int'}, + 'bus': {'dtype': 'str'}, + 'generator': {'dtype': 'str'}, + }, + 'lookups': { + 'zone': {'over': 'generator', 'into': 'bus'}, + 'season': {'over': 'generator', 'dtype': 'str'}, + }, + 'parameters': {'load': {'dims': ['snapshot']}}, + 'variables': {'p': {'foreach': ['snapshot', 'generator']}}, + 'constraints': {'c': {'foreach': ['snapshot'] if foreach is None else foreach, 'expression': expression}}, + } - `zone` deliberately targets a dim `p` does *not* carry: grouping into - one it already has needs that dim twice, which is its own error. - `season` is a label space over the same dim, for the refusals below. - """ - foreach = ['snapshot'] if foreach is None else foreach # an explicit [] is a scalar constraint - return to_spec( - { - 'dimensions': { - 'snapshot': {'dtype': 'int'}, - 'bus': {'dtype': 'str'}, - 'generator': {'dtype': 'str'}, - }, - 'lookups': { - 'zone': {'over': 'generator', 'into': 'bus'}, - 'season': {'over': 'generator', 'dtype': 'str'}, - }, - 'parameters': {'load': {'dims': ['snapshot']}}, - 'variables': {'p': {'foreach': ['snapshot', 'generator']}}, - 'constraints': {'c': {'foreach': foreach, 'expression': expression}}, - } - ) + +class TestDimensionKwargs: + """A dim kwarg that names nothing is a silent no-op, not an error — `sum(p, over=snapshto)` used to load.""" @pytest.mark.parametrize( ('expression', 'fragments'), @@ -143,7 +91,7 @@ def _schema(expression: str, foreach: list[str] | None = None) -> Spec: ), pytest.param( 'sum(p, by=zne) == load', - ('does not name a lookup', "Lookups: ['zone']"), + ('does not name a lookup', "Did you mean 'zone'?"), id='by-lookup-typo', ), pytest.param( @@ -154,10 +102,9 @@ def _schema(expression: str, foreach: list[str] | None = None) -> Spec: ], ) def test_a_dim_kwarg_typo_is_rejected(self, expression, fragments): - with pytest.raises(LanguageError) as exc: - self._schema(expression) + message = _refusal(_kwarg_model(expression)) for fragment in fragments: - assert fragment in str(exc.value) + assert fragment in message @pytest.mark.parametrize( ('expression', 'foreach'), @@ -173,7 +120,7 @@ def test_a_dim_kwarg_typo_is_rejected(self, expression, fragments): ], ) def test_declared_dimensions_still_pass(self, expression, foreach): - self._schema(expression, foreach) + to_spec(_kwarg_model(expression, foreach)) @pytest.mark.parametrize( 'expression', @@ -187,10 +134,9 @@ def test_a_label_space_is_refused_wherever_by_needs_a_target(self, expression): """Every `by=` but `position`'s reaches a target dimension: `sum` and `at` to place terms on it, `shift` so a named `offset=` may vary per group. A label space targets nothing, so all three refuse it and name the promotion (#280).""" - with pytest.raises(LanguageError) as exc: - self._schema(expression, ['snapshot', 'bus']) - assert 'is a label space' in str(exc.value), 'the refusal names the kind, not just the name' - assert 'season_of' in str(exc.value), 'and it spells the promotion out' + message = _refusal(_kwarg_model(expression, ['snapshot', 'bus'])) + assert 'is a label space' in message, 'the refusal names the kind, not just the name' + assert 'season_of' in message, 'and it spells the promotion out' def test_macro_formals_are_not_mistaken_for_dimensions(self): """A formal in a dim position is legal inside the template body.""" @@ -226,7 +172,8 @@ class TestArithmeticDtype: """ @staticmethod - def _schema(dtype: str, expression: str) -> Spec: + def _schema_with_typed_a(dtype: str, expression: str) -> Spec: + """`SMALL_MODEL` plus a parameter `a` of *dtype*, standing in the constraint *expression*.""" return _schema( **{ 'parameters.a': {'dims': ['g'], 'dtype': dtype}, @@ -247,11 +194,11 @@ def _schema(dtype: str, expression: str) -> Spec: ) def test_a_label_or_a_flag_is_not_a_value(self, dtype, expression): with pytest.raises(LanguageError, match=f'declared dtype: {dtype}'): - self._schema(dtype, expression) + self._schema_with_typed_a(dtype, expression) @pytest.mark.parametrize('dtype', ['float', 'int']) def test_a_number_is(self, dtype): - self._schema(dtype, 'a * p <= c') + self._schema_with_typed_a(dtype, 'a * p <= c') @pytest.mark.parametrize( ('dtype', 'where'), @@ -284,10 +231,7 @@ def test_absent_and_zero_are_the_unstable_surface(self, top): assert _schema(**top).version == 0 def test_an_unknown_version_is_refused_not_interpreted(self): - with pytest.raises(LanguageError) as exc: - _schema(version=1) - - message = str(exc.value) + message = _refusal(version=1) assert 'declares version 1' in message assert 'understands [0]' in message, 'the error has to say what this reader can read' assert 'Upgrade math_spec' in message, 'and what to do about it' @@ -332,7 +276,9 @@ class TestPositionResolves: ids=['first', 'first of each period', 'first of each season, by a label space'], ) def test_it_resolves(self, mask: str, position: int, by: str | None): - node = where_of(mask, Namespace.of(POSITION_SCHEMA), 'the mask') + resolved = where_of(mask, Namespace.of(POSITION_SCHEMA), 'the mask') + assert resolved is not None + node = resolved.root assert isinstance(node, DimensionPositionNode) assert node.name == 'snapshot' assert node.position == position @@ -361,6 +307,51 @@ class TestRulesDecidedWithoutData: @pytest.mark.parametrize( ('patch', 'fragments'), [ + pytest.param( + {'constraints': {'cap': {'foreach': ['g'], 'expression': 'nope <= c'}}}, + ("'nope' not found", "Constraint 'cap'", 'c'), + id='an-unknown-name-in-a-constraint', + ), + pytest.param( + {'constraints': {'cap': {'foreach': ['g'], 'expression': 'p + c'}}}, + ('exactly one comparison',), + id='a-constraint-without-a-comparison', + ), + pytest.param( + {'objective': {'expression': 'sum(p, over=g) <= 5'}}, + ('must not contain a comparison',), + id='an-objective-with-a-comparison', + ), + pytest.param( + {'constraints': {'cap': {'foreach': ['g'], 'expression': 'c <= 1'}}}, + ('decides nothing', "Constraint 'cap'", "'c <= 1'"), + id='a-comparison-with-no-variable-in-it', + ), + pytest.param( + {'constraints': {'cap': {'foreach': ['g'], 'expression': 'p * p * p <= c'}}}, + ("Constraint 'cap'", 'this product is degree 3'), + id='a-cubic-constraint', + ), + pytest.param( + {'objective': {'expression': 'sum(p ** 2, over=g)'}}, + ('The objective', '`**` is not in the language over variables'), + id='a-variable-under-a-power', + ), + pytest.param( + {'expressions': {'sq': 'p * p'}}, + ("Named expression 'sq'", 'which is degree 2'), + id='a-quadratic-named-expression', + ), + pytest.param( + {'constraints': {'cap': {'foreach': ['g'], 'where': 'c >', 'expression': 'p <= c'}}}, + ('Failed to parse where string',), + id='a-malformed-where-string', + ), + pytest.param( + {'constraints': {'cap': {'foreach': ['g'], 'where': 'not_a_param > 0', 'expression': 'p <= c'}}}, + ("'not_a_param' not found",), + id='an-unknown-name-in-a-where', + ), pytest.param( {'sos': {'s': {'variable': 'p', 'over': 'z', 'type': 1}}}, ("undeclared dimension 'z'",), @@ -576,15 +567,15 @@ class TestRulesDecidedWithoutData: ], ) def test_a_rule_decided_without_data(self, patch, fragments): - with pytest.raises(LanguageError) as exc: - _schema(**patch) + message = _refusal(**patch) for fragment in fragments: - assert fragment in str(exc.value) + assert fragment in message class TestTheFrontDoor: def test_a_list_of_models_is_not_a_model(self): - with pytest.raises(TypeError, match='one file, one dict or one Spec, never a list'): + """Composition is Python's, not the file's (#30) — and the refusal is the package's own, so the CLI's one except catches it.""" + with pytest.raises(SchemaError, match='one file, one dict or one Spec, never a list'): to_spec([DISPATCH_MODEL, DISPATCH_MODEL]) def test_a_loaded_model_passes_through_as_itself(self): @@ -608,8 +599,8 @@ def test_to_yaml_reproduces_the_model(self): def test_an_empty_list_survives_the_round_trip(self): """`foreach: []` is a scalar declaration, not an absence — stripping it would put the variable on every dim it names.""" model = _schema(**{'variables.p.foreach': []}) - assert model.to_dict()['variables']['p']['foreach'] == [] - assert to_spec(model.to_dict()).variables['p'].foreach == [] + assert model.to_dict()['variables']['p']['foreach'] == [], 'the empty frame is written out, not dropped' + assert to_spec(model.to_dict()).variables['p'].foreach == [], 'and reads back as the scalar it declares' def test_an_empty_section_is_not_written(self): written = to_spec(DISPATCH_MODEL).to_yaml() @@ -635,15 +626,21 @@ def test_a_default_is_written_out_and_an_absence_is_not(self): OPENING = {'opening': {'when': 'position(snapshot) == 0', 'expression': 'p_max'}} +def _headroom(block: dict[str, Any]) -> dict[str, Any]: + """`CASED_BASE` with *block* as its one named expression, `headroom`.""" + return {**copy.deepcopy(CASED_BASE), 'expressions': {'headroom': block}} + + def _cased(cases: dict[str, Any] | None = None, **block: Any) -> dict[str, Any]: - """`CASED_BASE` with one cased expression named `headroom`.""" - declared = { - 'foreach': ['snapshot', 'generator'], - 'cases': OPENING if cases is None else cases, - 'otherwise': 0, - **block, - } - return {**copy.deepcopy(CASED_BASE), 'expressions': {'headroom': declared}} + """`_headroom` over a cased block: `OPENING` or *cases*, an `otherwise:` of 0, and *block* on top.""" + return _headroom( + { + 'foreach': ['snapshot', 'generator'], + 'cases': OPENING if cases is None else cases, + 'otherwise': 0, + **block, + } + ) class TestExpressionCases: @@ -651,22 +648,37 @@ class TestExpressionCases: def test_a_cased_expression_loads(self): block = to_spec(_cased()).expressions['headroom'] - assert list(block.cases) == ['opening'] + assert list(block.cases) == ['opening'], 'the one case, under the name the file gave it' assert block.otherwise == '0' + @pytest.mark.parametrize( + ('when', 'fragment'), + [ + pytest.param( + 'position(snapshot) == 0 OR True', + 'no other arm can hold anywhere', + id='folds-to-every-row', + ), + pytest.param('False', 'this arm never applies', id='admits-no-row'), + ], + ) + def test_an_arm_the_data_cannot_decide_is_refused(self, when: str, fragment: str): + """A mask that folds to a literal is not a case, and the refusal names the rewrite.""" + model = _cased(cases={'opening': {'when': when, 'expression': 'p_max'}}) + with pytest.raises(SchemaError, match=fragment): + to_spec(model) + def test_it_round_trips(self): """The mapping form goes back out as it came in, `otherwise:` and all.""" schema = to_spec(_cased(description='what is spare')) assert to_spec(schema.to_dict()).to_yaml() == schema.to_yaml() def test_the_fallback_is_written_as_the_bare_value(self): - """`otherwise:` carries nothing but its value, so a mapping around it would be ceremony. - - The same shorthand `expressions:` itself takes, and a number is how a - constant region is spelled — YAML reads `otherwise: 0` as an int. - """ + """`otherwise:` carries nothing but its value, so a mapping around it would be ceremony.""" written = to_spec(_cased()).to_dict()['expressions']['headroom'] - assert written['cases'] == {'opening': {'when': 'position(snapshot) == 0', 'expression': 'p_max'}} + assert written['cases'] == {'opening': {'when': 'position(snapshot) == 0', 'expression': 'p_max'}}, ( + 'a case goes back out as the mapping it came in as' + ) assert written['otherwise'] == '0' @pytest.mark.parametrize( @@ -697,9 +709,8 @@ def test_the_fallback_is_written_as_the_bare_value(self): ], ) def test_the_two_forms_do_not_mix(self, block: dict[str, Any], fragment: str): - model = {**copy.deepcopy(CASED_BASE), 'expressions': {'headroom': block}} with pytest.raises(SchemaError, match=re.escape(fragment)): - to_spec(model) + to_spec(_headroom(block)) @pytest.mark.parametrize( ('cases', 'message'), @@ -717,12 +728,7 @@ def test_the_two_forms_do_not_mix(self, block: dict[str, Any], fragment: str): ], ) def test_the_schema_itself_states_the_shape_of_a_case(self, cases: dict[str, Any], message: str): - """Every case says where it applies, and a block carries one — both the closed schema's own error. - - Neither is a rule this module writes a sentence for: `when:` is - required, and an empty `cases:` beside an `otherwise:` is the one value - everywhere that `expression:` already says. - """ + """Every case says where it applies, and a block carries one — both the closed schema's own error.""" with pytest.raises(SchemaError, match=re.escape(message)): to_spec(_cased(cases)) @@ -741,7 +747,9 @@ def test_cases_spelled_apart_load(self): 'gas': {'when': "generator == 'gas'", 'expression': 'p_max'}, 'opening': {'when': "generator != 'gas' and position(snapshot) == 0", 'expression': 'p_max * 2'}, } - assert list(to_spec(_cased(cases)).expressions['headroom'].cases) == ['gas', 'opening'] + assert list(to_spec(_cased(cases)).expressions['headroom'].cases) == ['gas', 'opening'], ( + 'both cases load, in the order the file wrote them' + ) def test_a_pair_that_cannot_be_decided_is_refused_as_an_overlap_is(self): """`snapshot` declares no `values:`, so 0 and -1 are one row on a one-member axis.""" @@ -764,7 +772,9 @@ def test_a_case_may_not_widen_the_frame(self): def test_a_when_may_not_test_a_dim_outside_the_frame(self): """The same rule a variable's or a constraint's mask is held to.""" - with pytest.raises(DimensionError, match="'snapshot', which is not in the frame"): + with pytest.raises( + DimensionError, match=r"where-dimension 'snapshot' reads dims \['snapshot'\] outside the frame" + ): to_spec(_cased(foreach=['generator'])) def test_an_unknown_name_in_a_case_is_a_load_error(self): @@ -798,10 +808,7 @@ def test_a_fault_in_an_arm_names_the_declaration_and_is_reported_once(self): name: {'foreach': ['snapshot', 'generator'], 'expression': f'p <= headroom + {n}'} for n, name in enumerate(('cap', 'floor')) } - with pytest.raises(SchemaError) as caught: - to_spec(model) - - message = str(caught.value) + message = _refusal(model) assert message.count("'nope' not found") == 1, 'two constraints read it; the fault is reported once' assert "Named expression 'headroom', case 'opening'" in message assert 'Constraint' not in message, "the arm is the declaration's, not the use site's" @@ -814,10 +821,7 @@ def test_the_fallback_is_not_named_as_a_case(self): """ model = _cased(otherwise='nope') model['constraints'] = {'cap': {'foreach': ['snapshot', 'generator'], 'expression': 'p <= headroom'}} - with pytest.raises(SchemaError) as caught: - to_spec(model) - - message = str(caught.value) + message = _refusal(model) assert "Named expression 'headroom', otherwise: 'nope' not found" in message assert "case 'otherwise'" not in message, 'the fallback is not one of the cases' @@ -848,3 +852,69 @@ def test_a_boolean_is_still_not_an_expression(self): """`true` is not arithmetic, and an error naming the type reads better than one naming `'True'`.""" with pytest.raises(SchemaError, match='valid string'): _schema(**{'expressions.always': {'expression': True}}) + + +class TestADeclarationIsNamed: + """A declaration's key must be a name the expression grammar could write. + + Nothing checked it, so `parameters: {'': {...}}` loaded, and a piecewise + block naming it under `points:` had its mask silently dropped — + `if mask:` in the expansion read a declared parameter as "this block + masks nothing", and the weights came out unmasked. Every unwritable name + has the same shape: a declaration no expression can reach, in a language + whose promise is that the file decides. + """ + + @pytest.mark.parametrize( + 'name', + [ + pytest.param('', id='empty'), + pytest.param(' ', id='a-space'), + pytest.param('a b', id='two-words'), + pytest.param('1x', id='leading-digit'), + pytest.param('a-b', id='a-hyphen'), + pytest.param('a.b', id='a-dot'), + ], + ) + @pytest.mark.parametrize( + 'section', + [ + 'dimensions', + 'lookups', + 'parameters', + 'variables', + 'expressions', + 'macros', + 'constraints', + 'piecewise', + 'sos', + ], + ) + def test_a_name_no_expression_could_write_is_refused(self, section: str, name: str): + declarations: dict[str, Any] = { + 'dimensions': {'dtype': 'str'}, + 'lookups': {'over': 'g', 'into': 'h'}, + 'parameters': {'dims': ['g']}, + 'variables': {'foreach': ['g']}, + 'expressions': {'expression': 'c'}, + 'macros': {'args': ['x'], 'template': 'x * 2'}, + 'constraints': {'foreach': ['g'], 'expression': 'p <= c'}, + 'piecewise': {'over': 'g', 'links': [['p', 'c'], ['q', 'c']], 'method': 'convex'}, + 'sos': {'variable': 'p', 'over': 'g', 'type': 1}, + } + model = copy.deepcopy(SMALL_MODEL) + model.setdefault(section, {})[name] = declarations[section] + with pytest.raises(LanguageError, match='is not a name'): + to_spec(model) + + def test_the_message_names_the_rewrite(self): + model = copy.deepcopy(SMALL_MODEL) + model['parameters']['a b'] = {'dims': ['g']} + message = _refusal(model) + assert "'a b'" in message, 'the offending name is quoted' + assert 'letter or an underscore' in message, ( + 'the message says what a name may be, not only that this is not one' + ) + + def test_an_ordinary_name_still_loads(self): + assert 'headroom_2' in _schema(**{'parameters.headroom_2': {'dims': ['g']}}).parameters diff --git a/tests/test_yaml_loading.py b/tests/test_yaml_loading.py index c9778aec..083dd745 100644 --- a/tests/test_yaml_loading.py +++ b/tests/test_yaml_loading.py @@ -6,8 +6,12 @@ from __future__ import annotations +import importlib + import pytest +import yaml +import math_spec._yaml as _yaml_module from math_spec._yaml import read_yaml from math_spec.errors import SchemaError from math_spec.validation import to_spec @@ -46,22 +50,22 @@ def test_only_true_and_false_are_booleans(tmp_path): """YAML 1.1 resolved these to bools, so the declaration the file names is not the one that reaches the schema — ``no`` is Norway.""" path = _write(tmp_path, _BOOLISH_DIMS) - assert list(read_yaml(path)['dimensions']) == _BOOLISH + assert list(read_yaml(path)['dimensions']) == _BOOLISH, 'every boolish word is a str key, in file order' def test_the_harness_reads_a_model_the_way_the_product_does(tmp_path): """``raw_of`` read YAML 1.1, so a dimension named ``no`` reached a test's schema as ``False`` while ``to_spec`` saw the string.""" path = _write(tmp_path, _BOOLISH_DIMS) - assert raw_of(path) == read_yaml(path) - assert list(raw_of(_BOOLISH_DIMS)['dimensions']) == _BOOLISH + assert raw_of(path) == read_yaml(path), 'a path is read by the same loader the product uses' + assert list(raw_of(_BOOLISH_DIMS)['dimensions']) == _BOOLISH, 'and so is a YAML string' def test_real_booleans_still_parse(tmp_path): """The narrowed resolver keeps 1.2's `true`/`false` as booleans, not labels.""" path = _write(tmp_path, 'flags:\n a: true\n b: false\n') - assert read_yaml(path)['flags'] == {'a': True, 'b': False} + assert read_yaml(path)['flags'] == {'a': True, 'b': False}, 'the two spellings YAML 1.2 keeps are still bools' def test_the_loader_yields_plain_types(tmp_path): @@ -102,20 +106,21 @@ def test_a_merge_key_override_is_not_a_duplicate(tmp_path): 'variables:\n p:\n <<: *d\n foreach: [generator]\n', ) - assert read_yaml(path)['variables']['p']['foreach'] == ['generator'] + assert read_yaml(path)['variables']['p']['foreach'] == ['generator'], 'the explicit key wins over the merged one' -def test_a_non_mapping_document_is_a_load_error(tmp_path): +@pytest.mark.parametrize( + 'text', [pytest.param('- a\n- b\n', id='a-sequence'), pytest.param('just a string\n', id='a-scalar')] +) +def test_a_non_mapping_document_is_a_load_error(tmp_path, text): """Otherwise `Spec(**raw)` raises a bare TypeError about `**`.""" - for text in ('- a\n- b\n', 'just a string\n'): - path = _write(tmp_path, text) - with pytest.raises(SchemaError, match='must be a mapping of sections'): - to_spec(path) + with pytest.raises(SchemaError, match='must be a mapping of sections'): + to_spec(_write(tmp_path, text)) def test_an_empty_file_is_an_empty_model(tmp_path): - assert read_yaml(_write(tmp_path, '')) == {} - assert read_yaml(_write(tmp_path, '# only a comment\n')) == {} + assert read_yaml(_write(tmp_path, '')) == {}, 'no document is an empty mapping rather than None' + assert read_yaml(_write(tmp_path, '# only a comment\n')) == {}, 'and so is a document of comments alone' def test_a_complex_key_is_refused_in_our_tree(tmp_path): @@ -134,3 +139,44 @@ def test_two_merge_keys_accumulate(tmp_path): ) with pytest.raises(SchemaError, match="unknown key 'anchors'"): to_spec(path) + + +#: One document per rule this module owns, plus the two 1.1 coercions it keeps +#: on purpose — what the two scanners have to agree about. +_EVERY_RULE = ( + _BOOLISH_DIMS + + 'flags: {a: true, b: false}\n' + + 'kept: {stamp: 2024-01-01, sexagesimal: 12:30}\n' + + 'merged:\n base: &b {dtype: int}\n use:\n <<: *b\n dtype: str\n' + + 'nested: [{a: 1}, [2, 3], null, 4.5, "quoted"]\n' +) + + +@pytest.mark.skipif(not hasattr(yaml, 'CSafeLoader'), reason='this PyYAML has no libyaml scanner to compare against') +def test_both_scanners_read_a_file_the_same_way(tmp_path, monkeypatch): + """The loader takes libyaml's scanner where the install has one, and PyYAML's own otherwise — a difference no model may be able to see. + + Only the faster one runs in CI, so the fallback is reachable here only by + hiding `CSafeLoader` and reloading the module. Both readings are compared + against each other rather than against a written-out expectation, so this + cannot drift from the rules the tests above pin. + """ + path = _write(tmp_path, _EVERY_RULE) + with_libyaml = read_yaml(path) + duplicate = _write(tmp_path, _EVERY_RULE + 'flags: {}\n', name='dup.yaml') + with pytest.raises(SchemaError) as fast: + read_yaml(duplicate) + + monkeypatch.delattr(yaml, 'CSafeLoader') + try: + importlib.reload(_yaml_module) + assert _yaml_module._StrictLoader.__mro__[1] is yaml.SafeLoader, 'the fallback base is what this test came for' + assert _yaml_module.read_yaml(path) == with_libyaml, 'the same document, scalar for scalar' + with pytest.raises(SchemaError) as slow: + _yaml_module.read_yaml(duplicate) + assert str(slow.value) == str(fast.value), 'and the same line for a duplicate key, since both carry marks' + finally: + monkeypatch.undo() + importlib.reload(_yaml_module) + + assert _yaml_module._StrictLoader.__mro__[1] is yaml.CSafeLoader, 'the module is left as the suite found it' diff --git a/tests/typesetting/fixtures.py b/tests/typesetting/fixtures.py index cd29e0cc..7a5dfa6b 100644 --- a/tests/typesetting/fixtures.py +++ b/tests/typesetting/fixtures.py @@ -10,7 +10,7 @@ from math_spec.typesetting import FORMATS -LATEX, TYPST = FORMATS['latex'], FORMATS['typst'] +LATEX = FORMATS['latex'] EVERY_FORMAT = pytest.mark.parametrize('fmt', list(FORMATS.values()), ids=list(FORMATS)) TYPST_SYMBOLS = { diff --git a/tests/typesetting/golden/__init__.py b/tests/typesetting/golden/__init__.py index cadcf362..a08f445d 100644 --- a/tests/typesetting/golden/__init__.py +++ b/tests/typesetting/golden/__init__.py @@ -13,6 +13,5 @@ def path_for(format_name: str) -> Path: - """The committed output for one format. Named for the format, not its - suffix, so two formats sharing a suffix could not collide silently.""" + """The committed output for one format.""" return DIRECTORY / f'{format_name}.out' diff --git a/tests/typesetting/golden/__main__.py b/tests/typesetting/golden/__main__.py index 800b29a7..d9f4e105 100644 --- a/tests/typesetting/golden/__main__.py +++ b/tests/typesetting/golden/__main__.py @@ -4,24 +4,21 @@ """Regenerate the committed golden output. - pixi run python -m tests.typesetting.golden - -Then **read the diff**. That is the review: a golden file is only worth having -if a change to it is looked at, and the reason the output is generated rather -than hand-written is that the diff is where human judgement belongs — at -review time, not at authoring time. +pixi run python -m tests.typesetting.golden """ from __future__ import annotations +from math_spec import to_spec from math_spec.typesetting import FORMATS, typeset from tests.typesetting.golden import MODEL, path_for def main() -> int: + model = to_spec(MODEL) for name, fmt in FORMATS.items(): path = path_for(name) - path.write_text(typeset(MODEL, fmt, standalone=True)) + path.write_text(typeset(model, fmt, standalone=True)) print(f'wrote {path}') return 0 diff --git a/tests/typesetting/test_cases.py b/tests/typesetting/test_cases.py index 91e2a159..24a35cb2 100644 --- a/tests/typesetting/test_cases.py +++ b/tests/typesetting/test_cases.py @@ -12,6 +12,7 @@ from math_spec import SchemaError, to_latex, to_spec, typeset from math_spec.piecewise import expand_piecewise +from math_spec.resolution import Namespace from math_spec.typesetting.symbols import chosen_expressions, printed_expressions from tests.fixtures import DISPATCH_MODEL as DISPATCH from tests.fixtures import override @@ -37,26 +38,37 @@ }, ) - -def _sections(rendered: str) -> list[str]: - """The section titles the render printed, in order.""" - return [title for title in ('Objective', 'Subject to', 'Definitions', 'Variable domains') if title in rendered] +#: One cased expression reached only through another's case. `opening_cost` has +#: no variable of its own — its route to one runs through `headroom`. +_NESTED = override( + CASED, + **{ + 'expressions.headroom.cases.opening.expression': 'p', + 'expressions.opening_cost.foreach': ['snapshot', 'generator'], + 'expressions.opening_cost.cases': { + 'opening': {'when': 'position(snapshot) == 0', 'expression': 'headroom * cost'}, + }, + 'expressions.opening_cost.otherwise': 0, + 'constraints.spare.expression': 'p <= opening_cost', + }, +) @EVERY_FORMAT def test_a_cased_expression_is_the_exception_that_keeps_its_name(fmt: Format): - """It prints once, as a definition, and its uses name it. - - The other way round — the block inlined at each use — is what the AST does + """The other way round — the block inlined at each use — is what the AST does and the wrong thing to print: a block three arms tall puts whatever follows - it beside its middle row. - """ + it beside its middle row.""" rendered = typeset(CASED, fmt, legend=False) - # counted indexed, because Typst spells a row label and an upright symbol - # the same way and only the symbol carries the dims indexed = fmt.subscript(fmt.upright('headroom'), ['t', 'g']) - assert rendered.count(indexed) == 2, 'one use and one definition, no more' - assert _sections(rendered) == ['Objective', 'Subject to', 'Definitions', 'Variable domains'] + assert rendered.count(indexed) == 2, ( + 'one use and one definition, no more — counted indexed, because Typst spells a row label and an upright ' + 'symbol the same way and only the symbol carries the dims' + ) + sections = [title for title in ('Objective', 'Subject to', 'Definitions', 'Variable domains') if title in rendered] + assert sections == ['Objective', 'Subject to', 'Definitions', 'Variable domains'], ( + 'the definition has a section of its own, after the constraints and before the domains' + ) @EVERY_FORMAT @@ -77,40 +89,29 @@ def test_a_declared_definition_prints_whether_or_not_a_row_names_it(fmt: Format) @EVERY_FORMAT -def test_a_case_is_given_when_its_values_are_however_its_regions_are_chosen(fmt: Format): - """A `when` mentioning a variable does not make the quantity one. - - The mask asks whether the variable *exists* at a coordinate, which the - model settles when it is built; only a value reaching one is a quantity the - solver returns. - """ - masked = override( - CASED, - **{'expressions.headroom.cases': {'running': {'when': 'p', 'expression': 'p_max'}}}, - ) - rendered = typeset(masked, fmt, legend=False) - assert fmt.upright('headroom') in rendered, 'every case is a parameter, so the quantity is given' - assert fmt.italic('headroom') not in rendered - - -@EVERY_FORMAT -def test_a_case_reaching_a_variable_is_chosen(fmt: Format): - """One case holding a variable is enough: the solver decides the quantity.""" - decided = override(CASED, **{'expressions.headroom.cases.opening.expression': 'p'}) - assert fmt.italic('headroom') in typeset(decided, fmt, legend=False) - - -@EVERY_FORMAT -def test_the_fallback_reaching_a_variable_is_chosen(fmt: Format): - """The `otherwise:` is a value of the quantity like any case's. - - `previous_status` in the commitment example is this shape and no other: its - two cases are a constant and a parameter, and the variable is in the - fallback alone. Read only the cases and the block prints upright, which - says the model was handed a quantity it in fact solves for. - """ - decided = override(CASED, **{'expressions.headroom.otherwise': 'p'}) - assert fmt.italic('headroom') in typeset(decided, fmt, legend=False) +@pytest.mark.parametrize( + ('patch', 'chosen'), + [ + pytest.param( + {'expressions.headroom.cases': {'running': {'when': 'p', 'expression': 'p_max'}}}, + False, + id='a-when-naming-a-variable-leaves-it-given', + ), + pytest.param({'expressions.headroom.cases.opening.expression': 'p'}, True, id='a-case-reaching-a-variable'), + pytest.param({'expressions.headroom.otherwise': 'p'}, True, id='the-fallback-reaching-a-variable'), + ], +) +def test_a_cased_expression_is_chosen_when_a_value_reaching_it_is(fmt: Format, patch: dict, chosen: bool): + """A `when` mentioning a variable does not make the quantity one: the mask + asks whether the variable *exists* at a coordinate, which the model settles + when it is built. Only a value reaching one is a quantity the solver + returns, and one case holding a variable is enough. The `otherwise:` is a + value of the quantity like any case's, so a walk reading only the cases + prints a solved quantity upright.""" + rendered = typeset(override(CASED, **patch), fmt, legend=False) + italic, upright = (fmt.subscript(face('headroom'), ['t', 'g']) for face in (fmt.italic, fmt.upright)) + assert (italic in rendered) is chosen, 'the quantity is chosen exactly when a value reaching it holds a variable' + assert (upright in rendered) is not chosen, 'and given otherwise, however its regions are chosen' @EVERY_FORMAT @@ -118,23 +119,9 @@ def test_a_definition_naming_another_one_prints_both(fmt: Format): """The cases are walked too, so the collection runs to a fixpoint.""" rendered = typeset(_NESTED, fmt, legend=False) assert fmt.italic('headroom') in rendered, 'the inner definition was reached through a case' - assert rendered.count(fmt.subscript(fmt.italic('opening_cost'), ['t', 'g'])) == 2 - - -#: One cased expression reached only through another's case. `opening_cost` has -#: no variable of its own — its route to one runs through `headroom`. -_NESTED = override( - CASED, - **{ - 'expressions.headroom.cases.opening.expression': 'p', - 'expressions.opening_cost.foreach': ['snapshot', 'generator'], - 'expressions.opening_cost.cases': { - 'opening': {'when': 'position(snapshot) == 0', 'expression': 'headroom * cost'}, - }, - 'expressions.opening_cost.otherwise': 0, - 'constraints.spare.expression': 'p <= opening_cost', - }, -) + assert rendered.count(fmt.subscript(fmt.italic('opening_cost'), ['t', 'g'])) == 2, ( + 'the outer definition and its one use' + ) def test_a_variable_reached_through_another_cased_expression_still_prints_chosen(): @@ -145,8 +132,9 @@ def test_a_variable_reached_through_another_cased_expression_still_prints_chosen one upright — a quantity the solver decides, set as one the model was handed. """ schema = expand_piecewise(to_spec(_NESTED)) - assert chosen_expressions(schema) == {'headroom', 'opening_cost'} - assert r'\mathit{opening\_cost}' in to_latex(_NESTED, legend=False) + assert chosen_expressions(schema, Namespace.of(schema)) == {'headroom', 'opening_cost'}, ( + 'the chain is followed to its end, so both are chosen' + ) def test_the_table_may_rename_a_cased_expression_but_not_a_plain_one(): diff --git a/tests/typesetting/test_cli.py b/tests/typesetting/test_cli.py index 923b47c7..2423e206 100644 --- a/tests/typesetting/test_cli.py +++ b/tests/typesetting/test_cli.py @@ -4,12 +4,6 @@ """The shell front — `python -m math_spec model.yaml`. -What is checked here is the *design* rather than argparse: that the typeset -verbs are read off `FORMATS` rather than listed twice, that `check` is the -language's verdict and its advice with no consumer installed, that the front -costs no dependency, that no verb binds data, and that a format nothing can -render is refused rather than written as an empty file. - `main` takes its argv and `parser` hands back the verbs, so none of this needs a subprocess or a scrape of help text. """ @@ -25,10 +19,11 @@ import math_spec.__main__ as front from math_spec.typesetting import FORMATS +from tests.fixtures import EXAMPLES from tests.typesetting import golden #: The golden model, not `examples/dispatch.yaml`: the CLI travels with the -#: renderer and this fixture travels with both, where the gallery stays. +#: renderer and this fixture travels with both. MODEL = str(golden.MODEL) @@ -67,7 +62,7 @@ def test_the_verbs_are_check_and_the_formats_and_nothing_else(): def test_check_prints_nothing_for_a_clean_file(capsys): - assert front.main(['check', str(Path(__file__).resolve().parents[2] / 'examples' / 'dispatch.yaml')]) == 0 + assert front.main(['check', str(EXAMPLES / 'dispatch.yaml')]) == 0 assert capsys.readouterr() == ('', ''), 'no advice, no output' @@ -75,14 +70,8 @@ def test_check_accepts_the_model_that_carries_every_construct(capsys): """The golden model loads, which is not the same as it being advice-free. It exercises every operator and every edge policy, so `check` accepting it - is the claim that the whole language loads through one door. It could not - be asked before #193, when the edge rules were decided in lowering alone - and `check` refused a file `to_spec` had just accepted. - - Silence is a different property and this model does not have it: `season` - is a lookup target nothing is indexed by, which `advice` says is a label - space. That is the fixture being deliberately odd, not a defect — hence the - neighbour above, on a model that is ordinary. + is the claim that the whole language loads through one door. `season` is a + lookup target nothing is indexed by, so it carries advice by design. """ assert front.main(['check', str(golden.MODEL)]) == 0, 'the whole language loads' out, err = capsys.readouterr() @@ -90,34 +79,32 @@ def test_check_accepts_the_model_that_carries_every_construct(capsys): assert 'is never an axis' in out, 'and the advice it carries is about `season`, not about an edge' -def test_check_prints_advice_and_does_not_fail(tmp_path, capsys): - model = tmp_path / 'm.yaml' - model.write_text(UNUSED_DIMENSION) - assert front.main(['check', str(model)]) == 0, 'advice is not a refusal' - out, err = capsys.readouterr() - assert "dimension 'spare' is never used" in out - assert err == '' +def _carries(stream: str, said: str) -> bool: + """*stream* mentions *said*, or is silent where *said* is empty.""" + return said in stream if said else stream == '' -def test_check_puts_a_refusal_on_stderr_with_exit_status_one(tmp_path, capsys): +@pytest.mark.parametrize( + ('yaml', 'status', 'out', 'err'), + [ + pytest.param(UNUSED_DIMENSION, 0, "dimension 'spare' is never used", '', id='advice'), + pytest.param( + UNUSED_DIMENSION.replace('sum(p * c)', 'sum(p * nope)'), 1, '', "'nope' not found", id='a-refusal' + ), + ], +) +def test_check_puts_advice_on_stdout_and_a_refusal_on_stderr(tmp_path, capsys, yaml, status, out, err): model = tmp_path / 'm.yaml' - model.write_text(UNUSED_DIMENSION.replace('sum(p * c)', 'sum(p * nope)')) - assert front.main(['check', str(model)]) == 1 - out, err = capsys.readouterr() - assert out == '' - assert "'nope' not found" in err + model.write_text(yaml) + assert front.main(['check', str(model)]) == status, 'advice is not a refusal, and a refusal is status one' + captured = capsys.readouterr() + assert _carries(captured.out, out), 'advice goes to stdout, and nothing else does' + assert _carries(captured.err, err), 'a refusal goes to stderr, and nothing else does' def test_the_shell_front_costs_no_dependency(): - """It is stdlib argparse over `typeset`, and that is a decision. - - A framework — typer, click — would buy nothing this front uses: it has one - command wearing three format names, generated by a loop, which is the shape - decorator-per-command ergonomics do not help. Behind an optional extra it - would be worse than stdlib, because `python -m math_spec latex` would stop - working on a bare install while the docs still showed it. If that trade is - ever worth making, this test is where the argument has to be made. - """ + """It is stdlib argparse over `typeset`, and that is a decision: an optional + extra would stop `python -m math_spec latex` working on a bare install.""" tree = ast.parse(Path(front.__file__).read_text()) roots = set() for node in ast.walk(tree): diff --git a/tests/typesetting/test_formats.py b/tests/typesetting/test_formats.py index b635c636..9c2a4f26 100644 --- a/tests/typesetting/test_formats.py +++ b/tests/typesetting/test_formats.py @@ -14,7 +14,7 @@ from math_spec.typesetting.format import OPERATOR_NAMES from tests.fixtures import DISPATCH_MODEL, override from tests.typesetting import golden -from tests.typesetting.fixtures import EVERY_FORMAT, TYPST, TYPST_SYMBOLS +from tests.typesetting.fixtures import EVERY_FORMAT, TYPST_SYMBOLS if TYPE_CHECKING: from pathlib import Path @@ -38,7 +38,7 @@ def test_markdown_keeps_names_out_of_the_math(): def test_typst_standalone_adds_page_setup(): assert to_typst(DISPATCH_MODEL, standalone=True).startswith('#set page') - assert not to_typst(DISPATCH_MODEL).startswith('#set page') + assert not to_typst(DISPATCH_MODEL).startswith('#set page'), 'a fragment carries no page setup' @pytest.fixture(scope='module') @@ -53,47 +53,53 @@ def test_typst_output_with_a_symbol_table_compiles(typst, tmp_path: Path): typst.compile(str(source), output=str(tmp_path / 'symbols.pdf')) -def test_a_description_of_every_special_compiles(typst, tmp_path: Path): - """Escapes that are *present* are not necessarily *right*, and only a - compiler says which. - - The golden model's description carries every character the notations - escape; this is the Typst half of that claim, and CI's `pdflatex` run over - the same file is the LaTeX half. - """ - source = tmp_path / 'specials.typ' - source.write_text(to_typst(golden.MODEL, standalone=True)) - typst.compile(str(source), output=str(tmp_path / 'specials.pdf')) - - def test_every_typst_operator_compiles(typst, tmp_path: Path): """Only a handful of operators appear in `examples/`; the rest would otherwise first fail on somebody's own model.""" probe = tmp_path / 'operators.typ' - probe.write_text('\n'.join(f'$ a {TYPST.operators[name]} b $' for name in sorted(OPERATOR_NAMES))) + probe.write_text('\n'.join(f'$ a {FORMATS["typst"].operators[name]} b $' for name in sorted(OPERATOR_NAMES))) typst.compile(str(probe), output=str(tmp_path / 'operators.pdf')) @EVERY_FORMAT -def test_the_model_description_opens_the_document(fmt: Format): +@pytest.mark.parametrize( + 'options', [pytest.param({}, id='with-a-legend'), pytest.param({'legend': False}, id='without-one')] +) +def test_the_model_description_opens_the_document(fmt: Format, options: dict): """What the file says it is, printed before anything it declares — and printed with `legend=False` too, since it is not a symbol table.""" described = override(DISPATCH_MODEL, description='least-cost dispatch of a generator fleet') - for options in ({}, {'legend': False}): - out = typeset(described, fmt, **options) - assert 'least-cost dispatch of a generator fleet' in out, f'missing with {options}' - assert out.index('least-cost dispatch') < out.index(fmt.operators['minimize']), 'it opens the document' + out = typeset(described, fmt, **options) + assert 'least-cost dispatch of a generator fleet' in out + assert out.index('least-cost dispatch') < out.index(fmt.operators['minimize']), 'it opens the document' assert 'least-cost dispatch' not in typeset(DISPATCH_MODEL, fmt), 'a model without one prints no empty paragraph' +# --------------------------------------------------------------------------- +# escaping — prose that each format would otherwise read as markup +# --------------------------------------------------------------------------- + + #: Every character the two typeset notations have to escape, in prose a #: modeller would plausibly write: the underscore in a coordinate's name is #: what #827 hit, on a description `examples/ports/pypsa_ac_dc.yaml` carried. -SPECIALS = r'flow to link_to, 100% & #1 costs $5 {net} ~ ^ \ *star* @ref