diff --git a/datajunction-server/datajunction_server/api/cubes.py b/datajunction-server/datajunction_server/api/cubes.py index 4c1bdf469..3602c6431 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,8 @@ 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 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 04876d498..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 @@ -55,6 +59,7 @@ class ReaggregateSpec: is an open dict, which has no automatic GraphQL mapping. """ + fn: strawberry.auto rules: strawberry.auto params: JSON | None = None @@ -76,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/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..1ea4f8ec2 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 { @@ -727,6 +736,7 @@ type Query { } type ReaggregateSpec { + fn: ReaggregationFunction rules: [DimensionReaggregateRule!]! params: JSON } @@ -740,6 +750,7 @@ enum ReaggregationFunction { FIRST_VALUE MIN MAX + TDIGEST } type SemanticEntity { diff --git a/datajunction-server/datajunction_server/construction/build_v3/builder.py b/datajunction-server/datajunction_server/construction/build_v3/builder.py index e9a6ca1e1..0276b93c2 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,8 @@ 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, + combiner_dialect: Dialect | None = None, ) -> BuildContext: """ Create and initialize a BuildContext with all setup done. @@ -269,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 @@ -291,6 +296,8 @@ 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 {}, lookback_window=lookback_window, @@ -377,6 +384,8 @@ 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, + combiner_dialect: Dialect | None = None, ) -> GeneratedMeasuresSQL: """ Build measures SQL for a set of metrics, dimensions, and filters. @@ -392,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. @@ -420,6 +431,8 @@ async def build_measures_sql( include_temporal_filters=include_temporal_filters, 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 5ca723234..4e135c5d0 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/combiners.py +++ b/datajunction-server/datajunction_server/construction/build_v3/combiners.py @@ -23,6 +23,10 @@ from datajunction_server.construction.build_v3.cte import ( 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 ( get_temporal_partitions, ) @@ -40,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 @@ -523,6 +530,8 @@ async def build_combiner_sql_from_preaggs( dimensions: list[str], filters: list[str] | None = None, dialect=None, + materialization_target: MaterializationTarget | None = None, + combiner_dialect: Dialect | None = None, ) -> tuple[ CombinedGrainGroupResult, list[PreAggSourceInfo], @@ -544,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: @@ -561,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 @@ -661,6 +672,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) @@ -827,6 +839,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. @@ -848,6 +861,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: @@ -862,20 +876,36 @@ 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 or original_gg.component_aliases.get(comp.name) == col.name ): merge_func = comp.merge + merge_args = comp.merge_args + component = comp break if merge_func: # Apply re-aggregation - agg_expr = ast.Function( - name=ast.Name(merge_func), - args=[col_ref], + assert component is not None + agg_expr: ast.Expression = build_merge_call( + merge_func, + merge_args, + col_ref, + ) + 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: @@ -905,7 +935,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 273d3cb4c..1dfa05718 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/decomposition.py +++ b/datajunction-server/datajunction_server/construction/build_v3/decomposition.py @@ -27,10 +27,12 @@ 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.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 -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 @@ -79,6 +81,7 @@ async def decompose_and_group_metrics( base_metric, nodes_cache=ctx.nodes, parent_map=ctx.parent_map, + dialect=ctx.combiner_dialect or ctx.dialect, ) all_decomposed[base_metric.name] = decomposed @@ -98,6 +101,7 @@ async def decompose_and_group_metrics( metric_node, nodes_cache=ctx.nodes, parent_map=ctx.parent_map, + dialect=ctx.combiner_dialect or ctx.dialect, ) all_decomposed[metric_name] = derived_decomposed else: @@ -121,6 +125,7 @@ async def decompose_and_group_metrics( metric_node, nodes_cache=ctx.nodes, parent_map=ctx.parent_map, + dialect=ctx.combiner_dialect or ctx.dialect, ) all_decomposed[metric_node.name] = decomposed @@ -146,6 +151,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 +168,9 @@ 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. This may differ from + the measures SQL dialect when Spark writes a Druid cube. + 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, @@ -206,7 +215,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. @@ -216,11 +228,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 @@ -230,25 +249,72 @@ 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_component_serialize(accumulated, component, materialization_target) + + +def apply_component_serialize( + expr: ast.Expression, + component: MetricComponent, + materialization_target: MaterializationTarget | None, +) -> ast.Expression: + """ + 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 + build, is unaffected. + """ + 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 = cached_parse( + f"SELECT {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: + # 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]) def get_base_metrics_for_derived(ctx: BuildContext, metric_node: Node) -> list[Node]: diff --git a/datajunction-server/datajunction_server/construction/build_v3/measures.py b/datajunction-server/datajunction_server/construction/build_v3/measures.py index 4e5ad98fb..cc4df71f8 100644 --- a/datajunction-server/datajunction_server/construction/build_v3/measures.py +++ b/datajunction-server/datajunction_server/construction/build_v3/measures.py @@ -27,7 +27,9 @@ ) from datajunction_server.construction.build_v3.decomposition import ( analyze_grain_groups, + apply_component_serialize, build_component_expression, + build_merge_call, merge_grain_groups, ) from datajunction_server.construction.build_v3.dimensions import ( @@ -78,10 +80,12 @@ 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 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__) @@ -157,10 +161,57 @@ 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: + # 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 + + 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 +228,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 +260,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 +1922,15 @@ 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: 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) @@ -1872,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, @@ -2140,7 +2215,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 +2237,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 +2482,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..18748d23f 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,11 @@ 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 alias_registry: AliasRegistry = field(default_factory=AliasRegistry) # Filter classification (populated early in setup) diff --git a/datajunction-server/datajunction_server/internal/cube_materializations.py b/datajunction-server/datajunction_server/internal/cube_materializations.py index f5531ff9d..6839ae7ec 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 @@ -370,7 +373,13 @@ 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 + # (`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/nodes.py b/datajunction-server/datajunction_server/internal/nodes.py index c0599cf2f..5e8f42fe6 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/decompose.py b/datajunction-server/datajunction_server/models/decompose.py index 351b2b527..36db3ec7c 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 @@ -91,6 +92,10 @@ class MetricComponent(BaseModel): {} 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 @@ -109,6 +114,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: diff --git a/datajunction-server/datajunction_server/models/materialization.py b/datajunction-server/datajunction_server/models/materialization.py index c1e8b45a6..e8f410591 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 @@ -32,6 +33,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", @@ -159,6 +176,29 @@ 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, + ) + 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/models/reaggregate.py b/datajunction-server/datajunction_server/models/reaggregate.py index 9ff2da446..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 @@ -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,12 +98,25 @@ 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 + @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/datajunction_server/sql/decompose.py b/datajunction-server/datajunction_server/sql/decompose.py index cb831479e..53df8f855 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,12 @@ AggregationRule, 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, + ReaggregationFunction, dimension_reaggregate_rules, is_supported_dimension_reaggregate_function, parse_reaggregate_spec, @@ -102,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): @@ -122,9 +147,30 @@ 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: - """Build the combiner expression from merged metric components.""" + def combine( + self, + components: list[MetricComponent], + func: ast.Function, + dialect: Dialect = Dialect.SPARK, + ) -> ast.Expression: + """ + 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. + """ # ============================================================================= @@ -134,6 +180,49 @@ def combine(self, components: list[MetricComponent]) -> ast.Expression: DECOMPOSITION_REGISTRY: dict[type, type[AggDecomposition] | None] = {} +# 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] + + +FAMILY_DECOMPOSITION_REGISTRY: dict[ + ReaggregationFunction, + FamilyDecompositionRegistration, +] = {} + + +def decomposes_family( + fn: ReaggregationFunction, + *, + aggregate_functions: tuple[type, ...], +): + """ + Register a decomposition for a reaggregation family. + + 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] = FamilyDecompositionRegistration( + decomposition=decomp_class, + aggregate_functions=frozenset(aggregate_functions), + ) + return decomp_class + + return decorator + + def decomposes(func_class: type): """Decorator to register a decomposition class for a function.""" @@ -162,7 +251,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 +268,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 +285,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 +302,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 +324,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 +341,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 +366,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 +403,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 +511,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 +524,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 +537,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 +555,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 +568,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 +581,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 +691,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 +704,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 +743,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 +818,25 @@ 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 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 = 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: return None @@ -935,14 +1126,30 @@ 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. 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, + 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 +1298,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, @@ -1403,9 +1613,36 @@ 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. - 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 +1658,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 @@ -1443,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: @@ -1505,7 +1746,7 @@ 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 @@ -1516,8 +1757,12 @@ def _attach_reaggregate_params( self._raise_unsupported_reaggregate_shape( "parameterized reaggregation requires an aggregating component", ) - for component in configurable: - component.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, @@ -1612,9 +1857,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 @@ -1628,6 +1874,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 @@ -1641,7 +1900,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 @@ -1721,11 +1980,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/tests/api/cubes_build_metrics_spec_test.py b/datajunction-server/tests/api/cubes_build_metrics_spec_test.py index e196fa151..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,50 @@ 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.""" + 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/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/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/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/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", ) diff --git a/datajunction-server/tests/api/metrics_test.py b/datajunction-server/tests/api/metrics_test.py index 90b2bfc0e..679a81dbe 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", @@ -528,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": [ { @@ -540,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": [ { @@ -552,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": [ { @@ -600,6 +611,88 @@ 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.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/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..8cbc1ec26 --- /dev/null +++ b/datajunction-server/tests/construction/build_v3/accumulate_type_test.py @@ -0,0 +1,221 @@ +""" +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 +from datajunction_server.sql.parsing import ast +from datajunction_server.sql.parsing import types as ct + + +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_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 + + 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 + ) + + 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 + + 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( + "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.""" + + 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_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" + + 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/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. 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/merge_args_test.py b/datajunction-server/tests/construction/build_v3/merge_args_test.py new file mode 100644 index 000000000..59e3f93a1 --- /dev/null +++ b/datajunction-server/tests/construction/build_v3/merge_args_test.py @@ -0,0 +1,139 @@ +""" +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)" + + 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.""" + + 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/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/serialize_target_test.py b/datajunction-server/tests/construction/build_v3/serialize_target_test.py new file mode 100644 index 000000000..88caa5e61 --- /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/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..00e3ab45a --- /dev/null +++ b/datajunction-server/tests/construction/build_v3/sketch_family_e2e_test.py @@ -0,0 +1,489 @@ +""" +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.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 +from datajunction_server.sql.decompose import ( + FAMILY_DECOMPOSITION_REGISTRY, + AggDecomposition, + ComponentDef, + decomposes_family, + make_func, +) +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 +from datajunction_server.utils import get_session + +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 # type: ignore[no-redef] +def infer_type( + digest: ct.ColumnType, + compression: ct.ColumnType, +) -> ct.BinaryType: + return ct.BinaryType() + + +@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 # type: ignore[no-redef] +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_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): + """ + 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 + """, + ) 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/models/cube_druid_sketch_spec_test.py b/datajunction-server/tests/models/cube_druid_sketch_spec_test.py new file mode 100644 index 000000000..8d96a267a --- /dev/null +++ b/datajunction-server/tests/models/cube_druid_sketch_spec_test.py @@ -0,0 +1,214 @@ +""" +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 + +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" + + 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/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 8d225f6ca..a8c5bcf84 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,48 @@ 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. + """ + 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. + """ + 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..e43dec537 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, @@ -25,9 +26,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, @@ -2834,11 +2841,8 @@ class TestReaggregateParams: @staticmethod def _spec(params): - """A spec carrying params.""" - return ReaggregateSpec.model_construct( - rules=[], - 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) @@ -2850,29 +2854,31 @@ 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): + """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(c.params == {"compression": 200} for c 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): """ - 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.""" @@ -2894,3 +2900,411 @@ 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, recording the call it + was handed. Tests the pipeline for dialect-specific extraction logic. + """ + + 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)) + + 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 +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. + """ + 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. + """ + 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. + """ + 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. + """ + 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) + + +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 == [] + + +# ============================================================================= +# 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, + aggregate_functions=(dj_functions.ApproxPercentile,), + ) + 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_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. + """ + 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. + """ + 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. + """ + 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 components[0].aggregation is not None + 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 == [] 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 () => {