diff --git a/datajunction-server/datajunction_server/alembic/versions/2026_08_24_0000-rg0001reaggregate_add_reaggregate_to_noderevision.py b/datajunction-server/datajunction_server/alembic/versions/2026_08_24_0000-rg0001reaggregate_add_reaggregate_to_noderevision.py new file mode 100644 index 000000000..58ffd9dcf --- /dev/null +++ b/datajunction-server/datajunction_server/alembic/versions/2026_08_24_0000-rg0001reaggregate_add_reaggregate_to_noderevision.py @@ -0,0 +1,27 @@ +""" +Add reaggregate column to noderevision + +Revision ID: rg0001reaggregate +Revises: ck0001results +Create Date: 2026-08-24 00:00:00.000000+00:00 +""" + +import sqlalchemy as sa +from alembic import op + +# revision identifiers, used by Alembic. +revision = "rg0001reaggregate" +down_revision = "ck0001results" +branch_labels = None +depends_on = None + + +def upgrade(): + op.add_column( + "noderevision", + sa.Column("reaggregate", sa.JSON(), nullable=True), + ) + + +def downgrade(): + op.drop_column("noderevision", "reaggregate") diff --git a/datajunction-server/datajunction_server/api/cubes.py b/datajunction-server/datajunction_server/api/cubes.py index 58e92f6a2..37dba17f1 100644 --- a/datajunction-server/datajunction_server/api/cubes.py +++ b/datajunction-server/datajunction_server/api/cubes.py @@ -14,6 +14,10 @@ _reorder_partition_column_last, build_combiner_sql_from_preaggs, ) +from datajunction_server.construction.build_v3.cube_matcher import ( + _metric_graph_has_reaggregate, + validate_cube_reaggregate_materialization, +) from datajunction_server.construction.build_v3.cte import strip_role_suffix from datajunction_server.construction.dimensions import build_dimensions_from_cube_query from datajunction_server.database.materialization import Materialization @@ -201,6 +205,38 @@ def _build_metrics_spec( return metrics +async def _validate_cube_reaggregate_materialization( + session: AsyncSession, + cube: Node, +) -> None: + """ + Validate materialization safety using full metric decomposition. + """ + if not cube.current: # pragma: no cover + return + + if not await _metric_graph_has_reaggregate( + session, + cube.current.cube_node_metrics, + ): + return + + from datajunction_server.construction.build_v3.builder import setup_build_context + + ctx = await setup_build_context( + session=session, + metrics=cube.current.cube_node_metrics, + dimensions=cube.current.cube_node_dimensions, + filters=cube.current.cube_filters or None, + dialect=Dialect.SPARK, + use_materialized=False, + ) + validate_cube_reaggregate_materialization( + cube.current, + decomposed_metrics=ctx.decomposed_metrics, + ) + + @router.get("/cubes", name="Get all Cubes") async def get_all_cubes( *, @@ -288,6 +324,7 @@ async def cube_materialization_info( message=f"Cube node `{name}` does not exist.", http_status_code=404, ) + await _validate_cube_reaggregate_materialization(session, node) temporal_partitions = node.current.temporal_partition_columns() # type: ignore if len(temporal_partitions) != 1: raise DJInvalidInputException( @@ -520,6 +557,7 @@ async def materialize_cube( message=f"Cube '{name}' has no current revision", http_status_code=HTTPStatus.NOT_FOUND, ) + await _validate_cube_reaggregate_materialization(session, node) cube_tps = cube_revision.temporal_partition_columns() diff --git a/datajunction-server/datajunction_server/api/graphql/dataloaders.py b/datajunction-server/datajunction_server/api/graphql/dataloaders.py index 882fc0ffc..bf56a2d0e 100644 --- a/datajunction-server/datajunction_server/api/graphql/dataloaders.py +++ b/datajunction-server/datajunction_server/api/graphql/dataloaders.py @@ -377,6 +377,7 @@ async def batch_load_extracted_measures( DBNodeRevision.id, DBNodeRevision.name, DBNodeRevision.query, + DBNodeRevision.reaggregate, ), ), ) diff --git a/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py b/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py index 536f54abd..873c39b1e 100644 --- a/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py +++ b/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py @@ -15,9 +15,15 @@ MetricComponent as MetricComponent_, ) from datajunction_server.models.node import MetricDirection as MetricDirection_ +from datajunction_server.models.reaggregate import ( + DimensionReaggregateRule as DimensionReaggregateRule_, + ReaggregateSpec as ReaggregateSpec_, + ReaggregationFunction as ReaggregationFunction_, +) MetricDirection = strawberry.enum(MetricDirection_) Aggregability = strawberry.enum(Aggregability_) +ReaggregationFunction = strawberry.enum(ReaggregationFunction_) @strawberry.type @@ -32,6 +38,17 @@ class Unit: abbreviation: str | None +@strawberry.experimental.pydantic.type( + model=DimensionReaggregateRule_, + all_fields=True, +) +class DimensionReaggregateRule: ... + + +@strawberry.experimental.pydantic.type(model=ReaggregateSpec_, all_fields=True) +class ReaggregateSpec: ... + + @strawberry.experimental.pydantic.type(model=AggregationRule_, all_fields=True) class AggregationRule: ... diff --git a/datajunction-server/datajunction_server/api/graphql/scalars/node.py b/datajunction-server/datajunction_server/api/graphql/scalars/node.py index e92b7ffd3..bba7260ea 100644 --- a/datajunction-server/datajunction_server/api/graphql/scalars/node.py +++ b/datajunction-server/datajunction_server/api/graphql/scalars/node.py @@ -31,7 +31,9 @@ ) from datajunction_server.api.graphql.scalars.metricmetadata import ( DecomposedMetric, + DimensionReaggregateRule, MetricMetadata, + ReaggregateSpec, ) from datajunction_server.api.graphql.scalars.user import User from datajunction_server.api.graphql.utils import extract_fields @@ -45,6 +47,7 @@ from datajunction_server.models.node import NodeMode as NodeMode_ from datajunction_server.models.node import NodeStatus as NodeStatus_ from datajunction_server.models.node import NodeType as NodeType_ +from datajunction_server.models.reaggregate import parse_reaggregate_spec from datajunction_server.sql.parsing.backends.antlr4 import ast, parse NodeType = strawberry.enum(NodeType_) @@ -410,6 +413,26 @@ def materializations( # Only metrics will have these fields required_dimensions: list[Column] | None = None + @strawberry.field + def reaggregate(self, root: DBNodeRevision) -> ReaggregateSpec | None: + """ + Metric reaggregation declaration. + """ + if root.type != NodeType.METRIC: + return None + spec = parse_reaggregate_spec(root.reaggregate) + if not spec: + return None + return ReaggregateSpec( + rules=[ + DimensionReaggregateRule( + dimension=rule.dimension, + fn=rule.fn, # type: ignore + ) + for rule in spec.rules + ], + ) + @strawberry.field def primary_key(self, root: DBNodeRevision) -> list[str]: """ diff --git a/datajunction-server/datajunction_server/api/graphql/schema.graphql b/datajunction-server/datajunction_server/api/graphql/schema.graphql index 6d442f339..7d7dd0b38 100644 --- a/datajunction-server/datajunction_server/api/graphql/schema.graphql +++ b/datajunction-server/datajunction_server/api/graphql/schema.graphql @@ -7,6 +7,7 @@ enum Aggregability { type AggregationRule { type: Aggregability! level: [String!] + reaggregate: DimensionReaggregateRule } type Attribute { @@ -149,6 +150,11 @@ type DimensionLink { defaultValue: String } +type DimensionReaggregateRule { + dimension: String! + fn: ReaggregationFunction! +} + type Engine { name: String! version: String! @@ -397,6 +403,7 @@ type NodeRevision { dimensionLinks: [DimensionLink!]! availability: AvailabilityState materializations: [MaterializationConfig!] + reaggregate: ReaggregateSpec primaryKey: [String!]! metricMetadata: MetricMetadata isDerivedMetric: Boolean! @@ -716,6 +723,21 @@ type Query { listNamespaces: [Namespace!]! } +type ReaggregateSpec { + rules: [DimensionReaggregateRule!]! +} + +enum ReaggregationFunction { + AUTO + NONE + SUM + AVG + LAST_VALUE + FIRST_VALUE + MIN + MAX +} + type SemanticEntity { name: String! diff --git a/datajunction-server/datajunction_server/construction/build_v3/builder.py b/datajunction-server/datajunction_server/construction/build_v3/builder.py index 403641d3c..a38deca37 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/builder.py +++ b/datajunction-server/datajunction_server/construction/build_v3/builder.py @@ -21,6 +21,7 @@ ) from datajunction_server.construction.build_v3.decomposition import ( decompose_and_group_metrics, + missing_reaggregate_dimensions, ) from datajunction_server.construction.build_v3.dimensions import parse_dimension_ref from datajunction_server.construction.build_v3.filters import ( @@ -318,6 +319,14 @@ async def setup_build_context( # Add dimensions referenced in metric expressions (e.g., LAG ORDER BY) add_dimensions_from_metric_expressions(ctx, ctx.decomposed_metrics) + output_dimensions_after_expression_scan = list(ctx.dimensions) + internal_reaggregate_dimensions = missing_reaggregate_dimensions( + ctx.decomposed_metrics.values(), + output_dimensions_after_expression_scan, + ) + for dimension in internal_reaggregate_dimensions: + if dimension not in ctx.dimensions: + ctx.dimensions.append(dimension) # A second load_nodes pass is needed when either: # 1. metric expressions introduced dimension nodes not yet in ctx.nodes, OR @@ -331,8 +340,15 @@ async def setup_build_context( } missing_dim_nodes = dim_roots_after - ctx.nodes.keys() internally_added_roots = dim_roots_after - dim_roots_before_load - if missing_dim_nodes or internally_added_roots: - await load_nodes(ctx) + try: + if ( + missing_dim_nodes + or internally_added_roots + or internal_reaggregate_dimensions + ): + await load_nodes(ctx) + finally: + ctx.dimensions = output_dimensions_after_expression_scan # Classify filters into dimension filters (WHERE) and metric filters (HAVING) # This MUST happen AFTER all nodes are loaded so we can correctly identify diff --git a/datajunction-server/datajunction_server/construction/build_v3/cube_matcher.py b/datajunction-server/datajunction_server/construction/build_v3/cube_matcher.py index 2314aa0b5..54924d563 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/cube_matcher.py +++ b/datajunction-server/datajunction_server/construction/build_v3/cube_matcher.py @@ -12,9 +12,16 @@ from sqlalchemy import and_, select from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import joinedload, load_only, noload, selectinload +from sqlalchemy.orm import aliased, joinedload, load_only, noload, selectinload -from datajunction_server.construction.build_v3.decomposition import is_derived_metric +from datajunction_server.construction.build_v3.decomposition import ( + _reaggregate_dimension_requested, + is_derived_metric, + missing_reaggregate_dimensions, +) +from datajunction_server.construction.build_v3.dimension_refs import ( + split_dimension_ref, +) from datajunction_server.construction.build_v3.dimensions import parse_dimension_ref from datajunction_server.construction.build_v3.filters import ( parse_and_resolve_filters, @@ -36,14 +43,19 @@ ) from datajunction_server.database.catalog import Catalog from datajunction_server.database.column import Column -from datajunction_server.database.node import Node, NodeRevision +from datajunction_server.database.node import Node, NodeRelationship, NodeRevision from datajunction_server.database.partition import Partition from datajunction_server.errors import DJInvalidInputException from datajunction_server.instrumentation.provider import timed from datajunction_server.models.decompose import Aggregability from datajunction_server.models.dialect import Dialect from datajunction_server.models.node_type import NodeType +from datajunction_server.models.reaggregate import ( + ReaggregationFunction, + dimension_reaggregate_rules, +) from datajunction_server.naming import amenable_name +from datajunction_server.sql.decompose import MetricComponentExtractor from datajunction_server.sql.parsing import ast logger = logging.getLogger(__name__) @@ -53,6 +65,317 @@ # parsed: coverage can't be proven, so cube matching must fail SAFE (reject the # cube and fall back) rather than fail open and emit invalid SQL. _FILTER_COVERAGE_UNKNOWN = object() +ReaggregateRequirement = tuple[str, str, ReaggregationFunction] + + +def _cube_dimension_covers_reaggregate_dimension( + protected_dimension: str, + cube_dimension: str, +) -> bool: + """ + Return whether a cube dimension materializes a protected dimension. + """ + if _reaggregate_dimension_requested(protected_dimension, [cube_dimension]): + return True + + protected_base, protected_role = split_dimension_ref(protected_dimension) + cube_base, cube_role = split_dimension_ref(cube_dimension) + if protected_role != cube_role: + return False + + # Deployment validation allows a metric to author the protected dimension as + # a bare parent column (the same form required_dimensions accepts). Cubes + # expose dimensions as fully-qualified element names, so compare the short + # column name for that narrow compatibility case. + if "." not in protected_base: + return protected_base == cube_base.rsplit(".", 1)[-1] + return False + + +def _missing_reaggregate_requirements( + requirements: list[ReaggregateRequirement], + cube_dims: set[str], +) -> list[ReaggregateRequirement]: + """ + Reaggregate requirements not covered by a cube's materialized dimensions. + """ + return [ + requirement + for requirement in requirements + if not any( + _cube_dimension_covers_reaggregate_dimension( + requirement[1], + cube_dim, + ) + for cube_dim in cube_dims + ) + ] + + +async def _metric_graph_has_reaggregate( + session: AsyncSession, + metrics: list[str], +) -> bool: + """ + Cheaply test whether requested metrics or metric ancestors use reaggregate. + """ + seen: set[str] = set() + pending = set(metrics) + ParentNode = aliased(Node) + + while pending: + batch = sorted(pending - seen) + if not batch: + break + seen.update(batch) + result = await session.execute( + select( + Node.name, + NodeRevision.reaggregate, + ParentNode.name.label("parent_name"), + ) + .select_from(Node) + .join( + NodeRevision, + and_( + Node.id == NodeRevision.node_id, + Node.current_version == NodeRevision.version, + ), + ) + .outerjoin(NodeRelationship, NodeRelationship.child_id == NodeRevision.id) + .outerjoin( + ParentNode, + and_( + NodeRelationship.parent_id == ParentNode.id, + ParentNode.type == NodeType.METRIC, + ), + ) + .where(Node.name.in_(batch)) + .where(Node.type == NodeType.METRIC), + ) + for _, reaggregate, parent_name in result.all(): + if reaggregate: + return True + if parent_name and parent_name not in seen: + pending.add(parent_name) + + return False + + +async def _reaggregate_requirements_for_metrics_if_needed( + session: AsyncSession, + metrics: list[str], + dimensions: list[str], +) -> list[ReaggregateRequirement]: + """ + Run full metric decomposition only when reaggregate is possible. + """ + if not await _metric_graph_has_reaggregate(session, metrics): + return [] + return await _reaggregate_requirements_for_metrics(session, metrics, dimensions) + + +def _reaggregate_requirements_for_cube_metrics( + cube: NodeRevision, + metrics: list[str], + requested_dimensions: list[str], +) -> list[ReaggregateRequirement]: + """ + Reaggregation requirements that are absent from the output grain. + """ + requirements: list[ReaggregateRequirement] = [] + metric_names = set(metrics) + for metric_revision in cube.metric_node_revisions(): + if ( + not metric_revision + or metric_revision.name not in metric_names + or not metric_revision.reaggregate + ): + continue + for rule in dimension_reaggregate_rules(metric_revision.reaggregate): + if not _reaggregate_dimension_requested( + rule.dimension, + requested_dimensions, + ): + requirements.append( + ( + metric_revision.name, + rule.dimension, + rule.fn, + ), + ) + return requirements + + +async def _reaggregate_requirements_for_metrics( + session: AsyncSession, + metrics: list[str], + requested_dimensions: list[str], +) -> list[ReaggregateRequirement]: + """ + Reaggregation requirements from full metric decomposition. + + Direct cube metadata only sees the metrics declared on the cube. Decomposition + also catches derived metrics whose base components carry reaggregate rules. + """ + requirements: list[ReaggregateRequirement] = [] + for metric_name in metrics: + extractor = await MetricComponentExtractor.from_node_name( + metric_name, + session, + ) + components, _ = await extractor.extract(session) + for component in components: + reaggregate = component.rule.reaggregate + if ( + reaggregate + and not _reaggregate_dimension_requested( + reaggregate.dimension, + requested_dimensions, + ) + and ( + metric_name, + reaggregate.dimension, + reaggregate.fn, + ) + not in requirements + ): + requirements.append( + ( + metric_name, + reaggregate.dimension, + reaggregate.fn, + ), + ) + return requirements + + +def _reaggregate_requirements_for_decomposed_metrics( + decomposed_metrics: dict[str, DecomposedMetricInfo], + metrics: list[str], + requested_dimensions: list[str], +) -> list[ReaggregateRequirement]: + """ + Reaggregation requirements from already-decomposed planner context. + """ + requirements: list[ReaggregateRequirement] = [] + metric_names = set(metrics) + for metric_name, decomposed in decomposed_metrics.items(): + if metric_name not in metric_names: + continue + for component in decomposed.components: + reaggregate = component.rule.reaggregate + if reaggregate and not _reaggregate_dimension_requested( + reaggregate.dimension, + requested_dimensions, + ): + requirement = ( + metric_name, + reaggregate.dimension, + reaggregate.fn, + ) + if requirement not in requirements: + requirements.append(requirement) + return requirements + + +def _reaggregate_dimensions_for_cube_metrics( + cube: NodeRevision, + metrics: list[str], + requested_dimensions: list[str], +) -> list[str]: + """ + Protected dimensions a cube must materialize for semi-additive rollups. + """ + required: list[str] = [] + for _, dimension, _ in _reaggregate_requirements_for_cube_metrics( + cube, + metrics, + requested_dimensions, + ): + if dimension not in required: + required.append(dimension) + return required + + +def _format_reaggregate_requirements( + requirements: list[ReaggregateRequirement], +) -> list[str]: + """Human-readable semi-additive requirements for error messages.""" + return [ + f"{metric} -> {dimension} ({function.value})" + for metric, dimension, function in requirements + ] + + +def validate_cube_covers_reaggregate_dimensions( + cube: NodeRevision, + metrics: list[str], + dimensions: list[str], + *, + decomposed_metrics: dict[str, DecomposedMetricInfo] | None = None, + additional_requirements: list[ReaggregateRequirement] | None = None, + usage: str = "Cube", + http_status_code: int = 422, +) -> None: + """ + Fail loud when a materialized cube cannot safely serve semi-additive metrics. + + V1 materialized cube usage is only safe for semi-additive metrics when the + cube itself retains the protected dimension. Otherwise querying the cube can + silently merge across that dimension with the wrong Druid aggregator. + """ + requirements = _reaggregate_requirements_for_cube_metrics( + cube, + metrics, + dimensions, + ) + if decomposed_metrics is not None: + for requirement in _reaggregate_requirements_for_decomposed_metrics( + decomposed_metrics, + metrics, + dimensions, + ): + if requirement not in requirements: + requirements.append(requirement) + for requirement in additional_requirements or []: + if requirement not in requirements: + requirements.append(requirement) + + cube_dims = set(cube.cube_dimensions()) + missing_requirements = _missing_reaggregate_requirements(requirements, cube_dims) + if not missing_requirements: + return + + raise DJInvalidInputException( + f"{usage} `{cube.name}` cannot safely use materialized semi-additive " + "metric(s) because it does not materialize protected dimension(s) " + f"{_format_reaggregate_requirements(missing_requirements)}. " + "Add the protected dimension to the cube, or run the query with " + "use_materialized=false.", + http_status_code=http_status_code, + ) + + +def validate_cube_reaggregate_materialization( + cube: NodeRevision, + decomposed_metrics: dict[str, DecomposedMetricInfo] | None = None, +) -> None: + """ + Guard Druid cube materialization from baking in semi-additive misaggregation. + + V1 does not materialize collapsed semi-additive metrics unless the cube grain + includes the protected dimension. This avoids generating a Druid metricsSpec + that sums first/last/min/max collapse inputs across the protected dimension. + """ + validate_cube_covers_reaggregate_dimensions( + cube, + cube.cube_node_metrics, + cube.cube_node_dimensions, + decomposed_metrics=decomposed_metrics, + usage="Cube", + http_status_code=400, + ) async def _required_filter_dimensions( @@ -190,7 +513,15 @@ async def find_matching_cube( result = await session.execute(statement) candidate_cubes = result.unique().scalars().all() - + reaggregate_requirements = ( + await _reaggregate_requirements_for_metrics_if_needed( + session, + metrics, + dimensions, + ) + if candidate_cubes + else [] + ) # Find the best matching cube (smallest grain that covers all dimensions) best_match: NodeRevision | None = None best_grain_size = float("inf") @@ -214,12 +545,32 @@ async def find_matching_cube( # Druid SQL referencing a missing column). cube_dims = set(cube_rev.cube_dimensions()) - if not required_dims.issubset(cube_dims): + missing_required_dims = required_dims - cube_dims + cube_reaggregate_requirements = _reaggregate_requirements_for_cube_metrics( + cube_rev, + metrics, + dimensions, + ) + for requirement in reaggregate_requirements: + if requirement not in cube_reaggregate_requirements: + cube_reaggregate_requirements.append(requirement) + missing_reaggregate_requirements = _missing_reaggregate_requirements( + cube_reaggregate_requirements, + cube_dims, + ) + if missing_required_dims: logger.debug( f"[BuildV3] Cube {cube_rev.name} dims {cube_dims} " f"don't cover required {required_dims}", ) continue + if missing_reaggregate_requirements: + logger.debug( + f"[BuildV3] Cube {cube_rev.name} dims {cube_dims} " + f"don't cover reaggregate requirements " + f"{_format_reaggregate_requirements(missing_reaggregate_requirements)}", + ) + continue # Found a match - prefer smallest grain (less roll-up work) if len(cube_dims) < best_grain_size: @@ -244,6 +595,7 @@ async def validate_pinned_cube_covers_filters( cube: NodeRevision, dimensions: list[str], filters: list[str] | None, + metrics: list[str] | None = None, ) -> None: """ Ensure an explicitly pinned cube covers every filtered dimension. @@ -278,6 +630,23 @@ async def validate_pinned_cube_covers_filters( http_status_code=422, ) + metric_names = metrics or [ + metric_revision.name + for metric_revision in cube.metric_node_revisions() + if metric_revision + ] + validate_cube_covers_reaggregate_dimensions( + cube, + metric_names, + dimensions, + additional_requirements=await _reaggregate_requirements_for_metrics_if_needed( + session, + metric_names, + dimensions, + ), + usage="Pinned cube", + ) + async def resolve_dialect_and_engine_for_metrics( session: AsyncSession, @@ -468,6 +837,13 @@ def build_sql_from_cube_impl( Returns: GeneratedSQL with the query and column metadata. """ + validate_cube_covers_reaggregate_dimensions( + cube, + ctx.metrics, + ctx.dimensions, + decomposed_metrics=ctx.decomposed_metrics, + ) + # Build synthetic GrainGroupSQL for cube table # This applies all filters in the cube CTE's WHERE clause synthetic_grain_group = build_synthetic_grain_group( @@ -646,13 +1022,27 @@ def build_synthetic_grain_group( # (filter-only dimensions were added by add_dimensions_from_filters() in setup_build_context) dimension_aliases: dict[str, str] = {} - # Add all dimensions (requested + filter-only). We need all dimensions + output_dimensions = [ + dim for dim in ctx.dimensions if dim not in ctx.filter_dimensions + ] + internal_reaggregate_dimensions = missing_reaggregate_dimensions( + decomposed_metrics.values(), + output_dimensions, + ) + + # Add all dimensions (requested + filter-only + internal semi-additive + # protected dimensions). We need all dimensions # in the cube SELECT for proper filter resolution. # dim_short_names holds the alias (short name) used everywhere outside the CTE. # dim_physical_names holds the actual column name in the Druid table (may differ). dim_short_names = [] dim_physical_names = [] - for dim_ref in ctx.dimensions: + dimension_to_alias: dict[str, str] = {} + dimension_refs = list(ctx.dimensions) + for dim_ref in internal_reaggregate_dimensions: + if dim_ref not in dimension_refs: + dimension_refs.append(dim_ref) + for dim_ref in dimension_refs: parsed_dim = parse_dimension_ref(dim_ref) short_name = parsed_dim.column_name if parsed_dim.role: @@ -668,6 +1058,7 @@ def build_synthetic_grain_group( # the WHERE is applied directly on the cube table, so we must reference # the physical column (e.g. common_DOT_..._DOT_dateint) not the alias. dimension_aliases[dim_ref] = physical_name + dimension_to_alias[dim_ref] = short_name grain_group_columns.append( ColumnMetadata( name=short_name, @@ -741,6 +1132,17 @@ def build_synthetic_grain_group( if metric_node and not is_derived_metric(ctx, metric_node): base_metrics.append(metric_name) + reaggregate_dimension_aliases: dict[str, str] = {} + internal_reaggregate_dimension_set = set(internal_reaggregate_dimensions) + for comp in all_components: + if ( + comp.rule.reaggregate + and comp.rule.reaggregate.dimension in internal_reaggregate_dimension_set + ): + reaggregate_dimension_aliases[comp.name] = dimension_to_alias[ + comp.rule.reaggregate.dimension + ] + # Create the synthetic GrainGroupSQL # Note: We use a placeholder parent_name since the cube combines multiple parents return GrainGroupSQL( @@ -751,6 +1153,7 @@ def build_synthetic_grain_group( metrics=base_metrics, # Only base metrics, not derived parent_name=cube.name, # Use cube name as parent component_aliases=component_aliases, + reaggregate_dimension_aliases=reaggregate_dimension_aliases, is_merged=False, components=all_components, dialect=ctx.dialect, diff --git a/datajunction-server/datajunction_server/construction/build_v3/decomposition.py b/datajunction-server/datajunction_server/construction/build_v3/decomposition.py index 8a3c1b56f..b7dfecdc6 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/decomposition.py +++ b/datajunction-server/datajunction_server/construction/build_v3/decomposition.py @@ -9,10 +9,15 @@ from __future__ import annotations +from collections import defaultdict +from collections.abc import Iterable from typing import cast from sqlalchemy.ext.asyncio import AsyncSession +from datajunction_server.construction.build_v3.dimension_refs import ( + split_dimension_ref, +) from datajunction_server.construction.build_v3.types import ( BuildContext, DecomposedMetricInfo, @@ -26,6 +31,7 @@ from datajunction_server.sql.decompose import MetricComponentExtractor from datajunction_server.sql.parsing import ast from datajunction_server.sql.parsing.backends.antlr4 import parse +from datajunction_server.utils import SEPARATOR async def decompose_and_group_metrics( @@ -322,6 +328,63 @@ def get_native_grain(node: Node) -> list[str]: return pk_columns +def _dimension_ref_column(ref: str) -> str: + """Return the column part of a fully qualified or local dimension ref.""" + base, _ = split_dimension_ref(ref) + return base.rsplit(SEPARATOR, 1)[-1] + + +def _reaggregate_dimension_requested( + protected_dimension: str, + requested_dimensions: list[str], +) -> bool: + """ + Whether the protected dimension is already in the requested output grain. + + Role-qualified dimensions only count when the role matches exactly. Coarser + dimensions on the same node do not count. + """ + protected_base, protected_role = split_dimension_ref(protected_dimension) + protected_col = _dimension_ref_column(protected_dimension) + for requested in requested_dimensions: + requested_base, requested_role = split_dimension_ref(requested) + if requested == protected_dimension: + return True + if requested_base == protected_base and requested_role == protected_role: + return True + if ( + protected_role is None + and requested_role is None + and requested == protected_col + ): + return True + return False + + +def missing_reaggregate_dimensions( + decomposed_metrics: Iterable[DecomposedMetricInfo], + requested_dimensions: list[str], +) -> list[str]: + """ + Protected dimensions needed as internal grain but absent from output grain. + """ + missing: list[str] = [] + for decomposed in decomposed_metrics: + for component in decomposed.components: + if not component.rule.reaggregate: + continue + dimension = component.rule.reaggregate.dimension + if ( + not _reaggregate_dimension_requested( + dimension, + requested_dimensions, + ) + and dimension not in missing + ): + missing.append(dimension) + return missing + + def analyze_grain_groups( metric_group: MetricGroup, requested_dimensions: list[str], @@ -341,7 +404,7 @@ def analyze_grain_groups( Args: metric_group: MetricGroup with decomposed metrics - requested_dimensions: Dimensions requested by user (column names only) + requested_dimensions: Dimension refs requested by the user Returns: List of GrainGroups, one per unique grain @@ -349,11 +412,15 @@ def analyze_grain_groups( parent_node = metric_group.parent_node # Group components by their effective grain - # Key: (aggregability, tuple of additional grain columns) + # Key: (aggregability, tuple of additional grain columns, tuple of internal dimensions) grain_buckets: dict[ - tuple[Aggregability, tuple[str, ...]], + tuple[Aggregability, tuple[str, ...], tuple[str, ...]], list[tuple[Node, MetricComponent]], ] = {} + reaggregate_component_dimensions: dict[ + tuple[Aggregability, tuple[str, ...], tuple[str, ...]], + dict[str, str], + ] = {} # Track non-decomposable metrics (those with no components) non_decomposable: list[DecomposedMetricInfo] = [] @@ -369,29 +436,50 @@ def analyze_grain_groups( agg_type = component.rule.type # Explicitly type the key to satisfy mypy - key: tuple[Aggregability, tuple[str, ...]] + key: tuple[Aggregability, tuple[str, ...], tuple[str, ...]] + reaggregate_dimension = ( + component.rule.reaggregate.dimension + if ( + component.rule.reaggregate + and not _reaggregate_dimension_requested( + component.rule.reaggregate.dimension, + requested_dimensions, + ) + ) + else None + ) + internal_dimensions = ( + (reaggregate_dimension,) if reaggregate_dimension else () + ) if agg_type == Aggregability.FULL: - # FULL: no additional grain columns needed - key = (Aggregability.FULL, ()) + # FULL: no additional grain columns needed, unless a semi-additive + # component must keep its protected dimension as internal grain. + key = (Aggregability.FULL, (), internal_dimensions) elif agg_type == Aggregability.LIMITED: # LIMITED: add level columns to grain level_cols = tuple(sorted(component.rule.level or [])) - key = (Aggregability.LIMITED, level_cols) + key = (Aggregability.LIMITED, level_cols, internal_dimensions) else: # NONE # NONE: use native grain (PK columns) native_grain = get_native_grain(parent_node) key = ( Aggregability.NONE, tuple(sorted(native_grain)), + internal_dimensions, ) # pragma: no cover if key not in grain_buckets: grain_buckets[key] = [] grain_buckets[key].append((decomposed.metric_node, component)) + if reaggregate_dimension: + reaggregate_component_dimensions.setdefault(key, {})[component.name] = ( + reaggregate_dimension + ) # Convert buckets to GrainGroup objects grain_groups = [] - for (agg_type, grain_cols), components in grain_buckets.items(): + for bucket_key, components in grain_buckets.items(): + agg_type, grain_cols, _internal_dimensions = bucket_key # Map each grain expression to its SQL alias. # LIMITED: alias comes from component.grain_alias (_make_component decides: # plain column maps to the bare name, complex expr maps to component.name). @@ -412,6 +500,10 @@ def analyze_grain_groups( grain_columns=list(grain_cols), components=components, grain_col_aliases=grain_col_aliases, + reaggregate_component_dimensions=reaggregate_component_dimensions.get( + bucket_key, + {}, + ), ), ) @@ -470,9 +562,10 @@ def merge_grain_groups(grain_groups: list[GrainGroup]) -> list[GrainGroup]: Returns: List of grain groups with compatible groups merged """ - from collections import defaultdict - - # Group by parent node name + # Group by parent node name first, then by the internal grain needed for + # semi-additive collapse. LIMITED/NONE groups must not be merged into a + # semi-additive FULL group, because their extra grain columns would make the + # protected-dimension bucket contain multiple rows and corrupt MAX_BY/MIN_BY. by_parent: dict[str, list[GrainGroup]] = defaultdict(list) for gg in grain_groups: by_parent[gg.parent_node.name].append(gg) @@ -483,6 +576,12 @@ def merge_grain_groups(grain_groups: list[GrainGroup]) -> list[GrainGroup]: if len(parent_groups) == 1: # Only one group for this parent - no merge needed merged_groups.append(parent_groups[0]) + elif any(group.reaggregate_component_dimensions for group in parent_groups): + # Keep semi-additive groups isolated. Merging a protected-dimension + # group with another grain can add rows inside the protected bucket, + # which makes MAX_BY/MIN_BY pick one lower-grain row instead of the + # already-aggregated value at that protected grain. + merged_groups.extend(parent_groups) else: # Multiple groups for same parent - merge them merged = _merge_parent_grain_groups(parent_groups) @@ -520,12 +619,16 @@ def _merge_parent_grain_groups(groups: list[GrainGroup]) -> GrainGroup: # Collect all components and track their original aggregabilities all_components: list[tuple[Node, MetricComponent]] = [] component_aggregabilities: dict[str, Aggregability] = {} + reaggregate_component_dimensions: dict[str, str] = {} for gg in groups: for metric_node, component in gg.components: all_components.append((metric_node, component)) # Track original aggregability for each component component_aggregabilities[component.name] = gg.aggregability + reaggregate_component_dimensions.update( + gg.reaggregate_component_dimensions, + ) # Carry over non-decomposable metrics from every contributing group. # Without this, merging a NONE group into a FULL/LIMITED neighbor @@ -559,5 +662,6 @@ def _merge_parent_grain_groups(groups: list[GrainGroup]) -> GrainGroup: is_merged=True, component_aggregabilities=component_aggregabilities, grain_col_aliases=merged_grain_col_aliases, + reaggregate_component_dimensions=reaggregate_component_dimensions, non_decomposable_metrics=all_non_decomposable, ) diff --git a/datajunction-server/datajunction_server/construction/build_v3/dimension_refs.py b/datajunction-server/datajunction_server/construction/build_v3/dimension_refs.py new file mode 100644 index 000000000..875dc3d54 --- /dev/null +++ b/datajunction-server/datajunction_server/construction/build_v3/dimension_refs.py @@ -0,0 +1,9 @@ +"""Low-level helpers for parsing semantic dimension references.""" + + +def split_dimension_ref(ref: str) -> tuple[str, str | None]: + """Split a dimension reference into its base reference and optional role.""" + if "[" not in ref: + return ref, None + base, role = ref.rsplit("[", 1) + return base, role.rstrip("]") diff --git a/datajunction-server/datajunction_server/construction/build_v3/dimensions.py b/datajunction-server/datajunction_server/construction/build_v3/dimensions.py index d1497a4cf..87b7a37c7 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/dimensions.py +++ b/datajunction-server/datajunction_server/construction/build_v3/dimensions.py @@ -17,6 +17,9 @@ get_base_metrics_for_derived, is_derived_metric, ) +from datajunction_server.construction.build_v3.dimension_refs import ( + split_dimension_ref, +) from datajunction_server.construction.build_v3.types import ( BuildContext, DimensionRef, @@ -65,13 +68,7 @@ def parse_dimension_ref(dim_ref: str) -> DimensionRef: """ from datajunction_server.errors import DJInvalidInputException - # Extract role if present - role = None - if "[" in dim_ref: - dim_part, role_part = dim_ref.rsplit("[", 1) - role = role_part.rstrip("]") - else: - dim_part = dim_ref + dim_part, role = split_dimension_ref(dim_ref) # Split into node and column parts = dim_part.rsplit(SEPARATOR, 1) diff --git a/datajunction-server/datajunction_server/construction/build_v3/loaders.py b/datajunction-server/datajunction_server/construction/build_v3/loaders.py index 934bf069b..9c744c900 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/loaders.py +++ b/datajunction-server/datajunction_server/construction/build_v3/loaders.py @@ -87,6 +87,7 @@ async def batch_load_nodes_with_dependencies( NodeRevision.name, NodeRevision.query, NodeRevision.schema_, + NodeRevision.reaggregate, NodeRevision.table, ), # NOTE: don't noload Column.attributes — Columns are identity- diff --git a/datajunction-server/datajunction_server/construction/build_v3/measures.py b/datajunction-server/datajunction_server/construction/build_v3/measures.py index f6912654c..bf6223bf1 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/measures.py +++ b/datajunction-server/datajunction_server/construction/build_v3/measures.py @@ -1078,6 +1078,8 @@ def build_select_ast( grain_columns: list[str] | None = None, grain_col_aliases: dict[str, str] | None = None, grain_col_specs: list[tuple[ast.Expression, str]] | None = None, + output_dimension_refs: set[str] | None = None, + internal_dimension_aliases: dict[str, str] | None = None, filters: list[str] | None = None, skip_aggregation: bool = False, ) -> tuple[ast.Query, list[str]]: @@ -1096,6 +1098,9 @@ def build_select_ast( the SQL alias to use in the generated SELECT. When provided, the alias is taken from this dict (keyed by the raw expression string from ``rule.level``). + output_dimension_refs: Optional set of dimension refs to expose as user-facing + dimensions. Other resolved dimensions are join/group-only. + internal_dimension_aliases: Dimension refs to project as private grain columns. filters: Optional list of filter strings to apply as WHERE clause. Filter strings can reference dimensions (e.g., "v3.product.category = 'Electronics'") or local columns (e.g., "status = 'active'"). @@ -1118,17 +1123,31 @@ def build_select_ast( dim_aliases, joins = build_dimension_joins(ctx, resolved_dimensions, main_alias) spark_hints = _collect_spark_hints(resolved_dimensions, dim_aliases) - # Add dimension columns to projection - # Filter-only dimensions are excluded from projection but included in GROUP BY + # Add dimension columns to projection. A filter-only dimension stays hidden + # unless it is also private grain required to collapse a reaggregate metric. + internal_dimension_aliases = internal_dimension_aliases or {} + effective_filter_dimensions = ctx.filter_dimensions - set( + internal_dimension_aliases, + ) for resolved_dim in resolved_dimensions: clean_alias = ctx.alias_registry.register(resolved_dim.original_ref) - if resolved_dim.original_ref in ctx.filter_dimensions: + if resolved_dim.original_ref in effective_filter_dimensions: + continue + if ( + output_dimension_refs is not None + and resolved_dim.original_ref not in output_dimension_refs + and resolved_dim.original_ref not in internal_dimension_aliases + ): continue + output_alias = internal_dimension_aliases.get( + resolved_dim.original_ref, + clean_alias, + ) col_expr = build_dimension_col_expr( resolved_dim, main_alias, dim_aliases, - clean_alias, + output_alias, ctx=ctx, resolved_dimensions=resolved_dimensions, parent_node_name=parent_node.name, @@ -1155,7 +1174,7 @@ def build_select_ast( projected_dim_col_names: set[str] = set() projected_dim_aliases: set[str] = set() for rd in resolved_dimensions: - if rd.original_ref not in ctx.filter_dimensions: + if rd.original_ref not in effective_filter_dimensions: projected_dim_col_names.add(rd.column_name) projected_dim_aliases.add(ctx.alias_registry.register(rd.original_ref)) grain_col_refs: list[ast.Column] = [] @@ -1199,7 +1218,7 @@ def build_select_ast( main_alias, grain_col_specs, projected_dim_col_names, - ctx.filter_dimensions, + effective_filter_dimensions, parent_node_name=parent_node.name, ) @@ -1690,6 +1709,7 @@ def build_grain_group_from_preagg( select_items: list[ast.Aliasable | ast.Expression | ast.Column] = [] columns: list[ColumnMetadata] = [] component_aliases: dict[str, str] = {} + reaggregate_dimension_aliases: dict[str, str] = {} metrics_covered: set[str] = set() unique_components: list[MetricComponent] = [] seen_components: set[str] = set() @@ -1787,6 +1807,14 @@ def build_grain_group_from_preagg( ), ) + for ( + component_name, + dimension_ref, + ) in grain_group.reaggregate_component_dimensions.items(): + reaggregate_dimension_aliases[component_name] = ctx.alias_registry.register( + dimension_ref, + ) + # Add measure columns with re-aggregation (or grain columns if no merge func) for metric_node, component in grain_group.components: metrics_covered.add(metric_node.name) @@ -1947,6 +1975,7 @@ def build_grain_group_from_preagg( metrics=list(metrics_covered), parent_name=parent_node.name, component_aliases=component_aliases, + reaggregate_dimension_aliases=reaggregate_dimension_aliases, is_merged=grain_group.is_merged, component_aggregabilities=grain_group.component_aggregabilities, components=unique_components, @@ -1960,6 +1989,7 @@ def build_grain_group_sql( grain_group: GrainGroup, resolved_dimensions: list[ResolvedDimension], components_per_metric: dict[str, int], + output_dimension_refs: set[str] | None = None, ) -> GrainGroupSQL: """ Build SQL for a single grain group. @@ -2012,6 +2042,21 @@ def build_grain_group_sql( # Track mapping from component name to actual SQL alias # This is needed for metrics SQL to correctly reference component columns component_aliases: dict[str, str] = {} + reaggregate_dimension_aliases: dict[str, str] = {} + internal_dimension_aliases: dict[str, str] = {} + + if output_dimension_refs: + for resolved_dim in resolved_dimensions: + if resolved_dim.original_ref in output_dimension_refs: + ctx.alias_registry.register(resolved_dim.original_ref) + + for ( + component_name, + dimension_ref, + ) in grain_group.reaggregate_component_dimensions.items(): + dimension_alias = ctx.alias_registry.register(dimension_ref) + reaggregate_dimension_aliases[component_name] = dimension_alias + internal_dimension_aliases[dimension_ref] = dimension_alias for metric_node, component in grain_group.components: metrics_covered.add(metric_node.name) @@ -2180,6 +2225,8 @@ def build_grain_group_sql( resolved_dimensions=effective_resolved_dimensions, parent_node=parent_node, grain_columns=pass_through_columns, + output_dimension_refs=output_dimension_refs, + internal_dimension_aliases=internal_dimension_aliases, filters=ctx.dimension_filters, # Use dimension_filters only (not metric_filters) skip_aggregation=True, # Don't add GROUP BY ) @@ -2203,6 +2250,8 @@ def build_grain_group_sql( grain_columns=effective_grain_columns, grain_col_aliases=grain_group.grain_col_aliases or None, grain_col_specs=grain_col_specs, # already parsed above + output_dimension_refs=output_dimension_refs, + internal_dimension_aliases=internal_dimension_aliases, filters=ctx.dimension_filters, # Use dimension_filters only (not metric_filters) skip_aggregation=skip_agg, ) @@ -2211,7 +2260,12 @@ def build_grain_group_sql( columns_metadata = [] # Add dimension columns (skip filter-only dimensions as they're not in projection) + effective_output_dimension_refs = output_dimension_refs or { + resolved_dim.original_ref for resolved_dim in resolved_dimensions + } for resolved_dim in resolved_dimensions: + if resolved_dim.original_ref not in effective_output_dimension_refs: + continue # Skip filter-only dimensions from column metadata if resolved_dim.original_ref in ctx.filter_dimensions: continue @@ -2268,6 +2322,30 @@ def build_grain_group_sql( ), ) + for dimension_ref, dimension_alias in internal_dimension_aliases.items(): + internal_resolved_dim = next( + (dim for dim in resolved_dimensions if dim.original_ref == dimension_ref), + None, + ) + if not internal_resolved_dim: + continue # pragma: no cover + dim_node = ctx.nodes.get(internal_resolved_dim.node_name) + col_type = ( + get_column_type(parent_node, internal_resolved_dim.column_name) + if internal_resolved_dim.is_local + else get_column_type(dim_node, internal_resolved_dim.column_name) + if dim_node + else "string" + ) + columns_metadata.append( + ColumnMetadata( + name=dimension_alias, + semantic_name=dimension_ref, + type=col_type, + semantic_type="dimension", + ), + ) + # Add metric component columns # All decomposed metrics are now treated as components - no special case for single-component for comp_alias, component, metric_node in component_metadata: @@ -2313,6 +2391,8 @@ def build_grain_group_sql( if grain_group.aggregability != Aggregability.NONE: # FULL/LIMITED: dimensions are part of the grain for resolved_dim in resolved_dimensions: + if resolved_dim.original_ref not in effective_output_dimension_refs: + continue # Skip filter-only dimensions from grain if resolved_dim.original_ref in ctx.filter_dimensions: continue @@ -2327,6 +2407,9 @@ def build_grain_group_sql( for gc_alias in effective_grain_aliases: if gc_alias not in full_grain: # pragma: no branch full_grain.append(gc_alias) + for dimension_alias in internal_dimension_aliases.values(): + if dimension_alias not in full_grain: + full_grain.append(dimension_alias) # Sort for deterministic output full_grain.sort() @@ -2339,6 +2422,7 @@ def build_grain_group_sql( metrics=list(metrics_covered), parent_name=grain_group.parent_node.name, component_aliases=component_aliases, + reaggregate_dimension_aliases=reaggregate_dimension_aliases, is_merged=grain_group.is_merged, component_aggregabilities=grain_group.component_aggregabilities, components=unique_components, @@ -2374,19 +2458,23 @@ def process_metric_group( for decomposed in metric_group.decomposed_metrics: components_per_metric[decomposed.metric_node.name] = len(decomposed.components) - # Analyze grain groups - split by aggregability - # Extract just the column names from dimensions for grain analysis - dim_column_names = [parse_dimension_ref(d).column_name for d in ctx.dimensions] - grain_groups = analyze_grain_groups(metric_group, dim_column_names) + # Filter-only dimensions must be resolved for WHERE clauses, but they are not + # part of the requested output grain. Treating one as output can suppress a + # reaggregate metric's private protected-dimension grain. + resolution_dimensions = list(ctx.dimensions) + output_dimensions = [ + dimension + for dimension in resolution_dimensions + if dimension not in ctx.filter_dimensions + ] + output_dimension_refs = set(output_dimensions) + grain_groups = analyze_grain_groups(metric_group, output_dimensions) # Merge compatible grain groups from same parent into single CTEs # This optimization reduces duplicate JOINs by outputting raw values # at finest grain, with aggregations applied in final SELECT grain_groups = merge_grain_groups(grain_groups) - # Resolve dimensions (find join paths) - shared across grain groups - resolved_dimensions = resolve_dimensions(ctx, parent_node) - # Build SQL for each grain group. Fan-out risk is flagged inside # build_grain_group_sql, where the full set of emitted join paths is known. grain_group_sqls: list[GrainGroupSQL] = [] @@ -2394,13 +2482,28 @@ def process_metric_group( # Reset alias registry for each grain group to avoid conflicts ctx.alias_registry = AliasRegistry() ctx._table_alias_counter = 0 - - grain_group_sql = build_grain_group_sql( - ctx, - grain_group, - resolved_dimensions, - components_per_metric, + original_dimensions = ctx.dimensions + internal_dimensions = [ + dimension_ref + for dimension_ref in dict.fromkeys( + grain_group.reaggregate_component_dimensions.values(), + ) + if dimension_ref not in output_dimension_refs + ] + ctx.dimensions = list( + dict.fromkeys(resolution_dimensions + internal_dimensions), ) + try: + resolved_dimensions = resolve_dimensions(ctx, parent_node) + grain_group_sql = build_grain_group_sql( + ctx, + grain_group, + resolved_dimensions, + components_per_metric, + output_dimension_refs=output_dimension_refs, + ) + finally: + ctx.dimensions = original_dimensions grain_group_sqls.append(grain_group_sql) return grain_group_sqls @@ -2514,6 +2617,33 @@ def find_parent_for_window_metric( else: return None, base_metrics # pragma: no cover + def find_base_metrics_for_window_parent( + metric_name: str, + visited: set[str], + ) -> set[str]: + """ + Recursively find grain-group metrics needed to compute a window parent. + + Window expressions may reference a derived metric directly, but the + window grain group still needs the underlying base metric components. + """ + if metric_name in visited: + return set() # pragma: no cover + visited.add(metric_name) + + for gg in existing_grain_groups: + if metric_name in gg.metrics: + return {metric_name} + + grain_group_metrics: set[str] = set() + for parent_name in ctx.parent_map.get(metric_name, []): + parent_node = ctx.nodes.get(parent_name) + if parent_node and parent_node.type.value == "metric": # pragma: no branch + grain_group_metrics.update( + find_base_metrics_for_window_parent(parent_name, visited), + ) + return grain_group_metrics + # Group window metrics by (ORDER BY grain, parent fact) # This ensures window metrics from different facts are processed separately # Key: (frozenset of grain cols, parent_name or "cross_fact") @@ -2525,13 +2655,18 @@ def find_parent_for_window_metric( parent_name, base_metrics = find_parent_for_window_metric(metric_name) # Use "cross_fact" as a marker for cross-fact window metrics parent_key = parent_name if parent_name else "cross_fact" + component_metrics = set() + for base_metric in base_metrics: + component_metrics.update( + find_base_metrics_for_window_parent(base_metric, set()), + ) group_key = (grain_key, parent_key) if group_key not in grain_parent_to_metrics: grain_parent_to_metrics[group_key] = [] grain_parent_to_base_metrics[group_key] = set() grain_parent_to_metrics[group_key].append(metric_name) - grain_parent_to_base_metrics[group_key].update(base_metrics) + grain_parent_to_base_metrics[group_key].update(component_metrics) additional_grain_groups: list[GrainGroupSQL] = [] diff --git a/datajunction-server/datajunction_server/construction/build_v3/metrics.py b/datajunction-server/datajunction_server/construction/build_v3/metrics.py index fb8cb3b28..36b5dbb4e 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/metrics.py +++ b/datajunction-server/datajunction_server/construction/build_v3/metrics.py @@ -50,13 +50,149 @@ ) from datajunction_server.errors import DJInvalidInputException from datajunction_server.models.decompose import Aggregability +from datajunction_server.models.dialect import Dialect from datajunction_server.models.node_type import NodeType +from datajunction_server.models.reaggregate import ReaggregationFunction from datajunction_server.sql.decompose import wrap_divisions_in_nullif from datajunction_server.sql.parsing import ast +from datajunction_server.sql.parsing.backends.antlr4 import parse logger = logging.getLogger(__name__) +def _build_reaggregate_collapse_expression( + collapse_function: ReaggregationFunction, + dialect: Dialect, + value_ref: ast.Expression, + protected_dim_ref: ast.Expression, +) -> ast.Function: + """ + Build the final aggregation for a semi-additive metric collapse. + """ + if collapse_function == ReaggregationFunction.LAST_VALUE: + function_name = "LATEST_BY" if dialect == Dialect.DRUID else "MAX_BY" + return ast.Function( + ast.Name(function_name), + args=[value_ref, protected_dim_ref], + ) + if collapse_function == ReaggregationFunction.FIRST_VALUE: + function_name = "EARLIEST_BY" if dialect == Dialect.DRUID else "MIN_BY" + return ast.Function( + ast.Name(function_name), + args=[value_ref, protected_dim_ref], + ) + if collapse_function == ReaggregationFunction.MIN: + return ast.Function(ast.Name("MIN"), args=[value_ref]) + if collapse_function == ReaggregationFunction.MAX: + return ast.Function(ast.Name("MAX"), args=[value_ref]) + + raise DJInvalidInputException( + f"Unsupported semi-additive collapse function: {collapse_function}", + ) + + +def _references_component(expr: ast.Node, component_name: str) -> bool: + """ + Return whether an expression references a decomposed metric component. + """ + for col in expr.find_all(ast.Column): + if col.name and col.name.name == component_name: + return True + return False + + +def _replace_reaggregate_merge_expression( + combiner_ast: ast.Expression, + component_name: str, + merge_function: str, + collapse_expr: ast.Expression, +) -> ast.Expression: + """ + Replace the component's merge aggregate in a combiner with a collapse expr. + + Semi-additive metrics still need the full decomposed combiner expression. + For example, ``SUM(balance) / 100`` should become + ``MAX_BY(balance_sum, date_id) / 100``, not just ``MAX_BY(...)``. + """ + expr_ast = ( + cast(ast.Expression, combiner_ast.child) + if isinstance(combiner_ast, ast.Alias) + else combiner_ast + ) + target_function = merge_function.upper() + + for func in list(expr_ast.find_all(ast.Function)): + if func.name.name.upper() == target_function and _references_component( + func, + component_name, + ): + if func is expr_ast: + return collapse_expr + if func.parent: + func.parent.replace(from_=func, to=collapse_expr) + return expr_ast + + raise DJInvalidInputException( + "Unsupported semi-additive metric shape: could not find the component " + "merge expression in the metric combiner.", + ) + + +def _source_dimension_alias( + source_dimension_aliases: dict[str, str], + dimension_ref: str, +) -> str | None: + """ + Return the source CTE's column alias for a dimension ref. + + Semantic refs, including their roles, must match exactly. + """ + return source_dimension_aliases.get(dimension_ref) + + +def _metric_refs_from_query(ctx: BuildContext, metric_name: str) -> set[str]: + """ + Return metric refs used by a metric's query expression. + """ + metric_node = ctx.nodes.get(metric_name) + if not metric_node: + return set() # pragma: no cover + + if not metric_node.current or not metric_node.current.query: + return set() # pragma: no cover + + query = parse(metric_node.current.query) + refs: set[str] = set() + for col in query.select.projection[0].find_all(ast.Column): + full_name = get_column_full_name(col) + if not full_name and hasattr(col, "identifier"): + full_name = col.identifier() + node = ctx.nodes.get(full_name) if full_name else None + if node and node.type == NodeType.METRIC: + refs.add(full_name) + return refs + + +def _metric_parent_refs(ctx: BuildContext, metric_name: str) -> list[str]: + """ + Return metric parents for a metric, using the parsed query as a fallback. + """ + refs: list[str] = [] + seen: set[str] = set() + for parent_name in ctx.parent_map.get(metric_name, []): + parent_node = ctx.nodes.get(parent_name) + if parent_node and parent_node.type == NodeType.METRIC: + refs.append(parent_name) + seen.add(parent_name) + + for parent_name in sorted(_metric_refs_from_query(ctx, metric_name)): + if parent_name not in seen: + refs.append(parent_name) + seen.add(parent_name) + + return refs + + def classify_filters( filters: list[str], ctx: BuildContext, @@ -207,6 +343,31 @@ def get_comp_aggregability(comp_name: str) -> Aggregability: comp = decomposed.components[0] orig_agg = get_comp_aggregability(comp.name) + if comp.rule.reaggregate and comp.name in gg.reaggregate_dimension_aliases: + _, col_name = comp_mappings[comp.name] + value_ref = make_column_ref(col_name, cte_alias) + protected_dim_ref = make_column_ref( + gg.reaggregate_dimension_aliases[comp.name], + cte_alias, + ) + if not comp.merge: # pragma: no cover + raise DJInvalidInputException( + "Unsupported semi-additive metric shape: missing merge function.", + ) + collapse_expr = _build_reaggregate_collapse_expression( + comp.rule.reaggregate.fn, + gg.dialect, + value_ref, + protected_dim_ref, + ) + combiner_ast = parse(f"SELECT {decomposed.combiner}").select.projection[0] + return _replace_reaggregate_merge_expression( + cast(ast.Expression, combiner_ast), + comp.name, + comp.merge, + collapse_expr, + ) + if orig_agg == Aggregability.LIMITED: _, col_name = comp_mappings[comp.name] col_ref = make_column_ref(col_name, cte_alias) @@ -413,6 +574,41 @@ def collect_and_build_ctes( return all_cte_asts, cte_aliases +def _validate_reaggregate_base_group_join_safety( + grain_groups: list[GrainGroupSQL], +) -> None: + """ + Fail closed for final joins that can duplicate non-idempotent aggregates. + + Semi-additive groups retain the protected dimension as private grain, so a + final join by only the requested output dimensions can fan out other CTEs. + LIMITED groups that were pre-aggregated to one row per output dimension use + MAX passthrough in the final select and are safe. Reaggregate groups are + idempotent under duplicate protected-dimension rows for their collapse. + Other separate groups must be rejected unless they were merged into the + protected grain earlier. + """ + if len(grain_groups) <= 1 or not any( + gg.reaggregate_dimension_aliases for gg in grain_groups + ): + return + + unsafe_groups = [ + gg.parent_name + for gg in grain_groups + if not gg.reaggregate_dimension_aliases + and not (gg.aggregability == Aggregability.LIMITED and gg.is_pre_aggregated) + ] + if not unsafe_groups: + return + + raise DJInvalidInputException( + "Semi-additive live SQL with multiple base grain groups is not " + "supported when another group can be fanned out before final " + f"aggregation. Unsafe parents: {sorted(set(unsafe_groups))}.", + ) + + def get_dimension_types( grain_groups: list[GrainGroupSQL], ) -> dict[str, str]: @@ -559,6 +755,7 @@ def process_base_metrics( """ component_refs: dict[str, ColumnRef] = {} metric_exprs: dict[str, MetricExprInfo] = {} + precollapsed_reaggregate_metrics: set[str] = set() # Collect all metrics in grain groups all_metrics: set[str] = set() @@ -601,6 +798,11 @@ def process_base_metrics( alias, gg, ) + if any( + comp.rule.reaggregate and comp.name in gg.reaggregate_dimension_aliases + for comp in decomposed.components + ): + precollapsed_reaggregate_metrics.add(metric_name) # Convert component mappings to ColumnRef objects for comp_name, (cte_alias, col_name) in comp_mappings.items(): component_refs[comp_name] = ColumnRef( @@ -621,6 +823,7 @@ def process_base_metrics( all_metrics=all_metrics, metric_exprs=metric_exprs, component_refs=component_refs, + precollapsed_reaggregate_metrics=precollapsed_reaggregate_metrics, ) @@ -1230,11 +1433,7 @@ def build_window_agg_cte_from_grain_group( # Find the base metrics that the window metrics reference base_metrics_needed: set[str] = set() for window_metric_name in window_grain_group.window_metrics_served: - parent_names = ctx.parent_map.get(window_metric_name, []) - for parent_name in parent_names: - parent_node = ctx.nodes.get(parent_name) - if parent_node and parent_node.type == NodeType.METRIC: # pragma: no branch - base_metrics_needed.add(parent_name) + base_metrics_needed.update(_metric_parent_refs(ctx, window_metric_name)) def is_derived_metric(metric_name: str) -> bool: """Check if a metric is derived (has metric parents, no direct components).""" @@ -1242,11 +1441,8 @@ def is_derived_metric(metric_name: str) -> bool: if not decomposed or not isinstance(decomposed, DecomposedMetricInfo): return False # pragma: no cover # Derived metrics have parent metrics but no direct components in grain groups - parent_metrics = ctx.parent_map.get(metric_name, []) - has_metric_parents = any( - ctx.nodes.get(p) and ctx.nodes.get(p).type == NodeType.METRIC # type: ignore - for p in parent_metrics - ) + parent_metrics = _metric_parent_refs(ctx, metric_name) + has_metric_parents = bool(parent_metrics) return has_metric_parents and not decomposed.components def get_metric_aggregation_expr( @@ -1284,7 +1480,7 @@ def get_metric_aggregation_expr( # Try to find matching metric parent_metric_name = None - for parent_name in ctx.parent_map.get(metric_name, []): + for parent_name in _metric_parent_refs(ctx, metric_name): parent_short = get_short_name(parent_name) if parent_short == short_col_name or col_name.endswith(parent_name): parent_metric_name = parent_name @@ -1300,21 +1496,11 @@ def get_metric_aggregation_expr( # Swap the column node with the parent's expression col_node.swap(parent_expr) else: - # Base metric: replace component references with CTE column refs - # Use base_grain_group's component_aliases since that's the source CTE - for col_node in combiner_ast.find_all(ast.Column): - col_full_name = ( - col_node.identifier() if hasattr(col_node, "identifier") else "" - ) - # Check if this matches a component alias - for ( - comp_name, - comp_alias, - ) in base_grain_group.component_aliases.items(): # pragma: no branch - if col_full_name == comp_name or col_full_name.endswith(comp_alias): - col_node.name = ast.Name(comp_alias) - col_node._table = ast.Table(ast.Name(source_cte_alias)) - break + combiner_ast, _ = build_base_metric_expression( + decomposed, + source_cte_alias, + base_grain_group, + ) return combiner_ast @@ -1326,13 +1512,13 @@ def get_metric_aggregation_expr( # Get aggregation expression (handles both base and derived metrics) combiner_ast = get_metric_aggregation_expr(base_metric_name, set()) - if not combiner_ast: + if combiner_ast is None: continue # pragma: no cover short_name = get_short_name(base_metric_name) - aliased = combiner_ast.set_alias(ast.Name(short_name)) # type: ignore - aliased.set_as(True) - projection.append(aliased) + metric_alias = ast.Alias(child=combiner_ast, alias=ast.Name(short_name)) + metric_alias.set_as(True) + projection.append(metric_alias) # Build FROM clause from_clause = ast.From( @@ -1354,6 +1540,8 @@ def build_window_agg_cte_from_base_metrics( window_grain_group: GrainGroupSQL, base_metrics_cte_alias: str, ctx: BuildContext, + source_dimension_aliases: dict[str, str], + precollapsed_reaggregate_metrics: set[str], decomposed_metrics: dict[str, DecomposedMetricInfo], ) -> ast.Query: """ @@ -1373,6 +1561,10 @@ def build_window_agg_cte_from_base_metrics( window_grain_group: Window grain group with metadata (dimensions, metrics served) base_metrics_cte_alias: Alias of the base_metrics CTE (typically "base_metrics") ctx: Build context + source_dimension_aliases: Semantic dimension refs mapped to the column + aliases actually projected by the base_metrics CTE + precollapsed_reaggregate_metrics: Semi-additive metrics whose protected + dimension was removed by an earlier collapse decomposed_metrics: Decomposed metric info Returns: @@ -1394,62 +1586,124 @@ def build_window_agg_cte_from_base_metrics( # Find the base metrics that the window metrics reference base_metrics_needed: set[str] = set() for window_metric_name in window_grain_group.window_metrics_served: - parent_names = ctx.parent_map.get(window_metric_name, []) - for parent_name in parent_names: - parent_node = ctx.nodes.get(parent_name) - if parent_node and parent_node.type == NodeType.METRIC: # pragma: no branch - base_metrics_needed.add(parent_name) - - # Build reaggregation expressions for each base metric - # Since we're selecting from base_metrics, we reference the metric columns directly + base_metrics_needed.update(_metric_parent_refs(ctx, window_metric_name)) + + window_dimension_aliases = { + col.name + for col in window_grain_group.columns + if col.semantic_type == "dimension" + } + + def build_base_metric_reaggregation_expr( + metric_name: str, + decomposed: DecomposedMetricInfo, + ) -> ast.Expression: + """ + Build a reaggregation expression for a base metric from base_metrics. + """ + short_name = get_short_name(metric_name) + metric_ref = make_column_ref(short_name, base_metrics_cte_alias) + + if len(decomposed.components) == 1: + comp = decomposed.components[0] + if comp.rule.reaggregate: + protected_dim_alias = _source_dimension_alias( + source_dimension_aliases, + comp.rule.reaggregate.dimension, + ) + if ( + protected_dim_alias + and protected_dim_alias not in window_dimension_aliases + ): + return _build_reaggregate_collapse_expression( + comp.rule.reaggregate.fn, + ctx.dialect, + metric_ref, + make_column_ref(protected_dim_alias, base_metrics_cte_alias), + ) + if ( + protected_dim_alias is None + and metric_name not in precollapsed_reaggregate_metrics + ): + raise DJInvalidInputException( + "Unsupported semi-additive metric shape: protected " + f"dimension '{comp.rule.reaggregate.dimension}' is not " + "projected by the base metrics CTE.", + ) + + if decomposed.aggregability == Aggregability.LIMITED: + raise DJInvalidInputException( + "Unsupported distinct metric reaggregation: metric " + f"'{metric_name}' cannot be collapsed from the base metrics CTE " + "because it no longer retains the distinct grain key required " + "for correct reaggregation.", + ) + + if decomposed.aggregability == Aggregability.NONE: # pragma: no cover + # NONE: non-additive (like AVG), need to recompute. + # For AVG, we need the raw sum and count, but those are in components. + # For now, just use the column directly (window function will handle it). + return metric_ref + + # base_metrics exposes metric columns, not raw component/grain columns. + # Reaggregate leaf metrics from those projected metric values. + return ast.Function(ast.Name("SUM"), args=[metric_ref]) + + def build_metric_reaggregation_expr( + metric_name: str, + visited: set[str], + ) -> ast.Expression | None: + """ + Build a reaggregation expression for base or derived metrics. + """ + if metric_name in visited: + return None # pragma: no cover + visited.add(metric_name) + + parent_refs = _metric_parent_refs(ctx, metric_name) + if not parent_refs: + decomposed = decomposed_metrics.get(metric_name) + if not decomposed: + return None # pragma: no cover + return build_base_metric_reaggregation_expr(metric_name, decomposed) + + metric_node = ctx.nodes.get(metric_name) + if not metric_node: + return None # pragma: no cover + if not metric_node.current or not metric_node.current.query: + return None # pragma: no cover + original_query = parse(metric_node.current.query) + expr_ast = deepcopy(original_query.select.projection[0]) + if isinstance(expr_ast, ast.Alias): + expr_ast = expr_ast.child # pragma: no cover + + for col_node in list(expr_ast.find_all(ast.Column)): + parent_name = get_column_full_name(col_node) + if not parent_name and hasattr(col_node, "identifier"): + parent_name = col_node.identifier() + if not parent_name or parent_name not in parent_refs: + continue + + parent_expr = build_metric_reaggregation_expr( + parent_name, + visited.copy(), + ) + if parent_expr is not None: + col_node.swap(parent_expr) + + wrap_divisions_in_nullif(cast(ast.Expression, expr_ast)) + return expr_ast # type: ignore + + # Build reaggregation expressions for each parent metric. for base_metric_name in sorted(base_metrics_needed): - decomposed = decomposed_metrics.get(base_metric_name) - if not decomposed: + agg_expr = build_metric_reaggregation_expr(base_metric_name, set()) + if agg_expr is None: continue # pragma: no cover short_name = get_short_name(base_metric_name) - - # For non-additive metrics (COUNT DISTINCT), we need COUNT(DISTINCT grain_col) - # For additive metrics, we can SUM/AVG the pre-computed column - if decomposed.aggregability == Aggregability.LIMITED: - # LIMITED: needs COUNT DISTINCT at this grain - # Find the grain column for COUNT DISTINCT - if decomposed.components and decomposed.components[0].rule.level: - grain_col = decomposed.components[0].rule.level[0] - # The grain column is in base_metrics as a dimension column - # Actually, for COUNT DISTINCT, the raw grain column should be in - # the base grain group, not base_metrics. We need to reference - # the original grain column from base_metrics. - # For now, use the metric column - this works because base_metrics - # computes COUNT DISTINCT at the fine grain, and we re-compute at coarser grain - col_ref = make_column_ref(grain_col, base_metrics_cte_alias) - agg_expr = ast.Function( - ast.Name("COUNT"), - args=[col_ref], - quantifier=ast.SetQuantifier.Distinct, - ) - else: # pragma: no cover - # Fallback: SUM the metric column - col_ref = make_column_ref(short_name, base_metrics_cte_alias) - agg_expr = ast.Function(ast.Name("SUM"), args=[col_ref]) - elif decomposed.aggregability == Aggregability.FULL: # pragma: no cover - # FULL: additive metric, use SUM - col_ref = make_column_ref(short_name, base_metrics_cte_alias) - agg_expr = ast.Function(ast.Name("SUM"), args=[col_ref]) - elif decomposed.aggregability == Aggregability.NONE: # pragma: no cover - # NONE: non-additive (like AVG), need to recompute - # For AVG, we need the raw sum and count, but those are in components - # For now, just use the column directly (window function will handle it) - col_ref = make_column_ref(short_name, base_metrics_cte_alias) - agg_expr = col_ref # type: ignore - else: # pragma: no cover - # Default: SUM - col_ref = make_column_ref(short_name, base_metrics_cte_alias) - agg_expr = ast.Function(ast.Name("SUM"), args=[col_ref]) - - aliased = agg_expr.set_alias(ast.Name(short_name)) # type: ignore - aliased.set_as(True) - projection.append(aliased) + metric_alias = ast.Alias(child=agg_expr, alias=ast.Name(short_name)) + metric_alias.set_as(True) + projection.append(metric_alias) # Build FROM clause - just base_metrics from_clause = ast.From( @@ -1760,6 +2014,7 @@ def generate_metrics_sql( base_grain_groups, skip_pre_agg=will_have_base_metrics_cte, ) + _validate_reaggregate_base_group_join_safety(base_grain_groups) # Build dimension info and projection # Filter out filter-only dimensions (they're needed for WHERE but not output) @@ -1995,6 +2250,8 @@ def collect_derived_dependencies(metric_name: str, visited: set[str]) -> None: wgg, source_cte_alias, # base_metrics CTE ctx, + dict(dim_info), + base_metrics_result.precollapsed_reaggregate_metrics, decomposed_metrics, ) else: diff --git a/datajunction-server/datajunction_server/construction/build_v3/preagg_matcher.py b/datajunction-server/datajunction_server/construction/build_v3/preagg_matcher.py index 575a22901..ec12e640e 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/preagg_matcher.py +++ b/datajunction-server/datajunction_server/construction/build_v3/preagg_matcher.py @@ -65,6 +65,25 @@ def get_required_measure_identities( } +def required_reaggregate_grain( + ctx: BuildContext, + node_rev_id: int, + grain_group: GrainGroup, +) -> set[str]: + """ + Protected dimensions a pre-agg must retain for final semi-additive collapse. + + ``reaggregate_component_dimensions`` is populated only when the requested + output grain omitted the protected dimension. In that case a pre-agg at the + output grain would have already collapsed away the key needed by MAX_BY / + MIN_BY / min / max in the metrics layer, so it is not substitutable. + """ + return { + canonical_dimension_ref(ctx, node_rev_id, dimension_ref) + for dimension_ref in grain_group.reaggregate_component_dimensions.values() + } + + def canonical_dimension_ref( ctx: BuildContext, node_rev_id: int, @@ -256,6 +275,12 @@ def find_matching_preagg( if not required_measures: return None + reaggregate_grain = required_reaggregate_grain( + ctx, + node_rev_id, + grain_group, + ) + requested_grain_set = { canonical_dimension_ref(ctx, node_rev_id, rdim.original_ref) for rdim in resolved_dimensions @@ -282,6 +307,14 @@ def find_matching_preagg( for ref in (preagg.grain_columns or []) } + if not reaggregate_grain.issubset(preagg_grain_set): + logger.debug( + f"[BuildV3] Pre-agg {preagg.id} grain {preagg_grain_set} " + f"doesn't retain semi-additive protected grain " + f"{reaggregate_grain}", + ) + continue + # Check grain compatibility. For additive measures the pre-agg may be at # the same or a finer grain (roll-up allowed → subset match). For # non-additive measures the pre-agg grain must exactly match the request. diff --git a/datajunction-server/datajunction_server/construction/build_v3/types.py b/datajunction-server/datajunction_server/construction/build_v3/types.py index faac2e34c..5be45c849 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/types.py +++ b/datajunction-server/datajunction_server/construction/build_v3/types.py @@ -316,6 +316,10 @@ class GrainGroupSQL: # Used by metrics SQL to correctly reference component columns component_aliases: dict[str, str] = field(default_factory=dict) + # Active semi-additive components that need collapse in the metrics layer. + # Maps component.name -> protected dimension column alias emitted by this CTE. + reaggregate_dimension_aliases: dict[str, str] = field(default_factory=dict) + # Merge tracking: when True, aggregations happen in final SELECT, not in CTE is_merged: bool = False @@ -617,6 +621,8 @@ class BaseMetricsResult: all_metrics: set[str] # All metric names in grain groups metric_exprs: dict[str, MetricExprInfo] # metric_name -> expression info component_refs: dict[str, ColumnRef] # component_name -> column reference + # Semi-additive base metrics collapsed before being projected by base_metrics. + precollapsed_reaggregate_metrics: set[str] = field(default_factory=set) @dataclass @@ -743,6 +749,10 @@ class GrainGroup: # between decompose.py component identifiers and measures.py SQL aliases. grain_col_aliases: dict[str, str] = field(default_factory=dict) + # Active semi-additive components that require an internal protected dimension + # in this grain group. Maps component.name -> protected dimension ref. + reaggregate_component_dimensions: dict[str, str] = field(default_factory=dict) + # Non-decomposable metrics that couldn't be broken into components # These need their raw metric expression applied in the final SELECT non_decomposable_metrics: list[DecomposedMetricInfo] = field(default_factory=list) diff --git a/datajunction-server/datajunction_server/database/node.py b/datajunction-server/datajunction_server/database/node.py index 1109ef9ba..4f3b13b67 100644 --- a/datajunction-server/datajunction_server/database/node.py +++ b/datajunction-server/datajunction_server/database/node.py @@ -87,6 +87,7 @@ from datajunction_server.models.custom_metadata import CustomMetadataFilter from datajunction_server.models.node_type import NodeType from datajunction_server.models.partition import PartitionType +from datajunction_server.models.reaggregate import parse_reaggregate_spec from datajunction_server.models.unit import ( AtomicUnit, CompoundUnit, @@ -661,6 +662,7 @@ async def to_spec(self, session: AsyncSession) -> NodeSpec: ) extra_kwargs.update( required_dimensions=required_dimensions_spec, + reaggregate=parse_reaggregate_spec(self.current.reaggregate), direction=self.current.metric_metadata.direction if self.current.metric_metadata else None, @@ -1704,6 +1706,14 @@ class NodeRevision( uselist=False, ) + # Declares how a metric should roll up across dimensions. + # Stored as JSON here and validated at the API/deployment boundaries. + reaggregate: Mapped[dict[str, Any] | None] = mapped_column( + JSON, + nullable=True, + default=None, + ) + # Filters that are always applied when generating SQL for a cube node cube_filters: Mapped[list[str] | None] = mapped_column( JSON, @@ -2059,6 +2069,12 @@ def extra_validation(self) -> None: "bound dimensions which are only for metrics.", ) + if self.type != NodeType.METRIC and self.reaggregate: + raise DJInvalidInputException( + f"Node {self.name} of type {self.type} cannot have " + "reaggregate settings which are only for metrics.", + ) + if self.type == NodeType.METRIC: self.check_metric() diff --git a/datajunction-server/datajunction_server/internal/client.py b/datajunction-server/datajunction_server/internal/client.py index 7964b4cd2..37c9d24ed 100644 --- a/datajunction-server/datajunction_server/internal/client.py +++ b/datajunction-server/datajunction_server/internal/client.py @@ -192,6 +192,11 @@ async def python_client_create_node( "required_dimensions": [ # type: ignore col.name for col in node.current.required_dimensions ], + **( + {"reaggregate": node.current.reaggregate} # type: ignore + if node.current.reaggregate + else {} + ), **( { "direction": ( # type: ignore diff --git a/datajunction-server/datajunction_server/internal/deployment/orchestrator.py b/datajunction-server/datajunction_server/internal/deployment/orchestrator.py index bb6ded24d..8d9678f54 100644 --- a/datajunction-server/datajunction_server/internal/deployment/orchestrator.py +++ b/datajunction-server/datajunction_server/internal/deployment/orchestrator.py @@ -154,6 +154,7 @@ NodeStatus, NodeType, ) +from datajunction_server.models.reaggregate import dump_reaggregate_spec from datajunction_server.models.unit import ( AtomicUnit, CompoundUnit, @@ -6008,6 +6009,9 @@ async def _create_node_revision( dependency_nodes, ) new_revision.required_dimensions = matched_columns + new_revision.reaggregate = dump_reaggregate_spec( + metric_spec.rendered_reaggregate, + ) return new_revision def _resolve_metric_unit( diff --git a/datajunction-server/datajunction_server/internal/deployment/validation.py b/datajunction-server/datajunction_server/internal/deployment/validation.py index daaf3156f..3e49480c4 100644 --- a/datajunction-server/datajunction_server/internal/deployment/validation.py +++ b/datajunction-server/datajunction_server/internal/deployment/validation.py @@ -19,7 +19,11 @@ ErrorCode, ) from datajunction_server.internal.deployment.type_inference import validate_node_query -from datajunction_server.internal.validation import validate_metric_query +from datajunction_server.internal.validation import ( + derived_metric_reaggregate_error, + invalid_reaggregate_dimension_references, + validate_metric_query, +) from datajunction_server.models.deployment import ( ColumnSpec, DimensionJoinLinkSpec, @@ -33,6 +37,9 @@ missing_join_on_message, ) from datajunction_server.models.node import NodeStatus, NodeType +from datajunction_server.models.reaggregate import ( + unsupported_dimension_reaggregate_functions, +) from datajunction_server.sql.dag import get_dimensions from datajunction_server.sql.parsing.ast import fast_parse_mode from datajunction_server.sql.parsing.backends.antlr4 import ast, parse_rule @@ -157,7 +164,7 @@ async def validate(self, node_specs: list[NodeSpec]) -> list[NodeValidationResul specs_needing_parse.append(spec) spec_indices.append(i) - # Pre-fetch dimension nodes for required_dimensions validation + # Pre-fetch dimension nodes for metric dimension-reference validation await self._prefetch_required_dimension_nodes(specs_needing_parse) # Pre-fetch dimension sets for cross-fact derived metric validation @@ -523,6 +530,10 @@ def validate_query_node( if req_dim_error is not None: errors.append(req_dim_error) + reaggregate_error = self._check_reaggregate_dimension(spec) + if reaggregate_error is not None: + errors.append(reaggregate_error) + cross_fact_error = self._check_cross_fact_dimensions(spec) if cross_fact_error is not None: errors.append(cross_fact_error) # pragma: no cover @@ -688,9 +699,9 @@ async def _prefetch_required_dimension_nodes( specs: list[NodeSpec], ) -> None: """ - Collect all dimension node names referenced via required_dimensions across the - batch, fetch any not already in dependency_nodes in a single DB query, and - store the combined map in self._all_dim_nodes. + Collect all dimension node names referenced via metric-only dimension + declarations across the batch, fetch any not already in dependency_nodes in a + single DB query, and store the combined map in self._all_dim_nodes. """ req_dim_node_names: set[str] = set() for spec in specs: @@ -700,6 +711,12 @@ async def _prefetch_required_dimension_nodes( if SEPARATOR in req_dim: dim_node_name = req_dim.rsplit(SEPARATOR, 1)[0] req_dim_node_names.add(dim_node_name) + reaggregate = getattr(spec, "rendered_reaggregate", None) + if reaggregate: + for rule in reaggregate.rules: + if SEPARATOR in rule.dimension: + dim_node_name = rule.dimension.rsplit(SEPARATOR, 1)[0] + req_dim_node_names.add(dim_node_name) self._all_dim_nodes = dict(self.context.dependency_nodes) @@ -813,6 +830,68 @@ def _check_required_dimensions(self, spec: NodeSpec) -> DJError | None: debug={"invalid_required_dimensions": list(invalid)}, ) + def _check_reaggregate_dimension(self, spec: NodeSpec) -> DJError | None: + """ + Validate that reaggregate rule dimensions resolve to real columns. + """ + reaggregate = getattr(spec, "rendered_reaggregate", None) + if not reaggregate or not reaggregate.rules: + return None + + metric_type_error = derived_metric_reaggregate_error( + spec.rendered_name, + spec.node_type == NodeType.METRIC + and spec.query_ast is not None + and spec.query_ast.select.from_ is None, + reaggregate, + ) + if metric_type_error: + return metric_type_error + + invalid_functions = unsupported_dimension_reaggregate_functions(reaggregate) + if invalid_functions: + return DJError( + code=ErrorCode.INVALID_ARGUMENTS_TO_FUNCTION, + message=( + "Node definition contains unsupported dimension " + "reaggregate function(s)." + ), + debug={ + "invalid_reaggregate_functions": invalid_functions, + }, + ) + + dep_names = self.context.node_graph.get(spec.rendered_name, []) + parent_columns = [ + col + for dep_name in dep_names + for dep_node in [self.context.dependency_nodes.get(dep_name)] + if dep_node and dep_node.current + for col in dep_node.current.columns + ] + + reaggregate_dimensions = [rule.dimension for rule in reaggregate.rules] + invalid, _ = _resolve_required_dimensions( + reaggregate_dimensions, + parent_columns, + self._all_dim_nodes, + ) + invalid.update( + invalid_reaggregate_dimension_references(reaggregate_dimensions), + ) + + if not invalid: + return None + + return DJError( + code=ErrorCode.INVALID_COLUMN, + message=( + "Node definition contains references to columns as " + "reaggregate dimensions that are not on parent nodes." + ), + debug={"invalid_reaggregate_dimensions": list(invalid)}, + ) + async def _prefetch_metric_dimensions( self, specs: list[NodeSpec], diff --git a/datajunction-server/datajunction_server/internal/namespaces.py b/datajunction-server/datajunction_server/internal/namespaces.py index d3eb5f10b..32bae088f 100644 --- a/datajunction-server/datajunction_server/internal/namespaces.py +++ b/datajunction-server/datajunction_server/internal/namespaces.py @@ -1229,6 +1229,7 @@ def _metric_project_config(node: Node, namespace_requested: str) -> dict: "query": node.current.query, "tags": [tag.name for tag in node.tags], "required_dimensions": [dim.name for dim in node.current.required_dimensions], + "reaggregate": node.current.reaggregate, "direction": ( node.current.metric_metadata.direction.name.lower() if node.current.metric_metadata and node.current.metric_metadata.direction @@ -1642,6 +1643,14 @@ async def get_node_specs_for_export( ) for required_dim in metric_spec.required_dimensions ] + if metric_spec.reaggregate: + for rule in metric_spec.reaggregate.rules: + rule.dimension = _inject_prefix_for_cube_ref( + rule.dimension, + namespace, + parent_namespace, + namespace_suffixes, + ) if node_spec.node_type in ( NodeType.SOURCE, NodeType.TRANSFORM, diff --git a/datajunction-server/datajunction_server/internal/nodes.py b/datajunction-server/datajunction_server/internal/nodes.py index f65ac0610..1d5aa1f86 100644 --- a/datajunction-server/datajunction_server/internal/nodes.py +++ b/datajunction-server/datajunction_server/internal/nodes.py @@ -103,6 +103,10 @@ fold_change_tiers, version_change_tier, ) +from datajunction_server.models.decompose import ( + AggregationRule as DecomposeAggregationRule, + MetricComponent, +) from datajunction_server.models.dimensionlink import ( JoinLinkInput, JoinType, @@ -128,6 +132,7 @@ ) from datajunction_server.models.node_type import NodeType from datajunction_server.models.query import QueryCreate +from datajunction_server.models.reaggregate import dump_reaggregate_spec from datajunction_server.models.table_metadata import TableMetadata, TableOwner from datajunction_server.service_clients import QueryServiceClient from datajunction_server.sql.dag import ( @@ -590,6 +595,11 @@ async def create_node_revision( query=data.query, mode=data.mode, required_dimensions=data.required_dimensions or [], + reaggregate=( + dump_reaggregate_spec(data.reaggregate) + if node_type == NodeType.METRIC + else None + ), created_by_id=current_user.id, custom_metadata=data.custom_metadata, ) @@ -896,15 +906,10 @@ async def _derive_frozen_measures_impl( session=session, name=measure.name, ) + if frozen_measure: + _raise_if_frozen_measure_conflicts(frozen_measure, measure) if not frozen_measure and measure.aggregation: - frozen_measure = FrozenMeasure( - name=measure.name, - upstream_revision_id=upstream_revision_id, - expression=measure.expression, - aggregation=measure.aggregation, - rule=measure.rule, - used_by_node_revisions=[], - ) + frozen_measure = _new_frozen_measure(measure, upstream_revision_id) session.add(frozen_measure) if frozen_measure: frozen_measure.used_by_node_revisions.append(node_revision) @@ -1024,19 +1029,64 @@ async def derive_frozen_measures_bulk( if frozen_measure is None: if not measure.aggregation: continue - frozen_measure = FrozenMeasure( - name=measure.name, - upstream_revision_id=upstream_revision_id, - expression=measure.expression, - aggregation=measure.aggregation, - rule=measure.rule, - used_by_node_revisions=[], - ) + frozen_measure = _new_frozen_measure(measure, upstream_revision_id) session.add(frozen_measure) fm_by_name[measure.name] = frozen_measure + else: + _raise_if_frozen_measure_conflicts(frozen_measure, measure) frozen_measure.used_by_node_revisions.append(rev) +def _aggregation_rule_identity(rule: DecomposeAggregationRule) -> dict[str, Any]: + """Return the stable JSON shape used for frozen-measure rule comparison.""" + return rule.model_dump(mode="json", exclude_none=True, exclude={"reaggregate"}) + + +def _frozen_measure_rule(rule: DecomposeAggregationRule) -> DecomposeAggregationRule: + """ + Return the metric-independent rule persisted on a shared frozen measure. + """ + return rule.model_copy(update={"reaggregate": None}) + + +def _new_frozen_measure( + measure: MetricComponent, + upstream_revision_id: int, +) -> FrozenMeasure: + """Construct a shared frozen measure without metric-level policy.""" + if not measure.aggregation: # pragma: no cover + raise ValueError("Frozen measures require an aggregation") + return FrozenMeasure( + name=measure.name, + upstream_revision_id=upstream_revision_id, + expression=measure.expression, + aggregation=measure.aggregation, + rule=_frozen_measure_rule(measure.rule), + used_by_node_revisions=[], + ) + + +def _raise_if_frozen_measure_conflicts( + frozen_measure: FrozenMeasure, + measure: MetricComponent, +) -> None: + """ + Prevent component-name collisions from reusing a different frozen measure. + """ + if ( + frozen_measure.expression == measure.expression + and frozen_measure.aggregation == measure.aggregation + and _aggregation_rule_identity(frozen_measure.rule) + == _aggregation_rule_identity(measure.rule) + ): + return + + raise DJInvalidInputException( + f"Frozen measure `{measure.name}` already exists with a different " + "expression, aggregation, or aggregation rule.", + ) + + async def save_node( session: AsyncSession, node_revision: NodeRevision, @@ -1138,6 +1188,7 @@ async def copy_to_new_node( table=old_revision.table, required_dimensions=list(old_revision.required_dimensions), metric_metadata=old_revision.metric_metadata, + reaggregate=old_revision.reaggregate, cube_elements=list(old_revision.cube_elements), cube_filters=old_revision.cube_filters, status=old_revision.status, @@ -2515,6 +2566,7 @@ def copy_existing_node_revision(old_revision: NodeRevision, current_user: User): status=old_revision.status, required_dimensions=list(old_revision.required_dimensions), metric_metadata=old_revision.metric_metadata, + reaggregate=old_revision.reaggregate, dimension_links=[ DimensionLink( dimension_id=link.dimension_id, @@ -2737,8 +2789,20 @@ async def create_new_revision_from_existing( and {col.name for col in old_revision.required_dimensions} != set(data.required_dimensions) ) + reaggregate_was_set = bool( + data and "reaggregate" in data.model_fields_set, + ) + reaggregate_changes = ( + reaggregate_was_set + and old_revision.reaggregate + != dump_reaggregate_spec(data.reaggregate if data else None) + ) major_changes = ( - query_changes or column_changes or pk_changes or required_dim_changes + query_changes + or column_changes + or pk_changes + or required_dim_changes + or reaggregate_changes ) # If nothing has changed, do not create the new node revision @@ -2788,6 +2852,13 @@ async def create_new_revision_from_existing( if data and data.metric_metadata else old_revision.metric_metadata ), + reaggregate=( + dump_reaggregate_spec(data.reaggregate) + if data and data.reaggregate is not None + else None + if reaggregate_was_set + else old_revision.reaggregate + ), dimension_links=[ DimensionLink( dimension_id=link.dimension_id, diff --git a/datajunction-server/datajunction_server/internal/sql.py b/datajunction-server/datajunction_server/internal/sql.py index fd20ec0ca..b635b7509 100644 --- a/datajunction-server/datajunction_server/internal/sql.py +++ b/datajunction-server/datajunction_server/internal/sql.py @@ -244,6 +244,7 @@ async def generate_metrics_sql( matched_cube, dimensions, merged_filters, + metrics, ) # Auto-resolve dialect if not explicitly provided. diff --git a/datajunction-server/datajunction_server/internal/validation.py b/datajunction-server/datajunction_server/internal/validation.py index 6eb1af905..632297422 100644 --- a/datajunction-server/datajunction_server/internal/validation.py +++ b/datajunction-server/datajunction_server/internal/validation.py @@ -11,6 +11,7 @@ from datajunction_server.errors import ( DJError, DJException, + DJInvalidInputException, DJInvalidMetricQueryException, ErrorCode, ) @@ -23,12 +24,50 @@ from datajunction_server.models.base import labelize from datajunction_server.models.node import NodeRevisionBase, NodeStatus from datajunction_server.models.node_type import NodeType +from datajunction_server.models.reaggregate import ( + ReaggregateSpec, + parse_reaggregate_spec, + unsupported_dimension_reaggregate_functions, +) from datajunction_server.sql.parsing import ast from datajunction_server.sql.parsing.backends.antlr4 import SqlSyntaxError, parse from datajunction_server.sql.parsing.backends.exceptions import DJParseException from datajunction_server.sql.parsing.types import ListType, MapType, StructType +def invalid_reaggregate_dimension_references(dimensions: list[str]) -> set[str]: + """Return reaggregate dimensions that are not fully qualified.""" + from datajunction_server.construction.build_v3.dimensions import ( + parse_dimension_ref, + ) + + invalid: set[str] = set() + for dimension in dimensions: + try: + parse_dimension_ref(dimension) + except DJInvalidInputException: + invalid.add(dimension) + return invalid + + +def derived_metric_reaggregate_error( + metric_name: str, + is_derived_metric: bool, + reaggregate: ReaggregateSpec | dict | None, +) -> DJError | None: + """Reject metric-level reaggregation declarations on derived metrics.""" + reaggregate_spec = parse_reaggregate_spec(reaggregate) + if not (is_derived_metric and reaggregate_spec and reaggregate_spec.rules): + return None + return DJError( + code=ErrorCode.INVALID_METRIC, + message=( + "Reaggregate declarations are only supported on base metrics. " + f"Derived metric `{metric_name}` declares reaggregate." + ), + ) + + def _reparse_parent_column_types(dependencies_map: dict) -> None: """Re-parse string column types on parent nodes before type inference. @@ -357,11 +396,38 @@ async def validate_node_data( parent_columns, ) node_validator.required_dimensions = matched_bound_columns + reaggregate_spec = parse_reaggregate_spec(validated_node.reaggregate) + invalid_reaggregate_dimensions: set[str] = set() + invalid_reaggregate_functions: list[str] = [] + if reaggregate_spec and reaggregate_spec.rules: + reaggregate_dimensions = [rule.dimension for rule in reaggregate_spec.rules] + ( + invalid_reaggregate_dimensions, + _, + ) = await find_required_dimensions( + session, + reaggregate_dimensions, + parent_columns, + ) + invalid_reaggregate_dimensions.update( + invalid_reaggregate_dimension_references(reaggregate_dimensions), + ) + invalid_reaggregate_functions = unsupported_dimension_reaggregate_functions( + reaggregate_spec, + ) except MissingGreenlet: invalid_required_dimensions = set() + invalid_reaggregate_dimensions = set() + invalid_reaggregate_functions = [] node_validator.required_dimensions = [] - if missing_parents_map or type_inference_failures or invalid_required_dimensions: + if ( + missing_parents_map + or type_inference_failures + or invalid_required_dimensions + or invalid_reaggregate_dimensions + or invalid_reaggregate_functions + ): # update status node_validator.status = NodeStatus.INVALID # build errors @@ -410,10 +476,46 @@ async def validate_node_data( if invalid_required_dimensions else [] ) + invalid_reaggregate_dimensions_error = ( + [ + DJError( + code=ErrorCode.INVALID_COLUMN, + message=( + "Node definition contains references to columns as " + "reaggregate dimensions that are not on parent nodes." + ), + debug={ + "invalid_reaggregate_dimensions": list( + invalid_reaggregate_dimensions, + ), + }, + ), + ] + if invalid_reaggregate_dimensions + else [] + ) + invalid_reaggregate_functions_error = ( + [ + DJError( + code=ErrorCode.INVALID_ARGUMENTS_TO_FUNCTION, + message=( + "Node definition contains unsupported dimension " + "reaggregate function(s)." + ), + debug={ + "invalid_reaggregate_functions": invalid_reaggregate_functions, + }, + ), + ] + if invalid_reaggregate_functions + else [] + ) errors = ( missing_parents_error + type_inference_error + invalid_required_dimensions_error + + invalid_reaggregate_dimensions_error + + invalid_reaggregate_functions_error ) node_validator.errors.extend(errors) @@ -547,6 +649,14 @@ async def validate_node_data_v2( # --- Step 4: classify parents (SHARED with deployment) --- is_derived_metric = is_metric and query_ast.select.from_ is None + reaggregate_metric_error = derived_metric_reaggregate_error( + validated_node.name, + is_derived_metric, + validated_node.reaggregate, + ) + if reaggregate_metric_error: + node_validator.status = NodeStatus.INVALID + node_validator.errors.append(reaggregate_metric_error) parents, missing = classify_parents( is_derived_metric, candidates, @@ -705,10 +815,34 @@ async def validate_node_data_v2( parent_columns, ) node_validator.required_dimensions = matched_bound_columns + reaggregate_spec = parse_reaggregate_spec(validated_node.reaggregate) + invalid_reaggregate_dimensions: set[str] = set() + invalid_reaggregate_functions: list[str] = [] + if reaggregate_spec and reaggregate_spec.rules: + reaggregate_dimensions = [rule.dimension for rule in reaggregate_spec.rules] + ( + invalid_reaggregate_dimensions, + _, + ) = await find_required_dimensions( + session, + reaggregate_dimensions, + parent_columns, + ) + invalid_reaggregate_dimensions.update( + invalid_reaggregate_dimension_references(reaggregate_dimensions), + ) + invalid_reaggregate_functions = unsupported_dimension_reaggregate_functions( + reaggregate_spec, + ) # --- Step 12: final error assembly for missing parents + invalid required # dims (matches legacy code shapes). - if node_validator.missing_parents_map or invalid_required_dimensions: + if ( + node_validator.missing_parents_map + or invalid_required_dimensions + or invalid_reaggregate_dimensions + or invalid_reaggregate_functions + ): node_validator.status = NodeStatus.INVALID if node_validator.missing_parents_map: node_validator.errors.append( @@ -740,6 +874,34 @@ async def validate_node_data_v2( }, ), ) + if invalid_reaggregate_dimensions: + node_validator.errors.append( + DJError( + code=ErrorCode.INVALID_COLUMN, + message=( + "Node definition contains references to columns as " + "reaggregate dimensions that are not on parent nodes." + ), + debug={ + "invalid_reaggregate_dimensions": list( + invalid_reaggregate_dimensions, + ), + }, + ), + ) + if invalid_reaggregate_functions: + node_validator.errors.append( + DJError( + code=ErrorCode.INVALID_ARGUMENTS_TO_FUNCTION, + message=( + "Node definition contains unsupported dimension " + "reaggregate function(s)." + ), + debug={ + "invalid_reaggregate_functions": invalid_reaggregate_functions, + }, + ), + ) return node_validator diff --git a/datajunction-server/datajunction_server/models/decompose.py b/datajunction-server/datajunction_server/models/decompose.py index 5e5655838..a8ce2df4d 100644 --- a/datajunction-server/datajunction_server/models/decompose.py +++ b/datajunction-server/datajunction_server/models/decompose.py @@ -12,6 +12,7 @@ from pydantic import BaseModel from datajunction_server.enum import StrEnum +from datajunction_server.models.reaggregate import DimensionReaggregateRule class Aggregability(StrEnum): @@ -54,6 +55,7 @@ class AggregationRule(BaseModel): type: Aggregability = Aggregability.NONE level: list[str] | None = None + reaggregate: DimensionReaggregateRule | None = None class MetricComponent(BaseModel): diff --git a/datajunction-server/datajunction_server/models/deployment.py b/datajunction-server/datajunction_server/models/deployment.py index ab1cf5a11..7fd5a66b4 100644 --- a/datajunction-server/datajunction_server/models/deployment.py +++ b/datajunction-server/datajunction_server/models/deployment.py @@ -40,6 +40,7 @@ NodeType, ) from datajunction_server.models.partition import Granularity, PartitionType +from datajunction_server.models.reaggregate import ReaggregateSpec from datajunction_server.models.semantic_fingerprint import SemanticFingerprintValue from datajunction_server.models.unit import ( Unit, @@ -1037,6 +1038,7 @@ class MetricSpec(NodeSpec): # Excluded from serialization so it's never exported. columns: list[ColumnSpec] | None = Field(default=None, exclude=True) required_dimensions: list[str] | None = None # Field(default_factory=list) + reaggregate: ReaggregateSpec | None = None direction: MetricDirection | None = None unit_enum: MetricUnit | None = Field(default=None, exclude=True) # Structured unit form at the metric level — peer of `unit_enum`. @@ -1053,6 +1055,7 @@ class MetricSpec(NodeSpec): "columns": ChangeTier.NONE, # Required dimensions constrain which queries the metric can answer. "required_dimensions": ChangeTier.MAJOR, + "reaggregate": ChangeTier.MAJOR, # Everything below is presentation metadata on the metric's single output # column -- the same set the PATCH path already treats as minor via # `metric_metadata` in `create_new_revision_from_existing`. @@ -1190,6 +1193,29 @@ def diff(self, other: "NodeSpec") -> list[str]: changed.remove("required_dimensions") return changed + @property + def rendered_reaggregate(self) -> ReaggregateSpec | None: + """ + Reaggregate spec with `${prefix}` resolved to this spec's namespace. + """ + if not self.reaggregate: + return None + rules = [ + rule.model_copy( + update={ + "dimension": render_prefixes(rule.dimension, self.namespace) + if "${prefix}" in rule.dimension + else rule.dimension, + }, + ) + for rule in self.reaggregate.rules + ] + return self.reaggregate.model_copy( + update={ + "rules": rules, + }, + ) + def model_dump(self, **kwargs): # pragma: no cover base = super().model_dump(**kwargs) base["unit"] = self.unit @@ -1229,6 +1255,7 @@ def __eq__(self, other: object) -> bool: other.canonical_required_dimensions, preserve_order=False, ) + and self.rendered_reaggregate == other.rendered_reaggregate and eq_or_fallback(self.direction, other.direction, MetricDirection.NEUTRAL) and self._normalized_unit() == other._normalized_unit() and self.significant_digits == other.significant_digits diff --git a/datajunction-server/datajunction_server/models/metric.py b/datajunction-server/datajunction_server/models/metric.py index 2ec7e50b6..43813ccec 100644 --- a/datajunction-server/datajunction_server/models/metric.py +++ b/datajunction-server/datajunction_server/models/metric.py @@ -16,6 +16,7 @@ MetricMetadataOutput, ) from datajunction_server.models.query import ColumnMetadata, V3ColumnMetadata +from datajunction_server.models.reaggregate import ReaggregateSpec from datajunction_server.models.sql import ScanEstimate, TranspiledSQL from datajunction_server.models.unit import unit_to_dict from datajunction_server.sql.decompose import MetricComponentExtractor @@ -50,6 +51,7 @@ class Metric(BaseModel): # going forward. `None` when no unit is set, regardless of input shape. unit: dict | None = None required_dimensions: list[str] + reaggregate: ReaggregateSpec | None = None # Whether the metric is a single aggregation call (a "measure") that can map # 1:1 to a column in an externally-built pre-aggregation table. Derived/ratio @@ -104,6 +106,7 @@ async def parse_node( node.current.columns[0].unit if node.current.columns else None, ), required_dimensions=[dim.name for dim in node.current.required_dimensions], + reaggregate=node.current.reaggregate, is_measure=node.current.is_measure, incompatible_druid_functions=incompatible_druid_functions, measures=measures, diff --git a/datajunction-server/datajunction_server/models/node.py b/datajunction-server/datajunction_server/models/node.py index e60ff6b67..13da677a4 100644 --- a/datajunction-server/datajunction_server/models/node.py +++ b/datajunction-server/datajunction_server/models/node.py @@ -33,6 +33,7 @@ from datajunction_server.models.materialization import MaterializationConfigOutput from datajunction_server.models.node_type import NodeNameOutput, NodeType from datajunction_server.models.partition import PartitionOutput +from datajunction_server.models.reaggregate import ReaggregateSpec from datajunction_server.models.tag import TagMinimum, TagOutput from datajunction_server.models.unit import ( AtomicUnit, @@ -878,6 +879,7 @@ class MetricNodeFields(BaseModel): """ required_dimensions: list[str] | None = None + reaggregate: ReaggregateSpec | None = None metric_metadata: MetricMetadataInput | None = None @@ -1033,6 +1035,7 @@ class NodeRevisionOutput(BaseModel): materializations: list[MaterializationConfigOutput] parents: list[NodeNameOutput] metric_metadata: MetricMetadataOutput | None = None + reaggregate: ReaggregateSpec | None = None dimension_links: list[LinkDimensionOutput] | None = None custom_metadata: dict | None = None @@ -1065,6 +1068,7 @@ class NodeOutput(GenericNodeOutputModel): materializations: list[MaterializationConfigOutput] parents: list[NodeNameOutput] metric_metadata: MetricMetadataOutput | None = None + reaggregate: ReaggregateSpec | None = None dimension_links: list[LinkDimensionOutput] = Field(default_factory=list) created_at: UTCDatetime created_by: UserNameOnly | None = None diff --git a/datajunction-server/datajunction_server/models/reaggregate.py b/datajunction-server/datajunction_server/models/reaggregate.py new file mode 100644 index 000000000..1cd755d70 --- /dev/null +++ b/datajunction-server/datajunction_server/models/reaggregate.py @@ -0,0 +1,114 @@ +"""Models for metric reaggregation declarations.""" + +from pydantic import BaseModel, ConfigDict, Field + +from datajunction_server.enum import StrEnum + + +class ReaggregationFunction(StrEnum): + """ + Supported functions for metric reaggregation. + """ + + AUTO = "auto" + NONE = "none" + SUM = "sum" + AVG = "avg" + LAST_VALUE = "last_value" + FIRST_VALUE = "first_value" + MIN = "min" + MAX = "max" + + +DIMENSION_REAGGREGATE_FUNCTIONS = frozenset( + { + ReaggregationFunction.LAST_VALUE, + ReaggregationFunction.FIRST_VALUE, + ReaggregationFunction.MIN, + ReaggregationFunction.MAX, + }, +) + + +def is_supported_dimension_reaggregate_function( + function: ReaggregationFunction, +) -> bool: + """ + Return whether a function can collapse across a protected dimension. + """ + return function in DIMENSION_REAGGREGATE_FUNCTIONS + + +def unsupported_dimension_reaggregate_functions( + spec: "ReaggregateSpec | dict | None", +) -> list[str]: + """ + Unsupported dimension-specific collapse functions in a reaggregate spec. + """ + reaggregate = parse_reaggregate_spec(spec) + if not reaggregate: + return [] + return sorted( + { + rule.fn.value + for rule in reaggregate.rules + if not is_supported_dimension_reaggregate_function(rule.fn) + }, + ) + + +class DimensionReaggregateRule(BaseModel): + """ + Dimension-specific metric reaggregation rule. + """ + + model_config = ConfigDict(extra="forbid") + + dimension: str + fn: ReaggregationFunction + + +class ReaggregateSpec(BaseModel): + """ + Declaration for how a metric rolls up from its accumulation grain. + """ + + model_config = ConfigDict(extra="forbid") + + rules: list[DimensionReaggregateRule] = Field(default_factory=list) + + +def dump_reaggregate_spec( + spec: ReaggregateSpec | dict | None, +) -> dict | None: + """ + Return a JSON-serializable reaggregate spec dictionary. + """ + if spec is None: + return None + if isinstance(spec, ReaggregateSpec): + return spec.model_dump(mode="json") + return ReaggregateSpec.model_validate(spec).model_dump(mode="json") + + +def parse_reaggregate_spec( + spec: ReaggregateSpec | dict | None, +) -> ReaggregateSpec | None: + """ + Return a validated reaggregate spec model. + """ + if spec is None: + return None + if isinstance(spec, ReaggregateSpec): + return spec + return ReaggregateSpec.model_validate(spec) + + +def dimension_reaggregate_rules( + spec: ReaggregateSpec | dict | None, +) -> list[DimensionReaggregateRule]: + """ + Return the dimension-specific reaggregation rules from a spec. + """ + reaggregate = parse_reaggregate_spec(spec) + return reaggregate.rules if reaggregate else [] diff --git a/datajunction-server/datajunction_server/semantic_fingerprints/v1.py b/datajunction-server/datajunction_server/semantic_fingerprints/v1.py index 08cc05e83..d40f99c16 100644 --- a/datajunction-server/datajunction_server/semantic_fingerprints/v1.py +++ b/datajunction-server/datajunction_server/semantic_fingerprints/v1.py @@ -45,7 +45,7 @@ "primary_key", "query", ), - MetricSpec: ("query", "required_dimensions"), + MetricSpec: ("query", "required_dimensions", "reaggregate"), CubeSpec: ("metrics", "dimensions", "filters", "columns"), } diff --git a/datajunction-server/datajunction_server/sql/decompose.py b/datajunction-server/datajunction_server/sql/decompose.py index 05b6d7dea..a52a148fc 100644 --- a/datajunction-server/datajunction_server/sql/decompose.py +++ b/datajunction-server/datajunction_server/sql/decompose.py @@ -10,12 +10,19 @@ from sqlalchemy.orm import aliased from datajunction_server.database.node import Node, NodeRelationship, NodeRevision +from datajunction_server.errors import DJInvalidInputException from datajunction_server.models.decompose import ( Aggregability, AggregationRule, MetricComponent, ) from datajunction_server.models.node_type import NodeType +from datajunction_server.models.reaggregate import ( + ReaggregateSpec, + dimension_reaggregate_rules, + is_supported_dimension_reaggregate_function, + parse_reaggregate_spec, +) from datajunction_server.naming import amenable_col_names from datajunction_server.sql import functions as dj_functions from datajunction_server.sql.parsing.backends.antlr4 import ast, parse @@ -746,6 +753,7 @@ class BaseMetricData: name: str query: str + reaggregate: ReaggregateSpec | dict | None = None @dataclass @@ -937,7 +945,10 @@ async def extract( if not metric_data.is_derived else parse(base_metric.query) ) - base_components, derived_ast = self._extract_base(base_ast) + base_components, derived_ast = self._extract_base( + base_ast, + parse_reaggregate_spec(base_metric.reaggregate), + ) for comp in base_components: if comp.name not in components_tracker: @@ -991,6 +1002,12 @@ def _build_metric_data_from_cache( ] if metric_parents: + if metric_node.current.reaggregate: + raise DJInvalidInputException( + "Semi-additive metrics must be base metrics in V1. " + f"Derived metric `{metric_node.name}` declares reaggregate.", + ) + # Derived metric - base metrics are the parent metrics return MetricData( query=metric_node.current.query, @@ -999,6 +1016,7 @@ def _build_metric_data_from_cache( BaseMetricData( name=parent_name, query=nodes_cache[parent_name].current.query, + reaggregate=nodes_cache[parent_name].current.reaggregate, ) for parent_name in metric_parents if nodes_cache[parent_name].current @@ -1014,6 +1032,7 @@ def _build_metric_data_from_cache( BaseMetricData( name=metric_node.name, query=metric_node.current.query, + reaggregate=metric_node.current.reaggregate, ), ], ) @@ -1030,6 +1049,7 @@ async def _load_metric_data(self, session: AsyncSession) -> MetricData: select( Node.name.label("parent_name"), NodeRevision.query.label("parent_query"), + NodeRevision.reaggregate.label("parent_reaggregate"), ) .select_from(NodeRelationship) .join(Node, NodeRelationship.parent_id == Node.id) @@ -1050,6 +1070,7 @@ async def _load_metric_data(self, session: AsyncSession) -> MetricData: this_metric_stmt = ( select( NodeRevision.query, + NodeRevision.reaggregate, Node.name, ) .join(Node, NodeRevision.node_id == Node.id) @@ -1062,12 +1083,22 @@ async def _load_metric_data(self, session: AsyncSession) -> MetricData: this_row = this_result.one() if parent_rows: + if this_row.reaggregate: + raise DJInvalidInputException( + "Semi-additive metrics must be base metrics in V1. " + f"Derived metric `{this_row.name}` declares reaggregate.", + ) + # Derived metric - base metrics are the parents return MetricData( query=this_row.query, is_derived=True, base_metrics=[ - BaseMetricData(name=row.parent_name, query=row.parent_query) + BaseMetricData( + name=row.parent_name, + query=row.parent_query, + reaggregate=row.parent_reaggregate, + ) for row in parent_rows ], ) @@ -1076,12 +1107,19 @@ async def _load_metric_data(self, session: AsyncSession) -> MetricData: return MetricData( query=this_row.query, is_derived=False, - base_metrics=[BaseMetricData(name=this_row.name, query=this_row.query)], + base_metrics=[ + BaseMetricData( + name=this_row.name, + query=this_row.query, + reaggregate=this_row.reaggregate, + ), + ], ) def _extract_base( self, query_ast: ast.Query, + reaggregate: ReaggregateSpec | None = None, ) -> tuple[list[MetricComponent], ast.Query]: """ Extract components from a base metric by decomposing aggregations. @@ -1113,6 +1151,11 @@ def _extract_base( # If any aggregation is non-decomposable, abort decomposition # entirely — the metric is non-decomposable as a whole. if any(get_decomposition(dj_fn) is None for _, dj_fn in agg_funcs): + if dimension_reaggregate_rules(reaggregate): + self._raise_unsupported_reaggregate_shape( + "dimension-specific reaggregation requires a " + "decomposable aggregation", + ) return [], query_ast for func, dj_function in agg_funcs: @@ -1135,8 +1178,55 @@ def _extract_base( for proj in query_ast.select.projection: wrap_divisions_in_nullif(cast(ast.Expression, proj)) + if reaggregate is not None and dimension_reaggregate_rules(reaggregate): + self._attach_reaggregate_spec(components, reaggregate) + return components, query_ast + def _attach_reaggregate_spec( + self, + components: list[MetricComponent], + reaggregate: ReaggregateSpec, + ) -> None: + """ + Attach a dimension-specific reaggregation rule to one ordinary measure. + """ + dimension_rules = dimension_reaggregate_rules(reaggregate) + if len(dimension_rules) != 1: + self._raise_unsupported_reaggregate_shape( + "dimension-specific reaggregation requires exactly one rule", + ) + + if not is_supported_dimension_reaggregate_function(dimension_rules[0].fn): + self._raise_unsupported_reaggregate_shape( + f"unsupported dimension reaggregation function " + f"`{dimension_rules[0].fn.value}`", + ) + + if len(components) != 1: + self._raise_unsupported_reaggregate_shape( + "dimension-specific reaggregation requires exactly one component", + ) + + component = components[0] + if ( + component.rule.type != Aggregability.FULL + or component.aggregation is None + or component.merge is None + ): + self._raise_unsupported_reaggregate_shape( + "dimension-specific reaggregation requires one non-distinct " + "fully-aggregatable component", + ) + + component.rule.reaggregate = dimension_rules[0] + + @staticmethod + def _raise_unsupported_reaggregate_shape(reason: str) -> None: + raise DJInvalidInputException( + f"Unsupported reaggregate metric shape: {reason}.", + ) + def _substitute_metric_references( self, query_ast: ast.Query, diff --git a/datajunction-server/tests/api/cubes_test.py b/datajunction-server/tests/api/cubes_test.py index 613202c65..9bc13a8c6 100644 --- a/datajunction-server/tests/api/cubes_test.py +++ b/datajunction-server/tests/api/cubes_test.py @@ -12,7 +12,10 @@ from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession -from datajunction_server.api.cubes import _resolve_cube_partition_output_column +from datajunction_server.api.cubes import ( + _resolve_cube_partition_output_column, + _validate_cube_reaggregate_materialization, +) from datajunction_server.construction.build_v3.combiners import ( PreAggSourceInfo, TemporalPartitionInfo, @@ -26,6 +29,7 @@ ErrorCode, ) from datajunction_server.models.cube import CubeElementMetadata +from datajunction_server.models.dialect import Dialect from datajunction_server.models.materialization import ( MaterializationInfo, MaterializationStrategy, @@ -39,6 +43,69 @@ from tests.sql.utils import compare_query_strings +@pytest.mark.asyncio +async def test_cube_reaggregate_validation_skips_decomposition_without_reaggregate( + mocker, +): + """Ordinary cube materializations avoid full metric decomposition.""" + session = mocker.MagicMock(spec=AsyncSession) + cube = mocker.MagicMock() + cube.current.cube_node_metrics = ["default.total_revenue"] + has_reaggregate = mocker.patch( + "datajunction_server.api.cubes._metric_graph_has_reaggregate", + new=mocker.AsyncMock(return_value=False), + ) + setup_build_context = mocker.patch( + "datajunction_server.construction.build_v3.builder.setup_build_context", + new=mocker.AsyncMock(), + ) + + await _validate_cube_reaggregate_materialization(session, cube) + + has_reaggregate.assert_awaited_once_with( + session, + ["default.total_revenue"], + ) + setup_build_context.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_cube_reaggregate_validation_decomposes_reaggregate_graph(mocker): + """Reaggregate metric graphs still receive full safety validation.""" + session = mocker.MagicMock(spec=AsyncSession) + cube = mocker.MagicMock() + cube.current.cube_node_metrics = ["default.balance_index"] + cube.current.cube_node_dimensions = ["default.product.category"] + cube.current.cube_filters = [] + mocker.patch( + "datajunction_server.api.cubes._metric_graph_has_reaggregate", + new=mocker.AsyncMock(return_value=True), + ) + context = mocker.MagicMock() + setup_build_context = mocker.patch( + "datajunction_server.construction.build_v3.builder.setup_build_context", + new=mocker.AsyncMock(return_value=context), + ) + validate = mocker.patch( + "datajunction_server.api.cubes.validate_cube_reaggregate_materialization", + ) + + await _validate_cube_reaggregate_materialization(session, cube) + + setup_build_context.assert_awaited_once_with( + session=session, + metrics=["default.balance_index"], + dimensions=["default.product.category"], + filters=None, + dialect=Dialect.SPARK, + use_materialized=False, + ) + validate.assert_called_once_with( + cube.current, + decomposed_metrics=context.decomposed_metrics, + ) + + def test_cube_partition_output_prefers_exact_unqualified_match(): """An exact bare semantic match wins even when a role appears first.""" role_column = V3ColumnMetadata( @@ -3588,6 +3655,7 @@ async def test_cube_materialization_metadata( "name": "count_c8e42e74", "rule": { "level": None, + "reaggregate": None, "type": "full", }, }, @@ -3599,6 +3667,7 @@ async def test_cube_materialization_metadata( "name": "discount_sum_30b84e6c", "rule": { "level": None, + "reaggregate": None, "type": "full", }, }, @@ -3610,6 +3679,7 @@ async def test_cube_materialization_metadata( "name": "price_count_935e7117", "rule": { "level": None, + "reaggregate": None, "type": "full", }, }, @@ -3621,6 +3691,7 @@ async def test_cube_materialization_metadata( "name": "price_discount_sum_e4ba5456", "rule": { "level": None, + "reaggregate": None, "type": "full", }, }, @@ -3632,6 +3703,7 @@ async def test_cube_materialization_metadata( "name": "price_sum_935e7117", "rule": { "level": None, + "reaggregate": None, "type": "full", }, }, @@ -3643,6 +3715,7 @@ async def test_cube_materialization_metadata( "name": "repair_order_id_count_bd241964", "rule": { "level": None, + "reaggregate": None, "type": "full", }, }, @@ -3654,6 +3727,7 @@ async def test_cube_materialization_metadata( "name": "total_repair_cost_sum_67874507", "rule": { "level": None, + "reaggregate": None, "type": "full", }, }, @@ -3773,6 +3847,7 @@ async def test_cube_materialization_metadata( "name": "price_sum_252381cf", "rule": { "level": None, + "reaggregate": None, "type": "full", }, }, @@ -4025,7 +4100,7 @@ async def test_cube_materialization_metadata( "grain_alias": None, "aggregation": "COUNT", "merge": "SUM", - "rule": {"type": "full", "level": None}, + "rule": {"type": "full", "level": None, "reaggregate": None}, }, { "name": "discount_sum_30b84e6c", @@ -4033,7 +4108,7 @@ async def test_cube_materialization_metadata( "grain_alias": None, "aggregation": "SUM", "merge": "SUM", - "rule": {"type": "full", "level": None}, + "rule": {"type": "full", "level": None, "reaggregate": None}, }, { "name": "price_count_935e7117", @@ -4041,7 +4116,7 @@ async def test_cube_materialization_metadata( "grain_alias": None, "aggregation": "COUNT", "merge": "SUM", - "rule": {"type": "full", "level": None}, + "rule": {"type": "full", "level": None, "reaggregate": None}, }, { "name": "price_discount_sum_e4ba5456", @@ -4049,7 +4124,7 @@ async def test_cube_materialization_metadata( "grain_alias": None, "aggregation": "SUM", "merge": "SUM", - "rule": {"type": "full", "level": None}, + "rule": {"type": "full", "level": None, "reaggregate": None}, }, { "name": "price_sum_935e7117", @@ -4057,7 +4132,7 @@ async def test_cube_materialization_metadata( "grain_alias": None, "aggregation": "SUM", "merge": "SUM", - "rule": {"type": "full", "level": None}, + "rule": {"type": "full", "level": None, "reaggregate": None}, }, { "name": "repair_order_id_count_bd241964", @@ -4065,7 +4140,7 @@ async def test_cube_materialization_metadata( "grain_alias": None, "aggregation": "COUNT", "merge": "SUM", - "rule": {"type": "full", "level": None}, + "rule": {"type": "full", "level": None, "reaggregate": None}, }, { "name": "total_repair_cost_sum_67874507", @@ -4073,7 +4148,7 @@ async def test_cube_materialization_metadata( "grain_alias": None, "aggregation": "SUM", "merge": "SUM", - "rule": {"type": "full", "level": None}, + "rule": {"type": "full", "level": None, "reaggregate": None}, }, { "name": "price_sum_252381cf", @@ -4081,7 +4156,7 @@ async def test_cube_materialization_metadata( "grain_alias": None, "aggregation": "SUM", "merge": "SUM", - "rule": {"type": "full", "level": None}, + "rule": {"type": "full", "level": None, "reaggregate": None}, }, ], "timestamp_column": "hire_date", @@ -4804,6 +4879,133 @@ class TestCubeMaterializeV2SuccessPaths: and mock the query service client to test the full endpoint flow. """ + @pytest.fixture(autouse=True) + def _preserve_app_dependency_overrides( + self, + client_with_repairs_cube: AsyncClient, + ): + """Protect module-scoped client overrides from function-scoped clients.""" + original_overrides = dict(client_with_repairs_cube.app.dependency_overrides) + yield + client_with_repairs_cube.app.dependency_overrides.clear() + client_with_repairs_cube.app.dependency_overrides.update(original_overrides) + + @pytest.mark.asyncio + async def test_materialize_reaggregate_cube_without_protected_dimension_fails( + self, + client_with_build_v3: AsyncClient, + mocker, + ): + """Druid cube materialization refuses to bake in collapsed + semi-additive metrics without the protected dimension.""" + response = await client_with_build_v3.post( + "/nodes/metric/", + json={ + "name": "v3.daily_balance", + "description": "Semi-additive balance measured by order date", + "query": "SELECT SUM(line_total) FROM v3.order_details", + "mode": "published", + "reaggregate": { + "rules": [ + { + "dimension": "v3.date.date_id[order]", + "fn": "last_value", + }, + ], + }, + }, + ) + assert response.status_code in (200, 201), response.json() + + cube_name = "v3.test_daily_balance_materialization_guard" + response = await client_with_build_v3.post( + "/nodes/cube/", + json={ + "name": cube_name, + "metrics": ["v3.daily_balance"], + "dimensions": ["v3.product.category"], + "mode": "published", + "description": "Unsafe semi-additive materialization", + }, + ) + assert response.status_code == 201, response.json() + + combiner = mocker.patch( + "datajunction_server.api.cubes.build_combiner_sql_from_preaggs", + ) + + response = await client_with_build_v3.post( + f"/cubes/{cube_name}/materialize", + json={"strategy": "full", "schedule": "0 0 * * *"}, + ) + + assert response.status_code == 400, response.json() + assert "protected dimension" in response.json()["message"] + combiner.assert_not_called() + + @pytest.mark.asyncio + async def test_materialize_derived_reaggregate_cube_without_protected_dimension_fails( + self, + client_with_build_v3: AsyncClient, + mocker, + ): + """Druid cube materialization also refuses derived metrics whose base + metric has an omitted semi-additive protected dimension.""" + response = await client_with_build_v3.post( + "/nodes/metric/", + json={ + "name": "v3.daily_balance", + "description": "Semi-additive balance measured by order date", + "query": "SELECT SUM(line_total) FROM v3.order_details", + "mode": "published", + "reaggregate": { + "rules": [ + { + "dimension": "v3.date.date_id[order]", + "fn": "last_value", + }, + ], + }, + }, + ) + assert response.status_code in (200, 201), response.json() + response = await client_with_build_v3.post( + "/nodes/metric/", + json={ + "name": "v3.daily_balance_index", + "description": "Derived semi-additive balance index", + "query": "SELECT 10.0 / v3.daily_balance", + "mode": "published", + }, + ) + assert response.status_code in (200, 201), response.json() + + cube_name = "v3.test_daily_balance_index_materialization_guard" + response = await client_with_build_v3.post( + "/nodes/cube/", + json={ + "name": cube_name, + "metrics": ["v3.daily_balance_index"], + "dimensions": ["v3.product.category"], + "mode": "published", + "description": "Unsafe derived semi-additive materialization", + }, + ) + assert response.status_code == 201, response.json() + + combiner = mocker.patch( + "datajunction_server.api.cubes.build_combiner_sql_from_preaggs", + ) + + response = await client_with_build_v3.post( + f"/cubes/{cube_name}/materialize", + json={"strategy": "full", "schedule": "0 0 * * *"}, + ) + + assert response.status_code == 400, response.json() + assert "protected dimension" in response.json()["message"] + combiner.assert_not_called() + @pytest.mark.asyncio async def test_materialize_cube_full_strategy_success( self, @@ -5596,6 +5798,12 @@ async def test_materialize_cube_returns_metric_combiners( mock_qs_client.materialize_cube_v2.return_value = mocker.MagicMock( urls=["http://workflow/test-cube"], ) + had_query_service_override = get_query_service_client in ( + client.app.dependency_overrides + ) + original_query_service_override = client.app.dependency_overrides.get( + get_query_service_client, + ) client.app.dependency_overrides[get_query_service_client] = lambda: ( mock_qs_client ) @@ -5783,8 +5991,12 @@ async def test_materialize_cube_returns_metric_combiners( ] finally: # Clean up dependency override - if get_query_service_client in client.app.dependency_overrides: - del client.app.dependency_overrides[get_query_service_client] + if had_query_service_override: + client.app.dependency_overrides[get_query_service_client] = ( + original_query_service_override + ) + else: + client.app.dependency_overrides.pop(get_query_service_client, None) class TestCubeDeactivateSuccessPaths: @@ -6020,6 +6232,12 @@ async def test_deactivate_cube_uses_stored_workflow_names( urls=["http://workflow/cube-workflow"], workflow_names=["cube_wf_name_1"], ) + had_query_service_override = get_query_service_client in ( + client.app.dependency_overrides + ) + original_query_service_override = client.app.dependency_overrides.get( + get_query_service_client, + ) client.app.dependency_overrides[get_query_service_client] = lambda: ( mock_qs_client ) @@ -6047,8 +6265,12 @@ async def test_deactivate_cube_uses_stored_workflow_names( # The old path should NOT have been called mock_qs_client.deactivate_cube_workflow.assert_not_called() finally: - if get_query_service_client in client.app.dependency_overrides: - del client.app.dependency_overrides[get_query_service_client] + if had_query_service_override: + client.app.dependency_overrides[get_query_service_client] = ( + original_query_service_override + ) + else: + client.app.dependency_overrides.pop(get_query_service_client, None) class TestCubeBackfillSuccessPaths: diff --git a/datajunction-server/tests/api/graphql/resolvers/test_node_resolver.py b/datajunction-server/tests/api/graphql/resolvers/test_node_resolver.py index 6a47982db..6d22da2e9 100644 --- a/datajunction-server/tests/api/graphql/resolvers/test_node_resolver.py +++ b/datajunction-server/tests/api/graphql/resolvers/test_node_resolver.py @@ -70,6 +70,84 @@ def test_columns_resolver_filters_by_attribute(): assert result[0].name == "user_id" +def test_reaggregate_resolver_returns_none_for_non_metric(): + """ + Test that the reaggregate resolver only resolves for metric revisions. + """ + from datajunction_server.api.graphql.scalars.node import NodeRevision + from datajunction_server.database import NodeRevision as DBNodeRevision + from datajunction_server.models.node_type import NodeType + + db_node_revision = DBNodeRevision( + name="source_rev", + type=NodeType.SOURCE, + version="1", + reaggregate={ + "rules": [ + { + "dimension": "default.date_dim.date", + "fn": "last_value", + }, + ], + }, + ) + + result = NodeRevision.reaggregate(NodeRevision, root=db_node_revision) + + assert result is None + + +def test_reaggregate_resolver_returns_none_for_empty_metric_spec(): + """ + Test that the reaggregate resolver returns None for metrics without a spec. + """ + from datajunction_server.api.graphql.scalars.node import NodeRevision + from datajunction_server.database import NodeRevision as DBNodeRevision + from datajunction_server.models.node_type import NodeType + + db_node_revision = DBNodeRevision( + name="metric_rev", + type=NodeType.METRIC, + version="1", + reaggregate=None, + ) + + result = NodeRevision.reaggregate(NodeRevision, root=db_node_revision) + + assert result is None + + +def test_reaggregate_resolver_returns_metric_spec(): + """ + Test that the reaggregate resolver exposes metric reaggregation declarations. + """ + from datajunction_server.api.graphql.scalars.node import NodeRevision + from datajunction_server.database import NodeRevision as DBNodeRevision + from datajunction_server.models.node_type import NodeType + from datajunction_server.models.reaggregate import ReaggregationFunction + + db_node_revision = DBNodeRevision( + name="metric_rev", + type=NodeType.METRIC, + version="1", + reaggregate={ + "rules": [ + { + "dimension": "default.date_dim.date", + "fn": "last_value", + }, + ], + }, + ) + + result = NodeRevision.reaggregate(NodeRevision, root=db_node_revision) + + assert result is not None + assert len(result.rules) == 1 + assert result.rules[0].dimension == "default.date_dim.date" + assert result.rules[0].fn == ReaggregationFunction.LAST_VALUE + + @pytest.mark.asyncio @patch("datajunction_server.api.graphql.resolvers.nodes.DBNode.get_by_name") @patch("datajunction_server.api.graphql.resolvers.nodes.load_node_options") diff --git a/datajunction-server/tests/api/measures_test.py b/datajunction-server/tests/api/measures_test.py index 15cb18343..4d3c0d943 100644 --- a/datajunction-server/tests/api/measures_test.py +++ b/datajunction-server/tests/api/measures_test.py @@ -292,6 +292,7 @@ async def test_list_frozen_measures( "name": "repair_order_id_count_bd241964", "rule": { "level": None, + "reaggregate": None, "type": "full", }, "upstream_revision": { diff --git a/datajunction-server/tests/api/metrics_test.py b/datajunction-server/tests/api/metrics_test.py index b6d939b76..f96c50eae 100644 --- a/datajunction-server/tests/api/metrics_test.py +++ b/datajunction-server/tests/api/metrics_test.py @@ -3,6 +3,7 @@ """ from unittest.mock import patch +from uuid import uuid4 import pytest import pytest_asyncio @@ -468,6 +469,7 @@ async def test_read_metrics(module__client_with_roads: AsyncClient) -> None: "merge": "SUM", "rule": { "level": None, + "reaggregate": None, "type": "full", }, }, @@ -479,6 +481,7 @@ async def test_read_metrics(module__client_with_roads: AsyncClient) -> None: "name": "count_c8e42e74", "rule": { "level": None, + "reaggregate": None, "type": "full", }, }, @@ -494,6 +497,102 @@ async def test_read_metrics(module__client_with_roads: AsyncClient) -> None: assert data["custom_metadata"] is None +@pytest.mark.asyncio +async def test_metric_reaggregate_roundtrip_and_validation( + client_with_roads: AsyncClient, +) -> None: + """ + Semi-additive declarations round-trip and validate their protected dimension. + """ + metric_name = f"default.reaggregate_repairs_{uuid4().hex}" + response = await client_with_roads.post( + "/nodes/metric/", + json={ + "name": metric_name, + "description": "Repair orders with semi-additive declaration", + "query": "SELECT COUNT(repair_order_id) FROM default.repair_orders_fact", + "mode": "published", + "reaggregate": { + "rules": [ + { + "dimension": "default.repair_orders_fact.repair_order_id", + "fn": "last_value", + }, + ], + }, + }, + ) + assert response.status_code in (200, 201), response.json() + assert response.json()["reaggregate"] == { + "rules": [ + { + "dimension": "default.repair_orders_fact.repair_order_id", + "fn": "last_value", + }, + ], + } + + response = await client_with_roads.get(f"/nodes/{metric_name}/") + assert response.status_code == 200 + assert response.json()["reaggregate"] == { + "rules": [ + { + "dimension": "default.repair_orders_fact.repair_order_id", + "fn": "last_value", + }, + ], + } + + response = await client_with_roads.get(f"/metrics/{metric_name}/") + assert response.status_code == 200 + assert response.json()["reaggregate"] == { + "rules": [ + { + "dimension": "default.repair_orders_fact.repair_order_id", + "fn": "last_value", + }, + ], + } + + invalid_metric_name = f"default.invalid_reaggregate_{uuid4().hex}" + response = await client_with_roads.post( + "/nodes/metric/", + json={ + "name": invalid_metric_name, + "description": "Invalid semi-additive declaration", + "query": "SELECT COUNT(repair_order_id) FROM default.repair_orders_fact", + "mode": "published", + "reaggregate": { + "rules": [ + { + "dimension": "default.repair_orders_fact.nope", + "fn": "last_value", + }, + ], + }, + }, + ) + assert response.status_code == 422 + assert response.json() == { + "message": "Node definition contains references to columns as " + "reaggregate dimensions that are not on parent nodes.", + "errors": [ + { + "code": "INVALID_COLUMN", + "message": "Node definition contains references to columns as " + "reaggregate dimensions that are not on parent nodes.", + "debug": { + "invalid_reaggregate_dimensions": [ + "default.repair_orders_fact.nope", + ], + }, + "context": "", + }, + ], + "warnings": [], + } + + @pytest_asyncio.fixture(scope="module") async def module__current_user(module__session: AsyncSession) -> User: """ diff --git a/datajunction-server/tests/api/preaggregations_test.py b/datajunction-server/tests/api/preaggregations_test.py index 72a8b9804..fdb2f54f5 100644 --- a/datajunction-server/tests/api/preaggregations_test.py +++ b/datajunction-server/tests/api/preaggregations_test.py @@ -741,6 +741,7 @@ async def test_get_preagg_by_id(self, client_with_preaggs): "source_column": None, "rule": { "level": None, + "reaggregate": None, "type": "full", }, "used_by_metrics": [ @@ -808,6 +809,7 @@ async def test_get_preagg_by_id(self, client_with_preaggs): "source_column": None, "rule": { "level": None, + "reaggregate": None, "type": "full", }, "used_by_metrics": [ diff --git a/datajunction-server/tests/api/sql_v2_test.py b/datajunction-server/tests/api/sql_v2_test.py index 42c58eb52..3ecd5ccdb 100644 --- a/datajunction-server/tests/api/sql_v2_test.py +++ b/datajunction-server/tests/api/sql_v2_test.py @@ -1203,7 +1203,11 @@ async def create_metric_distinct_single_column(client: AsyncClient): "grain_alias": "hard_hat_id", "merge": None, "name": "hard_hat_id", - "rule": {"level": ["hard_hat_id"], "type": "limited"}, + "rule": { + "level": ["hard_hat_id"], + "reaggregate": None, + "type": "limited", + }, }, ] assert metric_data["derived_expression"] == "COUNT( DISTINCT hard_hat_id)" @@ -1230,7 +1234,11 @@ async def create_metric_distinct_expression(client: AsyncClient): "grain_alias": "hard_hat_id_distinct_0291ee39", "merge": None, "name": "hard_hat_id_distinct_0291ee39", - "rule": {"level": ["IF(hard_hat_id = 1, 1, 0)"], "type": "limited"}, + "rule": { + "level": ["IF(hard_hat_id = 1, 1, 0)"], + "reaggregate": None, + "type": "limited", + }, }, ] assert ( @@ -1627,6 +1635,7 @@ async def test_metric_definitions_with_nonjoinable_dimensions( "name": "default_DOT_local_hard_hats_2_DOT_hard_hat_id_sum_bf8a8419", "rule": { "level": None, + "reaggregate": None, "type": "full", }, }, @@ -1691,6 +1700,7 @@ async def test_metric_definitions_with_single_joinable_dimensions( "name": "contact_name", "rule": { "level": ["default.municipality_dim.contact_name"], + "reaggregate": None, "type": "limited", }, }, @@ -1799,6 +1809,7 @@ async def test_metric_definition_with_multiple_joinable_dimensions( "IF(default.hard_hat.state = 'NY', " "default.hard_hat.first_name, NULL)", ], + "reaggregate": None, "type": "limited", }, }, diff --git a/datajunction-server/tests/construction/build_v3/cube_matcher_test.py b/datajunction-server/tests/construction/build_v3/cube_matcher_test.py index 2a7b5e5a5..d34782f24 100644 --- a/datajunction-server/tests/construction/build_v3/cube_matcher_test.py +++ b/datajunction-server/tests/construction/build_v3/cube_matcher_test.py @@ -30,12 +30,49 @@ from datajunction_server.construction.build_v3.utils import ( extract_filter_dimension_refs, ) +from datajunction_server.database.node import Node from datajunction_server.errors import DJInvalidInputException from datajunction_server.models.decompose import Aggregability from datajunction_server.models.dialect import Dialect from tests.construction.build_v3 import assert_sql_equal +async def _create_daily_balance_metric(client_with_build_v3): + """Create a semi-additive metric protected by order date.""" + response = await client_with_build_v3.post( + "/nodes/metric/", + json={ + "name": "v3.daily_balance", + "description": "Semi-additive balance measured by order date", + "query": "SELECT SUM(line_total) FROM v3.order_details", + "mode": "published", + "reaggregate": { + "rules": [ + { + "dimension": "v3.date.date_id[order]", + "fn": "last_value", + }, + ], + }, + }, + ) + assert response.status_code in (200, 201), response.json() + + +async def _create_daily_balance_index_metric(client_with_build_v3): + """Create a derived metric that depends on the semi-additive balance.""" + response = await client_with_build_v3.post( + "/nodes/metric/", + json={ + "name": "v3.daily_balance_index", + "description": "Derived semi-additive balance index", + "query": "SELECT 10.0 / v3.daily_balance", + "mode": "published", + }, + ) + assert response.status_code in (200, 201), response.json() + + def test_materialized_dimension_lookup_preserves_roles(): """Two roles with one short name resolve to distinct physical columns.""" cube = SimpleNamespace( @@ -91,6 +128,34 @@ async def test_resolve_dialect_requires_metric_or_dimension(): ) +@pytest.mark.asyncio +async def test_pinned_cube_validation_uses_loaded_metric_revisions(monkeypatch): + """Pinned validation must not lazy-load each metric revision's node.""" + + async def no_reaggregate_requirements(*_args, **_kwargs): + return [] + + monkeypatch.setattr( + "datajunction_server.construction.build_v3.cube_matcher." + "_reaggregate_requirements_for_metrics_if_needed", + no_reaggregate_requirements, + ) + cube = SimpleNamespace( + name="v3.test_cube", + cube_dimensions=lambda: [], + metric_node_revisions=lambda: [ + SimpleNamespace(name="v3.total_revenue", reaggregate=None), + ], + ) + + await validate_pinned_cube_covers_filters( + session=None, # type: ignore[arg-type] + cube=cube, + dimensions=[], + filters=[], + ) + + class TestExtractFilterDimensionRefs: """Unit tests for the shared filter-dimension extraction helper. @@ -1168,6 +1233,197 @@ def _boom(_filters): ) assert "could not be parsed" in str(exc.value) + @pytest.mark.asyncio + async def test_reaggregate_cube_missing_protected_dimension_not_matched( + self, + client_with_build_v3, + session, + ): + """A materialized cube that dropped a semi-additive protected dimension + is not safe to auto-route to.""" + await _create_daily_balance_metric(client_with_build_v3) + response = await client_with_build_v3.post( + "/nodes/cube/", + json={ + "name": "v3.test_daily_balance_category_cube", + "metrics": ["v3.daily_balance"], + "dimensions": ["v3.product.category"], + "mode": "published", + "description": "Category cube missing semi-additive protected grain", + }, + ) + assert response.status_code == 201, response.json() + + response = await client_with_build_v3.post( + "/data/v3.test_daily_balance_category_cube/availability/", + json={ + "catalog": "default", + "schema_": "analytics", + "table": "daily_balance_category_cube", + "valid_through_ts": int(time.time() * 1000), + }, + ) + assert response.status_code == 200, response.json() + + result = await find_matching_cube( + session, + metrics=["v3.daily_balance"], + dimensions=["v3.product.category"], + ) + assert result is None + + @pytest.mark.asyncio + async def test_reaggregate_unsafe_cube_falls_back_to_live_sql( + self, + client_with_build_v3, + session, + ): + """An unsafe materialized cube is ignored, leaving the live path usable.""" + await _create_daily_balance_metric(client_with_build_v3) + response = await client_with_build_v3.post( + "/nodes/cube/", + json={ + "name": "v3.test_daily_balance_live_fallback_cube", + "metrics": ["v3.daily_balance"], + "dimensions": ["v3.product.category"], + "mode": "published", + "description": "Unsafe cube for live fallback", + }, + ) + assert response.status_code == 201, response.json() + + response = await client_with_build_v3.post( + "/data/v3.test_daily_balance_live_fallback_cube/availability/", + json={ + "catalog": "default", + "schema_": "analytics", + "table": "daily_balance_live_fallback_cube", + "valid_through_ts": int(time.time() * 1000), + }, + ) + assert response.status_code == 200, response.json() + + result = await build_metrics_sql( + session=session, + metrics=["v3.daily_balance"], + dimensions=["v3.product.category"], + ) + + assert result.cube_name is None + assert "daily_balance_live_fallback_cube" not in result.sql + assert "MAX_BY(" in result.sql + assert "date_id_order" in result.sql + + @pytest.mark.asyncio + async def test_derived_reaggregate_cube_missing_protected_dimension_not_matched( + self, + client_with_build_v3, + session, + ): + """Derived metric cubes inherit base semi-additive materialization needs.""" + await _create_daily_balance_metric(client_with_build_v3) + await _create_daily_balance_index_metric(client_with_build_v3) + response = await client_with_build_v3.post( + "/nodes/cube/", + json={ + "name": "v3.test_daily_balance_index_category_cube", + "metrics": ["v3.daily_balance_index"], + "dimensions": ["v3.product.category"], + "mode": "published", + "description": "Derived cube missing semi-additive protected grain", + }, + ) + assert response.status_code == 201, response.json() + + response = await client_with_build_v3.post( + "/data/v3.test_daily_balance_index_category_cube/availability/", + json={ + "catalog": "default", + "schema_": "analytics", + "table": "daily_balance_index_category_cube", + "valid_through_ts": int(time.time() * 1000), + }, + ) + assert response.status_code == 200, response.json() + + result = await find_matching_cube( + session, + metrics=["v3.daily_balance_index"], + dimensions=["v3.product.category"], + ) + assert result is None + + @pytest.mark.asyncio + async def test_pinned_reaggregate_cube_missing_protected_dimension_raises( + self, + client_with_build_v3, + session, + ): + """Pinned cube validation also blocks unsafe semi-additive materialization.""" + await _create_daily_balance_metric(client_with_build_v3) + response = await client_with_build_v3.post( + "/nodes/cube/", + json={ + "name": "v3.test_daily_balance_pinned_unsafe_cube", + "metrics": ["v3.daily_balance"], + "dimensions": ["v3.product.category"], + "mode": "published", + "description": "Pinned unsafe semi-additive cube", + }, + ) + assert response.status_code == 201, response.json() + cube_node = await Node.get_cube_by_name( + session, + "v3.test_daily_balance_pinned_unsafe_cube", + ) + assert cube_node and cube_node.current + + with pytest.raises(DJInvalidInputException) as exc: + await validate_pinned_cube_covers_filters( + session, + cube_node.current, + dimensions=["v3.product.category"], + filters=[], + metrics=["v3.daily_balance"], + ) + assert "protected dimension" in str(exc.value) + + @pytest.mark.asyncio + async def test_pinned_derived_reaggregate_cube_missing_protected_dimension_raises( + self, + client_with_build_v3, + session, + ): + """Pinned derived metric cubes also inherit base semi-additive needs.""" + await _create_daily_balance_metric(client_with_build_v3) + await _create_daily_balance_index_metric(client_with_build_v3) + response = await client_with_build_v3.post( + "/nodes/cube/", + json={ + "name": "v3.test_daily_balance_index_pinned_unsafe_cube", + "metrics": ["v3.daily_balance_index"], + "dimensions": ["v3.product.category"], + "mode": "published", + "description": "Pinned unsafe derived semi-additive cube", + }, + ) + assert response.status_code == 201, response.json() + cube_node = await Node.get_cube_by_name( + session, + "v3.test_daily_balance_index_pinned_unsafe_cube", + ) + assert cube_node and cube_node.current + + with pytest.raises(DJInvalidInputException) as exc: + await validate_pinned_cube_covers_filters( + session, + cube_node.current, + dimensions=["v3.product.category"], + filters=[], + metrics=["v3.daily_balance_index"], + ) + assert "protected dimension" in str(exc.value) + @pytest.mark.asyncio async def test_pinned_cube_non_druid_dialect_skips_validation( self, diff --git a/datajunction-server/tests/construction/build_v3/decomposition_test.py b/datajunction-server/tests/construction/build_v3/decomposition_test.py index 5a8a6af39..7c5981a69 100644 --- a/datajunction-server/tests/construction/build_v3/decomposition_test.py +++ b/datajunction-server/tests/construction/build_v3/decomposition_test.py @@ -5,9 +5,11 @@ import pytest from datajunction_server.construction.build_v3.decomposition import ( + _reaggregate_dimension_requested, get_base_metrics_for_derived, is_derived_metric, ) +from datajunction_server.construction.build_v3.metrics import _source_dimension_alias from datajunction_server.construction.build_v3.types import BuildContext from datajunction_server.models.node_type import NodeType @@ -223,6 +225,58 @@ def test_metric_with_only_dimension_parent_not_derived(self): assert result is False +@pytest.mark.parametrize( + ("protected_dimension", "requested_dimensions", "expected"), + [ + ("v3.date.date_id[order]", ["v3.date.date_id[order]"], True), + ("v3.date.date_id[order]", ["v3.date.date_id[order]]"], True), + ("v3.date.date_id[order]", ["v3.date.date_id"], False), + ("v3.date.date_id[order]", ["v3.date.date_id[ship]"], False), + ("v3.date.date_id[order]", ["date_id"], False), + ("v3.date.date_id", ["v3.date.date_id"], True), + ("v3.date.date_id", ["date_id"], True), + ("v3.date.date_id", ["v3.date.date_id[order]"], False), + ], +) +def test_reaggregate_dimension_requested_is_role_sensitive( + protected_dimension, + requested_dimensions, + expected, +): + """A role-less ref must not satisfy a role-qualified protected dimension.""" + assert ( + _reaggregate_dimension_requested(protected_dimension, requested_dimensions) + is expected + ) + + +def test_source_dimension_alias_does_not_fall_back_for_roled_refs(): + """A roled protected dimension must collapse by that role's physical alias.""" + source_dimension_aliases = { + "v3.date.date_id": "date_id", + "v3.date.date_id[ship]": "date_id_ship", + } + + assert ( + _source_dimension_alias( + source_dimension_aliases, + "v3.date.date_id[order]", + ) + is None + ) + assert ( + _source_dimension_alias( + source_dimension_aliases, + "v3.date.date_id[ship]", + ) + == "date_id_ship" + ) + assert ( + _source_dimension_alias(source_dimension_aliases, "v3.date.date_id") + == "date_id" + ) + + @pytest.mark.asyncio async def test_decomposition_with_dimension_parent_integration( module__client_with_build_v3, diff --git a/datajunction-server/tests/construction/build_v3/dimension_refs_test.py b/datajunction-server/tests/construction/build_v3/dimension_refs_test.py new file mode 100644 index 000000000..2d8a2f44a --- /dev/null +++ b/datajunction-server/tests/construction/build_v3/dimension_refs_test.py @@ -0,0 +1,24 @@ +"""Tests for low-level dimension reference parsing.""" + +import pytest + +from datajunction_server.construction.build_v3.dimension_refs import ( + split_dimension_ref, +) + + +@pytest.mark.parametrize( + ("ref", "expected"), + [ + ("v3.date.date_id", ("v3.date.date_id", None)), + ("v3.date.date_id[order]", ("v3.date.date_id", "order")), + ("date_id[order]", ("date_id", "order")), + ( + "v3.date.date_id[customer->registration]", + ("v3.date.date_id", "customer->registration"), + ), + ], +) +def test_split_dimension_ref(ref, expected): + """Qualified and bare refs share the same role parsing.""" + assert split_dimension_ref(ref) == expected diff --git a/datajunction-server/tests/construction/build_v3/metrics_sql_test.py b/datajunction-server/tests/construction/build_v3/metrics_sql_test.py index 1624493e0..8dab1b2c7 100644 --- a/datajunction-server/tests/construction/build_v3/metrics_sql_test.py +++ b/datajunction-server/tests/construction/build_v3/metrics_sql_test.py @@ -1,3 +1,5 @@ +import time + import pytest from . import assert_sql_equal @@ -284,90 +286,770 @@ async def test_basic_metrics_sql(self, client_with_build_v3): }, ) - # Should return 200 OK with SQL - assert response.status_code == 200 - result = response.json() - assert "sql" in result - assert result["sql"] - assert "SELECT" in result["sql"].upper() + # Should return 200 OK with SQL + assert response.status_code == 200 + result = response.json() + assert "sql" in result + assert result["sql"] + assert "SELECT" in result["sql"].upper() + + @pytest.mark.asyncio + async def test_explicit_dialect_is_honored(self, client_with_build_v3): + """ + An explicit ``dialect`` param bypasses auto-resolution: the shared + ``generate_metrics_sql`` helper uses the caller-supplied dialect instead + of resolving one from cube availability / the metric's catalog. + """ + response = await client_with_build_v3.get( + "/sql/metrics/v3/", + params={ + "metrics": ["v3.total_revenue"], + "dimensions": ["v3.order_details.status"], + "dialect": "trino", + }, + ) + + assert response.status_code == 200, response.json() + assert response.json()["dialect"] == "trino" + + @pytest.mark.asyncio + async def test_simple_single_metric(self, client_with_build_v3): + """ + Test metrics SQL for a single simple metric (SUM). + + Even for single grain groups, the unified generate_metrics_sql + wraps the result in a grain group CTE (e.g., order_details_0) for consistency. + """ + response = await client_with_build_v3.get( + "/sql/metrics/v3/", + params={ + "metrics": ["v3.total_revenue"], + "dimensions": ["v3.order_details.status"], + }, + ) + + assert response.status_code == 200, response.json() + result = response.json() + + # Should have SQL output with shared CTEs, grain group wrapper, + # and re-aggregation in final SELECT (always applied for consistency) + assert_sql_equal( + result["sql"], + """ + WITH + v3_order_details AS ( + SELECT o.status, oi.quantity * oi.unit_price AS line_total + FROM default.v3.orders o + JOIN default.v3.order_items oi ON o.order_id = oi.order_id + ), + order_details_0 AS ( + SELECT t1.status, SUM(t1.line_total) line_total_sum_e1f61696 + FROM v3_order_details t1 + GROUP BY t1.status + ) + SELECT order_details_0.status AS status, + SUM(order_details_0.line_total_sum_e1f61696) AS total_revenue + FROM order_details_0 + GROUP BY order_details_0.status + """, + ) + + # Should have columns (names match SQL AS aliases) + assert result["columns"] == [ + { + "name": "status", + "type": "string", + "semantic_entity": "v3.order_details.status", + "semantic_type": "dimension", + }, + { + "name": "total_revenue", + "type": "double", + "semantic_entity": "v3.total_revenue", + "semantic_type": "metric", + }, + ] + + @staticmethod + async def _create_daily_balance_metric(client_with_build_v3): + response = await client_with_build_v3.post( + "/nodes/metric/", + json={ + "name": "v3.daily_balance", + "description": "Semi-additive balance measured by order date", + "query": "SELECT SUM(line_total) FROM v3.order_details", + "mode": "published", + "reaggregate": { + "rules": [ + { + "dimension": "v3.date.date_id[order]", + "fn": "last_value", + }, + ], + }, + }, + ) + assert response.status_code in (200, 201), response.json() + + @staticmethod + async def _create_scaled_daily_balance_metric(client_with_build_v3): + response = await client_with_build_v3.post( + "/nodes/metric/", + json={ + "name": "v3.scaled_daily_balance", + "description": "Scaled semi-additive balance measured by order date", + "query": "SELECT SUM(line_total) / 100 FROM v3.order_details", + "mode": "published", + "reaggregate": { + "rules": [ + { + "dimension": "v3.date.date_id[order]", + "fn": "last_value", + }, + ], + }, + }, + ) + assert response.status_code in (200, 201), response.json() + + @staticmethod + async def _create_first_daily_balance_metric(client_with_build_v3): + response = await client_with_build_v3.post( + "/nodes/metric/", + json={ + "name": "v3.first_daily_balance", + "description": "Semi-additive balance measured by first order date", + "query": "SELECT SUM(line_total) FROM v3.order_details", + "mode": "published", + "reaggregate": { + "rules": [ + { + "dimension": "v3.date.date_id[order]", + "fn": "first_value", + }, + ], + }, + }, + ) + assert response.status_code in (200, 201), response.json() + + @staticmethod + async def _create_daily_balance_index_metric(client_with_build_v3): + response = await client_with_build_v3.post( + "/nodes/metric/", + json={ + "name": "v3.daily_balance_index", + "description": "Derived semi-additive balance index", + "query": "SELECT 10.0 / v3.daily_balance", + "mode": "published", + }, + ) + assert response.status_code in (200, 201), response.json() + + @staticmethod + async def _create_wow_daily_balance_index_metric(client_with_build_v3): + response = await client_with_build_v3.post( + "/nodes/metric/", + json={ + "name": "v3.wow_daily_balance_index", + "description": "Week-over-week daily balance index change", + "query": """ + SELECT + (v3.daily_balance_index - LAG(v3.daily_balance_index, 1) + OVER (ORDER BY v3.date.week[order])) + / NULLIF( + LAG(v3.daily_balance_index, 1) + OVER (ORDER BY v3.date.week[order]), + 0 + ) * 100 + """, + "mode": "published", + }, + ) + assert response.status_code in (200, 201), response.json() + + @staticmethod + async def _create_product_balance_metric(client_with_build_v3): + response = await client_with_build_v3.post( + "/nodes/metric/", + json={ + "name": "v3.product_balance", + "description": "Semi-additive balance protected by product", + "query": "SELECT SUM(line_total) FROM v3.order_details", + "mode": "published", + "reaggregate": { + "rules": [ + { + "dimension": "v3.product.product_id", + "fn": "last_value", + }, + ], + }, + }, + ) + assert response.status_code in (200, 201), response.json() + + @staticmethod + async def _create_product_balance_index_metric(client_with_build_v3): + response = await client_with_build_v3.post( + "/nodes/metric/", + json={ + "name": "v3.product_balance_index", + "description": "Derived semi-additive product balance index", + "query": "SELECT 10.0 / v3.product_balance", + "mode": "published", + }, + ) + assert response.status_code in (200, 201), response.json() + + @staticmethod + async def _create_wow_product_balance_index_metric(client_with_build_v3): + response = await client_with_build_v3.post( + "/nodes/metric/", + json={ + "name": "v3.wow_product_balance_index", + "description": "Week-over-week product balance index change", + "query": """ + SELECT + (v3.product_balance_index - LAG(v3.product_balance_index, 1) + OVER (ORDER BY v3.date.week[order])) + / NULLIF( + LAG(v3.product_balance_index, 1) + OVER (ORDER BY v3.date.week[order]), + 0 + ) * 100 + """, + "mode": "published", + }, + ) + assert response.status_code in (200, 201), response.json() + + @pytest.mark.asyncio + async def test_reaggregate_collapses_when_protected_dimension_omitted( + self, + client_with_build_v3, + ): + """A semi-additive metric keeps its protected dimension as private grain.""" + await self._create_daily_balance_metric(client_with_build_v3) + + response = await client_with_build_v3.get( + "/sql/metrics/v3/", + params={ + "metrics": ["v3.daily_balance"], + "dimensions": ["v3.product.category"], + "use_materialized": "false", + }, + ) + assert response.status_code == 200, response.json() + + assert_sql_equal( + response.json()["sql"], + """ + WITH + v3_order_details AS ( + SELECT o.order_date, oi.product_id, oi.quantity * oi.unit_price AS line_total + FROM default.v3.orders o + JOIN default.v3.order_items oi ON o.order_id = oi.order_id + ), + v3_product AS ( + SELECT product_id, category + FROM default.v3.products + ), + order_details_0 AS ( + SELECT t2.category, + t1.order_date AS date_id_order, + SUM(t1.line_total) AS line_total_sum_e1f61696 + FROM v3_order_details t1 + LEFT OUTER JOIN v3_product t2 ON t1.product_id = t2.product_id + GROUP BY t2.category, t1.order_date + ) + SELECT order_details_0.category AS category, + MAX_BY( + order_details_0.line_total_sum_e1f61696, + order_details_0.date_id_order + ) AS daily_balance + FROM order_details_0 + GROUP BY order_details_0.category + """, + ) + assert response.json()["columns"] == [ + { + "name": "category", + "type": "string", + "semantic_entity": "v3.product.category", + "semantic_type": "dimension", + }, + { + "name": "daily_balance", + "type": "double", + "semantic_entity": "v3.daily_balance", + "semantic_type": "metric", + }, + ] + + @pytest.mark.asyncio + async def test_reaggregate_filter_keeps_protected_dimension_as_private_grain( + self, + client_with_build_v3, + ): + """Filtering protected grain must not make a reaggregate metric additive.""" + await self._create_daily_balance_metric(client_with_build_v3) + + response = await client_with_build_v3.get( + "/sql/metrics/v3/", + params={ + "metrics": ["v3.daily_balance"], + "dimensions": ["v3.product.category"], + "filters": ["v3.date.date_id[order] >= 20260101"], + "use_materialized": "false", + }, + ) + assert response.status_code == 200, response.json() + + sql = response.json()["sql"] + normalized_sql = " ".join(sql.split()) + assert "WHERE o.order_date >= 20260101" in normalized_sql + assert "t1.order_date" in normalized_sql + assert "date_id_order" in normalized_sql + assert "GROUP BY t2.category, t1.order_date" in normalized_sql + assert "MAX_BY(" in sql + assert ( + "SUM(order_details_0.line_total_sum_e1f61696) AS daily_balance" + not in normalized_sql + ) + assert [column["name"] for column in response.json()["columns"]] == [ + "category", + "daily_balance", + ] + + @pytest.mark.asyncio + async def test_reaggregate_with_limited_metric_keeps_protected_grain( + self, + client_with_build_v3, + ): + """Merged metric queries must not add limited grain to the collapse CTE.""" + await self._create_daily_balance_metric(client_with_build_v3) + + response = await client_with_build_v3.get( + "/sql/metrics/v3/", + params={ + "metrics": ["v3.daily_balance", "v3.order_count"], + "dimensions": ["v3.product.category"], + "use_materialized": "false", + }, + ) + assert response.status_code == 200, response.json() + + sql = response.json()["sql"] + normalized_sql = " ".join(sql.split()) + assert "order_details_0 AS" in sql + assert "GROUP BY t2.category, t1.order_date" in normalized_sql + assert "GROUP BY t2.category, t1.order_date, t1.order_id" not in normalized_sql + assert "order_details_1_agg AS" in sql + assert "COUNT( DISTINCT order_id)" in normalized_sql + assert "MAX_BY(" in sql + assert "order_details_0.date_id_order" in sql + + @pytest.mark.asyncio + async def test_reaggregate_preserves_single_component_combiner_wrapper( + self, + client_with_build_v3, + ): + """Semi-additive collapse preserves arithmetic around the aggregate.""" + await self._create_scaled_daily_balance_metric(client_with_build_v3) + + response = await client_with_build_v3.get( + "/sql/metrics/v3/", + params={ + "metrics": ["v3.scaled_daily_balance"], + "dimensions": ["v3.product.category"], + "use_materialized": "false", + }, + ) + assert response.status_code == 200, response.json() + + sql = response.json()["sql"] + normalized_sql = " ".join(sql.split()) + assert "MAX_BY(" in sql + assert "order_details_0.date_id_order" in sql + assert "/ 100 AS scaled_daily_balance" in normalized_sql + assert ( + "MAX_BY(order_details_0.line_total_sum_e1f61696, order_details_0.date_id_order) / 100" + in normalized_sql + ) + + @pytest.mark.asyncio + async def test_reaggregate_druid_uses_latest_by_for_last_value( + self, + client_with_build_v3, + ): + """Druid live SQL uses Druid-native collapse functions.""" + await self._create_daily_balance_metric(client_with_build_v3) + + response = await client_with_build_v3.get( + "/sql/metrics/v3/", + params={ + "metrics": ["v3.daily_balance"], + "dimensions": ["v3.product.category"], + "dialect": "druid", + "use_materialized": "false", + }, + ) + assert response.status_code == 200, response.json() + assert response.json()["dialect"] == "druid" + + assert_sql_equal( + response.json()["sql"], + """ + WITH + v3_order_details AS ( + SELECT o.order_date, oi.product_id, oi.quantity * oi.unit_price AS line_total + FROM default.v3.orders o + JOIN default.v3.order_items oi ON o.order_id = oi.order_id + ), + v3_product AS ( + SELECT product_id, category + FROM default.v3.products + ), + order_details_0 AS ( + SELECT t2.category, + t1.order_date AS date_id_order, + SUM(t1.line_total) AS line_total_sum_e1f61696 + FROM v3_order_details t1 + LEFT OUTER JOIN v3_product t2 ON t1.product_id = t2.product_id + GROUP BY t2.category, t1.order_date + ) + SELECT order_details_0.category AS category, + LATEST_BY( + order_details_0.line_total_sum_e1f61696, + order_details_0.date_id_order + ) AS daily_balance + FROM order_details_0 + GROUP BY order_details_0.category + """, + ) + assert "MAX_BY" not in response.json()["sql"] + assert "ARG_MAX" not in response.json()["sql"] + + @pytest.mark.asyncio + async def test_reaggregate_auto_routed_druid_uses_latest_by( + self, + client_with_build_v3, + ): + """Auto-routed Druid cube SQL uses Druid-native collapse functions.""" + await self._create_daily_balance_metric(client_with_build_v3) + + response = await client_with_build_v3.post( + "/nodes/cube/", + json={ + "name": "v3.daily_balance_category_cube", + "metrics": ["v3.daily_balance"], + "dimensions": ["v3.product.category"], + "mode": "published", + "description": "Category-only decoy cube", + }, + ) + assert response.status_code == 201, response.json() + + response = await client_with_build_v3.post( + "/data/v3.daily_balance_category_cube/availability/", + json={ + "catalog": "default", + "schema_": "analytics", + "table": "daily_balance_category_cube", + "valid_through_ts": int(time.time() * 1000), + }, + ) + assert response.status_code == 200, response.json() + + response = await client_with_build_v3.post( + "/nodes/cube/", + json={ + "name": "v3.daily_balance_cube", + "metrics": ["v3.daily_balance"], + "dimensions": [ + "v3.product.category", + "v3.date.date_id[order]", + ], + "mode": "published", + "description": "Daily balance cube at protected grain", + }, + ) + assert response.status_code == 201, response.json() + + response = await client_with_build_v3.post( + "/data/v3.daily_balance_cube/availability/", + json={ + "catalog": "default", + "schema_": "analytics", + "table": "daily_balance_cube", + "valid_through_ts": int(time.time() * 1000), + }, + ) + assert response.status_code == 200, response.json() + + response = await client_with_build_v3.get( + "/sql/metrics/v3/", + params={ + "metrics": ["v3.daily_balance"], + "dimensions": ["v3.product.category"], + }, + ) + assert response.status_code == 200, response.json() + assert response.json()["dialect"] == "druid" + + sql = response.json()["sql"] + assert "FROM daily_balance_cube" in sql + assert "daily_balance_category_cube" not in sql + assert "LATEST_BY(" in sql + assert "date_id_order" in sql + assert "MAX_BY" not in sql + assert "ARG_MAX" not in sql + + @pytest.mark.asyncio + async def test_reaggregate_druid_uses_earliest_by_for_first_value( + self, + client_with_build_v3, + ): + """Druid first-value collapse renders as EARLIEST_BY.""" + await self._create_first_daily_balance_metric(client_with_build_v3) + + response = await client_with_build_v3.get( + "/sql/metrics/v3/", + params={ + "metrics": ["v3.first_daily_balance"], + "dimensions": ["v3.product.category"], + "dialect": "druid", + "use_materialized": "false", + }, + ) + assert response.status_code == 200, response.json() + + sql = response.json()["sql"] + assert "EARLIEST_BY(" in sql + assert "MIN_BY" not in sql + assert "ARG_MIN" not in sql + + @pytest.mark.asyncio + async def test_reaggregate_derived_metric_uses_collapsed_base( + self, + client_with_build_v3, + ): + """Derived metrics inline semi-additive collapse in denominator position.""" + await self._create_daily_balance_metric(client_with_build_v3) + await self._create_daily_balance_index_metric(client_with_build_v3) + + response = await client_with_build_v3.get( + "/sql/metrics/v3/", + params={ + "metrics": ["v3.daily_balance_index"], + "dimensions": ["v3.product.category"], + "use_materialized": "false", + }, + ) + assert response.status_code == 200, response.json() + + sql = response.json()["sql"] + assert "MAX_BY(" in sql + assert "order_details_0.date_id_order" in sql + assert "MAX_BY" in sql.split(" AS daily_balance_index")[0] + assert "10.0 / NULLIF(MAX_BY(" in sql + assert "10.0 / NULLIF(SUM(" not in sql + + @pytest.mark.asyncio + async def test_reaggregate_derived_metric_preserves_requested_protected_dimension( + self, + client_with_build_v3, + ): + """Derived metrics use normal aggregation when the protected dimension is requested.""" + await self._create_daily_balance_metric(client_with_build_v3) + await self._create_daily_balance_index_metric(client_with_build_v3) + + response = await client_with_build_v3.get( + "/sql/metrics/v3/", + params={ + "metrics": ["v3.daily_balance_index"], + "dimensions": ["v3.date.date_id[order]"], + "use_materialized": "false", + }, + ) + assert response.status_code == 200, response.json() + + sql = response.json()["sql"] + assert "10.0 / NULLIF(SUM(" in sql + assert "date_id_order" in sql + assert "MAX_BY(" not in sql + + @pytest.mark.asyncio + async def test_reaggregate_nested_window_metric_reaggregates_with_collapse( + self, + client_with_build_v3, + ): + """Window reaggregation of derived metrics uses semi-additive parents.""" + await self._create_daily_balance_metric(client_with_build_v3) + await self._create_daily_balance_index_metric(client_with_build_v3) + await self._create_wow_daily_balance_index_metric(client_with_build_v3) + + response = await client_with_build_v3.get( + "/sql/metrics/v3/", + params={ + "metrics": ["v3.wow_daily_balance_index"], + "dimensions": ["v3.product.category"], + "use_materialized": "false", + }, + ) + assert response.status_code == 200, response.json() + + sql = response.json()["sql"] + assert "base_metrics AS" in sql + assert "MAX_BY(" in sql + assert "order_details_0.date_id_order" in sql + assert "10.0 / NULLIF(MAX_BY(" in sql + assert "LAG(base_metrics.daily_balance_index, 1)" in sql + + @pytest.mark.asyncio + async def test_reaggregate_window_grain_reaggregation_uses_collapse( + self, + client_with_build_v3, + ): + """Window aggregation CTEs collapse semi-additive derived parents.""" + await self._create_product_balance_metric(client_with_build_v3) + await self._create_product_balance_index_metric(client_with_build_v3) + await self._create_wow_product_balance_index_metric(client_with_build_v3) + + response = await client_with_build_v3.get( + "/sql/metrics/v3/", + params={ + "metrics": ["v3.wow_product_balance_index"], + "dimensions": ["v3.date.date_id[order]"], + "use_materialized": "false", + }, + ) + assert response.status_code == 200, response.json() + + sql = response.json()["sql"] + assert "order_details_week_agg AS" in sql + assert "MAX_BY(" in sql + assert "product_id" in sql + assert "10.0 / NULLIF(MAX_BY(" in sql + assert "10.0 / NULLIF(SUM(" not in sql + assert "LAG(order_details_week_agg.product_balance_index, 1)" in sql + + @pytest.mark.asyncio + async def test_reaggregate_uses_normal_aggregation_when_protected_dimension_requested( + self, + client_with_build_v3, + ): + """Requesting the protected dimension means there is nothing to collapse.""" + await self._create_daily_balance_metric(client_with_build_v3) + + response = await client_with_build_v3.get( + "/sql/metrics/v3/", + params={ + "metrics": ["v3.daily_balance"], + "dimensions": ["v3.date.date_id[order]"], + "use_materialized": "false", + }, + ) + assert response.status_code == 200, response.json() + + assert_sql_equal( + response.json()["sql"], + """ + WITH + v3_order_details AS ( + SELECT o.order_date, oi.quantity * oi.unit_price AS line_total + FROM default.v3.orders o + JOIN default.v3.order_items oi ON o.order_id = oi.order_id + ), + order_details_0 AS ( + SELECT t1.order_date AS date_id_order, + SUM(t1.line_total) AS line_total_sum_e1f61696 + FROM v3_order_details t1 + GROUP BY t1.order_date + ) + SELECT order_details_0.date_id_order AS date_id_order, + SUM(order_details_0.line_total_sum_e1f61696) AS daily_balance + FROM order_details_0 + GROUP BY order_details_0.date_id_order + """, + ) + assert "MAX_BY" not in response.json()["sql"] @pytest.mark.asyncio - async def test_explicit_dialect_is_honored(self, client_with_build_v3): - """ - An explicit ``dialect`` param bypasses auto-resolution: the shared - ``generate_metrics_sql`` helper uses the caller-supplied dialect instead - of resolving one from cube availability / the metric's catalog. - """ + async def test_reaggregate_collapses_when_roleless_base_dimension_requested( + self, + client_with_build_v3, + ): + """A role-less dimension request does not satisfy a roled protected grain.""" + await self._create_daily_balance_metric(client_with_build_v3) + response = await client_with_build_v3.get( "/sql/metrics/v3/", params={ - "metrics": ["v3.total_revenue"], - "dimensions": ["v3.order_details.status"], - "dialect": "trino", + "metrics": ["v3.daily_balance"], + "dimensions": ["v3.date.date_id"], + "use_materialized": "false", }, ) - assert response.status_code == 200, response.json() - assert response.json()["dialect"] == "trino" + + sql = response.json()["sql"] + assert "MAX_BY(" in sql + assert "order_details_0.date_id_order" in sql + assert ( + "SUM(order_details_0.line_total_sum_e1f61696) AS daily_balance" not in sql + ) @pytest.mark.asyncio - async def test_simple_single_metric(self, client_with_build_v3): - """ - Test metrics SQL for a single simple metric (SUM). + async def test_reaggregate_collapses_with_coarser_time_dimension( + self, + client_with_build_v3, + ): + """Coarser time output still collapses across the protected date grain.""" + await self._create_daily_balance_metric(client_with_build_v3) - Even for single grain groups, the unified generate_metrics_sql - wraps the result in a grain group CTE (e.g., order_details_0) for consistency. - """ response = await client_with_build_v3.get( "/sql/metrics/v3/", params={ - "metrics": ["v3.total_revenue"], - "dimensions": ["v3.order_details.status"], + "metrics": ["v3.daily_balance"], + "dimensions": ["v3.date.month[order]"], + "use_materialized": "false", }, ) - assert response.status_code == 200, response.json() - result = response.json() - # Should have SQL output with shared CTEs, grain group wrapper, - # and re-aggregation in final SELECT (always applied for consistency) assert_sql_equal( - result["sql"], + response.json()["sql"], """ WITH + v3_date AS ( + SELECT date_id, month + FROM default.v3.dates + ), v3_order_details AS ( - SELECT o.status, oi.quantity * oi.unit_price AS line_total + SELECT o.order_date, oi.quantity * oi.unit_price AS line_total FROM default.v3.orders o JOIN default.v3.order_items oi ON o.order_id = oi.order_id ), order_details_0 AS ( - SELECT t1.status, SUM(t1.line_total) line_total_sum_e1f61696 + SELECT t2.month AS month_order, + COALESCE(t1.order_date, t2.date_id) AS date_id_order, + SUM(t1.line_total) AS line_total_sum_e1f61696 FROM v3_order_details t1 - GROUP BY t1.status + LEFT OUTER JOIN v3_date t2 ON t1.order_date = t2.date_id + GROUP BY t2.month, COALESCE(t1.order_date, t2.date_id) ) - SELECT order_details_0.status AS status, - SUM(order_details_0.line_total_sum_e1f61696) AS total_revenue + SELECT order_details_0.month_order AS month_order, + MAX_BY( + order_details_0.line_total_sum_e1f61696, + order_details_0.date_id_order + ) AS daily_balance FROM order_details_0 - GROUP BY order_details_0.status + GROUP BY order_details_0.month_order """, ) - # Should have columns (names match SQL AS aliases) - assert result["columns"] == [ - { - "name": "status", - "type": "string", - "semantic_entity": "v3.order_details.status", - "semantic_type": "dimension", - }, - { - "name": "total_revenue", - "type": "double", - "semantic_entity": "v3.total_revenue", - "semantic_type": "metric", - }, - ] - @pytest.mark.asyncio async def test_multiple_metrics_same_grain(self, client_with_build_v3): """ @@ -3936,12 +4618,8 @@ async def test_cross_fact_window_metric_with_finer_grain( - Window metric orders by weekly grain (v3.date.week) - Metric is cross-fact (conversion_rate = order_count / visitor_count) - Expected behavior: - 1. Grain groups are built at daily grain (include date_id AND week) - 2. base_metrics CTE combines facts with FULL OUTER JOIN at daily grain - 3. A window aggregation CTE reaggregates base_metrics to weekly grain - 4. Window CTE applies LAG on the weekly-aggregated data - 5. Final SELECT joins daily base_metrics with weekly window results + Reaggregation must fail because base_metrics no longer retains the distinct + grain keys needed to collapse order_count and visitor_count safely. """ # Create the metric locally for this test response = await client_with_build_v3.post( @@ -3972,77 +4650,10 @@ async def test_cross_fact_window_metric_with_finer_grain( }, ) - assert response.status_code == 200, response.json() - result = response.json() - - # Verify the SQL has the expected structure with reaggregation CTE - # The key is that there should be a CTE that aggregates from base_metrics - # to weekly grain before applying the LAG window function - sql = result["sql"] - assert_sql_equal( - sql, - """ - WITH - v3_date AS ( - SELECT date_id, - week - FROM default.v3.dates - ), - v3_order_details AS ( - SELECT o.order_id, - o.order_date, - oi.product_id - FROM default.v3.orders o JOIN default.v3.order_items oi ON o.order_id = oi.order_id - ), - v3_product AS ( - SELECT product_id, - category - FROM default.v3.products - ), - v3_page_views_enriched AS ( - SELECT customer_id, - page_date, - product_id - FROM default.v3.page_views - ), - order_details_0 AS ( - SELECT COALESCE(t1.order_date, t3.date_id) AS date_id, - t2.category, - t3.week, - t1.order_id - FROM v3_order_details t1 LEFT OUTER JOIN v3_product t2 ON t1.product_id = t2.product_id - LEFT OUTER JOIN v3_date t3 ON t1.order_date = t3.date_id - GROUP BY COALESCE(t1.order_date, t3.date_id), t2.category, t3.week, t1.order_id - ), - page_views_enriched_0 AS ( - SELECT COALESCE(t1.page_date, t3.date_id) AS date_id, - t2.category, - t3.week, - t1.customer_id - FROM v3_page_views_enriched t1 LEFT OUTER JOIN v3_product t2 ON t1.product_id = t2.product_id - LEFT OUTER JOIN v3_date t3 ON t1.page_date = t3.date_id - GROUP BY COALESCE(t1.page_date, t3.date_id), t2.category, t3.week, t1.customer_id - ), - base_metrics AS ( - SELECT COALESCE(order_details_0.date_id, page_views_enriched_0.date_id) AS date_id, - COALESCE(order_details_0.category, page_views_enriched_0.category) AS category, - COALESCE(order_details_0.week, page_views_enriched_0.week) AS week, - COUNT( DISTINCT order_details_0.order_id) AS order_count, - COUNT( DISTINCT page_views_enriched_0.customer_id) AS visitor_count, - CAST(COUNT( DISTINCT order_details_0.order_id) AS DOUBLE) / NULLIF(COUNT( DISTINCT page_views_enriched_0.customer_id), 0) AS conversion_rate - FROM order_details_0 FULL OUTER JOIN page_views_enriched_0 ON order_details_0.date_id = page_views_enriched_0.date_id AND order_details_0.category = page_views_enriched_0.category AND order_details_0.week = page_views_enriched_0.week - GROUP BY 1, 2, 3 - ) - - SELECT base_metrics.date_id AS date_id, - base_metrics.category AS category, - base_metrics.week AS week, - (base_metrics.conversion_rate - LAG(base_metrics.conversion_rate, 1) OVER ( PARTITION BY base_metrics.category - ORDER BY base_metrics.week) ) / NULLIF(LAG(base_metrics.conversion_rate, 1) OVER ( PARTITION BY base_metrics.category - ORDER BY base_metrics.week) , 0) * 100 AS wow_conversion_rate_change - FROM base_metrics - """, - ) + assert response.status_code == 422, response.json() + message = response.json()["message"] + assert "Unsupported distinct metric reaggregation" in message + assert "no longer retains the distinct grain key" in message @pytest.mark.asyncio async def test_cross_fact_window_on_derived_metric(self, client_with_build_v3): @@ -4055,8 +4666,8 @@ async def test_cross_fact_window_on_derived_metric(self, client_with_build_v3): - pages_per_session = page_view_count / visitor_count (derived from page_views) - wow_efficiency_ratio_change = LAG(efficiency_ratio, 1) OVER (ORDER BY week) - This hits lines 1029-1055 in metrics.py where derived metrics are expanded - by replacing column references with parent metric expressions. + Expanding the derived metric reaches LIMITED leaf metrics, which cannot be + collapsed after base_metrics discards their distinct grain keys. """ # Create the metric locally for this test response = await client_with_build_v3.post( @@ -4082,80 +4693,10 @@ async def test_cross_fact_window_on_derived_metric(self, client_with_build_v3): }, ) - assert response.status_code == 200, response.json() - result = response.json() - sql = result["sql"] - assert_sql_equal( - sql, - """ - WITH - v3_date AS ( - SELECT date_id, - week - FROM default.v3.dates - ), - v3_order_details AS ( - SELECT o.order_id, - o.order_date, - oi.product_id, - oi.quantity * oi.unit_price AS line_total - FROM default.v3.orders o JOIN default.v3.order_items oi ON o.order_id = oi.order_id - ), - v3_product AS ( - SELECT product_id, - category - FROM default.v3.products - ), - v3_page_views_enriched AS ( - SELECT view_id, - session_id, - page_date, - product_id - FROM default.v3.page_views - ), - order_details_0 AS ( - SELECT COALESCE(t1.order_date, t3.date_id) AS date_id, - t2.category, - t3.week, - t1.order_id, - SUM(t1.line_total) line_total_sum_e1f61696 - FROM v3_order_details t1 LEFT OUTER JOIN v3_product t2 ON t1.product_id = t2.product_id - LEFT OUTER JOIN v3_date t3 ON t1.order_date = t3.date_id - GROUP BY COALESCE(t1.order_date, t3.date_id), t2.category, t3.week, t1.order_id - ), - page_views_enriched_0 AS ( - SELECT COALESCE(t1.page_date, t3.date_id) AS date_id, - t2.category, - t3.week, - t1.session_id, - COUNT(t1.view_id) view_id_count_f41e2db4 - FROM v3_page_views_enriched t1 LEFT OUTER JOIN v3_product t2 ON t1.product_id = t2.product_id - LEFT OUTER JOIN v3_date t3 ON t1.page_date = t3.date_id - GROUP BY COALESCE(t1.page_date, t3.date_id), t2.category, t3.week, t1.session_id - ), - base_metrics AS ( - SELECT COALESCE(order_details_0.date_id, page_views_enriched_0.date_id) AS date_id, - COALESCE(order_details_0.category, page_views_enriched_0.category) AS category, - COALESCE(order_details_0.week, page_views_enriched_0.week) AS week, - COUNT( DISTINCT order_details_0.order_id) AS order_count, - SUM(page_views_enriched_0.view_id_count_f41e2db4) AS page_view_count, - COUNT( DISTINCT page_views_enriched_0.session_id) AS session_count, - SUM(order_details_0.line_total_sum_e1f61696) AS total_revenue, - SUM(order_details_0.line_total_sum_e1f61696) / NULLIF(COUNT( DISTINCT order_details_0.order_id), 0) AS avg_order_value, - SUM(order_details_0.line_total_sum_e1f61696) / NULLIF(COUNT( DISTINCT order_details_0.order_id), 0) / NULLIF(SUM(page_views_enriched_0.view_id_count_f41e2db4) / NULLIF(COUNT( DISTINCT page_views_enriched_0.session_id), 0), 0) AS efficiency_ratio, - SUM(page_views_enriched_0.view_id_count_f41e2db4) / NULLIF(COUNT( DISTINCT page_views_enriched_0.session_id), 0) AS pages_per_session - FROM order_details_0 FULL OUTER JOIN page_views_enriched_0 ON order_details_0.date_id = page_views_enriched_0.date_id AND order_details_0.category = page_views_enriched_0.category AND order_details_0.week = page_views_enriched_0.week - GROUP BY 1, 2, 3 - ) - - SELECT base_metrics.date_id AS date_id, - base_metrics.category AS category, - base_metrics.week AS week, - (base_metrics.efficiency_ratio - LAG(base_metrics.efficiency_ratio, 1) OVER ( PARTITION BY base_metrics.category - ORDER BY base_metrics.week) ) / NULLIF(LAG(base_metrics.efficiency_ratio, 1) OVER ( PARTITION BY base_metrics.category - ORDER BY base_metrics.week) , 0) * 100 AS wow_efficiency_ratio_change - FROM base_metrics""", - ) + assert response.status_code == 422, response.json() + message = response.json()["message"] + assert "Unsupported distinct metric reaggregation" in message + assert "no longer retains the distinct grain key" in message @pytest.mark.asyncio async def test_cross_fact_window_on_base_metrics(self, client_with_build_v3): @@ -4168,10 +4709,12 @@ async def test_cross_fact_window_on_base_metrics(self, client_with_build_v3): - visitor_count is a base metric from page_views_enriched - Both are in grain groups (not derived metrics) - This should trigger build_window_agg_cte_from_base_metrics because: + This triggers build_window_agg_cte_from_base_metrics because: 1. It's cross-fact (order_count + visitor_count span multiple facts) 2. The base metrics ARE in grain groups (unlike derived metrics) 3. Window ORDER BY grain (week) is coarser than requested grain (date_id) + + The request must fail rather than summing the finalized distinct counts. """ # Create the metric locally for this test response = await client_with_build_v3.post( @@ -4199,84 +4742,10 @@ async def test_cross_fact_window_on_base_metrics(self, client_with_build_v3): }, ) - assert response.status_code == 200, response.json() - result = response.json() - sql = result["sql"] - assert_sql_equal( - sql, - """ - WITH - v3_date AS ( - SELECT date_id, - week - FROM default.v3.dates - ), - v3_order_details AS ( - SELECT o.order_id, - o.order_date, - oi.product_id - FROM default.v3.orders o JOIN default.v3.order_items oi ON o.order_id = oi.order_id - ), - v3_product AS ( - SELECT product_id, - category - FROM default.v3.products - ), - v3_page_views_enriched AS ( - SELECT customer_id, - page_date, - product_id - FROM default.v3.page_views - ), - order_details_0 AS ( - SELECT COALESCE(t1.order_date, t3.date_id) AS date_id, - t2.category, - t3.week, - t1.order_id - FROM v3_order_details t1 LEFT OUTER JOIN v3_product t2 ON t1.product_id = t2.product_id - LEFT OUTER JOIN v3_date t3 ON t1.order_date = t3.date_id - GROUP BY COALESCE(t1.order_date, t3.date_id), t2.category, t3.week, t1.order_id - ), - page_views_enriched_0 AS ( - SELECT COALESCE(t1.page_date, t3.date_id) AS date_id, - t2.category, - t3.week, - t1.customer_id - FROM v3_page_views_enriched t1 LEFT OUTER JOIN v3_product t2 ON t1.product_id = t2.product_id - LEFT OUTER JOIN v3_date t3 ON t1.page_date = t3.date_id - GROUP BY COALESCE(t1.page_date, t3.date_id), t2.category, t3.week, t1.customer_id - ), - base_metrics AS ( - SELECT COALESCE(order_details_0.date_id, page_views_enriched_0.date_id) AS date_id, - COALESCE(order_details_0.category, page_views_enriched_0.category) AS category, - COALESCE(order_details_0.week, page_views_enriched_0.week) AS week, - COUNT( DISTINCT order_details_0.order_id) AS order_count, - COUNT( DISTINCT page_views_enriched_0.customer_id) AS visitor_count - FROM order_details_0 FULL OUTER JOIN page_views_enriched_0 ON order_details_0.date_id = page_views_enriched_0.date_id AND order_details_0.category = page_views_enriched_0.category AND order_details_0.week = page_views_enriched_0.week - GROUP BY 1, 2, 3 - ), - order_details_week_agg AS ( - SELECT base_metrics.category AS category, - base_metrics.week AS week, - COUNT( DISTINCT base_metrics.order_id) AS order_count, - COUNT( DISTINCT base_metrics.customer_id) AS visitor_count - FROM base_metrics - GROUP BY base_metrics.category, base_metrics.week - ), - order_details_week AS ( - SELECT order_details_week_agg.category AS category, - order_details_week_agg.week AS week, - (order_details_week_agg.order_count + order_details_week_agg.visitor_count) - LAG(order_details_week_agg.order_count + order_details_week_agg.visitor_count, 1) OVER ( PARTITION BY order_details_week_agg.category - ORDER BY order_details_week_agg.week) AS wow_order_and_visitor_change - FROM order_details_week_agg - ) - SELECT base_metrics.date_id AS date_id, - base_metrics.category AS category, - base_metrics.week AS week, - order_details_week.wow_order_and_visitor_change AS wow_order_and_visitor_change - FROM base_metrics LEFT OUTER JOIN order_details_week ON base_metrics.category = order_details_week.category AND base_metrics.week = order_details_week.week - """, - ) + assert response.status_code == 422, response.json() + message = response.json()["message"] + assert "Unsupported distinct metric reaggregation" in message + assert "no longer retains the distinct grain key" in message @pytest.mark.asyncio async def test_multi_fact_window_metrics_same_grain(self, client_with_build_v3): diff --git a/datajunction-server/tests/construction/build_v3/preagg_substitution_test.py b/datajunction-server/tests/construction/build_v3/preagg_substitution_test.py index 0edd7753e..f8a045a70 100644 --- a/datajunction-server/tests/construction/build_v3/preagg_substitution_test.py +++ b/datajunction-server/tests/construction/build_v3/preagg_substitution_test.py @@ -86,6 +86,28 @@ async def _fake_columns(*args, **kwargs): del client.app.dependency_overrides[get_query_service_client] +async def _create_daily_balance_metric(client_with_build_v3): + """Create a semi-additive metric protected by order date.""" + response = await client_with_build_v3.post( + "/nodes/metric/", + json={ + "name": "v3.daily_balance", + "description": "Semi-additive balance measured by order date", + "query": "SELECT SUM(line_total) FROM v3.order_details", + "mode": "published", + "reaggregate": { + "rules": [ + { + "dimension": "v3.date.date_id[order]", + "fn": "last_value", + }, + ], + }, + }, + ) + assert response.status_code in (200, 201), response.json() + + class TestExternalPreAggRouting: """Queries route to externally-registered pre-agg tables via source_column.""" @@ -198,6 +220,98 @@ async def test_external_preagg_rolls_up_additive(self, client_with_build_v3): """, ) + @pytest.mark.asyncio + async def test_external_preagg_retaining_reaggregate_dimension_rolls_up( + self, + client_with_build_v3, + ): + """A semi-additive metric can read a fine pre-agg that still has the + protected dimension, then collapse in the final metrics query.""" + await _create_daily_balance_metric(client_with_build_v3) + await _register_external_preagg( + client_with_build_v3, + metrics=["v3.daily_balance"], + dimensions=[ + "v3.product.category", + "v3.date.date_id[order]", + ], + table_ref={ + "catalog": "default", + "schema": "analytics", + "table": "daily_balance_by_category_day", + "valid_through_ts": 20250101, + }, + measure_columns={"v3.daily_balance": "balance_sum"}, + dimension_columns={"v3.date.date_id[order]": "order_day"}, + table_columns={ + "category": "string", + "order_day": "int", + "balance_sum": "double", + }, + ) + + response = await client_with_build_v3.get( + "/sql/metrics/v3/", + params={ + "metrics": ["v3.daily_balance"], + "dimensions": ["v3.product.category"], + }, + ) + assert response.status_code == 200, response.json() + assert_sql_equal( + response.json()["sql"], + """ + WITH order_details_0 AS ( + SELECT category, + order_day date_id_order, + SUM(balance_sum) balance_sum + FROM default.analytics.daily_balance_by_category_day + GROUP BY category, order_day + ) + SELECT order_details_0.category AS category, + MAX_BY(order_details_0.balance_sum, order_details_0.date_id_order) + AS daily_balance + FROM order_details_0 + GROUP BY order_details_0.category + """, + ) + + @pytest.mark.asyncio + async def test_external_preagg_missing_reaggregate_dimension_is_not_used( + self, + client_with_build_v3, + ): + """A semi-additive metric must not read a pre-agg that dropped the key + required by the final collapse.""" + await _create_daily_balance_metric(client_with_build_v3) + await _register_external_preagg( + client_with_build_v3, + metrics=["v3.daily_balance"], + dimensions=["v3.product.category"], + table_ref={ + "catalog": "default", + "schema": "analytics", + "table": "daily_balance_by_category", + "valid_through_ts": 20250101, + }, + measure_columns={"v3.daily_balance": "balance_sum"}, + table_columns={"category": "string", "balance_sum": "double"}, + ) + + response = await client_with_build_v3.get( + "/sql/metrics/v3/", + params={ + "metrics": ["v3.daily_balance"], + "dimensions": ["v3.product.category"], + }, + ) + assert response.status_code == 200, response.json() + + sql = response.json()["sql"] + assert "default.analytics.daily_balance_by_category" not in sql + assert "MAX_BY(" in sql + assert "date_id_order" in sql + @pytest.mark.asyncio async def test_external_non_additive_not_rolled_up(self, client_with_build_v3): """A non-additive measure (COUNT DISTINCT) does not roll up to a coarser diff --git a/datajunction-server/tests/construction/build_v3/reaggregate_coverage_test.py b/datajunction-server/tests/construction/build_v3/reaggregate_coverage_test.py new file mode 100644 index 000000000..7b83a2126 --- /dev/null +++ b/datajunction-server/tests/construction/build_v3/reaggregate_coverage_test.py @@ -0,0 +1,846 @@ +"""Focused coverage tests for build-v3 reaggregation helpers.""" + +from types import SimpleNamespace + +import pytest + +from datajunction_server.construction.build_v3.cube_matcher import ( + _cube_dimension_covers_reaggregate_dimension, + _metric_graph_has_reaggregate, + _reaggregate_dimensions_for_cube_metrics, + _reaggregate_requirements_for_cube_metrics, + _reaggregate_requirements_for_metrics, + _reaggregate_requirements_for_decomposed_metrics, + _reaggregate_requirements_for_metrics_if_needed, + build_synthetic_grain_group, +) +from datajunction_server.construction.build_v3 import ( + cube_matcher as cube_matcher_module, +) +from datajunction_server.construction.build_v3 import builder as builder_module +from datajunction_server.construction.build_v3 import metrics as metrics_module +from datajunction_server.construction.build_v3.builder import setup_build_context +from datajunction_server.construction.build_v3.measures import ( + build_grain_group_sql, + build_select_ast, + build_window_metric_grain_groups, +) +from datajunction_server.construction.build_v3.metrics import ( + _build_reaggregate_collapse_expression, + _references_component, + _replace_reaggregate_merge_expression, + _metric_parent_refs, + build_window_agg_cte_from_base_metrics, + generate_metrics_sql, +) +from datajunction_server.construction.build_v3.types import ( + BuildContext, + ColumnMetadata, + DecomposedMetricInfo, + GeneratedMeasuresSQL, + GrainGroup, + GrainGroupSQL, + ResolvedDimension, +) +from datajunction_server.errors import DJInvalidInputException +from datajunction_server.models.decompose import ( + Aggregability, + AggregationRule, + MetricComponent, +) +from datajunction_server.models.dialect import Dialect +from datajunction_server.models.node import NodeType +from datajunction_server.models.reaggregate import ( + DimensionReaggregateRule, + ReaggregationFunction, +) +from datajunction_server.sql.parsing import ast +from datajunction_server.sql.parsing.backends.antlr4 import parse +from datajunction_server.sql.parsing.types import BigIntType, DoubleType, StringType + + +def _metric_node(name: str, query: str = "SELECT SUM(value) FROM test.parent"): + return SimpleNamespace( + name=name, + type=NodeType.METRIC, + current=SimpleNamespace( + query=query, + columns=[SimpleNamespace(type=DoubleType())], + ), + ) + + +def _source_node(name: str = "test.parent"): + return SimpleNamespace( + name=name, + type=NodeType.SOURCE, + current=SimpleNamespace( + catalog=SimpleNamespace(name="default"), + schema_="analytics", + table="parent", + columns=[ + SimpleNamespace(name="id", type=BigIntType()), + SimpleNamespace(name="value", type=DoubleType()), + SimpleNamespace(name="region", type=StringType()), + SimpleNamespace(name="date_id", type=BigIntType()), + ], + ), + ) + + +def _semi_additive_component( + name: str = "balance_sum", + dimension: str = "v3.date.date_id", + fn: ReaggregationFunction = ReaggregationFunction.LAST_VALUE, +) -> MetricComponent: + return MetricComponent( + name=name, + expression="value", + aggregation="SUM", + rule=AggregationRule( + type=Aggregability.FULL, + reaggregate=DimensionReaggregateRule( + dimension=dimension, + fn=fn, + ), + ), + ) + + +def _decomposed_metric( + metric_name: str, + component: MetricComponent | None = None, + query: str = "SELECT SUM(value) FROM test.parent", +) -> DecomposedMetricInfo: + return DecomposedMetricInfo( + metric_node=_metric_node(metric_name, query), + components=[component or _semi_additive_component()], + aggregability=Aggregability.FULL, + combiner=query, + derived_ast=parse(query), + ) + + +def test_reaggregate_collapse_expression_handles_min_max_and_rejects_unsupported(): + """MIN/MAX collapse ignore the protected dimension; unsupported functions fail.""" + value_ref = ast.Column(name=ast.Name("balance")) + protected_ref = ast.Column(name=ast.Name("date_id")) + + assert ( + str( + _build_reaggregate_collapse_expression( + ReaggregationFunction.MIN, + Dialect.SPARK, + value_ref, + protected_ref, + ), + ) + == "MIN(balance)" + ) + assert ( + str( + _build_reaggregate_collapse_expression( + ReaggregationFunction.MAX, + Dialect.SPARK, + value_ref, + protected_ref, + ), + ) + == "MAX(balance)" + ) + + with pytest.raises(DJInvalidInputException, match="Unsupported semi-additive"): + _build_reaggregate_collapse_expression( + ReaggregationFunction.SUM, + Dialect.SPARK, + value_ref, + protected_ref, + ) + + +def test_references_component_checks_all_columns_and_handles_no_match(): + """Component reference detection handles misses before matches.""" + expression = parse( + "SELECT SUM(other_value) + SUM(balance_sum) FROM base", + ).select.projection[0] + + assert _references_component(expression, "balance_sum") + assert not _references_component(expression, "missing_component") + + +def test_replace_reaggregate_merge_expression_skips_nonmatching_functions(): + """Combiner replacement skips unrelated functions before replacing the match.""" + combiner = parse( + "SELECT SUM(other_value) + SUM(balance_sum) FROM base", + ).select.projection[0] + collapse_expr = ast.Function( + ast.Name("MAX"), + args=[ast.Column(name=ast.Name("balance_sum"))], + ) + + rendered = str( + _replace_reaggregate_merge_expression( + combiner, + "balance_sum", + "SUM", + collapse_expr, + ), + ) + + assert "SUM(other_value)" in rendered + assert "MAX(balance_sum)" in rendered + + +def test_replace_reaggregate_merge_expression_rejects_missing_component(): + """Combiner replacement fails if the expected merge function is absent.""" + combiner = parse("SELECT AVG(balance_sum) FROM base").select.projection[0] + + with pytest.raises(DJInvalidInputException, match="could not find"): + _replace_reaggregate_merge_expression( + combiner, + "balance_sum", + "SUM", + ast.Function( + ast.Name("MAX"), + args=[ast.Column(name=ast.Name("balance_sum"))], + ), + ) + + +def test_replace_reaggregate_merge_expression_skips_parentless_nested_match(): + """A matching child without a parent is skipped before the final error.""" + + class ParentlessExpression(ast.Expression): + def __str__(self): + return "parentless" + + @property + def children(self): + yield ast.Function( + ast.Name("SUM"), + args=[ast.Column(name=ast.Name("balance_sum"))], + ) + + with pytest.raises(DJInvalidInputException, match="could not find"): + _replace_reaggregate_merge_expression( + ParentlessExpression(), + "balance_sum", + "SUM", + ast.Function( + ast.Name("MAX"), + args=[ast.Column(name=ast.Name("balance_sum"))], + ), + ) + + +def test_metric_parent_refs_adds_metric_refs_discovered_from_query(): + """Parsed-query metric references supplement the cached parent map.""" + balance = _metric_node("balance") + extra = _metric_node("extra_metric") + derived = _metric_node( + "derived", + "SELECT balance + extra_metric + raw_col FROM test.parent", + ) + ctx = BuildContext(session=SimpleNamespace(), metrics=[], dimensions=[]) + ctx.nodes = { + "balance": balance, + "extra_metric": extra, + "derived": derived, + } + ctx.parent_map = {"derived": ["balance"]} + + assert _metric_parent_refs(ctx, "derived") == ["balance", "extra_metric"] + + +def test_window_agg_from_base_metrics_collapses_semi_additive_derived_parent( + monkeypatch, +): + """Cross-fact window reaggregation replaces derived metric parents recursively.""" + monkeypatch.setattr(metrics_module, "get_column_full_name", lambda _col: "") + balance = _metric_node("balance") + extra = _metric_node("extra_metric") + derived = _metric_node( + "derived", + "SELECT balance + extra_metric + raw_col FROM test.parent", + ) + window = _metric_node("window_metric") + ctx = BuildContext( + session=SimpleNamespace(), + metrics=["window_metric"], + dimensions=["category"], + ) + ctx.nodes = { + "balance": balance, + "extra_metric": extra, + "derived": derived, + "window_metric": window, + } + ctx.parent_map = { + "window_metric": ["derived"], + "derived": ["balance"], + } + window_group = GrainGroupSQL( + query=parse("SELECT category FROM base_metrics"), + columns=[ + ColumnMetadata( + name="category", + semantic_name="category", + type="string", + semantic_type="dimension", + ), + ], + grain=["category"], + aggregability=Aggregability.FULL, + metrics=[], + parent_name="cross_fact", + is_window_grain_group=True, + window_metrics_served=["window_metric"], + ) + decomposed_metrics = { + "balance": _decomposed_metric( + "balance", + _semi_additive_component(dimension="v3.date.date_id[order]"), + ), + } + query = build_window_agg_cte_from_base_metrics( + window_group, + "base_metrics", + ctx, + {"v3.date.date_id[order]": "date_id_order"}, + set(), + decomposed_metrics, + ) + + rendered = str(query) + assert "MAX_BY(base_metrics.balance, base_metrics.date_id_order)" in rendered + assert "extra_metric" in rendered + assert "raw_col" in rendered + + with pytest.raises( + DJInvalidInputException, + match="protected dimension 'v3.date.date_id\\[order\\]' is not projected", + ): + build_window_agg_cte_from_base_metrics( + window_group, + "base_metrics", + ctx, + {}, + set(), + decomposed_metrics, + ) + + collapsed_query = build_window_agg_cte_from_base_metrics( + window_group, + "base_metrics", + ctx, + {}, + {"balance"}, + decomposed_metrics, + ) + assert "SUM(base_metrics.balance)" in str(collapsed_query) + + +@pytest.mark.parametrize("component_count", [1, 2]) +def test_window_agg_from_base_metrics_rejects_limited_leaf_metric( + component_count: int, +): + """Distinct metrics cannot be summed after their grain key is discarded.""" + visitor_count = _decomposed_metric( + "visitor_count", + MetricComponent( + name="visitor_count", + expression="customer_id", + aggregation="COUNT", + rule=AggregationRule( + type=Aggregability.LIMITED, + level=["customer_id"], + ), + ), + ) + if component_count == 2: + visitor_count.components.append( + MetricComponent( + name="order_total", + expression="order_amount", + aggregation="SUM", + rule=AggregationRule(type=Aggregability.FULL), + ), + ) + visitor_count.aggregability = Aggregability.LIMITED + window_group = GrainGroupSQL( + query=parse("SELECT category FROM base_metrics"), + columns=[ + ColumnMetadata( + name="category", + semantic_name="category", + type="string", + semantic_type="dimension", + ), + ], + grain=["category"], + aggregability=Aggregability.FULL, + metrics=[], + parent_name="cross_fact", + is_window_grain_group=True, + window_metrics_served=["window_metric"], + ) + ctx = BuildContext( + session=SimpleNamespace(), + metrics=["window_metric"], + dimensions=["category"], + ) + ctx.nodes = { + "visitor_count": visitor_count.metric_node, + "window_metric": _metric_node("window_metric"), + } + ctx.parent_map = {"window_metric": ["visitor_count"]} + + with pytest.raises( + DJInvalidInputException, + match="no longer retains the distinct grain key", + ): + build_window_agg_cte_from_base_metrics( + window_group, + "base_metrics", + ctx, + {}, + set(), + {"visitor_count": visitor_count}, + ) + + +def test_generate_metrics_sql_rejects_multiple_base_grain_groups_with_reaggregate(): + """Live SQL does not allow protected-dimension fanout across base groups.""" + ctx = BuildContext(session=SimpleNamespace(), metrics=["balance"], dimensions=[]) + grain_groups = [ + GrainGroupSQL( + query=parse("SELECT category, balance_sum FROM fact_a"), + columns=[], + grain=["category"], + aggregability=Aggregability.FULL, + metrics=["balance"], + parent_name="fact_a", + reaggregate_dimension_aliases={"balance_sum": "date_id"}, + ), + GrainGroupSQL( + query=parse("SELECT category, orders FROM fact_b"), + columns=[], + grain=["category"], + aggregability=Aggregability.FULL, + metrics=["orders"], + parent_name="fact_b", + ), + ] + measures_result = GeneratedMeasuresSQL( + grain_groups=grain_groups, + dialect=Dialect.SPARK, + requested_dimensions=[], + ctx=ctx, + ) + + with pytest.raises(DJInvalidInputException, match="multiple base grain groups"): + generate_metrics_sql(ctx, measures_result, {}) + + +def test_cube_reaggregate_requirements_skip_satisfied_dimensions_and_dedupe(): + """Cube requirement helpers skip requested protected dims and de-duplicate.""" + metric_revision = SimpleNamespace( + name="balance", + reaggregate={ + "rules": [ + { + "dimension": "v3.date.date_id[order]", + "fn": "last_value", + }, + ], + }, + ) + cube = SimpleNamespace(metric_node_revisions=lambda: [metric_revision]) + + assert ( + _reaggregate_requirements_for_cube_metrics( + cube, + ["balance"], + ["v3.date.date_id[order]"], + ) + == [] + ) + + duplicate_cube = SimpleNamespace( + metric_node_revisions=lambda: [metric_revision, metric_revision], + ) + assert _reaggregate_dimensions_for_cube_metrics( + duplicate_cube, + ["balance"], + ["v3.product.category"], + ) == ["v3.date.date_id[order]"] + + +def test_cube_dimension_coverage_accepts_bare_protected_parent_column(): + """Cube coverage accepts bare parent-column protected dimensions.""" + assert _cube_dimension_covers_reaggregate_dimension( + "order_date", + "v3.order_details.order_date", + ) + assert not _cube_dimension_covers_reaggregate_dimension( + "order_date", + "v3.order_details.ship_date", + ) + assert not _cube_dimension_covers_reaggregate_dimension( + "order_date", + "v3.order_details.order_date[ship]", + ) + assert not _cube_dimension_covers_reaggregate_dimension( + "v3.date.date_id", + "v3.order_details.order_date", + ) + + +@pytest.mark.asyncio +async def test_metric_graph_has_reaggregate_skips_empty_metric_list(): + """An empty metric list does not touch the database.""" + + class Session: + async def execute(self, _stmt): + raise AssertionError("empty metric graph should not query") + + assert not await _metric_graph_has_reaggregate(Session(), []) + + +@pytest.mark.asyncio +async def test_cube_reaggregate_requirements_skip_full_extract_without_reaggregate( + monkeypatch, +): + """No reaggregate in the metric graph skips full component extraction.""" + + async def fail_full_extract(*_args, **_kwargs): + raise AssertionError("full metric extraction should not run") + + class Result: + def all(self): + return [("v3.total_revenue", None, None)] + + class Session: + async def execute(self, _stmt): + return Result() + + monkeypatch.setattr( + cube_matcher_module, + "_reaggregate_requirements_for_metrics", + fail_full_extract, + ) + + assert ( + await _reaggregate_requirements_for_metrics_if_needed( + Session(), + ["v3.total_revenue"], + ["v3.product.category"], + ) + == [] + ) + + +@pytest.mark.asyncio +async def test_cube_reaggregate_requirements_checks_derived_metric_ancestors( + monkeypatch, +): + """Derived metrics still use full extraction when an ancestor has reaggregate.""" + expected = [("v3.balance_index", "v3.date.date_id", ReaggregationFunction.MAX)] + + async def full_extract(*_args, **_kwargs): + return expected + + class Result: + def __init__(self, rows): + self.rows = rows + + def all(self): + return self.rows + + class Session: + def __init__(self): + self.calls = 0 + + async def execute(self, _stmt): + self.calls += 1 + if self.calls == 1: + return Result([("v3.balance_index", None, "v3.daily_balance")]) + return Result([("v3.daily_balance", {"rules": []}, None)]) + + monkeypatch.setattr( + cube_matcher_module, + "_reaggregate_requirements_for_metrics", + full_extract, + ) + + assert ( + await _reaggregate_requirements_for_metrics_if_needed( + Session(), + ["v3.balance_index"], + ["v3.product.category"], + ) + == expected + ) + + +@pytest.mark.asyncio +async def test_reaggregate_requirements_skip_components_without_reaggregate( + monkeypatch, +): + """Full extraction ignores ordinary components before collecting reaggregate.""" + + class Extractor: + async def extract(self, _session): + return [ + MetricComponent( + name="ordinary_sum", + expression="value", + aggregation="SUM", + rule=AggregationRule(type=Aggregability.FULL), + ), + _semi_additive_component( + "balance_sum", + "v3.date.date_id", + ReaggregationFunction.MAX, + ), + ], None + + class FakeMetricComponentExtractor: + @classmethod + async def from_node_name(cls, _metric_name, _session): + return Extractor() + + monkeypatch.setattr( + cube_matcher_module, + "MetricComponentExtractor", + FakeMetricComponentExtractor, + ) + + assert await _reaggregate_requirements_for_metrics( + SimpleNamespace(), + ["balance"], + ["v3.product.category"], + ) == [("balance", "v3.date.date_id", ReaggregationFunction.MAX)] + + +def test_decomposed_reaggregate_requirements_dedupe_duplicate_components(): + """Duplicate component rules produce one cube materialization requirement.""" + rule = DimensionReaggregateRule( + dimension="v3.date.date_id", + fn=ReaggregationFunction.LAST_VALUE, + ) + components = [ + MetricComponent( + name=f"balance_sum_{idx}", + expression="value", + aggregation="SUM", + rule=AggregationRule(type=Aggregability.FULL, reaggregate=rule), + ) + for idx in range(2) + ] + decomposed = _decomposed_metric("balance", components[0]) + decomposed.components = components + + assert _reaggregate_requirements_for_decomposed_metrics( + {"balance": decomposed}, + ["balance"], + [], + ) == [("balance", "v3.date.date_id", ReaggregationFunction.LAST_VALUE)] + + +def test_build_synthetic_grain_group_skips_duplicate_internal_dimension(): + """Filter-only protected dims can already be present in ctx.dimensions.""" + ctx = BuildContext( + session=SimpleNamespace(), + metrics=["balance"], + dimensions=["v3.date.date_id"], + dialect=Dialect.DRUID, + ) + ctx.filter_dimensions = {"v3.date.date_id"} + ctx.nodes = {"balance": _metric_node("balance")} + ctx.parent_map = {"balance": ["test.parent"]} + cube = SimpleNamespace( + name="daily_balance_cube", + availability=SimpleNamespace(table="daily_balance_cube"), + materializations=[], + ) + + grain_group = build_synthetic_grain_group( + ctx, + {"balance": _decomposed_metric("balance")}, + cube, + ) + + assert grain_group.grain == ["date_id"] + assert [ + col.name for col in grain_group.columns if col.semantic_type == "dimension" + ] == [ + "date_id", + ] + + +def test_build_select_ast_skips_non_output_dimensions(): + """Resolved dimensions outside output refs are join/group-only.""" + ctx = BuildContext(session=SimpleNamespace(), metrics=[], dimensions=[]) + parent = _source_node() + resolved_dimensions = [ + ResolvedDimension( + original_ref="test.parent.region", + node_name="test.parent", + column_name="region", + role=None, + join_path=None, + is_local=True, + ), + ResolvedDimension( + original_ref="test.parent.date_id", + node_name="test.parent", + column_name="date_id", + role=None, + join_path=None, + is_local=True, + ), + ] + + query, _ = build_select_ast( + ctx, + metric_expressions=[], + resolved_dimensions=resolved_dimensions, + parent_node=parent, + output_dimension_refs={"test.parent.region"}, + ) + + rendered = str(query) + assert "t1.region" in rendered + assert "t1.date_id AS date_id" not in rendered + + +def test_build_grain_group_sql_handles_internal_alias_already_in_grain(): + """The internal reaggregate alias can match a requested dimension alias.""" + ctx = BuildContext( + session=SimpleNamespace(), + metrics=["balance"], + dimensions=["test.parent.date_id"], + use_materialized=False, + ) + parent = _source_node() + metric = _metric_node("balance") + component = _semi_additive_component(dimension="test.parent.date_id") + resolved_dimension = ResolvedDimension( + original_ref="test.parent.date_id", + node_name="test.parent", + column_name="date_id", + role=None, + join_path=None, + is_local=True, + ) + grain_group = GrainGroup( + parent_node=parent, + aggregability=Aggregability.FULL, + grain_columns=[], + components=[(metric, component)], + reaggregate_component_dimensions={component.name: "test.parent.date_id"}, + ) + + result = build_grain_group_sql( + ctx, + grain_group, + [resolved_dimension], + {"balance": 1}, + output_dimension_refs={"test.parent.date_id"}, + ) + + assert result.grain == ["date_id"] + + +def test_build_window_metric_grain_groups_skips_missing_base_components(): + """Window grain planning skips parent groups with no decomposed components.""" + ctx = BuildContext(session=SimpleNamespace(), metrics=["window"], dimensions=[]) + ctx.parent_map = {"window": ["missing_base"]} + existing_group = GrainGroupSQL( + query=parse("SELECT missing_base FROM test_parent_0"), + columns=[], + grain=[], + aggregability=Aggregability.FULL, + metrics=["missing_base"], + parent_name="test.parent", + ) + + assert ( + build_window_metric_grain_groups( + ctx, + {"window": {"v3.date.week"}}, + [existing_group], + {}, + ) + == [] + ) + + +def test_build_window_metric_grain_groups_skips_when_parent_node_missing(): + """Window grain planning skips groups whose source parent is not loaded.""" + ctx = BuildContext(session=SimpleNamespace(), metrics=["window"], dimensions=[]) + ctx.nodes = {"missing_base": _metric_node("missing_base")} + ctx.parent_map = {"window": ["missing_base"]} + existing_group = GrainGroupSQL( + query=parse("SELECT missing_base FROM missing_parent_0"), + columns=[], + grain=[], + aggregability=Aggregability.FULL, + metrics=["missing_base"], + parent_name="missing.parent", + ) + + assert ( + build_window_metric_grain_groups( + ctx, + {"window": {"v3.date.week"}}, + [existing_group], + { + "window": _decomposed_metric("window"), + "missing_base": _decomposed_metric("missing_base"), + }, + ) + == [] + ) + + +@pytest.mark.asyncio +async def test_setup_build_context_skips_duplicate_internal_reaggregate_dimension( + monkeypatch, +): + """setup_build_context does not duplicate internal reaggregate dimensions.""" + load_calls = [] + + async def fake_load_nodes(ctx): + load_calls.append(list(ctx.dimensions)) + ctx.nodes = {"balance": _metric_node("balance")} + + async def fake_decompose_and_group_metrics(ctx): + return [], {"balance": _decomposed_metric("balance")} + + monkeypatch.setattr(builder_module, "load_nodes", fake_load_nodes) + monkeypatch.setattr( + builder_module, + "decompose_and_group_metrics", + fake_decompose_and_group_metrics, + ) + monkeypatch.setattr( + builder_module, + "add_dimensions_from_metric_expressions", + lambda _ctx, _decomposed_metrics: None, + ) + monkeypatch.setattr( + builder_module, + "missing_reaggregate_dimensions", + lambda _decomposed_metrics, _requested_dimensions: ["v3.date.date_id"], + ) + + ctx = await setup_build_context( + session=SimpleNamespace(), + metrics=["balance"], + dimensions=["v3.date.date_id"], + ) + + assert ctx.dimensions == ["v3.date.date_id"] + assert load_calls == [["v3.date.date_id"], ["v3.date.date_id"]] diff --git a/datajunction-server/tests/internal/deployment/orchestration_test.py b/datajunction-server/tests/internal/deployment/orchestration_test.py index 5effcbd1e..f1f1a00fa 100644 --- a/datajunction-server/tests/internal/deployment/orchestration_test.py +++ b/datajunction-server/tests/internal/deployment/orchestration_test.py @@ -30,6 +30,7 @@ from datajunction_server.internal.deployment.orchestrator import ( DeploymentOrchestrator, DeploymentPlan, + DeploymentTimer, ResourceRegistry, column_changed, tag_needs_update, @@ -86,6 +87,26 @@ def mock_deployment_context(current_user: User): return context +def test_deployment_timer_logs_unaccounted_overhead(monkeypatch, caplog): + """ + Deployment timing summary includes unaccounted overhead when it is meaningful. + """ + timer = DeploymentTimer() + timer.record("validate", 10, "1 node") + monkeypatch.setattr( + "datajunction_server.internal.deployment.orchestrator.time.perf_counter", + lambda: timer._start + 0.2, + ) + + with caplog.at_level( + "INFO", + logger="datajunction_server.internal.deployment.orchestrator", + ): + timer.log_summary("default", "deployment-1") + + assert "(unaccounted overhead)" in caplog.text + + @pytest.fixture def sample_deployment_spec(): """Sample deployment specification for testing""" diff --git a/datajunction-server/tests/internal/deployment/validation_test.py b/datajunction-server/tests/internal/deployment/validation_test.py index ebc6cda33..54fc181eb 100644 --- a/datajunction-server/tests/internal/deployment/validation_test.py +++ b/datajunction-server/tests/internal/deployment/validation_test.py @@ -824,6 +824,132 @@ async def test_invalid_short_name_not_in_parent_columns( assert err.debug is not None assert "no_such_col" in err.debug["invalid_required_dimensions"] + @pytest.mark.asyncio + async def test_valid_reaggregate_dimension_full_path( + self, + session: AsyncSession, + parent_node: Node, + ): + """reaggregate.rules[].dimension full-path column found in dim nodes is valid.""" + dim_node = self._make_dim_node("test.dim", ["dateint"]) + context = self._make_context(session, parent_node) + spec = MetricSpec( + name="test.metric", + query="SELECT SUM(value) FROM test.parent", + reaggregate={ + "rules": [ + { + "dimension": "test.dim.dateint", + "fn": "last_value", + }, + ], + }, + ) + validator = NodeSpecBulkValidator(context) + validator._all_dim_nodes = {**context.dependency_nodes, "test.dim": dim_node} + result = validator.validate_query_node(spec) + + assert result.status == NodeStatus.VALID + error_codes = [e.code for e in result.errors] + assert ErrorCode.INVALID_COLUMN not in error_codes + + @pytest.mark.asyncio + async def test_reaggregate_on_derived_metric_is_invalid( + self, + session: AsyncSession, + parent_node: Node, + ): + """Deployment validation rejects metric-level policy on derived metrics.""" + context = self._make_context(session, parent_node) + spec = MetricSpec( + name="test.metric", + query="SELECT test.base_metric * 2", + reaggregate={ + "rules": [ + { + "dimension": "test.date.date_id", + "fn": "last_value", + }, + ], + }, + ) + validator = NodeSpecBulkValidator(context) + + result = validator.validate_query_node(spec) + + assert result.status == NodeStatus.INVALID + error = next( + error for error in result.errors if error.code == ErrorCode.INVALID_METRIC + ) + assert "only supported on base metrics" in error.message + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "dimension", + ["test.dim.nonexistent_col", "value"], + ) + async def test_invalid_reaggregate_dimension_column( + self, + session: AsyncSession, + parent_node: Node, + dimension: str, + ): + """Reaggregate dimensions must be qualified and resolve to a column.""" + dim_node = self._make_dim_node("test.dim", ["dateint"]) + context = self._make_context(session, parent_node) + spec = MetricSpec( + name="test.metric", + query="SELECT SUM(value) FROM test.parent", + reaggregate={ + "rules": [ + { + "dimension": dimension, + "fn": "last_value", + }, + ], + }, + ) + validator = NodeSpecBulkValidator(context) + validator._all_dim_nodes = {**context.dependency_nodes, "test.dim": dim_node} + result = validator.validate_query_node(spec) + + assert result.status == NodeStatus.INVALID + err = next(e for e in result.errors if e.code == ErrorCode.INVALID_COLUMN) + assert err.debug is not None + assert dimension in err.debug["invalid_reaggregate_dimensions"] + + @pytest.mark.asyncio + async def test_invalid_reaggregate_function( + self, + session: AsyncSession, + parent_node: Node, + ): + """unsupported dimension-specific reaggregate functions are invalid.""" + context = self._make_context(session, parent_node) + spec = MetricSpec( + name="test.metric", + query="SELECT SUM(value) FROM test.parent", + reaggregate={ + "rules": [ + { + "dimension": "test.parent.id", + "fn": "sum", + }, + ], + }, + ) + validator = NodeSpecBulkValidator(context) + validator._all_dim_nodes = dict(context.dependency_nodes) + result = validator.validate_query_node(spec) + + assert result.status == NodeStatus.INVALID + err = next( + e + for e in result.errors + if e.code == ErrorCode.INVALID_ARGUMENTS_TO_FUNCTION + ) + assert err.debug == {"invalid_reaggregate_functions": ["sum"]} + @pytest.mark.asyncio async def test_no_required_dimensions_is_noop( self, @@ -897,6 +1023,64 @@ async def test_prefetch_fetches_dim_node_from_db( col_names = {c.name for c in fetched.current.columns} assert "region" in col_names + @pytest.mark.asyncio + async def test_prefetch_fetches_reaggregate_dim_node_from_db( + self, + session: AsyncSession, + parent_node: Node, + dim_node_in_db: Node, + ): + """_prefetch_required_dimension_nodes includes reaggregate dimensions.""" + context = self._make_context(session, parent_node) + spec = MetricSpec( + name="test.metric", + query="SELECT SUM(value) FROM test.parent", + reaggregate={ + "rules": [ + { + "dimension": "test.external_dim.region", + "fn": "last_value", + }, + ], + }, + ) + validator = NodeSpecBulkValidator(context) + + await validator._prefetch_required_dimension_nodes([spec]) + + assert "test.external_dim" in validator._all_dim_nodes + fetched = validator._all_dim_nodes["test.external_dim"] + assert fetched.current is not None + assert {c.name for c in fetched.current.columns} == {"region"} + + @pytest.mark.asyncio + async def test_prefetch_skips_short_reaggregate_dimension( + self, + session: AsyncSession, + parent_node: Node, + ): + """Short-form reaggregate dimensions do not trigger dimension-node prefetch.""" + context = self._make_context(session, parent_node) + spec = MetricSpec( + name="test.metric", + query="SELECT SUM(value) FROM test.parent", + reaggregate={ + "rules": [ + { + "dimension": "value", + "fn": "last_value", + }, + ], + }, + ) + validator = NodeSpecBulkValidator(context) + + await validator._prefetch_required_dimension_nodes([spec]) + + assert set(validator._all_dim_nodes.keys()) == set( + context.dependency_nodes.keys(), + ) + @pytest.mark.asyncio async def test_prefetch_with_no_required_dimensions( self, diff --git a/datajunction-server/tests/internal/namespaces_test.py b/datajunction-server/tests/internal/namespaces_test.py index 1c563211d..839f99086 100644 --- a/datajunction-server/tests/internal/namespaces_test.py +++ b/datajunction-server/tests/internal/namespaces_test.py @@ -3,6 +3,7 @@ """ from datetime import UTC, datetime +from types import SimpleNamespace import pytest from ruamel.yaml import YAML @@ -22,6 +23,7 @@ _merge_list_with_key, _merge_yaml_preserving_comments, create_or_reactivate_namespace, + get_node_specs_for_export, node_spec_to_yaml, provision_namespace_boundary, ) @@ -1108,6 +1110,81 @@ def test_no_yaml_document_start_marker(self): "query: SELECT SUM(rev) FROM ns.transforms.t", ] + def test_metric_reaggregate_exports_dimension_rules(self): + """metric reaggregate exports the rule list without null future fields""" + spec = MetricSpec( + name="ns.metrics.daily_balance", + node_type=NodeType.METRIC, + query="SELECT SUM(balance) FROM ns.transforms.accounts", + reaggregate={ + "rules": [ + { + "dimension": "${prefix}v3.date.date_id[order]", + "fn": "last_value", + }, + ], + }, + ) + assert node_spec_to_yaml(spec).splitlines() == [ + "name: ns.metrics.daily_balance", + "node_type: metric", + "mode: published", + "query: SELECT SUM(balance) FROM ns.transforms.accounts", + "reaggregate:", + " rules:", + " - dimension: ${prefix}v3.date.date_id[order]", + " fn: last_value", + ] + + @pytest.mark.asyncio + async def test_metric_reaggregate_export_injects_dimension_prefix( + self, + monkeypatch, + ): + """namespace export parameterizes reaggregate rule dimensions.""" + spec = MetricSpec( + name="demo.branch.metrics.daily_balance", + node_type=NodeType.METRIC, + query="SELECT SUM(balance) FROM demo.branch.transforms.accounts", + reaggregate={ + "rules": [ + { + "dimension": "demo.branch.dims.date.date_id[order]", + "fn": "last_value", + }, + ], + }, + ) + + async def to_spec(_session): + return spec + + fake_node = SimpleNamespace( + name="demo.branch.metrics.daily_balance", + current=SimpleNamespace(parents=[], status="valid"), + to_spec=to_spec, + ) + + async def namespace_get(_session, _namespace, raise_if_not_exists=False): + return SimpleNamespace(parent_namespace="demo.main") + + async def list_all_nodes(_session, _namespace, options=None): + return [fake_node] + + monkeypatch.setattr(NodeNamespace, "get", namespace_get) + monkeypatch.setattr(NodeNamespace, "list_all_nodes", list_all_nodes) + + exported = await get_node_specs_for_export(SimpleNamespace(), "demo.branch") + + assert exported[0].name == "${prefix}metrics.daily_balance" + assert exported[0].query == ( + "SELECT SUM(balance) FROM ${prefix}transforms.accounts" + ) + assert ( + exported[0].reaggregate.rules[0].dimension + == "${prefix}dims.date.date_id[order]" + ) + def test_multiline_query_uses_literal_block_style(self): """multiline queries are serialized with |- literal block style""" spec = MetricSpec( diff --git a/datajunction-server/tests/internal/node_validation_test.py b/datajunction-server/tests/internal/node_validation_test.py index 5711a6ea7..42503a60a 100644 --- a/datajunction-server/tests/internal/node_validation_test.py +++ b/datajunction-server/tests/internal/node_validation_test.py @@ -1019,6 +1019,195 @@ async def test_validate_node_data_v2_flags_invalid_required_dimensions( ), [(e.code, e.message) for e in validator.errors] +@pytest.mark.asyncio +@pytest.mark.parametrize( + "dimension", + ["test.v2_reagg_dim.ghost_col", "id"], +) +async def test_validate_node_data_v2_flags_invalid_reaggregate_dimensions( + session: AsyncSession, + user: User, + dimension: str, +): + """reaggregate dimensions must be qualified and resolve to parent columns.""" + from datajunction_server.errors import ErrorCode + from datajunction_server.internal.validation import validate_node_data_v2 + + source = Node( + name="test.v2_reagg_dim_parent", + type=NodeType.SOURCE, + created_by_id=user.id, + current_version="v1.0", + ) + source_rev = NodeRevision( + name="test.v2_reagg_dim_parent", + display_name="reaggregate dim parent", + type=NodeType.SOURCE, + query=None, + status=NodeStatus.VALID, + version="v1.0", + node=source, + columns=[Column(name="id", type=ct.BigIntType(), order=0)], + created_by_id=user.id, + ) + dim = Node( + name="test.v2_reagg_dim", + type=NodeType.DIMENSION, + created_by_id=user.id, + current_version="v1.0", + ) + dim_rev = NodeRevision( + name="test.v2_reagg_dim", + display_name="tiny reaggregate dim", + type=NodeType.DIMENSION, + query="SELECT 1 AS id", + status=NodeStatus.VALID, + version="v1.0", + node=dim, + columns=[Column(name="id", type=ct.BigIntType(), order=0)], + created_by_id=user.id, + ) + session.add_all([source, source_rev, dim, dim_rev]) + await session.commit() + + child = NodeRevision( + name="test.v2_reagg_dim_child", + display_name="reaggregate dim child", + type=NodeType.METRIC, + query="SELECT SUM(id) FROM test.v2_reagg_dim_parent", + status=NodeStatus.VALID, + reaggregate={ + "rules": [ + { + "dimension": dimension, + "fn": "last_value", + }, + ], + }, + ) + validator = await validate_node_data_v2(child, session) + + assert validator.status == NodeStatus.INVALID + assert any( + err.code == ErrorCode.INVALID_COLUMN + and "reaggregate dimensions" in err.message + and dimension in err.debug["invalid_reaggregate_dimensions"] + for err in validator.errors + ), [(e.code, e.message) for e in validator.errors] + + +@pytest.mark.asyncio +async def test_validate_node_data_v2_rejects_reaggregate_on_derived_metric( + session: AsyncSession, + user: User, +): + """Derived metrics cannot declare their own reaggregation policy.""" + from datajunction_server.errors import ErrorCode + from datajunction_server.internal.validation import validate_node_data_v2 + + base = Node( + name="test.v2_base_balance", + type=NodeType.METRIC, + created_by_id=user.id, + current_version="v1.0", + ) + base_revision = NodeRevision( + name=base.name, + display_name="base balance", + type=NodeType.METRIC, + query="SELECT SUM(balance) FROM test.balance_source", + status=NodeStatus.VALID, + version="v1.0", + node=base, + columns=[Column(name=base.name, type=ct.DoubleType(), order=0)], + created_by_id=user.id, + ) + session.add_all([base, base_revision]) + await session.commit() + + derived = NodeRevision( + name="test.v2_double_balance", + display_name="double balance", + type=NodeType.METRIC, + query="SELECT test.v2_base_balance * 2", + status=NodeStatus.VALID, + reaggregate={ + "rules": [ + { + "dimension": "test.date.date_id", + "fn": "last_value", + }, + ], + }, + ) + + validator = await validate_node_data_v2(derived, session) + + assert validator.status == NodeStatus.INVALID + assert any( + error.code == ErrorCode.INVALID_METRIC + and "only supported on base metrics" in error.message + for error in validator.errors + ) + + +@pytest.mark.asyncio +async def test_validate_node_data_v2_flags_unsupported_reaggregate_function( + session: AsyncSession, + user: User, +): + """dimension-specific reaggregate only accepts supported collapse functions.""" + from datajunction_server.errors import ErrorCode + from datajunction_server.internal.validation import validate_node_data_v2 + + source = Node( + name="test.v2_reagg_fn_parent", + type=NodeType.SOURCE, + created_by_id=user.id, + current_version="v1.0", + ) + source_rev = NodeRevision( + name="test.v2_reagg_fn_parent", + display_name="reaggregate function parent", + type=NodeType.SOURCE, + query=None, + status=NodeStatus.VALID, + version="v1.0", + node=source, + columns=[ + Column(name="id", type=ct.BigIntType(), order=0), + Column(name="order_date", type=ct.BigIntType(), order=1), + ], + created_by_id=user.id, + ) + session.add_all([source, source_rev]) + await session.commit() + + child = NodeRevision( + name="test.v2_reagg_fn_child", + display_name="reaggregate function child", + type=NodeType.METRIC, + query="SELECT SUM(id) FROM test.v2_reagg_fn_parent", + status=NodeStatus.VALID, + reaggregate={ + "rules": [ + { + "dimension": "test.v2_reagg_fn_parent.order_date", + "fn": "sum", + }, + ], + }, + ) + validator = await validate_node_data_v2(child, session) + + assert validator.status == NodeStatus.INVALID + assert any( + err.code == ErrorCode.INVALID_ARGUMENTS_TO_FUNCTION + and err.debug == {"invalid_reaggregate_functions": ["sum"]} + for err in validator.errors + ), [(e.code, e.message, e.debug) for e in validator.errors] + + @pytest.mark.asyncio async def test_validate_node_data_v2_cross_fact_metrics_no_shared_dims( session: AsyncSession, diff --git a/datajunction-server/tests/internal/nodes/derive_frozen_measures_test.py b/datajunction-server/tests/internal/nodes/derive_frozen_measures_test.py index 292be4521..9b8754c97 100644 --- a/datajunction-server/tests/internal/nodes/derive_frozen_measures_test.py +++ b/datajunction-server/tests/internal/nodes/derive_frozen_measures_test.py @@ -1,9 +1,6 @@ """ -Unit tests for ``derive_frozen_measures_bulk`` — the batched derivation path -used by the deployment orchestrator. Exercises the cache-construction branches -(derived-metric expansion, deep-chain iterative expansion) and edge cases -(empty list, shared measures across metrics) that aren't reachable through -the per-metric ``derive_frozen_measures`` entry point. +Unit tests for frozen-measure derivation. Exercises both the per-metric path used +by direct node creation and the batched path used by the deployment orchestrator. """ import pytest @@ -16,7 +13,17 @@ from datajunction_server.database.measure import FrozenMeasure from datajunction_server.database.node import Node, NodeRevision from datajunction_server.database.user import OAuthProvider, User -from datajunction_server.internal.nodes import derive_frozen_measures_bulk +from datajunction_server.errors import DJInvalidInputException +from datajunction_server.internal.nodes import ( + _derive_frozen_measures_impl, + _raise_if_frozen_measure_conflicts, + derive_frozen_measures_bulk, +) +from datajunction_server.models.decompose import ( + Aggregability, + AggregationRule, + MetricComponent, +) from datajunction_server.models.node import NodeStatus from datajunction_server.models.node_type import NodeType @@ -62,6 +69,7 @@ async def _make_metric( name: str, query: str, parents: list[Node], + reaggregate: dict[str, object] | None = None, ) -> Node: node = Node( name=name, @@ -75,6 +83,7 @@ async def _make_metric( type=NodeType.METRIC, version="v1.0", query=query, + reaggregate=reaggregate, status=NodeStatus.VALID, parents=parents, created_by_id=user.id, @@ -123,6 +132,153 @@ async def test_base_metric_populates_derived_expression_and_measure( assert any(fm.aggregation == "SUM" for fm in metric.current.frozen_measures) +@pytest.mark.asyncio +async def test_frozen_measure_name_collision_allows_reaggregate_only_difference( + session: AsyncSession, + user: User, +): + """Metric-level reaggregate metadata does not change measure identity.""" + src = await _make_source( + session, + user, + "src_reaggregate_collision", + [Column(name="semi_amount", type=ct.DoubleType(), order=0)], + ) + reaggregate = await _make_metric( + session, + user, + "m.semi_amount_eod", + "SELECT SUM(semi_amount) FROM src_reaggregate_collision", + [src], + reaggregate={ + "rules": [ + { + "dimension": "default.date_dim.date", + "fn": "last_value", + }, + ], + }, + ) + additive = await _make_metric( + session, + user, + "m.semi_amount_total", + "SELECT SUM(semi_amount) FROM src_reaggregate_collision", + [src], + ) + await session.refresh(reaggregate, ["current"]) + await session.refresh(additive, ["current"]) + + await derive_frozen_measures_bulk(session, [additive.current.id]) + await session.commit() + + await session.refresh(additive.current, ["frozen_measures"]) + assert all(fm.rule.reaggregate is None for fm in additive.current.frozen_measures) + + await derive_frozen_measures_bulk(session, [reaggregate.current.id]) + await session.commit() + await session.refresh(reaggregate.current, ["frozen_measures"]) + + assert {fm.name for fm in additive.current.frozen_measures} == { + fm.name for fm in reaggregate.current.frozen_measures + } + + +@pytest.mark.asyncio +async def test_reaggregate_rule_is_not_persisted_on_shared_frozen_measure( + session: AsyncSession, + user: User, +): + """FrozenMeasure.rule stays metric-independent even for semi-additive metrics.""" + src = await _make_source( + session, + user, + "src_reaggregate_storage", + [Column(name="semi_amount", type=ct.DoubleType(), order=0)], + ) + reaggregate = await _make_metric( + session, + user, + "m.semi_amount_snapshot", + "SELECT SUM(semi_amount) FROM src_reaggregate_storage", + [src], + reaggregate={ + "rules": [ + { + "dimension": "default.date_dim.date", + "fn": "last_value", + }, + ], + }, + ) + await session.refresh(reaggregate, ["current"]) + + await derive_frozen_measures_bulk(session, [reaggregate.current.id]) + await session.commit() + + await session.refresh(reaggregate.current, ["frozen_measures"]) + assert all( + fm.rule.reaggregate is None for fm in reaggregate.current.frozen_measures + ) + + +@pytest.mark.asyncio +async def test_direct_derivation_does_not_persist_reaggregate_rule( + session: AsyncSession, + user: User, +): + """Direct metric creation persists only metric-independent measure rules.""" + src = await _make_source( + session, + user, + "src_direct_reaggregate_storage", + [Column(name="balance", type=ct.DoubleType(), order=0)], + ) + metric = await _make_metric( + session, + user, + "m.current_balance", + "SELECT SUM(balance) FROM src_direct_reaggregate_storage", + [src], + reaggregate={ + "rules": [ + { + "dimension": "default.date_dim.date", + "fn": "last_value", + }, + ], + }, + ) + await session.refresh(metric, ["current"]) + + frozen_measures = await _derive_frozen_measures_impl(metric.current.id, session) + await session.commit() + + assert metric.current.reaggregate is not None + assert frozen_measures + assert all(fm.rule.reaggregate is None for fm in frozen_measures) + + +def test_frozen_measure_conflict_rejects_different_measure_identity(): + """FrozenMeasure name collisions fail when the metric-independent rule differs.""" + frozen_measure = FrozenMeasure( + name="amount_sum", + upstream_revision_id=1, + expression="amount", + aggregation="SUM", + rule=AggregationRule(type=Aggregability.FULL), + ) + measure = MetricComponent( + name="amount_sum", + expression="discounted_amount", + aggregation="SUM", + rule=AggregationRule(type=Aggregability.FULL), + ) + + with pytest.raises(DJInvalidInputException, match="already exists"): + _raise_if_frozen_measure_conflicts(frozen_measure, measure) + + @pytest.mark.asyncio async def test_derived_metric_expands_parent_cache( session: AsyncSession, diff --git a/datajunction-server/tests/models/deployment_test.py b/datajunction-server/tests/models/deployment_test.py index 78c47bf34..ea9347970 100644 --- a/datajunction-server/tests/models/deployment_test.py +++ b/datajunction-server/tests/models/deployment_test.py @@ -1956,7 +1956,7 @@ def fingerprint(spec: NodeSpec) -> SemanticFingerprint: "source": "71dcbc388988c2bdd850670427710384687b58565ee38ca392dc220adfed868d", "transform": "978e692880c7bcfb1bd78ece85895a1ec1558e85377b064a8dac3f3719cff2a5", "dimension": "f7b3c87a61fdadf9997432fd9334befdf43f2488874ef555e3f7d4c4ba86e3e1", - "metric": "a3fde7af5dbd00d194805af33fc213bca52244c7cdedf1f7363ec52d2f6d4116", + "metric": "f0b5356d75b1f6a39a2809581a9167b9e1298a41980ee8e83040cf6bd0996851", "cube": "9b0a56d974d1e3769bc2db94e2cfbae7a6a4839f664eebd2a4387ef112ceea81", } @@ -2103,6 +2103,38 @@ def test_semantic_fingerprint_renders_prefixes_and_normalizes_sql(): assert fingerprint(explicit) != fingerprint(implicit) +def test_metric_spec_equality_compares_rendered_reaggregate(): + """Parameterized reaggregate dimensions should not cause perpetual deploys.""" + parameterized = MetricSpec( + namespace="analytics", + name="daily_balance", + query="SELECT SUM(balance) FROM analytics.daily_balances", + reaggregate={ + "rules": [ + { + "dimension": "${prefix}date.date_id[order]", + "fn": "last_value", + }, + ], + }, + ) + rendered = MetricSpec( + namespace="analytics", + name="daily_balance", + query="SELECT SUM(balance) FROM analytics.daily_balances", + reaggregate={ + "rules": [ + { + "dimension": "analytics.date.date_id[order]", + "fn": "last_value", + }, + ], + }, + ) + + assert parameterized == rendered + + def test_semantic_diff_and_fingerprint_share_change_rules(): original = TransformSpec(name="node", query="SELECT id AS value FROM source") formatted = TransformSpec( @@ -2256,6 +2288,7 @@ def test_metric_presentation_fields_preserve_semantic_fingerprint(): "query", "columns", "required_dimensions", + "reaggregate", } ) assert all( diff --git a/datajunction-server/tests/models/node_test.py b/datajunction-server/tests/models/node_test.py index 1489d3928..3a630811c 100644 --- a/datajunction-server/tests/models/node_test.py +++ b/datajunction-server/tests/models/node_test.py @@ -227,6 +227,30 @@ def test_extra_validation() -> None: "bound dimensions which are only for metrics." ) + node = Node(name="A", type=NodeType.TRANSFORM, current_version="1") + node_revision = NodeRevision( + name=node.name, + type=node.type, + node=node, + version="1", + query="SELECT * FROM B", + reaggregate={ + "rules": [ + { + "dimension": "B.date_id", + "fn": "last_value", + }, + ], + }, + ) + with pytest.raises(Exception) as excinfo: + node_revision.extra_validation() + + assert str(excinfo.value) == ( + "Node A of type transform cannot have " + "reaggregate settings which are only for metrics." + ) + def test_merging_availability_simple_no_partitions() -> None: """ diff --git a/datajunction-server/tests/models/reaggregate_test.py b/datajunction-server/tests/models/reaggregate_test.py new file mode 100644 index 000000000..60cc67b9f --- /dev/null +++ b/datajunction-server/tests/models/reaggregate_test.py @@ -0,0 +1,133 @@ +"""Tests for reaggregation models.""" + +import pytest +from pydantic import ValidationError + +from datajunction_server.models.reaggregate import ( + DimensionReaggregateRule, + ReaggregateSpec, + ReaggregationFunction, + dimension_reaggregate_rules, + dump_reaggregate_spec, + parse_reaggregate_spec, + unsupported_dimension_reaggregate_functions, +) + + +def test_dump_reaggregate_spec_from_dict(): + """ + Reaggregate specs passed as dictionaries are validated and serialized. + """ + assert dump_reaggregate_spec( + { + "rules": [ + { + "dimension": "default.date_dim.date", + "fn": "last_value", + }, + ], + }, + ) == { + "rules": [ + { + "dimension": "default.date_dim.date", + "fn": "last_value", + }, + ], + } + + +def test_dump_reaggregate_spec_from_model(): + """ + Reaggregate spec models are serialized to JSON-compatible dictionaries. + """ + assert dump_reaggregate_spec( + ReaggregateSpec( + rules=[ + DimensionReaggregateRule( + dimension="default.date_dim.date", + fn=ReaggregationFunction.LAST_VALUE, + ), + ], + ), + ) == { + "rules": [ + { + "dimension": "default.date_dim.date", + "fn": "last_value", + }, + ], + } + + +def test_dump_reaggregate_spec_none(): + """ + Missing reaggregate specs pass through unchanged. + """ + assert dump_reaggregate_spec(None) is None + + +def test_parse_reaggregate_spec_from_model(): + """ + Parsed model inputs are returned unchanged. + """ + spec = ReaggregateSpec( + rules=[ + DimensionReaggregateRule( + dimension="default.date_dim.date", + fn=ReaggregationFunction.LAST_VALUE, + ), + ], + ) + + assert parse_reaggregate_spec(spec) is spec + + +@pytest.mark.parametrize( + "unknown_field", + [ + {"fn": "last_value"}, + {"weight": "default.orders.quantity"}, + { + "rules": [ + { + "dimension": "default.date_dim.date", + "fn": "last_value", + "weight": "default.orders.quantity", + }, + ], + }, + ], +) +def test_parse_reaggregate_spec_rejects_unknown_fields(unknown_field: dict): + """Unknown fields are rejected at both reaggregate model levels.""" + with pytest.raises(ValidationError): + parse_reaggregate_spec(unknown_field) + + +def test_dimension_reaggregate_rules_empty_for_none(): + """ + Empty reaggregate specs have no dimension-specific rules. + """ + assert dimension_reaggregate_rules(None) == [] + + +def test_unsupported_dimension_reaggregate_functions_handles_empty_and_invalid(): + """ + Unsupported dimension collapse functions are reported from parsed specs. + """ + assert unsupported_dimension_reaggregate_functions(None) == [] + assert unsupported_dimension_reaggregate_functions( + { + "rules": [ + { + "dimension": "default.date_dim.date", + "fn": "sum", + }, + { + "dimension": "default.date_dim.date", + "fn": "last_value", + }, + ], + }, + ) == ["sum"] diff --git a/datajunction-server/tests/sql/decompose_test.py b/datajunction-server/tests/sql/decompose_test.py index b2c93644b..f76004e77 100644 --- a/datajunction-server/tests/sql/decompose_test.py +++ b/datajunction-server/tests/sql/decompose_test.py @@ -2,11 +2,14 @@ Tests for ``datajunction_server.sql.decompose``. """ +from types import SimpleNamespace + import pytest import pytest_asyncio from sqlalchemy.ext.asyncio import AsyncSession from datajunction_server.database.node import Node, NodeRelationship, NodeRevision +from datajunction_server.errors import DJInvalidInputException from datajunction_server.models.cube_materialization import ( Aggregability, AggregationRule, @@ -14,6 +17,10 @@ ) from datajunction_server.models.engine import Dialect from datajunction_server.models.node_type import NodeType +from datajunction_server.models.reaggregate import ( + DimensionReaggregateRule, + ReaggregationFunction, +) from datajunction_server.sql import functions as dj_functions from datajunction_server.sql.decompose import ( DUPLICATION_INVARIANT_AGGREGATIONS, @@ -62,7 +69,12 @@ async def create_metric(session: AsyncSession, current_user, parent_node): """Fixture to create a metric node with a query.""" created_metrics: list[NodeRevision] = [] - async def _create(query: str, name: str | None = None, parent=None): + async def _create( + query: str, + name: str | None = None, + parent=None, + reaggregate: dict[str, object] | None = None, + ): parent_to_use = parent if parent else parent_node metric_name = name or f"test_metric_{len(created_metrics)}" @@ -81,6 +93,7 @@ async def _create(query: str, name: str | None = None, parent=None): name=metric_name, type=NodeType.METRIC, query=query, + reaggregate=reaggregate, created_by_id=current_user.id, ) session.add(metric_rev) @@ -120,6 +133,105 @@ async def test_simple_sum(session: AsyncSession, create_metric): ) +@pytest.mark.asyncio +async def test_reaggregate_sum_attaches_spec(session: AsyncSession, create_metric): + """ + Semi-additive declarations attach to the single extracted base measure. + """ + metric_rev = await create_metric( + "SELECT SUM(account_balance) FROM parent_node", + reaggregate={ + "rules": [ + { + "dimension": "default.date_dim.date", + "fn": "last_value", + }, + ], + }, + ) + + extractor = MetricComponentExtractor(metric_rev.id) + measures, derived_sql = await extractor.extract(session) + + expected_spec = DimensionReaggregateRule( + dimension="default.date_dim.date", + fn=ReaggregationFunction.LAST_VALUE, + ) + assert measures == [ + MetricComponent( + name="account_balance_sum_8e611a76", + expression="account_balance", + aggregation="SUM", + merge="SUM", + rule=AggregationRule( + type=Aggregability.FULL, + reaggregate=expected_spec, + ), + ), + ] + assert_sql_equal( + str(derived_sql), + "SELECT SUM(account_balance_sum_8e611a76) FROM parent_node", + ) + + +@pytest.mark.asyncio +async def test_reaggregate_multiple_rules_not_supported( + session: AsyncSession, + create_metric, +): + """ + V1 semi-additive metrics support one dimension-specific rule. + """ + metric_rev = await create_metric( + "SELECT SUM(account_balance) FROM parent_node", + reaggregate={ + "rules": [ + { + "dimension": "default.date_dim.date", + "fn": "last_value", + }, + { + "dimension": "default.region_dim.region", + "fn": "first_value", + }, + ], + }, + ) + + extractor = MetricComponentExtractor(metric_rev.id) + with pytest.raises(DJInvalidInputException, match="exactly one rule"): + await extractor.extract(session) + + +@pytest.mark.asyncio +async def test_reaggregate_unsupported_dimension_function_not_supported( + session: AsyncSession, + create_metric, +): + """ + Dimension-specific reaggregate accepts only collapse-safe functions. + """ + metric_rev = await create_metric( + "SELECT SUM(account_balance) FROM parent_node", + reaggregate={ + "rules": [ + { + "dimension": "default.date_dim.date", + "fn": "sum", + }, + ], + }, + ) + + extractor = MetricComponentExtractor(metric_rev.id) + with pytest.raises( + DJInvalidInputException, + match="unsupported dimension reaggregation function", + ): + await extractor.extract(session) + + @pytest.mark.asyncio async def test_sum_with_cast(session: AsyncSession, create_metric): """ @@ -854,6 +966,50 @@ async def test_unsupported_aggregation_function(session: AsyncSession, create_me ) +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("query", "match"), + [ + ( + "SELECT AVG(sales_amount) FROM parent_node", + "exactly one component", + ), + ( + "SELECT COUNT(DISTINCT user_id) FROM parent_node", + "one non-distinct fully-aggregatable component", + ), + ( + "SELECT MEDIAN(sales_amount) FROM parent_node", + "decomposable aggregation", + ), + ], +) +async def test_unsupported_reaggregate_shapes( + session: AsyncSession, + create_metric, + query: str, + match: str, +): + """ + V1 semi-additive metrics only support one ordinary decomposable component. + """ + metric_rev = await create_metric( + query, + reaggregate={ + "rules": [ + { + "dimension": "default.date_dim.date", + "fn": "last_value", + }, + ], + }, + ) + + extractor = MetricComponentExtractor(metric_rev.id) + with pytest.raises(DJInvalidInputException, match=match): + await extractor.extract(session) + + @pytest.mark.asyncio async def test_count_if(session: AsyncSession, create_metric): """ @@ -1539,7 +1695,12 @@ async def create_derived_metric(clean_session: AsyncSession, clean_current_user) session = clean_session current_user = clean_current_user - async def _create(name: str, query: str, base_metric_nodes: list[Node]): + async def _create( + name: str, + query: str, + base_metric_nodes: list[Node], + reaggregate: dict[str, object] | None = None, + ): metric_node = Node( name=name, type=NodeType.METRIC, @@ -1555,6 +1716,7 @@ async def _create(name: str, query: str, base_metric_nodes: list[Node]): name=name, type=NodeType.METRIC, query=query, + reaggregate=reaggregate, created_by_id=current_user.id, ) session.add(metric_rev) @@ -1571,6 +1733,116 @@ async def _create(name: str, query: str, base_metric_nodes: list[Node]): return _create +@pytest.mark.asyncio +async def test_derived_metric_reaggregate_not_supported_db_load( + clean_session: AsyncSession, + create_base_metric, + create_derived_metric, +): + """ + Semi-additive declarations are currently limited to base metrics. + """ + session = clean_session + balance_node, _ = await create_base_metric( + "default.balance", + "SELECT SUM(account_balance) FROM parent_node", + ) + _, derived_rev = await create_derived_metric( + "default.double_balance", + "SELECT default.balance * 2", + [balance_node], + reaggregate={ + "rules": [ + { + "dimension": "default.date_dim.date", + "fn": "last_value", + }, + ], + }, + ) + + extractor = MetricComponentExtractor(derived_rev.id) + with pytest.raises( + DJInvalidInputException, + match="Derived metric `default.double_balance` declares reaggregate", + ): + await extractor.extract(session) + + +def test_derived_metric_reaggregate_not_supported_cache_load(): + """ + Cache-based extraction enforces the same base-metric-only constraint. + """ + reaggregate = { + "rules": [ + { + "dimension": "default.date_dim.date", + "fn": "last_value", + }, + ], + } + base_node = SimpleNamespace( + name="default.balance", + type=NodeType.METRIC, + current=SimpleNamespace( + query="SELECT SUM(account_balance) FROM parent_node", + reaggregate=None, + ), + ) + derived_node = SimpleNamespace( + name="default.double_balance", + type=NodeType.METRIC, + current=SimpleNamespace( + query="SELECT default.balance * 2", + reaggregate=reaggregate, + ), + ) + + extractor = MetricComponentExtractor(0) + with pytest.raises( + DJInvalidInputException, + match="Derived metric `default.double_balance` declares reaggregate", + ): + extractor._build_metric_data_from_cache( + derived_node, + { + "default.balance": base_node, + "default.double_balance": derived_node, + }, + {"default.double_balance": ["default.balance"]}, + ) + + +def test_normalize_aliases_leaves_other_namespaces_qualified(): + """ + Alias normalization only strips the primary parent alias from column refs. + """ + extractor = MetricComponentExtractor(0) + query = parse("SELECT p.value, other.value FROM parent_node p") + + normalized = extractor._normalize_aliases(query) + + assert_sql_equal(str(normalized), "SELECT value, other.value FROM parent_node") + + +def test_substitute_metric_references_ignores_non_metric_columns(): + """ + Derived metric substitution leaves ordinary columns in place. + """ + extractor = MetricComponentExtractor(0) + query = parse("SELECT default.balance + raw_value FROM parent_node") + + substituted = extractor._substitute_metric_references( + query, + {"default.balance": ([], "SUM(balance)")}, + ) + + assert_sql_equal( + str(substituted), + "SELECT SUM(balance) + raw_value FROM parent_node", + ) + + @pytest.mark.asyncio async def test_extract_derived_metric_revenue_per_order( clean_session: AsyncSession, diff --git a/datajunction-ui/src/app/pages/AddEditNodePage/ReaggregateFields.jsx b/datajunction-ui/src/app/pages/AddEditNodePage/ReaggregateFields.jsx new file mode 100644 index 000000000..eb1344376 --- /dev/null +++ b/datajunction-ui/src/app/pages/AddEditNodePage/ReaggregateFields.jsx @@ -0,0 +1,79 @@ +/** + * Semi-additive metric controls. + */ +import { ErrorMessage, Field, useField, useFormikContext } from 'formik'; +import { useContext, useEffect, useMemo, useState } from 'react'; +import DJClientContext from '../../providers/djclient'; +import { labelize } from '../../../utils/form'; +import { FormikSelect } from './FormikSelect'; + +const SEMI_ADDITIVE_FUNCTIONS = ['last_value', 'first_value', 'min', 'max']; + +export const ReaggregateFields = () => { + const djClient = useContext(DJClientContext).DataJunctionAPI; + const { values } = useFormikContext(); + const [dimensionField] = useField('reaggregate_dimension'); + const [dimensionOptions, setDimensionOptions] = useState([]); + + useEffect(() => { + const fetchData = async () => { + if (values.upstream_node) { + const data = await djClient.node(values.upstream_node); + setDimensionOptions( + data.columns.map(col => ({ + value: `${values.upstream_node}.${col.name}`, + label: `${values.upstream_node}.${col.name}`, + })), + ); + } else { + setDimensionOptions([]); + } + }; + fetchData().catch(console.error); + }, [djClient, values.upstream_node]); + + const selectOptions = useMemo(() => { + if ( + !dimensionField.value || + dimensionOptions.some(option => option.value === dimensionField.value) + ) { + return dimensionOptions; + } + return [ + { value: dimensionField.value, label: dimensionField.value }, + ...dimensionOptions, + ]; + }, [dimensionField.value, dimensionOptions]); + + return ( +
+
+ + + + + +
+
+ + + + + {SEMI_ADDITIVE_FUNCTIONS.map(func => ( + + ))} + +
+
+ ); +}; diff --git a/datajunction-ui/src/app/pages/AddEditNodePage/__tests__/AddEditNodePageFormSuccess.test.jsx b/datajunction-ui/src/app/pages/AddEditNodePage/__tests__/AddEditNodePageFormSuccess.test.jsx index f399c4b8d..c66a9cdb8 100644 --- a/datajunction-ui/src/app/pages/AddEditNodePage/__tests__/AddEditNodePageFormSuccess.test.jsx +++ b/datajunction-ui/src/app/pages/AddEditNodePage/__tests__/AddEditNodePageFormSuccess.test.jsx @@ -156,6 +156,101 @@ describe('AddEditNodePage submission succeeded', () => { ); }, 60000); + it('for creating a semi-additive metric node', async () => { + const mockDjClient = initializeMockDJClient(); + mockDjClient.DataJunctionAPI.createNode.mockReturnValue({ + status: 200, + json: { name: 'default.some_test_metric' }, + }); + + mockDjClient.DataJunctionAPI.tagsNode.mockReturnValue({ + status: 200, + json: { message: 'Success' }, + }); + + mockDjClient.DataJunctionAPI.node.mockResolvedValue({ + columns: [{ name: 'order_date' }], + }); + + mockDjClient.DataJunctionAPI.listTags.mockReturnValue([ + { name: 'purpose', display_name: 'Purpose' }, + { name: 'intent', display_name: 'Intent' }, + ]); + + mockDjClient.DataJunctionAPI.listMetricMetadata.mockReturnValue( + mocks.metricMetadata, + ); + + mockDjClient.DataJunctionAPI.whoami.mockReturnValue({ + id: 123, + username: 'test_user', + }); + + const element = testElement(mockDjClient); + const { getByTestId } = renderCreateMetric(element); + + await userEvent.type( + screen.getByLabelText('Display Name *'), + 'Some Test Metric', + ); + const selectUpstream = screen.getByTestId('select-upstream-node'); + fireEvent.keyDown(selectUpstream.firstChild, { key: 'ArrowDown' }); + fireEvent.click(screen.getByText('default.repair_orders')); + + await waitFor(() => { + expect(mockDjClient.DataJunctionAPI.node).toHaveBeenCalledWith( + 'default.repair_orders', + ); + }); + + const selectReaggregateDimension = getByTestId( + 'select-semi-additive-dimension', + ); + fireEvent.keyDown(selectReaggregateDimension.firstChild, { + key: 'ArrowDown', + }); + fireEvent.click( + await screen.findByText('default.repair_orders.order_date'), + ); + fireEvent.change(screen.getByLabelText('Semi-Additive Type'), { + target: { value: 'last_value' }, + }); + + await userEvent.type( + screen.getByLabelText('Aggregate Expression *'), + 'SUM(balance)', + ); + await userEvent.click(screen.getByText('Create metric')); + + await waitFor( + () => { + expect(mockDjClient.DataJunctionAPI.createNode).toBeCalledWith( + 'metric', + 'default.some_test_metric', + 'Some Test Metric', + '', + 'SELECT SUM(balance) \n FROM default.repair_orders', + 'published', + 'default', + null, + undefined, + undefined, + undefined, + null, + { + rules: [ + { + dimension: 'default.repair_orders.order_date', + fn: 'last_value', + }, + ], + }, + ); + }, + { timeout: 10000 }, + ); + }, 60000); + it('for editing a transform or dimension node', async () => { const mockDjClient = initializeMockDJClient(); @@ -288,4 +383,159 @@ describe('AddEditNodePage submission succeeded', () => { ).toBeInTheDocument(); }); }, 1000000); + + it('for editing a semi-additive metric node', async () => { + const mockDjClient = initializeMockDJClient(); + + mockDjClient.DataJunctionAPI.getNodeForEditing.mockReturnValue({ + ...mocks.mockGetMetricNode, + current: { + ...mocks.mockGetMetricNode.current, + reaggregate: { + rules: [ + { + dimension: 'v3.date.date_id[order]', + fn: 'LAST_VALUE', + }, + ], + }, + }, + }); + mockDjClient.DataJunctionAPI.patchNode = vi.fn(); + mockDjClient.DataJunctionAPI.patchNode.mockReturnValue({ + status: 201, + json: { name: 'default.num_repair_orders', type: 'metric' }, + }); + + mockDjClient.DataJunctionAPI.tagsNode.mockReturnValue({ + status: 200, + json: { message: 'Success' }, + }); + + mockDjClient.DataJunctionAPI.listTags.mockReturnValue([ + { name: 'purpose', display_name: 'Purpose' }, + { name: 'intent', display_name: 'Intent' }, + ]); + + mockDjClient.DataJunctionAPI.whoami.mockReturnValue({ + id: 123, + username: 'test_user', + }); + + const element = testElement(mockDjClient); + renderEditNode(element); + + await waitFor(() => { + expect(screen.getByLabelText('Semi-Additive Type')).toHaveValue( + 'last_value', + ); + expect(screen.getByText('v3.date.date_id[order]')).toBeInTheDocument(); + }); + + await userEvent.type(screen.getByLabelText('Description'), '!!!'); + await userEvent.click(screen.getByText('Save')); + + await waitFor(() => { + expect(mockDjClient.DataJunctionAPI.patchNode).toBeCalledWith( + 'default.num_repair_orders', + 'Default: Num Repair Orders', + 'Number of repair orders!!!', + 'SELECT count(repair_order_id) \n FROM default.repair_orders', + 'published', + ['repair_order_id', 'country'], + 'neutral', + 'unitless', + 5, + [], + ['dj'], + { key1: 'value1', key2: 'value2' }, + { + rules: [ + { + dimension: 'v3.date.date_id[order]', + fn: 'last_value', + }, + ], + }, + ); + }); + }, 1000000); + + it('for clearing semi-additive fields from a metric node', async () => { + const mockDjClient = initializeMockDJClient(); + + mockDjClient.DataJunctionAPI.getNodeForEditing.mockReturnValue({ + ...mocks.mockGetMetricNode, + current: { + ...mocks.mockGetMetricNode.current, + reaggregate: { + rules: [ + { + dimension: 'v3.date.date_id[order]', + fn: 'LAST_VALUE', + }, + ], + }, + }, + }); + mockDjClient.DataJunctionAPI.patchNode = vi.fn(); + mockDjClient.DataJunctionAPI.patchNode.mockReturnValue({ + status: 201, + json: { name: 'default.num_repair_orders', type: 'metric' }, + }); + + mockDjClient.DataJunctionAPI.tagsNode.mockReturnValue({ + status: 200, + json: { message: 'Success' }, + }); + + mockDjClient.DataJunctionAPI.listTags.mockReturnValue([ + { name: 'purpose', display_name: 'Purpose' }, + { name: 'intent', display_name: 'Intent' }, + ]); + + mockDjClient.DataJunctionAPI.whoami.mockReturnValue({ + id: 123, + username: 'test_user', + }); + + const element = testElement(mockDjClient); + renderEditNode(element); + + await waitFor(() => { + expect(screen.getByLabelText('Semi-Additive Type')).toHaveValue( + 'last_value', + ); + expect(screen.getByText('v3.date.date_id[order]')).toBeInTheDocument(); + }); + + const clearReaggregateDimension = screen + .getByTestId('select-semi-additive-dimension') + .querySelector('.ReaggregateDimension__clear-indicator'); + expect(clearReaggregateDimension).toBeInTheDocument(); + fireEvent.mouseDown(clearReaggregateDimension, { button: 0 }); + fireEvent.change(screen.getByLabelText('Semi-Additive Type'), { + target: { value: '' }, + }); + + await userEvent.click(screen.getByText('Save')); + + await waitFor(() => { + expect(mockDjClient.DataJunctionAPI.patchNode).toBeCalledWith( + 'default.num_repair_orders', + 'Default: Num Repair Orders', + 'Number of repair orders', + 'SELECT count(repair_order_id) \n FROM default.repair_orders', + 'published', + ['repair_order_id', 'country'], + 'neutral', + 'unitless', + 5, + [], + ['dj'], + { key1: 'value1', key2: 'value2' }, + null, + ); + }); + }, 1000000); }); diff --git a/datajunction-ui/src/app/pages/AddEditNodePage/index.jsx b/datajunction-ui/src/app/pages/AddEditNodePage/index.jsx index 2dc0acb22..fcb0f3088 100644 --- a/datajunction-ui/src/app/pages/AddEditNodePage/index.jsx +++ b/datajunction-ui/src/app/pages/AddEditNodePage/index.jsx @@ -27,6 +27,7 @@ import { RequiredDimensionsSelect } from './RequiredDimensionsSelect'; import LoadingIcon from '../../icons/LoadingIcon'; import { ColumnsSelect } from './ColumnsSelect'; import { CustomMetadataField } from './CustomMetadataField'; +import { ReaggregateFields } from './ReaggregateFields'; class Action { static Add = new Action('add'); @@ -71,6 +72,9 @@ export function AddEditNodePage({ extensions = {} }) { mode: 'published', owners: [], custom_metadata: '', + reaggregate_dimension: '', + reaggregate_function: '', + had_reaggregate: false, }; const validator = values => { @@ -92,6 +96,14 @@ export function AddEditNodePage({ extensions = {} }) { if (values.type !== 'metric' && !values.query) { errors.query = 'Required'; } + if (values.type === 'metric') { + if (values.reaggregate_dimension && !values.reaggregate_function) { + errors.reaggregate_function = 'Required'; + } + if (values.reaggregate_function && !values.reaggregate_dimension) { + errors.reaggregate_dimension = 'Required'; + } + } return errors; }; @@ -161,8 +173,28 @@ export function AddEditNodePage({ extensions = {} }) { return `SELECT ${aggregateExpression}`; }; + const buildReaggregateSpec = values => { + if (values.reaggregate_dimension && values.reaggregate_function) { + return { + rules: [ + { + dimension: values.reaggregate_dimension, + fn: values.reaggregate_function, + }, + ], + }; + } + return values.had_reaggregate ? null : undefined; + }; + + const firstReaggregateRule = reaggregate => reaggregate?.rules?.[0]; + + const normalizeReaggregateFunction = reaggregateFunction => + reaggregateFunction ? reaggregateFunction.toLowerCase() : ''; + const createNode = async (values, setStatus) => { - const { status, json } = await djClient.createNode( + const reaggregate = buildReaggregateSpec(values); + const createNodeArgs = [ nodeType, values.name, values.display_name, @@ -180,7 +212,11 @@ export function AddEditNodePage({ extensions = {} }) { ? values.required_dimensions : undefined, parseCustomMetadata(values.custom_metadata), - ); + ]; + if (reaggregate !== undefined) { + createNodeArgs.push(reaggregate); + } + const { status, json } = await djClient.createNode(...createNodeArgs); if (status === 200 || status === 201) { if (values.tags) { await djClient.tagsNode(values.name, values.tags); @@ -201,7 +237,8 @@ export function AddEditNodePage({ extensions = {} }) { }; const patchNode = async (values, setStatus) => { - const { status, json } = await djClient.patchNode( + const reaggregate = buildReaggregateSpec(values); + const patchNodeArgs = [ values.name, values.display_name, values.description, @@ -219,7 +256,11 @@ export function AddEditNodePage({ extensions = {} }) { : undefined, values.owners, parseCustomMetadata(values.custom_metadata), - ); + ]; + if (reaggregate !== undefined) { + patchNodeArgs.push(reaggregate); + } + const { status, json } = await djClient.patchNode(...patchNodeArgs); const tagsResponse = await djClient.tagsNode( values.name, values.tags.map(tag => tag), @@ -296,6 +337,12 @@ export function AddEditNodePage({ extensions = {} }) { required_dimensions: node.current.requiredDimensions.map( dim => dim.name, ), + reaggregate_dimension: + firstReaggregateRule(node.current.reaggregate)?.dimension || '', + reaggregate_function: normalizeReaggregateFunction( + firstReaggregateRule(node.current.reaggregate)?.fn, + ), + had_reaggregate: Boolean(node.current.reaggregate), upstream_node: '', // Derived metrics have no upstream node aggregate_expression: derivedExpression, }; @@ -310,6 +357,12 @@ export function AddEditNodePage({ extensions = {} }) { required_dimensions: node.current.requiredDimensions.map( dim => dim.name, ), + reaggregate_dimension: + firstReaggregateRule(node.current.reaggregate)?.dimension || '', + reaggregate_function: normalizeReaggregateFunction( + firstReaggregateRule(node.current.reaggregate)?.fn, + ), + had_reaggregate: Boolean(node.current.reaggregate), upstream_node: nonMetricParent?.name || '', aggregate_expression: node.current.metricMetadata?.expression, }; @@ -360,6 +413,9 @@ export function AddEditNodePage({ extensions = {} }) { 'metric_direction', 'significant_digits', 'required_dimensions', + 'reaggregate_dimension', + 'reaggregate_function', + 'had_reaggregate', 'owners', 'custom_metadata', ]; @@ -495,16 +551,21 @@ export function AddEditNodePage({ extensions = {} }) { if (action === Action.Edit) { const data = await getExistingNodeData(name); runValidityChecks(data, setNode, setMessage); - updateFieldsWithNodeData( - data, - setFieldValue, - setNode, - setSelectTags, - setSelectPrimaryKey, - setSelectUpstreamNode, - setSelectRequiredDims, - setSelectOwners, - ); + if ( + data.message === undefined && + nodeCanBeEdited(data.type) + ) { + updateFieldsWithNodeData( + data, + setFieldValue, + setNode, + setSelectTags, + setSelectPrimaryKey, + setSelectUpstreamNode, + setSelectRequiredDims, + setSelectOwners, + ); + } } }; fetchData().catch(console.error); @@ -560,7 +621,10 @@ export function AddEditNodePage({ extensions = {} }) { /> )} {(nodeType === 'metric' || node.type === 'metric') && ( - + <> + + + )} diff --git a/datajunction-ui/src/app/pages/AddEditNodePage/styles.css b/datajunction-ui/src/app/pages/AddEditNodePage/styles.css index f18307426..1eb0b846c 100644 --- a/datajunction-ui/src/app/pages/AddEditNodePage/styles.css +++ b/datajunction-ui/src/app/pages/AddEditNodePage/styles.css @@ -139,6 +139,16 @@ margin: 0 !important; } +.node-builder .ReaggregateFields { + gap: 12px; +} + +.node-builder .ReaggregateFields > * { + flex: 1; + min-width: 0; + margin: 0 !important; +} + /* Override legacy form-scoped element styling for inputs/selects/buttons. * Scope to the specific wrapper classes so we DON'T accidentally restyle * the hidden that lives inside react-select's Control — that input diff --git a/datajunction-ui/src/app/pages/NodePage/NodeInfoTab.jsx b/datajunction-ui/src/app/pages/NodePage/NodeInfoTab.jsx index cd18f4fb7..bf5667944 100644 --- a/datajunction-ui/src/app/pages/NodePage/NodeInfoTab.jsx +++ b/datajunction-ui/src/app/pages/NodePage/NodeInfoTab.jsx @@ -10,6 +10,23 @@ import { labelize } from '../../../utils/form'; SyntaxHighlighter.registerLanguage('sql', sql); foundation.hljs['padding'] = '2rem'; +const dimensionNodePath = dimension => { + if (!dimension) { + return null; + } + const dimensionWithoutRole = dimension.replace(/\[[^\]]+\]$/, ''); + const parts = dimensionWithoutRole.split('.'); + if (parts.length <= 1) { + return null; + } + return parts.slice(0, -1).join('.'); +}; + +const reaggregateFunctionLabel = func => + func ? labelize(func.toLowerCase()) : null; + +const firstReaggregateRule = reaggregate => reaggregate?.rules?.[0]; + // interface MetricInfo { // name: string; // current: MetricRevision; @@ -59,6 +76,7 @@ export default function NodeInfoTab({ node }) { expression: metric.current.metricMetadata?.expression, incompatible_druid_functions: metric.current.metricMetadata?.incompatibleDruidFunctions || [], + reaggregate: metric.current.reaggregate, }); }; if (node.type === 'metric') { @@ -267,6 +285,39 @@ export default function NodeInfoTab({ node }) { {metricInfo?.metric_metadata?.significantDigits || 'None'}

+
+
Semi-Additive
+

+ {firstReaggregateRule(metricInfo?.reaggregate) ? ( + <> + {reaggregateFunctionLabel( + firstReaggregateRule(metricInfo.reaggregate).fn, + )} + {' on '} + {dimensionNodePath( + firstReaggregateRule(metricInfo.reaggregate).dimension, + ) ? ( + + {firstReaggregateRule(metricInfo.reaggregate).dimension} + + ) : ( + firstReaggregateRule(metricInfo.reaggregate).dimension + )} + + ) : ( + 'None' + )} +

+
) : ( diff --git a/datajunction-ui/src/app/pages/NodePage/__tests__/NodePage.test.jsx b/datajunction-ui/src/app/pages/NodePage/__tests__/NodePage.test.jsx index 7fe47c5ff..66e676b95 100644 --- a/datajunction-ui/src/app/pages/NodePage/__tests__/NodePage.test.jsx +++ b/datajunction-ui/src/app/pages/NodePage/__tests__/NodePage.test.jsx @@ -388,6 +388,10 @@ describe('', () => { screen.getByRole('dialog', { name: 'RequiredDimensions' }), ).toHaveTextContent(''); + expect( + screen.getByRole('dialog', { name: 'Reaggregate' }), + ).toHaveTextContent('None'); + expect( screen.getByRole('dialog', { name: 'DisplayName' }), ).toHaveTextContent('Default: Num Repair Orders'); @@ -414,6 +418,50 @@ describe('', () => { expect(container.getElementsByClassName('language-sql')).toMatchSnapshot(); }, 60000); + it('renders semi-additive metric information with a dimension link', async () => { + const djClient = mockDJClient(); + djClient.DataJunctionAPI.node.mockReturnValue(mocks.mockMetricNode); + djClient.DataJunctionAPI.getMetric.mockResolvedValue({ + ...mocks.mockMetricNodeJson, + current: { + ...mocks.mockMetricNodeJson.current, + reaggregate: { + rules: [ + { + dimension: 'v3.date.date_id[order]', + fn: 'LAST_VALUE', + }, + ], + }, + }, + }); + const element = ( + + + + ); + render( + + + + + , + ); + + const reaggregate = await screen.findByRole('dialog', { + name: 'Reaggregate', + }); + await waitFor(() => { + expect(reaggregate).toHaveTextContent( + 'Last Value on v3.date.date_id[order]', + ); + }); + + expect( + screen.getByRole('link', { name: 'v3.date.date_id[order]' }), + ).toHaveAttribute('href', '/nodes/v3.date'); + }, 60000); + it('hides Edit and shows the read-only badge for a node in a read-only (flat git) namespace', async () => { const djClient = mockDJClient(); djClient.DataJunctionAPI.node.mockReturnValue(mocks.mockMetricNode); diff --git a/datajunction-ui/src/app/services/DJService.js b/datajunction-ui/src/app/services/DJService.js index cec600653..213777c20 100644 --- a/datajunction-ui/src/app/services/DJService.js +++ b/datajunction-ui/src/app/services/DJService.js @@ -520,6 +520,14 @@ export const DataJunctionAPI = { requiredDimensions { name } + reaggregate { + fn + weight + rules { + dimension + fn + } + } mode customMetadata } @@ -608,6 +616,14 @@ export const DataJunctionAPI = { requiredDimensions { name } + reaggregate { + fn + weight + rules { + dimension + fn + } + } } } } @@ -747,6 +763,7 @@ export const DataJunctionAPI = { metric_unit, required_dimensions, custom_metadata, + reaggregate, ) { const metricMetadata = metric_direction || metric_unit @@ -767,6 +784,7 @@ export const DataJunctionAPI = { metric_metadata: metricMetadata, required_dimensions: required_dimensions, custom_metadata: custom_metadata, + reaggregate: reaggregate, }; // Remove undefined fields to avoid sending them to the API Object.keys(requestBody).forEach( @@ -797,6 +815,7 @@ export const DataJunctionAPI = { required_dimensions, owners, custom_metadata, + reaggregate, ) { try { const metricMetadata = @@ -818,6 +837,7 @@ export const DataJunctionAPI = { required_dimensions: required_dimensions, owners: owners, custom_metadata: custom_metadata, + reaggregate: reaggregate, }; // Remove undefined fields to avoid sending them to the API Object.keys(requestBody).forEach( diff --git a/datajunction-ui/src/app/services/__tests__/DJService.test.jsx b/datajunction-ui/src/app/services/__tests__/DJService.test.jsx index 0335179ca..3c6cfb5d3 100644 --- a/datajunction-ui/src/app/services/__tests__/DJService.test.jsx +++ b/datajunction-ui/src/app/services/__tests__/DJService.test.jsx @@ -182,6 +182,59 @@ describe('DataJunctionAPI', () => { }); }); + it('calls createNode with semi-additive metadata correctly', async () => { + fetch.mockResponseOnce(JSON.stringify({})); + await DataJunctionAPI.createNode( + 'metric', + 'default.daily_balance', + 'Daily Balance', + 'Daily balance', + 'SELECT sum(balance) FROM default.accounts', + 'published', + 'default', + null, + undefined, + undefined, + undefined, + null, + { + rules: [ + { + dimension: 'v3.date.date_id[order]', + fn: 'last_value', + }, + ], + }, + ); + expect(fetch).toHaveBeenCalledWith(`${DJ_URL}/nodes/metric`, { + method: 'POST', + headers: { + 'Content-Type': 'application/json', + }, + body: JSON.stringify({ + name: 'default.daily_balance', + display_name: 'Daily Balance', + description: 'Daily balance', + query: 'SELECT sum(balance) FROM default.accounts', + mode: 'published', + namespace: 'default', + primary_key: null, + metric_metadata: null, + required_dimensions: undefined, + custom_metadata: null, + reaggregate: { + rules: [ + { + dimension: 'v3.date.date_id[order]', + fn: 'last_value', + }, + ], + }, + }), + credentials: 'include', + }); + }); + it('calls patchNode correctly', async () => { const sampleArgs = [ 'name', @@ -233,6 +286,65 @@ describe('DataJunctionAPI', () => { }); }); + it('calls patchNode with semi-additive metadata correctly', async () => { + fetch.mockResponseOnce(JSON.stringify({})); + await DataJunctionAPI.patchNode( + 'default.daily_balance', + 'Daily Balance', + 'Daily balance', + 'SELECT sum(balance) FROM default.accounts', + 'published', + null, + 'neutral', + 'unitless', + 5, + [], + ['dj'], + null, + { + rules: [ + { + dimension: 'v3.date.date_id[order]', + fn: 'last_value', + }, + ], + }, + ); + expect(fetch).toHaveBeenCalledWith( + `${DJ_URL}/nodes/default.daily_balance`, + { + method: 'PATCH', + headers: { + 'Content-Type': 'application/json', + }, + body: JSON.stringify({ + display_name: 'Daily Balance', + description: 'Daily balance', + query: 'SELECT sum(balance) FROM default.accounts', + mode: 'published', + primary_key: null, + metric_metadata: { + direction: 'neutral', + unit: 'unitless', + significant_digits: 5, + }, + required_dimensions: [], + owners: ['dj'], + custom_metadata: null, + reaggregate: { + rules: [ + { + dimension: 'v3.date.date_id[order]', + fn: 'last_value', + }, + ], + }, + }), + credentials: 'include', + }, + ); + }); + it('calls createCube correctly', async () => { const sampleArgs = [ 'default.node_name', @@ -1297,6 +1409,7 @@ describe('DataJunctionAPI', () => { }), ); await DataJunctionAPI.getMetric('default.num_repair_orders'); + const requestBody = JSON.parse(fetch.mock.calls[0][1].body); expect(fetch).toHaveBeenCalledWith( `${DJ_URL}/graphql`, expect.objectContaining({ @@ -1307,6 +1420,13 @@ describe('DataJunctionAPI', () => { }, }), ); + expect(requestBody.variables).toEqual({ + name: 'default.num_repair_orders', + }); + expect(requestBody.query).toContain('reaggregate'); + expect(requestBody.query).toContain('dimension'); + expect(requestBody.query).toContain('rules'); + expect(requestBody.query).toContain('fn'); }); it('calls notebookExportCube correctly', async () => { @@ -1924,7 +2044,13 @@ describe('DataJunctionAPI', () => { ); const result = await DataJunctionAPI.getNodeForEditing('default.node1'); + const requestBody = JSON.parse(fetch.mock.calls[0][1].body); expect(result).toHaveProperty('name', 'default.node1'); + expect(requestBody.variables).toEqual({ name: 'default.node1' }); + expect(requestBody.query).toContain('reaggregate'); + expect(requestBody.query).toContain('dimension'); + expect(requestBody.query).toContain('rules'); + expect(requestBody.query).toContain('fn'); }); it('returns null when getNodeForEditing finds no nodes', async () => { diff --git a/docker-compose.yml b/docker-compose.yml index 82020adac..845bf7a2f 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -106,7 +106,13 @@ services: stdin_open: true volumes: - ./datajunction-ui:/usr/src/app/ - - ./datajunction-ui/node_modules:/usr/src/app/node_modules + # Keep platform-specific packages inside Docker. `nocopy` prevents the + # host's macOS node_modules from seeding this Linux volume. + - type: volume + source: djui_node_modules + target: /usr/src/app/node_modules + volume: + nocopy: true environment: - NODE_OPTIONS=--max-old-space-size=4096 - WATCHPACK_POLLING=true @@ -327,3 +333,4 @@ services: volumes: postgres_metadata: + djui_node_modules: