diff --git a/src/math_spec/program.py b/src/math_spec/program.py index 7567aa09..52426fff 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -267,10 +267,11 @@ class GroupSum(Expression): ``walks`` says, per relation, which columns are consumed, which produced and which joined on, and is the one fact the node holds: ``coordinate`` - names the relations, ``over`` is the dims every walk consumes and ``into`` - the dims they produce, in walk order, so that several coordinates are - one grouping into a product of targets, consumed in a single join. The - result replaces every dim in ``over`` with every dim in ``into``. The + names the relations, ``over`` is the dims every walk consumes, ``into`` + the dims they produce, in walk order, and ``joined`` the dims they join + on, so that several coordinates are one grouping into a product of + targets, consumed in a single join. The result replaces every dim in + ``over`` with every dim in ``into`` and keeps every dim in ``joined``. The join keys on the consumed columns and every joined column, and on a produced column too where the operand already carries its dimension. """ @@ -290,6 +291,11 @@ def over(self) -> tuple[str, ...]: def into(self) -> tuple[str, ...]: return tuple(dim for walk in self.walks for dim in walk.produced_dims) + @property + def joined(self) -> tuple[str, ...]: + """The dims the walks join on, each once — the key columns neither consumed nor produced, which the operand carries.""" + return _joined_dims(self.walks) + @dataclass(frozen=True) class At(Expression): @@ -301,7 +307,7 @@ class At(Expression): (``Walk.is_function_read``). The join fans out, many ``over`` tuples sharing one ``into`` tuple — at each coordinate of the joined columns, which the operand carries and the result keeps. As on - :class:`GroupSum`, ``walks`` is the fact and the three are read off it. + :class:`GroupSum`, ``walks`` is the fact and the four are read off it. """ operand: ExpressionNode @@ -319,6 +325,16 @@ def over(self) -> tuple[str, ...]: def into(self) -> tuple[str, ...]: return tuple(dim for walk in self.walks for dim in walk.consumed_dims) + @property + def joined(self) -> tuple[str, ...]: + """The dims the walks join on, each once — the key columns neither consumed nor produced, which the operand carries.""" + return _joined_dims(self.walks) + + +def _joined_dims(walks: tuple[Walk, ...]) -> tuple[str, ...]: + """The dims *walks* join on, each once, in walk order — the rule :attr:`GroupSum.joined` and :attr:`At.joined` share.""" + return tuple(dict.fromkeys(dim for walk in walks for dim in walk.joined_dims)) + @dataclass(frozen=True) class Translate(Expression): diff --git a/tests/test_lowering.py b/tests/test_lowering.py index 84ddb1f8..300833a4 100644 --- a/tests/test_lowering.py +++ b/tests/test_lowering.py @@ -521,9 +521,12 @@ def test_a_relation_lowers_with_the_walk_each_call_takes(): 'a grouped sum names the column it consumes, the one it produces and the one it joins on' ) assert isinstance(zonal, GroupSum) - assert (zonal.over, zonal.into, zonal.coordinate) == (('generator',), ('zone',), ('zone_of',)), ( - 'the dims a consumer reads are read off the walk' - ) + assert (zonal.over, zonal.into, zonal.joined, zonal.coordinate) == ( + ('generator',), + ('zone',), + ('snapshot',), + ('zone_of',), + ), 'the dims a consumer reads are read off the walk' assert program.constraints['history'].lhs == GroupSum( Variable('p'), walks=(Walk(declared, ('snapshot',), ('zone',), ('generator',)),) ), 'the same table walked from its other key column' @@ -532,8 +535,8 @@ def test_a_relation_lowers_with_the_walk_each_call_takes(): 'and its adjoint consumes the value column and produces the key column' ) assert isinstance(priced, At) - assert (priced.over, priced.into) == (('generator',), ('zone',)), ( - 'an at produces the fine dims and consumes the coarse' + assert (priced.over, priced.into, priced.joined) == (('generator',), ('zone',), ('snapshot',)), ( + 'an at produces the fine dims, consumes the coarse, and joins on the rest of the key' ) p_where = program.variable('p').where assert p_where is not None