From 079ba21adf851fe69e97d54bb717a8492bd63fc4 Mon Sep 17 00:00:00 2001 From: Robin Davis Date: Fri, 18 Sep 2026 22:01:35 -0700 Subject: [PATCH 01/17] Make metric decomposition dialect-aware Pass the output dialect and source aggregation call through metric decomposition so dialect-specific combiners can be rendered correctly. Add family-based registration so sketch decompositions can be selected explicitly through a metric's reaggregation spec. --- .../construction/build_v3/decomposition.py | 11 +- .../datajunction_server/models/reaggregate.py | 26 +- .../datajunction_server/sql/decompose.py | 225 ++++++++++-- .../tests/models/reaggregate_test.py | 38 +- .../tests/sql/decompose_test.py | 331 ++++++++++++++++++ 5 files changed, 600 insertions(+), 31 deletions(-) diff --git a/datajunction-server/datajunction_server/construction/build_v3/decomposition.py b/datajunction-server/datajunction_server/construction/build_v3/decomposition.py index 273d3cb4c..313d44aac 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/decomposition.py +++ b/datajunction-server/datajunction_server/construction/build_v3/decomposition.py @@ -27,6 +27,7 @@ from datajunction_server.database.node import Node from datajunction_server.errors import DJInvalidInputException from datajunction_server.models.decompose import Aggregability, MetricComponent +from datajunction_server.models.dialect import Dialect from datajunction_server.models.node_type import NodeType from datajunction_server.sql.decompose import MetricComponentExtractor from datajunction_server.sql.parsing import ast @@ -79,6 +80,7 @@ async def decompose_and_group_metrics( base_metric, nodes_cache=ctx.nodes, parent_map=ctx.parent_map, + dialect=ctx.dialect, ) all_decomposed[base_metric.name] = decomposed @@ -98,6 +100,7 @@ async def decompose_and_group_metrics( metric_node, nodes_cache=ctx.nodes, parent_map=ctx.parent_map, + dialect=ctx.dialect, ) all_decomposed[metric_name] = derived_decomposed else: @@ -121,6 +124,7 @@ async def decompose_and_group_metrics( metric_node, nodes_cache=ctx.nodes, parent_map=ctx.parent_map, + dialect=ctx.dialect, ) all_decomposed[metric_node.name] = decomposed @@ -146,6 +150,7 @@ async def decompose_metric( *, nodes_cache: dict[str, Node] | None = None, parent_map: dict[str, list[str]] | None = None, + dialect: Dialect = Dialect.SPARK, ) -> DecomposedMetricInfo: """ Decompose a metric into its constituent components. @@ -162,6 +167,10 @@ async def decompose_metric( parent_map, avoids database queries by using cached data. parent_map: Optional dict of child_name -> list of parent_names. Required if nodes_cache is provided. + dialect: Dialect the combiner is rendered for. ``decompose_and_group_metrics`` + passes the build's resolved dialect, which for a materialized cube is + the one taken from its availability catalog -- Druid, in practice. + Defaults to Spark for direct callers. Returns: DecomposedMetricInfo with components, combiner expression, and aggregability @@ -172,7 +181,7 @@ async def decompose_metric( ) # Use the MetricComponentExtractor with optional cache - extractor = MetricComponentExtractor(metric_node.current.id) + extractor = MetricComponentExtractor(metric_node.current.id, dialect=dialect) components, derived_ast = await extractor.extract( session, nodes_cache=nodes_cache, diff --git a/datajunction-server/datajunction_server/models/reaggregate.py b/datajunction-server/datajunction_server/models/reaggregate.py index 9ff2da446..14ad0b759 100644 --- a/datajunction-server/datajunction_server/models/reaggregate.py +++ b/datajunction-server/datajunction_server/models/reaggregate.py @@ -20,6 +20,10 @@ class ReaggregationFunction(StrEnum): FIRST_VALUE = "first_value" MIN = "min" MAX = "max" + # Quantile sketch family. Unlike the functions above it does not describe a + # rollup arithmetic; it selects how a quantile metric is accumulated, merged + # and read back, which is what makes percentiles pre-aggregatable at all. + TDIGEST = "tdigest" DIMENSION_REAGGREGATE_FUNCTIONS = frozenset( @@ -32,6 +36,23 @@ class ReaggregationFunction(StrEnum): ) +PARAMETERIZED_REAGGREGATE_FUNCTIONS: frozenset[ReaggregationFunction] = frozenset( + { + # compression: centroids retained, trading sketch size for tail accuracy. + ReaggregationFunction.TDIGEST, + }, +) + + +def is_parameterized_reaggregate_function( + function: ReaggregationFunction | None, +) -> bool: + """ + Return whether a function accepts tuning parameters in `params`. + """ + return function in PARAMETERIZED_REAGGREGATE_FUNCTIONS + + def is_supported_dimension_reaggregate_function( function: ReaggregationFunction, ) -> bool: @@ -77,13 +98,16 @@ class ReaggregateSpec(BaseModel): model_config = ConfigDict(extra="forbid") + # A sketch family selects an alternate decomposition for the metric's + # aggregate expression. Ordinary dimension-specific behavior remains in + # `rules` below. + fn: ReaggregationFunction | None = None 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, ) -> dict | None: diff --git a/datajunction-server/datajunction_server/sql/decompose.py b/datajunction-server/datajunction_server/sql/decompose.py index cb831479e..b26c0e81b 100644 --- a/datajunction-server/datajunction_server/sql/decompose.py +++ b/datajunction-server/datajunction_server/sql/decompose.py @@ -3,7 +3,7 @@ import hashlib from abc import ABC, abstractmethod from dataclasses import dataclass -from typing import cast +from typing import Any, cast from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession @@ -16,9 +16,11 @@ AggregationRule, MetricComponent, ) +from datajunction_server.models.dialect import Dialect from datajunction_server.models.node_type import NodeType from datajunction_server.models.reaggregate import ( ReaggregateSpec, + ReaggregationFunction, dimension_reaggregate_rules, is_supported_dimension_reaggregate_function, parse_reaggregate_spec, @@ -122,8 +124,22 @@ class AggDecomposition(ABC): def components(self) -> list[ComponentDef]: """Define the components needed for this aggregation.""" + def __init__(self, params: dict[str, Any] | None = None): + """ + Args: + params: Tuning parameters from the metric's ``reaggregate.params``. + Empty for the registry-by-function decompositions, which take + none; a sketch family reads its accuracy setting from here. + """ + self.params = params or {} + @abstractmethod - def combine(self, components: list[MetricComponent]) -> ast.Expression: + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ) -> ast.Expression: """Build the combiner expression from merged metric components.""" @@ -134,6 +150,30 @@ def combine(self, components: list[MetricComponent]) -> ast.Expression: DECOMPOSITION_REGISTRY: dict[type, type[AggDecomposition] | None] = {} +# Decompositions selected by the metric's declared reaggregation family rather +# than by its aggregation function. A family entry wins over the by-function +# registry, which is what lets `APPROX_PERCENTILE(x, 0.95)` keep its spelling and +# still decompose, but only for metrics that opted in with `reaggregate.fn`. +FAMILY_DECOMPOSITION_REGISTRY: dict[ReaggregationFunction, type[AggDecomposition]] = {} + + +def decomposes_family(fn: ReaggregationFunction): + """ + Register a decomposition for a reaggregation family. + + Exported for downstream deployments, which supply the engine-specific + functions a family needs. OSS registers none, so a metric declaring an + unregistered family falls through to the by-function registry and keeps the + aggregability it has today. + """ + + def decorator(decomp_class: type[AggDecomposition]): + FAMILY_DECOMPOSITION_REGISTRY[fn] = decomp_class + return decomp_class + + return decorator + + def decomposes(func_class: type): """Decorator to register a decomposition class for a function.""" @@ -162,7 +202,12 @@ class SumDecomposition(AggDecomposition): def components(self) -> list[ComponentDef]: return [ComponentDef("_sum", "SUM", "SUM")] - def combine(self, components: list[MetricComponent]): + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ): return make_func("SUM", components[0].name) @@ -174,7 +219,12 @@ class MaxDecomposition(AggDecomposition): def components(self) -> list[ComponentDef]: return [ComponentDef("_max", "MAX", "MAX")] - def combine(self, components: list[MetricComponent]): + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ): return make_func("MAX", components[0].name) @@ -186,7 +236,12 @@ class MinDecomposition(AggDecomposition): def components(self) -> list[ComponentDef]: return [ComponentDef("_min", "MIN", "MIN")] - def combine(self, components: list[MetricComponent]): + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ): return make_func("MIN", components[0].name) @@ -198,7 +253,12 @@ class AnyValueDecomposition(AggDecomposition): def components(self) -> list[ComponentDef]: return [ComponentDef("_any_value", "ANY_VALUE", "ANY_VALUE")] - def combine(self, components: list[MetricComponent]): + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ): return make_func("ANY_VALUE", components[0].name) @@ -215,7 +275,12 @@ class CountDecomposition(AggDecomposition): def components(self) -> list[ComponentDef]: return [ComponentDef("_count", "COUNT", "SUM")] - def combine(self, components: list[MetricComponent]): + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ): return make_func("SUM", components[0].name) @@ -227,7 +292,12 @@ class CountIfDecomposition(AggDecomposition): def components(self) -> list[ComponentDef]: return [ComponentDef("_count_if", "COUNT_IF", "SUM")] - def combine(self, components: list[MetricComponent]): + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ): return make_func("SUM", components[0].name) @@ -247,7 +317,12 @@ def components(self) -> list[ComponentDef]: ComponentDef("_count", "COUNT", "SUM"), ] - def combine(self, components: list[MetricComponent]): + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ): return ast.BinaryOp( op=ast.BinaryOpKind.Divide, left=make_func("SUM", components[0].name), @@ -279,7 +354,12 @@ def components(self) -> list[ComponentDef]: ), ] - def combine(self, components: list[MetricComponent]): + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ): return make_func( "hll_sketch_estimate", make_func("hll_union_agg", components[0].name), @@ -382,7 +462,12 @@ def _make_var_samp(self, components: list[MetricComponent]) -> ast.Expression: class VarPopDecomposition(VarianceDecompositionBase): """Population variance: E[X²] - E[X]²""" - def combine(self, components: list[MetricComponent]): + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ): return self._make_var_pop(components) @@ -390,7 +475,12 @@ def combine(self, components: list[MetricComponent]): class VarSampDecomposition(VarianceDecompositionBase): """Sample variance with Bessel's correction.""" - def combine(self, components: list[MetricComponent]): + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ): return self._make_var_samp(components) @@ -398,7 +488,12 @@ def combine(self, components: list[MetricComponent]): class VarianceDecomposition(VarianceDecompositionBase): """VARIANCE (alias for VAR_SAMP in Spark).""" - def combine(self, components: list[MetricComponent]): + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ): return self._make_var_samp(components) # pragma: no cover @@ -411,7 +506,12 @@ def combine(self, components: list[MetricComponent]): class StddevPopDecomposition(VarianceDecompositionBase): """Population standard deviation: sqrt(VAR_POP)""" - def combine(self, components: list[MetricComponent]): + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ): return make_func("SQRT", self._make_var_pop(components)) @@ -419,7 +519,12 @@ def combine(self, components: list[MetricComponent]): class StddevSampDecomposition(VarianceDecompositionBase): """Sample standard deviation: sqrt(VAR_SAMP)""" - def combine(self, components: list[MetricComponent]): + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ): return make_func("SQRT", self._make_var_samp(components)) @@ -427,7 +532,12 @@ def combine(self, components: list[MetricComponent]): class StddevDecomposition(VarianceDecompositionBase): """STDDEV (alias for STDDEV_SAMP in Spark).""" - def combine(self, components: list[MetricComponent]): + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ): return make_func("SQRT", self._make_var_samp(components)) # pragma: no cover @@ -532,7 +642,12 @@ def _make_covar_samp(self, components: list[MetricComponent]) -> ast.Expression: class CovarPopDecomposition(CovarianceDecompositionBase): """Population covariance: E[XY] - E[X]*E[Y]""" - def combine(self, components: list[MetricComponent]): + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ): return self._make_covar_pop(components) @@ -540,7 +655,12 @@ def combine(self, components: list[MetricComponent]): class CovarSampDecomposition(CovarianceDecompositionBase): """Sample covariance with Bessel's correction.""" - def combine(self, components: list[MetricComponent]): + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ): return self._make_covar_samp(components) @@ -574,7 +694,12 @@ def components(self) -> list[ComponentDef]: ComponentDef("_count", "COUNT({0})", "SUM", arg_index=0), ] - def combine(self, components: list[MetricComponent]) -> ast.Expression: + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ) -> ast.Expression: """ Build CORR: COVAR(X,Y) / (STDDEV(X) * STDDEV(Y)) @@ -644,8 +769,27 @@ def combine(self, components: list[MetricComponent]) -> ast.Expression: # ============================================================================= -def get_decomposition(func_class: type) -> AggDecomposition | None: - """Get decomposition instance for a function class, or None if not decomposable.""" +def get_decomposition( + func_class: type, + reaggregate: ReaggregateSpec | None = None, +) -> AggDecomposition | None: + """ + Get the decomposition for an aggregation, or None if not decomposable. + + A metric that declares a reaggregation family gets that family's + decomposition, which is how an opt-in sketch overrides the default handling + of its aggregation function. Everything else resolves by function class as + before. + + A declared family with nothing registered for it falls through rather than + raising: OSS ships no family implementations, so the metric simply keeps the + aggregability it would have had. + """ + if reaggregate is not None and reaggregate.fn is not None: + family_class = FAMILY_DECOMPOSITION_REGISTRY.get(reaggregate.fn) + if family_class is not None: + return family_class(params=reaggregate.params) + decomp_class = DECOMPOSITION_REGISTRY.get(func_class) if decomp_class is None: return None @@ -935,14 +1079,29 @@ class MetricComponentExtractor: For derived metrics: collects components from base metrics and substitutes references. """ - def __init__(self, node_revision_id: int): + def __init__( + self, + node_revision_id: int, + dialect: Dialect = Dialect.SPARK, + ): """ Extract metric components from a specific metric revision. Args: node_revision_id: ID of the metric node revision + dialect: Dialect the combiner will be rendered for. Combiners are + mostly dialect-neutral, but a sketch family whose engines expose + different function shapes needs the target -- Druid fuses a + t-digest's merge and combine, where Spark and Trino keep them + separate. + + Defaults to Spark, which is right for the callers that render + for display or for frozen measures. Only the build_v3 path, + which knows the cube's actual target engine, passes anything + else. """ self._node_revision_id = node_revision_id + self._dialect = dialect @classmethod async def from_node_name( # pragma: no cover @@ -1091,7 +1250,10 @@ async def extract( if is_parent_derived and parent_revision_id: # Recursively extract the derived metric (inline expansion) - parent_extractor = MetricComponentExtractor(parent_revision_id) + parent_extractor = MetricComponentExtractor( + parent_revision_id, + dialect=self._dialect, + ) base_components, derived_ast = await parent_extractor.extract( session, nodes_cache=nodes_cache, @@ -1405,7 +1567,15 @@ 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): + # + # The spec has to be passed here, not just to `_decompose` below: + # a family-gated metric decomposes precisely because it declared a + # family, and its aggregation function has no entry of its own. Ask + # without the spec and every such metric aborts at this gate and + # never reaches the family registry at all. + if any( + get_decomposition(dj_fn, reaggregate) is None for _, dj_fn in agg_funcs + ): if dimension_reaggregate_rules(reaggregate): self._raise_unsupported_reaggregate_shape( "dimension-specific reaggregation requires a " @@ -1421,7 +1591,7 @@ def _extract_base( return [], query_ast for func, dj_function in agg_funcs: - result = self._decompose(func, dj_function, query_ast) + result = self._decompose(func, dj_function, query_ast, reaggregate) if result: # pragma: no branch # Apply combiner to AST func.parent.replace(from_=func, to=result.combiner) # type: ignore @@ -1612,9 +1782,10 @@ def _decompose( func: ast.Function, dj_function: type, query_ast: ast.Query, + reaggregate: ReaggregateSpec | None = None, ) -> DecompositionResult | None: """Decompose an aggregation function using the registry.""" - decomposition = get_decomposition(dj_function) + decomposition = get_decomposition(dj_function, reaggregate) if decomposition is None: # pragma: no cover # Defensive: ``_extract_base`` filters non-decomposable @@ -1641,7 +1812,7 @@ def _decompose( quantifier=ast.SetQuantifier.Distinct, ) else: - combiner_ast = decomposition.combine(components) + combiner_ast = decomposition.combine(components, func, self._dialect) # Decomposed AVG / variance / stddev / covariance all build # SUM(...) / SUM(count)-style combiners where the denominator # can legitimately be 0. Wrap to produce NULL rather than diff --git a/datajunction-server/tests/models/reaggregate_test.py b/datajunction-server/tests/models/reaggregate_test.py index 8d225f6ca..93e6f552e 100644 --- a/datajunction-server/tests/models/reaggregate_test.py +++ b/datajunction-server/tests/models/reaggregate_test.py @@ -3,12 +3,14 @@ import pytest from pydantic import ValidationError +from datajunction_server.models import reaggregate as reaggregate_module from datajunction_server.models.reaggregate import ( DimensionReaggregateRule, ReaggregateSpec, ReaggregationFunction, dimension_reaggregate_rules, dump_reaggregate_spec, + is_parameterized_reaggregate_function, parse_reaggregate_spec, unsupported_dimension_reaggregate_functions, ) @@ -28,6 +30,7 @@ def test_dump_reaggregate_spec_from_dict(): ], }, ) == { + "fn": None, "params": None, "rules": [ { @@ -52,6 +55,7 @@ def test_dump_reaggregate_spec_from_model(): ], ), ) == { + "fn": None, "params": None, "rules": [ { @@ -88,7 +92,6 @@ def test_parse_reaggregate_spec_from_model(): @pytest.mark.parametrize( "unknown_field", [ - {"fn": "last_value"}, {"weight": "default.orders.quantity"}, { "rules": [ @@ -143,11 +146,42 @@ def test_empty_params_allowed(): assert ReaggregateSpec(params={}).params == {} +def test_params_accepted_for_parameterized_function(monkeypatch): + """ + Registering a function as parameterized enables `params` assignment. + """ + monkeypatch.setattr( + reaggregate_module, + "PARAMETERIZED_REAGGREGATE_FUNCTIONS", + frozenset({ReaggregationFunction.AVG}), + ) + spec = ReaggregateSpec(fn=ReaggregationFunction.AVG, params={"compression": 200}) + assert spec.params == {"compression": 200} + assert dump_reaggregate_spec(spec)["params"] == {"compression": 200} + + +def test_only_sketch_families_are_parameterized(): + """ + Only a sketch family takes tuning parameters. + + The rollup functions are fully specified by their name -- there is nothing + to tune about a SUM. A sketch is not: a t-digest still needs a compression, + which is why `params` is accepted for it and rejected everywhere else. + """ + parameterized = { + fn for fn in ReaggregationFunction if is_parameterized_reaggregate_function(fn) + } + + assert parameterized == {ReaggregationFunction.TDIGEST} + + def test_params_round_trip_through_parse(): """ `params` survives dict -> model -> dict without loss. """ - spec = parse_reaggregate_spec({"params": {"compression": 200}}) + spec = parse_reaggregate_spec( + {"fn": "tdigest", "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 8f8db41c6..66a3a4106 100644 --- a/datajunction-server/tests/sql/decompose_test.py +++ b/datajunction-server/tests/sql/decompose_test.py @@ -25,9 +25,15 @@ ) from datajunction_server.sql import functions as dj_functions from datajunction_server.sql.decompose import ( + DECOMPOSITION_REGISTRY, + FAMILY_DECOMPOSITION_REGISTRY, DUPLICATION_INVARIANT_AGGREGATIONS, + AggDecomposition, + ComponentDef, MetricComponentExtractor, + decomposes_family, get_decomposition, + make_func, is_duplication_invariant, is_metric_duplication_sensitive, safe_denominator, @@ -2894,3 +2900,328 @@ def test_params_without_an_aggregating_component_is_rejected(self): self._spec({"compression": 200}), ) assert "requires an aggregating component" in str(excinfo.value) + + +# ============================================================================= +# Dialect-aware combiners +# ============================================================================= + + +class _DialectAwareSum(AggDecomposition): + """ + A SUM whose combiner renders differently per dialect, and which records the + call it was handed. + + Stands in for a quantile sketch, the first family that actually needs this. + Druid exposes no scalar extractor for ``COMPLEX``, so a + t-digest's merge and combine fuse into one ``TDIGEST_QUANTILE(col, fraction)`` + call while Spark and Trino keep them separate -- a difference a rename table + cannot express. The fraction comes from the authored call, which is why + ``func`` is passed alongside the dialect. + """ + + seen: list[tuple["ast.Function", Dialect]] = [] + + @property + def components(self) -> list[ComponentDef]: + return [ComponentDef(suffix="_sum", accumulate="SUM", merge="SUM")] + + def combine(self, components, func, dialect=Dialect.SPARK): + type(self).seen.append((func, dialect)) + name = "druid_combiner" if dialect == Dialect.DRUID else "spark_combiner" + return make_func(name, components[0].name) + + +@pytest.fixture +def dialect_aware_sum(): + """Swap SUM's decomposition for one that branches on dialect, then restore.""" + original = DECOMPOSITION_REGISTRY.get(dj_functions.Sum) + _DialectAwareSum.seen = [] + DECOMPOSITION_REGISTRY[dj_functions.Sum] = _DialectAwareSum + try: + yield _DialectAwareSum + finally: + DECOMPOSITION_REGISTRY[dj_functions.Sum] = original + + +def test_combine_renders_per_dialect(): + """A decomposition can key its combiner off the target dialect.""" + components = [ + MetricComponent( + name="price_sum_abc", + expression="price", + aggregation="SUM", + merge="SUM", + rule=AggregationRule(), + ), + ] + func = parse("SELECT SUM(price) FROM parent_node").select.projection[0] + decomposition = _DialectAwareSum() + + spark = decomposition.combine(components, func, Dialect.SPARK) + druid = decomposition.combine(components, func, Dialect.DRUID) + + assert "spark_combiner" in str(spark) + assert "druid_combiner" in str(druid) + assert str(spark) != str(druid) + + +def test_combine_defaults_to_spark(): + """ + Omitting the dialect yields the Spark rendering. + + The default is what keeps the seven non-build_v3 extractor call sites + working untouched; they render for display and for frozen measures, where + Spark is the right answer. + """ + components = [ + MetricComponent( + name="price_sum_abc", + expression="price", + aggregation="SUM", + merge="SUM", + rule=AggregationRule(), + ), + ] + func = parse("SELECT SUM(price) FROM parent_node").select.projection[0] + + assert "spark_combiner" in str(_DialectAwareSum().combine(components, func)) + + +def test_combine_receives_the_originating_call(): + """ + The combiner is handed the authored ``ast.Function``, arguments included. + + This is what makes the quantile fraction reachable: the ``0.5`` in + ``APPROX_PERCENTILE(col, 0.5)`` lives only on the call, and before this the + combiner saw nothing but the merged components. Asserted on the arguments + rather than on identity, since that is the part a combiner needs. + """ + components = [ + MetricComponent( + name="latency_p50", + expression="latency_ms", + aggregation="SUM", + merge="SUM", + rule=AggregationRule(), + ), + ] + func = parse( + "SELECT APPROX_PERCENTILE(latency_ms, 0.5) FROM parent_node", + ).select.projection[0] + + decomposition = _DialectAwareSum() + decomposition.seen = [] + decomposition.combine(components, func, Dialect.DRUID) + + seen_func, seen_dialect = _DialectAwareSum.seen[-1] + assert seen_dialect == Dialect.DRUID + assert str(seen_func.args[1]) == "0.5" + + +def test_extractor_defaults_to_spark_dialect(): + """An extractor built without a dialect renders for Spark.""" + assert MetricComponentExtractor(1)._dialect == Dialect.SPARK + + +@pytest.mark.asyncio +async def test_extractor_threads_dialect_into_combiner( + session: AsyncSession, + create_metric, + dialect_aware_sum, +): + """ + The dialect handed to the extractor reaches the combiner. + + End-to-end over the real extraction path rather than a direct ``combine`` + call, because the plumbing between them is the part that was missing. + """ + metric_rev = await create_metric("SELECT SUM(price) FROM parent_node") + + _, spark_sql = await MetricComponentExtractor(metric_rev.id).extract(session) + _, druid_sql = await MetricComponentExtractor( + metric_rev.id, + dialect=Dialect.DRUID, + ).extract(session) + + assert "spark_combiner" in str(spark_sql) + assert "druid_combiner" in str(druid_sql) + + +@pytest.mark.asyncio +async def test_extractor_dialect_does_not_leak_between_instances( + session: AsyncSession, + create_metric, + dialect_aware_sum, +): + """ + Dialect is per-extractor configuration, not shared state. + + ``get_decomposition`` builds a fresh decomposition per call and the dialect + rides the extractor, so a Druid extraction must not change what a later + default extraction renders. + """ + metric_rev = await create_metric("SELECT SUM(price) FROM parent_node") + + await MetricComponentExtractor(metric_rev.id, dialect=Dialect.DRUID).extract( + session, + ) + _, after = await MetricComponentExtractor(metric_rev.id).extract(session) + + assert "spark_combiner" in str(after) + assert "druid_combiner" not in str(after) + + + + +# ============================================================================= +# Family-gated decompositions +# ============================================================================= + + +@pytest.fixture +def registered_family(): + """ + Register a decomposition for a reaggregation family, then remove it. + + OSS ships no family implementations -- they live downstream with the + engine-specific functions -- so a test has to supply one to exercise the + selection path at all. + """ + + @decomposes_family(ReaggregationFunction.TDIGEST) + class _FamilyDecomposition(AggDecomposition): + @property + def components(self) -> list[ComponentDef]: + return [ + ComponentDef( + suffix="_sketch", + accumulate=f"build_sketch({{}}, {self.params.get('compression', 1)})", + merge="merge_sketch", + ), + ] + + def combine(self, components, func, dialect=Dialect.SPARK): + return make_func("read_sketch", components[0].name) + + try: + yield _FamilyDecomposition + finally: + FAMILY_DECOMPOSITION_REGISTRY.pop(ReaggregationFunction.TDIGEST, None) + + +def test_declared_family_overrides_the_function_registry(registered_family): + """ + A metric declaring a family gets that family's decomposition. + + This is what lets `APPROX_PERCENTILE(x, 0.95)` keep its spelling and still + decompose: the aggregation function is unchanged, the family does the work. + """ + spec = ReaggregateSpec(fn=ReaggregationFunction.TDIGEST) + + decomposition = get_decomposition(dj_functions.ApproxPercentile, spec) + + assert isinstance(decomposition, registered_family) + + +def test_family_receives_its_tuning_params(registered_family): + """ + `reaggregate.params` reaches the decomposition, and thus the accumulate. + + Without this the declared compression would be recorded for reuse identity + and Druid ingestion but silently ignored by the SQL that builds the sketch. + """ + spec = ReaggregateSpec( + fn=ReaggregationFunction.TDIGEST, + params={"compression": 500}, + ) + + decomposition = get_decomposition(dj_functions.ApproxPercentile, spec) + + assert decomposition.params == {"compression": 500} + assert decomposition.components[0].accumulate == "build_sketch({}, 500)" + + +def test_undeclared_metric_is_untouched_by_a_registered_family(registered_family): + """ + A metric without `reaggregate` keeps the by-function decomposition. + + Registering a family must not change metrics that never opted in -- that + separation is the whole point of gating on the spec. + """ + assert get_decomposition(dj_functions.ApproxPercentile) is None + assert isinstance( + get_decomposition(dj_functions.Sum), + type(get_decomposition(dj_functions.Sum)), + ) + assert get_decomposition(dj_functions.Sum) is not None + + +def test_unregistered_family_falls_through(): + """ + Declaring a family nothing implements degrades rather than raising. + + OSS has no implementations, so `fn: tdigest` there simply leaves the metric + with the aggregability it already had. + """ + spec = ReaggregateSpec(fn=ReaggregationFunction.TDIGEST) + + assert get_decomposition(dj_functions.ApproxPercentile, spec) is None + assert get_decomposition(dj_functions.Sum, spec) is not None + + +@pytest.mark.asyncio +async def test_family_gated_percentile_decomposes_through_the_extractor( + session: AsyncSession, + create_metric, + registered_family, +): + """ + A family-gated metric survives the decomposability gate. + + The regression this guards: `_extract_base` aborts early if any aggregation + is non-decomposable, and that check used to ask `get_decomposition(fn)` + without the declared family. A family-gated metric is precisely one whose + aggregation function has no entry of its own, so every one of them aborted + there and fell through to native grain -- the family registry was correct + and simply never reached. + + Testing `get_decomposition` directly cannot catch this, and neither can a + family test over `SUM`, which passes the gate on its own merits. It needs a + non-decomposable function driven through the extractor. + """ + metric_rev = await create_metric( + "SELECT APPROX_PERCENTILE(price, 0.95) FROM parent_node", + reaggregate={"fn": "tdigest", "params": {"compression": 200}}, + ) + + components, combiner = await MetricComponentExtractor(metric_rev.id).extract( + session, + ) + + assert components, "family-gated percentile produced no components" + assert components[0].merge == "merge_sketch" + assert "build_sketch" in components[0].aggregation + assert "read_sketch" in str(combiner) + + +@pytest.mark.asyncio +async def test_percentile_without_a_family_still_does_not_decompose( + session: AsyncSession, + create_metric, + registered_family, +): + """ + The gate still stops metrics that did not opt in. + + The fix must not make every percentile decomposable -- only the ones that + declared the family. A registered family that nobody references changes + nothing. + """ + metric_rev = await create_metric( + "SELECT APPROX_PERCENTILE(price, 0.95) FROM parent_node", + ) + + components, _ = await MetricComponentExtractor(metric_rev.id).extract(session) + + assert components == [] From 5edc21e08da48e0c8efecf4fb4025b2f233b257e Mon Sep 17 00:00:00 2001 From: Robin Davis Date: Fri, 18 Sep 2026 22:01:55 -0700 Subject: [PATCH 02/17] Serialize sketches for materialization targets Allow components to serialize accumulated values for specific materialization targets and declare the resulting type. Add multi-argument type inference and fixed merge arguments so t-digest measures retain the correct representation and compression when written to Druid. --- .../python/datajunction/client.py | 26 +- .../python/datajunction/deployment.py | 25 +- .../datajunction_server/api/sql.py | 65 +---- .../datajunction_server/config.py | 36 +-- .../construction/build_v3/builder.py | 4 + .../construction/build_v3/combiners.py | 10 +- .../construction/build_v3/decomposition.py | 86 ++++-- .../construction/build_v3/measures.py | 81 +++++- .../construction/build_v3/types.py | 4 + .../database/preaggregation.py | 59 +---- .../internal/cube_materializations.py | 10 +- .../internal/materializations.py | 66 +---- .../datajunction_server/internal/nodes.py | 20 +- .../datajunction_server/models/decompose.py | 92 +++---- .../datajunction_server/models/deployment.py | 77 +----- .../models/materialization.py | 16 ++ .../datajunction_server/sql/decompose.py | 36 ++- .../datajunction_server/sql/parsing/ast.py | 23 -- datajunction-server/tests/api/cubes_test.py | 64 +++++ datajunction-server/tests/api/metrics_test.py | 8 + .../tests/api/preaggregations_test.py | 8 + datajunction-server/tests/api/sql_v2_test.py | 20 ++ .../build_v3/accumulate_type_test.py | 152 +++++++++++ .../construction/build_v3/merge_args_test.py | 129 +++++++++ .../build_v3/serialize_target_test.py | 244 ++++++++++++++++++ .../tests/database/preaggregation_test.py | 12 +- .../nodes/derive_frozen_measures_test.py | 15 +- .../models/cube_druid_sketch_spec_test.py | 196 ++++++++++++++ .../tests/models/reaggregate_test.py | 4 - .../tests/sql/decompose_test.py | 111 +++++--- 30 files changed, 1176 insertions(+), 523 deletions(-) create mode 100644 datajunction-server/tests/construction/build_v3/accumulate_type_test.py create mode 100644 datajunction-server/tests/construction/build_v3/merge_args_test.py create mode 100644 datajunction-server/tests/construction/build_v3/serialize_target_test.py create mode 100644 datajunction-server/tests/models/cube_druid_sketch_spec_test.py diff --git a/datajunction-clients/python/datajunction/client.py b/datajunction-clients/python/datajunction/client.py index ef9867264..f236ca422 100644 --- a/datajunction-clients/python/datajunction/client.py +++ b/datajunction-clients/python/datajunction/client.py @@ -207,31 +207,7 @@ def plan( include_temporal_filters: bool = False, lookback_window: str | None = None, ): - """ - Returns a query execution plan for the given metrics and dimensions. - - The plan shows: - - grain_groups: How metrics are grouped and their intermediate SQL - - metric_formulas: How each metric combines its components - - requested_dimensions: The dimensions being queried - - This is useful for understanding how DJ decomposes metrics into - atomic aggregations and how multiple fact tables are joined together. - - Args: - metrics: List of metric names to include - dimensions: List of dimensions to group by - filters: List of filter expressions - cube: Optional cube node name. When provided, the cube's stored - filters are automatically prepended to the query filters. - dialect: SQL dialect (e.g., 'spark', 'trino'). Defaults to engine dialect. - use_materialized: Whether to use materialized tables when available - include_temporal_filters: Whether to include temporal partition filters. - Only applies if the metrics and dimensions resolve to a cube with - temporal partitions. - lookback_window: Lookback window for temporal filters (e.g., '3 DAY', - '1 WEEK'). Only applicable when include_temporal_filters is True. - """ + """Returns a query execution plan for the given metrics and dimensions.""" params: dict = { "metrics": metrics, "dimensions": dimensions or [], diff --git a/datajunction-clients/python/datajunction/deployment.py b/datajunction-clients/python/datajunction/deployment.py index e8ec4fe54..0ba52c62a 100644 --- a/datajunction-clients/python/datajunction/deployment.py +++ b/datajunction-clients/python/datajunction/deployment.py @@ -352,30 +352,7 @@ def build_codeowners( default_owner: str | None = None, exclude_dirs: list[str] | None = None, ) -> int: - """ - Generate a CODEOWNERS file from the owners fields in DJ node YAML files. - - Walks base_dir recursively, reads every *.yaml file (skipping dj.yaml), - and maps each file path to its owners list. Files with no owners are - omitted. Paths in the output are relative to base_dir and prefixed with - / so GitHub resolves them from the repo root. - - If github_api_url is provided (and GITHUB_TOKEN / github_token_env is set), - email addresses in owners fields are resolved to GitHub usernames via the - search API. Unresolvable emails are emitted as-is with a warning comment. - - ``default_owner``, when set, emits a leading ``* `` rule so - unmatched files (and any excluded directories) fall through to it. Because - CODEOWNERS is last-match-wins, the per-file rules below it still take - precedence for the files they name. - - ``exclude_dirs`` lists directories (relative to base_dir) whose nodes are - NOT given per-file owners — they fall through to ``default_owner`` instead. - This is for machine-generated trees (e.g. ``nodes/generated``) that have no - individual human owner and should be team-owned as a block. - - Returns the number of per-file entries written (excludes the default rule). - """ + """Generate a CODEOWNERS file from the owners fields in DJ node YAML files.""" base = Path(base_dir).resolve() excluded_dirs = [(base / d).resolve() for d in (exclude_dirs or [])] diff --git a/datajunction-server/datajunction_server/api/sql.py b/datajunction-server/datajunction_server/api/sql.py index c2216df2a..52e462b43 100644 --- a/datajunction-server/datajunction_server/api/sql.py +++ b/datajunction-server/datajunction_server/api/sql.py @@ -222,36 +222,7 @@ async def get_measures_sql_v3( session: AsyncSession = Depends(get_session), current_user: User = Depends(get_current_user), ) -> MeasuresSQLResponse: - """ - Generate pre-aggregated measures SQL for the requested metrics. - - Measures SQL represents the first stage of metric computation - it decomposes - each metric into its atomic aggregation components (e.g., SUM(amount), COUNT(*)) - and produces SQL that computes these components at the requested dimensional grain. - - Metrics are separated into grain groups, which represent sets of metrics that can be - computed together at a common grain. Each grain group produces its own SQL query, which - can be materialized independently to produce intermediate tables that are then queried - to compute final metric values. - - Returns: - One or more `GrainGroupSQL` objects, each containing: - - SQL query computing metric components at the specified grain - - Column metadata with semantic types - - Component details for downstream re-aggregation - - Args: - cube: Optional cube node name. When provided, the cube's stored filters are - automatically prepended to the query filters. - use_materialized: If True (default), use materialized tables when available. - Set to False when generating SQL for materialization refresh to avoid - circular references. - include_temporal_filters: If True, checks if metrics+dimensions resolve to - a cube with temporal partitions, and applies partition filters if so. - lookback_window: Lookback window for temporal filters when applicable. - - See also: `/sql/metrics/v3/` for the final combined query with metric expressions. - """ + """Generate pre-aggregated measures SQL for the requested metrics.""" merged_filters = list(filters) cube_node = None if cube: @@ -445,38 +416,7 @@ async def get_combined_measures_sql_v3( session: AsyncSession = Depends(get_session), current_user: User = Depends(get_current_user), ) -> CombinedMeasuresSQLResponse: - """ - Generate combined pre-aggregated measures SQL for the requested metrics. - - This endpoint combines multiple grain groups into a single SQL query using - FULL OUTER JOIN on shared dimensions. Dimension columns are wrapped with - COALESCE to handle NULLs from non-matching rows. - - This is useful for: - - Druid cube materialization where a single combined table is needed - - Simplifying downstream queries that need data from multiple fact tables - - Pre-computing joined aggregations for dashboards - - The combined SQL contains: - - CTEs for each grain group's pre-aggregated data - - FULL OUTER JOIN between grain groups on shared dimensions - - COALESCE on dimension columns to handle NULL values - - All measure columns from all grain groups - - Args: - metrics: List of metric names to include - dimensions: List of dimensions to group by (the grain) - filters: Optional filters to apply - use_preagg_tables: If False (default), compute from scratch using source tables. - If True, read from pre-aggregation tables. - - Returns: - Combined SQL query with column metadata and grain information. - - See also: - - `/sql/measures/v3/` for individual grain group queries - - `/sql/metrics/v3/` for final metric computations with combiner expressions - """ + """Generate combined pre-aggregated measures SQL for the requested metrics.""" _t0 = time.monotonic() if use_preagg_tables: # Generate SQL that reads from pre-agg tables (deterministic names) @@ -714,7 +654,6 @@ async def get_metrics_sql_v3( """ if not metrics and not cube: raise DJInvalidInputException("At least one metric is required") - # Shared metrics-SQL core (cube pinning, cube_filters prepend, dialect # auto-resolve, build_metrics_sql, and the build-latency metrics + [SQL] log). # Also used by the semantic-layer endpoint. diff --git a/datajunction-server/datajunction_server/config.py b/datajunction-server/datajunction_server/config.py index 39941bca4..31ee27a8e 100644 --- a/datajunction-server/datajunction_server/config.py +++ b/datajunction-server/datajunction_server/config.py @@ -48,41 +48,7 @@ class DatabaseConfig(BaseModel): class QueryClientConfig(BaseModel): - """ - Configuration for query service clients. - - Set via environment variables using double-underscore delimiters, e.g.:: - - QUERY_CLIENT__TYPE=bigquery - QUERY_CLIENT__CONNECTION__PROJECT=my-gcp-project - - Supported client types and their required connection parameters: - - **http** (default):: - - QUERY_CLIENT__TYPE=http - QUERY_CLIENT__CONNECTION__URI=http://djqs:8001 - - **snowflake**:: - - QUERY_CLIENT__TYPE=snowflake - QUERY_CLIENT__CONNECTION__ACCOUNT=my-account - QUERY_CLIENT__CONNECTION__USER=my-user - QUERY_CLIENT__CONNECTION__PASSWORD=my-password - - **bigquery**:: - - QUERY_CLIENT__TYPE=bigquery - QUERY_CLIENT__CONNECTION__PROJECT=my-gcp-project - # Optional: path to a service account JSON key file - QUERY_CLIENT__CONNECTION__CREDENTIALS_PATH=/path/to/service-account.json - # Optional: BigQuery location (e.g. US, EU) - QUERY_CLIENT__CONNECTION__LOCATION=US - - When multiple DJ catalogs map to different GCP projects, set the engine URI - on the DJ engine to ``bigquery://my-gcp-project`` — the project is resolved - from the engine URI first, falling back to ``PROJECT`` above. - """ + """Configuration for query service clients.""" # Type of query client: 'http', 'snowflake', 'bigquery', 'databricks', 'trino', etc. type: str = "http" diff --git a/datajunction-server/datajunction_server/construction/build_v3/builder.py b/datajunction-server/datajunction_server/construction/build_v3/builder.py index e9a6ca1e1..9e34d7a24 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/builder.py +++ b/datajunction-server/datajunction_server/construction/build_v3/builder.py @@ -59,6 +59,7 @@ from datajunction_server.errors import DJError, DJInvalidInputException, ErrorCode from datajunction_server.instrumentation import events from datajunction_server.models.dialect import Dialect +from datajunction_server.models.materialization import MaterializationTarget from datajunction_server.models.partition import PartitionType from datajunction_server.sql.parsing import ast from datajunction_server.sql.parsing.backends.antlr4 import parse @@ -252,6 +253,7 @@ async def setup_build_context( include_temporal_filters: bool = False, lookback_window: str | None = None, matched_cube: NodeRevision | None = None, + materialization_target: MaterializationTarget | None = None, ) -> BuildContext: """ Create and initialize a BuildContext with all setup done. @@ -377,6 +379,7 @@ async def build_measures_sql( lookback_window: str | None = None, query_parameters: dict[str, Any] | None = None, matched_cube: NodeRevision | None = None, + materialization_target: MaterializationTarget | None = None, ) -> GeneratedMeasuresSQL: """ Build measures SQL for a set of metrics, dimensions, and filters. @@ -420,6 +423,7 @@ async def build_measures_sql( include_temporal_filters=include_temporal_filters, lookback_window=lookback_window, matched_cube=matched_cube, + materialization_target=materialization_target, ) # Build grain groups from context diff --git a/datajunction-server/datajunction_server/construction/build_v3/combiners.py b/datajunction-server/datajunction_server/construction/build_v3/combiners.py index 5ca723234..7c5d710cf 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/combiners.py +++ b/datajunction-server/datajunction_server/construction/build_v3/combiners.py @@ -23,6 +23,9 @@ from datajunction_server.construction.build_v3.cte import ( process_metric_combiner_expression, ) +from datajunction_server.construction.build_v3.decomposition import ( + build_merge_call, +) from datajunction_server.construction.build_v3.preagg_matcher import ( get_temporal_partitions, ) @@ -862,20 +865,19 @@ def _build_grain_group_from_preagg_table( # Find the component to get the merge function merge_func = None + merge_args: list[str] = [] for comp in original_gg.components: if ( # pragma: no branch comp.name == col.name or original_gg.component_aliases.get(comp.name) == col.name ): merge_func = comp.merge + merge_args = comp.merge_args break if merge_func: # Apply re-aggregation - agg_expr = ast.Function( - name=ast.Name(merge_func), - args=[col_ref], - ) + agg_expr = build_merge_call(merge_func, merge_args, col_ref) aliased = ast.Alias(child=agg_expr, alias=ast.Name(col.name)) select_items.append(aliased) else: diff --git a/datajunction-server/datajunction_server/construction/build_v3/decomposition.py b/datajunction-server/datajunction_server/construction/build_v3/decomposition.py index 313d44aac..960625e90 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/decomposition.py +++ b/datajunction-server/datajunction_server/construction/build_v3/decomposition.py @@ -28,6 +28,7 @@ from datajunction_server.errors import DJInvalidInputException from datajunction_server.models.decompose import Aggregability, MetricComponent from datajunction_server.models.dialect import Dialect +from datajunction_server.models.materialization import MaterializationTarget from datajunction_server.models.node_type import NodeType from datajunction_server.sql.decompose import MetricComponentExtractor from datajunction_server.sql.parsing import ast @@ -215,7 +216,10 @@ async def decompose_metric( ) -def build_component_expression(component: MetricComponent) -> ast.Expression: +def build_component_expression( + component: MetricComponent, + materialization_target: MaterializationTarget | None = None, +) -> ast.Expression: """ Build the accumulate expression AST for a metric component. @@ -225,11 +229,18 @@ def build_component_expression(component: MetricComponent) -> ast.Expression: Note: Templates may be pre-expanded (e.g., "SUM(POWER(match_score, 2))") by the decomposition phase, so we detect this by checking for parentheses without template placeholders. + + When ``materialization_target`` is one the component declares a ``serialize`` + conversion for, the accumulated expression is wrapped in it. That happens + only while writing a measures table: a query-time build passes no target, so + the unwrapped expression stands. """ if not component.aggregation: # pragma: no cover # No aggregation - just return the expression as a column return ast.Column(name=ast.Name(component.expression)) + accumulated: ast.Expression + # Check if it's an unexpanded template with {} if "{" in component.aggregation: # pragma: no cover # Template like "SUM(POWER({}, 2))" - expand it @@ -239,25 +250,67 @@ def build_component_expression(component: MetricComponent) -> ast.Expression: if isinstance(expr_ast, ast.Alias): expr_ast = expr_ast.child expr_ast.clear_parent() - return cast(ast.Expression, expr_ast) + accumulated = cast(ast.Expression, expr_ast) # Check if it's a pre-expanded template (contains parentheses, like "SUM(POWER(x, 2))") # vs a simple function name (like "SUM") - if "(" in component.aggregation: + elif "(" in component.aggregation: # Pre-expanded template - parse it directly as a complete expression expr_ast = parse(f"SELECT {component.aggregation}").select.projection[0] if isinstance(expr_ast, ast.Alias): expr_ast = expr_ast.child # pragma: no cover expr_ast.clear_parent() - return cast(ast.Expression, expr_ast) + accumulated = cast(ast.Expression, expr_ast) else: # Simple function name like "SUM" - build SUM(expression) arg_expr = parse(f"SELECT {component.expression}").select.projection[0] - func = ast.Function( + accumulated = ast.Function( name=ast.Name(component.aggregation), args=[cast(ast.Expression, arg_expr)], ) - return func + + # All accumulate shapes must receive the target conversion. + return _apply_serialize(accumulated, component, materialization_target) + + +def _apply_serialize( + expr: ast.Expression, + component: MetricComponent, + materialization_target: MaterializationTarget | None, +) -> ast.Expression: + """ + Wrap an accumulated expression in the component's serialize conversion. + + Returns the expression untouched unless the component declares a conversion + and names this target -- so every existing component, and every query-time + build, is unaffected. + """ + if not component.serializes_for(materialization_target): + return expr + wrapped = parse( + f"SELECT {component.serialize.replace('{}', str(expr))}", + ).select.projection[0] + wrapped.clear_parent() + return cast(ast.Expression, wrapped) + + +def build_merge_call( + merge: str, + merge_args: list[str], + arg: ast.Expression, +) -> ast.Function: + """ + Build a component's Phase 2 merge call, including any fixed tuning arguments. + + Components without `merge_args` -- every non-sketch aggregation -- produce + the same single-argument call as before. + """ + extra: list[ast.Expression] = [] + for literal in merge_args: + parsed = parse(f"SELECT {literal}").select.projection[0] + parsed.clear_parent() + extra.append(cast(ast.Expression, parsed)) + return ast.Function(name=ast.Name(merge), args=[arg, *extra]) def get_base_metrics_for_derived(ctx: BuildContext, metric_node: Node) -> list[Node]: @@ -275,12 +328,9 @@ def collect_bases(node: Node): return visited.add(node.name) - # Recurse through metric parents. A node with any metric parent is itself - # derived; a node with no metric parent is a base metric -- regardless of - # whether its data source is a fact/transform or a dimension node (a base - # metric can be defined directly on a dimension node). Classifying by - # "has a metric parent" avoids mistaking a dimension-sourced base metric - # for a bare required-dimension reference, which would drop its grain group. + # Recurse through metric parents. A node with any metric parent is derived; + # a node with no metric parent is a base metric. Checking for metric parents + # avoids dropping base metrics that are defined directly on dimension nodes. has_metric_parent = False for parent_name in ctx.parent_map.get(node.name, []): parent = ctx.nodes.get(parent_name) @@ -589,8 +639,7 @@ def merge_grain_groups(grain_groups: list[GrainGroup]) -> list[GrainGroup]: """ # 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. + # semi-additive FULL group, as extra grain columns would 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) @@ -604,8 +653,7 @@ def merge_grain_groups(grain_groups: list[GrainGroup]) -> list[GrainGroup]: 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. + # corrupting the already-aggregated value for MAX_BY/MIN_BY. merged_groups.extend(parent_groups) else: # Multiple groups for same parent - merge them @@ -655,10 +703,8 @@ def _merge_parent_grain_groups(groups: list[GrainGroup]) -> GrainGroup: gg.reaggregate_component_dimensions, ) - # Carry over non-decomposable metrics from every contributing group. - # Without this, merging a NONE group into a FULL/LIMITED neighbor - # silently drops the non-decomposable metric expressions and the - # response loses those metrics entirely. + # Carry over non-decomposable metrics from every contributing group to prevent them + # from being silently dropped when merging a NONE group into a FULL/LIMITED neighbor. all_non_decomposable: list[DecomposedMetricInfo] = [] for gg in groups: all_non_decomposable.extend(gg.non_decomposable_metrics) diff --git a/datajunction-server/datajunction_server/construction/build_v3/measures.py b/datajunction-server/datajunction_server/construction/build_v3/measures.py index 4e5ad98fb..0607e69f6 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/measures.py +++ b/datajunction-server/datajunction_server/construction/build_v3/measures.py @@ -28,6 +28,7 @@ from datajunction_server.construction.build_v3.decomposition import ( analyze_grain_groups, build_component_expression, + build_merge_call, merge_grain_groups, ) from datajunction_server.construction.build_v3.dimensions import ( @@ -78,6 +79,7 @@ from datajunction_server.internal.scan_estimation import calculate_scan_estimate from datajunction_server.models.decompose import Aggregability, MetricComponent from datajunction_server.models.node_type import NodeType +from datajunction_server.models.materialization import MaterializationTarget from datajunction_server.sql.functions import function_registry from datajunction_server.sql.parsing import ast from datajunction_server.sql.parsing import types as ct @@ -157,10 +159,51 @@ def _parse_type_string(type_str: str | None) -> ct.ColumnType | None: return _TYPE_STRING_MAP.get(normalized) +def _multi_argument_accumulate_types( + aggregation: str, + parent_node: Node | None, +) -> list[ct.ColumnType] | None: + """ + Argument types of an accumulate whose outermost call takes more than one. + + Type inference feeds a component exactly one input type, which is right for + every aggregation that takes one argument -- including the templated ones, + whose outermost call is still a single-argument ``SUM``. A sketch breaks + that assumption: ``nflx_tdigest(latency_ms, CAST(200.0 AS DOUBLE))`` needs + two, so calling ``infer_type`` with one raises ``TypeError``, and the caller + quietly falls back to the metric's own type. That is how a struct-valued + sketch column comes to be recorded as ``double``. + + Returns None whenever this does not apply or an argument cannot be typed, + so the single-argument path is left exactly as it was. + """ + if parent_node is None: + return None + call = parse(f"SELECT {aggregation}").select.projection[0] + if not isinstance(call, ast.Function) or len(call.args) < 2: + return None + + arg_types: list[ct.ColumnType] = [] + for arg in call.args: + if isinstance(arg, ast.Column): + # Only the columns need the parent to resolve; literals and casts + # carry their own type without being bound to a table. + resolved = _parse_type_string( + get_column_type(parent_node, str(arg.alias_or_name.name)), + ) + else: + resolved = arg.type + if resolved is None: + return None + arg_types.append(resolved) + return arg_types + + def infer_component_type( component: MetricComponent, metric_type: str, parent_node: Node | None = None, + materialization_target: MaterializationTarget | None = None, ) -> str: """ Infer the SQL type of a metric component based on its aggregation function. @@ -177,6 +220,14 @@ def infer_component_type( Returns: The inferred SQL type string """ + # A serialized column stores the converted type, and the family declares it. + # Checked first: inference below cannot see the conversion, and for a + # multi-argument sketch accumulate it falls back to `metric_type` entirely, + # so leaving this to inference records a type that makes the Druid + # aggregator lookup miss and drop the measure with no error. + if component.serializes_for(materialization_target) and component.serialize_type: + return component.serialize_type + if not component.aggregation: return metric_type # pragma: no cover @@ -201,17 +252,22 @@ def infer_component_type( col_type_str = get_column_type(parent_node, component.expression) input_type = _parse_type_string(col_type_str) + multi_arg_types = _multi_argument_accumulate_types(agg_str, parent_node) + try: - if input_type: + if multi_arg_types is not None: + result_type = func_class.infer_type(*multi_arg_types) + elif input_type: result_type = func_class.infer_type(input_type) else: # pragma: no cover # Fallback: try with a generic ColumnType result_type = func_class.infer_type(ct.ColumnType("unknown", "unknown")) - return str(result_type) except (TypeError, NotImplementedError, AttributeError): # Function may require more specific types - fall back to metric type return metric_type + return str(result_type) + def _get_filter_column_name_for_dimension( resolved_dim: ResolvedDimension, @@ -1858,9 +1914,10 @@ def build_grain_group_from_preagg( # If no merge function, output column directly (e.g., grain column for LIMITED) # Otherwise, apply the merge function for re-aggregation if component.merge: - agg_expr = ast.Function( - name=ast.Name(component.merge), - args=[_preagg_column(scan_name, scan_alias)], + agg_expr = build_merge_call( + component.merge, + component.merge_args, + _preagg_column(scan_name, scan_alias), ) aliased = ast.Alias(child=agg_expr, alias=ast.Name(output_alias)) select_items.append(aliased) @@ -2140,7 +2197,10 @@ def build_grain_group_sql( # FULL: apply aggregation at finest grain, will be re-aggregated in final SELECT # Always use component.name for consistency - no special case for single-component component_alias = component.name - expr_ast = build_component_expression(component) + expr_ast = build_component_expression( + component, + ctx.materialization_target, + ) component_expressions.append((component_alias, expr_ast)) component_metadata.append( (component_alias, component, metric_node), @@ -2159,7 +2219,7 @@ def build_grain_group_sql( # Always use component.name for consistency - no special case for single-component component_alias = component.name - expr_ast = build_component_expression(component) + expr_ast = build_component_expression(component, ctx.materialization_target) component_expressions.append((component_alias, expr_ast)) component_metadata.append((component_alias, component, metric_node)) @@ -2404,7 +2464,12 @@ def build_grain_group_sql( ColumnMetadata( name=ctx.alias_registry.get_alias(comp_alias) or comp_alias, semantic_name=f"{metric_node.name}:{component.name}", - type=infer_component_type(component, metric_type, parent_node), + type=infer_component_type( + component, + metric_type, + parent_node, + ctx.materialization_target, + ), semantic_type="metric_component", ), ) diff --git a/datajunction-server/datajunction_server/construction/build_v3/types.py b/datajunction-server/datajunction_server/construction/build_v3/types.py index baac1e4f0..9f50ef097 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/types.py +++ b/datajunction-server/datajunction_server/construction/build_v3/types.py @@ -13,6 +13,7 @@ from datajunction_server.errors import DJInvalidInputException, DJWarning from datajunction_server.models.decompose import Aggregability, MetricComponent from datajunction_server.models.dialect import Dialect +from datajunction_server.models.materialization import MaterializationTarget from datajunction_server.models.node_type import NodeType from datajunction_server.sql.parsing import ast from datajunction_server.sql.parsing.ast import to_sql @@ -43,6 +44,9 @@ class BuildContext: dimensions: list[str] filters: list[str] = field(default_factory=list) dialect: Dialect = Dialect.SPARK + # Set only when this build is producing a measures table for materialization. + # None for query-time builds, which never serialize. + materialization_target: MaterializationTarget | None = None alias_registry: AliasRegistry = field(default_factory=AliasRegistry) # Filter classification (populated early in setup) diff --git a/datajunction-server/datajunction_server/database/preaggregation.py b/datajunction-server/datajunction_server/database/preaggregation.py index 937dd051d..f7265d724 100644 --- a/datajunction-server/datajunction_server/database/preaggregation.py +++ b/datajunction-server/datajunction_server/database/preaggregation.py @@ -222,36 +222,7 @@ def compute_preagg_hash( class PreAggregation(Base): - """ - First-class pre-aggregation entity that can be shared across cubes. - - A pre-aggregation represents a materialized grouping of measures at a specific grain, - enabling efficient metric calculations by pre-computing aggregations. - - Pre-aggregations are ALWAYS created by DJ (via /preaggs/plan endpoint) from - metrics + dimensions. Users never manually construct them - this ensures - consistency between DJ-managed (Flow A) and user-managed (Flow B) materialization. - - Key concepts: - - `node_revision`: The specific node revision this pre-agg is based on - - `grain_columns`: Fully qualified dimension references that define the aggregation level - - `measures`: Full MetricComponent info for matching and re-aggregation - - `sql`: The generated SQL for materializing this pre-agg - - `grain_group_hash`: Hash of (node_revision_id + sorted(grain_columns)) for grouping - - Measure format (MetricComponent): - - name: Column name in materialized table - - expression: The raw SQL expression - - expr_hash: Hash of expression for identity matching - - aggregation: Phase 1 function (e.g., "SUM") - - merge: Phase 2 re-aggregation function - - rule: Aggregation rules (type, level) - - Availability tracking: - - Materialization status is tracked via AvailabilityState - - Flow A: DJ's query service posts availability after materialization - - Flow B: User's query service posts to /preaggs/{id}/availability/ - """ + """First-class pre-aggregation entity that can be shared across cubes.""" __tablename__ = "pre_aggregation" @@ -490,33 +461,7 @@ async def find_matching( grain_columns: list[str], measure_identities: set[str], ) -> PreAggregation | None: - """ - Find the row this exact declaration already occupies, if any. - - This is the upsert's identity check -- "is this the same declaration I am - about to write?" -- so it matches the uniqueness key exactly: same - revision and grain (via ``grain_group_hash``) and the SAME set of measure - identities, not merely a covering one. - - Covering was wrong here, in two compounding ways. Callers replace the - matched row's contents wholesale, so a narrow declaration could match a - wider pre-agg and silently strip measures off it, breaking routing for - whatever metric depended on the dropped ones. And ``preagg_hash`` is - UNIQUE over exactly ``(node_revision_id, grain_columns, - measure_identities)`` and frozen at insert, so a covering match could - hand back a row whose stored hash no longer described its own contents. - A declaration that genuinely differs now gets its own row instead. - - Identity tokens rather than bare expression hashes, because the latter - made SUM- and MAX-backed pre-aggs look like one row, so registering - either silently overwrote the other. - - Contrast ``find_latest_for_node``, which asks the other question -- "what - did this declaration look like before?" -- and does want covering. - - Returns: - Matching PreAggregation if found, None otherwise - """ + """Find the row this exact declaration already occupies, if any.""" grain_group_hash = compute_grain_group_hash(node_revision_id, grain_columns) candidates = await cls.get_by_grain_group_hash(session, grain_group_hash) diff --git a/datajunction-server/datajunction_server/internal/cube_materializations.py b/datajunction-server/datajunction_server/internal/cube_materializations.py index f5531ff9d..aa529c30b 100644 --- a/datajunction-server/datajunction_server/internal/cube_materializations.py +++ b/datajunction-server/datajunction_server/internal/cube_materializations.py @@ -24,7 +24,10 @@ UpsertCubeMaterialization, ) from datajunction_server.models.dialect import Dialect -from datajunction_server.models.materialization import MaterializationStrategy +from datajunction_server.models.materialization import ( + MaterializationStrategy, + MaterializationTarget, +) from datajunction_server.models.node_type import NodeNameVersion from datajunction_server.models.partition import Granularity from datajunction_server.models.query import ColumnMetadata @@ -371,6 +374,11 @@ async def build_cube_materialization( filters=(current_revision.cube_filters or []) + extra_filters, dialect=Dialect.SPARK, use_materialized=True, + # Measures SQL runs in Spark for every target, so the dialect above + # cannot say where the output is headed. Druid is the only cube target + # (`UpsertCubeMaterialization.job` is Literal["druid_cube"]), and this is + # the one caller that knows it. + materialization_target=MaterializationTarget.DRUID, ) measures_queries = sorted( [ diff --git a/datajunction-server/datajunction_server/internal/materializations.py b/datajunction-server/datajunction_server/internal/materializations.py index 41025d10e..3f2fab8c6 100644 --- a/datajunction-server/datajunction_server/internal/materializations.py +++ b/datajunction-server/datajunction_server/internal/materializations.py @@ -472,36 +472,6 @@ async def reconcile_declared_materializations( Returns what each declared block resolved to, and the materializations the blocks together superseded. - - A cube may declare more than one block -- typically an `incremental_time` build - for freshness beside a periodic `full` rebuild that corrects late-arriving data - -- and each is reconciled the same way. They cannot collide: the row name is - derived from job, strategy and partition, the job is fixed for a cube and the - strategies are unique by validation, so every declared block owns a row of its - own. - - Mirrors what `POST /nodes/{name}/materialization/` does -- build the config, then - update the row of the same name in place rather than inserting a second one, - since `(name, node_revision_id)` is unique -- with one deliberate difference. The - endpoint decides "unchanged" on `config` alone, but `schedule` and `strategy` are - columns rather than config keys, so by that test a schedule-only edit compares - equal and is silently dropped. Rescheduling is the whole point of a declared - block, so all three are compared here. - - The declared blocks describe *the* materializations for their cube, so every - other active row on the revision is deactivated. The rule is "any active row - whose name no block built" rather than "any row of the same job type", because a - cube's materializations are all writing one Druid datasource and a full rebuild - replaces that datasource wholesale -- so a legacy `druid_measures_cube` row is - just as much a competing writer as a second `druid_cube` row, and matching on job - type would leave it running. The name is what carries the difference: it is - derived from job, strategy and partition, so declaring a strategy the cube was - not already materialized with builds a differently named row, and without this - the cube ends up with two live workflows deleting each other's data. - - A cube planner row is the one thing left alone. It writes a datasource of its own - and DJ cannot rebuild it from a declared block, so superseding it would stop a - workflow nothing here can replace. """ # Snapshotted before the builds, each of which sets the new materialization's # backref and so appends it to this very collection -- searching afterwards would @@ -582,27 +552,7 @@ async def swap_cube_materializations( ) -> CubeMaterializationSwap | None: """ Rebuild a cube's materializations against a new revision and retire the old ones. - - Materializations belong to a single `NodeRevision` and availability is scoped to - the revision encoded in the materialized table name, so without this a new cube - revision -- including a metadata-only one -- has neither: the cube silently falls - back to live queries while the superseded revision's workflow keeps posting - availability for a table built from the old definition. Every new revision - therefore swaps, no matter how insignificant the change was. - - Rebuilt rather than copied: a stored config embeds the cube version plus combiner - SQL and a Druid spec derived from the old definition, so the new revision goes - back through `create_new_materialization`. The old materialization is the default - source of the user's intent because it is often the only record of it -- a cube - materialized through the UI has no YAML to read it from. A cube that does declare - `materialization:` passes it as `declared`, which wins, so a push that edits both - a metric and the schedule rebuilds with the new schedule rather than the old. - - `previous_table_usable` is the caller's answer to whether the superseded - revision's materialized table is still valid (`is_non_trivial_cube_change` - inverted), recorded on the history event so an operator can tell whether the - rebuild can adopt the existing data or needs a fresh build and backfill. - + Touches only DJ-side state, and returns the query service work still owed -- `None` when the cube had nothing materialized and there is no work at all. The caller commits and then hands the result to `apply_cube_materialization_swap`, so @@ -648,18 +598,8 @@ async def swap_cube_materializations( try: upsert = _upsert_from_materialization(materialization) if declared and isinstance(upsert, UpsertCubeMaterialization): - # The declared block wins over the recovered intent: a cube that - # declares `materialization:` has its config in the repo, so a - # rebuild triggered by the same deploy must build what the YAML now - # says. Only cube materializations can be declared; anything else - # keeps what was recovered. - # - # Which block, when the cube declares several: the one naming the - # strategy this row was built with, since that is what identifies a - # declared entry. A row whose strategy nothing declares falls to the - # first block, which is what a cube declaring exactly one has always - # done -- and is the only sensible answer, since a rebuild has to - # produce something for a row that is being retired either way. + # The declared block wins over the recovered intent. + # If there are several, match by strategy, otherwise use the first. block = next( ( candidate diff --git a/datajunction-server/datajunction_server/internal/nodes.py b/datajunction-server/datajunction_server/internal/nodes.py index c0599cf2f..4bbb5960b 100644 --- a/datajunction-server/datajunction_server/internal/nodes.py +++ b/datajunction-server/datajunction_server/internal/nodes.py @@ -2524,14 +2524,7 @@ async def _propagate_update_downstream( cache.delete(upstream_cache_key) if downstream.type == NodeType.CUBE: - # Any tier rebuilds, and the churn is deliberate. Narrowing this by - # comparing the upstream's resolved columns was rejected: a query edit - # can move a filter, a join or a CASE threshold while leaving every - # column and type identical, and each changes every row the cube serves. - # Nothing short of reading the SQL tells those apart, so a changed query - # makes anything built from it suspect. Only NONE is skipped, the one - # case where DJ knows nothing material happened. - # + # Any tier rebuilds except NONE. # A rebuild can fail, and one cube's failure must not cost the remaining # downstreams theirs. if change_tier is not ChangeTier.NONE: @@ -4429,17 +4422,6 @@ async def revalidate_node( ] # A tier rather than a version, so propagation can hand the same value to # `bump_version` for downstream cubes. - # - # Any column change is major. An addition looks harmless -- nothing could - # already reference a column that did not exist -- but the query produced it, - # and a query edit can move a filter or a join while leaving the rest of the - # projection identical. Demoting additions to MINOR buys nothing anyway: the - # cube rebuild below skips only NONE, so a minor bump rebuilds all the same. - # - # `order_fixed` earns no tier. DJ filling in a missing projection index is its - # own bookkeeping, not a change to the node, so a revision would describe - # nothing and any tier above NONE would rebuild every cube below. It is applied - # to the current revision in place instead, below. change_tier = fold_change_tiers( [ ChangeTier.MAJOR diff --git a/datajunction-server/datajunction_server/models/decompose.py b/datajunction-server/datajunction_server/models/decompose.py index 351b2b527..bc7db2658 100644 --- a/datajunction-server/datajunction_server/models/decompose.py +++ b/datajunction-server/datajunction_server/models/decompose.py @@ -11,9 +11,10 @@ from typing import Any -from pydantic import BaseModel +from pydantic import BaseModel, Field from datajunction_server.enum import StrEnum +from datajunction_server.models.materialization import MaterializationTarget from datajunction_server.models.reaggregate import DimensionReaggregateRule @@ -64,34 +65,7 @@ class AggregationRule(BaseModel): class MetricComponent(BaseModel): - """ - A reusable, named building block of a metric definition. - - A MetricComponent represents a SQL expression that can serve as an input - to building a metric. It supports a two-phase aggregation model: - - - Phase 1 (Accumulate): Build from raw data using `aggregation` - Can be a function name ("SUM") or a template ("SUM(POWER({}, 2))") - - - Phase 2 (Merge): Combine pre-aggregated values using `merge` function - Examples: SUM, SUM (for COUNT), hll_union_agg - - For most aggregations, accumulate and merge use the same function (SUM → SUM). - For COUNT, merge is SUM (sum up the counts). - For HLL sketches, they differ: hll_sketch_estimate vs hll_union_agg. - - The final expression combining merged components is specified in - DecomposedMetric.combiner. - - Attributes: - name: A unique name for the component, derived from its expression. - expression: The raw SQL expression (column/value) being aggregated. - aggregation: Function name or template for Phase 1. Simple cases use - just the name ("SUM"), complex cases use templates with - {} placeholder ("SUM(POWER({}, 2))"). - merge: The function name for combining pre-aggregated values (Phase 2). - rule: Aggregation rules defining how/when the component can be aggregated. - """ + """A reusable, named building block of a metric definition.""" name: str expression: str @@ -109,6 +83,41 @@ class MetricComponent(BaseModel): # `measure_identity_token` -- because a sketch built at one accuracy must not # satisfy a query asking for another. params: dict[str, Any] | None = None + # Fixed arguments appended after the column in the Phase 2 merge call, as SQL + # literals: `nflx_tdigest_agg(col, 200.0)` is merge="nflx_tdigest_agg" plus + # merge_args=["200.0"]. `merge` stays a bare function name because two things + # match on it -- the Druid aggregator lookup key and the semi-additive rewrite + # in `_replace_reaggregate_merge_expression` -- so the tuning cannot be folded + # into it as a template. Rendered from `params` by the decomposition; `params` + # remains the identity record, this is the emission form. + merge_args: list[str] = Field(default_factory=list) + + # Conversion applied when writing this component to a materialized table for + # a target listed in `serialize_targets` -- see `ComponentDef.serialize`. The + # value is a template expanded against the accumulated expression. + serialize: str | None = None + # A list rather than a set: this model is dumped straight to JSON in cube + # materialization configs, and `json.dumps` cannot encode a set. StrEnum + # members are fine -- they subclass `str`. + serialize_targets: list[MaterializationTarget] = Field(default_factory=list) + # Column type `serialize` produces. See `ComponentDef.serialize_type`. + serialize_type: str | None = None + + def serializes_for(self, materialization_target: Any | None) -> bool: + """ + Whether this component converts its accumulated value for `target`. + + Two things key off this and they must agree: the SQL that writes the + column, and the type recorded for it. If they disagree the measures + table holds one representation while the catalog claims another, and the + Druid aggregator lookup -- which keys on the column type -- silently + finds nothing and drops the measure. + """ + return bool( + self.serialize + and materialization_target is not None + and materialization_target in self.serialize_targets, + ) @property def normalized_aggregation(self) -> str: @@ -139,30 +148,7 @@ class PreAggMeasure(MetricComponent): class DecomposedMetric(BaseModel): - """ - A metric decomposed into its constituent components with a combining expression. - - This is the result of decomposing a metric query. It specifies: - - components: The measures needed for pre-aggregation - - combiner: How to combine merged components into the final metric value - - derived_query: The full SQL query using the combiner - - Examples: - SUM metric: - components: [{name: "revenue_sum", aggregation: "SUM", merge: "SUM"}] - combiner: "SUM(revenue_sum)" - - AVG metric: - components: [ - {name: "revenue_sum", aggregation: "SUM", merge: "SUM"}, - {name: "revenue_count", aggregation: "COUNT", merge: "SUM"} - ] - combiner: "SUM(revenue_sum) / SUM(revenue_count)" - - APPROX_COUNT_DISTINCT metric (uses Spark function names): - components: [{name: "user_hll", aggregation: "hll_sketch_agg", merge: "hll_union"}] - combiner: "hll_sketch_estimate(hll_union(user_hll))" - """ + """A metric decomposed into its constituent components with a combining expression.""" components: list[MetricComponent] combiner: str # Expression combining merged components into final value diff --git a/datajunction-server/datajunction_server/models/deployment.py b/datajunction-server/datajunction_server/models/deployment.py index 9b01113a1..1e7daa9a2 100644 --- a/datajunction-server/datajunction_server/models/deployment.py +++ b/datajunction-server/datajunction_server/models/deployment.py @@ -218,37 +218,10 @@ class HierarchySpec(NamespacedSpec): class PreAggSpec(NamespacedSpec): """ Specification for an externally-built pre-aggregation table adopted at deploy - time (equivalent to POST /preaggs/register). ``name`` is a stable handle used - for reconciliation and availability callbacks. Metric/dimension references may - use ``${prefix}`` or be fully qualified; they are rendered against the - deployment namespace. - - Every metric and every dimension is declared together with the physical - column of the external table that holds it, as a map:: - - metrics: - ${prefix}paid_members: paid_members_sum - dimensions: - ${prefix}country_dim.country_iso: country - ${prefix}date_dim.utc_date: utc_date - - Both maps require a value for every key -- including a dimension whose - physical column happens to match its DJ column name, which is written out - rather than left empty. An optional value would make the map not really a - mapping, a trailing colon is easy to write by accident, and spelling the - physical name out documents the table in the file that declares it. - - What goes under ``metrics`` are the measures the table stores. A derived - metric (a ratio of two others, say) is not listed and cannot be: it has no - column of its own. It is covered anyway, because both registration and - query-time matching work on decomposed measure identities rather than metric - names, so any metric that decomposes into the stored measures resolves to - this table. - - The earlier four-field form -- ``metrics``/``dimensions`` as lists alongside - separate ``measure_columns``/``dimension_columns`` maps -- is no longer - accepted, and a spec still using it is rejected with a message describing - what to write instead. + time. ``name`` is a stable handle used for reconciliation. + + Every metric and dimension is declared together with its physical column + in the external table. """ # Metric/dimension reference -> the physical column of the external table @@ -363,16 +336,7 @@ class PartitionSpec(BaseModel): class MaterializationSpec(BaseModel): - """ - Declarative materialization config for a cube. - - Deliberately carries only what the author decides: when to build, how, how far - back to look, how much of history to serve, how long the result is kept, and the - Druid, Spark and platform settings their own site needs. The backend that runs it - (and everything it derives -- the measures queries, combiner SQL, the rest of the - Druid spec, output tables) is DJ's choice, so no `job` field is exposed here and - the block stays portable if that choice changes. - """ + """Declarative materialization config for a cube.""" schedule: str strategy: MaterializationStrategy = MaterializationStrategy.INCREMENTAL_TIME @@ -1316,20 +1280,11 @@ class CubeSpec(NodeSpec): filters: list[str] | None = None columns: list[ColumnSpec] | None = None # Tri-state. A spec materializes the cube on the schedule it names; the `none` - # sentinel tears down whatever the cube has materialized; absent -- and null, - # which is what serializing an absent field produces -- means the cube's - # materialization is not managed here and whatever exists is left alone. Only a - # value can carry intent through serialization, so removal is spelled - # `materialization: none` rather than inferred from a key being present. + # sentinel tears down whatever the cube has materialized; absent means the cube's + # materialization is not managed here and whatever exists is left alone. # - # A list declares more than one, which a cube legitimately needs: an - # `incremental_time` build for freshness alongside a periodic `full` rebuild - # that corrects late-arriving data, out-of-order events and dimension - # backfills. `strategy` is what tells two entries apart -- `job` is not - # authorable (see `MaterializationSpec`) and everything else is a knob rather - # than an identity -- so two entries sharing one are rejected below. The scalar - # form stays valid and means exactly what it always did; nothing has to be - # rewritten as a one-element list. + # A list declares more than one (e.g. incremental + full). `strategy` is what tells + # two entries apart, so two entries sharing one are rejected below. materialization: ( MaterializationSpec | list[MaterializationSpec] | MaterializationAction | None ) = None @@ -1346,18 +1301,8 @@ class CubeSpec(NodeSpec): # Only user-authored partition config is compared here (see __eq__); # everything else about cube columns is auto-derived. "columns": ChangeTier.MAJOR, - # Materialization is not part of a cube's definition. Configuring one - # through `POST /nodes/{name}/materialization/` has never created a node - # revision, and a YAML-declared schedule has to behave the same way, or - # the same edit would cut a version through one door and not the other. - # Because `materialization` is also excluded from `__eq__`, a - # materialization-only edit leaves the cube in the deployment's skip list - # and never reaches the classifier at all; NONE records that intent rather - # than describing a reachable code path. Nothing is lost by the exclusion: - # reconciling the declared block against the persisted config is a separate - # concern that runs over every declared cube whether or not the node - # changed, which is also what catches a materialization changed outside - # YAML. + # Materialization is not part of a cube's definition, so changing it + # does not mint a new node revision. "materialization": ChangeTier.NONE, } diff --git a/datajunction-server/datajunction_server/models/materialization.py b/datajunction-server/datajunction_server/models/materialization.py index c1e8b45a6..0dfed5393 100644 --- a/datajunction-server/datajunction_server/models/materialization.py +++ b/datajunction-server/datajunction_server/models/materialization.py @@ -32,6 +32,22 @@ if TYPE_CHECKING: from datajunction_server.database.node import NodeRevision + +class MaterializationTarget(StrEnum): + """ + Where a measures table is being written. + + Distinct from ``Dialect``: measures SQL executes in Spark for a Druid cube + and an Iceberg pre-agg alike, so the dialect cannot tell the two apart. Only + the target can, and some components must be written differently depending on + it -- a sketch whose in-engine representation is not the one the destination + reads needs converting on the way out. + """ + + DRUID = "druid" + ICEBERG = "iceberg" + + DRUID_AGG_MAPPING = { ("bigint", "sum"): "longSum", ("int", "sum"): "longSum", diff --git a/datajunction-server/datajunction_server/sql/decompose.py b/datajunction-server/datajunction_server/sql/decompose.py index b26c0e81b..70868a83c 100644 --- a/datajunction-server/datajunction_server/sql/decompose.py +++ b/datajunction-server/datajunction_server/sql/decompose.py @@ -17,6 +17,7 @@ MetricComponent, ) from datajunction_server.models.dialect import Dialect +from datajunction_server.models.materialization import MaterializationTarget from datajunction_server.models.node_type import NodeType from datajunction_server.models.reaggregate import ( ReaggregateSpec, @@ -104,7 +105,29 @@ class ComponentDef: suffix: str accumulate: str merge: str + # Fixed arguments appended after the column in the merge call, as SQL literals. + # `nflx_tdigest_agg(col, 200.0)` is merge="nflx_tdigest_agg" with + # merge_args=("200.0",). Kept separate from `merge` rather than templated into + # it because the Druid aggregator mapping and the semi-additive rewrite both + # match on the bare function name. Omitting a required tuning argument is not + # an error the engine reports: `nflx_tdigest_agg` has a one-argument overload + # that silently falls back to a near-useless compression and collapses the + # digest to a single centroid at the mean. + merge_args: tuple[str, ...] = () arg_index: int | None = 0 # Which arg to use, or None for multi-arg templates + # Conversion applied to the accumulated value when writing a measures table + # for a target that cannot read the in-engine representation. A template, as + # `accumulate` is: "nflx_tdigest_sketch({})". Applied only at materialization + # time and only for the targets in `serialize_targets`, so query-time SQL and + # every other destination keep the unwrapped value. + serialize: str | None = None + # Targets that need `serialize`. Empty means it never applies. + serialize_targets: tuple[MaterializationTarget, ...] = () + # Column type the conversion produces, e.g. "binary". Declared rather than + # inferred: a sketch accumulate takes more arguments than type inference + # feeds it, so inference falls back to the metric type and would silently + # record the wrong type for the column. + serialize_type: str | None = None class AggDecomposition(ABC): @@ -140,7 +163,14 @@ def combine( func: ast.Function, dialect: Dialect = Dialect.SPARK, ) -> ast.Expression: - """Build the combiner expression from merged metric components.""" + """ + Build the combiner expression from merged metric components. + + Returns Spark-canonical SQL; translation to other dialects happens in + the transpilation layer, so re-spelling function names, array literals + or index bases here applies them twice. Branch on `dialect` only for an + engine needing a structurally different expression. + """ # ============================================================================= @@ -1892,11 +1922,15 @@ def _make_component( expression=expression, aggregation=None if is_distinct else accumulate_expr, merge=None if is_distinct else comp_def.merge, + merge_args=[] if is_distinct else list(comp_def.merge_args), rule=AggregationRule( type=Aggregability.LIMITED if is_distinct else Aggregability.FULL, level=[str(a) for a in func.args] if is_distinct else None, ), grain_alias=grain_alias, + serialize=comp_def.serialize, + serialize_targets=list(comp_def.serialize_targets), + serialize_type=comp_def.serialize_type, ) def _expand_template(self, template: str, args: list) -> str: diff --git a/datajunction-server/datajunction_server/sql/parsing/ast.py b/datajunction-server/datajunction_server/sql/parsing/ast.py index a4a49d51a..3c2570f19 100644 --- a/datajunction-server/datajunction_server/sql/parsing/ast.py +++ b/datajunction-server/datajunction_server/sql/parsing/ast.py @@ -1680,29 +1680,6 @@ async def add_ref_column( """ Add column referenced from this table. Returns True if the table has the column and False otherwise. - - This function handles the following cases: - - Regular columns. For example: - (1) non-aliased columns - `SELECT country_id AS country1 FROM countries` should match the `country_id` - column in the table `countries` - (2) aliased columns - `SELECT C.country_id AS country1 FROM countries C` should match the `country_id` - column in the table `countries` with the column namespace/table alias `C` - - Struct columns. For example: - (1) non-aliased struct columns - `countries` has column `identifiers` with type: - STRUCT - `SELECT identifiers.country_name AS name FROM countries` should match the - `identifier` -> `country_name` column in the table `countries` - (2) aliased struct columns - `countries` has column `identifiers` with type: - STRUCT - `SELECT C.identifiers.country_name AS name FROM countries C` should match the - `identifier` -> `country_name` column in the table `countries` with the column namespace/ - table alias `C` """ if not self._columns: if ctx is None: diff --git a/datajunction-server/tests/api/cubes_test.py b/datajunction-server/tests/api/cubes_test.py index 410233d61..c6d40f21a 100644 --- a/datajunction-server/tests/api/cubes_test.py +++ b/datajunction-server/tests/api/cubes_test.py @@ -3652,6 +3652,10 @@ async def test_cube_materialization_metadata( "expression": "*", "grain_alias": None, "params": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "merge": "SUM", "name": "count_c8e42e74", "rule": { @@ -3666,6 +3670,10 @@ async def test_cube_materialization_metadata( "expression": "if(discount > 0.0, 1, 0)", "grain_alias": None, "params": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "merge": "SUM", "name": "discount_sum_30b84e6c", "rule": { @@ -3680,6 +3688,10 @@ async def test_cube_materialization_metadata( "expression": "price", "grain_alias": None, "params": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "merge": "SUM", "name": "price_count_935e7117", "rule": { @@ -3694,6 +3706,10 @@ async def test_cube_materialization_metadata( "expression": "price * discount", "grain_alias": None, "params": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "merge": "SUM", "name": "price_discount_sum_e4ba5456", "rule": { @@ -3708,6 +3724,10 @@ async def test_cube_materialization_metadata( "expression": "price", "grain_alias": None, "params": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "merge": "SUM", "name": "price_sum_935e7117", "rule": { @@ -3722,6 +3742,10 @@ async def test_cube_materialization_metadata( "expression": "repair_order_id", "grain_alias": None, "params": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "merge": "SUM", "name": "repair_order_id_count_bd241964", "rule": { @@ -3736,6 +3760,10 @@ async def test_cube_materialization_metadata( "expression": "total_repair_cost", "grain_alias": None, "params": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "merge": "SUM", "name": "total_repair_cost_sum_67874507", "rule": { @@ -3858,6 +3886,10 @@ async def test_cube_materialization_metadata( "expression": "price", "grain_alias": None, "params": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "merge": "SUM", "name": "price_sum_252381cf", "rule": { @@ -4115,6 +4147,10 @@ async def test_cube_materialization_metadata( "expression": "*", "grain_alias": None, "params": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "aggregation": "COUNT", "merge": "SUM", "rule": { @@ -4129,6 +4165,10 @@ async def test_cube_materialization_metadata( "expression": "if(discount > 0.0, 1, 0)", "grain_alias": None, "params": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "aggregation": "SUM", "merge": "SUM", "rule": { @@ -4143,6 +4183,10 @@ async def test_cube_materialization_metadata( "expression": "price", "grain_alias": None, "params": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "aggregation": "COUNT", "merge": "SUM", "rule": { @@ -4157,6 +4201,10 @@ async def test_cube_materialization_metadata( "expression": "price * discount", "grain_alias": None, "params": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "aggregation": "SUM", "merge": "SUM", "rule": { @@ -4171,6 +4219,10 @@ async def test_cube_materialization_metadata( "expression": "price", "grain_alias": None, "params": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "aggregation": "SUM", "merge": "SUM", "rule": { @@ -4185,6 +4237,10 @@ async def test_cube_materialization_metadata( "expression": "repair_order_id", "grain_alias": None, "params": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "aggregation": "COUNT", "merge": "SUM", "rule": { @@ -4199,6 +4255,10 @@ async def test_cube_materialization_metadata( "expression": "total_repair_cost", "grain_alias": None, "params": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "aggregation": "SUM", "merge": "SUM", "rule": { @@ -4213,6 +4273,10 @@ async def test_cube_materialization_metadata( "expression": "price", "grain_alias": None, "params": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": 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 90b2bfc0e..91402b042 100644 --- a/datajunction-server/tests/api/metrics_test.py +++ b/datajunction-server/tests/api/metrics_test.py @@ -465,6 +465,10 @@ async def test_read_metrics(module__client_with_roads: AsyncClient) -> None: "aggregation": "SUM", "expression": "if(discount > 0.0, 1, 0)", "grain_alias": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "params": None, "name": "discount_sum_30b84e6c", "merge": "SUM", @@ -479,6 +483,10 @@ async def test_read_metrics(module__client_with_roads: AsyncClient) -> None: "aggregation": "COUNT", "expression": "*", "grain_alias": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "params": None, "merge": "SUM", "name": "count_c8e42e74", diff --git a/datajunction-server/tests/api/preaggregations_test.py b/datajunction-server/tests/api/preaggregations_test.py index 5d73a6b99..7c484ab76 100644 --- a/datajunction-server/tests/api/preaggregations_test.py +++ b/datajunction-server/tests/api/preaggregations_test.py @@ -736,6 +736,10 @@ async def test_get_preagg_by_id(self, client_with_preaggs): "expr_hash": "83632b779d87", "expression": "line_total", "grain_alias": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "params": None, "merge": "SUM", "name": "line_total_sum_e1f61696", @@ -806,6 +810,10 @@ async def test_get_preagg_by_id(self, client_with_preaggs): "expr_hash": "221d2a4bfdae", "expression": "quantity", "grain_alias": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "params": None, "merge": "SUM", "name": "quantity_sum_06b64d2e", diff --git a/datajunction-server/tests/api/sql_v2_test.py b/datajunction-server/tests/api/sql_v2_test.py index 659f3fb91..945d070f1 100644 --- a/datajunction-server/tests/api/sql_v2_test.py +++ b/datajunction-server/tests/api/sql_v2_test.py @@ -1201,6 +1201,10 @@ async def create_metric_distinct_single_column(client: AsyncClient): "aggregation": None, "expression": "hard_hat_id", "grain_alias": "hard_hat_id", + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "params": None, "merge": None, "name": "hard_hat_id", @@ -1234,6 +1238,10 @@ 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", + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "params": None, "merge": None, "name": "hard_hat_id_distinct_0291ee39", @@ -1635,6 +1643,10 @@ async def test_metric_definitions_with_nonjoinable_dimensions( "aggregation": "SUM", "expression": "default.local_hard_hats_2.hard_hat_id", "grain_alias": None, + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "params": None, "merge": "SUM", "name": "default_DOT_local_hard_hats_2_DOT_hard_hat_id_sum_bf8a8419", @@ -1702,6 +1714,10 @@ async def test_metric_definitions_with_single_joinable_dimensions( "aggregation": None, "expression": "default.municipality_dim.contact_name", "grain_alias": "contact_name", + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "params": None, "merge": None, "name": "contact_name", @@ -1810,6 +1826,10 @@ 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", + "merge_args": [], + "serialize": None, + "serialize_targets": [], + "serialize_type": None, "params": None, "merge": None, "name": "default_DOT_hard_hat_DOT_state_default_DOT_hard_hat_DOT_first_name_distinct_1a99d6a7", diff --git a/datajunction-server/tests/construction/build_v3/accumulate_type_test.py b/datajunction-server/tests/construction/build_v3/accumulate_type_test.py new file mode 100644 index 000000000..43c048ca4 --- /dev/null +++ b/datajunction-server/tests/construction/build_v3/accumulate_type_test.py @@ -0,0 +1,152 @@ +""" +Tests for typing an accumulate whose outermost call takes several arguments. + +Type inference feeds a component exactly one input type. That is right for +every aggregation DJ ships, including the templated ones -- ``SUM(POWER(x, 2))`` +is still a single-argument ``SUM`` on the outside. A sketch is the first shape +that breaks it: ``nflx_tdigest(latency_ms, CAST(200.0 AS DOUBLE))`` needs two, +so inference raised ``TypeError``, the caller swallowed it, and the column was +recorded as the metric's own type. + +That mattered beyond cosmetics. The recorded type is persisted on the +pre-aggregation row and forwarded to the query service, which uses it to +create the materialized table -- so a struct-valued sketch column was being +described as ``double``. + +``POWER`` stands in for a sketch accumulate here: OSS registers no sketch +families, and it is the available two-argument function whose inferred result +differs from the fallback. +""" + +from types import SimpleNamespace + +import pytest + +from datajunction_server.construction.build_v3.measures import ( + _multi_argument_accumulate_types, + infer_component_type, +) +from datajunction_server.models.decompose import AggregationRule, MetricComponent +from datajunction_server.models.materialization import MaterializationTarget + + +def _parent(column_type: str = "double"): + return SimpleNamespace( + current=SimpleNamespace( + columns=[SimpleNamespace(name="latency_ms", type=column_type)], + ), + ) + + +def _component(aggregation: str, **kwargs) -> MetricComponent: + return MetricComponent( + name="latency_acc_abc123", + expression="latency_ms", + aggregation=aggregation, + merge="SUM", + rule=AggregationRule(), + **kwargs, + ) + + +class TestMultiArgumentTypes: + """The helper resolves every argument, or declines entirely.""" + + def test_resolves_a_column_and_a_literal(self): + types = _multi_argument_accumulate_types("POWER(latency_ms, 2)", _parent()) + assert [str(t) for t in types] == ["double", "int"] + + def test_resolves_a_cast_without_binding_it_to_a_table(self): + # Literals and casts carry their own type; only columns need the parent. + types = _multi_argument_accumulate_types( + "POWER(latency_ms, CAST(200.0 AS DOUBLE))", + _parent(), + ) + assert [str(t) for t in types] == ["double", "double"] + + def test_declines_without_a_parent_to_resolve_columns_against(self): + assert _multi_argument_accumulate_types("POWER(latency_ms, 2)", None) is None + + def test_declines_for_a_bare_function_name(self): + # "SUM" is a name, not a call -- the overwhelmingly common shape. + assert _multi_argument_accumulate_types("SUM", _parent()) is None + + def test_declines_for_a_single_argument_call(self): + # Including templated ones, whose outermost call still takes one arg. + assert _multi_argument_accumulate_types("SUM(POWER(latency_ms, 2))", _parent()) is None + + def test_declines_when_a_column_type_is_unrecognized(self): + # Better to fall back than to invent a type for the stored schema. + assert ( + _multi_argument_accumulate_types( + "POWER(latency_ms, 2)", + _parent("some_unknown_type"), + ) + is None + ) + + +class TestInferredColumnType: + """What the measures column ends up recorded as.""" + + def test_multi_argument_accumulate_is_typed_from_its_arguments(self): + # The regression: this used to return the metric type, because + # POWER.infer_type(double) raises and the error was swallowed. + assert ( + infer_component_type( + _component("POWER(latency_ms, 2)"), + "bigint", + _parent(), + ) + == "double" + ) + + def test_single_argument_accumulate_is_unchanged(self): + assert ( + infer_component_type(_component("SUM"), "bigint", _parent()) == "double" + ) + + def test_unresolvable_multi_argument_accumulate_falls_back(self): + assert ( + infer_component_type( + _component("POWER(latency_ms, 2)"), + "bigint", + _parent("some_unknown_type"), + ) + == "bigint" + ) + + @pytest.mark.parametrize("target", [None, MaterializationTarget.ICEBERG]) + def test_unconverted_targets_get_the_argument_derived_type(self, target): + assert ( + infer_component_type( + _component("POWER(latency_ms, 2)"), + "bigint", + _parent(), + target, + ) + == "double" + ) + + def test_declared_serialize_type_still_wins_for_its_target(self): + """ + The conversion type takes precedence over the accumulated one. + + Druid stores the converted representation, so the column must be + described as that -- the aggregator lookup keys on it. + """ + component = _component( + "POWER(latency_ms, 2)", + serialize="to_druid_bytes({})", + serialize_targets=[MaterializationTarget.DRUID], + serialize_type="binary", + ) + assert ( + infer_component_type( + component, + "bigint", + _parent(), + MaterializationTarget.DRUID, + ) + == "binary" + ) diff --git a/datajunction-server/tests/construction/build_v3/merge_args_test.py b/datajunction-server/tests/construction/build_v3/merge_args_test.py new file mode 100644 index 000000000..e044b4361 --- /dev/null +++ b/datajunction-server/tests/construction/build_v3/merge_args_test.py @@ -0,0 +1,129 @@ +""" +Tests for fixed merge arguments on a component's Phase 2 call. + +Sketch families need tuning passed to the merge as well as the accumulate -- +`nflx_tdigest_agg(digest, compression)` -- but `merge` itself has to stay a bare +function name, because the Druid aggregator mapping and the semi-additive +rewrite both match on it. `merge_args` carries the rest. + +The omission this guards against is silent: a sketch merge called without its +compression is accepted by the engine and returns a digest collapsed to a single +centroid, so every quantile comes back equal to the mean. +""" + +from datajunction_server.construction.build_v3.decomposition import build_merge_call +from datajunction_server.models.decompose import ( + AggregationRule, + Aggregability, + MetricComponent, +) +from datajunction_server.sql.decompose import ComponentDef +from datajunction_server.sql.parsing import ast + + +def _column(name: str = "latency_tdigest_ab12") -> ast.Column: + return ast.Column(name=ast.Name(name)) + + +class TestBuildMergeCall: + """The helper both emission sites use.""" + + def test_no_args_is_the_plain_single_argument_call(self): + # Every non-sketch aggregation takes this path and must be unchanged. + expr = build_merge_call("SUM", [], _column("revenue_sum_ab12")) + assert str(expr) == "SUM(revenue_sum_ab12)" + + def test_one_fixed_argument_is_appended_after_the_column(self): + expr = build_merge_call( + "nflx_tdigest_agg", + ["CAST(200.0 AS DOUBLE)"], + _column(), + ) + assert str(expr) == ( + "nflx_tdigest_agg(latency_tdigest_ab12, CAST(200.0 AS DOUBLE))" + ) + + def test_several_fixed_arguments_keep_their_order(self): + expr = build_merge_call("sketch_merge", ["12", "'HLL_4'"], _column("s_ab12")) + assert str(expr) == "sketch_merge(s_ab12, 12, 'HLL_4')" + + def test_the_column_is_always_the_first_argument(self): + expr = build_merge_call("f", ["1"], _column("c")) + assert isinstance(expr, ast.Function) + assert str(expr.args[0]) == "c" + + def test_the_merge_name_is_preserved_verbatim(self): + # Matched on elsewhere, so casing must not be normalized. + expr = build_merge_call("nflx_tdigest_agg", ["1.0"], _column("c")) + assert expr.name.name == "nflx_tdigest_agg" + + def test_arguments_are_parsed_not_pasted(self): + # A literal arrives as an expression node, not raw text, so it renders + # through the same dialect machinery as the rest of the tree. + expr = build_merge_call("f", ["1 + 2"], _column("c")) + assert isinstance(expr.args[1], ast.Expression) + assert str(expr) == "f(c, 1 + 2)" + + +class TestComponentDefDefaults: + """Existing decompositions must be untouched by the new field.""" + + def test_merge_args_defaults_to_empty(self): + assert ComponentDef(suffix="_sum", accumulate="SUM", merge="SUM").merge_args == () + + def test_merge_args_is_a_tuple_so_the_default_cannot_be_mutated(self): + first = ComponentDef(suffix="_sum", accumulate="SUM", merge="SUM") + second = ComponentDef(suffix="_cnt", accumulate="COUNT", merge="SUM") + assert first.merge_args == second.merge_args == () + assert isinstance(first.merge_args, tuple) + + def test_merge_args_is_declared_on_the_component_def(self): + comp = ComponentDef( + suffix="_tdigest", + accumulate="nflx_tdigest({}, CAST(200.0 AS DOUBLE))", + merge="nflx_tdigest_agg", + merge_args=("CAST(200.0 AS DOUBLE)",), + ) + assert comp.merge == "nflx_tdigest_agg" + assert comp.merge_args == ("CAST(200.0 AS DOUBLE)",) + + +class TestMetricComponentField: + """`merge_args` survives the model, which is dumped to JSON in cube configs.""" + + def _component(self, merge_args): + return MetricComponent( + name="latency_tdigest_ab12", + expression="latency_ms", + aggregation="nflx_tdigest", + merge="nflx_tdigest_agg", + merge_args=merge_args, + rule=AggregationRule(type=Aggregability.FULL), + ) + + def test_defaults_to_empty_list(self): + component = MetricComponent( + name="revenue_sum_ab12", + expression="revenue", + aggregation="SUM", + merge="SUM", + rule=AggregationRule(type=Aggregability.FULL), + ) + assert component.merge_args == [] + + def test_round_trips_through_json(self): + # A list rather than a tuple: these models are serialized straight into + # materialization configs. + component = self._component(["CAST(200.0 AS DOUBLE)"]) + assert component.model_dump()["merge_args"] == ["CAST(200.0 AS DOUBLE)"] + + def test_feeds_the_merge_call(self): + component = self._component(["CAST(200.0 AS DOUBLE)"]) + expr = build_merge_call( + component.merge, + component.merge_args, + _column(component.name), + ) + assert str(expr) == ( + "nflx_tdigest_agg(latency_tdigest_ab12, CAST(200.0 AS DOUBLE))" + ) diff --git a/datajunction-server/tests/construction/build_v3/serialize_target_test.py b/datajunction-server/tests/construction/build_v3/serialize_target_test.py new file mode 100644 index 000000000..ef5bbcf76 --- /dev/null +++ b/datajunction-server/tests/construction/build_v3/serialize_target_test.py @@ -0,0 +1,244 @@ +""" +Tests for the materialization-target serialize hook. + +Some components must be written to destinations in a representation different +from their accumulated form. The fork keys off ``MaterializationTarget`` to +determine if the expression needs to be serialized during measures table writing. +""" + +import pytest + +from types import SimpleNamespace + +from datajunction_server.construction.build_v3.decomposition import ( + build_component_expression, +) +from datajunction_server.construction.build_v3.measures import infer_component_type +from datajunction_server.models.decompose import AggregationRule, MetricComponent +from datajunction_server.models.materialization import MaterializationTarget + + +def _component( + aggregation: str, + serialize: str | None = None, + serialize_targets: list[MaterializationTarget] | None = None, +) -> MetricComponent: + return MetricComponent( + name="latency_td_abc123", + expression="latency_ms", + aggregation=aggregation, + merge="nflx_tdigest_agg", + rule=AggregationRule(), + serialize=serialize, + serialize_targets=serialize_targets or [], + ) + + +DRUID = [MaterializationTarget.DRUID] + + +class TestSerializeIsApplied: + """When the component declares a conversion and the target asks for it.""" + + def test_simple_accumulate_is_wrapped(self): + """A bare function-name accumulate gets the conversion.""" + component = _component( + "nflx_tdigest", + serialize="nflx_tdigest_sketch({})", + serialize_targets=DRUID, + ) + + expr = build_component_expression(component, MaterializationTarget.DRUID) + + assert str(expr) == "nflx_tdigest_sketch(nflx_tdigest(latency_ms))" + + def test_template_accumulate_is_wrapped(self): + """ + A pre-expanded template accumulate gets the conversion too. + """ + component = _component( + "nflx_tdigest(latency_ms, 200)", + serialize="nflx_tdigest_sketch({})", + serialize_targets=DRUID, + ) + + expr = build_component_expression(component, MaterializationTarget.DRUID) + + assert str(expr) == "nflx_tdigest_sketch(nflx_tdigest(latency_ms, 200))" + + +class TestSerializeIsSkipped: + """Every path that must leave the accumulated expression untouched.""" + + def test_no_target_means_no_conversion(self): + """ + A query-time build passes no target and gets the unwrapped expression. + """ + component = _component( + "nflx_tdigest(latency_ms, 200)", + serialize="nflx_tdigest_sketch({})", + serialize_targets=DRUID, + ) + + expr = build_component_expression(component) + + assert str(expr) == "nflx_tdigest(latency_ms, 200)" + + def test_different_target_means_no_conversion(self): + """ + An Iceberg pre-agg keeps the in-engine representation. + """ + component = _component( + "nflx_tdigest(latency_ms, 200)", + serialize="nflx_tdigest_sketch({})", + serialize_targets=DRUID, + ) + + expr = build_component_expression(component, MaterializationTarget.ICEBERG) + + assert str(expr) == "nflx_tdigest(latency_ms, 200)" + + def test_component_without_serialize_is_untouched(self): + """ + A component declaring no conversion is unaffected even for Druid. + """ + component = _component("SUM") + + expr = build_component_expression(component, MaterializationTarget.DRUID) + + assert str(expr) == "SUM(latency_ms)" + + @pytest.mark.parametrize( + "target", + [None, MaterializationTarget.DRUID, MaterializationTarget.ICEBERG], + ) + def test_plain_sum_is_stable_across_targets(self, target): + """ + An ordinary measure renders identically whatever the target. + """ + expr = build_component_expression(_component("SUM"), target) + + assert str(expr) == "SUM(latency_ms)" + + +class TestSerializeDeclaration: + """The declaration itself.""" + + def test_targets_default_to_empty(self): + """ + A component declaring a conversion but no targets never applies it. + + Defaulting to "no targets" rather than "all targets" prevents + half-finished declarations from silently rewriting measures tables. + """ + component = _component( + "nflx_tdigest(latency_ms, 200)", + serialize="nflx_tdigest_sketch({})", + ) + + expr = build_component_expression(component, MaterializationTarget.DRUID) + + assert str(expr) == "nflx_tdigest(latency_ms, 200)" + + +class TestSerializedColumnType: + """ + The recorded column type must match the representation actually written. + + These are set in different places (SQL by ``build_component_expression``, + type by ``infer_component_type``). A disagreement is silent and could cause + Druid ingestion to miss the correct aggregator mappings. + """ + + @staticmethod + def _parent(): + return SimpleNamespace( + current=SimpleNamespace( + columns=[SimpleNamespace(name="latency_ms", type="double")], + ), + ) + + def _sketch_component(self, serialize_type: str | None = "binary"): + """ + A component that converts on the way to Druid. + """ + component = _component( + "SUM", + serialize="to_druid_bytes({})", + serialize_targets=DRUID, + ) + component.serialize_type = serialize_type + return component + + def test_druid_target_records_the_converted_type(self): + """The declared conversion type wins for the target that converts.""" + assert ( + infer_component_type( + self._sketch_component(), + "double", + self._parent(), + MaterializationTarget.DRUID, + ) + == "binary" + ) + + @pytest.mark.parametrize( + "target", + [None, MaterializationTarget.ICEBERG], + ) + def test_other_targets_keep_the_accumulated_type(self, target): + """Query-time and Iceberg store the unconverted value, so keep its type.""" + assert ( + infer_component_type( + self._sketch_component(), + "double", + self._parent(), + target, + ) + == "double" + ) + + def test_conversion_without_a_declared_type_falls_through(self): + """ + Declaring a conversion but no type leaves inference in charge. + """ + assert ( + infer_component_type( + self._sketch_component(serialize_type=None), + "double", + self._parent(), + MaterializationTarget.DRUID, + ) + == "double" + ) + + def test_ordinary_component_is_unaffected(self): + """A component with no conversion types exactly as it did before.""" + assert ( + infer_component_type( + _component("SUM"), + "double", + self._parent(), + MaterializationTarget.DRUID, + ) + == "double" + ) + + +class TestSerializesFor: + """The single predicate both the SQL and the type side consult.""" + + @pytest.mark.parametrize( + ("serialize", "targets", "target", "expected"), + [ + ("f({})", DRUID, MaterializationTarget.DRUID, True), + ("f({})", DRUID, MaterializationTarget.ICEBERG, False), + ("f({})", DRUID, None, False), + ("f({})", [], MaterializationTarget.DRUID, False), + (None, DRUID, MaterializationTarget.DRUID, False), + ], + ) + def test_predicate(self, serialize, targets, target, expected): + component = _component("SUM", serialize=serialize, serialize_targets=targets) + + assert component.serializes_for(target) is expected diff --git a/datajunction-server/tests/database/preaggregation_test.py b/datajunction-server/tests/database/preaggregation_test.py index fbe61eaea..a88bac361 100644 --- a/datajunction-server/tests/database/preaggregation_test.py +++ b/datajunction-server/tests/database/preaggregation_test.py @@ -719,14 +719,10 @@ async def test_find_matching_no_candidates( class TestMeasureIdentityParams: - """ - Tuning parameters participate in measure identity. - """ + """Tuning parameters participate in measure identity.""" def test_existing_tokens_are_unchanged(self): - """ - A component without params produces the token byte-for-byte. - """ + """A component without params produces the identical token.""" 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" @@ -742,9 +738,7 @@ def test_differing_params_differ(self): 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. - """ + """Order and type variations (e.g. 200 vs 200.0) do not change the identity token.""" assert measure_identity_token( "abc123", "TDIGEST_AGG", 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 eb92ce22c..3a7539787 100644 --- a/datajunction-server/tests/internal/nodes/derive_frozen_measures_test.py +++ b/datajunction-server/tests/internal/nodes/derive_frozen_measures_test.py @@ -374,13 +374,8 @@ def test_frozen_measure_conflict_rejects_different_measure_identity(): 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. + A name collision fails when only tuning parameters differ, preventing + a metric from silently reusing a sketch built with a different accuracy. """ frozen_measure = FrozenMeasure( name="latency_tdigest", @@ -404,11 +399,7 @@ def test_frozen_measure_conflict_rejects_different_tuning_params(): 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. + Identical params are reusable; absent-vs-empty is not a conflict. """ frozen_measure = FrozenMeasure( name="latency_tdigest", diff --git a/datajunction-server/tests/models/cube_druid_sketch_spec_test.py b/datajunction-server/tests/models/cube_druid_sketch_spec_test.py new file mode 100644 index 000000000..09a3bb067 --- /dev/null +++ b/datajunction-server/tests/models/cube_druid_sketch_spec_test.py @@ -0,0 +1,196 @@ +""" +Regression tests for sketch measures reaching Druid ingestion. +""" + +from datajunction_server.api.cubes import _build_metrics_spec +from datajunction_server.models.cube_materialization import CombineMaterialization +from datajunction_server.models.decompose import AggregationRule, MetricComponent +from datajunction_server.models.node_type import NodeNameVersion +from datajunction_server.models.partition import Granularity +from datajunction_server.models.query import ColumnMetadata + +TIMESTAMP_COLUMN = ColumnMetadata(name="order_date", type="int") + + +def _measure( + name: str, + aggregation: str | None, + merge: str | None = None, +) -> MetricComponent: + """A measure component as the decomposer would emit it.""" + return MetricComponent( + name=name, + expression=name, + aggregation=aggregation, + merge=merge, + rule=AggregationRule(), + ) + + +def _combiner( + measures: list[MetricComponent], + measure_columns: list[ColumnMetadata], +) -> CombineMaterialization: + """A combiner stage carrying real measures, ready to build a Druid spec.""" + return CombineMaterialization( + node=NodeNameVersion(name="default.repairs_cube", version="v1.0"), + columns=[TIMESTAMP_COLUMN, *measure_columns], + grain=["order_date"], + dimensions=["order_date"], + measures=measures, + timestamp_column="order_date", + timestamp_format="yyyyMMdd", + granularity=Granularity.DAY, + ) + + +class TestCombinerSketchMetricsSpec: + """``CombineMaterialization.metrics_spec`` with measures that carry sketches.""" + + def test_hll_measure_maps_to_hll_sketch_merge(self): + """ + An HLL sketch measure reaches Druid as an HLLSketchMerge aggregator. + """ + combiner = _combiner( + measures=[ + _measure("customer_id_hll_23002251", "hll_sketch_agg", "hll_union_agg"), + ], + measure_columns=[ + ColumnMetadata(name="customer_id_hll_23002251", type="binary"), + ], + ) + + assert combiner.metrics_spec() == [ + { + "fieldName": "customer_id_hll_23002251", + "name": "customer_id_hll_23002251", + "type": "HLLSketchMerge", + "lgK": 12, + "tgtHllType": "HLL_4", + }, + ] + + def test_hll_measure_is_neither_dropped_nor_downgraded(self): + """ + The sketch measure survives as a sketch aggregator. + """ + combiner = _combiner( + measures=[ + _measure("customer_id_hll_23002251", "hll_sketch_agg", "hll_union_agg"), + ], + measure_columns=[ + ColumnMetadata(name="customer_id_hll_23002251", type="binary"), + ], + ) + + spec = combiner.metrics_spec() + + assert len(spec) == 1, "sketch measure was dropped from the metricsSpec" + assert spec[0]["type"] != "longSum" + + def test_plain_sum_measure_maps_to_long_sum(self): + """A non-sketch measure still maps as before.""" + combiner = _combiner( + measures=[_measure("total_repair_cost", "SUM", "SUM")], + measure_columns=[ + ColumnMetadata(name="total_repair_cost", type="bigint"), + ], + ) + + assert combiner.metrics_spec() == [ + { + "fieldName": "total_repair_cost", + "name": "total_repair_cost", + "type": "longSum", + }, + ] + + def test_merge_phase_takes_precedence_over_accumulate(self): + """ + Pins which phase the aggregator lookup uses: ``merge``, not ``aggregation``. + Druid ingests already-accumulated partials, so the merge function is the + correct phase to key off. + """ + combiner = _combiner( + measures=[_measure("order_count", "COUNT", "SUM")], + measure_columns=[ColumnMetadata(name="order_count", type="double")], + ) + + assert combiner.metrics_spec() == [ + { + "fieldName": "order_count", + "name": "order_count", + "type": "doubleSum", + }, + ] + + def test_percentile_measure_is_dropped_silently(self): + """ + A percentile measure disappears from the ingestion spec without error. + """ + combiner = _combiner( + measures=[_measure("p95_latency", "approx_percentile")], + measure_columns=[ColumnMetadata(name="p95_latency", type="double")], + ) + + assert combiner.metrics_spec() == [] + + +class TestMetricsSpecBuildersAgree: + """ + The two Druid metricsSpec builders, compared directly. + + ``CombineMaterialization.metrics_spec`` and ``api.cubes._build_metrics_spec`` + map measures for different entry points. They share ``get_druid_aggregator_spec`` + so the aggregator type and family config agree. + """ + + def test_builders_agree_on_hll_aggregator_type(self): + """Both resolve an HLL sketch measure to the same aggregator type.""" + measure = _measure( + "customer_id_hll_23002251", + "hll_sketch_agg", + "hll_union_agg", + ) + measure_column = ColumnMetadata( + name="customer_id_hll_23002251", + type="binary", + ) + + v3_spec = _combiner([measure], [measure_column]).metrics_spec() + cubes_api_spec = _build_metrics_spec([measure_column], [measure], {}) + + assert v3_spec[0]["type"] == cubes_api_spec[0]["type"] == "HLLSketchMerge" + + def test_builders_agree_on_hll_tuning_config(self): + """ + Both builders attach the same HLL tuning config. + """ + measure = _measure( + "customer_id_hll_23002251", + "hll_sketch_agg", + "hll_union_agg", + ) + measure_column = ColumnMetadata( + name="customer_id_hll_23002251", + type="binary", + ) + + v3_entry = _combiner([measure], [measure_column]).metrics_spec()[0] + cubes_api_entry = _build_metrics_spec([measure_column], [measure], {})[0] + + assert v3_entry["lgK"] == cubes_api_entry["lgK"] == 12 + assert v3_entry["tgtHllType"] == cubes_api_entry["tgtHllType"] == "HLL_4" + + def test_builders_disagree_on_unmappable_measures(self): + """ + An unmappable measure is dropped by one builder and defaulted by the other. + """ + measure = _measure("p95_latency", "approx_percentile") + measure_column = ColumnMetadata(name="p95_latency", type="double") + + v3_spec = _combiner([measure], [measure_column]).metrics_spec() + cubes_api_spec = _build_metrics_spec([measure_column], [measure], {}) + + assert v3_spec == [] + assert cubes_api_spec[0]["type"] == "longSum" diff --git a/datajunction-server/tests/models/reaggregate_test.py b/datajunction-server/tests/models/reaggregate_test.py index 93e6f552e..8f5701481 100644 --- a/datajunction-server/tests/models/reaggregate_test.py +++ b/datajunction-server/tests/models/reaggregate_test.py @@ -163,10 +163,6 @@ def test_params_accepted_for_parameterized_function(monkeypatch): def test_only_sketch_families_are_parameterized(): """ Only a sketch family takes tuning parameters. - - The rollup functions are fully specified by their name -- there is nothing - to tune about a SUM. A sketch is not: a t-digest still needs a compression, - which is why `params` is accepted for it and rejected everywhere else. """ parameterized = { fn for fn in ReaggregationFunction if is_parameterized_reaggregate_function(fn) diff --git a/datajunction-server/tests/sql/decompose_test.py b/datajunction-server/tests/sql/decompose_test.py index 66a3a4106..9bd07a48d 100644 --- a/datajunction-server/tests/sql/decompose_test.py +++ b/datajunction-server/tests/sql/decompose_test.py @@ -17,6 +17,7 @@ MetricComponent, ) from datajunction_server.models.engine import Dialect +from datajunction_server.models.materialization import MaterializationTarget from datajunction_server.models.node_type import NodeType from datajunction_server.models.reaggregate import ( DimensionReaggregateRule, @@ -2909,15 +2910,8 @@ def test_params_without_an_aggregating_component_is_rejected(self): class _DialectAwareSum(AggDecomposition): """ - A SUM whose combiner renders differently per dialect, and which records the - call it was handed. - - Stands in for a quantile sketch, the first family that actually needs this. - Druid exposes no scalar extractor for ``COMPLEX``, so a - t-digest's merge and combine fuse into one ``TDIGEST_QUANTILE(col, fraction)`` - call while Spark and Trino keep them separate -- a difference a rename table - cannot express. The fraction comes from the authored call, which is why - ``func`` is passed alongside the dialect. + A SUM whose combiner renders differently per dialect, recording the call it + was handed. Tests the pipeline for dialect-specific extraction logic. """ seen: list[tuple["ast.Function", Dialect]] = [] @@ -2928,8 +2922,18 @@ def components(self) -> list[ComponentDef]: def combine(self, components, func, dialect=Dialect.SPARK): type(self).seen.append((func, dialect)) - name = "druid_combiner" if dialect == Dialect.DRUID else "spark_combiner" - return make_func(name, components[0].name) + + combiners = { + Dialect.DRUID: self._combine_druid, + Dialect.SPARK: self._combine_spark, + } + return combiners.get(dialect, self._combine_spark)(components) + + def _combine_druid(self, components): + return make_func("druid_combiner", components[0].name) + + def _combine_spark(self, components): + return make_func("spark_combiner", components[0].name) @pytest.fixture @@ -2969,10 +2973,6 @@ def test_combine_renders_per_dialect(): def test_combine_defaults_to_spark(): """ Omitting the dialect yields the Spark rendering. - - The default is what keeps the seven non-build_v3 extractor call sites - working untouched; they render for display and for frozen measures, where - Spark is the right answer. """ components = [ MetricComponent( @@ -2991,11 +2991,6 @@ def test_combine_defaults_to_spark(): def test_combine_receives_the_originating_call(): """ The combiner is handed the authored ``ast.Function``, arguments included. - - This is what makes the quantile fraction reachable: the ``0.5`` in - ``APPROX_PERCENTILE(col, 0.5)`` lives only on the call, and before this the - combiner saw nothing but the merged components. Asserted on the arguments - rather than on identity, since that is the part a combiner needs. """ components = [ MetricComponent( @@ -3032,9 +3027,6 @@ async def test_extractor_threads_dialect_into_combiner( ): """ The dialect handed to the extractor reaches the combiner. - - End-to-end over the real extraction path rather than a direct ``combine`` - call, because the plumbing between them is the part that was missing. """ metric_rev = await create_metric("SELECT SUM(price) FROM parent_node") @@ -3056,10 +3048,6 @@ async def test_extractor_dialect_does_not_leak_between_instances( ): """ Dialect is per-extractor configuration, not shared state. - - ``get_decomposition`` builds a fresh decomposition per call and the dialect - rides the extractor, so a Druid extraction must not change what a later - default extraction renders. """ metric_rev = await create_metric("SELECT SUM(price) FROM parent_node") @@ -3072,6 +3060,66 @@ async def test_extractor_dialect_does_not_leak_between_instances( assert "druid_combiner" not in str(after) +class _SerializingSum(AggDecomposition): + """A SUM that declares a Druid-only serialize conversion.""" + + @property + def components(self) -> list[ComponentDef]: + return [ + ComponentDef( + suffix="_sum", + accumulate="SUM", + merge="SUM", + serialize="to_druid_bytes({})", + serialize_targets=(MaterializationTarget.DRUID,), + ), + ] + + def combine(self, components, func, dialect=Dialect.SPARK): + return make_func("SUM", components[0].name) + + +@pytest.fixture +def serializing_sum(): + """Swap SUM's decomposition for one declaring a serialize, then restore.""" + original = DECOMPOSITION_REGISTRY.get(dj_functions.Sum) + DECOMPOSITION_REGISTRY[dj_functions.Sum] = _SerializingSum + try: + yield + finally: + DECOMPOSITION_REGISTRY[dj_functions.Sum] = original + + +@pytest.mark.asyncio +async def test_serialize_propagates_from_component_def( + session: AsyncSession, + create_metric, + serializing_sum, +): + """ + A ``serialize`` declared on ``ComponentDef`` reaches the extracted component. + """ + metric_rev = await create_metric("SELECT SUM(price) FROM parent_node") + + components, _ = await MetricComponentExtractor(metric_rev.id).extract(session) + + assert len(components) == 1 + assert components[0].serialize == "to_druid_bytes({})" + assert components[0].serialize_targets == [MaterializationTarget.DRUID] + + +@pytest.mark.asyncio +async def test_components_without_serialize_declare_none( + session: AsyncSession, + create_metric, +): + """An ordinary decomposition leaves the conversion unset.""" + metric_rev = await create_metric("SELECT SUM(price) FROM parent_node") + + components, _ = await MetricComponentExtractor(metric_rev.id).extract(session) + + assert components[0].serialize is None + assert components[0].serialize_targets == [] # ============================================================================= @@ -3127,9 +3175,6 @@ def test_declared_family_overrides_the_function_registry(registered_family): def test_family_receives_its_tuning_params(registered_family): """ `reaggregate.params` reaches the decomposition, and thus the accumulate. - - Without this the declared compression would be recorded for reuse identity - and Druid ingestion but silently ignored by the SQL that builds the sketch. """ spec = ReaggregateSpec( fn=ReaggregationFunction.TDIGEST, @@ -3145,9 +3190,6 @@ def test_family_receives_its_tuning_params(registered_family): def test_undeclared_metric_is_untouched_by_a_registered_family(registered_family): """ A metric without `reaggregate` keeps the by-function decomposition. - - Registering a family must not change metrics that never opted in -- that - separation is the whole point of gating on the spec. """ assert get_decomposition(dj_functions.ApproxPercentile) is None assert isinstance( @@ -3160,9 +3202,6 @@ def test_undeclared_metric_is_untouched_by_a_registered_family(registered_family def test_unregistered_family_falls_through(): """ Declaring a family nothing implements degrades rather than raising. - - OSS has no implementations, so `fn: tdigest` there simply leaves the metric - with the aggregability it already had. """ spec = ReaggregateSpec(fn=ReaggregationFunction.TDIGEST) From 0071c495d38f2458da4e2f23ef614ae247569c06 Mon Sep 17 00:00:00 2001 From: Robin Davis Date: Tue, 22 Sep 2026 11:21:35 -0700 Subject: [PATCH 03/17] Restore docstrings flagged in review --- .../datajunction_server/config.py | 36 +++++++++++- .../datajunction_server/models/decompose.py | 58 ++++++++++++++++++- 2 files changed, 91 insertions(+), 3 deletions(-) diff --git a/datajunction-server/datajunction_server/config.py b/datajunction-server/datajunction_server/config.py index 31ee27a8e..39941bca4 100644 --- a/datajunction-server/datajunction_server/config.py +++ b/datajunction-server/datajunction_server/config.py @@ -48,7 +48,41 @@ class DatabaseConfig(BaseModel): class QueryClientConfig(BaseModel): - """Configuration for query service clients.""" + """ + Configuration for query service clients. + + Set via environment variables using double-underscore delimiters, e.g.:: + + QUERY_CLIENT__TYPE=bigquery + QUERY_CLIENT__CONNECTION__PROJECT=my-gcp-project + + Supported client types and their required connection parameters: + + **http** (default):: + + QUERY_CLIENT__TYPE=http + QUERY_CLIENT__CONNECTION__URI=http://djqs:8001 + + **snowflake**:: + + QUERY_CLIENT__TYPE=snowflake + QUERY_CLIENT__CONNECTION__ACCOUNT=my-account + QUERY_CLIENT__CONNECTION__USER=my-user + QUERY_CLIENT__CONNECTION__PASSWORD=my-password + + **bigquery**:: + + QUERY_CLIENT__TYPE=bigquery + QUERY_CLIENT__CONNECTION__PROJECT=my-gcp-project + # Optional: path to a service account JSON key file + QUERY_CLIENT__CONNECTION__CREDENTIALS_PATH=/path/to/service-account.json + # Optional: BigQuery location (e.g. US, EU) + QUERY_CLIENT__CONNECTION__LOCATION=US + + When multiple DJ catalogs map to different GCP projects, set the engine URI + on the DJ engine to ``bigquery://my-gcp-project`` — the project is resolved + from the engine URI first, falling back to ``PROJECT`` above. + """ # Type of query client: 'http', 'snowflake', 'bigquery', 'databricks', 'trino', etc. type: str = "http" diff --git a/datajunction-server/datajunction_server/models/decompose.py b/datajunction-server/datajunction_server/models/decompose.py index bc7db2658..36db3ec7c 100644 --- a/datajunction-server/datajunction_server/models/decompose.py +++ b/datajunction-server/datajunction_server/models/decompose.py @@ -65,7 +65,38 @@ class AggregationRule(BaseModel): class MetricComponent(BaseModel): - """A reusable, named building block of a metric definition.""" + """ + A reusable, named building block of a metric definition. + + A MetricComponent represents a SQL expression that can serve as an input + to building a metric. It supports a two-phase aggregation model: + + - Phase 1 (Accumulate): Build from raw data using `aggregation` + Can be a function name ("SUM") or a template ("SUM(POWER({}, 2))") + + - Phase 2 (Merge): Combine pre-aggregated values using `merge` function + Examples: SUM, SUM (for COUNT), hll_union_agg + + For most aggregations, accumulate and merge use the same function (SUM → SUM). + For COUNT, merge is SUM (sum up the counts). + For HLL sketches, they differ: hll_sketch_estimate vs hll_union_agg. + + The final expression combining merged components is specified in + DecomposedMetric.combiner. + + Attributes: + name: A unique name for the component, derived from its expression. + expression: The raw SQL expression (column/value) being aggregated. + aggregation: Function name or template for Phase 1. Simple cases use + just the name ("SUM"), complex cases use templates with + {} placeholder ("SUM(POWER({}, 2))"). + merge: The function name for combining pre-aggregated values (Phase 2). + rule: Aggregation rules defining how/when the component can be aggregated. + + The sketch-backed fields added alongside these -- params, merge_args, + serialize, serialize_targets, serialize_type -- are documented inline + below, where the reasoning sits next to the declaration. + """ name: str expression: str @@ -148,7 +179,30 @@ class PreAggMeasure(MetricComponent): class DecomposedMetric(BaseModel): - """A metric decomposed into its constituent components with a combining expression.""" + """ + A metric decomposed into its constituent components with a combining expression. + + This is the result of decomposing a metric query. It specifies: + - components: The measures needed for pre-aggregation + - combiner: How to combine merged components into the final metric value + - derived_query: The full SQL query using the combiner + + Examples: + SUM metric: + components: [{name: "revenue_sum", aggregation: "SUM", merge: "SUM"}] + combiner: "SUM(revenue_sum)" + + AVG metric: + components: [ + {name: "revenue_sum", aggregation: "SUM", merge: "SUM"}, + {name: "revenue_count", aggregation: "COUNT", merge: "SUM"} + ] + combiner: "SUM(revenue_sum) / SUM(revenue_count)" + + APPROX_COUNT_DISTINCT metric (uses Spark function names): + components: [{name: "user_hll", aggregation: "hll_sketch_agg", merge: "hll_union"}] + combiner: "hll_sketch_estimate(hll_union(user_hll))" + """ components: list[MetricComponent] combiner: str # Expression combining merged components into final value From 9872ffcd1a90c555a9f5a3a871e6640a9f17835f Mon Sep 17 00:00:00 2001 From: Robin Davis Date: Tue, 22 Sep 2026 11:46:51 -0700 Subject: [PATCH 04/17] Exercise a registered sketch family through the SQL endpoints --- .../datajunction_server/sql/decompose.py | 11 +- .../build_v3/sketch_family_e2e_test.py | 429 ++++++++++++++++++ 2 files changed, 435 insertions(+), 5 deletions(-) create mode 100644 datajunction-server/tests/construction/build_v3/sketch_family_e2e_test.py diff --git a/datajunction-server/datajunction_server/sql/decompose.py b/datajunction-server/datajunction_server/sql/decompose.py index 70868a83c..2173f1b34 100644 --- a/datajunction-server/datajunction_server/sql/decompose.py +++ b/datajunction-server/datajunction_server/sql/decompose.py @@ -1119,11 +1119,12 @@ def __init__( Args: node_revision_id: ID of the metric node revision - dialect: Dialect the combiner will be rendered for. Combiners are - mostly dialect-neutral, but a sketch family whose engines expose - different function shapes needs the target -- Druid fuses a - t-digest's merge and combine, where Spark and Trino keep them - separate. + dialect: Dialect the combiner will be rendered for. Every dialect + gets a combiner -- whether one exists at all is a property of + the aggregation function, not the engine -- but its shape can + differ: a sketch family whose engines expose different function + shapes needs the target, as Druid fuses a t-digest's merge and + combine where Spark and Trino keep them separate. Defaults to Spark, which is right for the callers that render for display or for frozen measures. Only the build_v3 path, diff --git a/datajunction-server/tests/construction/build_v3/sketch_family_e2e_test.py b/datajunction-server/tests/construction/build_v3/sketch_family_e2e_test.py new file mode 100644 index 000000000..95e7ca8f4 --- /dev/null +++ b/datajunction-server/tests/construction/build_v3/sketch_family_e2e_test.py @@ -0,0 +1,429 @@ +""" +A sketch family, end to end, through the HTTP endpoints. + +The rest of the suite tests this machinery in pieces -- `merge_args_test` on +the merge call, `accumulate_type_test` on inferred types, `serialize_target_test` +on the write-time conversion, `reaggregate_test` on the spec. Each one holds a +`ComponentDef` or a `MetricComponent` in its hand and asserts about it. None of +them puts a family behind the server and asks what SQL comes out. + +This does. One metric, one grain, followed from raw rows to the number the +engine returns: + + accumulate build_digest(line_total, 200.0) from source, per partition + merge merge_digest(digest, 200.0) from the pre-agg table + combine digest_quantiles(digest, ...)[0] Spark / Trino + DIGEST_QUANTILE(digest, 0.95) Druid, merge fused in + +OSS registers no family of its own -- they live downstream with the +engine-specific functions -- so the fixture below supplies one, the same way +`decompose_test.registered_family` does for the unit tests. That makes this +file the executable form of the extension contract: a downstream deployment +has to register *both* a decomposition and the functions it names, and if +either half is missing the failure is a 500 out of the measures endpoint, not +something a type checker catches. + +`serialize` is the one hook not asserted here. It applies when writing a +measures table for a target rather than when reading, so no query endpoint +renders it; `serialize_target_test` and `cube_druid_sketch_spec_test` cover it. +""" + +import re + +import pytest +from httpx import AsyncClient + +from datajunction_server.models.dialect import Dialect +from datajunction_server.models.materialization import MaterializationTarget +from datajunction_server.models.reaggregate import ReaggregationFunction +from datajunction_server.sql.decompose import ( + FAMILY_DECOMPOSITION_REGISTRY, + AggDecomposition, + ComponentDef, + decomposes_family, + make_func, +) +from datajunction_server.sql.functions import Function, function_registry +from datajunction_server.sql.parsing import ast +from datajunction_server.sql.parsing import types as ct + +from . import assert_sql_equal + +METRIC = "v3.p95_line_total" +DIMENSIONS = ["v3.product.category"] +COMPRESSION = 200 +PREAGG_TABLE = "analytics.preaggs.v3_digest_by_cat" + + +# --------------------------------------------------------------------------- +# A stand-in sketch family: four functions and a decomposition that uses them. +# Shaped after a real t-digest -- an accumulate that takes a tuning argument, a +# merge that must repeat it, and a combiner whose form differs on Druid -- so +# the parts of the machinery that only a sketch exercises are actually hit. +# --------------------------------------------------------------------------- + + +class _DigestFunction(Function): + """Base class, so the fixture can find these to register and unregister.""" + + +class BuildDigest(_DigestFunction): + """build_digest(col, compression) -> digest. Phase 1.""" + + is_aggregation = True + dialects = [Dialect.SPARK, Dialect.TRINO] + + +class MergeDigest(_DigestFunction): + """merge_digest(digest, compression) -> digest. Phase 2.""" + + is_aggregation = True + dialects = [Dialect.SPARK, Dialect.TRINO] + + +class DigestQuantiles(_DigestFunction): + """digest_quantiles(digest, quantiles) -> array. Phase 3, non-Druid.""" + + is_aggregation = False + dialects = [Dialect.SPARK, Dialect.TRINO] + + +class DigestQuantile(_DigestFunction): + """DIGEST_QUANTILE(digest, q) -> double. Phase 2+3 fused, Druid only.""" + + is_aggregation = True + dialects = [Dialect.DRUID] + + +@BuildDigest.register +def infer_type(col: ct.ColumnType, compression: ct.ColumnType) -> ct.BinaryType: + return ct.BinaryType() + + +@MergeDigest.register +def infer_type(digest: ct.ColumnType, compression: ct.ColumnType) -> ct.BinaryType: + return ct.BinaryType() + + +@DigestQuantiles.register +def infer_type(digest: ct.ColumnType, quantiles: ct.ColumnType) -> ct.ListType: + return ct.ListType(element_type=ct.DoubleType()) + + +@DigestQuantile.register +def infer_type(digest: ct.ColumnType, quantile: ct.ColumnType) -> ct.DoubleType: + return ct.DoubleType() + + +@pytest.fixture +def digest_family(): + """ + Register the family and its functions, then put the registries back. + + Both registries are process-global, so leaving either populated would leak + into every later test in the session. + """ + saved: dict[str, type | None] = {} + for cls in _DigestFunction.__subclasses__(): + snake_cased = re.sub(r"(? int: + return int(self.params.get("compression", 100)) + + @property + def compression_literal(self) -> str: + return f"CAST({self.compression}.0 AS DOUBLE)" + + @property + def components(self) -> list[ComponentDef]: + return [ + ComponentDef( + # Compression rides in the suffix so two differently tuned + # sketches over the same column get distinct columns. + suffix=f"_digest_c{self.compression}", + accumulate=f"build_digest({{}}, {self.compression_literal})", + merge="merge_digest", + merge_args=(self.compression_literal,), + serialize="digest_to_bytes({})", + serialize_targets=(MaterializationTarget.DRUID,), + serialize_type="binary", + ), + ] + + def combine(self, components, func, dialect=Dialect.SPARK): + column, fraction = components[0].name, func.args[1] + if dialect == Dialect.DRUID: + return make_func("DIGEST_QUANTILE", column, fraction) + return ast.Subscript( + expr=make_func( + "digest_quantiles", + column, + make_func("array", fraction), + ), + index=ast.Number(0), + ) + + try: + yield _DigestDecomposition + finally: + FAMILY_DECOMPOSITION_REGISTRY.pop(ReaggregationFunction.TDIGEST, None) + for key, previous in saved.items(): + if previous is None: + function_registry.pop(key, None) + else: + function_registry[key] = previous + + +@pytest.fixture +async def percentile_metric(client_with_build_v3: AsyncClient, digest_family): + """ + A percentile metric that opts into the family. + + The aggregation is a plain APPROX_PERCENTILE -- nothing about the query + mentions a sketch. Only `reaggregate.fn` selects the family, which is the + point of the family registry: the metric keeps its readable definition. + """ + response = await client_with_build_v3.post( + "/nodes/metric/", + json={ + "name": METRIC, + "description": "95th percentile line total", + "query": ( + "SELECT APPROX_PERCENTILE(line_total, 0.95) FROM v3.order_details" + ), + "mode": "published", + "reaggregate": { + "fn": "tdigest", + "params": {"compression": COMPRESSION}, + }, + }, + ) + assert response.status_code in (200, 201), response.json() + return client_with_build_v3 + + +async def _measures(client: AsyncClient, dialect: str = "spark") -> dict: + response = await client.get( + "/sql/measures/v3/", + params={"metrics": [METRIC], "dimensions": DIMENSIONS, "dialect": dialect}, + ) + assert response.status_code == 200, response.json() + return response.json() + + +async def _publish_preagg(client: AsyncClient) -> None: + response = await client.post( + "/preaggs/plan", + json={ + "metrics": [METRIC], + "dimensions": DIMENSIONS, + "strategy": "full", + "schedule": "0 0 * * *", + }, + ) + assert response.status_code == 201, response.text + + catalog, schema, table = PREAGG_TABLE.split(".") + response = await client.post( + f"/preaggs/{response.json()['preaggs'][0]['id']}/availability/", + json={ + "catalog": catalog, + "schema": schema, + "table": table, + "valid_through_ts": 1704067200, + }, + ) + assert response.status_code == 200, response.text + + +@pytest.mark.asyncio +async def test_accumulate_renders_the_tuning_param(percentile_metric): + """ + Phase 1: `params` reaches the generated SQL as a real argument. + + `compression` is declared once on the metric's reaggregate spec and has to + survive into the accumulate call. If it were dropped the query would still + run -- `build_digest` has a one-argument form -- and quietly build sketches + at the wrong accuracy, so this is worth pinning on the SQL rather than on + the spec. + """ + (group,) = (await _measures(percentile_metric))["grain_groups"] + + assert_sql_equal( + group["sql"], + f""" + WITH v3_order_details AS ( + SELECT 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 + ) + SELECT t2.category, + build_digest(t1.line_total, CAST({COMPRESSION}.0 AS DOUBLE)) + line_total_digest_c{COMPRESSION}_e1f61696 + FROM v3_order_details t1 + LEFT OUTER JOIN v3_product t2 ON t1.product_id = t2.product_id + GROUP BY t2.category + """, + ) + + +@pytest.mark.asyncio +async def test_component_identity_carries_the_compression(percentile_metric): + """ + The column name encodes the tuning, and the recorded type is the sketch's. + + Both matter downstream: the name keeps a sketch built at one accuracy from + satisfying a query that asked for another, and the type is what the + materialized table is created from. + """ + (group,) = (await _measures(percentile_metric))["grain_groups"] + (component,) = group["components"] + (column,) = [c for c in group["columns"] if c["semantic_type"] != "dimension"] + + assert component["name"] == f"line_total_digest_c{COMPRESSION}_e1f61696" + assert component["merge"] == "merge_digest" + assert component["aggregability"] == "full" + assert column["type"] == "binary" + + +@pytest.mark.asyncio +async def test_merge_from_preagg_repeats_the_tuning_param(percentile_metric): + """ + Phase 2: reading stored sketches back, with `merge_args` applied. + + The accumulate is gone and the source CTEs with it -- a single scan of the + pre-agg table. The compression has to appear a second time here: `merge` is + a bare function name (the Druid aggregator lookup and the semi-additive + rewrite both match on it), so the argument cannot be folded into it and + travels as `merge_args` instead. + """ + await _publish_preagg(percentile_metric) + (group,) = (await _measures(percentile_metric))["grain_groups"] + + assert_sql_equal( + group["sql"], + f""" + SELECT category, + merge_digest(line_total_digest_c{COMPRESSION}_e1f61696, + CAST({COMPRESSION}.0 AS DOUBLE)) + line_total_digest_c{COMPRESSION}_e1f61696 + FROM {PREAGG_TABLE} + GROUP BY category + """, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "dialect, expected", + [ + ( + "spark", + f"digest_quantiles(line_total_digest_c{COMPRESSION}_e1f61696, " + "array(0.95))[0]", + ), + ( + "druid", + f"DIGEST_QUANTILE(line_total_digest_c{COMPRESSION}_e1f61696, 0.95)", + ), + ], +) +async def test_combiner_takes_the_shape_the_engine_needs( + percentile_metric, + dialect, + expected, +): + """ + Phase 3: the combiner, and the one place the dialect legitimately matters. + + Druid has no scalar reader for the sketch, so it fuses merge and combine + into one aggregating call; Spark and Trino read the merged column with a + separate extractor. That is a difference in shape, which no transpiler can + derive -- unlike function spelling or array index base, which the + transpilation layer handles and a decomposition must therefore leave alone. + + The requested quantile is authored in the metric query, not the reaggregate + spec, so it has to survive decomposition and reappear here. That is what + lets p50/p95/p99 share one sketch column. + """ + payload = await _measures(percentile_metric, dialect=dialect) + (formula,) = payload["metric_formulas"] + + assert formula["combiner"] == expected + + +@pytest.mark.asyncio +async def test_metrics_layer_recomputes_from_source(percentile_metric): + """ + Final metrics layer: recomputes from source, pre-agg or not. + + Pinning current behaviour, not endorsing it. The pre-agg published here is + the same one `test_merge_from_preagg_repeats_the_tuning_param` reads, and + the measures layer picks it up. The metrics layer never consults the + matcher: `build_metrics_sql` takes the cube branch only when + `use_materialized and dialect == Dialect.DRUID`, and otherwise builds grain + groups from source. + + So no sketch appears below at all -- a bare APPROX_PERCENTILE off the fact + -- and the combiner asserted above never reaches rendered SQL on this path. + Defensible if scanning source on Trino is cheap enough, but it does mean + the sketch columns are written for Druid's benefit alone. + """ + await _publish_preagg(percentile_metric) + + response = await percentile_metric.get( + "/sql/", + params={"metrics": [METRIC], "dimensions": DIMENSIONS}, + ) + assert response.status_code == 200, response.json() + + assert_sql_equal( + response.json()["sql"], + """ + WITH v3_DOT_order_details AS ( + SELECT o.order_id, + oi.line_number, + o.customer_id, + o.order_date, + o.from_location_id, + o.to_location_id, + o.status, + oi.product_id, + oi.quantity, + oi.unit_price, + oi.quantity * oi.unit_price AS line_total + FROM v3.orders AS o + JOIN v3.order_items AS oi ON o.order_id = oi.order_id + ), + v3_DOT_product AS ( + SELECT v3_DOT_src_products.product_id, + v3_DOT_src_products.name, + v3_DOT_src_products.category, + v3_DOT_src_products.subcategory, + v3_DOT_src_products.price + FROM v3.products AS v3_DOT_src_products + ), + v3_DOT_order_details_metrics AS ( + SELECT v3_DOT_product.category v3_DOT_product_DOT_category, + APPROX_PERCENTILE(v3_DOT_order_details.line_total, 0.95) + v3_DOT_p95_line_total + FROM v3_DOT_order_details + LEFT JOIN v3_DOT_product + ON v3_DOT_order_details.product_id = v3_DOT_product.product_id + GROUP BY v3_DOT_product.category + ) + SELECT v3_DOT_order_details_metrics.v3_DOT_product_DOT_category, + v3_DOT_order_details_metrics.v3_DOT_p95_line_total + FROM v3_DOT_order_details_metrics + """, + ) From bbce3fa4b8e272577bf4ea000c180c8baf273631 Mon Sep 17 00:00:00 2001 From: Beto Dealmeida Date: Thu, 24 Sep 2026 13:29:29 -0300 Subject: [PATCH 05/17] Fix sketch decomposition rebase compatibility --- .../construction/build_v3/decomposition.py | 7 +++++-- .../datajunction_server/construction/build_v3/measures.py | 6 +++++- .../datajunction_server/internal/materializations.py | 2 +- .../datajunction_server/models/deployment.py | 2 +- .../datajunction_server/models/reaggregate.py | 1 + .../tests/construction/build_v3/serialize_target_test.py | 2 +- .../tests/models/cube_druid_sketch_spec_test.py | 2 +- datajunction-server/tests/sql/decompose_test.py | 6 +++--- 8 files changed, 18 insertions(+), 10 deletions(-) diff --git a/datajunction-server/datajunction_server/construction/build_v3/decomposition.py b/datajunction-server/datajunction_server/construction/build_v3/decomposition.py index 960625e90..e78caf645 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/decomposition.py +++ b/datajunction-server/datajunction_server/construction/build_v3/decomposition.py @@ -287,8 +287,11 @@ def _apply_serialize( """ if not component.serializes_for(materialization_target): return expr + serialize = component.serialize + if serialize is None: # pragma: no cover - guaranteed by serializes_for + return expr wrapped = parse( - f"SELECT {component.serialize.replace('{}', str(expr))}", + f"SELECT {serialize.replace('{}', str(expr))}", ).select.projection[0] wrapped.clear_parent() return cast(ast.Expression, wrapped) @@ -329,7 +332,7 @@ def collect_bases(node: Node): visited.add(node.name) # Recurse through metric parents. A node with any metric parent is derived; - # a node with no metric parent is a base metric. Checking for metric parents + # a node with no metric parent is a base metric. Checking for metric parents # avoids dropping base metrics that are defined directly on dimension nodes. has_metric_parent = False for parent_name in ctx.parent_map.get(node.name, []): diff --git a/datajunction-server/datajunction_server/construction/build_v3/measures.py b/datajunction-server/datajunction_server/construction/build_v3/measures.py index 0607e69f6..5f24d6365 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/measures.py +++ b/datajunction-server/datajunction_server/construction/build_v3/measures.py @@ -185,6 +185,7 @@ def _multi_argument_accumulate_types( arg_types: list[ct.ColumnType] = [] for arg in call.args: + resolved: ct.ColumnType | None if isinstance(arg, ast.Column): # Only the columns need the parent to resolve; literals and casts # carry their own type without being bound to a table. @@ -192,7 +193,10 @@ def _multi_argument_accumulate_types( get_column_type(parent_node, str(arg.alias_or_name.name)), ) else: - resolved = arg.type + inferred = arg.type + if isinstance(inferred, list): + return None + resolved = inferred if resolved is None: return None arg_types.append(resolved) diff --git a/datajunction-server/datajunction_server/internal/materializations.py b/datajunction-server/datajunction_server/internal/materializations.py index 3f2fab8c6..be7816e6b 100644 --- a/datajunction-server/datajunction_server/internal/materializations.py +++ b/datajunction-server/datajunction_server/internal/materializations.py @@ -552,7 +552,7 @@ async def swap_cube_materializations( ) -> CubeMaterializationSwap | None: """ Rebuild a cube's materializations against a new revision and retire the old ones. - + Touches only DJ-side state, and returns the query service work still owed -- `None` when the cube had nothing materialized and there is no work at all. The caller commits and then hands the result to `apply_cube_materialization_swap`, so diff --git a/datajunction-server/datajunction_server/models/deployment.py b/datajunction-server/datajunction_server/models/deployment.py index 1e7daa9a2..e49fc6c0b 100644 --- a/datajunction-server/datajunction_server/models/deployment.py +++ b/datajunction-server/datajunction_server/models/deployment.py @@ -219,7 +219,7 @@ class PreAggSpec(NamespacedSpec): """ Specification for an externally-built pre-aggregation table adopted at deploy time. ``name`` is a stable handle used for reconciliation. - + Every metric and dimension is declared together with its physical column in the external table. """ diff --git a/datajunction-server/datajunction_server/models/reaggregate.py b/datajunction-server/datajunction_server/models/reaggregate.py index 14ad0b759..2b6b2cb4f 100644 --- a/datajunction-server/datajunction_server/models/reaggregate.py +++ b/datajunction-server/datajunction_server/models/reaggregate.py @@ -108,6 +108,7 @@ class ReaggregateSpec(BaseModel): # materialization adapter validates the supported keys for its aggregator. params: dict[str, Any] | None = None + def dump_reaggregate_spec( spec: ReaggregateSpec | dict | None, ) -> dict | None: diff --git a/datajunction-server/tests/construction/build_v3/serialize_target_test.py b/datajunction-server/tests/construction/build_v3/serialize_target_test.py index ef5bbcf76..88caa5e61 100644 --- a/datajunction-server/tests/construction/build_v3/serialize_target_test.py +++ b/datajunction-server/tests/construction/build_v3/serialize_target_test.py @@ -128,7 +128,7 @@ def test_targets_default_to_empty(self): """ A component declaring a conversion but no targets never applies it. - Defaulting to "no targets" rather than "all targets" prevents + Defaulting to "no targets" rather than "all targets" prevents half-finished declarations from silently rewriting measures tables. """ component = _component( diff --git a/datajunction-server/tests/models/cube_druid_sketch_spec_test.py b/datajunction-server/tests/models/cube_druid_sketch_spec_test.py index 09a3bb067..0a26916fc 100644 --- a/datajunction-server/tests/models/cube_druid_sketch_spec_test.py +++ b/datajunction-server/tests/models/cube_druid_sketch_spec_test.py @@ -141,7 +141,7 @@ class TestMetricsSpecBuildersAgree: The two Druid metricsSpec builders, compared directly. ``CombineMaterialization.metrics_spec`` and ``api.cubes._build_metrics_spec`` - map measures for different entry points. They share ``get_druid_aggregator_spec`` + map measures for different entry points. They share ``get_druid_aggregator_spec`` so the aggregator type and family config agree. """ diff --git a/datajunction-server/tests/sql/decompose_test.py b/datajunction-server/tests/sql/decompose_test.py index 9bd07a48d..a9634671a 100644 --- a/datajunction-server/tests/sql/decompose_test.py +++ b/datajunction-server/tests/sql/decompose_test.py @@ -2922,16 +2922,16 @@ def components(self) -> list[ComponentDef]: def combine(self, components, func, dialect=Dialect.SPARK): type(self).seen.append((func, dialect)) - + combiners = { Dialect.DRUID: self._combine_druid, Dialect.SPARK: self._combine_spark, } return combiners.get(dialect, self._combine_spark)(components) - + def _combine_druid(self, components): return make_func("druid_combiner", components[0].name) - + def _combine_spark(self, components): return make_func("spark_combiner", components[0].name) From cbf947be559d3ed88f7fb8d6f8b91c24c237a160 Mon Sep 17 00:00:00 2001 From: Beto Dealmeida Date: Thu, 24 Sep 2026 14:18:28 -0300 Subject: [PATCH 06/17] Propagate materialization target into build context --- .../construction/build_v3/builder.py | 1 + .../construction/build_v3/build_v3_test.py | 17 +++++++++++++++++ 2 files changed, 18 insertions(+) diff --git a/datajunction-server/datajunction_server/construction/build_v3/builder.py b/datajunction-server/datajunction_server/construction/build_v3/builder.py index 9e34d7a24..79f45f7d8 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/builder.py +++ b/datajunction-server/datajunction_server/construction/build_v3/builder.py @@ -293,6 +293,7 @@ async def setup_build_context( dimensions=list(dimensions), filters=filters or [], dialect=dialect, + materialization_target=materialization_target, use_materialized=use_materialized, temporal_partition_columns=temporal_partition_columns or {}, lookback_window=lookback_window, diff --git a/datajunction-server/tests/construction/build_v3/build_v3_test.py b/datajunction-server/tests/construction/build_v3/build_v3_test.py index 5366d697a..1856b314d 100644 --- a/datajunction-server/tests/construction/build_v3/build_v3_test.py +++ b/datajunction-server/tests/construction/build_v3/build_v3_test.py @@ -14,6 +14,7 @@ setup_build_context, ) from datajunction_server.errors import DJInvalidInputException +from datajunction_server.models.materialization import MaterializationTarget from datajunction_server.sql.parsing import ast from . import assert_sql_equal, get_first_grain_group @@ -56,6 +57,22 @@ def test_unsupported_type_raises(self): _normalize_query_param_value("p", [1, 2, 3]) +@pytest.mark.asyncio +async def test_setup_build_context_preserves_materialization_target( + client_with_build_v3, + session, +): + """The destination reaches the context used to render and type measures.""" + ctx = await setup_build_context( + session=session, + metrics=["v3.total_revenue"], + dimensions=[], + materialization_target=MaterializationTarget.DRUID, + ) + + assert ctx.materialization_target == MaterializationTarget.DRUID + + class TestInnerCTEFlattening: """ Tests for flattening inner CTEs within transforms. From 39b6c389ab57383cb12666247088c82896d79a2f Mon Sep 17 00:00:00 2001 From: Beto Dealmeida Date: Thu, 24 Sep 2026 14:23:58 -0300 Subject: [PATCH 07/17] Restore documentation lost during rebase --- .../python/datajunction/client.py | 26 ++++++- .../python/datajunction/deployment.py | 25 +++++- .../datajunction_server/api/sql.py | 65 +++++++++++++++- .../construction/build_v3/decomposition.py | 21 +++-- .../database/preaggregation.py | 59 +++++++++++++- .../internal/materializations.py | 64 ++++++++++++++- .../datajunction_server/internal/nodes.py | 20 ++++- .../datajunction_server/models/deployment.py | 77 ++++++++++++++++--- .../datajunction_server/sql/parsing/ast.py | 23 ++++++ .../nodes/derive_frozen_measures_test.py | 15 +++- 10 files changed, 365 insertions(+), 30 deletions(-) diff --git a/datajunction-clients/python/datajunction/client.py b/datajunction-clients/python/datajunction/client.py index f236ca422..ef9867264 100644 --- a/datajunction-clients/python/datajunction/client.py +++ b/datajunction-clients/python/datajunction/client.py @@ -207,7 +207,31 @@ def plan( include_temporal_filters: bool = False, lookback_window: str | None = None, ): - """Returns a query execution plan for the given metrics and dimensions.""" + """ + Returns a query execution plan for the given metrics and dimensions. + + The plan shows: + - grain_groups: How metrics are grouped and their intermediate SQL + - metric_formulas: How each metric combines its components + - requested_dimensions: The dimensions being queried + + This is useful for understanding how DJ decomposes metrics into + atomic aggregations and how multiple fact tables are joined together. + + Args: + metrics: List of metric names to include + dimensions: List of dimensions to group by + filters: List of filter expressions + cube: Optional cube node name. When provided, the cube's stored + filters are automatically prepended to the query filters. + dialect: SQL dialect (e.g., 'spark', 'trino'). Defaults to engine dialect. + use_materialized: Whether to use materialized tables when available + include_temporal_filters: Whether to include temporal partition filters. + Only applies if the metrics and dimensions resolve to a cube with + temporal partitions. + lookback_window: Lookback window for temporal filters (e.g., '3 DAY', + '1 WEEK'). Only applicable when include_temporal_filters is True. + """ params: dict = { "metrics": metrics, "dimensions": dimensions or [], diff --git a/datajunction-clients/python/datajunction/deployment.py b/datajunction-clients/python/datajunction/deployment.py index 0ba52c62a..e8ec4fe54 100644 --- a/datajunction-clients/python/datajunction/deployment.py +++ b/datajunction-clients/python/datajunction/deployment.py @@ -352,7 +352,30 @@ def build_codeowners( default_owner: str | None = None, exclude_dirs: list[str] | None = None, ) -> int: - """Generate a CODEOWNERS file from the owners fields in DJ node YAML files.""" + """ + Generate a CODEOWNERS file from the owners fields in DJ node YAML files. + + Walks base_dir recursively, reads every *.yaml file (skipping dj.yaml), + and maps each file path to its owners list. Files with no owners are + omitted. Paths in the output are relative to base_dir and prefixed with + / so GitHub resolves them from the repo root. + + If github_api_url is provided (and GITHUB_TOKEN / github_token_env is set), + email addresses in owners fields are resolved to GitHub usernames via the + search API. Unresolvable emails are emitted as-is with a warning comment. + + ``default_owner``, when set, emits a leading ``* `` rule so + unmatched files (and any excluded directories) fall through to it. Because + CODEOWNERS is last-match-wins, the per-file rules below it still take + precedence for the files they name. + + ``exclude_dirs`` lists directories (relative to base_dir) whose nodes are + NOT given per-file owners — they fall through to ``default_owner`` instead. + This is for machine-generated trees (e.g. ``nodes/generated``) that have no + individual human owner and should be team-owned as a block. + + Returns the number of per-file entries written (excludes the default rule). + """ base = Path(base_dir).resolve() excluded_dirs = [(base / d).resolve() for d in (exclude_dirs or [])] diff --git a/datajunction-server/datajunction_server/api/sql.py b/datajunction-server/datajunction_server/api/sql.py index 52e462b43..c2216df2a 100644 --- a/datajunction-server/datajunction_server/api/sql.py +++ b/datajunction-server/datajunction_server/api/sql.py @@ -222,7 +222,36 @@ async def get_measures_sql_v3( session: AsyncSession = Depends(get_session), current_user: User = Depends(get_current_user), ) -> MeasuresSQLResponse: - """Generate pre-aggregated measures SQL for the requested metrics.""" + """ + Generate pre-aggregated measures SQL for the requested metrics. + + Measures SQL represents the first stage of metric computation - it decomposes + each metric into its atomic aggregation components (e.g., SUM(amount), COUNT(*)) + and produces SQL that computes these components at the requested dimensional grain. + + Metrics are separated into grain groups, which represent sets of metrics that can be + computed together at a common grain. Each grain group produces its own SQL query, which + can be materialized independently to produce intermediate tables that are then queried + to compute final metric values. + + Returns: + One or more `GrainGroupSQL` objects, each containing: + - SQL query computing metric components at the specified grain + - Column metadata with semantic types + - Component details for downstream re-aggregation + + Args: + cube: Optional cube node name. When provided, the cube's stored filters are + automatically prepended to the query filters. + use_materialized: If True (default), use materialized tables when available. + Set to False when generating SQL for materialization refresh to avoid + circular references. + include_temporal_filters: If True, checks if metrics+dimensions resolve to + a cube with temporal partitions, and applies partition filters if so. + lookback_window: Lookback window for temporal filters when applicable. + + See also: `/sql/metrics/v3/` for the final combined query with metric expressions. + """ merged_filters = list(filters) cube_node = None if cube: @@ -416,7 +445,38 @@ async def get_combined_measures_sql_v3( session: AsyncSession = Depends(get_session), current_user: User = Depends(get_current_user), ) -> CombinedMeasuresSQLResponse: - """Generate combined pre-aggregated measures SQL for the requested metrics.""" + """ + Generate combined pre-aggregated measures SQL for the requested metrics. + + This endpoint combines multiple grain groups into a single SQL query using + FULL OUTER JOIN on shared dimensions. Dimension columns are wrapped with + COALESCE to handle NULLs from non-matching rows. + + This is useful for: + - Druid cube materialization where a single combined table is needed + - Simplifying downstream queries that need data from multiple fact tables + - Pre-computing joined aggregations for dashboards + + The combined SQL contains: + - CTEs for each grain group's pre-aggregated data + - FULL OUTER JOIN between grain groups on shared dimensions + - COALESCE on dimension columns to handle NULL values + - All measure columns from all grain groups + + Args: + metrics: List of metric names to include + dimensions: List of dimensions to group by (the grain) + filters: Optional filters to apply + use_preagg_tables: If False (default), compute from scratch using source tables. + If True, read from pre-aggregation tables. + + Returns: + Combined SQL query with column metadata and grain information. + + See also: + - `/sql/measures/v3/` for individual grain group queries + - `/sql/metrics/v3/` for final metric computations with combiner expressions + """ _t0 = time.monotonic() if use_preagg_tables: # Generate SQL that reads from pre-agg tables (deterministic names) @@ -654,6 +714,7 @@ async def get_metrics_sql_v3( """ if not metrics and not cube: raise DJInvalidInputException("At least one metric is required") + # Shared metrics-SQL core (cube pinning, cube_filters prepend, dialect # auto-resolve, build_metrics_sql, and the build-latency metrics + [SQL] log). # Also used by the semantic-layer endpoint. diff --git a/datajunction-server/datajunction_server/construction/build_v3/decomposition.py b/datajunction-server/datajunction_server/construction/build_v3/decomposition.py index e78caf645..4e6773bea 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/decomposition.py +++ b/datajunction-server/datajunction_server/construction/build_v3/decomposition.py @@ -331,9 +331,12 @@ def collect_bases(node: Node): return visited.add(node.name) - # Recurse through metric parents. A node with any metric parent is derived; - # a node with no metric parent is a base metric. Checking for metric parents - # avoids dropping base metrics that are defined directly on dimension nodes. + # Recurse through metric parents. A node with any metric parent is itself + # derived; a node with no metric parent is a base metric -- regardless of + # whether its data source is a fact/transform or a dimension node (a base + # metric can be defined directly on a dimension node). Classifying by + # "has a metric parent" avoids mistaking a dimension-sourced base metric + # for a bare required-dimension reference, which would drop its grain group. has_metric_parent = False for parent_name in ctx.parent_map.get(node.name, []): parent = ctx.nodes.get(parent_name) @@ -642,7 +645,8 @@ def merge_grain_groups(grain_groups: list[GrainGroup]) -> list[GrainGroup]: """ # 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, as extra grain columns would corrupt MAX_BY/MIN_BY. + # 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) @@ -656,7 +660,8 @@ def merge_grain_groups(grain_groups: list[GrainGroup]) -> list[GrainGroup]: 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, - # corrupting the already-aggregated value for MAX_BY/MIN_BY. + # 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 @@ -706,8 +711,10 @@ def _merge_parent_grain_groups(groups: list[GrainGroup]) -> GrainGroup: gg.reaggregate_component_dimensions, ) - # Carry over non-decomposable metrics from every contributing group to prevent them - # from being silently dropped when merging a NONE group into a FULL/LIMITED neighbor. + # Carry over non-decomposable metrics from every contributing group. + # Without this, merging a NONE group into a FULL/LIMITED neighbor + # silently drops the non-decomposable metric expressions and the + # response loses those metrics entirely. all_non_decomposable: list[DecomposedMetricInfo] = [] for gg in groups: all_non_decomposable.extend(gg.non_decomposable_metrics) diff --git a/datajunction-server/datajunction_server/database/preaggregation.py b/datajunction-server/datajunction_server/database/preaggregation.py index f7265d724..937dd051d 100644 --- a/datajunction-server/datajunction_server/database/preaggregation.py +++ b/datajunction-server/datajunction_server/database/preaggregation.py @@ -222,7 +222,36 @@ def compute_preagg_hash( class PreAggregation(Base): - """First-class pre-aggregation entity that can be shared across cubes.""" + """ + First-class pre-aggregation entity that can be shared across cubes. + + A pre-aggregation represents a materialized grouping of measures at a specific grain, + enabling efficient metric calculations by pre-computing aggregations. + + Pre-aggregations are ALWAYS created by DJ (via /preaggs/plan endpoint) from + metrics + dimensions. Users never manually construct them - this ensures + consistency between DJ-managed (Flow A) and user-managed (Flow B) materialization. + + Key concepts: + - `node_revision`: The specific node revision this pre-agg is based on + - `grain_columns`: Fully qualified dimension references that define the aggregation level + - `measures`: Full MetricComponent info for matching and re-aggregation + - `sql`: The generated SQL for materializing this pre-agg + - `grain_group_hash`: Hash of (node_revision_id + sorted(grain_columns)) for grouping + + Measure format (MetricComponent): + - name: Column name in materialized table + - expression: The raw SQL expression + - expr_hash: Hash of expression for identity matching + - aggregation: Phase 1 function (e.g., "SUM") + - merge: Phase 2 re-aggregation function + - rule: Aggregation rules (type, level) + + Availability tracking: + - Materialization status is tracked via AvailabilityState + - Flow A: DJ's query service posts availability after materialization + - Flow B: User's query service posts to /preaggs/{id}/availability/ + """ __tablename__ = "pre_aggregation" @@ -461,7 +490,33 @@ async def find_matching( grain_columns: list[str], measure_identities: set[str], ) -> PreAggregation | None: - """Find the row this exact declaration already occupies, if any.""" + """ + Find the row this exact declaration already occupies, if any. + + This is the upsert's identity check -- "is this the same declaration I am + about to write?" -- so it matches the uniqueness key exactly: same + revision and grain (via ``grain_group_hash``) and the SAME set of measure + identities, not merely a covering one. + + Covering was wrong here, in two compounding ways. Callers replace the + matched row's contents wholesale, so a narrow declaration could match a + wider pre-agg and silently strip measures off it, breaking routing for + whatever metric depended on the dropped ones. And ``preagg_hash`` is + UNIQUE over exactly ``(node_revision_id, grain_columns, + measure_identities)`` and frozen at insert, so a covering match could + hand back a row whose stored hash no longer described its own contents. + A declaration that genuinely differs now gets its own row instead. + + Identity tokens rather than bare expression hashes, because the latter + made SUM- and MAX-backed pre-aggs look like one row, so registering + either silently overwrote the other. + + Contrast ``find_latest_for_node``, which asks the other question -- "what + did this declaration look like before?" -- and does want covering. + + Returns: + Matching PreAggregation if found, None otherwise + """ grain_group_hash = compute_grain_group_hash(node_revision_id, grain_columns) candidates = await cls.get_by_grain_group_hash(session, grain_group_hash) diff --git a/datajunction-server/datajunction_server/internal/materializations.py b/datajunction-server/datajunction_server/internal/materializations.py index be7816e6b..41025d10e 100644 --- a/datajunction-server/datajunction_server/internal/materializations.py +++ b/datajunction-server/datajunction_server/internal/materializations.py @@ -472,6 +472,36 @@ async def reconcile_declared_materializations( Returns what each declared block resolved to, and the materializations the blocks together superseded. + + A cube may declare more than one block -- typically an `incremental_time` build + for freshness beside a periodic `full` rebuild that corrects late-arriving data + -- and each is reconciled the same way. They cannot collide: the row name is + derived from job, strategy and partition, the job is fixed for a cube and the + strategies are unique by validation, so every declared block owns a row of its + own. + + Mirrors what `POST /nodes/{name}/materialization/` does -- build the config, then + update the row of the same name in place rather than inserting a second one, + since `(name, node_revision_id)` is unique -- with one deliberate difference. The + endpoint decides "unchanged" on `config` alone, but `schedule` and `strategy` are + columns rather than config keys, so by that test a schedule-only edit compares + equal and is silently dropped. Rescheduling is the whole point of a declared + block, so all three are compared here. + + The declared blocks describe *the* materializations for their cube, so every + other active row on the revision is deactivated. The rule is "any active row + whose name no block built" rather than "any row of the same job type", because a + cube's materializations are all writing one Druid datasource and a full rebuild + replaces that datasource wholesale -- so a legacy `druid_measures_cube` row is + just as much a competing writer as a second `druid_cube` row, and matching on job + type would leave it running. The name is what carries the difference: it is + derived from job, strategy and partition, so declaring a strategy the cube was + not already materialized with builds a differently named row, and without this + the cube ends up with two live workflows deleting each other's data. + + A cube planner row is the one thing left alone. It writes a datasource of its own + and DJ cannot rebuild it from a declared block, so superseding it would stop a + workflow nothing here can replace. """ # Snapshotted before the builds, each of which sets the new materialization's # backref and so appends it to this very collection -- searching afterwards would @@ -553,6 +583,26 @@ async def swap_cube_materializations( """ Rebuild a cube's materializations against a new revision and retire the old ones. + Materializations belong to a single `NodeRevision` and availability is scoped to + the revision encoded in the materialized table name, so without this a new cube + revision -- including a metadata-only one -- has neither: the cube silently falls + back to live queries while the superseded revision's workflow keeps posting + availability for a table built from the old definition. Every new revision + therefore swaps, no matter how insignificant the change was. + + Rebuilt rather than copied: a stored config embeds the cube version plus combiner + SQL and a Druid spec derived from the old definition, so the new revision goes + back through `create_new_materialization`. The old materialization is the default + source of the user's intent because it is often the only record of it -- a cube + materialized through the UI has no YAML to read it from. A cube that does declare + `materialization:` passes it as `declared`, which wins, so a push that edits both + a metric and the schedule rebuilds with the new schedule rather than the old. + + `previous_table_usable` is the caller's answer to whether the superseded + revision's materialized table is still valid (`is_non_trivial_cube_change` + inverted), recorded on the history event so an operator can tell whether the + rebuild can adopt the existing data or needs a fresh build and backfill. + Touches only DJ-side state, and returns the query service work still owed -- `None` when the cube had nothing materialized and there is no work at all. The caller commits and then hands the result to `apply_cube_materialization_swap`, so @@ -598,8 +648,18 @@ async def swap_cube_materializations( try: upsert = _upsert_from_materialization(materialization) if declared and isinstance(upsert, UpsertCubeMaterialization): - # The declared block wins over the recovered intent. - # If there are several, match by strategy, otherwise use the first. + # The declared block wins over the recovered intent: a cube that + # declares `materialization:` has its config in the repo, so a + # rebuild triggered by the same deploy must build what the YAML now + # says. Only cube materializations can be declared; anything else + # keeps what was recovered. + # + # Which block, when the cube declares several: the one naming the + # strategy this row was built with, since that is what identifies a + # declared entry. A row whose strategy nothing declares falls to the + # first block, which is what a cube declaring exactly one has always + # done -- and is the only sensible answer, since a rebuild has to + # produce something for a row that is being retired either way. block = next( ( candidate diff --git a/datajunction-server/datajunction_server/internal/nodes.py b/datajunction-server/datajunction_server/internal/nodes.py index 4bbb5960b..c0599cf2f 100644 --- a/datajunction-server/datajunction_server/internal/nodes.py +++ b/datajunction-server/datajunction_server/internal/nodes.py @@ -2524,7 +2524,14 @@ async def _propagate_update_downstream( cache.delete(upstream_cache_key) if downstream.type == NodeType.CUBE: - # Any tier rebuilds except NONE. + # Any tier rebuilds, and the churn is deliberate. Narrowing this by + # comparing the upstream's resolved columns was rejected: a query edit + # can move a filter, a join or a CASE threshold while leaving every + # column and type identical, and each changes every row the cube serves. + # Nothing short of reading the SQL tells those apart, so a changed query + # makes anything built from it suspect. Only NONE is skipped, the one + # case where DJ knows nothing material happened. + # # A rebuild can fail, and one cube's failure must not cost the remaining # downstreams theirs. if change_tier is not ChangeTier.NONE: @@ -4422,6 +4429,17 @@ async def revalidate_node( ] # A tier rather than a version, so propagation can hand the same value to # `bump_version` for downstream cubes. + # + # Any column change is major. An addition looks harmless -- nothing could + # already reference a column that did not exist -- but the query produced it, + # and a query edit can move a filter or a join while leaving the rest of the + # projection identical. Demoting additions to MINOR buys nothing anyway: the + # cube rebuild below skips only NONE, so a minor bump rebuilds all the same. + # + # `order_fixed` earns no tier. DJ filling in a missing projection index is its + # own bookkeeping, not a change to the node, so a revision would describe + # nothing and any tier above NONE would rebuild every cube below. It is applied + # to the current revision in place instead, below. change_tier = fold_change_tiers( [ ChangeTier.MAJOR diff --git a/datajunction-server/datajunction_server/models/deployment.py b/datajunction-server/datajunction_server/models/deployment.py index e49fc6c0b..9b01113a1 100644 --- a/datajunction-server/datajunction_server/models/deployment.py +++ b/datajunction-server/datajunction_server/models/deployment.py @@ -218,10 +218,37 @@ class HierarchySpec(NamespacedSpec): class PreAggSpec(NamespacedSpec): """ Specification for an externally-built pre-aggregation table adopted at deploy - time. ``name`` is a stable handle used for reconciliation. - - Every metric and dimension is declared together with its physical column - in the external table. + time (equivalent to POST /preaggs/register). ``name`` is a stable handle used + for reconciliation and availability callbacks. Metric/dimension references may + use ``${prefix}`` or be fully qualified; they are rendered against the + deployment namespace. + + Every metric and every dimension is declared together with the physical + column of the external table that holds it, as a map:: + + metrics: + ${prefix}paid_members: paid_members_sum + dimensions: + ${prefix}country_dim.country_iso: country + ${prefix}date_dim.utc_date: utc_date + + Both maps require a value for every key -- including a dimension whose + physical column happens to match its DJ column name, which is written out + rather than left empty. An optional value would make the map not really a + mapping, a trailing colon is easy to write by accident, and spelling the + physical name out documents the table in the file that declares it. + + What goes under ``metrics`` are the measures the table stores. A derived + metric (a ratio of two others, say) is not listed and cannot be: it has no + column of its own. It is covered anyway, because both registration and + query-time matching work on decomposed measure identities rather than metric + names, so any metric that decomposes into the stored measures resolves to + this table. + + The earlier four-field form -- ``metrics``/``dimensions`` as lists alongside + separate ``measure_columns``/``dimension_columns`` maps -- is no longer + accepted, and a spec still using it is rejected with a message describing + what to write instead. """ # Metric/dimension reference -> the physical column of the external table @@ -336,7 +363,16 @@ class PartitionSpec(BaseModel): class MaterializationSpec(BaseModel): - """Declarative materialization config for a cube.""" + """ + Declarative materialization config for a cube. + + Deliberately carries only what the author decides: when to build, how, how far + back to look, how much of history to serve, how long the result is kept, and the + Druid, Spark and platform settings their own site needs. The backend that runs it + (and everything it derives -- the measures queries, combiner SQL, the rest of the + Druid spec, output tables) is DJ's choice, so no `job` field is exposed here and + the block stays portable if that choice changes. + """ schedule: str strategy: MaterializationStrategy = MaterializationStrategy.INCREMENTAL_TIME @@ -1280,11 +1316,20 @@ class CubeSpec(NodeSpec): filters: list[str] | None = None columns: list[ColumnSpec] | None = None # Tri-state. A spec materializes the cube on the schedule it names; the `none` - # sentinel tears down whatever the cube has materialized; absent means the cube's - # materialization is not managed here and whatever exists is left alone. + # sentinel tears down whatever the cube has materialized; absent -- and null, + # which is what serializing an absent field produces -- means the cube's + # materialization is not managed here and whatever exists is left alone. Only a + # value can carry intent through serialization, so removal is spelled + # `materialization: none` rather than inferred from a key being present. # - # A list declares more than one (e.g. incremental + full). `strategy` is what tells - # two entries apart, so two entries sharing one are rejected below. + # A list declares more than one, which a cube legitimately needs: an + # `incremental_time` build for freshness alongside a periodic `full` rebuild + # that corrects late-arriving data, out-of-order events and dimension + # backfills. `strategy` is what tells two entries apart -- `job` is not + # authorable (see `MaterializationSpec`) and everything else is a knob rather + # than an identity -- so two entries sharing one are rejected below. The scalar + # form stays valid and means exactly what it always did; nothing has to be + # rewritten as a one-element list. materialization: ( MaterializationSpec | list[MaterializationSpec] | MaterializationAction | None ) = None @@ -1301,8 +1346,18 @@ class CubeSpec(NodeSpec): # Only user-authored partition config is compared here (see __eq__); # everything else about cube columns is auto-derived. "columns": ChangeTier.MAJOR, - # Materialization is not part of a cube's definition, so changing it - # does not mint a new node revision. + # Materialization is not part of a cube's definition. Configuring one + # through `POST /nodes/{name}/materialization/` has never created a node + # revision, and a YAML-declared schedule has to behave the same way, or + # the same edit would cut a version through one door and not the other. + # Because `materialization` is also excluded from `__eq__`, a + # materialization-only edit leaves the cube in the deployment's skip list + # and never reaches the classifier at all; NONE records that intent rather + # than describing a reachable code path. Nothing is lost by the exclusion: + # reconciling the declared block against the persisted config is a separate + # concern that runs over every declared cube whether or not the node + # changed, which is also what catches a materialization changed outside + # YAML. "materialization": ChangeTier.NONE, } diff --git a/datajunction-server/datajunction_server/sql/parsing/ast.py b/datajunction-server/datajunction_server/sql/parsing/ast.py index 3c2570f19..a4a49d51a 100644 --- a/datajunction-server/datajunction_server/sql/parsing/ast.py +++ b/datajunction-server/datajunction_server/sql/parsing/ast.py @@ -1680,6 +1680,29 @@ async def add_ref_column( """ Add column referenced from this table. Returns True if the table has the column and False otherwise. + + This function handles the following cases: + + Regular columns. For example: + (1) non-aliased columns + `SELECT country_id AS country1 FROM countries` should match the `country_id` + column in the table `countries` + (2) aliased columns + `SELECT C.country_id AS country1 FROM countries C` should match the `country_id` + column in the table `countries` with the column namespace/table alias `C` + + Struct columns. For example: + (1) non-aliased struct columns + `countries` has column `identifiers` with type: + STRUCT + `SELECT identifiers.country_name AS name FROM countries` should match the + `identifier` -> `country_name` column in the table `countries` + (2) aliased struct columns + `countries` has column `identifiers` with type: + STRUCT + `SELECT C.identifiers.country_name AS name FROM countries C` should match the + `identifier` -> `country_name` column in the table `countries` with the column namespace/ + table alias `C` """ if not self._columns: if ctx is None: 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 3a7539787..eb92ce22c 100644 --- a/datajunction-server/tests/internal/nodes/derive_frozen_measures_test.py +++ b/datajunction-server/tests/internal/nodes/derive_frozen_measures_test.py @@ -374,8 +374,13 @@ def test_frozen_measure_conflict_rejects_different_measure_identity(): def test_frozen_measure_conflict_rejects_different_tuning_params(): """ - A name collision fails when only tuning parameters differ, preventing - a metric from silently reusing a sketch built with a different accuracy. + 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", @@ -399,7 +404,11 @@ def test_frozen_measure_conflict_rejects_different_tuning_params(): def test_frozen_measure_reuse_allows_matching_params(): """ - Identical params are reusable; absent-vs-empty is not a conflict. + 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", From 4474da0c71b9adc2e7af4ab8ebd06a9040552f6d Mon Sep 17 00:00:00 2001 From: Beto Dealmeida Date: Thu, 24 Sep 2026 16:16:41 -0300 Subject: [PATCH 08/17] Fix sketch decomposition CI checks --- .../api/graphql/scalars/metricmetadata.py | 1 + .../api/graphql/scalars/node.py | 1 + .../api/graphql/schema.graphql | 2 ++ .../graphql/resolvers/test_node_resolver.py | 2 ++ .../build_v3/accumulate_type_test.py | 9 ++++---- .../construction/build_v3/merge_args_test.py | 4 +++- .../build_v3/sketch_family_e2e_test.py | 23 +++++++++++++------ .../tests/sql/decompose_test.py | 1 + 8 files changed, 31 insertions(+), 12 deletions(-) diff --git a/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py b/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py index 04876d498..911642f3c 100644 --- a/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py +++ b/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py @@ -55,6 +55,7 @@ class ReaggregateSpec: is an open dict, which has no automatic GraphQL mapping. """ + fn: strawberry.auto rules: strawberry.auto params: JSON | None = None diff --git a/datajunction-server/datajunction_server/api/graphql/scalars/node.py b/datajunction-server/datajunction_server/api/graphql/scalars/node.py index a80d2692b..9abcdc978 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( + fn=spec.fn, params=spec.params, rules=[ DimensionReaggregateRule( diff --git a/datajunction-server/datajunction_server/api/graphql/schema.graphql b/datajunction-server/datajunction_server/api/graphql/schema.graphql index 4edfd79fd..3f1620a8e 100644 --- a/datajunction-server/datajunction_server/api/graphql/schema.graphql +++ b/datajunction-server/datajunction_server/api/graphql/schema.graphql @@ -727,6 +727,7 @@ type Query { } type ReaggregateSpec { + fn: ReaggregationFunction rules: [DimensionReaggregateRule!]! params: JSON } @@ -740,6 +741,7 @@ enum ReaggregationFunction { FIRST_VALUE MIN MAX + TDIGEST } type SemanticEntity { 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 6d22da2e9..79d59185f 100644 --- a/datajunction-server/tests/api/graphql/resolvers/test_node_resolver.py +++ b/datajunction-server/tests/api/graphql/resolvers/test_node_resolver.py @@ -131,6 +131,7 @@ def test_reaggregate_resolver_returns_metric_spec(): type=NodeType.METRIC, version="1", reaggregate={ + "fn": "tdigest", "rules": [ { "dimension": "default.date_dim.date", @@ -143,6 +144,7 @@ def test_reaggregate_resolver_returns_metric_spec(): result = NodeRevision.reaggregate(NodeRevision, root=db_node_revision) assert result is not None + assert result.fn == ReaggregationFunction.TDIGEST assert len(result.rules) == 1 assert result.rules[0].dimension == "default.date_dim.date" assert result.rules[0].fn == ReaggregationFunction.LAST_VALUE diff --git a/datajunction-server/tests/construction/build_v3/accumulate_type_test.py b/datajunction-server/tests/construction/build_v3/accumulate_type_test.py index 43c048ca4..c26a6fc94 100644 --- a/datajunction-server/tests/construction/build_v3/accumulate_type_test.py +++ b/datajunction-server/tests/construction/build_v3/accumulate_type_test.py @@ -73,7 +73,10 @@ def test_declines_for_a_bare_function_name(self): def test_declines_for_a_single_argument_call(self): # Including templated ones, whose outermost call still takes one arg. - assert _multi_argument_accumulate_types("SUM(POWER(latency_ms, 2))", _parent()) is None + assert ( + _multi_argument_accumulate_types("SUM(POWER(latency_ms, 2))", _parent()) + is None + ) def test_declines_when_a_column_type_is_unrecognized(self): # Better to fall back than to invent a type for the stored schema. @@ -102,9 +105,7 @@ def test_multi_argument_accumulate_is_typed_from_its_arguments(self): ) def test_single_argument_accumulate_is_unchanged(self): - assert ( - infer_component_type(_component("SUM"), "bigint", _parent()) == "double" - ) + assert infer_component_type(_component("SUM"), "bigint", _parent()) == "double" def test_unresolvable_multi_argument_accumulate_falls_back(self): assert ( diff --git a/datajunction-server/tests/construction/build_v3/merge_args_test.py b/datajunction-server/tests/construction/build_v3/merge_args_test.py index e044b4361..ab3e78ddb 100644 --- a/datajunction-server/tests/construction/build_v3/merge_args_test.py +++ b/datajunction-server/tests/construction/build_v3/merge_args_test.py @@ -69,7 +69,9 @@ class TestComponentDefDefaults: """Existing decompositions must be untouched by the new field.""" def test_merge_args_defaults_to_empty(self): - assert ComponentDef(suffix="_sum", accumulate="SUM", merge="SUM").merge_args == () + assert ( + ComponentDef(suffix="_sum", accumulate="SUM", merge="SUM").merge_args == () + ) def test_merge_args_is_a_tuple_so_the_default_cannot_be_mutated(self): first = ComponentDef(suffix="_sum", accumulate="SUM", merge="SUM") diff --git a/datajunction-server/tests/construction/build_v3/sketch_family_e2e_test.py b/datajunction-server/tests/construction/build_v3/sketch_family_e2e_test.py index 95e7ca8f4..98f49f84c 100644 --- a/datajunction-server/tests/construction/build_v3/sketch_family_e2e_test.py +++ b/datajunction-server/tests/construction/build_v3/sketch_family_e2e_test.py @@ -100,18 +100,27 @@ def infer_type(col: ct.ColumnType, compression: ct.ColumnType) -> ct.BinaryType: return ct.BinaryType() -@MergeDigest.register -def infer_type(digest: ct.ColumnType, compression: ct.ColumnType) -> ct.BinaryType: +@MergeDigest.register # type: ignore[no-redef] +def infer_type( + digest: ct.ColumnType, + compression: ct.ColumnType, +) -> ct.BinaryType: return ct.BinaryType() -@DigestQuantiles.register -def infer_type(digest: ct.ColumnType, quantiles: ct.ColumnType) -> ct.ListType: +@DigestQuantiles.register # type: ignore[no-redef] +def infer_type( + digest: ct.ColumnType, + quantiles: ct.ColumnType, +) -> ct.ListType: return ct.ListType(element_type=ct.DoubleType()) -@DigestQuantile.register -def infer_type(digest: ct.ColumnType, quantile: ct.ColumnType) -> ct.DoubleType: +@DigestQuantile.register # type: ignore[no-redef] +def infer_type( + digest: ct.ColumnType, + quantile: ct.ColumnType, +) -> ct.DoubleType: return ct.DoubleType() @@ -288,7 +297,7 @@ async def test_component_identity_carries_the_compression(percentile_metric): """ (group,) = (await _measures(percentile_metric))["grain_groups"] (component,) = group["components"] - (column,) = [c for c in group["columns"] if c["semantic_type"] != "dimension"] + (column,) = (c for c in group["columns"] if c["semantic_type"] != "dimension") assert component["name"] == f"line_total_digest_c{COMPRESSION}_e1f61696" assert component["merge"] == "merge_digest" diff --git a/datajunction-server/tests/sql/decompose_test.py b/datajunction-server/tests/sql/decompose_test.py index a9634671a..c2645bdaa 100644 --- a/datajunction-server/tests/sql/decompose_test.py +++ b/datajunction-server/tests/sql/decompose_test.py @@ -3240,6 +3240,7 @@ async def test_family_gated_percentile_decomposes_through_the_extractor( assert components, "family-gated percentile produced no components" assert components[0].merge == "merge_sketch" + assert components[0].aggregation is not None assert "build_sketch" in components[0].aggregation assert "read_sketch" in str(combiner) From 0bce6ed5c7c31b1090208d24f719f1947fc58de1 Mon Sep 17 00:00:00 2001 From: Beto Dealmeida Date: Thu, 24 Sep 2026 16:42:06 -0300 Subject: [PATCH 09/17] Update reaggregate API roundtrip expectation --- datajunction-server/tests/api/metrics_test.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/datajunction-server/tests/api/metrics_test.py b/datajunction-server/tests/api/metrics_test.py index 91402b042..c6bab38ce 100644 --- a/datajunction-server/tests/api/metrics_test.py +++ b/datajunction-server/tests/api/metrics_test.py @@ -536,6 +536,7 @@ async def test_metric_reaggregate_roundtrip_and_validation( ) assert response.status_code in (200, 201), response.json() assert response.json()["reaggregate"] == { + "fn": None, "params": None, "rules": [ { @@ -548,6 +549,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"] == { + "fn": None, "params": None, "rules": [ { @@ -560,6 +562,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"] == { + "fn": None, "params": None, "rules": [ { From d2160351aee8ab0bdd0e11f39ddae4468f4f6442 Mon Sep 17 00:00:00 2001 From: Beto Dealmeida Date: Thu, 24 Sep 2026 17:03:25 -0300 Subject: [PATCH 10/17] Cover ambiguous accumulate argument types --- .../build_v3/accumulate_type_test.py | 19 +++++++++++++++++++ 1 file changed, 19 insertions(+) diff --git a/datajunction-server/tests/construction/build_v3/accumulate_type_test.py b/datajunction-server/tests/construction/build_v3/accumulate_type_test.py index c26a6fc94..e97cb8457 100644 --- a/datajunction-server/tests/construction/build_v3/accumulate_type_test.py +++ b/datajunction-server/tests/construction/build_v3/accumulate_type_test.py @@ -28,6 +28,8 @@ ) from datajunction_server.models.decompose import AggregationRule, MetricComponent from datajunction_server.models.materialization import MaterializationTarget +from datajunction_server.sql.parsing import ast +from datajunction_server.sql.parsing import types as ct def _parent(column_type: str = "double"): @@ -88,6 +90,23 @@ def test_declines_when_a_column_type_is_unrecognized(self): is None ) + def test_declines_when_an_argument_has_multiple_possible_types(self, monkeypatch): + """Table-valued/lambda-style results cannot type a scalar argument.""" + call = ast.Function( + name=ast.Name("test_accumulate"), + args=[ + SimpleNamespace(type=[ct.IntegerType(), ct.StringType()]), + SimpleNamespace(type=ct.IntegerType()), + ], + ) + parsed = SimpleNamespace(select=SimpleNamespace(projection=[call])) + monkeypatch.setattr( + "datajunction_server.construction.build_v3.measures.parse", + lambda _: parsed, + ) + + assert _multi_argument_accumulate_types("ignored", _parent()) is None + class TestInferredColumnType: """What the measures column ends up recorded as.""" From 91fe63b894f5bffb98fc63f50e88209d5389cbef Mon Sep 17 00:00:00 2001 From: Beto Dealmeida Date: Thu, 24 Sep 2026 17:21:49 -0300 Subject: [PATCH 11/17] Make materialization test order-independent --- .../tests/api/materializations_test.py | 24 ++++++++++++------- 1 file changed, 16 insertions(+), 8 deletions(-) diff --git a/datajunction-server/tests/api/materializations_test.py b/datajunction-server/tests/api/materializations_test.py index ca305b167..4995099b0 100644 --- a/datajunction-server/tests/api/materializations_test.py +++ b/datajunction-server/tests/api/materializations_test.py @@ -1082,7 +1082,11 @@ async def test_spark_sql_full( expected_query = load_expected_file("spark_sql.full.partition.query.sql") args, _ = module__query_service_client.materialize.call_args_list[-1] # type: ignore assert str(parse(args[0].query)) == str(parse(expected_query)) - materialization_with_partitions = data["materializations"][1] + materialization_with_partitions = next( + materialization + for materialization in data["materializations"] + if materialization["name"] == "spark_sql__full__birth_date__country" + ) del materialization_with_partitions["config"]["query"] expected_config = load_expected_file("spark_sql.full.partition.config.json") expected_config["node_revision_id"] = mock.ANY @@ -1093,15 +1097,19 @@ async def test_spark_sql_full( "/nodes/default.hard_hat/materializations/", ) materializations = response.json() - materializations[0]["config"]["query"] = mock.ANY - materializations[0]["node_revision_id"] = mock.ANY - assert materializations[0] == load_expected_file( + materializations_by_name = { + materialization["name"]: materialization for materialization in materializations + } + unpartitioned = materializations_by_name["spark_sql__full"] + unpartitioned["config"]["query"] = mock.ANY + unpartitioned["node_revision_id"] = mock.ANY + assert unpartitioned == load_expected_file( "spark_sql.full.materializations.json", ) - materializations = response.json() - materializations[1]["config"]["query"] = mock.ANY - materializations[1]["node_revision_id"] = mock.ANY - assert materializations[1] == load_expected_file( + partitioned = materializations_by_name["spark_sql__full__birth_date__country"] + partitioned["config"]["query"] = mock.ANY + partitioned["node_revision_id"] = mock.ANY + assert partitioned == load_expected_file( "spark_sql.full.partition.materializations.json", ) From bb2a43582df4764b8a5d90aa9a88c57bfdd875f8 Mon Sep 17 00:00:00 2001 From: Robin Davis Date: Fri, 25 Sep 2026 17:44:46 -0700 Subject: [PATCH 12/17] Fix sketch decomposition review findings --- .../datajunction_server/api/cubes.py | 17 ++++ .../api/graphql/scalars/metricmetadata.py | 8 ++ .../api/graphql/schema.graphql | 11 ++- .../construction/build_v3/combiners.py | 31 +++++- .../construction/build_v3/decomposition.py | 6 +- .../construction/build_v3/measures.py | 42 +++++--- .../datajunction_server/internal/nodes.py | 8 +- .../models/cube_materialization.py | 23 ++++- .../models/materialization.py | 22 +++++ .../datajunction_server/sql/decompose.py | 98 ++++++++++++++----- .../api/cubes_build_metrics_spec_test.py | 22 +++++ .../graphql/scalars/test_metric_component.py | 43 ++++++++ datajunction-server/tests/api/metrics_test.py | 38 +++++++ .../build_v3/accumulate_type_test.py | 26 +++++ .../construction/build_v3/combiners_test.py | 40 ++++++++ .../build_v3/preagg_substitution_test.py | 66 +++++++++++++ .../build_v3/sketch_family_e2e_test.py | 11 ++- .../models/cube_druid_sketch_spec_test.py | 18 ++++ .../tests/sql/decompose_test.py | 74 ++++++++++---- .../AddEditNodePageFormSuccess.test.jsx | 46 +++++++-- .../src/app/pages/AddEditNodePage/index.jsx | 27 +++++ datajunction-ui/src/app/services/DJService.js | 3 + .../app/services/__tests__/DJService.test.jsx | 3 + 23 files changed, 603 insertions(+), 80 deletions(-) create mode 100644 datajunction-server/tests/api/graphql/scalars/test_metric_component.py diff --git a/datajunction-server/datajunction_server/api/cubes.py b/datajunction-server/datajunction_server/api/cubes.py index 4c1bdf469..f98035522 100644 --- a/datajunction-server/datajunction_server/api/cubes.py +++ b/datajunction-server/datajunction_server/api/cubes.py @@ -66,6 +66,7 @@ Granularity, MaterializationJobTypeEnum, MaterializationStrategy, + MaterializationTarget, get_druid_aggregator_spec, ) from datajunction_server.models.metric import TranslatedSQL @@ -189,6 +190,21 @@ def _build_metrics_spec( if component else None ) + if ( + metric_spec is None + and component is not None + and ( + component.params + or component.serializes_for(MaterializationTarget.DRUID) + ) + ): + raise DJInvalidInputException( + message=( + f"No Druid aggregator is registered for sketch measure " + f"`{col.name}` with column type `{col.type}` and merge " + f"function `{component.merge or component.aggregation}`." + ), + ) # 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 { @@ -570,6 +586,7 @@ async def materialize_cube( dimensions=cube_revision.cube_node_dimensions, filters=cube_revision.cube_filters or None, dialect=Dialect.SPARK, + materialization_target=MaterializationTarget.DRUID, ) except Exception as e: # pragma: no cover raise DJInvalidInputException( # pragma: no cover diff --git a/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py b/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py index 911642f3c..1633f7f98 100644 --- a/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py +++ b/datajunction-server/datajunction_server/api/graphql/scalars/metricmetadata.py @@ -16,6 +16,9 @@ MetricComponent as MetricComponent_, ) from datajunction_server.models.node import MetricDirection as MetricDirection_ +from datajunction_server.models.materialization import ( + MaterializationTarget as MaterializationTarget_, +) from datajunction_server.models.reaggregate import ( DimensionReaggregateRule as DimensionReaggregateRule_, ReaggregateSpec as ReaggregateSpec_, @@ -25,6 +28,7 @@ MetricDirection = strawberry.enum(MetricDirection_) Aggregability = strawberry.enum(Aggregability_) ReaggregationFunction = strawberry.enum(ReaggregationFunction_) +MaterializationTarget = strawberry.enum(MaterializationTarget_) @strawberry.type @@ -77,9 +81,13 @@ class MetricComponent: expression: strawberry.auto aggregation: strawberry.auto merge: strawberry.auto + merge_args: strawberry.auto rule: strawberry.auto grain_alias: strawberry.auto params: JSON | None = None + serialize: strawberry.auto + serialize_targets: strawberry.auto + serialize_type: strawberry.auto @strawberry.experimental.pydantic.type(model=DecomposedMetric_, all_fields=True) diff --git a/datajunction-server/datajunction_server/api/graphql/schema.graphql b/datajunction-server/datajunction_server/api/graphql/schema.graphql index 3f1620a8e..2fcd0f150 100644 --- a/datajunction-server/datajunction_server/api/graphql/schema.graphql +++ b/datajunction-server/datajunction_server/api/graphql/schema.graphql @@ -288,6 +288,11 @@ type MaterializationPlan { units: [MaterializationUnit!]! } +enum MaterializationTarget { + DRUID + ICEBERG +} + type MaterializationUnit { upstream: VersionedRef! grainDimensions: [VersionedRef!]! @@ -304,6 +309,10 @@ type MetricComponent { rule: AggregationRule! grainAlias: String params: JSON + mergeArgs: [String!]! + serialize: String + serializeTargets: [MaterializationTarget!]! + serializeType: String } enum MetricDirection { @@ -802,4 +811,4 @@ type User { type VersionedRef { name: String! version: String! -} \ No newline at end of file +} diff --git a/datajunction-server/datajunction_server/construction/build_v3/combiners.py b/datajunction-server/datajunction_server/construction/build_v3/combiners.py index 7c5d710cf..2e2b54a93 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/combiners.py +++ b/datajunction-server/datajunction_server/construction/build_v3/combiners.py @@ -24,6 +24,7 @@ process_metric_combiner_expression, ) from datajunction_server.construction.build_v3.decomposition import ( + apply_component_serialize, build_merge_call, ) from datajunction_server.construction.build_v3.preagg_matcher import ( @@ -43,7 +44,10 @@ from datajunction_server.models.column import SemanticType from datajunction_server.models.decompose import MetricComponent from datajunction_server.models.dialect import Dialect -from datajunction_server.models.materialization import MaterializationStrategy +from datajunction_server.models.materialization import ( + MaterializationStrategy, + MaterializationTarget, +) from datajunction_server.models.query import V3ColumnMetadata from datajunction_server.sql.parsing import ast from datajunction_server.sql.parsing.ast import render_for_dialect, to_sql @@ -526,6 +530,7 @@ async def build_combiner_sql_from_preaggs( dimensions: list[str], filters: list[str] | None = None, dialect=None, + materialization_target: MaterializationTarget | None = None, ) -> tuple[ CombinedGrainGroupResult, list[PreAggSourceInfo], @@ -664,6 +669,7 @@ async def build_combiner_sql_from_preaggs( preagg_gg = _build_grain_group_from_preagg_table( gg, full_table_ref, + materialization_target, ) preagg_grain_groups.append(preagg_gg) @@ -830,6 +836,7 @@ def _get_projection_name(proj: ast.Node) -> str | None: def _build_grain_group_from_preagg_table( original_gg: GrainGroupSQL, preagg_table_ref: str, + materialization_target: MaterializationTarget | None = None, ) -> GrainGroupSQL: """ Build a GrainGroupSQL that reads from a pre-agg table. @@ -851,6 +858,7 @@ def _build_grain_group_from_preagg_table( # Build SELECT columns select_items: list[ast.Aliasable | ast.Expression | ast.Column] = [] group_by_cols: list[str] = [] + serialized_types: dict[str, str] = {} # Add dimension columns for grain_col in original_gg.grain: @@ -866,6 +874,7 @@ def _build_grain_group_from_preagg_table( # Find the component to get the merge function merge_func = None merge_args: list[str] = [] + component: MetricComponent | None = None for comp in original_gg.components: if ( # pragma: no branch comp.name == col.name @@ -873,11 +882,27 @@ def _build_grain_group_from_preagg_table( ): merge_func = comp.merge merge_args = comp.merge_args + component = comp break if merge_func: # Apply re-aggregation - agg_expr = build_merge_call(merge_func, merge_args, col_ref) + agg_expr: ast.Expression = build_merge_call( + merge_func, + merge_args, + col_ref, + ) + if component is not None: + agg_expr = apply_component_serialize( + agg_expr, + component, + materialization_target, + ) + if ( + component.serializes_for(materialization_target) + and component.serialize_type + ): + serialized_types[col.name] = component.serialize_type aliased = ast.Alias(child=agg_expr, alias=ast.Name(col.name)) select_items.append(aliased) else: @@ -907,7 +932,7 @@ def _build_grain_group_from_preagg_table( ColumnMetadata( name=col.name, semantic_name=col.semantic_name, - type=col.type, + type=serialized_types.get(col.name, col.type), semantic_type=col.semantic_type, ) for col in original_gg.columns diff --git a/datajunction-server/datajunction_server/construction/build_v3/decomposition.py b/datajunction-server/datajunction_server/construction/build_v3/decomposition.py index 4e6773bea..276efd8eb 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/decomposition.py +++ b/datajunction-server/datajunction_server/construction/build_v3/decomposition.py @@ -270,16 +270,16 @@ def build_component_expression( ) # All accumulate shapes must receive the target conversion. - return _apply_serialize(accumulated, component, materialization_target) + return apply_component_serialize(accumulated, component, materialization_target) -def _apply_serialize( +def apply_component_serialize( expr: ast.Expression, component: MetricComponent, materialization_target: MaterializationTarget | None, ) -> ast.Expression: """ - Wrap an accumulated expression in the component's serialize conversion. + Wrap an accumulated or merged expression in the component's conversion. Returns the expression untouched unless the component declares a conversion and names this target -- so every existing component, and every query-time diff --git a/datajunction-server/datajunction_server/construction/build_v3/measures.py b/datajunction-server/datajunction_server/construction/build_v3/measures.py index 5f24d6365..cc4df71f8 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/measures.py +++ b/datajunction-server/datajunction_server/construction/build_v3/measures.py @@ -27,6 +27,7 @@ ) from datajunction_server.construction.build_v3.decomposition import ( analyze_grain_groups, + apply_component_serialize, build_component_expression, build_merge_call, merge_grain_groups, @@ -84,6 +85,7 @@ from datajunction_server.sql.parsing import ast from datajunction_server.sql.parsing import types as ct from datajunction_server.sql.parsing.backends.antlr4 import parse +from datajunction_server.sql.parsing.backends.exceptions import DJParseException from datajunction_server.utils import SEPARATOR _logger = logging.getLogger(__name__) @@ -185,19 +187,21 @@ def _multi_argument_accumulate_types( arg_types: list[ct.ColumnType] = [] for arg in call.args: - resolved: ct.ColumnType | None - if isinstance(arg, ast.Column): - # Only the columns need the parent to resolve; literals and casts - # carry their own type without being bound to a table. - resolved = _parse_type_string( - get_column_type(parent_node, str(arg.alias_or_name.name)), - ) - else: - inferred = arg.type - if isinstance(inferred, list): - return None - resolved = inferred - if resolved is None: + # A parsed arithmetic or nested-function argument still has unbound + # columns. Resolve them before asking the whole expression for its type. + if isinstance(arg, ast.Expression): + for column in arg.find_all(ast.Column): + column_type = _parse_type_string( + get_column_type(parent_node, column.name.name), + ) + if column_type is None: + return None + column.add_type(column_type) + try: + resolved = arg.type + except (DJParseException, TypeError, AttributeError): + return None + if resolved is None or isinstance(resolved, list): return None arg_types.append(resolved) return arg_types @@ -1918,11 +1922,16 @@ def build_grain_group_from_preagg( # If no merge function, output column directly (e.g., grain column for LIMITED) # Otherwise, apply the merge function for re-aggregation if component.merge: - agg_expr = build_merge_call( + agg_expr: ast.Expression = build_merge_call( component.merge, component.merge_args, _preagg_column(scan_name, scan_alias), ) + agg_expr = apply_component_serialize( + agg_expr, + component, + ctx.materialization_target, + ) aliased = ast.Alias(child=agg_expr, alias=ast.Name(output_alias)) select_items.append(aliased) else: @@ -1933,6 +1942,11 @@ def build_grain_group_from_preagg( # Get type from pre-agg columns col_type = preagg.get_column_type(measure_col, default="double") + if ( + component.serializes_for(ctx.materialization_target) + and component.serialize_type + ): + col_type = component.serialize_type columns.append( ColumnMetadata( name=output_alias, diff --git a/datajunction-server/datajunction_server/internal/nodes.py b/datajunction-server/datajunction_server/internal/nodes.py index c0599cf2f..7ffe410f0 100644 --- a/datajunction-server/datajunction_server/internal/nodes.py +++ b/datajunction-server/datajunction_server/internal/nodes.py @@ -2873,11 +2873,9 @@ async def create_new_revision_from_existing( 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) - ) + reaggregate_changes = reaggregate_was_set and dump_reaggregate_spec( + old_revision.reaggregate + ) != dump_reaggregate_spec(data.reaggregate if data else None) major_changes = ( query_changes or column_changes diff --git a/datajunction-server/datajunction_server/models/cube_materialization.py b/datajunction-server/datajunction_server/models/cube_materialization.py index 23abfbd2b..30dd95560 100644 --- a/datajunction-server/datajunction_server/models/cube_materialization.py +++ b/datajunction-server/datajunction_server/models/cube_materialization.py @@ -19,6 +19,7 @@ CoverageSpec, MaterializationJobTypeEnum, MaterializationStrategy, + MaterializationTarget, SparkSpec, get_druid_aggregator_spec, ) @@ -385,19 +386,31 @@ 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 - specs = ( - get_druid_aggregator_spec( + specs = [] + for measure in self.measures: + spec = get_druid_aggregator_spec( column_name=measure.name, column_type=column_mapping.get(measure.name), aggregation=measure.aggregation, merge=measure.merge, params=measure.params, ) - for measure in self.measures - ) + if spec is None and ( + measure.params or measure.serializes_for(MaterializationTarget.DRUID) + ): + raise DJInvalidInputException( + message=( + f"No Druid aggregator is registered for sketch measure " + f"`{measure.name}` with column type " + f"`{column_mapping.get(measure.name)}` and merge function " + f"`{measure.merge or measure.aggregation}`." + ), + ) + if spec is not None: + specs.append(spec) # 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] + return specs @computed_field # type: ignore[misc] @property diff --git a/datajunction-server/datajunction_server/models/materialization.py b/datajunction-server/datajunction_server/models/materialization.py index 0dfed5393..87e756436 100644 --- a/datajunction-server/datajunction_server/models/materialization.py +++ b/datajunction-server/datajunction_server/models/materialization.py @@ -1,6 +1,7 @@ """Models for materialization""" import enum +import math import re from collections.abc import Callable from datetime import date, timedelta @@ -175,6 +176,27 @@ def get_druid_aggregator_spec( ) ), ) + for key, value in params.items(): + default = family_config[key] # unknown keys were rejected above + if isinstance(default, (int, float)) and not isinstance(default, bool): + valid_type = isinstance(value, (int, float)) and not isinstance( + value, bool + ) + if valid_type: + try: + valid_type = math.isfinite(value) + except OverflowError: + # Huge JSON integers cannot be represented as a Druid number. + valid_type = False + else: + valid_type = isinstance(value, type(default)) + if not valid_type: + raise DJInvalidInputException( + message=( + f"Druid aggregator `{aggregator}` parameter `{key}` must " + f"have the same kind of value as its default `{default}`." + ), + ) spec: dict[str, Any] = { "fieldName": column_name, diff --git a/datajunction-server/datajunction_server/sql/decompose.py b/datajunction-server/datajunction_server/sql/decompose.py index 2173f1b34..affd72b9d 100644 --- a/datajunction-server/datajunction_server/sql/decompose.py +++ b/datajunction-server/datajunction_server/sql/decompose.py @@ -180,25 +180,44 @@ def combine( DECOMPOSITION_REGISTRY: dict[type, type[AggDecomposition] | None] = {} -# Decompositions selected by the metric's declared reaggregation family rather -# than by its aggregation function. A family entry wins over the by-function -# registry, which is what lets `APPROX_PERCENTILE(x, 0.95)` keep its spelling and -# still decompose, but only for metrics that opted in with `reaggregate.fn`. -FAMILY_DECOMPOSITION_REGISTRY: dict[ReaggregationFunction, type[AggDecomposition]] = {} +# A declared family overrides the by-function registry only for the aggregate +# functions named at registration. This lets an opted-in APPROX_PERCENTILE +# decompose without changing unrelated aggregates in the same metric. +@dataclass(frozen=True) +class FamilyDecompositionRegistration: + """A family implementation and the aggregate functions it can replace.""" + decomposition: type[AggDecomposition] + aggregate_functions: frozenset[type] -def decomposes_family(fn: ReaggregationFunction): + +FAMILY_DECOMPOSITION_REGISTRY: dict[ + ReaggregationFunction, + FamilyDecompositionRegistration, +] = {} + + +def decomposes_family( + fn: ReaggregationFunction, + *, + aggregate_functions: tuple[type, ...], +): """ Register a decomposition for a reaggregation family. - Exported for downstream deployments, which supply the engine-specific - functions a family needs. OSS registers none, so a metric declaring an - unregistered family falls through to the by-function registry and keeps the - aggregability it has today. + Downstream deployments supply the engine-specific functions and explicitly + name the aggregate calls this family can replace. Other aggregates in the + same metric retain their ordinary decomposition. """ + if not aggregate_functions: + raise ValueError("A decomposition family must name aggregate functions") + def decorator(decomp_class: type[AggDecomposition]): - FAMILY_DECOMPOSITION_REGISTRY[fn] = decomp_class + FAMILY_DECOMPOSITION_REGISTRY[fn] = FamilyDecompositionRegistration( + decomposition=decomp_class, + aggregate_functions=frozenset(aggregate_functions), + ) return decomp_class return decorator @@ -806,19 +825,17 @@ def get_decomposition( """ Get the decomposition for an aggregation, or None if not decomposable. - A metric that declares a reaggregation family gets that family's - decomposition, which is how an opt-in sketch overrides the default handling - of its aggregation function. Everything else resolves by function class as - before. + A declared family overrides the default decomposition only for aggregate + functions it registered. Everything else resolves by function class. A declared family with nothing registered for it falls through rather than raising: OSS ships no family implementations, so the metric simply keeps the aggregability it would have had. """ if reaggregate is not None and reaggregate.fn is not None: - family_class = FAMILY_DECOMPOSITION_REGISTRY.get(reaggregate.fn) - if family_class is not None: - return family_class(params=reaggregate.params) + family = FAMILY_DECOMPOSITION_REGISTRY.get(reaggregate.fn) + if family is not None and func_class in family.aggregate_functions: + return family.decomposition(params=reaggregate.params) decomp_class = DECOMPOSITION_REGISTRY.get(func_class) if decomp_class is None: @@ -1596,6 +1613,25 @@ def _extract_base( if dj_function and dj_function.is_aggregation: agg_funcs.append((func, dj_function)) + family = ( + FAMILY_DECOMPOSITION_REGISTRY.get(reaggregate.fn) + if reaggregate is not None and reaggregate.fn is not None + else None + ) + if ( + family + and reaggregate is not None + and reaggregate.fn is not None + and not any( + dj_function in family.aggregate_functions + for _, dj_function in agg_funcs + ) + ): + self._raise_unsupported_reaggregate_shape( + f"family `{reaggregate.fn.value}` does not support any " + "aggregation in this metric", + ) + # If any aggregation is non-decomposable, abort decomposition # entirely — the metric is non-decomposable as a whole. # @@ -1644,7 +1680,11 @@ 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: + if ( + reaggregate is not None + and reaggregate.params + and reaggregate.fn not in FAMILY_DECOMPOSITION_REGISTRY + ): self._attach_reaggregate_params(components, reaggregate) if fixed_grain is not None: @@ -1713,12 +1753,11 @@ def _attach_reaggregate_params( for component in components if component.aggregation is not None and component.merge is not None ] - if not configurable: + if len(configurable) != 1 or len(components) != 1: self._raise_unsupported_reaggregate_shape( - "parameterized reaggregation requires an aggregating component", + "parameterized reaggregation requires exactly one aggregating component", ) - for component in configurable: - component.params = dict(reaggregate.params or {}) + configurable[0].params = dict(reaggregate.params or {}) def _attach_reaggregate_spec( self, @@ -1830,6 +1869,19 @@ def _decompose( self._make_component(func, comp_def, query_ast) for comp_def in decomposition.components ] + family = ( + FAMILY_DECOMPOSITION_REGISTRY.get(reaggregate.fn) + if reaggregate is not None and reaggregate.fn is not None + else None + ) + if ( + family + and reaggregate is not None + and dj_function in family.aggregate_functions + and reaggregate.params + ): + for component in components: + component.params = dict(reaggregate.params) # Build combiner AST is_distinct = func.quantifier == ast.SetQuantifier.Distinct 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 e196fa151..e7e3b628b 100644 --- a/datajunction-server/tests/api/cubes_build_metrics_spec_test.py +++ b/datajunction-server/tests/api/cubes_build_metrics_spec_test.py @@ -588,6 +588,28 @@ def test_register_without_config(self): finally: DRUID_AGG_MAPPING.pop(("bigint", "test_plain_merge"), None) + @pytest.mark.parametrize("bad_value", ["broken", True, float("nan"), 10**1000]) + def test_rejects_invalid_compression(self, bad_value): + """Invalid tuning values fail before a Druid ingestion job is submitted.""" + try: + register_druid_aggregator( + column_type="binary", + merge_func="test_checked_tdigest_agg", + aggregator="testCheckedTDigestSketch", + default_config={"compression": 200}, + ) + with pytest.raises(DJInvalidInputException, match="compression"): + get_druid_aggregator_spec( + column_name="latency_tdigest", + column_type="binary", + aggregation="test_checked_tdigest_agg", + merge="test_checked_tdigest_agg", + params={"compression": bad_value}, + ) + finally: + DRUID_AGG_MAPPING.pop(("binary", "test_checked_tdigest_agg"), None) + DRUID_SKETCH_CONFIG.pop("testCheckedTDigestSketch", None) + def test_conflicting_defaults_are_rejected(self): """ Re-registering an aggregator with different defaults raises an error. diff --git a/datajunction-server/tests/api/graphql/scalars/test_metric_component.py b/datajunction-server/tests/api/graphql/scalars/test_metric_component.py new file mode 100644 index 000000000..113135f87 --- /dev/null +++ b/datajunction-server/tests/api/graphql/scalars/test_metric_component.py @@ -0,0 +1,43 @@ +"""GraphQL exposes the same sketch component metadata as REST.""" + +import strawberry + +from datajunction_server.api.graphql.scalars.metricmetadata import ( + MetricComponent as GraphQLMetricComponent, +) +from datajunction_server.models.decompose import AggregationRule, MetricComponent +from datajunction_server.models.materialization import MaterializationTarget + + +def test_sketch_component_fields_are_queryable(): + component = MetricComponent( + name="latency_digest", + expression="latency_ms", + aggregation="build_digest(latency_ms, 200)", + merge="merge_digest", + merge_args=["200"], + serialize="digest_to_bytes({})", + serialize_targets=[MaterializationTarget.DRUID], + serialize_type="binary", + rule=AggregationRule(), + ) + + @strawberry.type + class Query: + @strawberry.field + def metric_component(self) -> GraphQLMetricComponent: + return GraphQLMetricComponent.from_pydantic(component) # type: ignore[attr-defined] + + result = strawberry.Schema(query=Query).execute_sync( + "{ metricComponent { mergeArgs serialize serializeTargets serializeType } }", + ) + + assert result.errors is None + assert result.data == { + "metricComponent": { + "mergeArgs": ["200"], + "serialize": "digest_to_bytes({})", + "serializeTargets": ["DRUID"], + "serializeType": "binary", + }, + } diff --git a/datajunction-server/tests/api/metrics_test.py b/datajunction-server/tests/api/metrics_test.py index c6bab38ce..c5efa7660 100644 --- a/datajunction-server/tests/api/metrics_test.py +++ b/datajunction-server/tests/api/metrics_test.py @@ -611,6 +611,44 @@ async def test_metric_reaggregate_roundtrip_and_validation( } +@pytest.mark.asyncio +async def test_legacy_reaggregate_shape_does_not_force_major_revision( + client_with_roads: AsyncClient, + session: AsyncSession, +) -> None: + """A legacy rules-only JSON spec is equal to the new canonical shape.""" + metric_name = f"default.legacy_reaggregate_{uuid4().hex}" + rules = [ + { + "dimension": "default.repair_orders_fact.repair_order_id", + "fn": "last_value", + }, + ] + response = await client_with_roads.post( + "/nodes/metric/", + json={ + "name": metric_name, + "query": "SELECT COUNT(repair_order_id) FROM default.repair_orders_fact", + "mode": "published", + "reaggregate": {"rules": rules}, + }, + ) + assert response.status_code in (200, 201), response.json() + + node = await Node.get_by_name(session, metric_name) + assert node is not None and node.current is not None + node.current.reaggregate = {"rules": rules} + await session.commit() + original_version = node.current_version + + response = await client_with_roads.patch( + f"/nodes/{metric_name}/", + json={"reaggregate": {"rules": rules}}, + ) + assert response.status_code == 200, response.json() + assert response.json()["current_version"] == original_version + + @pytest_asyncio.fixture(scope="module") async def module__current_user(module__session: AsyncSession) -> User: """ diff --git a/datajunction-server/tests/construction/build_v3/accumulate_type_test.py b/datajunction-server/tests/construction/build_v3/accumulate_type_test.py index e97cb8457..54f991d41 100644 --- a/datajunction-server/tests/construction/build_v3/accumulate_type_test.py +++ b/datajunction-server/tests/construction/build_v3/accumulate_type_test.py @@ -66,6 +66,13 @@ def test_resolves_a_cast_without_binding_it_to_a_table(self): ) assert [str(t) for t in types] == ["double", "double"] + def test_resolves_columns_inside_an_expression(self): + types = _multi_argument_accumulate_types( + "POWER(latency_ms * 2, CAST(200.0 AS DOUBLE))", + _parent(), + ) + assert [str(t) for t in types] == ["double", "double"] + def test_declines_without_a_parent_to_resolve_columns_against(self): assert _multi_argument_accumulate_types("POWER(latency_ms, 2)", None) is None @@ -107,6 +114,15 @@ def test_declines_when_an_argument_has_multiple_possible_types(self, monkeypatch assert _multi_argument_accumulate_types("ignored", _parent()) is None + def test_declines_when_a_nested_column_cannot_be_typed(self): + assert ( + _multi_argument_accumulate_types( + "POWER(latency_ms * 2, CAST(200.0 AS DOUBLE))", + _parent("some_unknown_type"), + ) + is None + ) + class TestInferredColumnType: """What the measures column ends up recorded as.""" @@ -123,6 +139,16 @@ def test_multi_argument_accumulate_is_typed_from_its_arguments(self): == "double" ) + def test_multi_argument_accumulate_accepts_an_expression_input(self): + assert ( + infer_component_type( + _component("POWER(latency_ms * 2, CAST(200.0 AS DOUBLE))"), + "bigint", + _parent(), + ) + == "double" + ) + def test_single_argument_accumulate_is_unchanged(self): assert infer_component_type(_component("SUM"), "bigint", _parent()) == "double" diff --git a/datajunction-server/tests/construction/build_v3/combiners_test.py b/datajunction-server/tests/construction/build_v3/combiners_test.py index 34b3b81fa..c01862d56 100644 --- a/datajunction-server/tests/construction/build_v3/combiners_test.py +++ b/datajunction-server/tests/construction/build_v3/combiners_test.py @@ -40,6 +40,7 @@ MetricComponent, PreAggMeasure, ) +from datajunction_server.models.materialization import MaterializationTarget from datajunction_server.models.partition import Granularity, PartitionType from datajunction_server.models.query import V3ColumnMetadata from datajunction_server.sql.parsing import ast @@ -1015,6 +1016,45 @@ def test_build_grain_group_from_preagg_table(self): # Grain should be preserved assert result.grain == ["date_id"] + def test_preagg_table_merge_converts_for_druid(self): + """The cube combiner converts sketches and reports the stored type.""" + component = MetricComponent( + name="digest", + expression="amount", + aggregation="build_digest(amount, 200)", + merge="merge_digest", + merge_args=["200"], + serialize="digest_to_bytes({})", + serialize_targets=[MaterializationTarget.DRUID], + serialize_type="binary", + rule=AggregationRule(type=Aggregability.FULL), + ) + gg = _create_grain_group( + sql="SELECT date_id, build_digest(amount, 200) AS digest FROM orders GROUP BY date_id", + columns=[ + {"name": "date_id", "semantic_type": "dimension"}, + { + "name": "digest", + "semantic_type": "metric_component", + "type": "struct", + }, + ], + grain=["date_id"], + components=[component], + ) + + raw = _build_grain_group_from_preagg_table(gg, "wh.preaggs.t") + druid = _build_grain_group_from_preagg_table( + gg, + "wh.preaggs.t", + MaterializationTarget.DRUID, + ) + + assert "digest_to_bytes" not in str(raw.query) + assert raw.columns[1].type == "struct" + assert "digest_to_bytes(merge_digest(digest, 200))" in str(druid.query) + assert druid.columns[1].type == "binary" + def test_build_grain_group_from_preagg_table_no_merge_function(self): """ Measures without merge function should be selected directly. 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 f8a045a70..e15e07b54 100644 --- a/datajunction-server/tests/construction/build_v3/preagg_substitution_test.py +++ b/datajunction-server/tests/construction/build_v3/preagg_substitution_test.py @@ -36,6 +36,8 @@ MetricComponent, PreAggMeasure, ) +from datajunction_server.models.materialization import MaterializationTarget +from datajunction_server.models.query import V3ColumnMetadata from datajunction_server.utils import get_query_service_client from . import assert_sql_equal, get_first_grain_group @@ -3305,6 +3307,70 @@ def test_deduplicates_repeated_components(self): # Only one component should appear in output despite two in input assert len(result.components) == 1 + @pytest.mark.parametrize( + ("target", "expected_type", "serialized"), + [ + (None, "struct", False), + (MaterializationTarget.DRUID, "binary", True), + ], + ) + def test_sketch_merge_converts_for_druid(self, target, expected_type, serialized): + """A cube reusing an Iceberg pre-agg converts the merged sketch.""" + from datajunction_server.database.availabilitystate import AvailabilityState + + node = self._make_node() + component = MetricComponent( + name="rev_digest", + expression="revenue", + aggregation="build_digest(revenue, 200)", + merge="merge_digest", + merge_args=["200"], + serialize="digest_to_bytes({})", + serialize_targets=[MaterializationTarget.DRUID], + serialize_type="binary", + rule=AggregationRule(type=Aggregability.FULL), + ) + measure = PreAggMeasure( + **component.model_dump(), + expr_hash=compute_expression_hash("revenue"), + ) + preagg = PreAggregation( + node_revision_id=1, + grain_columns=[], + measures=[measure], + columns=[ + V3ColumnMetadata( + name="rev_digest", + type="struct", + semantic_name="test_node:rev_digest", + semantic_type="metric_component", + ), + ], + sql="SELECT 1", + grain_group_hash="abc", + preagg_hash="def", + availability=AvailabilityState( + catalog="wh", + schema_="preaggs", + table="tbl", + valid_through_ts=99999, + ), + ) + ctx = self._make_ctx() + ctx.materialization_target = target + + result = build_grain_group_from_preagg( + ctx, + self._make_grain_group(node, [(node, component)]), + preagg, + resolved_dimensions=[], + components_per_metric={}, + ) + + assert ("digest_to_bytes(" in str(result.query)) is serialized + assert "merge_digest(rev_digest, 200)" in str(result.query) + assert result.columns[0].type == expected_type + class TestPreAggFreshnessGating: """ diff --git a/datajunction-server/tests/construction/build_v3/sketch_family_e2e_test.py b/datajunction-server/tests/construction/build_v3/sketch_family_e2e_test.py index 98f49f84c..e6f274d0b 100644 --- a/datajunction-server/tests/construction/build_v3/sketch_family_e2e_test.py +++ b/datajunction-server/tests/construction/build_v3/sketch_family_e2e_test.py @@ -43,7 +43,11 @@ decomposes_family, make_func, ) -from datajunction_server.sql.functions import Function, function_registry +from datajunction_server.sql.functions import ( + ApproxPercentile, + Function, + function_registry, +) from datajunction_server.sql.parsing import ast from datajunction_server.sql.parsing import types as ct @@ -139,7 +143,10 @@ def digest_family(): saved[key] = function_registry.get(key) function_registry[key] = cls - @decomposes_family(ReaggregationFunction.TDIGEST) + @decomposes_family( + ReaggregationFunction.TDIGEST, + aggregate_functions=(ApproxPercentile,), + ) class _DigestDecomposition(AggDecomposition): @property def compression(self) -> int: diff --git a/datajunction-server/tests/models/cube_druid_sketch_spec_test.py b/datajunction-server/tests/models/cube_druid_sketch_spec_test.py index 0a26916fc..8d96a267a 100644 --- a/datajunction-server/tests/models/cube_druid_sketch_spec_test.py +++ b/datajunction-server/tests/models/cube_druid_sketch_spec_test.py @@ -2,9 +2,13 @@ Regression tests for sketch measures reaching Druid ingestion. """ +import pytest + from datajunction_server.api.cubes import _build_metrics_spec +from datajunction_server.errors import DJInvalidInputException from datajunction_server.models.cube_materialization import CombineMaterialization from datajunction_server.models.decompose import AggregationRule, MetricComponent +from datajunction_server.models.materialization import MaterializationTarget from datajunction_server.models.node_type import NodeNameVersion from datajunction_server.models.partition import Granularity from datajunction_server.models.query import ColumnMetadata @@ -194,3 +198,17 @@ def test_builders_disagree_on_unmappable_measures(self): assert v3_spec == [] assert cubes_api_spec[0]["type"] == "longSum" + + def test_unregistered_serialized_sketch_fails_in_both_builders(self): + """A sketch cannot be silently omitted or ingested as a numeric sum.""" + measure = _measure("p95_latency", "custom_tdigest_agg", "custom_tdigest_agg") + measure.params = {"compression": 200} + measure.serialize = "custom_tdigest_sketch({})" + measure.serialize_targets = [MaterializationTarget.DRUID] + measure.serialize_type = "binary" + column = ColumnMetadata(name="p95_latency", type="binary") + + with pytest.raises(DJInvalidInputException, match="No Druid aggregator"): + _combiner([measure], [column]).metrics_spec() + with pytest.raises(DJInvalidInputException, match="No Druid aggregator"): + _build_metrics_spec([column], [measure], {}) diff --git a/datajunction-server/tests/sql/decompose_test.py b/datajunction-server/tests/sql/decompose_test.py index c2645bdaa..b5dd79406 100644 --- a/datajunction-server/tests/sql/decompose_test.py +++ b/datajunction-server/tests/sql/decompose_test.py @@ -2842,10 +2842,7 @@ class TestReaggregateParams: @staticmethod def _spec(params): """A spec carrying params.""" - return ReaggregateSpec.model_construct( - rules=[], - params=params, - ) + return ReaggregateSpec(params=params) def test_params_reach_components(self): extractor = MetricComponentExtractor(1) @@ -2857,29 +2854,28 @@ def test_params_reach_components(self): {"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.""" + def test_params_without_a_family_reject_multi_component_metric(self): + """A generic parameter cannot be broadcast to AVG's SUM and COUNT.""" 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) + with pytest.raises(DJInvalidInputException, match="exactly one"): + extractor._extract_base( + parse("SELECT AVG(latency_ms) FROM t"), + self._spec({"compression": 200}), + ) def test_params_are_copied_not_shared(self): """ - Each component receives an independent copy of its params dict. + The component keeps a copy of the declaration's params dict. """ extractor = MetricComponentExtractor(1) params = {"compression": 200} components, _ = extractor._extract_base( - parse("SELECT AVG(latency_ms) FROM t"), + parse("SELECT SUM(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 + assert components[0].params is not params def test_no_params_leaves_components_untouched(self): """Omitting params applies no configuration to components.""" @@ -2900,7 +2896,7 @@ def test_params_without_an_aggregating_component_is_rejected(self): parse("SELECT COUNT(DISTINCT order_id) FROM t"), self._spec({"compression": 200}), ) - assert "requires an aggregating component" in str(excinfo.value) + assert "requires exactly one aggregating component" in str(excinfo.value) # ============================================================================= @@ -3137,7 +3133,10 @@ def registered_family(): selection path at all. """ - @decomposes_family(ReaggregationFunction.TDIGEST) + @decomposes_family( + ReaggregationFunction.TDIGEST, + aggregate_functions=(dj_functions.ApproxPercentile,), + ) class _FamilyDecomposition(AggDecomposition): @property def components(self) -> list[ComponentDef]: @@ -3172,6 +3171,47 @@ def test_declared_family_overrides_the_function_registry(registered_family): assert isinstance(decomposition, registered_family) +def test_declared_family_keeps_unrelated_aggregates(registered_family): + """A percentile family cannot replace a SUM in the same metric.""" + spec = ReaggregateSpec(fn=ReaggregationFunction.TDIGEST) + + assert not isinstance(get_decomposition(dj_functions.Sum, spec), registered_family) + + +def test_family_registration_requires_aggregate_functions(): + """An unrestricted family registration would affect every aggregate.""" + with pytest.raises(ValueError, match="must name aggregate functions"): + decomposes_family(ReaggregationFunction.TDIGEST, aggregate_functions=()) + + +def test_mixed_family_metric_attaches_params_only_to_family_component( + registered_family, +): + """SUM remains ordinary while the opted-in percentile carries tuning.""" + spec = ReaggregateSpec( + fn=ReaggregationFunction.TDIGEST, + params={"compression": 200}, + ) + components, _ = MetricComponentExtractor(0)._extract_base( + parse("SELECT APPROX_PERCENTILE(price, 0.95) + SUM(quantity) FROM t"), + spec, + ) + + by_merge = {component.merge: component for component in components} + assert by_merge["merge_sketch"].params == {"compression": 200} + assert by_merge["SUM"].params is None + + +def test_family_rejects_metric_without_a_supported_aggregate(registered_family): + """A mismatched declaration fails clearly instead of calling the combiner.""" + spec = ReaggregateSpec(fn=ReaggregationFunction.TDIGEST) + with pytest.raises(DJInvalidInputException, match="does not support"): + MetricComponentExtractor(0)._extract_base( + parse("SELECT SUM(price) FROM t"), + spec, + ) + + def test_family_receives_its_tuning_params(registered_family): """ `reaggregate.params` reaches the decomposition, and thus the accumulate. 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 c66a9cdb8..687d2b05c 100644 --- a/datajunction-ui/src/app/pages/AddEditNodePage/__tests__/AddEditNodePageFormSuccess.test.jsx +++ b/datajunction-ui/src/app/pages/AddEditNodePage/__tests__/AddEditNodePageFormSuccess.test.jsx @@ -449,14 +449,46 @@ describe('AddEditNodePage submission succeeded', () => { [], ['dj'], { key1: 'value1', key2: 'value2' }, - { - rules: [ - { - dimension: 'v3.date.date_id[order]', - fn: 'last_value', - }, - ], + ); + }); + }, 1000000); + + it('preserves a sketch family when editing an unrelated metric field', async () => { + const mockDjClient = initializeMockDJClient(); + mockDjClient.DataJunctionAPI.getNodeForEditing.mockReturnValue({ + ...mocks.mockGetMetricNode, + current: { + ...mocks.mockGetMetricNode.current, + reaggregate: { + fn: 'TDIGEST', + params: { compression: 200 }, + rules: [], }, + }, + }); + mockDjClient.DataJunctionAPI.patchNode = vi.fn().mockReturnValue({ + status: 201, + json: { name: 'default.num_repair_orders', type: 'metric' }, + }); + mockDjClient.DataJunctionAPI.tagsNode.mockReturnValue({ + status: 200, + json: { message: 'Success' }, + }); + mockDjClient.DataJunctionAPI.listTags.mockReturnValue([]); + mockDjClient.DataJunctionAPI.whoami.mockReturnValue({ + id: 123, + username: 'test_user', + }); + + renderEditNode(testElement(mockDjClient)); + await screen.findByLabelText('Description'); + await userEvent.type(screen.getByLabelText('Description'), '!!!'); + await userEvent.click(screen.getByText('Save')); + + await waitFor(() => { + expect(mockDjClient.DataJunctionAPI.patchNode).toHaveBeenCalledTimes(1); + expect(mockDjClient.DataJunctionAPI.patchNode.mock.calls[0]).toHaveLength( + 12, ); }); }, 1000000); diff --git a/datajunction-ui/src/app/pages/AddEditNodePage/index.jsx b/datajunction-ui/src/app/pages/AddEditNodePage/index.jsx index fcb0f3088..d3ecb6368 100644 --- a/datajunction-ui/src/app/pages/AddEditNodePage/index.jsx +++ b/datajunction-ui/src/app/pages/AddEditNodePage/index.jsx @@ -75,6 +75,7 @@ export function AddEditNodePage({ extensions = {} }) { reaggregate_dimension: '', reaggregate_function: '', had_reaggregate: false, + reaggregate_original: null, }; const validator = values => { @@ -174,8 +175,28 @@ export function AddEditNodePage({ extensions = {} }) { }; const buildReaggregateSpec = values => { + // The form only edits dimension rules. Preserve sketch-family settings when + // changing another field (or clearing a dimension rule). + const original = values.reaggregate_original; + const originalRule = original?.rules?.[0]; + if ( + original && + (originalRule?.dimension || '') === + (values.reaggregate_dimension || '') && + (originalRule?.fn?.toLowerCase() || '') === + (values.reaggregate_function || '') + ) { + // Omit unchanged metadata so a stale edit form cannot overwrite a newer + // family or tuning-parameter update made by another client. + return undefined; + } + const family = { + ...(original?.fn ? { fn: original.fn.toLowerCase() } : {}), + ...(original?.params != null ? { params: original.params } : {}), + }; if (values.reaggregate_dimension && values.reaggregate_function) { return { + ...family, rules: [ { dimension: values.reaggregate_dimension, @@ -184,6 +205,9 @@ export function AddEditNodePage({ extensions = {} }) { ], }; } + if (Object.keys(family).length) { + return family; + } return values.had_reaggregate ? null : undefined; }; @@ -343,6 +367,7 @@ export function AddEditNodePage({ extensions = {} }) { firstReaggregateRule(node.current.reaggregate)?.fn, ), had_reaggregate: Boolean(node.current.reaggregate), + reaggregate_original: node.current.reaggregate, upstream_node: '', // Derived metrics have no upstream node aggregate_expression: derivedExpression, }; @@ -363,6 +388,7 @@ export function AddEditNodePage({ extensions = {} }) { firstReaggregateRule(node.current.reaggregate)?.fn, ), had_reaggregate: Boolean(node.current.reaggregate), + reaggregate_original: node.current.reaggregate, upstream_node: nonMetricParent?.name || '', aggregate_expression: node.current.metricMetadata?.expression, }; @@ -416,6 +442,7 @@ export function AddEditNodePage({ extensions = {} }) { 'reaggregate_dimension', 'reaggregate_function', 'had_reaggregate', + 'reaggregate_original', 'owners', 'custom_metadata', ]; diff --git a/datajunction-ui/src/app/services/DJService.js b/datajunction-ui/src/app/services/DJService.js index 95ff088e2..0b9d8dc59 100644 --- a/datajunction-ui/src/app/services/DJService.js +++ b/datajunction-ui/src/app/services/DJService.js @@ -521,6 +521,8 @@ export const DataJunctionAPI = { name } reaggregate { + fn + params rules { dimension fn @@ -615,6 +617,7 @@ export const DataJunctionAPI = { name } reaggregate { + fn rules { dimension fn diff --git a/datajunction-ui/src/app/services/__tests__/DJService.test.jsx b/datajunction-ui/src/app/services/__tests__/DJService.test.jsx index 3c6cfb5d3..aee0ad6fb 100644 --- a/datajunction-ui/src/app/services/__tests__/DJService.test.jsx +++ b/datajunction-ui/src/app/services/__tests__/DJService.test.jsx @@ -1427,6 +1427,7 @@ describe('DataJunctionAPI', () => { expect(requestBody.query).toContain('dimension'); expect(requestBody.query).toContain('rules'); expect(requestBody.query).toContain('fn'); + expect(requestBody.query).not.toContain('weight'); }); it('calls notebookExportCube correctly', async () => { @@ -2051,6 +2052,8 @@ describe('DataJunctionAPI', () => { expect(requestBody.query).toContain('dimension'); expect(requestBody.query).toContain('rules'); expect(requestBody.query).toContain('fn'); + expect(requestBody.query).toContain('params'); + expect(requestBody.query).not.toContain('weight'); }); it('returns null when getNodeForEditing finds no nodes', async () => { From 6d26b7cf0aaae3acdd84690a465bde8efaf5615e Mon Sep 17 00:00:00 2001 From: Robin Davis Date: Sat, 26 Sep 2026 23:15:44 -0700 Subject: [PATCH 13/17] Fix Druid sketch combiners and legacy reaggregation params --- .../datajunction_server/api/cubes.py | 1 + .../construction/build_v3/builder.py | 8 ++++ .../construction/build_v3/combiners.py | 3 ++ .../construction/build_v3/decomposition.py | 17 +++---- .../construction/build_v3/types.py | 2 + .../internal/cube_materializations.py | 1 + .../datajunction_server/sql/decompose.py | 13 ++++-- .../construction/build_v3/merge_args_test.py | 8 ++++ .../build_v3/sketch_family_e2e_test.py | 44 +++++++++++++++++++ .../tests/sql/decompose_test.py | 17 +++---- 10 files changed, 94 insertions(+), 20 deletions(-) diff --git a/datajunction-server/datajunction_server/api/cubes.py b/datajunction-server/datajunction_server/api/cubes.py index f98035522..3602c6431 100644 --- a/datajunction-server/datajunction_server/api/cubes.py +++ b/datajunction-server/datajunction_server/api/cubes.py @@ -586,6 +586,7 @@ async def materialize_cube( dimensions=cube_revision.cube_node_dimensions, filters=cube_revision.cube_filters or None, dialect=Dialect.SPARK, + combiner_dialect=Dialect.DRUID, materialization_target=MaterializationTarget.DRUID, ) except Exception as e: # pragma: no cover diff --git a/datajunction-server/datajunction_server/construction/build_v3/builder.py b/datajunction-server/datajunction_server/construction/build_v3/builder.py index 79f45f7d8..0276b93c2 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/builder.py +++ b/datajunction-server/datajunction_server/construction/build_v3/builder.py @@ -254,6 +254,7 @@ async def setup_build_context( lookback_window: str | None = None, matched_cube: NodeRevision | None = None, materialization_target: MaterializationTarget | None = None, + combiner_dialect: Dialect | None = None, ) -> BuildContext: """ Create and initialize a BuildContext with all setup done. @@ -271,6 +272,8 @@ async def setup_build_context( dimensions: List of dimension names filters: Optional list of filter expressions dialect: SQL dialect for output + combiner_dialect: Optional dialect for metric combiners when they run + somewhere other than the measures SQL engine. use_materialized: Whether to use materialized tables include_temporal_filters: Whether to include temporal partition filters from cube lookback_window: Lookback window for temporal filters @@ -293,6 +296,7 @@ async def setup_build_context( dimensions=list(dimensions), filters=filters or [], dialect=dialect, + combiner_dialect=combiner_dialect, materialization_target=materialization_target, use_materialized=use_materialized, temporal_partition_columns=temporal_partition_columns or {}, @@ -381,6 +385,7 @@ async def build_measures_sql( query_parameters: dict[str, Any] | None = None, matched_cube: NodeRevision | None = None, materialization_target: MaterializationTarget | None = None, + combiner_dialect: Dialect | None = None, ) -> GeneratedMeasuresSQL: """ Build measures SQL for a set of metrics, dimensions, and filters. @@ -396,6 +401,8 @@ async def build_measures_sql( dimensions: List of dimension names (format: "node.column" or "node.column[role]") filters: Optional list of filter expressions dialect: SQL dialect for output + combiner_dialect: Optional dialect for metric combiners when they run + somewhere other than the measures SQL engine. use_materialized: If True (default), use materialized tables when available. Set to False when generating SQL for materialization refresh to avoid circular references. @@ -425,6 +432,7 @@ async def build_measures_sql( lookback_window=lookback_window, matched_cube=matched_cube, materialization_target=materialization_target, + combiner_dialect=combiner_dialect, ) # Build grain groups from context diff --git a/datajunction-server/datajunction_server/construction/build_v3/combiners.py b/datajunction-server/datajunction_server/construction/build_v3/combiners.py index 2e2b54a93..7c7ab4d33 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/combiners.py +++ b/datajunction-server/datajunction_server/construction/build_v3/combiners.py @@ -531,6 +531,7 @@ async def build_combiner_sql_from_preaggs( filters: list[str] | None = None, dialect=None, materialization_target: MaterializationTarget | None = None, + combiner_dialect: Dialect | None = None, ) -> tuple[ CombinedGrainGroupResult, list[PreAggSourceInfo], @@ -552,6 +553,7 @@ async def build_combiner_sql_from_preaggs( dimensions: List of dimension references filters: Optional filters dialect: SQL dialect + combiner_dialect: Optional dialect for the final metric combiners. Returns: Tuple of: @@ -569,6 +571,7 @@ async def build_combiner_sql_from_preaggs( filters=filters, dialect=dialect or Dialect.SPARK, use_materialized=False, # We'll manually reference pre-agg tables + combiner_dialect=combiner_dialect, ) if not result.grain_groups: # pragma: no cover diff --git a/datajunction-server/datajunction_server/construction/build_v3/decomposition.py b/datajunction-server/datajunction_server/construction/build_v3/decomposition.py index 276efd8eb..28ce6f539 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/decomposition.py +++ b/datajunction-server/datajunction_server/construction/build_v3/decomposition.py @@ -32,7 +32,7 @@ from datajunction_server.models.node_type import NodeType 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.sql.parsing.backends.antlr4 import cached_parse, parse from datajunction_server.utils import SEPARATOR @@ -81,7 +81,7 @@ async def decompose_and_group_metrics( base_metric, nodes_cache=ctx.nodes, parent_map=ctx.parent_map, - dialect=ctx.dialect, + dialect=ctx.combiner_dialect or ctx.dialect, ) all_decomposed[base_metric.name] = decomposed @@ -101,7 +101,7 @@ async def decompose_and_group_metrics( metric_node, nodes_cache=ctx.nodes, parent_map=ctx.parent_map, - dialect=ctx.dialect, + dialect=ctx.combiner_dialect or ctx.dialect, ) all_decomposed[metric_name] = derived_decomposed else: @@ -125,7 +125,7 @@ async def decompose_and_group_metrics( metric_node, nodes_cache=ctx.nodes, parent_map=ctx.parent_map, - dialect=ctx.dialect, + dialect=ctx.combiner_dialect or ctx.dialect, ) all_decomposed[metric_node.name] = decomposed @@ -168,9 +168,8 @@ async def decompose_metric( parent_map, avoids database queries by using cached data. parent_map: Optional dict of child_name -> list of parent_names. Required if nodes_cache is provided. - dialect: Dialect the combiner is rendered for. ``decompose_and_group_metrics`` - passes the build's resolved dialect, which for a materialized cube is - the one taken from its availability catalog -- Druid, in practice. + dialect: Dialect the combiner is rendered for. This may differ from + the measures SQL dialect when Spark writes a Druid cube. Defaults to Spark for direct callers. Returns: @@ -310,7 +309,9 @@ def build_merge_call( """ extra: list[ast.Expression] = [] for literal in merge_args: - parsed = parse(f"SELECT {literal}").select.projection[0] + # Family tuning literals recur across components and requests. The + # cached parser returns a deep copy, so parent links remain independent. + parsed = cached_parse(f"SELECT {literal}").select.projection[0] parsed.clear_parent() extra.append(cast(ast.Expression, parsed)) return ast.Function(name=ast.Name(merge), args=[arg, *extra]) diff --git a/datajunction-server/datajunction_server/construction/build_v3/types.py b/datajunction-server/datajunction_server/construction/build_v3/types.py index 9f50ef097..18748d23f 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/types.py +++ b/datajunction-server/datajunction_server/construction/build_v3/types.py @@ -44,6 +44,8 @@ class BuildContext: dimensions: list[str] filters: list[str] = field(default_factory=list) dialect: Dialect = Dialect.SPARK + # The measures may run in Spark while their metric combiner runs in Druid. + combiner_dialect: Dialect | None = None # Set only when this build is producing a measures table for materialization. # None for query-time builds, which never serialize. materialization_target: MaterializationTarget | None = None diff --git a/datajunction-server/datajunction_server/internal/cube_materializations.py b/datajunction-server/datajunction_server/internal/cube_materializations.py index aa529c30b..6839ae7ec 100644 --- a/datajunction-server/datajunction_server/internal/cube_materializations.py +++ b/datajunction-server/datajunction_server/internal/cube_materializations.py @@ -373,6 +373,7 @@ async def build_cube_materialization( dimensions=current_revision.cube_node_dimensions, filters=(current_revision.cube_filters or []) + extra_filters, dialect=Dialect.SPARK, + combiner_dialect=Dialect.DRUID, use_materialized=True, # Measures SQL runs in Spark for every target, so the dialect above # cannot say where the output is headed. Druid is the only cube target diff --git a/datajunction-server/datajunction_server/sql/decompose.py b/datajunction-server/datajunction_server/sql/decompose.py index affd72b9d..986026e05 100644 --- a/datajunction-server/datajunction_server/sql/decompose.py +++ b/datajunction-server/datajunction_server/sql/decompose.py @@ -1746,18 +1746,23 @@ def _attach_reaggregate_params( reaggregate: ReaggregateSpec, ) -> None: """ - Propagate tuning parameters from the reaggregate spec to aggregating components. + Attach tuning parameters only when they identify one unambiguous component. """ configurable = [ component for component in components if component.aggregation is not None and component.merge is not None ] - if len(configurable) != 1 or len(components) != 1: + # Older metrics could declare params on a multi-component expression + # (notably AVG). Those params never described SUM and COUNT separately; + # broadcasting them changes component identity and may select the wrong + # pre-aggregation. Keep such metrics queryable without propagating them. + if not configurable: self._raise_unsupported_reaggregate_shape( - "parameterized reaggregation requires exactly one aggregating component", + "parameterized reaggregation requires an aggregating component", ) - configurable[0].params = dict(reaggregate.params or {}) + if len(configurable) == 1 and len(components) == 1: + configurable[0].params = dict(reaggregate.params or {}) def _attach_reaggregate_spec( self, diff --git a/datajunction-server/tests/construction/build_v3/merge_args_test.py b/datajunction-server/tests/construction/build_v3/merge_args_test.py index ab3e78ddb..59e3f93a1 100644 --- a/datajunction-server/tests/construction/build_v3/merge_args_test.py +++ b/datajunction-server/tests/construction/build_v3/merge_args_test.py @@ -64,6 +64,14 @@ def test_arguments_are_parsed_not_pasted(self): assert isinstance(expr.args[1], ast.Expression) assert str(expr) == "f(c, 1 + 2)" + def test_cached_literal_ast_is_independent_between_calls(self): + """A caller's rewrite must not mutate the cached parse template.""" + first = build_merge_call("f", ["CAST(200.0 AS DOUBLE)"], _column("a")) + first.args[1] = ast.Number(999) + + second = build_merge_call("f", ["CAST(200.0 AS DOUBLE)"], _column("b")) + assert str(second) == "f(b, CAST(200.0 AS DOUBLE))" + class TestComponentDefDefaults: """Existing decompositions must be untouched by the new field.""" diff --git a/datajunction-server/tests/construction/build_v3/sketch_family_e2e_test.py b/datajunction-server/tests/construction/build_v3/sketch_family_e2e_test.py index e6f274d0b..00e3ab45a 100644 --- a/datajunction-server/tests/construction/build_v3/sketch_family_e2e_test.py +++ b/datajunction-server/tests/construction/build_v3/sketch_family_e2e_test.py @@ -33,6 +33,10 @@ import pytest from httpx import AsyncClient +from datajunction_server.construction.build_v3.builder import build_measures_sql +from datajunction_server.construction.build_v3.combiners import ( + build_combiner_sql_from_preaggs, +) from datajunction_server.models.dialect import Dialect from datajunction_server.models.materialization import MaterializationTarget from datajunction_server.models.reaggregate import ReaggregationFunction @@ -50,6 +54,7 @@ ) from datajunction_server.sql.parsing import ast from datajunction_server.sql.parsing import types as ct +from datajunction_server.utils import get_session from . import assert_sql_equal @@ -378,6 +383,45 @@ async def test_combiner_takes_the_shape_the_engine_needs( assert formula["combiner"] == expected +@pytest.mark.asyncio +async def test_druid_cube_measures_keep_spark_sql_and_use_druid_combiner( + percentile_metric, +): + """The legacy cube path runs measures in Spark but combines in Druid.""" + session = percentile_metric.app.dependency_overrides[get_session]() + result = await build_measures_sql( + session=session, + metrics=[METRIC], + dimensions=DIMENSIONS, + dialect=Dialect.SPARK, + combiner_dialect=Dialect.DRUID, + materialization_target=MaterializationTarget.DRUID, + ) + + assert "build_digest(" in result.grain_groups[0].sql + assert "digest_to_bytes(" in result.grain_groups[0].sql + assert "DIGEST_QUANTILE(" in str(result.decomposed_metrics[METRIC].derived_ast) + assert "digest_quantiles(" not in str(result.decomposed_metrics[METRIC].derived_ast) + + +@pytest.mark.asyncio +async def test_druid_cube_preagg_combiner_uses_druid_shape(percentile_metric): + """The pre-agg cube path must not carry a Spark-only combiner to Druid.""" + await _publish_preagg(percentile_metric) + session = percentile_metric.app.dependency_overrides[get_session]() + result, _, _ = await build_combiner_sql_from_preaggs( + session=session, + metrics=[METRIC], + dimensions=DIMENSIONS, + dialect=Dialect.SPARK, + combiner_dialect=Dialect.DRUID, + materialization_target=MaterializationTarget.DRUID, + ) + + assert "DIGEST_QUANTILE(" in result.metric_combiners[METRIC] + assert "digest_quantiles(" not in result.metric_combiners[METRIC] + + @pytest.mark.asyncio async def test_metrics_layer_recomputes_from_source(percentile_metric): """ diff --git a/datajunction-server/tests/sql/decompose_test.py b/datajunction-server/tests/sql/decompose_test.py index b5dd79406..47c5d922c 100644 --- a/datajunction-server/tests/sql/decompose_test.py +++ b/datajunction-server/tests/sql/decompose_test.py @@ -2854,14 +2854,15 @@ def test_params_reach_components(self): {"compression": 200}, ] - def test_params_without_a_family_reject_multi_component_metric(self): - """A generic parameter cannot be broadcast to AVG's SUM and COUNT.""" + def test_params_without_a_family_leave_multi_component_metric_untouched(self): + """Legacy AVG params cannot be broadcast to its SUM and COUNT.""" extractor = MetricComponentExtractor(1) - with pytest.raises(DJInvalidInputException, match="exactly one"): - extractor._extract_base( - parse("SELECT AVG(latency_ms) FROM t"), - self._spec({"compression": 200}), - ) + components, _ = extractor._extract_base( + parse("SELECT AVG(latency_ms) FROM t"), + self._spec({"compression": 200}), + ) + assert len(components) == 2 + assert all(component.params is None for component in components) def test_params_are_copied_not_shared(self): """ @@ -2896,7 +2897,7 @@ def test_params_without_an_aggregating_component_is_rejected(self): parse("SELECT COUNT(DISTINCT order_id) FROM t"), self._spec({"compression": 200}), ) - assert "requires exactly one aggregating component" in str(excinfo.value) + assert "requires an aggregating component" in str(excinfo.value) # ============================================================================= From 1e972deb5b533c36a696736f62ff3969d9eceb44 Mon Sep 17 00:00:00 2001 From: Beto Dealmeida Date: Mon, 28 Sep 2026 17:15:59 -0400 Subject: [PATCH 14/17] Address sketch decomposition review feedback --- .../construction/build_v3/decomposition.py | 2 +- datajunction-server/datajunction_server/internal/nodes.py | 2 +- .../datajunction_server/models/materialization.py | 4 +++- 3 files changed, 5 insertions(+), 3 deletions(-) diff --git a/datajunction-server/datajunction_server/construction/build_v3/decomposition.py b/datajunction-server/datajunction_server/construction/build_v3/decomposition.py index 28ce6f539..1dfa05718 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/decomposition.py +++ b/datajunction-server/datajunction_server/construction/build_v3/decomposition.py @@ -289,7 +289,7 @@ def apply_component_serialize( serialize = component.serialize if serialize is None: # pragma: no cover - guaranteed by serializes_for return expr - wrapped = parse( + wrapped = cached_parse( f"SELECT {serialize.replace('{}', str(expr))}", ).select.projection[0] wrapped.clear_parent() diff --git a/datajunction-server/datajunction_server/internal/nodes.py b/datajunction-server/datajunction_server/internal/nodes.py index 7ffe410f0..5e8f42fe6 100644 --- a/datajunction-server/datajunction_server/internal/nodes.py +++ b/datajunction-server/datajunction_server/internal/nodes.py @@ -2874,7 +2874,7 @@ async def create_new_revision_from_existing( data and "reaggregate" in data.model_fields_set, ) reaggregate_changes = reaggregate_was_set and dump_reaggregate_spec( - old_revision.reaggregate + old_revision.reaggregate, ) != dump_reaggregate_spec(data.reaggregate if data else None) major_changes = ( query_changes diff --git a/datajunction-server/datajunction_server/models/materialization.py b/datajunction-server/datajunction_server/models/materialization.py index 87e756436..e8f410591 100644 --- a/datajunction-server/datajunction_server/models/materialization.py +++ b/datajunction-server/datajunction_server/models/materialization.py @@ -176,11 +176,13 @@ def get_druid_aggregator_spec( ) ), ) + assert family_config is not None # params are non-empty and all keys are valid for key, value in params.items(): default = family_config[key] # unknown keys were rejected above if isinstance(default, (int, float)) and not isinstance(default, bool): valid_type = isinstance(value, (int, float)) and not isinstance( - value, bool + value, + bool, ) if valid_type: try: From 8731269119e9e4c6bec97bb9c7676f8357c43d42 Mon Sep 17 00:00:00 2001 From: Beto Dealmeida Date: Mon, 28 Sep 2026 18:09:49 -0400 Subject: [PATCH 15/17] Fix sketch decomposition coverage and schema checks --- .../api/graphql/schema.graphql | 2 +- .../construction/build_v3/combiners.py | 22 +++++++++--------- .../api/cubes_build_metrics_spec_test.py | 22 ++++++++++++++++++ .../build_v3/accumulate_type_test.py | 23 +++++++++++++++++++ 4 files changed, 57 insertions(+), 12 deletions(-) diff --git a/datajunction-server/datajunction_server/api/graphql/schema.graphql b/datajunction-server/datajunction_server/api/graphql/schema.graphql index 2fcd0f150..1ea4f8ec2 100644 --- a/datajunction-server/datajunction_server/api/graphql/schema.graphql +++ b/datajunction-server/datajunction_server/api/graphql/schema.graphql @@ -811,4 +811,4 @@ type User { type VersionedRef { name: String! version: String! -} +} \ No newline at end of file diff --git a/datajunction-server/datajunction_server/construction/build_v3/combiners.py b/datajunction-server/datajunction_server/construction/build_v3/combiners.py index 7c7ab4d33..4e135c5d0 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/combiners.py +++ b/datajunction-server/datajunction_server/construction/build_v3/combiners.py @@ -890,22 +890,22 @@ def _build_grain_group_from_preagg_table( if merge_func: # Apply re-aggregation + assert component is not None agg_expr: ast.Expression = build_merge_call( merge_func, merge_args, col_ref, ) - if component is not None: - agg_expr = apply_component_serialize( - agg_expr, - component, - materialization_target, - ) - if ( - component.serializes_for(materialization_target) - and component.serialize_type - ): - serialized_types[col.name] = component.serialize_type + agg_expr = apply_component_serialize( + agg_expr, + component, + materialization_target, + ) + if ( + component.serializes_for(materialization_target) + and component.serialize_type + ): + serialized_types[col.name] = component.serialize_type aliased = ast.Alias(child=agg_expr, alias=ast.Name(col.name)) select_items.append(aliased) else: 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 e7e3b628b..08b40fb5b 100644 --- a/datajunction-server/tests/api/cubes_build_metrics_spec_test.py +++ b/datajunction-server/tests/api/cubes_build_metrics_spec_test.py @@ -588,6 +588,28 @@ def test_register_without_config(self): finally: DRUID_AGG_MAPPING.pop(("bigint", "test_plain_merge"), None) + def test_accepts_non_numeric_config_value(self): + """String-valued tuning knobs are checked against string defaults.""" + try: + register_druid_aggregator( + column_type="binary", + merge_func="test_mode_agg", + aggregator="testModeSketch", + default_config={"mode": "compact"}, + ) + spec = get_druid_aggregator_spec( + column_name="sketch", + column_type="binary", + aggregation="test_mode", + merge="test_mode_agg", + params={"mode": "fast"}, + ) + assert spec is not None + assert spec["mode"] == "fast" + finally: + DRUID_AGG_MAPPING.pop(("binary", "test_mode_agg"), None) + DRUID_SKETCH_CONFIG.pop("testModeSketch", None) + @pytest.mark.parametrize("bad_value", ["broken", True, float("nan"), 10**1000]) def test_rejects_invalid_compression(self, bad_value): """Invalid tuning values fail before a Druid ingestion job is submitted.""" diff --git a/datajunction-server/tests/construction/build_v3/accumulate_type_test.py b/datajunction-server/tests/construction/build_v3/accumulate_type_test.py index 54f991d41..8cbc1ec26 100644 --- a/datajunction-server/tests/construction/build_v3/accumulate_type_test.py +++ b/datajunction-server/tests/construction/build_v3/accumulate_type_test.py @@ -114,6 +114,29 @@ def test_declines_when_an_argument_has_multiple_possible_types(self, monkeypatch assert _multi_argument_accumulate_types("ignored", _parent()) is None + def test_declines_when_argument_type_resolution_raises(self, monkeypatch): + """An argument that cannot infer its type leaves the caller on fallback.""" + + class UntypeableArgument: + @property + def type(self): + raise TypeError("cannot infer argument type") + + call = ast.Function( + name=ast.Name("test_accumulate"), + args=[ + UntypeableArgument(), + SimpleNamespace(type=ct.IntegerType()), + ], + ) + parsed = SimpleNamespace(select=SimpleNamespace(projection=[call])) + monkeypatch.setattr( + "datajunction_server.construction.build_v3.measures.parse", + lambda _: parsed, + ) + + assert _multi_argument_accumulate_types("ignored", _parent()) is None + def test_declines_when_a_nested_column_cannot_be_typed(self): assert ( _multi_argument_accumulate_types( From ada55fef6403001f1bfa1c483eeb8978da2c5e6d Mon Sep 17 00:00:00 2001 From: Beto Dealmeida Date: Tue, 29 Sep 2026 11:27:50 -0400 Subject: [PATCH 16/17] Reject ambiguous reaggregation params --- .../datajunction_server/sql/decompose.py | 12 ++++++------ .../tests/sql/decompose_test.py | 18 ++++++++++-------- 2 files changed, 16 insertions(+), 14 deletions(-) diff --git a/datajunction-server/datajunction_server/sql/decompose.py b/datajunction-server/datajunction_server/sql/decompose.py index 986026e05..53df8f855 100644 --- a/datajunction-server/datajunction_server/sql/decompose.py +++ b/datajunction-server/datajunction_server/sql/decompose.py @@ -1753,16 +1753,16 @@ def _attach_reaggregate_params( for component in components if component.aggregation is not None and component.merge is not None ] - # Older metrics could declare params on a multi-component expression - # (notably AVG). Those params never described SUM and COUNT separately; - # broadcasting them changes component identity and may select the wrong - # pre-aggregation. Keep such metrics queryable without propagating them. if not configurable: self._raise_unsupported_reaggregate_shape( "parameterized reaggregation requires an aggregating component", ) - if len(configurable) == 1 and len(components) == 1: - configurable[0].params = dict(reaggregate.params or {}) + if len(configurable) != 1 or len(components) != 1: + self._raise_unsupported_reaggregate_shape( + "parameterized reaggregation must resolve to exactly one measure; " + "use a registered reaggregation family for a multi-component metric", + ) + configurable[0].params = dict(reaggregate.params or {}) def _attach_reaggregate_spec( self, diff --git a/datajunction-server/tests/sql/decompose_test.py b/datajunction-server/tests/sql/decompose_test.py index 47c5d922c..7094932ab 100644 --- a/datajunction-server/tests/sql/decompose_test.py +++ b/datajunction-server/tests/sql/decompose_test.py @@ -2854,15 +2854,17 @@ def test_params_reach_components(self): {"compression": 200}, ] - def test_params_without_a_family_leave_multi_component_metric_untouched(self): - """Legacy AVG params cannot be broadcast to its SUM and COUNT.""" + def test_params_without_a_family_reject_multi_component_metric(self): + """Params cannot be assigned ambiguously to AVG's SUM and COUNT.""" extractor = MetricComponentExtractor(1) - components, _ = extractor._extract_base( - parse("SELECT AVG(latency_ms) FROM t"), - self._spec({"compression": 200}), - ) - assert len(components) == 2 - assert all(component.params is None for component in components) + with pytest.raises( + DJInvalidInputException, + match="must resolve to exactly one measure", + ): + extractor._extract_base( + parse("SELECT AVG(latency_ms) FROM t"), + self._spec({"compression": 200}), + ) def test_params_are_copied_not_shared(self): """ From a98c086d7b0e83f97356c269008415bca7a49297 Mon Sep 17 00:00:00 2001 From: Beto Dealmeida Date: Tue, 29 Sep 2026 11:52:23 -0400 Subject: [PATCH 17/17] Validate reaggregation params before persistence --- .../datajunction_server/models/reaggregate.py | 13 +++++- datajunction-server/tests/api/metrics_test.py | 44 +++++++++++++++++++ .../tests/models/deployment_test.py | 13 ++++++ .../tests/models/reaggregate_test.py | 10 +++++ .../tests/sql/decompose_test.py | 4 +- 5 files changed, 80 insertions(+), 4 deletions(-) diff --git a/datajunction-server/datajunction_server/models/reaggregate.py b/datajunction-server/datajunction_server/models/reaggregate.py index 2b6b2cb4f..9bed0e82e 100644 --- a/datajunction-server/datajunction_server/models/reaggregate.py +++ b/datajunction-server/datajunction_server/models/reaggregate.py @@ -1,8 +1,8 @@ """Models for metric reaggregation declarations.""" -from typing import Any +from typing import Any, Self -from pydantic import BaseModel, ConfigDict, Field +from pydantic import BaseModel, ConfigDict, Field, model_validator from datajunction_server.enum import StrEnum @@ -108,6 +108,15 @@ class ReaggregateSpec(BaseModel): # materialization adapter validates the supported keys for its aggregator. params: dict[str, Any] | None = None + @model_validator(mode="after") + def validate_params_function(self) -> Self: + """Require non-empty tuning parameters to name a parameterized family.""" + if self.params and not is_parameterized_reaggregate_function(self.fn): + raise ValueError( + "reaggregate.params requires a parameterized reaggregate.fn", + ) + return self + def dump_reaggregate_spec( spec: ReaggregateSpec | dict | None, diff --git a/datajunction-server/tests/api/metrics_test.py b/datajunction-server/tests/api/metrics_test.py index c5efa7660..679a81dbe 100644 --- a/datajunction-server/tests/api/metrics_test.py +++ b/datajunction-server/tests/api/metrics_test.py @@ -649,6 +649,50 @@ async def test_legacy_reaggregate_shape_does_not_force_major_revision( assert response.json()["current_version"] == original_version +@pytest.mark.asyncio +async def test_metric_reaggregate_params_validate_before_persistence( + client_with_roads: AsyncClient, + session: AsyncSession, +) -> None: + """Invalid params fail create and update before derivation is scheduled.""" + invalid_name = f"default.invalid_params_{uuid4().hex}" + payload = { + "name": invalid_name, + "query": "SELECT AVG(total_repair_cost) FROM default.repair_orders_fact", + "mode": "published", + "reaggregate": {"fn": "avg", "params": {"compression": 200}}, + } + + response = await client_with_roads.post("/nodes/metric/", json=payload) + assert response.status_code == 422, response.json() + assert "requires a parameterized reaggregate.fn" in str(response.json()) + assert await Node.get_by_name(session, invalid_name) is None + + valid_name = f"default.valid_before_invalid_patch_{uuid4().hex}" + response = await client_with_roads.post( + "/nodes/metric/", + json={ + "name": valid_name, + "query": "SELECT AVG(total_repair_cost) FROM default.repair_orders_fact", + "mode": "published", + }, + ) + assert response.status_code in (200, 201), response.json() + original_version = response.json()["current_version"] + + response = await client_with_roads.patch( + f"/nodes/{valid_name}/", + json={"reaggregate": payload["reaggregate"]}, + ) + assert response.status_code == 422, response.json() + assert "requires a parameterized reaggregate.fn" in str(response.json()) + + response = await client_with_roads.get(f"/nodes/{valid_name}/") + assert response.status_code == 200, response.json() + assert response.json()["current_version"] == original_version + assert response.json()["reaggregate"] is None + + @pytest_asyncio.fixture(scope="module") async def module__current_user(module__session: AsyncSession) -> User: """ diff --git a/datajunction-server/tests/models/deployment_test.py b/datajunction-server/tests/models/deployment_test.py index dbf5c9ae3..d3fe55a38 100644 --- a/datajunction-server/tests/models/deployment_test.py +++ b/datajunction-server/tests/models/deployment_test.py @@ -2502,6 +2502,19 @@ def test_required_dimensions_normalization_falls_back_for_invalid_query(): ] +def test_metric_spec_rejects_params_for_nonparameterized_reaggregation(): + """Deployment specs reject invalid tuning parameters during input parsing.""" + with pytest.raises( + ValidationError, + match="requires a parameterized reaggregate.fn", + ): + MetricSpec( + name="orders", + query="SELECT AVG(amount) FROM orders", + reaggregate={"fn": "avg", "params": {"compression": 200}}, + ) + + @pytest.mark.parametrize( "digest", ["a" * 63, "a" * 65, "A" * 64, "g" * 64], diff --git a/datajunction-server/tests/models/reaggregate_test.py b/datajunction-server/tests/models/reaggregate_test.py index 8f5701481..a8c5bcf84 100644 --- a/datajunction-server/tests/models/reaggregate_test.py +++ b/datajunction-server/tests/models/reaggregate_test.py @@ -146,6 +146,16 @@ def test_empty_params_allowed(): assert ReaggregateSpec(params={}).params == {} +@pytest.mark.parametrize("fn", [None, ReaggregationFunction.AVG]) +def test_nonempty_params_require_parameterized_function(fn): + """Tuning parameters are invalid without a family that interprets them.""" + with pytest.raises( + ValidationError, + match="requires a parameterized reaggregate.fn", + ): + ReaggregateSpec(fn=fn, params={"compression": 200}) + + def test_params_accepted_for_parameterized_function(monkeypatch): """ Registering a function as parameterized enables `params` assignment. diff --git a/datajunction-server/tests/sql/decompose_test.py b/datajunction-server/tests/sql/decompose_test.py index 7094932ab..e43dec537 100644 --- a/datajunction-server/tests/sql/decompose_test.py +++ b/datajunction-server/tests/sql/decompose_test.py @@ -2841,8 +2841,8 @@ class TestReaggregateParams: @staticmethod def _spec(params): - """A spec carrying params.""" - return ReaggregateSpec(params=params) + """A parameterized-family spec with no registered implementation.""" + return ReaggregateSpec(fn=ReaggregationFunction.TDIGEST, params=params) def test_params_reach_components(self): extractor = MetricComponentExtractor(1)