Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 32 additions & 1 deletion canonic/compiler/strategies/semi_additive.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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],
Expand All @@ -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))
Expand Down
55 changes: 53 additions & 2 deletions tests/compiler/test_semi_additive.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Loading