From a3fa151655a14f208226f689e509fedfcff58309 Mon Sep 17 00:00:00 2001 From: Robin Davis Date: Tue, 15 Sep 2026 16:43:42 -0700 Subject: [PATCH 1/4] Add parameters for sketch-backed reaggregation Persist reaggregation parameters on frozen measures and include them in measure identity and pre-aggregation matching. Thread the parameters into decomposition and Druid metric spec construction, with migration and API coverage. --- .gitignore | 3 + ...001params_add_params_to_frozen_measures.py | 29 ++ .../datajunction_server/api/cubes.py | 39 ++- .../api/graphql/scalars/metricmetadata.py | 34 +- .../api/graphql/scalars/node.py | 1 + .../api/graphql/schema.graphql | 2 + .../construction/build_v3/combiners.py | 1 + .../construction/build_v3/preagg_matcher.py | 40 ++- .../datajunction_server/database/measure.py | 6 + .../database/preaggregation.py | 57 +++- .../datajunction_server/internal/nodes.py | 25 +- .../internal/preaggregations.py | 14 +- .../models/cube_materialization.py | 29 +- .../datajunction_server/models/decompose.py | 8 + .../models/materialization.py | 136 +++++++- .../datajunction_server/models/reaggregate.py | 6 + .../datajunction_server/sql/decompose.py | 23 ++ .../datajunction_server/sql/functions.py | 8 +- .../api/cubes_build_metrics_spec_test.py | 302 ++++++++++++++++++ datajunction-server/tests/api/cubes_test.py | 16 + datajunction-server/tests/api/metrics_test.py | 5 + .../tests/api/preaggregations_test.py | 2 + datajunction-server/tests/api/sql_v2_test.py | 5 + .../build_v3/preagg_matcher_test.py | 15 +- .../tests/database/preaggregation_test.py | 78 +++++ .../nodes/derive_frozen_measures_test.py | 73 +++++ .../tests/models/reaggregate_test.py | 21 ++ .../tests/sql/decompose_test.py | 72 +++++ .../tests/sql/functions_test.py | 27 ++ 29 files changed, 986 insertions(+), 91 deletions(-) create mode 100644 datajunction-server/datajunction_server/alembic/versions/2026_09_15_0000-fm0001params_add_params_to_frozen_measures.py diff --git a/.gitignore b/.gitignore index dcfc4b40f5..c8450e9529 100644 --- a/.gitignore +++ b/.gitignore @@ -134,3 +134,6 @@ Untitled* postgres_metadata postgres_superset node_modules + +# Local Claude Code workspace state (large; never committed) +.claude/ diff --git a/datajunction-server/datajunction_server/alembic/versions/2026_09_15_0000-fm0001params_add_params_to_frozen_measures.py b/datajunction-server/datajunction_server/alembic/versions/2026_09_15_0000-fm0001params_add_params_to_frozen_measures.py new file mode 100644 index 0000000000..c00473a556 --- /dev/null +++ b/datajunction-server/datajunction_server/alembic/versions/2026_09_15_0000-fm0001params_add_params_to_frozen_measures.py @@ -0,0 +1,29 @@ +""" +Add params column to frozen_measures + +Adds a JSON column to persist tuning parameters (e.g., accuracy) for sketch-backed measures. + +Revision ID: fm0001params +Revises: rg0001reaggregate +Create Date: 2026-09-15 00:00:00.000000+00:00 +""" + +import sqlalchemy as sa +from alembic import op + +# revision identifiers, used by Alembic. +revision = "fm0001params" +down_revision = "rg0001reaggregate" +branch_labels = None +depends_on = None + + +def upgrade(): + op.add_column( + "frozen_measures", + sa.Column("params", sa.JSON(), nullable=True), + ) + + +def downgrade(): + op.drop_column("frozen_measures", "params") diff --git a/datajunction-server/datajunction_server/api/cubes.py b/datajunction-server/datajunction_server/api/cubes.py index 37dba17f1f..4c1bdf4695 100644 --- a/datajunction-server/datajunction_server/api/cubes.py +++ b/datajunction-server/datajunction_server/api/cubes.py @@ -14,11 +14,11 @@ _reorder_partition_column_last, build_combiner_sql_from_preaggs, ) +from datajunction_server.construction.build_v3.cte import strip_role_suffix 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 from datajunction_server.database.node import Node @@ -34,7 +34,6 @@ AccessDenialMode, get_access_checker, ) -from datajunction_server.models.access import ResourceAction from datajunction_server.internal.materializations import ( build_cube_materialization, stop_cube_materialization_workflows, @@ -44,6 +43,7 @@ get_single_cube_revision_metadata, ) from datajunction_server.internal.views import CubeViewNames, _build_view_body +from datajunction_server.models.access import ResourceAction from datajunction_server.models.cube import ( CubeRevisionMetadata, DimensionValue, @@ -63,11 +63,10 @@ ) from datajunction_server.models.dialect import Dialect from datajunction_server.models.materialization import ( - DRUID_AGG_MAPPING, - DRUID_SKETCH_TYPES, Granularity, MaterializationJobTypeEnum, MaterializationStrategy, + get_druid_aggregator_spec, ) from datajunction_server.models.metric import TranslatedSQL from datajunction_server.models.node_type import NodeNameVersion @@ -179,27 +178,25 @@ def _build_metrics_spec( if internal_name: component = component_by_name.get(internal_name) - druid_type = "longSum" # Default fallback - - if component: - # Use merge function for pre-aggregated data, fall back to aggregation - agg_func = component.merge or component.aggregation - if agg_func: - key = (col.type, agg_func.lower()) - if key in DRUID_AGG_MAPPING: - druid_type = DRUID_AGG_MAPPING[key] - - metric_spec = { + metric_spec = ( + get_druid_aggregator_spec( + column_name=col.name, + column_type=col.type, + aggregation=component.aggregation, + merge=component.merge, + params=component.params, + ) + if component + else None + ) + # Unmappable measures fall back to longSum here, since we're loading + # pre-aggregated data; the materialization config omits them instead. + metric_spec = metric_spec or { "fieldName": col.name, "name": col.name, - "type": druid_type, + "type": "longSum", } - # HLL sketches need additional configuration - if druid_type in DRUID_SKETCH_TYPES: - metric_spec["lgK"] = 12 # Log2 of K, controls precision (4-21) - metric_spec["tgtHllType"] = "HLL_4" # HLL_4, HLL_6, or HLL_8 - metrics.append(metric_spec) return metrics diff --git a/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py b/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py index 873c39b1e6..802cb6ce19 100644 --- a/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py +++ b/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py @@ -1,6 +1,7 @@ """Metric metadata scalars""" import strawberry +from strawberry.scalars import JSON from datajunction_server.models.cube_materialization import ( Aggregability as Aggregability_, @@ -45,16 +46,41 @@ class Unit: class DimensionReaggregateRule: ... -@strawberry.experimental.pydantic.type(model=ReaggregateSpec_, all_fields=True) -class ReaggregateSpec: ... +@strawberry.experimental.pydantic.type(model=ReaggregateSpec_) +class ReaggregateSpec: + """ + Metric reaggregation declaration. + + Fields are listed explicitly rather than via `all_fields` because `params` + is an open dict, which has no automatic GraphQL mapping. + """ + + fn: strawberry.auto + weight: strawberry.auto + rules: strawberry.auto + params: JSON | None = None @strawberry.experimental.pydantic.type(model=AggregationRule_, all_fields=True) class AggregationRule: ... -@strawberry.experimental.pydantic.type(model=MetricComponent_, all_fields=True) -class MetricComponent: ... +@strawberry.experimental.pydantic.type(model=MetricComponent_) +class MetricComponent: + """ + A single measure with accumulate/merge phases. + + Fields are listed explicitly rather than via `all_fields` because `params` + is an open dict, which has no automatic GraphQL mapping. + """ + + name: strawberry.auto + expression: strawberry.auto + aggregation: strawberry.auto + merge: strawberry.auto + rule: strawberry.auto + grain_alias: strawberry.auto + params: JSON | None = None @strawberry.experimental.pydantic.type(model=DecomposedMetric_, all_fields=True) diff --git a/datajunction-server/datajunction_server/api/graphql/scalars/node.py b/datajunction-server/datajunction_server/api/graphql/scalars/node.py index 6ac64b9e47..a80d2692bf 100644 --- a/datajunction-server/datajunction_server/api/graphql/scalars/node.py +++ b/datajunction-server/datajunction_server/api/graphql/scalars/node.py @@ -433,6 +433,7 @@ def reaggregate(self, root: DBNodeRevision) -> ReaggregateSpec | None: if not spec: return None return ReaggregateSpec( + params=spec.params, rules=[ DimensionReaggregateRule( dimension=rule.dimension, diff --git a/datajunction-server/datajunction_server/api/graphql/schema.graphql b/datajunction-server/datajunction_server/api/graphql/schema.graphql index af550a4efc..4edfd79fd7 100644 --- a/datajunction-server/datajunction_server/api/graphql/schema.graphql +++ b/datajunction-server/datajunction_server/api/graphql/schema.graphql @@ -303,6 +303,7 @@ type MetricComponent { merge: String rule: AggregationRule! grainAlias: String + params: JSON } enum MetricDirection { @@ -727,6 +728,7 @@ type Query { type ReaggregateSpec { rules: [DimensionReaggregateRule!]! + params: JSON } enum ReaggregationFunction { diff --git a/datajunction-server/datajunction_server/construction/build_v3/combiners.py b/datajunction-server/datajunction_server/construction/build_v3/combiners.py index c398577452..5ca723234e 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/combiners.py +++ b/datajunction-server/datajunction_server/construction/build_v3/combiners.py @@ -600,6 +600,7 @@ async def build_combiner_sql_from_preaggs( measure_identity_token( compute_expression_hash(m.expression), m.aggregation, + m.params, ) for m in gg.components if m.expression 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 ec12e640ef..4cdcdd3110 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/preagg_matcher.py +++ b/datajunction-server/datajunction_server/construction/build_v3/preagg_matcher.py @@ -19,6 +19,8 @@ from datajunction_server.database.preaggregation import ( PreAggregation, compute_expression_hash, + get_measure_identities, + measure_identity_token, ) from datajunction_server.errors import DJInvalidInputException from datajunction_server.models.decompose import Aggregability, MetricComponent @@ -42,10 +44,9 @@ def get_required_measure_identities( grain_group: GrainGroup, -) -> set[tuple[str, str]]: +) -> set[str]: """ - Identity of each measure a grain group needs: (expression hash, Phase-1 - aggregation). + Identity token of each measure a grain group needs. Matching is name-independent, so ``expr_hash`` stands in for identity -- but it covers the expression alone, making ``SUM(x)`` and ``MAX(x)`` hash alike @@ -55,11 +56,16 @@ def get_required_measure_identities( Phase-1 ``aggregation`` rather than ``merge``, because merge is not unique -- COUNT accumulates with COUNT but merges with SUM. + + Built through ``measure_identity_token`` rather than assembled here, so the + definition of identity lives in one place -- it also folds in sketch tuning + parameters, which a locally-built tuple would silently omit. """ return { - ( + measure_identity_token( compute_expression_hash(component.expression), - component.normalized_aggregation, + component.aggregation, + component.params, ) for _, component in grain_group.components } @@ -348,13 +354,9 @@ def find_matching_preagg( continue join_back = coverage - # Coverage check on (expression hash, Phase-1 aggregation) -- see - # get_required_measure_identities for why the aggregation is part of it. - preagg_measures = { - (measure.expr_hash, measure.normalized_aggregation) - for measure in preagg.measures - if measure.expr_hash - } + # Coverage check on measure identity -- see + # get_required_measure_identities for what that comprises. + preagg_measures = get_measure_identities(preagg.measures) if not required_measures.issubset(preagg_measures): logger.debug( f"[BuildV3] Pre-agg {preagg.id} measures {preagg_measures} " @@ -406,13 +408,21 @@ def get_preagg_measure_column( Returns: Column name in the pre-agg, or None if not found """ - target = ( + target = measure_identity_token( compute_expression_hash(component.expression), - component.normalized_aggregation, + component.aggregation, + component.params, ) for measure in preagg.measures: - if (measure.expr_hash, measure.normalized_aggregation) == target: + if ( + measure_identity_token( + measure.expr_hash, + measure.aggregation, + measure.params, + ) + == target + ): # Externally-registered pre-aggs bind the measure to a physical # column name that may differ from the DJ component name. return measure.source_column or measure.name diff --git a/datajunction-server/datajunction_server/database/measure.py b/datajunction-server/datajunction_server/database/measure.py index bfdb81c11f..32d0bc0796 100644 --- a/datajunction-server/datajunction_server/database/measure.py +++ b/datajunction-server/datajunction_server/database/measure.py @@ -119,6 +119,12 @@ class FrozenMeasure(Base): # How to aggregate the resolved expression (e.g., SUM, COUNT, AVG) aggregation: Mapped[str] + # Tuning parameters for `aggregation`, for sketch-backed measures whose + # function name alone doesn't pin down the accumulate (a t-digest still needs + # a compression). Part of the measure's identity: a sketch frozen at one + # accuracy must not be reused for a metric asking for another. + params: Mapped[dict | None] = mapped_column(JSON, nullable=True) + # Additivity or rollup rule - tells the planner if this measure can be summed, # needs special handling, or is non-additive. rule: Mapped[MeasureAggregationRule] = mapped_column(MeasureAggregationRuleType) diff --git a/datajunction-server/datajunction_server/database/preaggregation.py b/datajunction-server/datajunction_server/database/preaggregation.py index 2572ff4185..32ecfd1cc5 100644 --- a/datajunction-server/datajunction_server/database/preaggregation.py +++ b/datajunction-server/datajunction_server/database/preaggregation.py @@ -6,6 +6,7 @@ import json from datetime import UTC, datetime from functools import partial +from typing import Any import sqlalchemy as sa from sqlalchemy import ( @@ -80,16 +81,64 @@ def compute_grain_group_hash( return hashlib.md5(content.encode()).hexdigest() -def measure_identity_token(expr_hash: str, aggregation: str | None) -> str: +def canonical_params(params: dict[str, Any] | None) -> str: + """ + Stable string form of a component's tuning parameters, for identity. + + Keys are sorted and numeric values normalized, so ``{"compression": 200}`` + and ``{"compression": 200.0}`` cannot yield two identities for one sketch. + """ + if not params: + return "" + + def norm(value: Any) -> Any: + # 200 and 200.0 are the same compression; don't let the literal's + # spelling fork the identity. (bools are not floats, so they pass + # through untouched.) + if isinstance(value, float) and value.is_integer(): + return int(value) + return value + + # `default=str` so an exotic value can't raise from inside pre-agg matching. + # Params arrive as parsed YAML/JSON, so this should be unreachable; an opaque + # TypeError deep in the matcher is a bad way to find out otherwise. + return json.dumps( + {key: norm(params[key]) for key in sorted(params)}, + sort_keys=True, + separators=(",", ":"), + default=str, + ) + + +def measure_identity_token( + expr_hash: str, + aggregation: str | None, + params: dict[str, Any] | None = None, +) -> str: """ Canonical token identifying one measure: expression hash plus Phase-1 - aggregation. + aggregation, plus tuning parameters when the aggregation takes them. The hash alone is not an identity -- ``SUM(x)`` and ``MAX(x)`` hash alike -- and a partial is reusable only by a metric that accumulates the same way. Everything comparing or hashing measures goes through this. + + Parameters extend that reasoning to sketches: a t-digest accumulated at + ``compression=100`` is not interchangeable with one at ``compression=1000``, + though both are ``nflx_tdigest_agg`` over the same expression. The segment is + appended **only** when params are present, so tokens for the overwhelming + majority of components -- plain ``SUM``, ``COUNT``, ``MAX`` -- are byte-identical + to what this returned before params existed. Stored ``preagg_hash`` values are + built from these tokens, so that stability is a compatibility requirement, not + a nicety. + + Note the quantile *fractions* a metric asks for are deliberately not here. One + sketch column serves p50, p95 and p99; folding fractions in would fragment it + into three identical columns. """ - return f"{expr_hash}:{(aggregation or '').strip().upper()}" + base = f"{expr_hash}:{(aggregation or '').strip().upper()}" + suffix = canonical_params(params) + return f"{base}:{suffix}" if suffix else base def get_measure_identities(measures: list[PreAggMeasure]) -> set[str]: @@ -103,7 +152,7 @@ def get_measure_identities(measures: list[PreAggMeasure]) -> set[str]: Set of identity tokens """ return { - measure_identity_token(m.expr_hash, m.aggregation) + measure_identity_token(m.expr_hash, m.aggregation, m.params) for m in measures if m.expr_hash } diff --git a/datajunction-server/datajunction_server/internal/nodes.py b/datajunction-server/datajunction_server/internal/nodes.py index e3093d0801..c0599cf2fd 100644 --- a/datajunction-server/datajunction_server/internal/nodes.py +++ b/datajunction-server/datajunction_server/internal/nodes.py @@ -66,6 +66,7 @@ ) from datajunction_server.internal.access.authorization.context import AuthContext from datajunction_server.internal.caching.interface import Cache +from datajunction_server.internal.custom_metadata import validate_custom_metadata from datajunction_server.internal.history import ActivityType, EntityType from datajunction_server.internal.materializations import ( apply_cube_materialization_swap, @@ -96,6 +97,12 @@ UpsertCubeMaterialization, principal_refs, ) +from datajunction_server.models.decompose import ( + AggregationRule as DecomposeAggregationRule, +) +from datajunction_server.models.decompose import ( + MetricComponent, +) from datajunction_server.models.deployment import ( ChangeTier, CubeSpec, @@ -104,10 +111,6 @@ fold_change_tiers, version_change_tier, ) -from datajunction_server.models.decompose import ( - AggregationRule as DecomposeAggregationRule, - MetricComponent, -) from datajunction_server.models.dimensionlink import ( JoinLinkInput, JoinType, @@ -148,7 +151,6 @@ from datajunction_server.sql.parsing.ast import CompileContext from datajunction_server.sql.parsing.backends.antlr4 import parse, parse_rule from datajunction_server.typing import UTCDatetime -from datajunction_server.internal.custom_metadata import validate_custom_metadata from datajunction_server.utils import ( SEPARATOR, Version, @@ -1085,6 +1087,7 @@ def _new_frozen_measure( upstream_revision_id=upstream_revision_id, expression=measure.expression, aggregation=measure.aggregation, + params=measure.params, rule=_frozen_measure_rule(measure.rule), used_by_node_revisions=[], ) @@ -1100,6 +1103,7 @@ def _raise_if_frozen_measure_conflicts( if ( frozen_measure.expression == measure.expression and frozen_measure.aggregation == measure.aggregation + and (frozen_measure.params or None) == (measure.params or None) and _aggregation_rule_identity(frozen_measure.rule) == _aggregation_rule_identity(measure.rule) ): @@ -1107,7 +1111,7 @@ def _raise_if_frozen_measure_conflicts( raise DJInvalidInputException( f"Frozen measure `{measure.name}` already exists with a different " - "expression, aggregation, or aggregation rule.", + "expression, aggregation, aggregation rule, or tuning parameters.", ) @@ -1989,6 +1993,7 @@ async def cube_metric_component_identities( FrozenMeasure.name, FrozenMeasure.expression, FrozenMeasure.aggregation, + FrozenMeasure.params, ) .select_from(NodeRevisionFrozenMeasure) .join( @@ -2004,8 +2009,12 @@ async def cube_metric_component_identities( rows = (await session.execute(statement)).all() return { f"{metric_name}:{component_name}:" - + measure_identity_token(compute_expression_hash(expression), aggregation) - for metric_name, component_name, expression, aggregation in rows + + measure_identity_token( + compute_expression_hash(expression), + aggregation, + params, + ) + for metric_name, component_name, expression, aggregation, params in rows } diff --git a/datajunction-server/datajunction_server/internal/preaggregations.py b/datajunction-server/datajunction_server/internal/preaggregations.py index 956053fe69..e1347e9d1d 100644 --- a/datajunction-server/datajunction_server/internal/preaggregations.py +++ b/datajunction-server/datajunction_server/internal/preaggregations.py @@ -28,6 +28,7 @@ compute_grain_group_hash, compute_preagg_hash, get_measure_identities, + measure_identity_token, ) from datajunction_server.errors import DJInvalidInputException from datajunction_server.models.decompose import PreAggMeasure @@ -195,7 +196,7 @@ async def register_external_preaggregations( # identity -- (expression hash, Phase-1 aggregation) -- to the declared # column. Keying on the hash alone would collapse SUM(x) and MAX(x), # silently discarding one metric's declared column. - measure_identity_to_column: dict[tuple[str, str], str] = {} + measure_identity_to_column: dict[str, str] = {} for metric_name, physical_column in measure_columns.items(): node = await Node.get_by_name( session, @@ -224,9 +225,10 @@ async def register_external_preaggregations( # is_measure guarantees exactly one component. component = components[0] measure_identity_to_column[ - ( + measure_identity_token( compute_expression_hash(component.expression), - component.normalized_aggregation, + component.aggregation, + component.params, ) ] = physical_column @@ -311,7 +313,11 @@ async def register_external_preaggregations( grain_measures: list[PreAggMeasure] = [] for component in grain_group.components: expr_hash = compute_expression_hash(component.expression) - identity = (expr_hash, component.normalized_aggregation) + identity = measure_identity_token( + expr_hash, + component.aggregation, + component.params, + ) if identity not in measure_identity_to_column: raise DJInvalidInputException( message=( diff --git a/datajunction-server/datajunction_server/models/cube_materialization.py b/datajunction-server/datajunction_server/models/cube_materialization.py index 1e12ebb40f..23abfbd2b1 100644 --- a/datajunction-server/datajunction_server/models/cube_materialization.py +++ b/datajunction-server/datajunction_server/models/cube_materialization.py @@ -16,11 +16,11 @@ ) from datajunction_server.models.materialization import ( DEFAULT_CUBE_RETENTION, - DRUID_AGG_MAPPING, CoverageSpec, MaterializationJobTypeEnum, MaterializationStrategy, SparkSpec, + get_druid_aggregator_spec, ) from datajunction_server.models.node_type import NodeNameVersion from datajunction_server.models.partition import Granularity @@ -385,22 +385,19 @@ def metrics_spec(self) -> list[dict[str, Any]]: Returns the Druid metrics spec for ingestion """ column_mapping = {col.name: col.type for col in self.columns} # type: ignore - return [ - { - "fieldName": measure.name, - "name": measure.name, - "type": DRUID_AGG_MAPPING[ - (column_mapping[measure.name], measure.aggregation.lower()) - ], - } - for measure in self.measures - if measure.aggregation - and ( - column_mapping.get(measure.name), - measure.aggregation.lower(), + specs = ( + get_druid_aggregator_spec( + column_name=measure.name, + column_type=column_mapping.get(measure.name), + aggregation=measure.aggregation, + merge=measure.merge, + params=measure.params, ) - in DRUID_AGG_MAPPING - ] + for measure in self.measures + ) + # Unmappable measures are omitted, as they always have been here -- the + # cube API substitutes longSum instead. See get_druid_aggregator_spec. + return [spec for spec in specs if spec is not None] @computed_field # type: ignore[misc] @property diff --git a/datajunction-server/datajunction_server/models/decompose.py b/datajunction-server/datajunction_server/models/decompose.py index 4b64aeb646..351b2b527f 100644 --- a/datajunction-server/datajunction_server/models/decompose.py +++ b/datajunction-server/datajunction_server/models/decompose.py @@ -9,6 +9,8 @@ - DecomposedMetric: A metric broken into components + combiner expression """ +from typing import Any + from pydantic import BaseModel from datajunction_server.enum import StrEnum @@ -101,6 +103,12 @@ class MetricComponent(BaseModel): # ("order_id"); for complex expressions it is component.name so the identifier stays # valid and consistent with what decompose.py computed. grain_alias: str | None = None + # Tuning parameters for `aggregation`/`merge`, for sketch-backed components + # whose function name alone doesn't pin down the accumulate (a t-digest still + # needs a compression, a KLL a `k`). Part of measure identity -- see + # `measure_identity_token` -- because a sketch built at one accuracy must not + # satisfy a query asking for another. + params: dict[str, Any] | None = None @property def normalized_aggregation(self) -> str: diff --git a/datajunction-server/datajunction_server/models/materialization.py b/datajunction-server/datajunction_server/models/materialization.py index 421e6269ca..1e763e6367 100644 --- a/datajunction-server/datajunction_server/models/materialization.py +++ b/datajunction-server/datajunction_server/models/materialization.py @@ -54,8 +54,124 @@ ("binary", "hll_sketch_agg"): "HLLSketchMerge", } -# Aggregation types that need special handling (extra config parameters) -DRUID_SKETCH_TYPES = {"HLLSketchMerge"} +# Aggregator types that carry extra config in the Druid metricsSpec, mapped to +# their default configuration. A sketch aggregator is not fully specified by its +# type: an HLL needs a precision, a t-digest a compression, a KLL a `k`. Defaults +# apply when a metric declares no `reaggregate.params`. +DRUID_SKETCH_CONFIG: dict[str, dict[str, Any]] = { + "HLLSketchMerge": { + "lgK": 12, # Log2 of K, controls precision (4-21) + "tgtHllType": "HLL_4", # HLL_4, HLL_6, or HLL_8 + }, +} + + +def register_druid_aggregator( + column_type: str, + merge_func: str, + aggregator: str, + default_config: dict[str, Any] | None = None, +) -> None: + """ + Register a Druid ingestion aggregator for a (column type, merge function) pair. + + Args: + column_type: Measures-table column type, e.g. "binary" + merge_func: Phase-2 merge function name, e.g. "nflx_tdigest_agg" + aggregator: Druid aggregator type, e.g. "tDigestSketch" + default_config: Extra metricsSpec keys and their defaults, e.g. + ``{"compression": 200}``. Marks the aggregator as parameterized. + + Raises: + DJInvalidInputException: If the pair is registered to a different aggregator, + or the aggregator has different defaults. Re-registering identical defaults + is a no-op. + """ + key = (column_type, merge_func.lower()) + existing = DRUID_AGG_MAPPING.get(key) + if existing is not None and existing != aggregator: + raise DJInvalidInputException( + message=( + f"Druid aggregator for {key} is already registered as " + f"`{existing}`; refusing to replace it with `{aggregator}`." + ), + ) + new_config = dict(default_config) if default_config else None + existing_config = DRUID_SKETCH_CONFIG.get(aggregator) + if ( + existing_config is not None + and new_config is not None + and existing_config != new_config + ): + raise DJInvalidInputException( + message=( + f"Druid aggregator `{aggregator}` is already registered with " + f"different defaults ({existing_config}); refusing to replace " + f"them with {new_config}." + ), + ) + DRUID_AGG_MAPPING[key] = aggregator + if new_config: + DRUID_SKETCH_CONFIG[aggregator] = new_config + + +def get_druid_aggregator_spec( + column_name: str, + column_type: str | None, + aggregation: str | None, + merge: str | None, + params: dict[str, Any] | None = None, +) -> dict[str, Any] | None: + """ + Build a Druid metricsSpec entry for a measure column, or None if unmappable. + + The `merge` function takes precedence over `aggregation` since ingestion + operates on pre-aggregated partials. + + Returns None for unmappable measures, allowing callers to determine fallback behavior. + + Raises: + DJInvalidInputException: If `params` contains a key not declared by the aggregator. + """ + agg_func = merge or aggregation + if not agg_func or column_type is None: + return None + + aggregator = DRUID_AGG_MAPPING.get((column_type, agg_func.lower())) + if aggregator is None: + return None + + family_config = DRUID_SKETCH_CONFIG.get(aggregator) + if params: + # Restrict `params` to declared tuning knobs to prevent overwriting + # structural keys (like `type` or `fieldName`) and catch typos. + allowed = set(family_config or ()) + unknown = sorted(set(params) - allowed) + if unknown: + raise DJInvalidInputException( + message=( + f"Druid aggregator `{aggregator}` does not accept " + f"{', '.join(f'`{key}`' for key in unknown)}. " + + ( + f"Supported: {', '.join(sorted(allowed))}." + if allowed + else "It takes no parameters." + ) + ), + ) + + spec: dict[str, Any] = { + "fieldName": column_name, + "name": column_name, + "type": aggregator, + } + if family_config: + # Apply metric-specific params over family defaults. + spec.update(family_config) + if params: + spec.update(params) + return spec + # How long ingested cube data is kept in Druid. Without an explicit rule a datasource # inherits the cluster default, which is unrelated to the span the cube holds, so an @@ -458,16 +574,18 @@ def metrics_spec(self) -> dict: for measure in measure_group.measures if (measure.type.lower(), measure.agg.lower()) not in DRUID_AGG_MAPPING ] - return { - measure.name: { - "fieldName": measure.field_name, - "name": measure.field_name, - "type": DRUID_AGG_MAPPING[(measure.type.lower(), measure.agg.lower())], - } + # The V2 `Measure` model lacks a merge phase, so `agg` is used for lookup. + specs = { + measure.name: get_druid_aggregator_spec( + column_name=measure.field_name, + column_type=measure.type.lower(), + aggregation=measure.agg, + merge=None, + ) for measure_group in self.measures.values() # type: ignore for measure in measure_group.measures - if (measure.type.lower(), measure.agg.lower()) in DRUID_AGG_MAPPING } + return {name: spec for name, spec in specs.items() if spec is not None} def build_druid_spec(self, node_revision: "NodeRevision"): """ diff --git a/datajunction-server/datajunction_server/models/reaggregate.py b/datajunction-server/datajunction_server/models/reaggregate.py index 1cd755d706..9ff2da4465 100644 --- a/datajunction-server/datajunction_server/models/reaggregate.py +++ b/datajunction-server/datajunction_server/models/reaggregate.py @@ -1,5 +1,7 @@ """Models for metric reaggregation declarations.""" +from typing import Any + from pydantic import BaseModel, ConfigDict, Field from datajunction_server.enum import StrEnum @@ -77,6 +79,10 @@ class ReaggregateSpec(BaseModel): rules: list[DimensionReaggregateRule] = Field(default_factory=list) + # Tuning parameters for sketch-backed aggregation/merge functions. The + # materialization adapter validates the supported keys for its aggregator. + params: dict[str, Any] | None = None + def dump_reaggregate_spec( spec: ReaggregateSpec | dict | None, diff --git a/datajunction-server/datajunction_server/sql/decompose.py b/datajunction-server/datajunction_server/sql/decompose.py index 5bbb064b43..cb831479e9 100644 --- a/datajunction-server/datajunction_server/sql/decompose.py +++ b/datajunction-server/datajunction_server/sql/decompose.py @@ -1443,6 +1443,9 @@ def _extract_base( if reaggregate is not None and dimension_reaggregate_rules(reaggregate): self._attach_reaggregate_spec(components, reaggregate) + if reaggregate is not None and reaggregate.params: + self._attach_reaggregate_params(components, reaggregate) + if fixed_grain is not None: self._attach_fixed_grain_spec(components, fixed_grain, query_ast) @@ -1496,6 +1499,26 @@ def _raise_unsupported_fixed_grain_shape(reason: str) -> None: f"Unsupported fixed_grain metric shape: {reason}.", ) + def _attach_reaggregate_params( + self, + components: list[MetricComponent], + reaggregate: ReaggregateSpec, + ) -> None: + """ + Propagate tuning parameters from the reaggregate spec to aggregating components. + """ + configurable = [ + component + for component in components + if component.aggregation is not None and component.merge is not None + ] + if not configurable: + self._raise_unsupported_reaggregate_shape( + "parameterized reaggregation requires an aggregating component", + ) + for component in configurable: + component.params = dict(reaggregate.params or {}) + def _attach_reaggregate_spec( self, components: list[MetricComponent], diff --git a/datajunction-server/datajunction_server/sql/functions.py b/datajunction-server/datajunction_server/sql/functions.py index b194d46e4b..87c6891ae1 100644 --- a/datajunction-server/datajunction_server/sql/functions.py +++ b/datajunction-server/datajunction_server/sql/functions.py @@ -552,16 +552,16 @@ class ApproxPercentile(Function): def infer_type( col: ct.NumberType, percentage: ct.ListType, - accuracy: ct.NumberType | None, -) -> ct.DoubleType: + accuracy: ct.NumberType | None = None, +) -> ct.ListType: return ct.ListType(element_type=col.type) # type: ignore @ApproxPercentile.register def infer_type( col: ct.NumberType, - percentage: ct.FloatType, - accuracy: ct.NumberType | None, + percentage: ct.FloatingBase, + accuracy: ct.NumberType | None = None, ) -> ct.NumberType: return col.type # type: ignore diff --git a/datajunction-server/tests/api/cubes_build_metrics_spec_test.py b/datajunction-server/tests/api/cubes_build_metrics_spec_test.py index 64c14bb342..e196fa151e 100644 --- a/datajunction-server/tests/api/cubes_build_metrics_spec_test.py +++ b/datajunction-server/tests/api/cubes_build_metrics_spec_test.py @@ -4,7 +4,19 @@ This function builds Druid metricsSpec from measure columns and their components. """ +import pytest + from datajunction_server.api.cubes import _build_metrics_spec +from datajunction_server.models.cube_materialization import MetricMeasures +from datajunction_server.models.materialization import ( + DRUID_AGG_MAPPING, + DRUID_SKETCH_CONFIG, + DruidMeasuresCubeConfig, + Measure, + get_druid_aggregator_spec, + register_druid_aggregator, +) +from datajunction_server.errors import DJInvalidInputException from datajunction_server.models.decompose import AggregationRule, MetricComponent from datajunction_server.models.query import ColumnMetadata @@ -374,3 +386,293 @@ def test_preserves_column_order(self): result = _build_metrics_spec(columns, components, aliases) assert [r["name"] for r in result] == ["z_col", "a_col", "m_col"] + + +class TestDruidAggregatorSpec: + """ + Tests for the shared aggregator helper behind both metricsSpec builders. + """ + + def test_merge_is_preferred_over_accumulate(self): + """ + Ingestion merges stored partials, so the merge function keys the lookup. + """ + spec = get_druid_aggregator_spec( + column_name="orders_count", + column_type="bigint", + aggregation="COUNT", + merge="SUM", + ) + assert spec == { + "fieldName": "orders_count", + "name": "orders_count", + "type": "longSum", + } + + def test_falls_back_to_accumulate_without_merge(self): + """A component with no merge phase keys off its accumulate.""" + spec = get_druid_aggregator_spec( + column_name="revenue_sum", + column_type="double", + aggregation="SUM", + merge=None, + ) + assert spec is not None + assert spec["type"] == "doubleSum" + + def test_unmappable_returns_none(self): + """ + Return None for unmappable measures so callers determine fallback behavior. + """ + assert ( + get_druid_aggregator_spec( + column_name="x", + column_type="varchar", + aggregation="SUM", + merge="SUM", + ) + is None + ) + assert ( + get_druid_aggregator_spec( + column_name="x", + column_type=None, + aggregation="SUM", + merge="SUM", + ) + is None + ) + assert ( + get_druid_aggregator_spec( + column_name="x", + column_type="bigint", + aggregation=None, + merge=None, + ) + is None + ) + + def test_sketch_gets_default_config(self): + """An HLL sketch carries its precision defaults into the spec.""" + spec = get_druid_aggregator_spec( + column_name="accounts_hll", + column_type="binary", + aggregation="hll_sketch_agg", + merge="hll_union_agg", + ) + assert spec == { + "fieldName": "accounts_hll", + "name": "accounts_hll", + "type": "HLLSketchMerge", + "lgK": 12, + "tgtHllType": "HLL_4", + } + + def test_declared_params_override_defaults_partially(self): + """ + Metric-level params override defaults for specific keys while retaining others. + """ + spec = get_druid_aggregator_spec( + column_name="accounts_hll", + column_type="binary", + aggregation="hll_sketch_agg", + merge="hll_union_agg", + params={"lgK": 17}, + ) + assert spec is not None + assert spec["lgK"] == 17 + assert spec["tgtHllType"] == "HLL_4" + + def test_params_cannot_overwrite_structural_keys(self): + """ + `params` must not overwrite structural keys like `type` or `fieldName`. + """ + with pytest.raises(DJInvalidInputException) as excinfo: + get_druid_aggregator_spec( + column_name="accounts_hll", + column_type="binary", + aggregation="hll_sketch_agg", + merge="hll_union_agg", + params={"type": "doubleSum", "fieldName": "some_other_column"}, + ) + assert "does not accept" in str(excinfo.value) + assert "`fieldName`" in str(excinfo.value) + assert "`type`" in str(excinfo.value) + + def test_unknown_param_is_rejected(self): + """Unknown parameters are rejected loudly.""" + with pytest.raises(DJInvalidInputException) as excinfo: + get_druid_aggregator_spec( + column_name="accounts_hll", + column_type="binary", + aggregation="hll_sketch_agg", + merge="hll_union_agg", + params={"lgk": 17}, # lowercase k + ) + assert "lgk" in str(excinfo.value) + assert "lgK" in str(excinfo.value) # the supported spelling is named + + def test_params_rejected_for_unparameterized_aggregator(self): + """Params passed to an unparameterized aggregator raise an error.""" + with pytest.raises(DJInvalidInputException) as excinfo: + get_druid_aggregator_spec( + column_name="orders_count", + column_type="bigint", + aggregation="COUNT", + merge="SUM", + params={"compression": 200}, + ) + assert "takes no parameters" in str(excinfo.value) + + +class TestRegisterDruidAggregator: + """ + Tests for registering custom sketch families. + """ + + def test_register_and_use(self): + """A registered family is mappable and carries its default config.""" + try: + register_druid_aggregator( + column_type="binary", + merge_func="test_tdigest_agg", + aggregator="testTDigestSketch", + default_config={"compression": 200}, + ) + spec = get_druid_aggregator_spec( + column_name="latency_tdigest", + column_type="binary", + aggregation="test_tdigest", + merge="test_tdigest_agg", + ) + assert spec == { + "fieldName": "latency_tdigest", + "name": "latency_tdigest", + "type": "testTDigestSketch", + "compression": 200, + } + # A metric's own params win over the family default. + tuned = get_druid_aggregator_spec( + column_name="latency_tdigest", + column_type="binary", + aggregation="test_tdigest", + merge="test_tdigest_agg", + params={"compression": 1000}, + ) + assert tuned is not None + assert tuned["compression"] == 1000 + assert "testTDigestSketch" in DRUID_SKETCH_CONFIG + finally: + DRUID_AGG_MAPPING.pop(("binary", "test_tdigest_agg"), None) + DRUID_SKETCH_CONFIG.pop("testTDigestSketch", None) + + def test_register_without_config(self): + """A family with nothing to tune registers without extra config keys.""" + try: + register_druid_aggregator( + column_type="bigint", + merge_func="test_plain_merge", + aggregator="testPlainAgg", + ) + spec = get_druid_aggregator_spec( + column_name="c", + column_type="bigint", + aggregation="test_plain", + merge="test_plain_merge", + ) + assert spec == { + "fieldName": "c", + "name": "c", + "type": "testPlainAgg", + } + finally: + DRUID_AGG_MAPPING.pop(("bigint", "test_plain_merge"), None) + + def test_conflicting_defaults_are_rejected(self): + """ + Re-registering an aggregator with different defaults raises an error. + """ + try: + register_druid_aggregator( + column_type="binary", + merge_func="test_conflict_agg", + aggregator="testConflictSketch", + default_config={"compression": 100}, + ) + with pytest.raises(DJInvalidInputException) as excinfo: + register_druid_aggregator( + column_type="binary", + merge_func="test_conflict_agg", + aggregator="testConflictSketch", + default_config={"compression": 200}, + ) + assert "already registered with different defaults" in str(excinfo.value) + # Re-registering exactly what is there is a no-op. + register_druid_aggregator( + column_type="binary", + merge_func="test_conflict_agg", + aggregator="testConflictSketch", + default_config={"compression": 100}, + ) + assert DRUID_SKETCH_CONFIG["testConflictSketch"] == {"compression": 100} + finally: + DRUID_AGG_MAPPING.pop(("binary", "test_conflict_agg"), None) + DRUID_SKETCH_CONFIG.pop("testConflictSketch", None) + + def test_conflicting_aggregator_is_rejected(self): + """Two families cannot claim the same (column type, merge func) pair.""" + try: + register_druid_aggregator( + column_type="binary", + merge_func="test_claimed", + aggregator="testFirstSketch", + ) + with pytest.raises(DJInvalidInputException) as excinfo: + register_druid_aggregator( + column_type="binary", + merge_func="test_claimed", + aggregator="testSecondSketch", + ) + assert "already registered as" in str(excinfo.value) + finally: + DRUID_AGG_MAPPING.pop(("binary", "test_claimed"), None) + + def test_v2_measures_cube_config_gets_family_defaults(self): + """ + The V2 path uses the aggregator fallback with `merge=None` to pull correct specs. + """ + try: + register_druid_aggregator( + column_type="binary", + merge_func="test_v2_tdigest", + aggregator="testV2Sketch", + default_config={"compression": 200}, + ) + config = DruidMeasuresCubeConfig.model_construct( + dimensions=[], + measures={ + "some.metric": MetricMeasures.model_construct( + metric="some.metric", + combiner="x", + measures=[ + Measure( + name="latency_tdigest", + field_name="latency_tdigest", + agg="test_v2_tdigest", + type="binary", + ), + ], + ), + }, + ) + assert config.metrics_spec() == { + "latency_tdigest": { + "fieldName": "latency_tdigest", + "name": "latency_tdigest", + "type": "testV2Sketch", + "compression": 200, + }, + } + finally: + DRUID_AGG_MAPPING.pop(("binary", "test_v2_tdigest"), None) + DRUID_SKETCH_CONFIG.pop("testV2Sketch", None) diff --git a/datajunction-server/tests/api/cubes_test.py b/datajunction-server/tests/api/cubes_test.py index 28f3c993bf..410233d616 100644 --- a/datajunction-server/tests/api/cubes_test.py +++ b/datajunction-server/tests/api/cubes_test.py @@ -3651,6 +3651,7 @@ async def test_cube_materialization_metadata( "aggregation": "COUNT", "expression": "*", "grain_alias": None, + "params": None, "merge": "SUM", "name": "count_c8e42e74", "rule": { @@ -3664,6 +3665,7 @@ async def test_cube_materialization_metadata( "aggregation": "SUM", "expression": "if(discount > 0.0, 1, 0)", "grain_alias": None, + "params": None, "merge": "SUM", "name": "discount_sum_30b84e6c", "rule": { @@ -3677,6 +3679,7 @@ async def test_cube_materialization_metadata( "aggregation": "COUNT", "expression": "price", "grain_alias": None, + "params": None, "merge": "SUM", "name": "price_count_935e7117", "rule": { @@ -3690,6 +3693,7 @@ async def test_cube_materialization_metadata( "aggregation": "SUM", "expression": "price * discount", "grain_alias": None, + "params": None, "merge": "SUM", "name": "price_discount_sum_e4ba5456", "rule": { @@ -3703,6 +3707,7 @@ async def test_cube_materialization_metadata( "aggregation": "SUM", "expression": "price", "grain_alias": None, + "params": None, "merge": "SUM", "name": "price_sum_935e7117", "rule": { @@ -3716,6 +3721,7 @@ async def test_cube_materialization_metadata( "aggregation": "COUNT", "expression": "repair_order_id", "grain_alias": None, + "params": None, "merge": "SUM", "name": "repair_order_id_count_bd241964", "rule": { @@ -3729,6 +3735,7 @@ async def test_cube_materialization_metadata( "aggregation": "SUM", "expression": "total_repair_cost", "grain_alias": None, + "params": None, "merge": "SUM", "name": "total_repair_cost_sum_67874507", "rule": { @@ -3850,6 +3857,7 @@ async def test_cube_materialization_metadata( "aggregation": "SUM", "expression": "price", "grain_alias": None, + "params": None, "merge": "SUM", "name": "price_sum_252381cf", "rule": { @@ -4106,6 +4114,7 @@ async def test_cube_materialization_metadata( "name": "count_c8e42e74", "expression": "*", "grain_alias": None, + "params": None, "aggregation": "COUNT", "merge": "SUM", "rule": { @@ -4119,6 +4128,7 @@ async def test_cube_materialization_metadata( "name": "discount_sum_30b84e6c", "expression": "if(discount > 0.0, 1, 0)", "grain_alias": None, + "params": None, "aggregation": "SUM", "merge": "SUM", "rule": { @@ -4132,6 +4142,7 @@ async def test_cube_materialization_metadata( "name": "price_count_935e7117", "expression": "price", "grain_alias": None, + "params": None, "aggregation": "COUNT", "merge": "SUM", "rule": { @@ -4145,6 +4156,7 @@ async def test_cube_materialization_metadata( "name": "price_discount_sum_e4ba5456", "expression": "price * discount", "grain_alias": None, + "params": None, "aggregation": "SUM", "merge": "SUM", "rule": { @@ -4158,6 +4170,7 @@ async def test_cube_materialization_metadata( "name": "price_sum_935e7117", "expression": "price", "grain_alias": None, + "params": None, "aggregation": "SUM", "merge": "SUM", "rule": { @@ -4171,6 +4184,7 @@ async def test_cube_materialization_metadata( "name": "repair_order_id_count_bd241964", "expression": "repair_order_id", "grain_alias": None, + "params": None, "aggregation": "COUNT", "merge": "SUM", "rule": { @@ -4184,6 +4198,7 @@ async def test_cube_materialization_metadata( "name": "total_repair_cost_sum_67874507", "expression": "total_repair_cost", "grain_alias": None, + "params": None, "aggregation": "SUM", "merge": "SUM", "rule": { @@ -4197,6 +4212,7 @@ async def test_cube_materialization_metadata( "name": "price_sum_252381cf", "expression": "price", "grain_alias": None, + "params": None, "aggregation": "SUM", "merge": "SUM", "rule": { diff --git a/datajunction-server/tests/api/metrics_test.py b/datajunction-server/tests/api/metrics_test.py index d3313ecde0..90b2bfc0e4 100644 --- a/datajunction-server/tests/api/metrics_test.py +++ b/datajunction-server/tests/api/metrics_test.py @@ -465,6 +465,7 @@ async def test_read_metrics(module__client_with_roads: AsyncClient) -> None: "aggregation": "SUM", "expression": "if(discount > 0.0, 1, 0)", "grain_alias": None, + "params": None, "name": "discount_sum_30b84e6c", "merge": "SUM", "rule": { @@ -478,6 +479,7 @@ async def test_read_metrics(module__client_with_roads: AsyncClient) -> None: "aggregation": "COUNT", "expression": "*", "grain_alias": None, + "params": None, "merge": "SUM", "name": "count_c8e42e74", "rule": { @@ -526,6 +528,7 @@ async def test_metric_reaggregate_roundtrip_and_validation( ) assert response.status_code in (200, 201), response.json() assert response.json()["reaggregate"] == { + "params": None, "rules": [ { "dimension": "default.repair_orders_fact.repair_order_id", @@ -537,6 +540,7 @@ async def test_metric_reaggregate_roundtrip_and_validation( response = await client_with_roads.get(f"/nodes/{metric_name}/") assert response.status_code == 200 assert response.json()["reaggregate"] == { + "params": None, "rules": [ { "dimension": "default.repair_orders_fact.repair_order_id", @@ -548,6 +552,7 @@ async def test_metric_reaggregate_roundtrip_and_validation( response = await client_with_roads.get(f"/metrics/{metric_name}/") assert response.status_code == 200 assert response.json()["reaggregate"] == { + "params": None, "rules": [ { "dimension": "default.repair_orders_fact.repair_order_id", diff --git a/datajunction-server/tests/api/preaggregations_test.py b/datajunction-server/tests/api/preaggregations_test.py index 1c13e10e2b..5d73a6b993 100644 --- a/datajunction-server/tests/api/preaggregations_test.py +++ b/datajunction-server/tests/api/preaggregations_test.py @@ -736,6 +736,7 @@ async def test_get_preagg_by_id(self, client_with_preaggs): "expr_hash": "83632b779d87", "expression": "line_total", "grain_alias": None, + "params": None, "merge": "SUM", "name": "line_total_sum_e1f61696", "source_column": None, @@ -805,6 +806,7 @@ async def test_get_preagg_by_id(self, client_with_preaggs): "expr_hash": "221d2a4bfdae", "expression": "quantity", "grain_alias": None, + "params": None, "merge": "SUM", "name": "quantity_sum_06b64d2e", "source_column": None, diff --git a/datajunction-server/tests/api/sql_v2_test.py b/datajunction-server/tests/api/sql_v2_test.py index 704c6e9ef5..659f3fb919 100644 --- a/datajunction-server/tests/api/sql_v2_test.py +++ b/datajunction-server/tests/api/sql_v2_test.py @@ -1201,6 +1201,7 @@ async def create_metric_distinct_single_column(client: AsyncClient): "aggregation": None, "expression": "hard_hat_id", "grain_alias": "hard_hat_id", + "params": None, "merge": None, "name": "hard_hat_id", "rule": { @@ -1233,6 +1234,7 @@ async def create_metric_distinct_expression(client: AsyncClient): "aggregation": None, "expression": "IF(hard_hat_id = 1, 1, 0)", "grain_alias": "hard_hat_id_distinct_0291ee39", + "params": None, "merge": None, "name": "hard_hat_id_distinct_0291ee39", "rule": { @@ -1633,6 +1635,7 @@ async def test_metric_definitions_with_nonjoinable_dimensions( "aggregation": "SUM", "expression": "default.local_hard_hats_2.hard_hat_id", "grain_alias": None, + "params": None, "merge": "SUM", "name": "default_DOT_local_hard_hats_2_DOT_hard_hat_id_sum_bf8a8419", "rule": { @@ -1699,6 +1702,7 @@ async def test_metric_definitions_with_single_joinable_dimensions( "aggregation": None, "expression": "default.municipality_dim.contact_name", "grain_alias": "contact_name", + "params": None, "merge": None, "name": "contact_name", "rule": { @@ -1806,6 +1810,7 @@ async def test_metric_definition_with_multiple_joinable_dimensions( "expression": "IF(default.hard_hat.state = 'NY', default.hard_hat.first_name, " "NULL)", "grain_alias": "default_DOT_hard_hat_DOT_state_default_DOT_hard_hat_DOT_first_name_distinct_1a99d6a7", + "params": None, "merge": None, "name": "default_DOT_hard_hat_DOT_state_default_DOT_hard_hat_DOT_first_name_distinct_1a99d6a7", "rule": { diff --git a/datajunction-server/tests/construction/build_v3/preagg_matcher_test.py b/datajunction-server/tests/construction/build_v3/preagg_matcher_test.py index 294428b821..e88dfb1ec2 100644 --- a/datajunction-server/tests/construction/build_v3/preagg_matcher_test.py +++ b/datajunction-server/tests/construction/build_v3/preagg_matcher_test.py @@ -32,6 +32,7 @@ from datajunction_server.database.preaggregation import ( PreAggregation, compute_expression_hash, + measure_identity_token, ) from datajunction_server.database.user import User from datajunction_server.errors import DJInvalidInputException @@ -220,7 +221,7 @@ async def test_returns_identities_for_all_components( parent_node: Node, metric_node: Node, ): - """Should return an (expression hash, aggregation) pair per component.""" + """Should return one identity token per component.""" components = [ (metric_node, make_component("sum_revenue", "price * quantity")), (metric_node, make_component("sum_quantity", "quantity")), @@ -230,8 +231,8 @@ async def test_returns_identities_for_all_components( identities = get_required_measure_identities(grain_group) assert identities == { - (compute_expression_hash("price * quantity"), "SUM"), - (compute_expression_hash("quantity"), "SUM"), + measure_identity_token(compute_expression_hash("price * quantity"), "SUM"), + measure_identity_token(compute_expression_hash("quantity"), "SUM"), } @pytest.mark.asyncio @@ -264,7 +265,9 @@ async def test_deduplicates_same_expression_and_aggregation( identities = get_required_measure_identities(grain_group) - assert identities == {(compute_expression_hash("price * quantity"), "SUM")} + assert identities == { + measure_identity_token(compute_expression_hash("price * quantity"), "SUM"), + } @pytest.mark.asyncio async def test_same_expression_different_aggregation_are_distinct( @@ -294,8 +297,8 @@ async def test_same_expression_different_aggregation_are_distinct( identities = get_required_measure_identities(grain_group) assert identities == { - (compute_expression_hash("unit_price"), "SUM"), - (compute_expression_hash("unit_price"), "MAX"), + measure_identity_token(compute_expression_hash("unit_price"), "SUM"), + measure_identity_token(compute_expression_hash("unit_price"), "MAX"), } diff --git a/datajunction-server/tests/database/preaggregation_test.py b/datajunction-server/tests/database/preaggregation_test.py index 0364ab65f3..fbe61eaea4 100644 --- a/datajunction-server/tests/database/preaggregation_test.py +++ b/datajunction-server/tests/database/preaggregation_test.py @@ -716,3 +716,81 @@ async def test_find_matching_no_candidates( }, ) assert result is None + + +class TestMeasureIdentityParams: + """ + Tuning parameters participate in measure identity. + """ + + def test_existing_tokens_are_unchanged(self): + """ + A component without params produces the token byte-for-byte. + """ + assert measure_identity_token("abc123", "SUM") == "abc123:SUM" + assert measure_identity_token("abc123", "SUM", None) == "abc123:SUM" + assert measure_identity_token("abc123", "SUM", {}) == "abc123:SUM" + assert measure_identity_token("abc123", None) == "abc123:" + assert measure_identity_token("abc123", " count ") == "abc123:COUNT" + + def test_differing_params_differ(self): + """Two sketches over one column at different accuracies are distinct.""" + low = measure_identity_token("abc123", "TDIGEST_AGG", {"compression": 100}) + high = measure_identity_token("abc123", "TDIGEST_AGG", {"compression": 1000}) + assert low != high + # ...and neither collides with the un-parameterized token. + assert low != measure_identity_token("abc123", "TDIGEST_AGG") + + def test_params_are_order_and_spelling_insensitive(self): + """ + One sketch declaration yields one identity. Key order and type variations (e.g. 200 vs 200.0) do not change the identity token. + """ + assert measure_identity_token( + "abc123", + "TDIGEST_AGG", + {"compression": 200, "k": 12}, + ) == measure_identity_token( + "abc123", + "TDIGEST_AGG", + {"k": 12, "compression": 200}, + ) + assert measure_identity_token( + "abc123", + "TDIGEST_AGG", + {"compression": 200.0}, + ) == measure_identity_token("abc123", "TDIGEST_AGG", {"compression": 200}) + # Fractional values are preserved. + assert measure_identity_token( + "abc123", + "TDIGEST_AGG", + {"compression": 200.5}, + ) != measure_identity_token("abc123", "TDIGEST_AGG", {"compression": 200}) + + def test_get_measure_identities_carries_params(self): + """Stored measures round-trip their params into the identity set.""" + plain = make_measure("latency_sum", "latency_ms") + sketched = make_measure( + "latency_tdigest", + "latency_ms", + aggregation="TDIGEST_AGG", + merge="TDIGEST_AGG", + ) + sketched.params = {"compression": 200} + identities = get_measure_identities([plain, sketched]) + assert identities == { + measure_identity_token(compute_expression_hash("latency_ms"), "SUM"), + measure_identity_token( + compute_expression_hash("latency_ms"), + "TDIGEST_AGG", + {"compression": 200}, + ), + } + # Ensure the sketch identity distinguishes different compressions. + assert ( + measure_identity_token( + compute_expression_hash("latency_ms"), + "TDIGEST_AGG", + {"compression": 1000}, + ) + not in identities + ) 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 84f2f491f9..eb92ce22ce 100644 --- a/datajunction-server/tests/internal/nodes/derive_frozen_measures_test.py +++ b/datajunction-server/tests/internal/nodes/derive_frozen_measures_test.py @@ -372,6 +372,79 @@ def test_frozen_measure_conflict_rejects_different_measure_identity(): _raise_if_frozen_measure_conflicts(frozen_measure, measure) +def test_frozen_measure_conflict_rejects_different_tuning_params(): + """ + A name collision fails when only the sketch tuning parameters differ. + + Component names are hashed from the expression and its source, not from + params, so a p95 at compression=200 and one at compression=1000 over the + same column collide on name. Without this the second metric would silently + bind the frozen measure -- and the materialized sketch -- built at the + first one's accuracy. + """ + frozen_measure = FrozenMeasure( + name="latency_tdigest", + upstream_revision_id=1, + expression="latency_ms", + aggregation="TDIGEST_AGG", + params={"compression": 200}, + rule=AggregationRule(type=Aggregability.FULL), + ) + measure = MetricComponent( + name="latency_tdigest", + expression="latency_ms", + aggregation="TDIGEST_AGG", + params={"compression": 1000}, + rule=AggregationRule(type=Aggregability.FULL), + ) + + with pytest.raises(DJInvalidInputException, match="tuning parameters"): + _raise_if_frozen_measure_conflicts(frozen_measure, measure) + + +def test_frozen_measure_reuse_allows_matching_params(): + """ + Identical params are reusable, and absent-vs-empty is not a difference. + + Almost every measure has no params at all; a stored ``None`` meeting a + freshly-extracted ``{}`` must not read as a conflict and start rejecting + ordinary metrics. + """ + frozen_measure = FrozenMeasure( + name="latency_tdigest", + upstream_revision_id=1, + expression="latency_ms", + aggregation="TDIGEST_AGG", + params={"compression": 200}, + rule=AggregationRule(type=Aggregability.FULL), + ) + measure = MetricComponent( + name="latency_tdigest", + expression="latency_ms", + aggregation="TDIGEST_AGG", + params={"compression": 200}, + rule=AggregationRule(type=Aggregability.FULL), + ) + _raise_if_frozen_measure_conflicts(frozen_measure, measure) + + unparameterized = FrozenMeasure( + name="amount_sum", + upstream_revision_id=1, + expression="amount", + aggregation="SUM", + params=None, + rule=AggregationRule(type=Aggregability.FULL), + ) + empty_params = MetricComponent( + name="amount_sum", + expression="amount", + aggregation="SUM", + params={}, + rule=AggregationRule(type=Aggregability.FULL), + ) + _raise_if_frozen_measure_conflicts(unparameterized, empty_params) + + @pytest.mark.asyncio async def test_derived_metric_expands_parent_cache( session: AsyncSession, diff --git a/datajunction-server/tests/models/reaggregate_test.py b/datajunction-server/tests/models/reaggregate_test.py index 60cc67b9f8..8d225f6ca4 100644 --- a/datajunction-server/tests/models/reaggregate_test.py +++ b/datajunction-server/tests/models/reaggregate_test.py @@ -28,6 +28,7 @@ def test_dump_reaggregate_spec_from_dict(): ], }, ) == { + "params": None, "rules": [ { "dimension": "default.date_dim.date", @@ -51,6 +52,7 @@ def test_dump_reaggregate_spec_from_model(): ], ), ) == { + "params": None, "rules": [ { "dimension": "default.date_dim.date", @@ -131,3 +133,22 @@ def test_unsupported_dimension_reaggregate_functions_handles_empty_and_invalid() ], }, ) == ["sum"] + + +def test_empty_params_allowed(): + """ + An absent or empty `params` declaration is preserved. + """ + assert ReaggregateSpec().params is None + assert ReaggregateSpec(params={}).params == {} + + +def test_params_round_trip_through_parse(): + """ + `params` survives dict -> model -> dict without loss. + """ + spec = parse_reaggregate_spec({"params": {"compression": 200}}) + assert spec.params == {"compression": 200} + assert parse_reaggregate_spec(dump_reaggregate_spec(spec)).params == { + "compression": 200, + } diff --git a/datajunction-server/tests/sql/decompose_test.py b/datajunction-server/tests/sql/decompose_test.py index 96d2460af7..c1fe86c68c 100644 --- a/datajunction-server/tests/sql/decompose_test.py +++ b/datajunction-server/tests/sql/decompose_test.py @@ -20,6 +20,7 @@ from datajunction_server.models.node_type import NodeType from datajunction_server.models.reaggregate import ( DimensionReaggregateRule, + ReaggregateSpec, ReaggregationFunction, ) from datajunction_server.sql import functions as dj_functions @@ -2824,3 +2825,74 @@ def test_fixed_grain_refused_on_a_component_that_cannot_re_merge(kwargs): assert "fully-aggregatable" in str(exc.value) assert components[0].rule.fixed_grain is None + + +class TestReaggregateParams: + """ + Tests that sketch tuning parameters flow through to components. + """ + + @staticmethod + def _spec(params): + """A spec carrying params.""" + return ReaggregateSpec.model_construct( + fn=ReaggregationFunction.SUM, + weight=None, + rules=[], + params=params, + ) + + def test_params_reach_components(self): + extractor = MetricComponentExtractor(1) + components, _ = extractor._extract_base( + parse("SELECT SUM(latency_ms) FROM t"), + self._spec({"compression": 200}), + ) + assert [component.params for component in components] == [ + {"compression": 200}, + ] + + def test_params_reach_every_component_of_a_multi_component_metric(self): + """AVG decomposes to a sum and a count; both are the same sketch family.""" + extractor = MetricComponentExtractor(1) + components, _ = extractor._extract_base( + parse("SELECT AVG(latency_ms) FROM t"), + self._spec({"compression": 200}), + ) + assert len(components) == 2 + assert all(c.params == {"compression": 200} for c in components) + + def test_params_are_copied_not_shared(self): + """ + Each component receives an independent copy of its params dict. + """ + extractor = MetricComponentExtractor(1) + params = {"compression": 200} + components, _ = extractor._extract_base( + parse("SELECT AVG(latency_ms) FROM t"), + self._spec(params), + ) + params["compression"] = 999 + assert all(c.params == {"compression": 200} for c in components) + assert components[0].params is not components[1].params + + def test_no_params_leaves_components_untouched(self): + """Omitting params applies no configuration to components.""" + extractor = MetricComponentExtractor(1) + components, _ = extractor._extract_base( + parse("SELECT SUM(latency_ms) FROM t"), + self._spec(None), + ) + assert [component.params for component in components] == [None] + + def test_params_without_an_aggregating_component_is_rejected(self): + """ + Declaring params on a non-aggregating metric is rejected. + """ + extractor = MetricComponentExtractor(1) + with pytest.raises(DJInvalidInputException) as excinfo: + extractor._extract_base( + parse("SELECT COUNT(DISTINCT order_id) FROM t"), + self._spec({"compression": 200}), + ) + assert "requires an aggregating component" in str(excinfo.value) diff --git a/datajunction-server/tests/sql/functions_test.py b/datajunction-server/tests/sql/functions_test.py index 10a1290df2..964d2815b8 100644 --- a/datajunction-server/tests/sql/functions_test.py +++ b/datajunction-server/tests/sql/functions_test.py @@ -288,6 +288,33 @@ async def test_approx_percentile(session: AsyncSession): assert not exc.errors assert query_with_list.select.projection[0].type == ct.FloatType() # type: ignore + # The two-argument form -- `accuracy` is optional in Spark, and this is the + # spelling metric authors actually use. + query = parse("SELECT approx_percentile(10.0, 0.5)") + exc = DJException() + ctx = ast.CompileContext(session=session, exception=exc) + await query.compile(ctx) + assert not exc.errors + assert query.select.projection[0].type == ct.FloatType() # type: ignore + + query = parse("SELECT approx_percentile(10.0, array(0.5, 0.9))") + exc = DJException() + ctx = ast.CompileContext(session=session, exception=exc) + await query.compile(ctx) + assert not exc.errors + assert query.select.projection[0].type == ct.ListType( # type: ignore + element_type=ct.FloatType(), + ) + + # A double-typed percentage dispatches too: DoubleType and FloatType are + # siblings under FloatingBase, so a FloatType-only registration missed it. + query = parse("SELECT approx_percentile(10.0, CAST(0.5 AS DOUBLE))") + exc = DJException() + ctx = ast.CompileContext(session=session, exception=exc) + await query.compile(ctx) + assert not exc.errors + assert query.select.projection[0].type == ct.FloatType() # type: ignore + @pytest.mark.asyncio async def test_array(session: AsyncSession): From ab13b4e8c1a8c3be5c44a45731ce91a6121fbbff Mon Sep 17 00:00:00 2001 From: Robin Davis Date: Tue, 22 Sep 2026 12:23:55 -0700 Subject: [PATCH 2/4] Use a placeholder sketch function name in doc examples --- .../datajunction_server/database/preaggregation.py | 2 +- .../datajunction_server/models/materialization.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/datajunction-server/datajunction_server/database/preaggregation.py b/datajunction-server/datajunction_server/database/preaggregation.py index 32ecfd1cc5..937dd051d2 100644 --- a/datajunction-server/datajunction_server/database/preaggregation.py +++ b/datajunction-server/datajunction_server/database/preaggregation.py @@ -125,7 +125,7 @@ def measure_identity_token( Parameters extend that reasoning to sketches: a t-digest accumulated at ``compression=100`` is not interchangeable with one at ``compression=1000``, - though both are ``nflx_tdigest_agg`` over the same expression. The segment is + though both are ``acme_tdigest_agg`` over the same expression. The segment is appended **only** when params are present, so tokens for the overwhelming majority of components -- plain ``SUM``, ``COUNT``, ``MAX`` -- are byte-identical to what this returned before params existed. Stored ``preagg_hash`` values are diff --git a/datajunction-server/datajunction_server/models/materialization.py b/datajunction-server/datajunction_server/models/materialization.py index 1e763e6367..e4dc53c6f8 100644 --- a/datajunction-server/datajunction_server/models/materialization.py +++ b/datajunction-server/datajunction_server/models/materialization.py @@ -77,7 +77,7 @@ def register_druid_aggregator( Args: column_type: Measures-table column type, e.g. "binary" - merge_func: Phase-2 merge function name, e.g. "nflx_tdigest_agg" + merge_func: Phase-2 merge function name, e.g. "acme_tdigest_agg" aggregator: Druid aggregator type, e.g. "tDigestSketch" default_config: Extra metricsSpec keys and their defaults, e.g. ``{"compression": 200}``. Marks the aggregator as parameterized. From 15e346d757d227a7c3f9d379144bd0e3c7d832a6 Mon Sep 17 00:00:00 2001 From: Beto Dealmeida Date: Wed, 23 Sep 2026 21:37:08 -0300 Subject: [PATCH 3/4] Adapt reaggregation params to rules-only model --- ...6_09_15_0000-fm0001params_add_params_to_frozen_measures.py | 4 ++-- .../datajunction_server/api/graphql/scalars/metricmetadata.py | 2 -- .../datajunction_server/models/materialization.py | 2 +- datajunction-server/tests/sql/decompose_test.py | 2 -- 4 files changed, 3 insertions(+), 7 deletions(-) diff --git a/datajunction-server/datajunction_server/alembic/versions/2026_09_15_0000-fm0001params_add_params_to_frozen_measures.py b/datajunction-server/datajunction_server/alembic/versions/2026_09_15_0000-fm0001params_add_params_to_frozen_measures.py index c00473a556..09a6701631 100644 --- a/datajunction-server/datajunction_server/alembic/versions/2026_09_15_0000-fm0001params_add_params_to_frozen_measures.py +++ b/datajunction-server/datajunction_server/alembic/versions/2026_09_15_0000-fm0001params_add_params_to_frozen_measures.py @@ -4,7 +4,7 @@ Adds a JSON column to persist tuning parameters (e.g., accuracy) for sketch-backed measures. Revision ID: fm0001params -Revises: rg0001reaggregate +Revises: fg0001fixedgrain Create Date: 2026-09-15 00:00:00.000000+00:00 """ @@ -13,7 +13,7 @@ # revision identifiers, used by Alembic. revision = "fm0001params" -down_revision = "rg0001reaggregate" +down_revision = "fg0001fixedgrain" branch_labels = None depends_on = None diff --git a/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py b/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py index 802cb6ce19..04876d4980 100644 --- a/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py +++ b/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py @@ -55,8 +55,6 @@ class ReaggregateSpec: is an open dict, which has no automatic GraphQL mapping. """ - fn: strawberry.auto - weight: strawberry.auto rules: strawberry.auto params: JSON | None = None diff --git a/datajunction-server/datajunction_server/models/materialization.py b/datajunction-server/datajunction_server/models/materialization.py index e4dc53c6f8..c1e8b45a6f 100644 --- a/datajunction-server/datajunction_server/models/materialization.py +++ b/datajunction-server/datajunction_server/models/materialization.py @@ -56,7 +56,7 @@ # Aggregator types that carry extra config in the Druid metricsSpec, mapped to # their default configuration. A sketch aggregator is not fully specified by its -# type: an HLL needs a precision, a t-digest a compression, a KLL a `k`. Defaults +# kind: an HLL needs a precision, a t-digest a compression, a KLL a `k`. Defaults # apply when a metric declares no `reaggregate.params`. DRUID_SKETCH_CONFIG: dict[str, dict[str, Any]] = { "HLLSketchMerge": { diff --git a/datajunction-server/tests/sql/decompose_test.py b/datajunction-server/tests/sql/decompose_test.py index c1fe86c68c..8f8db41c61 100644 --- a/datajunction-server/tests/sql/decompose_test.py +++ b/datajunction-server/tests/sql/decompose_test.py @@ -2836,8 +2836,6 @@ class TestReaggregateParams: def _spec(params): """A spec carrying params.""" return ReaggregateSpec.model_construct( - fn=ReaggregationFunction.SUM, - weight=None, rules=[], params=params, ) From 85d77a8f303e4df15d75afffd765604dd1772538 Mon Sep 17 00:00:00 2001 From: Beto Dealmeida Date: Mon, 28 Sep 2026 17:41:09 -0400 Subject: [PATCH 4/4] Fix reaggregation migration ancestry --- ...6_09_15_0000-fm0001params_add_params_to_frozen_measures.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/datajunction-server/datajunction_server/alembic/versions/2026_09_15_0000-fm0001params_add_params_to_frozen_measures.py b/datajunction-server/datajunction_server/alembic/versions/2026_09_15_0000-fm0001params_add_params_to_frozen_measures.py index 09a6701631..5d56d7887e 100644 --- a/datajunction-server/datajunction_server/alembic/versions/2026_09_15_0000-fm0001params_add_params_to_frozen_measures.py +++ b/datajunction-server/datajunction_server/alembic/versions/2026_09_15_0000-fm0001params_add_params_to_frozen_measures.py @@ -4,7 +4,7 @@ Adds a JSON column to persist tuning parameters (e.g., accuracy) for sketch-backed measures. Revision ID: fm0001params -Revises: fg0001fixedgrain +Revises: dv0001typed Create Date: 2026-09-15 00:00:00.000000+00:00 """ @@ -13,7 +13,7 @@ # revision identifiers, used by Alembic. revision = "fm0001params" -down_revision = "fg0001fixedgrain" +down_revision = "dv0001typed" branch_labels = None depends_on = None