From 07ad518abd77b1a2528f2024e2998a2099aa07d6 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 16 Sep 2026 06:11:16 +0000 Subject: [PATCH] feat(program): a grouped sum and a pullback say which dimensions their walks join on `GroupSum.joined` and `At.joined` sit beside `over` and `into`: the dimensions of the key columns the walks neither consume nor produce, each once, which the operand carries and the result keeps. A consumer joining a relation's table reads all three off the node rather than deriving the third from the walks itself. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01DT1r52zPp9BdnrzFdp1xoY --- src/math_spec/program.py | 26 +++++++++++++++++++++----- tests/test_lowering.py | 13 ++++++++----- 2 files changed, 29 insertions(+), 10 deletions(-) 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