diff --git a/crates/types/src/ir/asap.rs b/crates/types/src/ir/asap.rs index faee35c5f..90b239463 100644 --- a/crates/types/src/ir/asap.rs +++ b/crates/types/src/ir/asap.rs @@ -17,11 +17,13 @@ pub const UNIMPLEMENTED_ASAP_OP: &str = "this ASAP operator is reserved: schema, accuracy, timing and export are not implemented"; #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] -pub enum ASAPOp { +// `#[serde(default)]` fields would otherwise make serde require `C: Default`. +#[serde(bound(deserialize = "C: Deserialize<'de>"))] +pub enum ASAPOp> { /// Summary aggregation. Output: grouping columns + one field carrying /// partial summary state per group, typed `family`. SummaryAgg { - child: Rc, + child: C, /// Which summary family realizes this aggregation. Never /// `FieldDataType::Plain`. family: FieldDataType, @@ -29,55 +31,55 @@ pub enum ASAPOp { reduction: Reduction, grouping: GroupingStrategy, #[serde(default)] - filter: Option, + filter: Option>, }, /// Read out a query result from built summary state. Output is a /// row-shaped schema. SummaryEstimate { - summary_input: Rc, + summary_input: C, query: SketchStatistic, }, /// Read an exact accumulator's state as its finalized value: the /// maintenance-to-read boundary before query-time operators. FinalizeExactAccumulator { - child: Rc, + child: C, }, /// Maintain the full declared population, including membership changes. MaintainPopulation { - child: Rc, + child: C, population: MaintainedPopulation, }, /// Read an aggregate or TopK prefix from the maintained population. EvaluatePopulation { - child: Rc, + child: C, evaluation: PopulationStatistic, }, // ── Reserved: migrated but unimplemented (§1.3 of the proposal) ── SummaryMerge { - children: Vec>, + children: Vec, }, SummarySubtract { - left: Rc, - right: Rc, + left: C, + right: C, }, SummaryDelete { - summary_input: Rc, + summary_input: C, key: ColumnId, }, SummaryJoin { - outer: Rc, - inner: Rc, + outer: C, + inner: C, key: ColumnId, family: FieldDataType, }, Extension { - child: Rc, + child: C, name: String, }, } -impl ASAPOp { - pub fn children(&self) -> Vec<&Rc> { +impl ASAPOp { + pub fn children(&self) -> Vec<&C> { use ASAPOp::*; match self { SummaryAgg { child, filter, .. } => { @@ -100,7 +102,9 @@ impl ASAPOp { } } - pub fn map_children(&self, mut f: impl FnMut(&Rc) -> Rc) -> Self { + /// `f` may change the reference type, e.g. from `Rc` to a + /// node id. + pub fn map_children(&self, mut f: impl FnMut(&C) -> D) -> ASAPOp { use ASAPOp::*; match self { SummaryAgg { @@ -180,7 +184,9 @@ impl ASAPOp { Extension { .. } => "Extension", } } +} +impl ASAPOp { /// Reserved variants that are migrated but not implemented. pub fn is_unimplemented(&self) -> bool { use ASAPOp::*; diff --git a/crates/types/src/ir/canonicalize.rs b/crates/types/src/ir/canonicalize.rs new file mode 100644 index 000000000..797cf61bc --- /dev/null +++ b/crates/types/src/ir/canonicalize.rs @@ -0,0 +1,1039 @@ +//! Post-lowering canonicalization of the operator DAG. +//! +//! Rewrites equivalent spellings of a query into one shape, so that common +//! sub-DAG sharing ([`super::cse`]) can match them by structural equality, +//! whichever front end produced them. +//! +//! ## Heavy-hitter promotion +//! +//! An additive-ranked "order by the aggregate, take the top k" is a +//! heavy-hitter represented by [`AggIntent::TopK`]. Front ends may emit it as +//! an ordinary `Limit { Sort { … Aggregate } }`; this pass promotes that shape +//! to the canonical +//! +//! ```text +//! Aggregate { reduction: Reduce(), measures: [TopK{k}], +//! child: Aggregate { measures: [Count | Sum], … } } +//! ``` +//! +//! Count supplies unit weights and Sum supplies value weights. Because the +//! match is positional, aliases do not affect it. Other ranked expressions +//! retain Sort + Limit. +//! +//! ## Subquery lowering +//! +//! EXISTS/NOT EXISTS and positive IN filter conjuncts may use semi/anti joins. +//! Scalar subqueries remain explicit: a cross join does not preserve their +//! zero-row NULL or multiple-row error semantics. All scalar plan references +//! participate in DAG traversal and canonicalization. + +use std::collections::HashMap; +use std::rc::Rc; + +use super::node::{Operator, OperatorNode}; +use super::non_asap::NonASAPOp; +use super::scalar::{ExprSemantics, Predicate, ProjectItem, ScalarExpr, SortKey}; +use crate::ir::operator_properties::{JoinKind, Reduction}; +use crate::ir::SchemaDerivationError; +use crate::pre_asap::agg_intent::{topk, AggIntent}; +use crate::pre_asap::expr_ir::{CompareOpKind, ScalarValue}; +use crate::types::AccuracyTarget; + +/// Rewrite the DAG under `root` into its canonical form (bottom-up). +/// Idempotent: an already-canonical DAG comes back as the same `Rc`. Only +/// nodes that change (or whose inputs change) are rebuilt; every untouched +/// sub-DAG keeps its pointer identity, and a shared sub-DAG that is rewritten +/// stays shared. +pub fn canonicalize(root: Rc) -> Result, SchemaDerivationError> { + canon(&root, &mut Memo::new()) +} + +/// Already-canonicalized nodes, by pointer. The recursion needs this table so +/// a shared sub-DAG is rewritten once and stays shared; `canonicalize` hides +/// it from callers. +type Memo = HashMap<*const OperatorNode, Rc>; + +fn canon( + node: &Rc, + memo: &mut Memo, +) -> Result, SchemaDerivationError> { + if let Some(done) = memo.get(&Rc::as_ptr(node)) { + return Ok(Rc::clone(done)); + } + + // A `Concat` asserting a caller-proven `discriminator_unique_key` (issue + // #228) had that key's `ColumnId`s resolved against exactly the first + // branch's output schema *as it stood before this pass ran*. The rewrites + // below can restructure that branch (anywhere within it) into a shape + // with a different output schema, which would leave those `ColumnId`s + // pointing at the wrong column, or out of bounds. Snapshot the schema the + // key was resolved against before recursing into the children. + let discriminator_branch_schema_before = match &node.operator { + Operator::NonASAP(NonASAPOp::Concat { + children, + discriminator_unique_key: Some(_), + }) => children.first().map(|c| c.schema.clone()), + _ => None, + }; + + // Bottom-up: canonicalize every child (operator inputs and the nodes read + // by scalar expressions) before matching at this node, so an inner + // heavy-hitter is promoted before an enclosing rewrite inspects it. + let mut rebuilt: Vec<(*const OperatorNode, Rc)> = Vec::new(); + let mut changed = false; + for child in node.operator.children() { + let new = canon(child, memo)?; + changed |= !Rc::ptr_eq(&new, child); + rebuilt.push((Rc::as_ptr(child), new)); + } + let mut current = if changed { + let rebuilt_child = |c: &Rc| { + rebuilt + .iter() + .find(|(ptr, _)| *ptr == Rc::as_ptr(c)) + .map_or_else(|| Rc::clone(c), |(_, new)| Rc::clone(new)) + }; + Rc::new(node.map_children(rebuilt_child)?) + } else { + Rc::clone(node) + }; + + // If the first branch's output schema moved out from under the asserted + // key, the key can no longer be trusted — drop it (never re-derive it by + // guessing at name/position). A wrong `unique_keys` claim is a wrong + // query answer, not a missed optimization, so any difference at all + // drops the key. + if let Operator::NonASAP(NonASAPOp::Concat { + children, + discriminator_unique_key: Some(_), + }) = ¤t.operator + { + let after = children.first().map(|c| &c.schema); + if discriminator_branch_schema_before.as_ref() != after { + current = OperatorNode::new_shared(crate::ir::Operator::NonASAP(NonASAPOp::Concat { + children: children.clone(), + discriminator_unique_key: None, + }))?; + } + } + + let current = apply_local_rules(current, memo)?; + + memo.insert(Rc::as_ptr(node), Rc::clone(¤t)); + Ok(current) +} + +/// Apply the local rewrite rules at `node` (whose inputs are already +/// canonical) until none matches. A `Filter` with several subquery +/// conjuncts sheds one per round. Each rule removes one `Limit{Sort}` idiom +/// or one subquery conjunct, so the loop terminates. +fn apply_local_rules( + mut current: Rc, + memo: &mut Memo, +) -> Result, SchemaDerivationError> { + loop { + let next = if let Some(next) = try_promote_additive_top_ranking(¤t)? { + next + } else if let Some(next) = try_lower_subquery_conjunct(¤t, memo)? { + next + } else { + break; + }; + current = next; + } + Ok(current) +} + +/// Recognise an additive-ranked +/// `Limit { Sort { [Project] Aggregate([Count | Sum]) } }` and rewrite it to +/// the canonical heavy-hitter `Aggregate([TopK])` over the explicit inner +/// aggregate. Returns `None` when the shape does not match. +fn try_promote_additive_top_ranking( + node: &OperatorNode, +) -> Result>, SchemaDerivationError> { + // Limit k, no offset (an OFFSET means "not the top k"). + let Some(NonASAPOp::Limit { + n: Some(k), + offset: 0, + partition_by: limit_partition, + child, + }) = node.non_asap() + else { + return Ok(None); + }; + // A single ordering key on a column. + let Some(NonASAPOp::Sort { + keys, + partition_by, + child: sort_child, + }) = child.non_asap() + else { + return Ok(None); + }; + // A per-group `Limit` must agree with its `Sort`'s partition: the + // ranking's partition is what the outer `TopK` groups by. + if !limit_partition.is_empty() && limit_partition != partition_by { + return Ok(None); + } + let [SortKey { + expr: ScalarExpr::Column(sort_col), + ascending, + .. + }] = keys.as_slice() + else { + return Ok(None); + }; + + // The ordered relation is an `Aggregate`, optionally behind a passthrough + // projection (a bare-column SELECT list). Map the sort key through the + // projection to the aggregate's own output column. + let (agg_node, ranked_col) = match sort_child.non_asap() { + Some(NonASAPOp::Project { cols, child, .. }) => { + let Some(ProjectItem { + expr: ScalarExpr::Column(underlying), + .. + }) = cols.get(*sort_col) + else { + return Ok(None); + }; + (child, *underlying) + } + _ => (sort_child, *sort_col), + }; + + // Exactly one aggregate, ranked by *its* output column — the measure sits + // at index `by.len()` (after the group keys). A `PerEntity` reduction has + // no `by` to rank a measure against, so it is a non-match. + let Some(NonASAPOp::Aggregate { + reduction, + measures, + child: aggregate_child, + .. + }) = agg_node.non_asap() + else { + return Ok(None); + }; + let Reduction::Reduce(by) = reduction else { + return Ok(None); + }; + let [ranked_agg] = measures.as_slice() else { + return Ok(None); + }; + if ranked_col != by.len() { + return Ok(None); + } + // The heavy-hitter decision — descending, over a measure with a realised + // heavy-hitter sketch — is the shared rule both front ends consult (issue + // #38). An ascending additive-ranked limit (bottom-k) stays generic. + if !topk::Ranking::from_aggregate(ranked_agg).is_supported(!ascending) { + return Ok(None); + } + // A direct Sum is a stream of additive observation weights. A Sum over a + // derived child such as Rate/Increase still needs exact reset-aware + // values to rerank sketch candidates, and the post-ASAP IR has no + // candidate-sidecar + exact-rerank node, so that shape keeps Sort + Limit. + if matches!(ranked_agg, AggIntent::Sum { .. }) + && matches!( + aggregate_child.non_asap(), + Some(NonASAPOp::Aggregate { .. }) + ) + { + return Ok(None); + } + // Count ranks unit updates; a direct Sum ranks weighted updates. + let accuracy = match ranked_agg { + AggIntent::Count { accuracy } => accuracy.clone(), + AggIntent::Sum { .. } => AccuracyTarget::Exact, + _ => unreachable!("additive ranking gate admitted a non-additive measure"), + }; + + // Outer heavy-hitter `TopK`, grouped by the ranking's partition (empty for + // a global `ORDER BY … LIMIT k`; the `by` labels for a partitioned `topk + // by`), over the unchanged inner additive aggregate. + OperatorNode::new_shared(crate::ir::Operator::NonASAP(NonASAPOp::Aggregate { + reduction: Reduction::by(partition_by.to_vec()), + measures: vec![AggIntent::TopK { k: *k, accuracy }], + output_names: Vec::new(), + filters: vec![], + having: None, + child: Rc::clone(agg_node), + })) + .map(Some) +} + +// ROW_NUMBER filters retain the window output. Eliminating it without a +// consumer-aware rewrite drops a visible column and invalidates outer scopes. + +/// `Predicate(true)`: the unconditional join predicate the SQL front end +/// emits for an uncorrelated `EXISTS` and for a `CROSS JOIN`. +fn always_true() -> Predicate { + Predicate(ScalarExpr::Literal(ScalarValue::Boolean(true))) +} + +/// True for the `Filter` conjuncts that `try_lower_subquery_conjunct` turns +/// into a join: `[NOT] EXISTS (s)` and a positive `x IN (s)`. `NOT IN` is +/// left alone (see the module docs), and so is an `IN` whose probe `x` +/// contains a scalar subquery, which would otherwise be moved into the join +/// predicate. +fn is_join_conjunct(conjunct: &ScalarExpr) -> bool { + match conjunct { + ScalarExpr::Exists { .. } => true, + ScalarExpr::InSubquery { + expr, + negated: false, + .. + } => !contains_scalar_subquery(expr), + _ => false, + } +} + +/// Lower one `[NOT] EXISTS (s)` / `x IN (s)` conjunct of a `Filter` to the +/// semi-/anti-join the SQL front end used to emit directly. The remaining +/// conjuncts stay in an outer `Filter` over the join: a semi/anti join's +/// output schema is the left's, so their column ids are unchanged. One +/// conjunct per call; the fixpoint loop picks up the next. +fn try_lower_subquery_conjunct( + node: &OperatorNode, + memo: &mut Memo, +) -> Result>, SchemaDerivationError> { + let Some(NonASAPOp::Filter { + pred: Predicate(pred), + child, + }) = node.non_asap() + else { + return Ok(None); + }; + let conjuncts = pred.conjuncts(); + let Some(idx) = conjuncts.iter().position(is_join_conjunct) else { + return Ok(None); + }; + let left_width = child.schema.fields.len(); + let (kind, subquery, join_pred) = match &conjuncts[idx] { + // Uncorrelated by construction (the IR's `Exists` carries no outer + // column references), so the join condition is unconditionally true. + ScalarExpr::Exists { subquery, negated } => { + let kind = if *negated { + JoinKind::Anti + } else { + JoinKind::Semi + }; + (kind, subquery, always_true()) + } + // `x = `, which sits right after the + // left's columns in the `left ++ right` scope the predicate resolves + // against. + ScalarExpr::InSubquery { expr, subquery, .. } => ( + JoinKind::Semi, + subquery, + Predicate(ScalarExpr::Compare { + left: expr.clone(), + op: CompareOpKind::Eq, + right: Box::new(ScalarExpr::Column(left_width)), + semantics: ExprSemantics::Sql, + }), + ), + _ => unreachable!("`is_join_conjunct` admitted a non-subquery conjunct"), + }; + let join = OperatorNode::new_shared(crate::ir::Operator::NonASAP(NonASAPOp::Join { + kind, + pred: join_pred, + left: Rc::clone(child), + right: canon(subquery, memo)?, + }))?; + let mut rest: Vec = conjuncts + .iter() + .enumerate() + .filter(|(i, _)| *i != idx) + .map(|(_, c)| c.clone()) + .collect(); + let rest = match rest.len() { + 0 => return Ok(Some(join)), + 1 => rest.remove(0), + _ => ScalarExpr::BoolAnd(rest), + }; + let out = OperatorNode::new_shared(crate::ir::Operator::NonASAP(NonASAPOp::Filter { + pred: Predicate(rest), + child: join, + }))?; + Ok(Some(out)) +} + +/// Whether `expr` reads a `ScalarSubquery` anywhere in its scalar tree. +fn contains_scalar_subquery(expr: &ScalarExpr) -> bool { + matches!(expr, ScalarExpr::ScalarSubquery(_)) + || expr.children().into_iter().any(contains_scalar_subquery) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::ir::operator_properties::WindowFuncKind; + use crate::ir::operator_properties::{ + ConcatDiscriminatorKey, GroupKeys, Source, WindowFrame, WindowFrameBound, + WindowFrameOffset, WindowFrameUnits, + }; + use crate::pre_asap::schema::{DataType, Field, Schema}; + + fn node(op: NonASAPOp) -> Rc { + Rc::new(OperatorNode::new(Operator::NonASAP(op)).expect("fixture derives a schema")) + } + + fn scan() -> Rc { + node(NonASAPOp::Scan { + source: Source::TimeSeries { metric: "m".into() }, + predicates: vec![], + schema: Schema::with_time_index( + vec![ + Field::plain("ts", DataType::Timestamp, false), + Field::plain("service", DataType::Utf8, false), + Field::plain("value", DataType::Float64, false), + ], + 0, + vec![], + ), + }) + } + + fn aggregate( + reduction: Reduction, + agg: AggIntent, + child: Rc, + ) -> Rc { + node(NonASAPOp::Aggregate { + reduction, + measures: vec![agg], + output_names: vec![], + filters: vec![], + having: None, + child, + }) + } + + fn count() -> AggIntent { + AggIntent::Count { + accuracy: AccuracyTarget::Exact, + } + } + + /// `Aggregate{ by: [1], [Count] }` over the scan — output cols `[service, count]`. + fn count_by_service() -> Rc { + aggregate(Reduction::by(vec![1]), count(), scan()) + } + + fn key(col: usize, ascending: bool) -> Vec { + vec![SortKey { + expr: ScalarExpr::Column(col), + ascending, + nulls_first: false, + }] + } + + fn desc(col: usize) -> Vec { + key(col, false) + } + + fn limit(n: usize, offset: usize, child: Rc) -> Rc { + node(NonASAPOp::Limit { + n: Some(n), + offset, + partition_by: GroupKeys::none(), + child, + }) + } + + fn sort(keys: Vec, child: Rc) -> Rc { + node(NonASAPOp::Sort { + keys, + partition_by: GroupKeys::none(), + child, + }) + } + + fn passthrough_project(child: Rc) -> Rc { + node(NonASAPOp::Project { + cols: vec![ + ProjectItem { + alias: None, + expr: ScalarExpr::Column(0), + }, + ProjectItem { + alias: Some("c".into()), + expr: ScalarExpr::Column(1), + }, + ], + qualifier: None, + child, + }) + } + + fn concat( + children: Vec>, + key: Option, + ) -> Rc { + node(NonASAPOp::Concat { + children, + discriminator_unique_key: key, + }) + } + + fn measures(n: &OperatorNode) -> &[AggIntent] { + match n.non_asap() { + Some(NonASAPOp::Aggregate { measures, .. }) => measures, + _ => &[], + } + } + + fn is_topk_over_count(n: &OperatorNode) -> bool { + let Some(NonASAPOp::Aggregate { + measures, child, .. + }) = n.non_asap() + else { + return false; + }; + matches!(measures.as_slice(), [AggIntent::TopK { k: 5, .. }]) + && matches!(self::measures(child), [AggIntent::Count { .. }]) + } + + #[test] + fn promotes_count_ranked_limit_sort() { + // Limit 5 { Sort DESC by count-col (1) { Aggregate[Count] by [1] } }. + let q = limit(5, 0, sort(desc(1), count_by_service())); + assert!(is_topk_over_count(&canonicalize(q).unwrap())); + } + + #[test] + fn promotes_through_a_passthrough_projection() { + // …with a `SELECT service, count` projection between the Sort and the Agg. + let q = limit(5, 0, sort(desc(1), passthrough_project(count_by_service()))); + assert!(is_topk_over_count(&canonicalize(q).unwrap())); + } + + #[test] + fn promoted_topk_reuses_the_inner_aggregate_node() { + // The inner aggregate is untouched, so the rewrite shares it rather + // than copying it. + let agg = count_by_service(); + let out = canonicalize(limit(5, 0, sort(desc(1), Rc::clone(&agg)))).unwrap(); + let Some(NonASAPOp::Aggregate { child, .. }) = out.non_asap() else { + panic!("expected TopK aggregate"); + }; + assert!(Rc::ptr_eq(child, &agg)); + } + + #[test] + fn is_idempotent() { + let q = limit(5, 0, sort(desc(1), count_by_service())); + let once = canonicalize(q).unwrap(); + let twice = canonicalize(Rc::clone(&once)).unwrap(); + assert!(Rc::ptr_eq(&once, &twice), "canonicalize must be idempotent"); + } + + #[test] + fn untouched_dag_is_returned_pointer_equal() { + // Nothing here matches a rewrite: a Concat of two projections over + // one shared aggregate. The root (and everything under it) must come + // back as the same `Rc`. + let agg = count_by_service(); + let q = concat( + vec![ + passthrough_project(Rc::clone(&agg)), + passthrough_project(Rc::clone(&agg)), + ], + None, + ); + let out = canonicalize(Rc::clone(&q)).unwrap(); + assert!(Rc::ptr_eq(&out, &q)); + } + + #[test] + fn rewritten_shared_subtree_stays_shared() { + // One promotable sub-DAG referenced twice is rewritten once. + let branch = limit(5, 0, sort(desc(1), count_by_service())); + let q = concat(vec![Rc::clone(&branch), Rc::clone(&branch)], None); + let out = canonicalize(q).unwrap(); + let Some(NonASAPOp::Concat { children, .. }) = out.non_asap() else { + panic!("expected Concat"); + }; + assert!(is_topk_over_count(&children[0])); + assert!(Rc::ptr_eq(&children[0], &children[1])); + } + + // ── Concat's discriminator_unique_key vs. canonicalize (issue #228) ── + // + // `discriminator_unique_key`'s `ColumnId`s were resolved against the + // first branch's *pre-canonicalize* output schema. The key is dropped + // whenever that branch's schema actually changed, and survives untouched + // otherwise. Never guessed at. + + fn discriminator_key(n: &OperatorNode) -> &Option { + match n.non_asap() { + Some(NonASAPOp::Concat { + discriminator_unique_key, + .. + }) => discriminator_unique_key, + _ => panic!("expected Concat"), + } + } + + #[test] + fn concat_discriminator_key_survives_canonicalize_when_first_branch_is_unaffected() { + // A plain `Aggregate` first branch matches neither rewrite trigger, + // so its schema is identical before and after canonicalize. + let q = concat( + vec![count_by_service(), count_by_service()], + Some(ConcatDiscriminatorKey::new(0, vec![1])), + ); + let out = canonicalize(Rc::clone(&q)).unwrap(); + assert!( + discriminator_key(&out).is_some(), + "an untouched first branch's discriminator key must survive canonicalize" + ); + assert!(Rc::ptr_eq(&out, &q)); + } + + #[test] + fn concat_discriminator_key_is_dropped_when_first_branch_gets_rewritten() { + // The first branch is exactly the heavy-hitter promotion trigger, so + // canonicalize rewrites it to `Aggregate{TopK}`, whose own output is + // a single column, not the original two (`[service, count]`). A key + // resolved against the 2-column shape must not survive pointing at + // the new 1-column schema. + let promotable_branch = limit(5, 0, sort(desc(1), count_by_service())); + let q = concat( + vec![promotable_branch, count_by_service()], + Some(ConcatDiscriminatorKey::new(0, vec![1])), + ); + let out = canonicalize(q).unwrap(); + let Some(NonASAPOp::Concat { + children, + discriminator_unique_key, + }) = out.non_asap() + else { + panic!("expected Concat"); + }; + assert!( + is_topk_over_count(&children[0]), + "the first branch is still promoted normally" + ); + assert!( + discriminator_unique_key.is_none(), + "a stale discriminator key must be dropped, never silently kept wrong" + ); + assert!( + out.schema.unique_keys.is_empty(), + "the dropped key leaves the schema" + ); + } + + #[test] + fn does_not_promote_ascending_sort() { + // Ascending = bottom-k: the Top-K ranking rule rejects it (needs + // descending), so it stays a generic Sort+Limit (issue #38). + let q = limit(5, 0, sort(key(1, true), count_by_service())); + assert!(!is_topk_over_count(&canonicalize(q).unwrap())); + } + + #[test] + fn does_not_promote_with_offset() { + let q = limit(5, 2, sort(desc(1), count_by_service())); + assert!(!is_topk_over_count(&canonicalize(q).unwrap())); + } + + #[test] + fn does_not_promote_ranking_by_a_group_key() { + // DESC by col 0 (the `service` group key), not the count → not a + // frequency heavy-hitter. + let q = limit(5, 0, sort(desc(0), count_by_service())); + assert!(!is_topk_over_count(&canonicalize(q).unwrap())); + } + + #[test] + fn does_not_promote_when_limit_partition_disagrees_with_sort() { + // A per-group Limit partitioned differently from its Sort is not the + // top-k shape. + let q = node(NonASAPOp::Limit { + n: Some(5), + offset: 0, + partition_by: GroupKeys::by(vec![0]), + child: sort(desc(1), count_by_service()), + }); + assert!(!is_topk_over_count(&canonicalize(q).unwrap())); + } + + #[test] + fn promotes_sum_ranked_limit_sort_as_weighted_heavy_hitter() { + let sum = aggregate(Reduction::by(vec![1]), AggIntent::Sum { col: None }, scan()); + let out = canonicalize(limit(5, 0, sort(desc(1), sum))).unwrap(); + let Some(NonASAPOp::Aggregate { + measures, child, .. + }) = out.non_asap() + else { + panic!("expected weighted TopK aggregate"); + }; + assert!(matches!( + measures.as_slice(), + [AggIntent::TopK { k: 5, .. }] + )); + assert!(matches!(self::measures(child), [AggIntent::Sum { .. }])); + } + + #[test] + fn keeps_sum_over_counter_reduction_as_exact_value_ranking() { + for counter in [AggIntent::Rate, AggIntent::Increase] { + let derived = aggregate(Reduction::PerEntity, counter, scan()); + let sum = aggregate( + Reduction::by(vec![1]), + AggIntent::Sum { col: None }, + derived, + ); + let out = canonicalize(limit(5, 0, sort(desc(1), sum))).unwrap(); + let Some(NonASAPOp::Limit { child, .. }) = out.non_asap() else { + panic!("expected Limit, got {out:?}"); + }; + let Some(NonASAPOp::Sort { child, .. }) = child.non_asap() else { + panic!("expected Sort under the Limit"); + }; + let Some(NonASAPOp::Aggregate { + measures, child, .. + }) = child.non_asap() + else { + panic!("expected Aggregate under the Sort"); + }; + assert!(matches!(measures.as_slice(), [AggIntent::Sum { .. }])); + assert!(matches!( + child.non_asap(), + Some(NonASAPOp::Aggregate { .. }) + )); + } + } + + // ── ROW_NUMBER() partitioned top-k (issue #24) ────────────────────────── + + /// A scan with `[ts, service, region, value]`. + fn scan4() -> Rc { + node(NonASAPOp::Scan { + source: Source::TimeSeries { metric: "m".into() }, + predicates: vec![], + schema: Schema::with_time_index( + vec![ + Field::plain("ts", DataType::Timestamp, false), + Field::plain("service", DataType::Utf8, false), + Field::plain("region", DataType::Utf8, false), + Field::plain("value", DataType::Float64, false), + ], + 0, + vec![], + ), + }) + } + + /// `Aggregate{ by: [1,2] (service, region), [agg] }` — output `[service, + /// region, ]` (3 cols), so a ROW_NUMBER over it appends `rn` at index 3. + fn grouped(agg: AggIntent) -> Rc { + aggregate(Reduction::by(vec![1, 2]), agg, scan4()) + } + + /// `ROW_NUMBER` ignores its frame clause; any concrete frame works. + fn rownumber_frame() -> WindowFrame { + WindowFrame { + units: WindowFrameUnits::Rows, + start_bound: WindowFrameBound::Preceding(WindowFrameOffset::Scalar(ScalarValue::Null)), + end_bound: WindowFrameBound::Following(WindowFrameOffset::Scalar(ScalarValue::Null)), + } + } + + /// `SQLWindowFunc{ RowNumber, PARTITION BY region(2), ORDER BY col(2) DESC } { agg }`. + fn rownumber_window(agg: Rc) -> Rc { + node(NonASAPOp::SQLWindowFunc { + func: WindowFuncKind::RowNumber, + args: vec![], + partition_by: GroupKeys::by(vec![2]), // region + order_by: vec![SortKey { + expr: ScalarExpr::Column(2), // the aggregate output column + ascending: false, + nulls_first: true, + }], + frame: Some(rownumber_frame()), + output_name: "rn".into(), + child: agg, + }) + } + + /// `Filter{ col <= 5 } { child }`. + fn filter_le_5(col: usize, child: Rc) -> Rc { + node(NonASAPOp::Filter { + pred: Predicate(ScalarExpr::Compare { + left: Box::new(ScalarExpr::Column(col)), + op: CompareOpKind::Le, + right: Box::new(ScalarExpr::Literal(ScalarValue::Int64(5))), + semantics: ExprSemantics::Sql, + }), + child, + }) + } + + /// `Filter{ rn(3) <= 5 } { ROW_NUMBER window { agg } }`. + fn rownumber_topk(agg: Rc) -> Rc { + filter_le_5(3, rownumber_window(agg)) + } + + #[test] + fn rownumber_count_topk_becomes_a_partitioned_heavy_hitter() { + let original = rownumber_topk(grouped(count())); + let out = canonicalize(Rc::clone(&original)).unwrap(); + assert_eq!(out.schema, original.schema); + assert!( + matches!(out.non_asap(),Some(NonASAPOp::Filter { child,.. }) if matches!(child.non_asap(),Some(NonASAPOp::SQLWindowFunc { .. }))) + ); + assert_idempotent(&out); + } + + #[test] + fn rownumber_avg_topk_becomes_a_partitioned_sort_limit() { + let original = rownumber_topk(grouped(AggIntent::Avg { col: None })); + let out = canonicalize(Rc::clone(&original)).unwrap(); + assert_eq!(out.schema, original.schema); + assert!( + matches!(out.non_asap(),Some(NonASAPOp::Filter { child,.. }) if matches!(child.non_asap(),Some(NonASAPOp::SQLWindowFunc { .. }))) + ); + assert_idempotent(&out); + } + + #[test] + fn filter_on_a_non_rownumber_column_is_left_alone() { + // `WHERE service_len <= 5` (col 0, not the rn window column) must not + // be mistaken for a top-k. + let q = filter_le_5(0, rownumber_window(grouped(count()))); + let out = canonicalize(Rc::clone(&q)).unwrap(); + assert!(Rc::ptr_eq(&out, &q), "left as the same Filter"); + } + + // ── Subquery lowering ─────────────────────────────────────────────────── + + fn filter_of(pred: ScalarExpr, child: Rc) -> Rc { + node(NonASAPOp::Filter { + pred: Predicate(pred), + child, + }) + } + + /// `SELECT service FROM scan` — a one-column subquery. + /// Scalar reads retain cardinality/null semantics and a shared producer. + #[test] + fn scalar_subqueries_remain_explicit_and_shared() { + let sub = one_column_subquery(); + let root = node(NonASAPOp::Project { + child: scan(), + qualifier: None, + cols: vec![ + ProjectItem { + alias: Some("a".into()), + expr: ScalarExpr::ScalarSubquery(Rc::clone(&sub)), + }, + ProjectItem { + alias: Some("b".into()), + expr: ScalarExpr::ScalarSubquery(Rc::clone(&sub)), + }, + ], + }); + let out = canonicalize(root).unwrap(); + let NonASAPOp::Project { cols, child, .. } = out.expect_non_asap() else { + panic!() + }; + assert!(matches!(child.expect_non_asap(), NonASAPOp::Scan { .. })); + for col in cols { + assert!(matches!(&col.expr,ScalarExpr::ScalarSubquery(node) if Rc::ptr_eq(node,&sub))); + } + assert!(out.schema.fields.iter().all(|f| f.nullable)); + assert_idempotent(&out); + } + + fn one_column_subquery() -> Rc { + node(NonASAPOp::Scan { + source: Source::Table { + table_ref: "sub".into(), + }, + predicates: vec![], + schema: Schema::new(vec![Field::plain("value", DataType::Float64, false)]), + }) + } + + fn exists(subquery: Rc, negated: bool) -> ScalarExpr { + ScalarExpr::Exists { subquery, negated } + } + + fn in_subquery(expr: ScalarExpr, subquery: Rc, negated: bool) -> ScalarExpr { + ScalarExpr::InSubquery { + expr: Box::new(expr), + subquery, + negated, + } + } + + /// `value(2) > 1`. + fn value_gt_1() -> ScalarExpr { + ScalarExpr::Compare { + left: Box::new(ScalarExpr::Column(2)), + op: CompareOpKind::Gt, + right: Box::new(ScalarExpr::Literal(ScalarValue::Int64(1))), + semantics: ExprSemantics::Sql, + } + } + + fn literal_true() -> ScalarExpr { + ScalarExpr::Literal(ScalarValue::Boolean(true)) + } + + fn join_parts( + n: &OperatorNode, + ) -> (JoinKind, &ScalarExpr, &Rc, &Rc) { + match n.non_asap() { + Some(NonASAPOp::Join { + kind, + pred: Predicate(pred), + left, + right, + }) => (kind.clone(), pred, left, right), + _ => panic!("expected a Join, got {n:?}"), + } + } + + fn assert_idempotent(once: &Rc) { + let twice = canonicalize(Rc::clone(once)).unwrap(); + assert!(Rc::ptr_eq(once, &twice), "canonicalize must be idempotent"); + } + + #[test] + fn exists_filter_becomes_semi_join() { + let (left, sub) = (scan(), one_column_subquery()); + let q = filter_of(exists(Rc::clone(&sub), false), Rc::clone(&left)); + let out = canonicalize(q).unwrap(); + let (kind, pred, l, r) = join_parts(&out); + assert_eq!(kind, JoinKind::Semi); + assert_eq!(*pred, literal_true()); + assert!(Rc::ptr_eq(l, &left) && Rc::ptr_eq(r, &sub)); + assert_eq!( + out.schema.fields, left.schema.fields, + "a semi join outputs the left's columns" + ); + assert_idempotent(&out); + } + + #[test] + fn not_exists_becomes_anti_join() { + let (left, sub) = (scan(), one_column_subquery()); + let q = filter_of(exists(Rc::clone(&sub), true), Rc::clone(&left)); + let out = canonicalize(q).unwrap(); + let (kind, pred, l, r) = join_parts(&out); + assert_eq!(kind, JoinKind::Anti); + assert_eq!(*pred, literal_true()); + assert!(Rc::ptr_eq(l, &left) && Rc::ptr_eq(r, &sub)); + assert_idempotent(&out); + } + + #[test] + fn in_subquery_becomes_semi_join_on_the_subquery_column() { + // `WHERE service IN (SELECT service …)` over a 3-column left: the + // subquery's column is `Column(3)` in the `left ++ right` scope. + let (left, sub) = (scan(), one_column_subquery()); + let q = filter_of( + in_subquery(ScalarExpr::Column(1), Rc::clone(&sub), false), + Rc::clone(&left), + ); + let out = canonicalize(q).unwrap(); + let (kind, pred, l, r) = join_parts(&out); + assert_eq!(kind, JoinKind::Semi); + assert_eq!( + *pred, + ScalarExpr::Compare { + left: Box::new(ScalarExpr::Column(1)), + op: CompareOpKind::Eq, + right: Box::new(ScalarExpr::Column(3)), + semantics: ExprSemantics::Sql, + } + ); + assert!(Rc::ptr_eq(l, &left) && Rc::ptr_eq(r, &sub)); + assert_eq!( + out.schema.fields, left.schema.fields, + "a semi join outputs the left's columns" + ); + assert_idempotent(&out); + } + + #[test] + fn exists_with_other_conjuncts_keeps_an_outer_filter() { + // `WHERE value > 1 AND EXISTS (…)` → Filter{ value > 1 }{ Semi }. + let (left, sub) = (scan(), one_column_subquery()); + let q = filter_of( + ScalarExpr::BoolAnd(vec![value_gt_1(), exists(Rc::clone(&sub), false)]), + Rc::clone(&left), + ); + let out = canonicalize(q).unwrap(); + let Some(NonASAPOp::Filter { + pred: Predicate(pred), + child, + }) = out.non_asap() + else { + panic!("expected an outer Filter, got {out:?}"); + }; + assert_eq!(*pred, value_gt_1()); + let (kind, _, l, r) = join_parts(child); + assert_eq!(kind, JoinKind::Semi); + assert!(Rc::ptr_eq(l, &left) && Rc::ptr_eq(r, &sub)); + assert_idempotent(&out); + } + + #[test] + fn two_subquery_conjuncts_become_nested_joins() { + // `WHERE EXISTS (a) AND service NOT EXISTS (b) AND value > 1` sheds + // one conjunct per round: Filter{ value > 1 }{ Anti{ Semi{ l, a }, b } }. + let (left, a, b) = (scan(), one_column_subquery(), one_column_subquery()); + let q = filter_of( + ScalarExpr::BoolAnd(vec![ + exists(Rc::clone(&a), false), + exists(Rc::clone(&b), true), + value_gt_1(), + ]), + Rc::clone(&left), + ); + let out = canonicalize(q).unwrap(); + let Some(NonASAPOp::Filter { + pred: Predicate(pred), + child, + }) = out.non_asap() + else { + panic!("expected an outer Filter, got {out:?}"); + }; + assert_eq!(*pred, value_gt_1()); + let (kind, _, inner, r) = join_parts(child); + assert_eq!(kind, JoinKind::Anti); + assert!(Rc::ptr_eq(r, &b)); + let (kind, _, l, r) = join_parts(inner); + assert_eq!(kind, JoinKind::Semi); + assert!(Rc::ptr_eq(l, &left) && Rc::ptr_eq(r, &a)); + assert_idempotent(&out); + } + + #[test] + fn not_in_subquery_is_left_alone() { + let q = filter_of( + in_subquery(ScalarExpr::Column(1), one_column_subquery(), true), + scan(), + ); + let out = canonicalize(Rc::clone(&q)).unwrap(); + assert!(Rc::ptr_eq(&out, &q), "NOT IN keeps its Filter"); + assert_idempotent(&out); + } + + #[test] + fn lifted_subquery_is_canonicalized() { + // The subquery is itself a promotable heavy-hitter; once lifted into + // the join it is canonical, so a second pass finds nothing to do. + let sub = limit(5, 0, sort(desc(1), count_by_service())); + let q = filter_of(exists(sub, false), scan()); + let out = canonicalize(q).unwrap(); + let (_, _, _, r) = join_parts(&out); + assert!(is_topk_over_count(r)); + assert_idempotent(&out); + } +} diff --git a/crates/types/src/ir/cse.rs b/crates/types/src/ir/cse.rs new file mode 100644 index 000000000..821da2395 --- /dev/null +++ b/crates/types/src/ir/cse.rs @@ -0,0 +1,773 @@ +//! Structural common-subexpression elimination over the unified operator IR: +//! bottom-up hash-consing of [`OperatorNode`] DAGs across a workload's roots. +//! +//! CSE only runs on already-bound, already-canonicalized plans — structural +//! matching is meaningless before canonicalization has converged +//! semantically-equivalent queries onto one shape. [`share_common_sub_dags`] +//! is the single entry point, run once per workload batch (a batch of one +//! still deduplicates a query's own repeated sub-DAGs, see below). +//! +//! ## Algorithm: classic hash-consing / value-numbering +//! +//! Bottom-up: every child is interned before its parent, so two parents whose +//! children were independently deduplicated down to the same `Rc`s are +//! structurally identical iff their own fields also match, without re-walking +//! the sub-DAGs. "Child" means everything [`OperatorNode::children`] returns: +//! the operator inputs *and* the operator nodes a scalar expression reads +//! (`PromqlScalarFromVector`, `ScalarSubquery`, `Exists`, `InSubquery`), so a +//! vector read by `scalar(v)` in two queries is shared like any other input. +//! The scalar expressions themselves stay opaque data on their owning node. +//! +//! ## Correctness: hash is a filter, `PartialEq` is the decision +//! +//! This is the one non-negotiable rule. A **false positive** here — two +//! sub-DAGs wrongly judged shareable — is a wrong query answer, not a missed +//! optimization: two different queries would read each other's data. +//! [`structural_hash`] (SipHash over a canonical serialization, no +//! collision-freedom guarantee) may only narrow the candidate set within one +//! bucket; the typed equality check on that bucket ([`same_node`]) is what +//! actually decides sharing, every time, no exceptions for "the hash probably +//! didn't collide." Equality is intentionally conservative: it recognizes +//! *exact* structural matches only, never "a stricter-accuracy summary could +//! also answer a looser request" (that subsumption question belongs to the +//! ASAP matcher, not here). +//! +//! ## Legality +//! +//! Structural equality is necessary but not sufficient. A non-ASAP node is +//! only ever *returned* as a match for another when its output has a provable +//! unique key (`Schema::has_unique_key()`): a producer's output can only be +//! shared across consumers when its row identity is stable across reads, so +//! an ungrouped aggregate, a `without(..)` grouping, a `Concat`/`SetOp` that +//! drops its keys, … is always inserted fresh even when it is structurally +//! identical to something already interned. An ASAP node (summary state and +//! its evaluations) has no such gate: equal operator, schema and guarantee make +//! it shareable, exactly as post-ASAP sharing decided before this IR. +//! +//! ## Single-query CSE falls out for free +//! +//! A repeated sub-expression within *one* query (the same grouped aggregate on +//! both `BinaryOp` branches) is deduplicated by the same bottom-up interning — +//! a workload of size one still interns bottom-up within that one DAG. + +use std::collections::hash_map::DefaultHasher; +use std::collections::HashMap; +use std::hash::{Hash, Hasher}; +use std::rc::Rc; + +use super::node::{Operator, OperatorNode}; + +/// [`structural_hash`]'s memoization cache: an already-hashed node's `Rc` +/// pointer to its hash. A fresh cache is always correct; what matters is +/// letting it persist across every node of one bottom-up pass rather than +/// starting a new one per call. The caller must keep every cached node alive +/// for the cache's lifetime, or a reused address would alias a stale entry. +pub type HashCache = HashMap<*const OperatorNode, u64>; + +/// The operator with every child (operator inputs and the operator nodes +/// referenced from its scalar expressions alike) replaced by `()`. What +/// remains is the node's own data: variant tag, scalar expressions, +/// parameters. +fn own_fields(node: &OperatorNode) -> Operator<()> { + node.operator.map_children(|_| ()) +} + +/// Coarse structural hash used only to bucket [`InternTable::intern`]'s +/// candidate search — never the sharing decision ([`same_node`] is). +/// +/// `OperatorNode` carries `f64`s (`ScalarValue::Float64`, quantile targets, +/// `ResultGuarantee` bounds, …), so it cannot derive `std::hash::Hash`. The +/// hash is SipHash over two parts: +/// +/// 1. the canonical JSON of [`own_fields`] plus `result_kind`, `schema`, +/// `guarantee` and `timing` — every field `PartialEq` compares except the +/// children. A scalar expression is serialized as data with each operator +/// node it reads replaced by `()`, so a reference to an +/// interned sub-DAG contributes nothing of its own here; +/// 2. for every child in [`OperatorNode::children`] order (operator inputs, +/// then scalar-referenced nodes), the child's own `structural_hash`, +/// memoized in `cache` by `Rc` pointer identity. +/// +/// Part 2 is what makes equal sub-DAGs hash equal whether they are reached +/// through an operator input or through a `scalar(v)`, and what keeps the +/// pass linear: a node is generally a DAG, and re-serializing a shared +/// descendant once per parent would cost `O(sub-DAG)` per node instead of +/// `O(1)` beyond the children's already-known hashes. A non-finite `f64` +/// serializes as `null`, merely widening one (still equality-checked) bucket. +pub fn structural_hash(node: &OperatorNode, cache: &mut HashCache) -> u64 { + fn child_hash(child: &Rc, cache: &mut HashCache) -> u64 { + let ptr = Rc::as_ptr(child); + if let Some(&h) = cache.get(&ptr) { + return h; + } + let h = structural_hash(child, cache); + cache.insert(ptr, h); + h + } + + let mut hasher = DefaultHasher::new(); + let own = ( + own_fields(node), + node.result_kind, + &node.schema, + &node.guarantee, + node.timing, + ); + serde_json::to_string(&own) + .unwrap_or_default() + .hash(&mut hasher); + for child in node.children() { + child_hash(child, cache).hash(&mut hasher); + } + hasher.finish() +} + +/// Numeric `PartialEq` alone conflates signed zeros. The serialized check is +/// additional evidence, never a replacement for typed equality (JSON maps +/// non-finite floats to `null`). Used for the guarantee, whose bounds are +/// floats a shared node must preserve bit-for-bit. +fn same_value(left: &T, right: &T) -> bool { + left == right + && match (serde_json::to_string(left), serde_json::to_string(right)) { + (Ok(left), Ok(right)) => left == right, + _ => false, + } +} + +/// Memo of child-pair comparisons already decided by [`same_node`], keyed by +/// pointer pair. Only interned (table-owned, hence alive) nodes are keys. +type EqMemo = HashMap<(*const OperatorNode, *const OperatorNode), bool>; + +/// The sharing decision: typed equality of two nodes. +/// +/// `OperatorNode`'s derived `PartialEq` would recurse into children by value +/// even when both sides hold the same `Rc` (`OperatorNode` is not `Eq`, so +/// `Rc` gets no pointer shortcut), expanding a shared diamond once per path. +/// Children are therefore compared by pointer first; only when the pointers +/// differ (an equal child that was not legal to share) are the values +/// compared, memoized per pair so a diamond is still walked once. +fn same_node(left: &OperatorNode, right: &OperatorNode, memo: &mut EqMemo) -> bool { + let (lc, rc) = (left.children(), right.children()); + if lc.len() != rc.len() { + return false; + } + let children_equal = lc.iter().zip(&rc).all(|(a, b)| { + if Rc::ptr_eq(a, b) { + return true; + } + let key = (Rc::as_ptr(a), Rc::as_ptr(b)); + if let Some(&eq) = memo.get(&key) { + return eq; + } + let eq = same_node(a, b, memo); + memo.insert(key, eq); + eq + }); + children_equal + && left.result_kind == right.result_kind + && left.schema == right.schema + && left.timing == right.timing + && same_value(&left.guarantee, &right.guarantee) + && same_value(&own_fields(left), &own_fields(right)) +} + +/// Bottom-up hash-consing table: structurally-equal, sharing-legal nodes +/// collapse onto one `Rc`. +/// +/// `buckets` is keyed by [`structural_hash`] — a coarse candidate filter +/// only. Every entry within one bucket is a full node kept around for the +/// [`same_node`] comparison that actually decides a match; a hash collision +/// between structurally different nodes just means a harmless linear scan of +/// a few extra candidates. +struct InternTable { + buckets: HashMap>>, + /// Persisted for the table's whole lifetime so hashing is `O(1)` per node + /// beyond its children; every cached node is owned by `buckets`. + hash_cache: HashCache, + eq_memo: EqMemo, +} + +impl InternTable { + fn new() -> Self { + Self { + buckets: HashMap::new(), + hash_cache: HashMap::new(), + eq_memo: HashMap::new(), + } + } + + /// Intern one node whose children are already interned: look it up by + /// [`structural_hash`], confirm with [`same_node`], and — only when + /// sharing is legal (module doc, "Legality") — return the existing `Rc` + /// instead of allocating a new one. + fn intern(&mut self, node: OperatorNode) -> Rc { + let hash = structural_hash(&node, &mut self.hash_cache); + // A node that is not legal to share is never *returned* as a match + // for something else; it still occupies a fresh slot in the bucket + // (harmless: later scans require legality of the new node too). + let reusable = node.is_asap() || node.schema.has_unique_key(); + let bucket = self.buckets.entry(hash).or_default(); + if reusable { + if let Some(existing) = bucket + .iter() + .find(|candidate| same_node(candidate, &node, &mut self.eq_memo)) + { + return Rc::clone(existing); + } + } + let rc = Rc::new(node); + bucket.push(Rc::clone(&rc)); + rc + } +} + +/// Count of *unique* nodes reachable from `root` (pointer identity, +/// following [`OperatorNode::children`]): the real size of the DAG, not a +/// tree-walk count that re-counts a shared descendant once per parent. +pub fn dag_node_count(root: &Rc) -> usize { + OperatorNode::reachable(root).len() +} + +/// Input pointer → (input `Rc`, interned result). The input `Rc` is retained +/// so its address cannot be freed and reused by a fresh allocation while the +/// memo still maps it. +type Visited = HashMap<*const OperatorNode, (Rc, Rc)>; + +/// Intern `node`'s children (recursively), then `node` itself. The rebuilt +/// node keeps `node`'s retained schema, result kind, guarantee and timing: +/// every child is replaced by an equal node, so each derived property stays +/// valid, and the result is `PartialEq`-equal to the input. +fn intern_bottom_up( + table: &mut InternTable, + visited: &mut Visited, + node: &Rc, +) -> Rc { + if let Some((_, interned)) = visited.get(&Rc::as_ptr(node)) { + return Rc::clone(interned); + } + let operator = node + .operator + .map_children(|child| intern_bottom_up(table, visited, child)); + let rebuilt = OperatorNode { + operator, + result_kind: node.result_kind, + schema: node.schema.clone(), + guarantee: node.guarantee.clone(), + timing: node.timing, + }; + let interned = table.intern(rebuilt); + visited.insert(Rc::as_ptr(node), (Rc::clone(node), Rc::clone(&interned))); + interned +} + +/// Share structurally-identical, sharing-legal sub-DAGs across a workload's +/// roots (or within one root). Every root's *value* is unchanged +/// (`PartialEq`-equal to its input) — only its internal `Rc` structure may +/// now alias another root's, or another part of its own DAG. A node already +/// reached through two paths is visited once. +/// +/// `Id` is caller-chosen — a workload entry's key, an index, a query name. +pub fn share_common_sub_dags( + roots: Vec<(Id, Rc)>, +) -> Vec<(Id, Rc)> { + let mut table = InternTable::new(); + let mut visited = Visited::new(); + roots + .into_iter() + .map(|(id, root)| (id, intern_bottom_up(&mut table, &mut visited, &root))) + .collect() +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::ir::asap::ASAPOp; + use crate::ir::operator_properties::{BinaryOpKind, GroupKeys, Reduction, Source}; + use crate::ir::ScalarExpr; + use crate::ir::{BinaryOperator, NonASAPOp}; + use crate::post_asap::guarantee::ResultGuarantee; + use crate::post_asap::sketch::{ + GroupingStrategy, SketchAlgorithm, SketchKind, SketchParams, SummaryUpdate, + }; + use crate::pre_asap::agg_intent::AggIntent; + use crate::pre_asap::expr_ir::{ColumnRef, CompareOpKind}; + use crate::pre_asap::schema::{DataType, Field, FieldDataType, Schema}; + + use crate::types::AccuracyTarget; + + fn node(op: NonASAPOp) -> Rc { + OperatorNode::new_shared(crate::ir::Operator::NonASAP(op)).unwrap() + } + + /// `[ts, service, value, latency]`, no unique key. + fn scan() -> Rc { + node(NonASAPOp::Scan { + source: Source::TimeSeries { metric: "m".into() }, + predicates: vec![], + schema: Schema::with_time_index( + vec![ + Field::plain("ts", DataType::Timestamp, false), + Field::plain("service", DataType::Utf8, false), + Field::plain("value", DataType::Float64, false), + Field::plain("latency", DataType::Float64, false), + ], + 0, + vec![], + ), + }) + } + + fn quantile_agg(by: Vec, col: Option, q: f64) -> Rc { + node(NonASAPOp::Aggregate { + reduction: Reduction::by(by), + measures: vec![AggIntent::Quantile { + col, + q, + accuracy: AccuracyTarget::Exact, + }], + output_names: vec![], + filters: vec![], + having: None, + child: scan(), + }) + } + + fn compare(lhs: Rc, rhs: Rc) -> Rc { + node(NonASAPOp::BinaryOp { + operator: BinaryOperator { + checked_relative_division: false, + checked_finite_division: false, + kind: BinaryOpKind::Compare(CompareOpKind::Eq), + vector_match: None, + }, + return_bool: false, + lhs, + rhs, + }) + } + + fn two_roots(a: Rc, b: Rc) -> (Rc, Rc) { + let shared = share_common_sub_dags(vec![("a", a), ("b", b)]); + let [(_, ra), (_, rb)] = shared.as_slice() else { + panic!("expected 2 roots"); + }; + (Rc::clone(ra), Rc::clone(rb)) + } + + #[test] + fn distinct_column_quantiles_do_not_merge() { + // Grouped (unique key present) so only the differing `col` blocks it. + let (ra, rb) = two_roots( + quantile_agg(vec![1], Some(2), 0.5), + quantile_agg(vec![1], Some(3), 0.5), + ); + assert!(!Rc::ptr_eq(&ra, &rb)); + assert_ne!(ra, rb); + } + + #[test] + fn no_unique_keys_means_no_merge_even_when_structurally_identical() { + let a = quantile_agg(vec![], Some(2), 0.9); + let b = quantile_agg(vec![], Some(2), 0.9); + assert_eq!(a, b, "fixture sanity: structurally equal"); + assert!( + !a.schema.has_unique_key(), + "fixture sanity: a global aggregate has no provable unique key" + ); + let (ra, rb) = two_roots(a, b); + assert!( + !Rc::ptr_eq(&ra, &rb), + "no unique key ⇒ never hoisted, even for an identical structural match" + ); + } + + #[test] + fn median_and_explicit_half_percentile_merge() { + // Two spellings that lower to the identical grouped `Quantile { q: 0.5 }`. + let (m, p) = two_roots( + quantile_agg(vec![1], Some(2), 0.5), + quantile_agg(vec![1], Some(2), 0.5), + ); + assert!(Rc::ptr_eq(&m, &p)); + } + + #[test] + fn single_query_shares_its_own_repeated_sub_dag() { + // One root with the same grouped aggregate on both branches, built as + // two separately-allocated sub-DAGs (no sharing yet). + let root = compare( + quantile_agg(vec![1], Some(2), 0.5), + quantile_agg(vec![1], Some(2), 0.5), + ); + let shared = share_common_sub_dags(vec![("q", root)]); + let [(_, root)] = shared.as_slice() else { + panic!("expected 1 root"); + }; + let Some(NonASAPOp::BinaryOp { lhs, rhs, .. }) = root.non_asap() else { + panic!("expected BinaryOp root, got {root:?}"); + }; + assert!(Rc::ptr_eq(lhs, rhs)); + } + + #[test] + fn shared_root_value_is_unchanged() { + let a = quantile_agg(vec![1], Some(2), 0.5) + .as_ref() + .clone() + .with_guarantee(Some(ResultGuarantee::exact("fixture"))); + let before = Rc::new(a); + let (ra, _) = two_roots(Rc::clone(&before), Rc::clone(&before)); + assert_eq!(ra.as_ref(), before.as_ref()); + assert!( + ra.guarantee.is_some(), + "retained properties survive the rebuild" + ); + } + + // ── scalar-referenced sub-DAGs ────────────────────────────────────── + + /// `vector(scalar(sum by (service) (up)))`. + fn scalar_of_vector() -> Rc { + let sum_up = node(NonASAPOp::Aggregate { + reduction: Reduction::by(vec![1]), + measures: vec![AggIntent::Sum { col: Some(2) }], + output_names: vec![], + filters: vec![], + having: None, + child: scan(), + }); + assert!(sum_up.schema.has_unique_key(), "fixture sanity"); + node(NonASAPOp::PromqlVectorFromScalar( + ScalarExpr::PromqlScalarFromVector(sum_up), + )) + } + + fn bridged_vector(root: &Rc) -> &Rc { + match root.non_asap() { + Some(NonASAPOp::PromqlVectorFromScalar(ScalarExpr::PromqlScalarFromVector(v))) => v, + other => panic!("expected vector(scalar(v)), got {other:?}"), + } + } + + #[test] + fn scalar_referenced_vector_is_shared_across_queries() { + let (ra, rb) = two_roots(scalar_of_vector(), scalar_of_vector()); + assert!( + Rc::ptr_eq(bridged_vector(&ra), bridged_vector(&rb)), + "the vector read by scalar(v) is a child and must be interned" + ); + assert!( + !Rc::ptr_eq(&ra, &rb), + "the scalar bridge itself has no unique key and stays separate" + ); + } + + #[test] + fn structural_hash_sees_through_a_scalar_reference() { + // Two equal bridges must hash equal whether or not their referenced + // vector is the same Rc — the reference contributes the vector's + // memoized hash, not its identity. + let a = scalar_of_vector(); + let b = scalar_of_vector(); + let mut cache = HashMap::new(); + assert_eq!( + structural_hash(&a, &mut cache), + structural_hash(&b, &mut cache) + ); + assert_eq!( + cache.len(), + 4, + "aggregate + scan cached once per root: {cache:?}" + ); + let other = node(NonASAPOp::PromqlVectorFromScalar( + ScalarExpr::PromqlScalarFromVector(quantile_agg(vec![1], Some(2), 0.5)), + )); + assert_ne!( + structural_hash(&a, &mut cache), + structural_hash(&other, &mut cache) + ); + } + + // ── ASAP nodes ────────────────────────────────────────────────────── + + fn summary_agg(alpha: f64, guarantee: Option) -> Rc { + let family = FieldDataType::Sketch( + SketchKind::new(SketchAlgorithm::DDSketch, SketchParams::DDSketch { alpha }), + GroupingStrategy::default(), + ); + let schema = Schema::lifted(vec![Field::new("state", family.clone(), false)], None); + assert!(!schema.has_unique_key(), "fixture sanity"); + Rc::new( + OperatorNode::with_schema( + Operator::ASAP(ASAPOp::SummaryAgg { + child: scan(), + family, + input: SummaryUpdate::column(ColumnRef::SampleValue), + reduction: Reduction::PerEntity, + grouping: GroupingStrategy::default(), + filter: None, + }), + schema, + ) + .with_guarantee(guarantee), + ) + } + + #[test] + fn asap_nodes_share_without_a_unique_key() { + let exact = || Some(ResultGuarantee::exact("fixture")); + let (ra, rb) = two_roots(summary_agg(0.01, exact()), summary_agg(0.01, exact())); + assert!(Rc::ptr_eq(&ra, &rb)); + assert!(ra.guarantee.is_some()); + } + + #[test] + fn asap_nodes_with_distinct_parameters_or_guarantees_are_not_shared() { + let exact = || Some(ResultGuarantee::exact("fixture")); + let (ra, rb) = two_roots(summary_agg(0.01, exact()), summary_agg(0.001, exact())); + assert!(!Rc::ptr_eq(&ra, &rb), "different sketch parameters"); + let (ra, rb) = two_roots(summary_agg(0.01, exact()), summary_agg(0.01, None)); + assert!( + !Rc::ptr_eq(&ra, &rb), + "an unknown guarantee never borrows an exact one" + ); + assert!(rb.guarantee.is_none()); + } + + #[test] + fn evaluations_share_their_producer_but_not_each_other() { + use crate::post_asap::sketch::SketchStatistic; + let evaluation = |q: f64| { + Rc::new(OperatorNode::with_schema( + Operator::ASAP(ASAPOp::SummaryEstimate { + summary_input: summary_agg(0.01, None), + query: SketchStatistic::Quantile { q }, + }), + Schema::lifted( + vec![Field::plain("quantile", DataType::Float64, false)], + None, + ), + )) + }; + let (p95, p99) = two_roots(evaluation(0.95), evaluation(0.99)); + let producer = |n: &Rc| Rc::clone(n.children()[0]); + assert!(!Rc::ptr_eq(&p95, &p99)); + assert!(Rc::ptr_eq(&producer(&p95), &producer(&p99))); + } + + // ── structural_hash (DAG-aware memoization) ───────────────────────── + + #[test] + fn structural_hash_is_stable_across_cache_states() { + let agg = quantile_agg(vec![1], Some(2), 0.5); + let mut cold = HashMap::new(); + let mut warm = HashMap::new(); + structural_hash(&scan(), &mut warm); + assert_eq!( + structural_hash(&agg, &mut cold), + structural_hash(&agg, &mut warm), + "hash must be independent of unrelated cache state" + ); + } + + #[test] + fn structural_hash_of_an_internally_shared_dag_matches_the_unshared_equivalent() { + let agg = quantile_agg(vec![1], Some(2), 0.5); + let shared_root = compare(Rc::clone(&agg), Rc::clone(&agg)); + let unshared_root = compare( + quantile_agg(vec![1], Some(2), 0.5), + quantile_agg(vec![1], Some(2), 0.5), + ); + assert_eq!( + structural_hash(&shared_root, &mut HashMap::new()), + structural_hash(&unshared_root, &mut HashMap::new()), + ); + } + + #[test] + fn structural_hash_memoizes_a_shared_descendant_exactly_once() { + let agg = quantile_agg(vec![1], Some(2), 0.5); + let root = compare(Rc::clone(&agg), Rc::clone(&agg)); + let mut cache = HashMap::new(); + structural_hash(&root, &mut cache); + assert_eq!( + cache.len(), + 2, + "one entry per unique node in the shared branch (Aggregate + Scan): {cache:?}" + ); + } + + // ── dag_node_count ─────────────────────────────────────────────────── + + #[test] + fn dag_node_count_is_the_naive_count_when_nothing_is_shared() { + assert_eq!(dag_node_count(&scan()), 1); + assert_eq!(dag_node_count(&quantile_agg(vec![1], Some(2), 0.5)), 2); + assert_eq!( + dag_node_count(&scalar_of_vector()), + 3, + "follows scalar references" + ); + } + + #[test] + fn dag_node_count_deduplicates_an_internally_shared_sub_dag() { + let root = compare( + quantile_agg(vec![1], Some(2), 0.5), + quantile_agg(vec![1], Some(2), 0.5), + ); + assert_eq!( + dag_node_count(&root), + 5, + "fixture sanity: nothing shared yet" + ); + let shared = share_common_sub_dags(vec![("q", root)]); + let [(_, root)] = shared.as_slice() else { + panic!("expected 1 root"); + }; + assert_eq!( + dag_node_count(root), + 3, + "BinaryOp + one Aggregate + its Scan" + ); + } + + #[test] + fn dag_node_count_deduplicates_across_two_workload_roots() { + let (ra, rb) = two_roots( + quantile_agg(vec![1], Some(2), 0.5), + quantile_agg(vec![1], Some(2), 0.5), + ); + assert!(Rc::ptr_eq(&ra, &rb), "fixture sanity: the two roots merged"); + assert_eq!(dag_node_count(&ra), 2); + assert_eq!(dag_node_count(&rb), 2); + } + + #[test] + fn dedup_gates_sharing_the_same_as_aggregate() { + // `Dedup { cols }` adds `cols` as a unique key, so two identical + // `Dedup`s merge even though their keyless `Scan`s could not. + let dedup = || { + node(NonASAPOp::Dedup { + cols: vec![1], + child: scan(), + }) + }; + let (ra, rb) = two_roots(dedup(), dedup()); + assert!(Rc::ptr_eq(&ra, &rb)); + } + + #[test] + fn group_keys_gate_still_prevented_when_partition_by_without_used() { + let without_agg = || { + node(NonASAPOp::Aggregate { + reduction: Reduction::Reduce(GroupKeys::without(vec![0])), + measures: vec![AggIntent::Count { + accuracy: AccuracyTarget::Exact, + }], + output_names: vec![], + filters: vec![], + having: None, + child: scan(), + }) + }; + let a = without_agg(); + assert!(!a.schema.has_unique_key()); + let (ra, rb) = two_roots(a, without_agg()); + assert!(!Rc::ptr_eq(&ra, &rb)); + } + + #[test] + fn already_shared_nodes_are_visited_once() { + // A diamond already present in the input stays one node and is not + // re-interned per path. + let agg = quantile_agg(vec![1], Some(2), 0.5); + let root = compare(Rc::clone(&agg), Rc::clone(&agg)); + let shared = share_common_sub_dags(vec![("q", root)]); + let Some(NonASAPOp::BinaryOp { lhs, rhs, .. }) = shared[0].1.non_asap() else { + panic!("expected BinaryOp root"); + }; + assert!(Rc::ptr_eq(lhs, rhs)); + assert_eq!(dag_node_count(&shared[0].1), 3); + } + + // Comparing a shareable node whose equal-but-unshareable children form a + // deep diamond must not expand the diamond once per path. The timeout is + // a coarse runaway guard, not a performance SLA. + #[test] + fn shared_diamond_does_not_expand_during_comparison() { + let (done, completion) = std::sync::mpsc::channel(); + let worker = std::thread::spawn(move || { + fn keyed_diamond() -> Rc { + // BinaryOp over a keyless scan has no unique key at any level, + // so none of the 24 levels is shareable; the `Dedup` on top is. + let mut current = scan(); + for _ in 0..24 { + current = compare(Rc::clone(¤t), current); + } + node(NonASAPOp::Dedup { + cols: vec![1], + child: current, + }) + } + let (ra, rb) = two_roots(keyed_diamond(), keyed_diamond()); + assert!(Rc::ptr_eq(&ra, &rb)); + done.send(()).unwrap(); + }); + completion + .recv_timeout(std::time::Duration::from_secs(5)) + .expect("comparison expanded the shared DAG"); + worker.join().unwrap(); + } + + /// A keyed (hence shareable) projection emitting the literal `value`. + fn keyed_literal(value: f64) -> Rc { + let keyed = node(NonASAPOp::Scan { + source: Source::TimeSeries { metric: "m".into() }, + predicates: vec![], + schema: Schema::with_time_index( + vec![ + Field::plain("ts", DataType::Timestamp, false), + Field::plain("service", DataType::Utf8, false), + ], + 0, + vec![vec![1]], + ), + }); + node(NonASAPOp::Project { + cols: vec![ + crate::ir::ProjectItem { + alias: None, + expr: ScalarExpr::Column(1), + }, + crate::ir::ProjectItem { + alias: Some("v".into()), + expr: ScalarExpr::literal_f64(value), + }, + ], + qualifier: None, + child: keyed, + }) + } + + /// Sharing preserves IEEE signed zero, and JSON's `null` encoding of + /// non-finite floats never becomes the equality decision. + #[test] + fn signed_zero_and_nonfinite_values_remain_distinct() { + assert!( + keyed_literal(0.0).schema.has_unique_key(), + "fixture is shareable" + ); + for (a, b) in [ + (0.0, -0.0), + (-0.0, 0.0), + (f64::INFINITY, f64::NEG_INFINITY), + (f64::NAN, f64::NAN), + ] { + let (ra, rb) = two_roots(keyed_literal(a), keyed_literal(b)); + assert!(!Rc::ptr_eq(&ra, &rb), "{a} and {b} must not be shared"); + } + let (ra, rb) = two_roots(keyed_literal(f64::INFINITY), keyed_literal(f64::INFINITY)); + assert!(Rc::ptr_eq(&ra, &rb)); + } +} diff --git a/crates/types/src/ir/flat.rs b/crates/types/src/ir/flat.rs new file mode 100644 index 000000000..434a3406f --- /dev/null +++ b/crates/types/src/ir/flat.rs @@ -0,0 +1,169 @@ +//! A shared operator DAG as a flat, serializable list of nodes. +//! +//! In memory, a node's children are `Rc`s, and a shared sub-DAG +//! is one `Rc` with several parents. Serializing that tree directly would +//! repeat every shared sub-DAG. [`flatten`] instead gives each distinct node +//! an index and writes its operator with every child, including the nodes +//! read by its scalar expressions, replaced by that index +//! ([`Operator`]). The operator and scalar types are the same ones +//! the planner uses; only the child reference type differs. + +use std::collections::HashMap; +use std::rc::Rc; + +use serde::{Deserialize, Serialize}; + +use super::node::{Operator, OperatorNode, OperatorResultKind}; +use super::query::QueryRoot; +use crate::post_asap::execution_data_state::ExecutionTiming; +use crate::post_asap::guarantee::ResultGuarantee; +use crate::pre_asap::schema::Schema; + +/// Index of a node in [`FlatDag::nodes`]. +pub type NodeId = usize; + +/// An [`OperatorNode`] whose children are [`NodeId`]s. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct FlatNode { + pub operator: Operator, + pub result_kind: OperatorResultKind, + pub schema: Schema, + pub guarantee: Option, + pub timing: Option, +} + +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct FlatDag { + /// Node `i` is `nodes[i]`. Children come before their parents. + pub nodes: Vec, + /// One root per input root, in input order. + pub roots: Vec>, +} + +/// Flatten the DAG reachable from `roots`. Each distinct `Rc` (by pointer) +/// becomes one node, so sharing is kept. Also returns the original node for +/// each id. +pub fn flatten(roots: &[QueryRoot]) -> (FlatDag, Vec>) { + let mut ids = HashMap::new(); + let mut nodes = Vec::new(); + let mut originals = Vec::new(); + let mut id_of = |node: &Rc| visit(node, &mut ids, &mut nodes, &mut originals); + let roots = roots + .iter() + .map(|root| match root { + QueryRoot::Operator(node) => QueryRoot::Operator(id_of(node)), + QueryRoot::Scalar(expr) => QueryRoot::Scalar(expr.map_operator_refs(&mut id_of)), + }) + .collect(); + (FlatDag { nodes, roots }, originals) +} + +fn visit( + node: &Rc, + ids: &mut HashMap<*const OperatorNode, NodeId>, + nodes: &mut Vec, + originals: &mut Vec>, +) -> NodeId { + if let Some(&id) = ids.get(&Rc::as_ptr(node)) { + return id; + } + let operator = node + .operator + .map_children(|child| visit(child, ids, nodes, originals)); + let id = nodes.len(); + nodes.push(FlatNode { + operator, + result_kind: node.result_kind, + schema: node.schema.clone(), + guarantee: node.guarantee.clone(), + timing: node.timing, + }); + originals.push(Rc::clone(node)); + ids.insert(Rc::as_ptr(node), id); + id +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::ir::{NonASAPOp, Predicate, ProjectItem, ScalarExpr}; + use crate::pre_asap::expr_ir::ScalarValue; + use crate::pre_asap::schema::{DataType, Field}; + + fn values() -> Rc { + OperatorNode::new_shared(Operator::NonASAP(NonASAPOp::Values { + rows: vec![vec![ScalarExpr::Literal(ScalarValue::Float64(1.0))]], + schema: Schema::lifted(vec![Field::plain("value", DataType::Float64, false)], None), + })) + .unwrap() + } + + fn filter(child: Rc) -> Rc { + OperatorNode::new_shared(Operator::NonASAP(NonASAPOp::Filter { + pred: Predicate(ScalarExpr::Literal(ScalarValue::Boolean(true))), + child, + })) + .unwrap() + } + + #[test] + fn a_shared_node_is_flattened_once() { + let shared = values(); + let a = filter(Rc::clone(&shared)); + let b = filter(Rc::clone(&shared)); + let (dag, originals) = flatten(&[QueryRoot::Operator(a), QueryRoot::Operator(b)]); + + assert_eq!(dag.nodes.len(), 3); + assert!(Rc::ptr_eq(&originals[0], &shared)); + assert_eq!( + dag.roots, + vec![QueryRoot::Operator(1), QueryRoot::Operator(2)] + ); + for root in [1, 2] { + assert_eq!(dag.nodes[root].operator.children(), vec![&0]); + } + } + + #[test] + fn scalar_references_become_node_ids() { + let v = values(); + let project = OperatorNode::new_shared(Operator::NonASAP(NonASAPOp::Project { + cols: vec![ProjectItem { + alias: Some("s".into()), + expr: ScalarExpr::ScalarSubquery(Rc::clone(&v)), + }], + qualifier: None, + child: Rc::clone(&v), + })) + .unwrap(); + let (dag, _) = flatten(&[QueryRoot::Operator(project)]); + + assert_eq!(dag.nodes.len(), 2); + let Operator::NonASAP(NonASAPOp::Project { cols, child, .. }) = &dag.nodes[1].operator + else { + panic!("expected a Project"); + }; + assert_eq!(*child, 0); + assert_eq!(cols[0].expr, ScalarExpr::ScalarSubquery(0)); + } + + #[test] + fn a_constant_scalar_root_has_no_nodes() { + let (dag, _) = flatten(&[QueryRoot::Scalar(ScalarExpr::literal_f64(42.0))]); + assert!(dag.nodes.is_empty()); + assert_eq!( + dag.roots, + vec![QueryRoot::Scalar(ScalarExpr::Literal( + ScalarValue::Float64(42.0) + ))] + ); + } + + #[test] + fn json_round_trips() { + let (dag, _) = flatten(&[QueryRoot::Operator(filter(values()))]); + assert!(dag.nodes.iter().all(|n| n.timing.is_none())); + let json = serde_json::to_string(&dag).unwrap(); + assert_eq!(serde_json::from_str::(&json).unwrap(), dag); + } +} diff --git a/crates/types/src/ir/mod.rs b/crates/types/src/ir/mod.rs index d20bfd075..e42ad375d 100644 --- a/crates/types/src/ir/mod.rs +++ b/crates/types/src/ir/mod.rs @@ -1,6 +1,6 @@ //! Unified operator and scalar representation from #511. -//! Graph algorithms are added in the next stack layer; legacy consumers -//! remain on their existing representation until the planner cutover. +//! Legacy consumers remain on their existing representation until the +//! planner cutover. pub mod aggregate_schema; pub mod asap; pub mod error; @@ -15,3 +15,7 @@ pub use node::{Operator, OperatorNode, OperatorResultKind}; pub use non_asap::{BinaryOperator, NonASAPOp, TimeRangeKind}; pub use query::QueryRoot; pub use scalar::{ExprSemantics, Predicate, ProjectItem, ScalarExpr, SortKey}; + +pub mod canonicalize; +pub mod cse; +pub mod flat; diff --git a/crates/types/src/ir/node.rs b/crates/types/src/ir/node.rs index d68f5aa03..a2c6c789c 100644 --- a/crates/types/src/ir/node.rs +++ b/crates/types/src/ir/node.rs @@ -30,26 +30,37 @@ pub enum OperatorResultKind { /// The operation a node performs: an ordinary query operator or an ASAP /// summary operator. Either category can consume the other's output. #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] -pub enum Operator { - NonASAP(NonASAPOp), - ASAP(ASAPOp), +pub enum Operator> { + NonASAP(NonASAPOp), + ASAP(ASAPOp), } -impl Operator { - pub fn children(&self) -> Vec<&Rc> { +impl Operator { + pub fn children(&self) -> Vec<&C> { match self { Operator::NonASAP(op) => op.children(), Operator::ASAP(op) => op.children(), } } - pub fn map_children(&self, f: impl FnMut(&Rc) -> Rc) -> Self { + /// `f` may change the reference type, e.g. from `Rc` to a + /// node id. + pub fn map_children(&self, f: impl FnMut(&C) -> D) -> Operator { match self { Operator::NonASAP(op) => Operator::NonASAP(op.map_children(f)), Operator::ASAP(op) => Operator::ASAP(op.map_children(f)), } } + pub fn kind_name(&self) -> &'static str { + match self { + Operator::NonASAP(op) => op.kind_name(), + Operator::ASAP(op) => op.kind_name(), + } + } +} + +impl Operator { pub fn output_schema(&self) -> Result { match self { Operator::NonASAP(op) => op.output_schema(), @@ -70,13 +81,6 @@ impl Operator { Operator::ASAP(op) => op.validate_inputs(), } } - - pub fn kind_name(&self) -> &'static str { - match self { - Operator::NonASAP(op) => op.kind_name(), - Operator::ASAP(op) => op.kind_name(), - } - } } /// A node of the logical DAG. Nodes are immutable and shared through `Rc`; diff --git a/crates/types/src/ir/non_asap.rs b/crates/types/src/ir/non_asap.rs index 0daf3a0ce..e1cd14687 100644 --- a/crates/types/src/ir/non_asap.rs +++ b/crates/types/src/ir/non_asap.rs @@ -52,35 +52,34 @@ pub enum TimeRangeKind { /// ordinary operator can read a summary evaluation and a summary can read any /// relational sub-DAG. #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] -pub enum NonASAPOp { +// `#[serde(default)]` fields would otherwise make serde require `C: Default`. +#[serde(bound(deserialize = "C: Deserialize<'de>"))] +pub enum NonASAPOp> { /// Leaf. `schema` is the binding schema every positional `ColumnId` in /// the tree indexes into; `predicates` are leaf-level row filters (PromQL /// label matchers, pushed-down `WHERE` conjuncts). Scan { source: Source, #[serde(default)] - predicates: Vec, + predicates: Vec>, schema: Schema, }, /// SQL `VALUES` rows, or the one empty row of a `SELECT` without `FROM`. /// Row expressions have no input-column scope. Values { - rows: Vec>, + rows: Vec>>, schema: Schema, }, /// σ — row-level filter. Output schema = child schema. - Filter { - pred: Predicate, - child: Rc, - }, + Filter { pred: Predicate, child: C }, /// π — projection. Project { - cols: Vec, + cols: Vec>, /// Re-qualifies every output column with this table alias (a derived /// table / inline view). `None` for an ordinary SELECT list. #[serde(default)] qualifier: Option, - child: Rc, + child: C, }, /// γ + α — grouping + aggregate intents. Aggregate { @@ -91,41 +90,38 @@ pub enum NonASAPOp { #[serde(default)] output_names: Vec, #[serde(default)] - filters: Vec>, + filters: Vec>>, #[serde(default)] - having: Option, - child: Rc, + having: Option>, + child: C, }, Join { kind: JoinKind, - pred: Predicate, - left: Rc, - right: Rc, + pred: Predicate, + left: C, + right: C, }, SetOp { kind: RelationalSetOpKind, all: bool, - left: Rc, - right: Rc, + left: C, + right: C, }, /// ⊕ — exact n-ary `UNION ALL` of union-compatible branches. The output /// schema is the first child's. Concat { - children: Vec>, + children: Vec, #[serde(default)] discriminator_unique_key: Option, }, /// δ — deduplication; empty `cols` = all columns. - Dedup { - cols: Vec, - child: Rc, - }, + Dedup { cols: Vec, child: C }, /// Order-by, per `partition_by` group when non-empty. Sort { - keys: Vec, + keys: Vec>, #[serde(default)] partition_by: GroupKeys, - child: Rc, + child: C, }, /// Row selection; `n = None` is offset-only. `partition_by` applies the /// limit per group (PromQL `topk by (..)`). @@ -134,7 +130,7 @@ pub enum NonASAPOp { offset: usize, #[serde(default)] partition_by: GroupKeys, - child: Rc, + child: C, }, /// Arithmetic / comparison / set composition of two operands (PromQL /// binary operators). Mixed scalar/vector operations use Project or Filter. @@ -144,68 +140,65 @@ pub enum NonASAPOp { /// filtering. Valid only for comparison operators. #[serde(default)] return_bool: bool, - lhs: Rc, - rhs: Rc, + lhs: C, + rhs: C, }, /// SQL analytic window function. Output schema = child schema + one /// column named `output_name`. SQLWindowFunc { func: WindowFuncKind, - args: Vec, + args: Vec>, partition_by: GroupKeys, - order_by: Vec, + order_by: Vec>, #[serde(default)] frame: Option, output_name: String, - child: Rc, + child: C, }, /// Temporal selection over a time-series input. TimeRange { range: Duration, kind: TimeRangeKind, - child: Rc, + child: C, }, /// PromQL `offset` / `@`: moves when `child` is evaluated. - TimeShift { - shift: TimeShift, - child: Rc, - }, + TimeShift { shift: TimeShift, child: C }, /// PromQL `vector(s)`: a label-less instant vector carrying a scalar. - PromqlVectorFromScalar(ScalarExpr), + PromqlVectorFromScalar(ScalarExpr), /// ρ — PromQL `label_replace` / `label_join`. PromqlRelabel { dst: String, - value: ScalarExpr, - child: Rc, + value: ScalarExpr, + child: C, }, /// PromQL `info(v, selector)` label enrichment. PromqlInfoEnrich { #[serde(default)] selector: Vec, - child: Rc, + child: C, }, /// PromQL `limitk` / `limit_ratio`. PromqlSeriesSample { #[serde(default)] by: GroupKeys, kind: SampleKind, - child: Rc, + child: C, }, /// PromQL subquery `[range:resolution]`. PromqlSubquery { range: Duration, #[serde(default)] resolution: Option, - child: Rc, + child: C, }, } -impl NonASAPOp { +impl NonASAPOp { /// The direct operator inputs, in field order, followed by the operator /// nodes referenced from this operator's scalar expressions. - pub fn children(&self) -> Vec<&Rc> { + pub fn children(&self) -> Vec<&C> { use NonASAPOp::*; - let mut out: Vec<&Rc> = match self { + let mut out: Vec<&C> = match self { Scan { .. } | Values { .. } | PromqlVectorFromScalar(_) => vec![], Filter { child, .. } | Project { child, .. } @@ -231,7 +224,7 @@ impl NonASAPOp { } /// Every scalar expression this operator owns. - pub fn scalar_exprs(&self) -> Vec<&ScalarExpr> { + pub fn scalar_exprs(&self) -> Vec<&ScalarExpr> { use NonASAPOp::*; match self { Scan { predicates, .. } => predicates.iter().map(|p| &p.0).collect(), @@ -268,20 +261,17 @@ impl NonASAPOp { /// Rebuild this operator with `f` applied to every child, including the /// operator nodes referenced from scalar expressions. Every other field - /// is cloned. - pub fn map_children(&self, mut f: impl FnMut(&Rc) -> Rc) -> Self { + /// is cloned. `f` may change the reference type, e.g. from + /// `Rc` to a node id. + pub fn map_children(&self, mut f: impl FnMut(&C) -> D) -> NonASAPOp { use NonASAPOp::*; - let mut map_scalar = |e: &ScalarExpr| e.map_operator_refs(&mut f); - fn map_pred(p: &Predicate, f: &mut impl FnMut(&ScalarExpr) -> ScalarExpr) -> Predicate { - Predicate(f(&p.0)) + fn map_pred(p: &Predicate, f: &mut impl FnMut(&C) -> D) -> Predicate { + Predicate(p.0.map_operator_refs(f)) } - fn map_keys( - keys: &[SortKey], - f: &mut impl FnMut(&ScalarExpr) -> ScalarExpr, - ) -> Vec { + fn map_keys(keys: &[SortKey], f: &mut impl FnMut(&C) -> D) -> Vec> { keys.iter() .map(|k| SortKey { - expr: f(&k.expr), + expr: k.expr.map_operator_refs(f), ascending: k.ascending, nulls_first: k.nulls_first, }) @@ -294,22 +284,19 @@ impl NonASAPOp { schema, } => Scan { source: source.clone(), - predicates: predicates - .iter() - .map(|p| map_pred(p, &mut map_scalar)) - .collect(), + predicates: predicates.iter().map(|p| map_pred(p, &mut f)).collect(), schema: schema.clone(), }, Values { rows, schema } => Values { rows: rows .iter() - .map(|row| row.iter().map(&mut map_scalar).collect()) + .map(|row| row.iter().map(|e| e.map_operator_refs(&mut f)).collect()) .collect(), schema: schema.clone(), }, - PromqlVectorFromScalar(e) => PromqlVectorFromScalar(map_scalar(e)), + PromqlVectorFromScalar(e) => PromqlVectorFromScalar(e.map_operator_refs(&mut f)), Filter { pred, child } => { - let pred = map_pred(pred, &mut map_scalar); + let pred = map_pred(pred, &mut f); Filter { pred, child: f(child), @@ -324,7 +311,7 @@ impl NonASAPOp { .iter() .map(|c| ProjectItem { alias: c.alias.clone(), - expr: map_scalar(&c.expr), + expr: c.expr.map_operator_refs(&mut f), }) .collect(); Project { @@ -343,9 +330,9 @@ impl NonASAPOp { } => { let filters = filters .iter() - .map(|p| p.as_ref().map(|p| map_pred(p, &mut map_scalar))) + .map(|p| p.as_ref().map(|p| map_pred(p, &mut f))) .collect(); - let having = having.as_ref().map(|p| map_pred(p, &mut map_scalar)); + let having = having.as_ref().map(|p| map_pred(p, &mut f)); Aggregate { reduction: reduction.clone(), measures: measures.clone(), @@ -361,7 +348,7 @@ impl NonASAPOp { left, right, } => { - let pred = map_pred(pred, &mut map_scalar); + let pred = map_pred(pred, &mut f); Join { kind: kind.clone(), pred, @@ -396,7 +383,7 @@ impl NonASAPOp { partition_by, child, } => { - let keys = map_keys(keys, &mut map_scalar); + let keys = map_keys(keys, &mut f); Sort { keys, partition_by: partition_by.clone(), @@ -434,8 +421,8 @@ impl NonASAPOp { output_name, child, } => { - let args = args.iter().map(&mut map_scalar).collect(); - let order_by = map_keys(order_by, &mut map_scalar); + let args = args.iter().map(|e| e.map_operator_refs(&mut f)).collect(); + let order_by = map_keys(order_by, &mut f); SQLWindowFunc { func: func.clone(), args, @@ -456,7 +443,7 @@ impl NonASAPOp { child: f(child), }, PromqlRelabel { dst, value, child } => { - let value = map_scalar(value); + let value = value.map_operator_refs(&mut f); PromqlRelabel { dst: dst.clone(), value, @@ -510,7 +497,9 @@ impl NonASAPOp { PromqlSubquery { .. } => "PromqlSubquery", } } +} +impl NonASAPOp { /// Output schema derived from this operator's parameters and its /// children's (already derived) schemas. pub fn output_schema(&self) -> Result { diff --git a/crates/types/src/ir/query.rs b/crates/types/src/ir/query.rs index 682c59465..67ec50359 100644 --- a/crates/types/src/ir/query.rs +++ b/crates/types/src/ir/query.rs @@ -4,9 +4,9 @@ use super::{OperatorNode, ScalarExpr}; use serde::{Deserialize, Serialize}; use std::rc::Rc; #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] -pub enum QueryRoot { - Operator(Rc), - Scalar(ScalarExpr), +pub enum QueryRoot> { + Operator(C), + Scalar(ScalarExpr), } impl From> for QueryRoot { fn from(node: Rc) -> Self { diff --git a/crates/types/src/ir/scalar.rs b/crates/types/src/ir/scalar.rs index c1a0b90a1..3d724b250 100644 --- a/crates/types/src/ir/scalar.rs +++ b/crates/types/src/ir/scalar.rs @@ -31,56 +31,56 @@ pub enum ExprSemantics { /// A scalar expression over the owning operator's input schema. Column /// references are positional [`ColumnId`]s. #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] -pub enum ScalarExpr { +pub enum ScalarExpr> { Column(ColumnId), Literal(ScalarValue), /// Unary minus. Negative { - expr: Box, + expr: Box>, semantics: ExprSemantics, }, Compare { - left: Box, + left: Box>, op: CompareOpKind, - right: Box, + right: Box>, semantics: ExprSemantics, }, /// Flat conjunction (logical AND). An empty list is vacuously true. - BoolAnd(Vec), + BoolAnd(Vec>), /// Flat disjunction (logical OR). An empty list is vacuously false. - BoolOr(Vec), - Not(Box), - IsNull(Box), - IsNotNull(Box), + BoolOr(Vec>), + Not(Box>), + IsNull(Box>), + IsNotNull(Box>), /// `CAST(expr AS to)`; `try_cast` for SQL `TRY_CAST` (NULL on failure). Cast { - expr: Box, + expr: Box>, to: DataType, try_cast: bool, }, /// `expr [NOT] IN (v1, v2, …)`. InList { - expr: Box, - list: Vec, + expr: Box>, + list: Vec>, negated: bool, }, /// Scalar function call, e.g. `LOWER(col)`, `ABS(x)`. FunctionCall { name: String, - args: Vec, + args: Vec>, }, Arithmetic { op: ArithmeticOpKind, - left: Box, - right: Box, + left: Box>, + right: Box>, semantics: ExprSemantics, }, /// SQL `CASE` (both searched and simple forms). `operand` present for the /// simple form (`CASE expr WHEN …`), absent for searched. Case { - operand: Option>, - branches: Vec<(ScalarExpr, ScalarExpr)>, - else_expr: Option>, + operand: Option>>, + branches: Vec<(ScalarExpr, ScalarExpr)>, + else_expr: Option>>, }, /// SQL `NOW()` / `CURRENT_TIMESTAMP`: the statement evaluation time. CurrentTimestamp, @@ -88,37 +88,37 @@ pub enum ScalarExpr { EvalTimestamp, /// PromQL `scalar(v)`: the single sample of an instant vector, NaN /// otherwise. The referenced operator is a real plan dependency. - PromqlScalarFromVector(Rc), + PromqlScalarFromVector(C), /// An uncorrelated SQL scalar subquery: one column; zero rows is NULL, /// more than one row is an error. - ScalarSubquery(Rc), + ScalarSubquery(C), /// SQL `[NOT] EXISTS (subquery)`. Exists { - subquery: Rc, + subquery: C, negated: bool, }, /// SQL `expr [NOT] IN (subquery)` over a one-column relation. InSubquery { - expr: Box, - subquery: Rc, + expr: Box>, + subquery: C, negated: bool, }, } /// A row-level filter predicate (WHERE clause / PromQL label matcher). #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] -pub struct Predicate(pub ScalarExpr); +pub struct Predicate>(pub ScalarExpr); /// One item in a SELECT projection list. #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] -pub struct ProjectItem { +pub struct ProjectItem> { pub alias: Option, - pub expr: ScalarExpr, + pub expr: ScalarExpr, } #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] -pub struct SortKey { - pub expr: ScalarExpr, +pub struct SortKey> { + pub expr: ScalarExpr, pub ascending: bool, pub nulls_first: bool, } @@ -131,9 +131,11 @@ impl ScalarExpr { pub fn column(id: ColumnId) -> Self { ScalarExpr::Column(id) } +} +impl ScalarExpr { /// If this expression is a `BoolAnd`, its elements; otherwise `self` alone. - pub fn conjuncts(&self) -> &[ScalarExpr] { + pub fn conjuncts(&self) -> &[ScalarExpr] { match self { ScalarExpr::BoolAnd(v) => v.as_slice(), _ => std::slice::from_ref(self), @@ -141,7 +143,7 @@ impl ScalarExpr { } /// If this expression is a `BoolOr`, its elements; otherwise `self` alone. - pub fn disjuncts(&self) -> &[ScalarExpr] { + pub fn disjuncts(&self) -> &[ScalarExpr] { match self { ScalarExpr::BoolOr(v) => v.as_slice(), _ => std::slice::from_ref(self), @@ -149,7 +151,7 @@ impl ScalarExpr { } /// The direct scalar sub-expressions. - pub fn children(&self) -> Vec<&ScalarExpr> { + pub fn children(&self) -> Vec<&ScalarExpr> { match self { ScalarExpr::Column(_) | ScalarExpr::Literal(_) @@ -198,13 +200,13 @@ impl ScalarExpr { /// The operator nodes this expression (transitively) reads: the explicit /// plan-reading variants. Every DAG traversal must follow these. - pub fn operator_refs(&self) -> Vec<&Rc> { + pub fn operator_refs(&self) -> Vec<&C> { let mut out = Vec::new(); self.collect_operator_refs(&mut out); out } - fn collect_operator_refs<'a>(&'a self, out: &mut Vec<&'a Rc>) { + fn collect_operator_refs<'a>(&'a self, out: &mut Vec<&'a C>) { match self { ScalarExpr::PromqlScalarFromVector(node) | ScalarExpr::ScalarSubquery(node) => { out.push(node) @@ -219,22 +221,17 @@ impl ScalarExpr { } /// Rebuild this expression with `f` applied to every operator node it - /// reads (recursively through scalar children). - pub fn map_operator_refs( - &self, - f: &mut impl FnMut(&Rc) -> Rc, - ) -> ScalarExpr { - fn map_box) -> Rc>( - e: &ScalarExpr, - f: &mut F, - ) -> Box { + /// reads (recursively through scalar children). `f` may change the + /// reference type, e.g. from `Rc` to a node id. + pub fn map_operator_refs(&self, f: &mut impl FnMut(&C) -> D) -> ScalarExpr { + fn map_box(e: &ScalarExpr, f: &mut impl FnMut(&C) -> D) -> Box> { Box::new(e.map_operator_refs(f)) } match self { - ScalarExpr::Column(_) - | ScalarExpr::Literal(_) - | ScalarExpr::CurrentTimestamp - | ScalarExpr::EvalTimestamp => self.clone(), + ScalarExpr::Column(c) => ScalarExpr::Column(*c), + ScalarExpr::Literal(v) => ScalarExpr::Literal(v.clone()), + ScalarExpr::CurrentTimestamp => ScalarExpr::CurrentTimestamp, + ScalarExpr::EvalTimestamp => ScalarExpr::EvalTimestamp, ScalarExpr::PromqlScalarFromVector(node) => ScalarExpr::PromqlScalarFromVector(f(node)), ScalarExpr::ScalarSubquery(node) => ScalarExpr::ScalarSubquery(f(node)), ScalarExpr::Exists { subquery, negated } => ScalarExpr::Exists { @@ -334,7 +331,9 @@ impl ScalarExpr { child.collect_columns(out); } } +} +impl ScalarExpr { /// Infer the `(DataType, nullable)` this expression produces against the /// input schema its owner evaluates it in. Unregistered functions are /// rejected. A reference to a field carrying summary state is an error: