diff --git a/src/math_spec/exclusivity.py b/src/math_spec/exclusivity.py index 19c0ed57..51f04b34 100644 --- a/src/math_spec/exclusivity.py +++ b/src/math_spec/exclusivity.py @@ -220,8 +220,10 @@ def _subject_of(node: TypedPredicateNode) -> Subject: return Subject('variable', name) case DimensionComparisonNode(name=name): return Subject('dim', name) - case DimensionPositionNode(name=name, by=by, group=group): - return Subject('rank', name, by, group) + case DimensionPositionNode(name=name, partition=partition): + if partition is None: + return Subject('rank', name) + return Subject('rank', name, partition.name, partition.produced) case LookupDefinedNode(name=name) | LookupComparisonNode(name=name): return Subject('lookup', name) case LookupPairComparisonNode(name=name, other=other): diff --git a/src/math_spec/program.py b/src/math_spec/program.py index bf8ebf4f..0909e8df 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -1149,20 +1149,17 @@ class DimensionComparisonNode: 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: ``walked`` is its - key column over ``name``, ``group`` the value columns the group is made - of, and ``dims`` the dimensions of its other key columns, which the frame - carries. + Both sides are integers, negative counting from the end. With a + ``partition`` the position is counted within each group the lookup makes, + walked as :class:`Translate` walks one: its consumed column is the key + column over ``name``, the group is its produced columns, and its joined + columns are the other key columns, whose dimensions the frame carries. """ name: str op: PredicateOperator position: int - by: str | None = None - walked: str | None = None - group: tuple[str, ...] = () - dims: tuple[str, ...] = () + partition: Walk | None = None @dataclass(frozen=True) @@ -1316,7 +1313,7 @@ def _atom_dims(atom: TypedPredicateNode) -> frozenset[str]: case DimensionComparisonNode(): return frozenset({atom.name}) case DimensionPositionNode(): - return frozenset({atom.name, *atom.dims}) + return frozenset({atom.name, *(atom.partition.joined_dims if atom.partition is not None else ())}) case LookupComparisonNode() | LookupPairComparisonNode() | LookupDefinedNode(): return frozenset(atom.dims) case _: diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index 2314b4d9..1dc49738 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -690,6 +690,15 @@ def _walk( f'{context}: {call}: from= and into= both name {both}, and a walk goes between two sets of columns.' ) return None + for kwarg, roles in (('from', from_roles), ('into', into_roles)): + dims = [shape.dim(r) for r in roles] + if shared := sorted({d for d in dims if dims.count(d) > 1}): + self.errors.append( + f'{context}: {call}: {kwarg}={list(roles)} names two columns over {shared}, and the operand ' + f'carries each dimension once, so nothing says which column its coordinate is read at. Walk ' + f'between columns over distinct dimensions.' + ) + return None joined = tuple(r for r in (shape.key or shape.roles) if r not in from_roles and r not in into_roles) walk = Walk(shape, from_roles, into_roles, joined) if not forward and not walk.is_function_read: @@ -881,10 +890,7 @@ def _position(self, node: UnresolvedPositionNode) -> DimensionPositionNode | Unr walk = self._partition_walk(node.by, 'position', node.dimension, node.into) if walk is None: return node - (walked,) = walk.consumed - return DimensionPositionNode( - node.dimension, node.op, node.position, node.by, walked, walk.produced, walk.joined_dims - ) + return DimensionPositionNode(node.dimension, node.op, node.position, walk) def _comparison(self, node: UnresolvedComparisonNode) -> WhereNode | UnresolvedWhereNode: """``name literal``, or the one structural form ``lookup lookup``.""" diff --git a/src/math_spec/typesetting/walk.py b/src/math_spec/typesetting/walk.py index 3bdd2503..ab677f16 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -297,10 +297,11 @@ def _value_read(self, name: str, column: str, ctx: _Context) -> str: def _position_group(self, node: DimensionPositionNode, ctx: _Context) -> str: """The group a grouped position counts within: the lookup's group columns read at the row's key.""" - assert node.by is not None - lk = self.schema.lookups[node.by] - keyed = self.format.joined([ctx.subscript(dict(lk.columns)[k]) for k in lk.keys], '') - reads = [self.format.apply(self._column(node.by, column, len(lk.values) == 1), keyed) for column in node.group] + assert node.partition is not None + walk = node.partition + keyed = self.format.joined([ctx.subscript(walk.dim(k)) for k in walk.key], '') + single = len(walk.values) == 1 + reads = [self.format.apply(self._column(walk.name, column, single), keyed) for column in walk.produced] return self._tuple(reads) def _tuple(self, reads: list[str]) -> str: @@ -587,7 +588,7 @@ def _where(self, node: WhereNode, ctx: _Context) -> tuple[str, int]: ) if isinstance(node, DimensionPositionNode): - grouping = None if node.by is None else self._position_group(node, ctx) + grouping = None if node.partition is None else self._position_group(node, ctx) 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}', comparison diff --git a/tests/test_validation.py b/tests/test_validation.py index d437a43a..d7cfc52f 100644 --- a/tests/test_validation.py +++ b/tests/test_validation.py @@ -455,7 +455,7 @@ def test_it_resolves(self, mask: str, position: int, by: str | None): assert isinstance(node, DimensionPositionNode) assert node.name == 'snapshot' assert node.position == position - assert node.by == by + assert (node.partition.name if node.partition is not None else None) == by @pytest.mark.parametrize( ('mask', 'fragments'), @@ -634,6 +634,14 @@ class TestRulesDecidedWithoutData: ("from= and into= both name ['h']",), id='a-from-list-overlapping-to', ), + pytest.param( + { + 'lookups.lz': {'over': {'g': 'g', 'h0': 'h', 'h1': 'h'}, 'key': 'g'}, + 'objective': {'expression': 'sum(sum(p, by=lz, from=[h0, h1], into=g))'}, + }, + ("from=['h0', 'h1'] names two columns over ['h'], and the operand carries each dimension once",), + id='a-from-list-naming-two-columns-over-one-dimension', + ), pytest.param( {'objective': {'expression': 'sum(shift(p, over=g, offset=1, edge=0, by=lk, from=g))'}}, (