From 5f65632886b23c195933415d2621a901312ff235 Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 22 Sep 2026 12:00:30 +0000 Subject: [PATCH] feat(program): a sum through a relation is a sum over a join `GroupSum` and `Lookup` are one node, `Join`, and `sum(by=)` lowers to the existing `Sum` over it, its `over` the dims the join drops; `at` lowers to the bare `Join`. The columns a call names move to `JoinColumns`. The YAML surface is unchanged. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01SWBcNGLjNH2i4AqRsfyxaN --- docs/about/relations-as-linear-maps.md | 4 +- docs/contributing.md | 14 +++---- src/math_spec/_expression_parser.py | 8 ++-- src/math_spec/advice.py | 6 +-- src/math_spec/boundedness.py | 5 +-- src/math_spec/dimensions.py | 10 ++--- src/math_spec/lowering.py | 13 ++++--- src/math_spec/program.py | 52 +++++++++++--------------- src/math_spec/resolution.py | 12 +++--- src/math_spec/separability.py | 30 ++++++++------- src/math_spec/typesetting/walk.py | 8 ++-- tests/test_lowering.py | 52 +++++++++++++------------- tests/test_parser.py | 4 +- tests/typesetting/test_golden.py | 2 +- 14 files changed, 108 insertions(+), 112 deletions(-) diff --git a/docs/about/relations-as-linear-maps.md b/docs/about/relations-as-linear-maps.md index 98d45d93..a3df2727 100644 --- a/docs/about/relations-as-linear-maps.md +++ b/docs/about/relations-as-linear-maps.md @@ -131,8 +131,8 @@ $`x`$ over generators and $`y`$ over buses, \langle M x, y \rangle = \sum_{b} y_b \sum_{g} \mathbf{1}_R(g, b)\, x_g = \sum_{g} x_g \sum_{b} \mathbf{1}_R(g, b)\, y_b = \langle x, M^{\mathsf{T}} y \rangle, ``` -which is why the program's `Lookup` is the same `Join` as its `GroupSum`, read -without the group-by. A bare relation has the same matrix without the +which is why the program lowers `at` to a `Join` node and `sum(by=)` to the +same `Join` under a `Sum`. A bare relation has the same matrix without the functional claim. A column of $`M`$ may hold several ones, so the sum fans out and no group is one row. That is why `at` through a bare relation is refused. diff --git a/docs/contributing.md b/docs/contributing.md index 6e9d5fe3..8a9c4fd9 100644 --- a/docs/contributing.md +++ b/docs/contributing.md @@ -107,13 +107,13 @@ suffix says which layer: A node names the operation, not the verb a file writes. One verb can lower to two nodes, so the file's spelling cannot decide the name. -| File verb | Node | What the node names | -| ------------------ | ----------- | ------------------------------ | -| `sum(over=)` | `Sum` | dims removed from the result | -| `sum(by=)` | `GroupSum` | a join and group-by | -| `at(by=)` | `Lookup` | a join with no group-by | -| `shift(along=)` | `Translate` | a re-index along one dimension | -| `sum_back(along=)` | `WindowSum` | a sum over a trailing window | +| File verb | Node | What the node names | +| ------------------ | ------------------- | ------------------------------------------ | +| `sum(over=)` | `Sum` | dims removed from the result | +| `sum(by=)` | `Sum` over a `Join` | a join, and the sum over the dims it drops | +| `at(by=)` | `Join` | a join with no sum over it | +| `shift(along=)` | `Translate` | a re-index along one dimension | +| `sum_back(along=)` | `WindowSum` | a sum over a trailing window | Nothing is abbreviated. diff --git a/src/math_spec/_expression_parser.py b/src/math_spec/_expression_parser.py index 32afc92b..656dd806 100644 --- a/src/math_spec/_expression_parser.py +++ b/src/math_spec/_expression_parser.py @@ -23,7 +23,7 @@ if TYPE_CHECKING: from collections.abc import Callable, Iterator, Mapping - from math_spec.program import Join, Partition, Predicate + from math_spec.program import JoinColumns, Partition, Predicate #: The relation a comparison may carry — the three an expression may be #: written with, which is what a constraint's sense is read off. @@ -133,12 +133,12 @@ def __str__(self) -> str: @dataclass(frozen=True) class JoinNode: - """A resolved ``by=`` on ``sum`` or ``at``: the relation, as the :class:`Join` the call names.""" + """A resolved ``by=`` on ``sum`` or ``at``: the relation, as the :class:`JoinColumns` the call names.""" - join: Join + columns: JoinColumns def __str__(self) -> str: - return self.join.name + return self.columns.name @dataclass(frozen=True) diff --git a/src/math_spec/advice.py b/src/math_spec/advice.py index aef61352..85ba75b7 100644 --- a/src/math_spec/advice.py +++ b/src/math_spec/advice.py @@ -14,7 +14,7 @@ from math_spec.boundedness import unbounded_notes from math_spec.errors import Advice from math_spec.lowering import to_program -from math_spec.program import GroupSum, Lookup, walk +from math_spec.program import Join, walk if TYPE_CHECKING: from collections.abc import Mapping @@ -74,6 +74,6 @@ def _grouped_axes(program: Program) -> set[str]: """ axes: set[str] = set() for node in walk(*program.roots): - if isinstance(node, GroupSum | Lookup): - axes.update(node.join.added_dims) + if isinstance(node, Join): + axes.update(node.columns.added_dims) return axes diff --git a/src/math_spec/boundedness.py b/src/math_spec/boundedness.py index 994d2646..89d8cb03 100644 --- a/src/math_spec/boundedness.py +++ b/src/math_spec/boundedness.py @@ -24,8 +24,7 @@ Divide, Dual, Expression, - GroupSum, - Lookup, + Join, Multiply, Negate, Parameter, @@ -162,7 +161,7 @@ def _record_signs(node: Expression, sign: Sign, signs: dict[str, Sign]) -> None: _record_signs(node.base, None, signs) _record_signs(node.exponent, None, signs) return - if isinstance(node, Sum | GroupSum | Lookup | Translate | WindowSum | Cases): + if isinstance(node, Sum | Join | Translate | WindowSum | Cases): for child in children(node): _record_signs(child, sign, signs) return diff --git a/src/math_spec/dimensions.py b/src/math_spec/dimensions.py index e8ba48d1..011a1fb9 100644 --- a/src/math_spec/dimensions.py +++ b/src/math_spec/dimensions.py @@ -43,7 +43,7 @@ from math_spec.program import ( DimensionComparison, DimensionPosition, - Join, + JoinColumns, Mask, ParameterComparison, ParameterDefined, @@ -157,7 +157,7 @@ def _sum_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spec, conte return inner - {summed.name} assert isinstance(by, JoinNode), 'resolution reads sum(by=) as a join' - join = by.join + join = by.columns if missing := sorted(set(join.dropped_dims) - inner): raise DimensionError( _not_carried( @@ -174,7 +174,7 @@ def _at_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spec, contex """``at`` is the join of ``sum(by=)`` with no group-by: it joins on the dims a sum groups by and groups by the ones a sum joins on.""" by = node.kwargs['by'] assert isinstance(by, JoinNode), 'resolution reads at(by=) as a join' - join = by.join + join = by.columns if absent := sorted(set(join.dropped_dims) - inner): raise DimensionError( f'{context}: at(by={join.name}) joins on ' @@ -208,7 +208,7 @@ def _translation_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spe return inner -def _join_dims(call: str, join: Join, inner: frozenset[str], context: str) -> frozenset[str]: +def _join_dims(call: str, join: JoinColumns, inner: frozenset[str], context: str) -> frozenset[str]: """The dims after *join*: the operand's, less the dims joined on, plus the dims grouped by. A column grouped by and not joined on brings its dim, so the operand does @@ -230,7 +230,7 @@ def _join_dims(call: str, join: Join, inner: frozenset[str], context: str) -> fr return (inner - set(join.joined_dims)) | set(join.grouped_dims) -def _check_joined(call: str, use: Join | Partition, inner: frozenset[str], context: str) -> None: +def _check_joined(call: str, use: JoinColumns | Partition, inner: frozenset[str], context: str) -> None: """The columns a call joins on are matched at their dimensions, so the operand carries every one, each once. Two joined columns over one dimension would match the operand's one diff --git a/src/math_spec/lowering.py b/src/math_spec/lowering.py index b7d0d439..76fa3284 100644 --- a/src/math_spec/lowering.py +++ b/src/math_spec/lowering.py @@ -252,9 +252,9 @@ def _cases(self, node: CasesNode) -> program.Cases: def sum(self, node: FunctionCallNode) -> program.Expression: """``sum(x)``, ``sum(x, over=d)`` or ``sum(x, by=relation)``. - Two program nodes under one surface verb: reducing a dim away and reducing it - *into* another are different relational shapes, so ``by=`` decides which - before anything else is read. + One node either way: a ``sum(by=)`` is a :class:`~math_spec.program.Sum` + over a :class:`~math_spec.program.Join`, its ``over`` the dims the join + drops, so ``by=`` decides only what the sum stands over. """ by_node = node.kwargs.get('by') operand = self.expr(node.args[0]) @@ -265,13 +265,14 @@ def sum(self, node: FunctionCallNode) -> program.Expression: assert isinstance(summed, DimensionNode), 'resolution refuses a over= that is not a dimension' return program.Sum(operand, (summed.name,)) assert isinstance(by_node, JoinNode), 'resolution reads sum(by=) as a join' - return program.GroupSum(operand, join=by_node.join) + columns = by_node.columns + return program.Sum(program.Join(operand, columns), columns.dropped_dims) def at(self, node: FunctionCallNode) -> program.Expression: - """``at(x, by=relation)`` — the join of :meth:`sum`'s ``by=`` form with no group-by.""" + """``at(x, by=relation)`` — the :class:`~math_spec.program.Join` of :meth:`sum`'s ``by=`` form, with no sum over it.""" by_node = node.kwargs['by'] assert isinstance(by_node, JoinNode), 'resolution reads at(by=) as a join' - return program.Lookup(self.expr(node.args[0]), join=by_node.join) + return program.Join(self.expr(node.args[0]), by_node.columns) def sum_back(self, node: FunctionCallNode) -> program.Expression: """``sum_back(x, along=d, window=w)`` — a trailing window along one dimension. diff --git a/src/math_spec/program.py b/src/math_spec/program.py index ffe3ef04..8d1023e5 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -62,11 +62,10 @@ 'FanIn', 'FirstOf', 'Footprint', - 'GroupSum', 'Increasing', 'Join', + 'JoinColumns', 'LastOf', - 'Lookup', 'Mask', 'MaskOf', 'Multiply', @@ -238,38 +237,32 @@ class Divide: @dataclass(frozen=True) class Sum: - """Sum ``operand`` over the named dims, removing them from the result.""" + """Sum ``operand`` over the named dims, removing them from the result. - operand: Expression - over: tuple[str, ...] - - -@dataclass(frozen=True) -class GroupSum: - """Sum ``operand`` through a relation: a join on ``join.joined``, then a group-by on ``join.grouped`` with a sum. - - The operand carries every dim joined on. The result drops the dims joined - on and not grouped by, keeps the ones both joined on and grouped by, and - gains the ones grouped by and not joined on. + ``sum(x, by=relation, over=a, into=b)`` lowers to a ``Sum`` over a + :class:`Join`, ``over`` naming the dims the join drops: the group-by is + this node, and the join is its operand. """ operand: Expression - join: Join + over: tuple[str, ...] @dataclass(frozen=True) -class Lookup: - """Read ``operand`` through a relation: the join of :class:`GroupSum` with no group-by. - - The grouped columns hold the relation's whole key, which the loader - checks, so each row of the result meets one row of the relation and - reads one value. The join fans out where several key tuples share the - values joined on, at each coordinate of the columns both joined on and - grouped by. +class Join: + """Join ``operand`` to a relation on the columns ``columns`` joins on, and carry the columns it groups by. + + The operand carries every dim joined on. The result keeps every dim the + operand carries and gains the dims grouped by and not joined on; the dims + joined on and not grouped by leave only under a :class:`Sum` over them. + ``at(x, by=relation, over=a, into=b)`` lowers to a bare ``Join``: the + grouped columns hold the relation's whole key, which the loader checks, so + each row of the result meets one row of the relation and reads one value. + The join fans out where several key tuples share the values joined on. """ operand: Expression - join: Join + columns: JoinColumns @dataclass(frozen=True) @@ -372,8 +365,7 @@ class Cases: | Power | Divide | Sum - | GroupSum - | Lookup + | Join | Translate | WindowSum | Cases @@ -386,13 +378,13 @@ def fan_in(expression: Expression) -> FanIn: For the absence rules, both classes other than ``'one-to-one'`` sum several input slots into an output row. """ - if isinstance(expression, (Sum, GroupSum)): + if isinstance(expression, Sum): return 'many-to-one' if isinstance(expression, WindowSum): return 'one-to-many' if isinstance( expression, - (Constant, Parameter, Variable, Dual, Negate, Add, Multiply, Power, Divide, Lookup, Translate, Cases), + (Constant, Parameter, Variable, Dual, Negate, Add, Multiply, Power, Divide, Join, Translate, Cases), ): return 'one-to-one' assert_never(expression) @@ -408,7 +400,7 @@ def children(expression: Expression) -> tuple[Expression, ...]: return (expression.numerator, expression.divisor) if isinstance(expression, Power): return (expression.base, expression.exponent) - if isinstance(expression, (Sum, GroupSum, Lookup, Translate, WindowSum)): + if isinstance(expression, (Sum, Join, Translate, WindowSum)): return (expression.operand,) if isinstance(expression, Cases): return tuple(region.value for region in expression.regions) @@ -462,7 +454,7 @@ def dim(self, role: str) -> str: @dataclass(frozen=True) -class Join: +class JoinColumns: """One relation as one call joins it: the columns joined on, and the columns grouped by. The declaration fixes no direction; the call does, and this is the one it diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index d5f416e5..bc642fed 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -68,7 +68,7 @@ BooleanLiteral, DimensionComparison, DimensionPosition, - Join, + JoinColumns, Mask, Not, Or, @@ -223,7 +223,7 @@ class Resolved: objective: The objective's expression, ``None`` where the file declares none. relations: Each relation's columns and key, as declared — the one - copy, which every :class:`~math_spec.program.Join` and + copy, which every :class:`~math_spec.program.JoinColumns` and :class:`~math_spec.program.Partition` in the trees holds. """ @@ -617,8 +617,8 @@ def _relation_ref( return value if partition is None else PartitionNode(partition) if not ({'over', 'into'} <= set(named)): return value # the call shape refused it already, with the wording that names the rewrite - join = self._join(name, operator, named['over'], named['into']) - return value if join is None else JoinNode(join) + columns = self._join(name, operator, named['over'], named['into']) + return value if columns is None else JoinNode(columns) def _role_name(self, value: ArithmeticNode, operator: str, key: str) -> tuple[str, ...] | None: """``over=`` or ``into=`` as the column names it must be — one bare name, or a bracketed list of them.""" @@ -637,7 +637,7 @@ def _join( operator: str, from_roles: tuple[str, ...], into_roles: tuple[str, ...], - ) -> Join | None: + ) -> JoinColumns | None: """How ``sum`` or ``at`` joins relation *name*, between the columns the call named. Both ends arrive written: the call shape refuses a call that leaves @@ -680,7 +680,7 @@ def _join( ) return None kept = tuple(r for r in shape.key if r not in from_roles and r not in into_roles) - join = Join(name, shape, (*from_roles, *kept), (*into_roles, *kept)) + join = JoinColumns(name, shape, (*from_roles, *kept), (*into_roles, *kept)) one_row_per_group = set(shape.key) <= set(join.grouped) if not forward and not one_row_per_group: self.errors.append( diff --git a/src/math_spec/separability.py b/src/math_spec/separability.py index c61fca43..e17f027a 100644 --- a/src/math_spec/separability.py +++ b/src/math_spec/separability.py @@ -11,8 +11,7 @@ from math_spec.program import ( Cases, DimensionPosition, - GroupSum, - Lookup, + Join, Mask, Reach, Separability, @@ -94,9 +93,20 @@ def waits_on(dimension: str, label: str, name: str, kind: Literal['offset', 'par if row is not None: rows[label] = row masks: list[Mask | None] = [mask] + grouped = { + id(node.operand) for node in walk(*nodes) if isinstance(node, Sum) and isinstance(node.operand, Join) + } for node in walk(*nodes): if isinstance(node, Cases): masks.extend(region.when for region in node.regions) + elif isinstance(node, Sum) and isinstance(node.operand, Join): + for dimension in node.over: + report( + 'coupled', + dimension, + label, + f'groups {dimension} into {", ".join(node.operand.columns.added_dims)} — window that dimension instead, or cut only at the group edges', + ) elif isinstance(node, Sum): if reductions_couple: for dimension in node.over: @@ -106,17 +116,11 @@ def waits_on(dimension: str, label: str, name: str, kind: Literal['offset', 'par label, f'sums over {dimension} — a rolling sum_back(window=n) windows, a total over the horizon does not', ) - elif isinstance(node, GroupSum): - for dimension in node.join.dropped_dims: - report( - 'coupled', - dimension, - label, - f'groups {dimension} into {", ".join(node.join.added_dims)} — window that dimension instead, or cut only at the group edges', - ) - elif isinstance(node, Lookup): - for dimension in node.join.dropped_dims: - waits_on(dimension, label, node.join.name, 'coordinate') + elif isinstance(node, Join): + if id(node) in grouped: + continue + for dimension in node.columns.dropped_dims: + waits_on(dimension, label, node.columns.name, 'coordinate') elif isinstance(node, (Translate, WindowSum)): dimension = node.along if node.wrap: diff --git a/src/math_spec/typesetting/walk.py b/src/math_spec/typesetting/walk.py index 90cad9b7..85503a03 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -40,7 +40,7 @@ BooleanLiteral, DimensionComparison, DimensionPosition, - Join, + JoinColumns, Mask, Not, Or, @@ -449,7 +449,7 @@ def _call(self, node: FunctionCallNode, ctx: _Context) -> tuple[str, int]: by = node.kwargs['by'] assert isinstance(by, JoinNode) outer = ctx - join = by.join + join = by.columns at = {r: outer.subscript(join.dim(r)) for r in join.grouped} for read in join.dropped: ctx = ctx.looked_up(join.dim(read), self._relation_read(join.name, at, read)) @@ -457,7 +457,7 @@ def _call(self, node: FunctionCallNode, ctx: _Context) -> tuple[str, int]: if (by := node.kwargs.get('by')) is not None: assert isinstance(by, JoinNode) - join = by.join + join = by.columns dummies: dict[str, str] = {} inner = ctx for d in join.dropped_dims: @@ -480,7 +480,7 @@ def _call(self, node: FunctionCallNode, ctx: _Context) -> tuple[str, int]: domain = self.format.joined(memberships, '') return self.format.summation(domain, self._reduction_body(node.args[0], inner)), _PRECEDENCE['+'] - def _grouping(self, join: Join, dummies: Mapping[str, str], ctx: _Context) -> list[str]: + def _grouping(self, join: JoinColumns, dummies: Mapping[str, str], ctx: _Context) -> list[str]: """The conditions a grouped sum's domain carries for one join: what it fixes of the row it joins on. A join fixes its relation's key either way, so each value column it diff --git a/tests/test_lowering.py b/tests/test_lowering.py index c08ca860..d71f7b54 100644 --- a/tests/test_lowering.py +++ b/tests/test_lowering.py @@ -34,9 +34,8 @@ Dual, Expression, Footprint, - GroupSum, Join, - Lookup, + JoinColumns, Mask, Multiply, Negate, @@ -87,7 +86,7 @@ #: `lk` as `sum` joins it: joined on the key, grouped by the value, no key column left unnamed. LK = RelationDeclaration((('g', 'g'), ('h', 'h')), ('g',)) LK2 = RelationDeclaration((('g', 'g'), ('z', 'z')), ('g',)) -LK_JOIN = Join('lk', LK, ('g',), ('h',)) +LK_JOIN = JoinColumns('lk', LK, ('g',), ('h',)) AT_BUS = RelationDeclaration((('g', 'g'), ('bus', 'bus')), ('g',)) #: `fixtures.SMALL_MODEL` plus a second relation and a per-entity @@ -413,13 +412,13 @@ def test_a_power_lowers_to_a_node_of_its_own(dispatch_schema): pytest.param('sum(q, over=h)', Sum(Variable('q'), ('h',)), id='an-over-sums-away-the-dim-it-names'), pytest.param( 'sum(p, by=lk, over=g, into=h)', - GroupSum(Variable('p'), join=LK_JOIN), - id='a-grouped-sum-names-the-dim-it-joins-on-and-the-one-it-groups-by', + Sum(Join(Variable('p'), LK_JOIN), ('g',)), + id='a-grouped-sum-is-a-sum-over-a-join-of-the-dim-the-join-drops', ), pytest.param( 'at(r, by=lk, over=h, into=g)', - Lookup(Variable('r'), join=Join('lk', LK, ('h',), ('g',))), - id='a-lookup-joins-the-same-table-the-other-way', + Join(Variable('r'), JoinColumns('lk', LK, ('h',), ('g',))), + id='an-at-is-the-same-join-the-other-way-with-no-sum-over-it', ), pytest.param( "shift(p, along=g, offset=1, edge='wrap')", @@ -547,29 +546,30 @@ def test_a_relation_lowers_with_the_join_each_call_names(): declared = RelationDeclaration(columns, ('generator', 'snapshot')) assert program.relations == {'zone_of': declared}, 'the relation sits once in the program, under its name' zonal = program.constraints['zonal'].lhs - assert zonal == GroupSum( - Variable('p'), join=Join('zone_of', declared, ('generator', 'snapshot'), ('zone', 'snapshot')) - ), ( - 'a grouped sum joins on the over= column and the unnamed key column, and groups by the into= column and that key column' + columns = JoinColumns('zone_of', declared, ('generator', 'snapshot'), ('zone', 'snapshot')) + assert zonal == Sum(Join(Variable('p'), columns), ('generator',)), ( + 'a grouped sum is a sum over a join: the join names the over= column and the unnamed key column as joined on, ' + 'the into= column and that key column as grouped by, and the sum stands over the dim the join drops' ) - assert isinstance(zonal, GroupSum) - assert (zonal.join.dropped_dims, zonal.join.added_dims, zonal.join.kept) == ( + assert isinstance(zonal, Sum) and isinstance(zonal.operand, Join) + assert (zonal.operand.columns.dropped_dims, zonal.operand.columns.added_dims, zonal.operand.columns.kept) == ( ('generator',), ('zone',), ('snapshot',), ), 'the dims a consumer reads are read off the join: dropped, added, and the key columns kept' - assert zonal.join.relation is program.relations['zone_of'], ( + assert zonal.operand.columns.relation is program.relations['zone_of'], ( 'the join holds the one declaration the program holds, not an equal copy built again' ) - assert program.constraints['history'].lhs == GroupSum( - Variable('p'), join=Join('zone_of', declared, ('snapshot', 'generator'), ('zone', 'generator')) + assert program.constraints['history'].lhs == Sum( + Join(Variable('p'), JoinColumns('zone_of', declared, ('snapshot', 'generator'), ('zone', 'generator'))), + ('snapshot',), ), 'the same table joined on its other key column' priced = program.constraints['priced'].rhs - assert priced == Lookup( - Parameter('price'), join=Join('zone_of', declared, ('zone', 'snapshot'), ('generator', 'snapshot')) - ), 'and the lookup joins on the value column and groups by the key column' - assert isinstance(priced, Lookup) - assert (priced.join.dropped_dims, priced.join.added_dims, priced.join.kept) == ( + assert priced == Join( + Parameter('price'), JoinColumns('zone_of', declared, ('zone', 'snapshot'), ('generator', 'snapshot')) + ), 'and an at is the bare join, on the value column, grouped by the key column' + assert isinstance(priced, Join) + assert (priced.columns.dropped_dims, priced.columns.added_dims, priced.columns.kept) == ( ('zone',), ('generator',), ('snapshot',), @@ -592,13 +592,13 @@ def test_a_binary_variable_lowers_to_a_binary_domain(): assert program.variables['dispatch'].domain == 'binary' -def test_a_divisor_under_a_lookup_is_still_named(): +def test_a_divisor_under_a_join_is_still_named(): """`children` has to descend through every node, or a refusal loses its name.""" quotient = Divide(Variable('x'), Parameter('rate')) component_of = RelationDeclaration((('flow', 'flow'), ('component', 'component')), ('flow',)) - looked_up = Lookup(quotient, join=Join('component_of', component_of, ('component',), ('flow',))) + looked_up = Join(quotient, JoinColumns('component_of', component_of, ('component',), ('flow',))) - assert divisor_parameters(looked_up) == frozenset({'rate'}), 'the walk descends through `Lookup`' + assert divisor_parameters(looked_up) == frozenset({'rate'}), 'the walk descends through `Join`' assert divisor_parameters(Sum(looked_up, ('flow',))) == frozenset({'rate'}), 'and through a `Sum` over it' @@ -671,8 +671,8 @@ def test_walk_is_the_node_column_of_walk_regions(): Power(Parameter('c'), Constant(2.0)): 'one-to-one', Divide(Variable('p'), Parameter('c')): 'one-to-one', Sum(Variable('p'), ('g',)): 'many-to-one', - GroupSum(Variable('p'), join=Join('at_bus', AT_BUS, ('g',), ('bus',))): 'many-to-one', - Lookup(Variable('p'), join=Join('at_bus', AT_BUS, ('bus',), ('g',))): 'one-to-one', + Sum(Join(Variable('p'), JoinColumns('at_bus', AT_BUS, ('g',), ('bus',))), ('g',)): 'many-to-one', + Join(Variable('p'), JoinColumns('at_bus', AT_BUS, ('bus',), ('g',))): 'one-to-one', Translate(Variable('p'), 't', offset=1, wrap=False, fill=0.0): 'one-to-one', WindowSum(Variable('p'), 't', width=2, wrap=False): 'one-to-many', Cases((Region(Mask(ParameterDefined('c', ('g',))), Variable('p')),)): 'one-to-one', diff --git a/tests/test_parser.py b/tests/test_parser.py index 1db9619a..11089ced 100644 --- a/tests/test_parser.py +++ b/tests/test_parser.py @@ -47,7 +47,7 @@ from math_spec.program import ( And, BooleanLiteral, - Join, + JoinColumns, Not, Or, Partition, @@ -554,7 +554,7 @@ def test_a_node_prints_as_the_file_writes_it(text, printed): pytest.param(DimensionNode('t'), 't', id='a-dimension'), pytest.param(DualNode('budget'), 'dual(budget)', id='a-dual'), pytest.param( - JoinNode(Join('zone_of', _ZONE_OF, ('u',), ('zone',))), + JoinNode(JoinColumns('zone_of', _ZONE_OF, ('u',), ('zone',))), 'zone_of', id='a-relation-as-a-call-joins-it', ), diff --git a/tests/typesetting/test_golden.py b/tests/typesetting/test_golden.py index ffe6f084..baf498df 100644 --- a/tests/typesetting/test_golden.py +++ b/tests/typesetting/test_golden.py @@ -152,7 +152,7 @@ def _rendered_trees() -> Iterator[object]: #: relation it reads are the facts a node carries rather than nodes. None is a #: member of any node union, so they are subtracted from what the tree walk #: finds rather than added to what the vocabulary declares. -CARRIERS = {'CaseArm', 'Join', 'Partition', 'RelationDeclaration'} +CARRIERS = {'CaseArm', 'JoinColumns', 'Partition', 'RelationDeclaration'} def test_the_golden_model_carries_every_node_kind_the_walk_renders():