From 2f00ed1d965197b973e468888e9ac2c584431edc Mon Sep 17 00:00:00 2001 From: Mirko <2361009+mischuh@users.noreply.github.com> Date: Sat, 18 Jul 2026 16:56:31 +0200 Subject: [PATCH] fix(compiler): partition semi-additive collapse by source grain MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ROW_NUMBER() was partitioned by the queried output dimensions instead of the source's grain, so a scalar query (no dimensions) ranked the entire table and rn=1 matched only one arbitrary row instead of one per entity — e.g. ending_inventory returned 1 instead of summing the latest snapshot across all vehicles. --- canonic/compiler/strategies/semi_additive.py | 33 +++++++++++- tests/compiler/test_semi_additive.py | 55 +++++++++++++++++++- 2 files changed, 85 insertions(+), 3 deletions(-) diff --git a/canonic/compiler/strategies/semi_additive.py b/canonic/compiler/strategies/semi_additive.py index 9380f5f..31dddfb 100644 --- a/canonic/compiler/strategies/semi_additive.py +++ b/canonic/compiler/strategies/semi_additive.py @@ -134,6 +134,23 @@ def _compile_semi_additive( ) collapse_alias, collapse_dim = collapse_dim_result + # Resolve the source's natural grain (minus collapse_dimension) — this is the + # partition key for "last/first per entity". It must not be derived from the + # queried output dimensions: a scalar query (no dimensions) still needs to dedupe + # per grain entity before summing, otherwise ROW_NUMBER() ranks the whole table + # and only one arbitrary row survives (SPEC §4.2). + grain_dims: list[tuple[str, Dimension]] = [] + for grain_col in source_obj.grain: + if grain_col == sa.collapse_dimension: + continue + grain_dim_result = _find_dimension(grain_col, sources_by_name, source_name, alias_to_source) + if grain_dim_result is None: + raise Unresolved( + f"semi_additive binding {queried_name!r}: grain column {grain_col!r} of " + f"source {source_name!r} is not declared as a dimension" + ) + grain_dims.append(grain_dim_result) + # Branch: is collapse_dimension among the grouped dimensions? grouped = {dim.name for _alias, dim in dimensions} collapsed = sa.collapse_dimension not in grouped @@ -157,6 +174,7 @@ def _compile_semi_additive( collapse_alias=collapse_alias, collapse_dim=collapse_dim, dimensions=dimensions, + grain_dims=grain_dims, where_conditions=where_conditions, join_edges=join_edges, sources_by_name=sources_by_name, @@ -190,6 +208,7 @@ def _build_semi_additive( collapse_alias: str, collapse_dim: Dimension, dimensions: list[tuple[str, Dimension]], + grain_dims: list[tuple[str, Dimension]], where_conditions: list[exp.Expression], join_edges: list[JoinEdge], sources_by_name: dict[str, SemanticSource], @@ -206,14 +225,26 @@ def _build_semi_additive( order_dir = "DESC" if collapse_agg is CollapseAgg.LAST else "ASC" # Inner CTE: project grouped dimensions + raw input columns + ROW_NUMBER window. + # The window partitions by the source's grain (minus collapse_dimension), not by + # the requested output dimensions — those may be a strict subset (or unrelated, + # via a join) of the entity key needed to dedupe "last per entity" correctly. dim_names = _dimension_output_names(dimensions) inner = exp.Select() inner_projections: list[exp.Expression] = [] - partition_exprs: list[exp.Expression] = [] + seen_names: set[str] = set() for (src, dim), name in zip(dimensions, dim_names, strict=True): expr = _dimension_expr(src, dim) inner_projections.append(_alias(expr, name)) + seen_names.add(name) + + partition_exprs: list[exp.Expression] = [] + grain_names = _dimension_output_names(grain_dims) + for (src, dim), name in zip(grain_dims, grain_names, strict=True): + expr = _dimension_expr(src, dim) partition_exprs.append(expr) + if name not in seen_names: + inner_projections.append(_alias(expr, name)) + seen_names.add(name) for input_col in _input_columns(measure): inner_projections.append(_alias(exp.column(input_col, table=owner), input_col)) diff --git a/tests/compiler/test_semi_additive.py b/tests/compiler/test_semi_additive.py index 2bbee05..69e54bb 100644 --- a/tests/compiler/test_semi_additive.py +++ b/tests/compiler/test_semi_additive.py @@ -151,18 +151,69 @@ def test_ac1_collapsed_uses_row_number( def test_ac1_scalar_no_dims(resolver: ContractResolver, inventory_source: SemanticSource) -> None: - """No dimensions: collapse across all rows, return the last snapshot's inventory.""" + """No dimensions: still partition by the source grain (warehouse_id), sum the last + snapshot per warehouse — not a single arbitrary row across the whole table (GH-119 + regression: an empty PARTITION BY let ROW_NUMBER() rank the entire table, so rn = 1 + matched only one of several tied-latest rows instead of one per entity). + """ result = compile( SemanticQuery(metrics=["ending_inventory"]), resolver, [inventory_source], ) _parse_ok(result.sql) - assert "ROW_NUMBER()" in result.sql.upper() + sql_upper = result.sql.upper() + assert "ROW_NUMBER()" in sql_upper + assert "PARTITION BY" in sql_upper + assert '"WAREHOUSE_ID"' in sql_upper assert result.partial_additive is not None assert result.partial_additive.collapsed is True +def test_ac1_scalar_no_dims_sums_across_entities(resolver: ContractResolver) -> None: + """Regression: scalar query over a multi-entity snapshot table must sum the latest + value per entity, not collapse to a single row (GH-119). + """ + import duckdb + + source = SemanticSource( + name="inventory_snapshots", + connection="warehouse_pg", + table="inventory_snapshots", + grain=["warehouse_id", "snapshot_date"], + columns=[ + Column(name="warehouse_id", type="string", nullable=False), + Column(name="snapshot_date", type="date", nullable=False), + Column(name="inventory_level", type="int", nullable=False), + ], + measures=[ + Measure(name="inventory_level", expr="sum(inventory_level)", additivity="additive") + ], + dimensions=[ + Dimension(name="warehouse_id", column="warehouse_id"), + Dimension(name="snapshot_date", column="snapshot_date", granularity="day"), + ], + ) + result = compile( + SemanticQuery(metrics=["ending_inventory"]), + resolver, + [source], + ) + + con = duckdb.connect() + con.execute( + "CREATE TABLE inventory_snapshots (warehouse_id TEXT, snapshot_date DATE, inventory_level INT)" + ) + con.execute( + "INSERT INTO inventory_snapshots VALUES " + "('w1', '2024-01-01', 10), ('w2', '2024-01-01', 20), " + "('w1', '2024-02-01', 5), ('w2', '2024-02-01', 8)" + ) + sql = result.sql.replace('"inventory_snapshots"', "inventory_snapshots") + rows = con.execute(sql).fetchall() + assert rows == [(13,)] # last-per-warehouse: w1=5 + w2=8, not a single tied row + + def test_ac1_resolved_by_alias( resolver: ContractResolver, inventory_source: SemanticSource ) -> None: