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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 18 additions & 0 deletions datajunction-server/datajunction_server/api/cubes.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,7 @@
Granularity,
MaterializationJobTypeEnum,
MaterializationStrategy,
MaterializationTarget,
get_druid_aggregator_spec,
)
from datajunction_server.models.metric import TranslatedSQL
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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,
Comment thread
betodealmeida marked this conversation as resolved.
)
except Exception as e: # pragma: no cover
raise DJInvalidInputException( # pragma: no cover
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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_,
Expand All @@ -25,6 +28,7 @@
MetricDirection = strawberry.enum(MetricDirection_)
Aggregability = strawberry.enum(Aggregability_)
ReaggregationFunction = strawberry.enum(ReaggregationFunction_)
MaterializationTarget = strawberry.enum(MaterializationTarget_)


@strawberry.type
Expand Down Expand Up @@ -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

Expand All @@ -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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
11 changes: 11 additions & 0 deletions datajunction-server/datajunction_server/api/graphql/schema.graphql
Original file line number Diff line number Diff line change
Expand Up @@ -288,6 +288,11 @@ type MaterializationPlan {
units: [MaterializationUnit!]!
}

enum MaterializationTarget {
DRUID
ICEBERG
}

type MaterializationUnit {
upstream: VersionedRef!
grainDimensions: [VersionedRef!]!
Expand All @@ -304,6 +309,10 @@ type MetricComponent {
rule: AggregationRule!
grainAlias: String
params: JSON
mergeArgs: [String!]!
serialize: String
serializeTargets: [MaterializationTarget!]!
serializeType: String
}

enum MetricDirection {
Expand Down Expand Up @@ -727,6 +736,7 @@ type Query {
}

type ReaggregateSpec {
fn: ReaggregationFunction
rules: [DimensionReaggregateRule!]!
params: JSON
}
Expand All @@ -740,6 +750,7 @@ enum ReaggregationFunction {
FIRST_VALUE
MIN
MAX
TDIGEST
}

type SemanticEntity {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.
Expand All @@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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.
Expand All @@ -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.
Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand All @@ -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
Expand Down Expand Up @@ -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],
Expand All @@ -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:
Expand All @@ -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
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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.
Expand All @@ -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:
Expand All @@ -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:
Expand Down Expand Up @@ -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
Expand Down
Loading
Loading