From 1d9be063f49b22c9ccd7caff6773c50e81ed93cc Mon Sep 17 00:00:00 2001 From: Selvomega Date: Sat, 26 Sep 2026 18:57:49 +0000 Subject: [PATCH] feat: per-measure FILTER predicates on Aggregate (#466) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Neither IR could give one aggregate function its own row predicate, so `count(CASE WHEN p THEN 1 END)` next to a plain `sum(x)` was rejected at lowering. `QueryExpr::Aggregate` and `ExactOperation::Aggregate` gain `filters: Vec>`, parallel to `measures` and positional against the child; `SummaryAgg` and its wire payload gain `filter`, and the executable DAG wire version goes to 6. The SQL front end fills the field from an explicit `FILTER (WHERE …)` (parsed through a GenericDialect wrapper that enables the clause), from `count(CASE WHEN p THEN x END)`, and from `count(expr)` over a nullable `expr`. Resolution, canonicalization, CSE, dependency collection and DAG export carry it. No binding rule applies a filter yet: every recognizer in asap-aware-mapping keeps a filtered Aggregate as `KeepPreAsap`, and the heavy-hitter promotion skips filtered rankings. Co-Authored-By: Claude Fable 5.1 --- .../src/accuracy/composition.rs | 1 + .../src/accuracy/reconciliation.rs | 8 +- crates/asap-aware-mapping/src/cost_model.rs | 2 + .../src/exact_composition.rs | 14 +- crates/asap-aware-mapping/src/explanation.rs | 1 + crates/asap-aware-mapping/src/grouping.rs | 4 + .../src/maintained_population.rs | 7 +- .../src/physical_plan_cost_model.rs | 1 + .../src/query_physical_lowering.rs | 11 +- crates/asap-aware-mapping/src/recurrence.rs | 3 + crates/asap-aware-mapping/src/replacement.rs | 48 ++++- crates/asap-aware-mapping/src/rewrite.rs | 28 ++- crates/asap-aware-mapping/src/rollup.rs | 11 +- .../src/summary_maintenance_cost/model.rs | 3 + .../src/summary_maintenance_lifecycle.rs | 4 + crates/devtools/src/bin/dag_export.rs | 1 + crates/frontend-metricsql/src/lib.rs | 1 + crates/frontend-promql/src/promql.rs | 7 + crates/frontend-sql/src/sql/dialect.rs | 112 +++++++++++ crates/frontend-sql/src/sql/mod.rs | 157 +++++++++++---- crates/frontend-sql/tests/pearson_corr.rs | 22 ++- crates/frontend-sql/tests/sql_lowering.rs | 179 ++++++++++++++++-- crates/integration-tests/tests/aggregate.rs | 1 + crates/integration-tests/tests/binary_op.rs | 3 + .../tests/exact_composition.rs | 3 + crates/integration-tests/tests/nested.rs | 3 + crates/integration-tests/tests/time_range.rs | 1 + crates/sql-function-catalog/src/lib.rs | 13 +- crates/types/src/dag_export.rs | 5 + crates/types/src/post_asap/cse.rs | 12 +- crates/types/src/post_asap/executable_dag.rs | 10 +- .../src/post_asap/execution_data_state.rs | 10 +- crates/types/src/post_asap/expr.rs | 13 ++ crates/types/src/pre_asap/canonicalize.rs | 31 +++ .../types/src/pre_asap/column_resolution.rs | 1 + crates/types/src/pre_asap/cse.rs | 39 +++- crates/types/src/pre_asap/mod.rs | 10 +- crates/types/src/pre_asap/query_expr.rs | 73 +++++++ crates/types/src/pre_asap/resolve.rs | 85 ++++++++- crates/types/src/pre_asap/schema_resolver.rs | 5 + .../architecture/physical-plan-integration.md | 7 + docs/develop_docs/pre-asap-ir.md | 37 +++- 42 files changed, 905 insertions(+), 82 deletions(-) create mode 100644 crates/frontend-sql/src/sql/dialect.rs diff --git a/crates/asap-aware-mapping/src/accuracy/composition.rs b/crates/asap-aware-mapping/src/accuracy/composition.rs index a284ab32..5c1a6e67 100644 --- a/crates/asap-aware-mapping/src/accuracy/composition.rs +++ b/crates/asap-aware-mapping/src/accuracy/composition.rs @@ -1036,6 +1036,7 @@ mod tests { measures: vec![intent], output_names: vec![], having: None, + filters: vec![], }; assert_eq!( DefaultAccuracyModel.exact_operation_rule(&operation(AggIntent::Rate)), diff --git a/crates/asap-aware-mapping/src/accuracy/reconciliation.rs b/crates/asap-aware-mapping/src/accuracy/reconciliation.rs index c5c09abf..d34dd4b3 100644 --- a/crates/asap-aware-mapping/src/accuracy/reconciliation.rs +++ b/crates/asap-aware-mapping/src/accuracy/reconciliation.rs @@ -153,7 +153,7 @@ use std::cmp::Ordering; use std::rc::Rc; use asap_types::pre_asap::agg_intent::AggIntent; -use asap_types::pre_asap::query_expr::{QueryExpr, Reduction}; +use asap_types::pre_asap::query_expr::{any_measure_filtered, QueryExpr, Reduction}; use asap_types::types::AccuracyTarget; use crate::replacement::{ @@ -186,6 +186,7 @@ fn bindable_accuracy_aggregate(node: &QueryExpr) -> Option Option Option<(MaintainedPopulation, PopulationReadou child, reduction: Reduction::Reduce(grouping), measures, + filters, having: None, .. } => { let [intent] = measures.as_slice() else { return None; }; + if any_measure_filtered(filters) { + return None; + } let (col, readout) = match intent { AggIntent::Quantile { q, col, .. } if q.is_finite() => { (*col, PopulationReadout::Quantile { q: *q }) diff --git a/crates/asap-aware-mapping/src/physical_plan_cost_model.rs b/crates/asap-aware-mapping/src/physical_plan_cost_model.rs index a8cb38d7..ffbe08fc 100644 --- a/crates/asap-aware-mapping/src/physical_plan_cost_model.rs +++ b/crates/asap-aware-mapping/src/physical_plan_cost_model.rs @@ -412,6 +412,7 @@ mod tests { accuracy: AccuracyTarget::Epsilon(0.01), }], output_names: vec![], + filters: vec![], having: None, child: Rc::new(QueryExpr::Scan { source: Source::Table { diff --git a/crates/asap-aware-mapping/src/query_physical_lowering.rs b/crates/asap-aware-mapping/src/query_physical_lowering.rs index bd2aa675..8b4ac801 100644 --- a/crates/asap-aware-mapping/src/query_physical_lowering.rs +++ b/crates/asap-aware-mapping/src/query_physical_lowering.rs @@ -283,11 +283,15 @@ pub fn lower_query_physical_dag( QueryExpr::Aggregate { reduction, measures, + filters, having, child, .. } => { - if having.is_some() || measures.is_empty() { + if having.is_some() + || asap_types::pre_asap::any_measure_filtered(filters) + || measures.is_empty() + { return Err(AnalyticalCostError::UnsupportedQueryOperator); } if matches!(reduction, asap_types::pre_asap::Reduction::PerEntity) { @@ -1457,6 +1461,7 @@ mod tests { reduction: Reduction::by(vec![]), measures: vec![AggIntent::PearsonCorr { left: 0, right: 1 }], output_names: vec!["r".into()], + filters: vec![], having: None, child: Rc::new(QueryExpr::Scan { source: source.clone(), @@ -1516,6 +1521,7 @@ mod tests { reduction: Reduction::by(vec![0]), measures: vec![AggIntent::Sum { col: Some(1) }], output_names: vec![], + filters: vec![], having: None, child: Rc::clone(&scan), }); @@ -2251,6 +2257,7 @@ mod tests { accuracy: AccuracyTarget::Exact, }], output_names: vec![], + filters: vec![], having: None, child: scan(), }); @@ -2331,6 +2338,7 @@ mod tests { reduction: Reduction::PerEntity, measures: vec![AggIntent::Absent], output_names: vec![], + filters: vec![], having: None, child: scan, }); @@ -2566,6 +2574,7 @@ mod tests { reduction: Reduction::PerEntity, measures: vec![AggIntent::Sum { col: None }], output_names: vec![], + filters: vec![], having: None, child: sample, }); diff --git a/crates/asap-aware-mapping/src/recurrence.rs b/crates/asap-aware-mapping/src/recurrence.rs index e01c10df..21214f52 100644 --- a/crates/asap-aware-mapping/src/recurrence.rs +++ b/crates/asap-aware-mapping/src/recurrence.rs @@ -820,6 +820,7 @@ mod tests { )), reduction: Reduction::by(vec![]), grouping: GroupingStrategy::default(), + filter: None, }, schema: SummarySchema { fields: vec![SummaryField { @@ -1162,6 +1163,7 @@ mod tests { reduction: QueryReduction::by(vec![2]), measures: vec![AggIntent::Sum { col: Some(1) }], output_names: vec![], + filters: vec![], having: None, child: Rc::new(labeled_scan()), } @@ -1385,6 +1387,7 @@ mod tests { reduction: QueryReduction::by(vec![]), measures: vec![AggIntent::Avg { col: None }], output_names: vec![], + filters: vec![], having: None, child: Rc::new(scan()), }; diff --git a/crates/asap-aware-mapping/src/replacement.rs b/crates/asap-aware-mapping/src/replacement.rs index e2ce61d1..0a940e16 100644 --- a/crates/asap-aware-mapping/src/replacement.rs +++ b/crates/asap-aware-mapping/src/replacement.rs @@ -360,6 +360,7 @@ use asap_types::post_asap::{AccuracyError, CompositionOperator, GuaranteeSource, use asap_types::pre_asap::agg_intent::{agg_is_mergeable, AggIntent}; use asap_types::pre_asap::cse::{share_common_subtrees, structural_hash, HashCache}; use asap_types::pre_asap::expr_ir::{ArithmeticOpKind, ColumnRef}; +use asap_types::pre_asap::query_expr::any_measure_filtered; use asap_types::pre_asap::query_expr::{ BinaryOpKind, Predicate, QueryExpr, QueryExprError, Reduction, }; @@ -1632,12 +1633,16 @@ fn exact_topk_over_temporal_values( reduction, measures, output_names: _, + filters, having: None, child, } = root.as_ref() else { return Ok(None); }; + if any_measure_filtered(filters) { + return Ok(None); + } let [AggIntent::TopK { k, accuracy: AccuracyTarget::Exact, @@ -2168,11 +2173,16 @@ fn keep_pre_asap_rc(expr: Rc) -> Result, RealizationE /// DAG assembly so their independently planned children remain visible. pub fn bindable_intent(node: &QueryExpr) -> Option<&AggIntent> { if let QueryExpr::Aggregate { - measures, having, .. + measures, + filters, + having, + .. } = node { if let ([intent], None) = (measures.as_slice(), having) { - return Some(intent); + if !any_measure_filtered(filters) { + return Some(intent); + } } } None @@ -2815,6 +2825,7 @@ fn construct_summary_agg( input: summary_input, reduction: physical_reduction, grouping: GroupingStrategy::default(), + filter: None, }, schema: state_schema, // Summary *state* carries no caller-visible guarantee; only a @@ -4668,6 +4679,7 @@ impl<'a> GlobalSelection<'a> { reduction, measures, output_names, + filters, having, child, } if query_time_nested_sum(target) => ( @@ -4676,6 +4688,7 @@ impl<'a> GlobalSelection<'a> { reduction: reduction.clone(), measures: measures.clone(), output_names: output_names.clone(), + filters: filters.clone(), having: having.clone(), }), ), @@ -4736,6 +4749,7 @@ impl<'a> GlobalSelection<'a> { fn query_time_nested_sum(target: &QueryExpr) -> bool { let QueryExpr::Aggregate { measures, + filters, having: None, child, .. @@ -4743,7 +4757,9 @@ fn query_time_nested_sum(target: &QueryExpr) -> bool { else { return false; }; - matches!(measures.as_slice(), [AggIntent::Sum { .. }]) && contains_aggregate(child) + !any_measure_filtered(filters) + && matches!(measures.as_slice(), [AggIntent::Sum { .. }]) + && contains_aggregate(child) } fn contains_aggregate(expr: &QueryExpr) -> bool { @@ -4785,6 +4801,7 @@ fn relink_agg_child(node: &Rc, new_child: &Rc) -> Rc { if Rc::ptr_eq(child, new_child) { return Rc::clone(node); @@ -4796,6 +4813,7 @@ fn relink_agg_child(node: &Rc, new_child: &Rc) -> Rc Option<(usize, Option)> { let QueryExpr::Aggregate { reduction, measures, + filters, having: None, child, .. @@ -83,6 +86,9 @@ fn avg_rewrite_target(node: &QueryExpr) -> Option<(usize, Option)> { else { return None; }; + if any_measure_filtered(filters) { + return None; + } let Reduction::Reduce(by) = reduction else { return None; }; @@ -133,6 +139,7 @@ pub(crate) fn temporal_average_components(root: &Rc) -> Option) -> Option) -> Option) -> Option> { reduction: reduction.clone(), measures: vec![AggIntent::Sum { col }], output_names: Vec::new(), + filters: vec![], having: None, child: Rc::clone(child), }); @@ -212,6 +224,7 @@ fn build_rewrite(root: &Rc) -> Option> { accuracy: AccuracyTarget::Exact, }], output_names: Vec::new(), + filters: vec![], having: None, child: Rc::clone(child), }); @@ -253,6 +266,7 @@ fn composed_aggregate_rewrite(root: &Rc) -> Option> { reduction: outer_reduction @ Reduction::Reduce(_), measures: outer_measures, output_names, + filters: outer_filters, having: None, child, } = root.as_ref() @@ -262,6 +276,7 @@ fn composed_aggregate_rewrite(root: &Rc) -> Option> { let QueryExpr::Aggregate { reduction: Reduction::PerEntity, measures: inner_measures, + filters: inner_filters, having: None, child: inner_child, .. @@ -269,6 +284,9 @@ fn composed_aggregate_rewrite(root: &Rc) -> Option> { else { return None; }; + if any_measure_filtered(outer_filters) || any_measure_filtered(inner_filters) { + return None; + } let ([outer], [inner]) = (outer_measures.as_slice(), inner_measures.as_slice()) else { return None; }; @@ -283,6 +301,7 @@ fn composed_aggregate_rewrite(root: &Rc) -> Option> { reduction: outer_reduction.clone(), measures: vec![composed], output_names: output_names.clone(), + filters: vec![], having: None, child: Rc::clone(inner_child), }); @@ -413,6 +432,7 @@ mod tests { reduction: Reduction::by(by), measures: vec![AggIntent::Avg { col }], output_names: vec![], + filters: vec![], having: None, child: Rc::new(child), } @@ -459,6 +479,7 @@ mod tests { reduction: Reduction::by(vec![2]), measures: vec![AggIntent::Sum { col: None }, AggIntent::Avg { col: None }], output_names: vec![], + filters: vec![], having: None, child: Rc::new(metric_scan(&["job"])), }); @@ -494,6 +515,7 @@ mod tests { reduction: Reduction::by(vec![2]), measures: vec![intent.clone()], output_names: vec![], + filters: vec![], having: None, child: Rc::new(metric_scan(&["job"])), }); @@ -514,6 +536,7 @@ mod tests { )), measures: vec![AggIntent::Avg { col: None }], output_names: vec![], + filters: vec![], having: None, child: Rc::new(metric_scan(&["job"])), }); @@ -528,6 +551,7 @@ mod tests { reduction: Reduction::PerEntity, measures: vec![AggIntent::Avg { col: None }], output_names: vec![], + filters: vec![], having: None, child: Rc::new(metric_scan(&[])), }); @@ -755,6 +779,7 @@ mod tests { reduction: Reduction::PerEntity, measures: vec![inner], output_names: vec![], + filters: vec![], having: None, child: Rc::new(QueryExpr::TimeRange { range: Duration::from_secs(300), @@ -767,6 +792,7 @@ mod tests { // Match the PromQL front end: an empty entry selects the intent's // canonical output name rather than an explicit alias. output_names: vec![String::new()], + filters: vec![], having: None, child: Rc::new(temporal), }) diff --git a/crates/asap-aware-mapping/src/rollup.rs b/crates/asap-aware-mapping/src/rollup.rs index cc173672..a8bd0312 100644 --- a/crates/asap-aware-mapping/src/rollup.rs +++ b/crates/asap-aware-mapping/src/rollup.rs @@ -101,7 +101,7 @@ use std::collections::HashSet; use std::rc::Rc; use asap_types::pre_asap::agg_intent::AggIntent; -use asap_types::pre_asap::query_expr::{GroupKeys, QueryExpr, Reduction}; +use asap_types::pre_asap::query_expr::{any_measure_filtered, GroupKeys, QueryExpr, Reduction}; use asap_types::pre_asap::schema::{ColumnId, Schema}; use asap_types::types::AccuracyTarget; @@ -120,6 +120,7 @@ fn bindable_grouped_aggregate( let QueryExpr::Aggregate { reduction, measures, + filters, having, child, .. @@ -130,6 +131,9 @@ fn bindable_grouped_aggregate( let ([intent], None) = (measures.as_slice(), having) else { return None; }; + if any_measure_filtered(filters) { + return None; + } let Reduction::Reduce(by) = reduction else { return None; }; @@ -367,6 +371,7 @@ fn build_rollup( reduction: Reduction::by(remapped_by), measures: vec![combinator], output_names: output_names.to_vec(), + filters: vec![], having: None, child: Rc::clone(finer), }; @@ -416,6 +421,7 @@ mod tests { reduction: Reduction::by(by), measures: vec![intent], output_names: vec![], + filters: vec![], having: None, child: Rc::clone(child), }) @@ -430,6 +436,7 @@ mod tests { reduction: Reduction::Reduce(GroupKeys::without(excluded)), measures: vec![intent], output_names: vec![], + filters: vec![], having: None, child: Rc::clone(child), }) @@ -740,6 +747,7 @@ mod tests { reduction: Reduction::by(vec![2]), measures: vec![AggIntent::Sum { col: Some(1) }], output_names: vec!["total_requests".into()], + filters: vec![], having: None, child: Rc::clone(&scan), }); @@ -868,6 +876,7 @@ mod tests { }, ], output_names: vec![], + filters: vec![], having: None, child: Rc::clone(&scan), }); diff --git a/crates/asap-aware-mapping/src/summary_maintenance_cost/model.rs b/crates/asap-aware-mapping/src/summary_maintenance_cost/model.rs index e339994a..46753516 100644 --- a/crates/asap-aware-mapping/src/summary_maintenance_cost/model.rs +++ b/crates/asap-aware-mapping/src/summary_maintenance_cost/model.rs @@ -2523,6 +2523,7 @@ mod tests { input: asap_types::post_asap::SummaryUpdate::column(ColumnRef::Wildcard), reduction: Reduction::by(vec![]), grouping: GroupingStrategy::PerSubpopulationInstance, + filter: None, }, schema: estimated.schema.clone(), guarantee: None, @@ -2890,6 +2891,7 @@ mod tests { input: asap_types::post_asap::SummaryUpdate::column(ColumnRef::Wildcard), reduction: Reduction::by(vec![]), grouping: GroupingStrategy::PerSubpopulationInstance, + filter: None, }, schema: schema.clone(), guarantee: None, @@ -3057,6 +3059,7 @@ mod tests { reduction: Reduction::by(vec![]), measures: vec![AggIntent::Sum { col: None }], output_names: vec![], + filters: vec![], having: None, child: scan, }) diff --git a/crates/asap-aware-mapping/src/summary_maintenance_lifecycle.rs b/crates/asap-aware-mapping/src/summary_maintenance_lifecycle.rs index 87ef8d03..c6a5ca4b 100644 --- a/crates/asap-aware-mapping/src/summary_maintenance_lifecycle.rs +++ b/crates/asap-aware-mapping/src/summary_maintenance_lifecycle.rs @@ -1583,6 +1583,7 @@ mod tests { reduction: Reduction::by(vec![]), measures: vec![AggIntent::Sum { col: None }], output_names: vec![], + filters: vec![], having: None, child: query_root(), }) @@ -1597,6 +1598,7 @@ mod tests { accuracy: AccuracyTarget::Epsilon(0.1), }], output_names: vec![], + filters: vec![], having: None, child: query_root(), }) @@ -1621,6 +1623,7 @@ mod tests { )), reduction: Reduction::by(vec![]), grouping: GroupingStrategy::default(), + filter: None, }, schema: SummarySchema { fields: vec![SummaryField { @@ -1646,6 +1649,7 @@ mod tests { )), reduction: Reduction::by(vec![]), grouping: GroupingStrategy::default(), + filter: None, }, schema: SummarySchema { fields: vec![SummaryField { diff --git a/crates/devtools/src/bin/dag_export.rs b/crates/devtools/src/bin/dag_export.rs index 0fb2fd8b..4ad58a9c 100644 --- a/crates/devtools/src/bin/dag_export.rs +++ b/crates/devtools/src/bin/dag_export.rs @@ -1650,6 +1650,7 @@ mod tests { accuracy: AccuracyTarget::Epsilon(0.1), }], output_names: vec![], + filters: vec![], having: None, child: Rc::new(QueryExpr::Scan { source: Source::Table { diff --git a/crates/frontend-metricsql/src/lib.rs b/crates/frontend-metricsql/src/lib.rs index 4d07f64d..26a034c4 100644 --- a/crates/frontend-metricsql/src/lib.rs +++ b/crates/frontend-metricsql/src/lib.rs @@ -276,6 +276,7 @@ fn aggregate(reduction: Reduction, intent: AggIntent, chil reduction, measures: vec![intent], output_names: vec![String::new()], + filters: vec![], having: None, child: Rc::new(child), } diff --git a/crates/frontend-promql/src/promql.rs b/crates/frontend-promql/src/promql.rs index f0cf907f..52a16fd7 100644 --- a/crates/frontend-promql/src/promql.rs +++ b/crates/frontend-promql/src/promql.rs @@ -581,6 +581,7 @@ fn mark_without(tree: Unresolved, without: bool) -> Unresolved { reduction, measures, output_names, + filters, having, child, } => { @@ -592,6 +593,7 @@ fn mark_without(tree: Unresolved, without: bool) -> Unresolved { reduction: Reduction::Reduce(GroupKeys::without(keys)), measures, output_names, + filters, having, child, } @@ -765,6 +767,7 @@ fn walk_histogram_quantiles(call: &Call) -> Result { // intent-keyed default) so `Concat` — which derives its schema // from the first branch — doesn't silently misdescribe the rest. output_names: vec!["value".into()], + filters: vec![], having: None, child: Rc::new(child), }; @@ -1524,6 +1527,7 @@ fn build(inner: Inner, keys: Vec, outer: Outer) -> Result accuracy: current_accuracy(), }], output_names: vec![], + filters: vec![], having: None, child: Rc::new(ranked_agg), }) @@ -1612,6 +1616,7 @@ fn windowed_aggregate( // PromQL's intent-keyed output names ("sum", "quantile_0_99", …) // instead. output_names: vec![String::new()], + filters: vec![], having: None, child: Rc::new(child), } @@ -1630,6 +1635,7 @@ fn outer_aggregate( reduction, measures: vec![intent], output_names: vec![String::new()], + filters: vec![], having: None, child: Rc::new(child), } @@ -1649,6 +1655,7 @@ fn per_series_aggregate( reduction, measures: vec![intent], output_names: vec![String::new()], + filters: vec![], having: None, child: Rc::new(child), } diff --git a/crates/frontend-sql/src/sql/dialect.rs b/crates/frontend-sql/src/sql/dialect.rs new file mode 100644 index 00000000..03d9253e --- /dev/null +++ b/crates/frontend-sql/src/sql/dialect.rs @@ -0,0 +1,112 @@ +//! The parser dialect for `SqlDialect::DataFusionSQL`. +//! +//! sqlparser's `GenericDialect` leaves `FILTER (WHERE …)` on aggregate calls +//! off (`supports_filter_during_aggregation`), and DataFusion only selects a +//! dialect by name — so `count(x) FILTER (WHERE p)` cannot reach the planner +//! through `SessionContext::sql`. This wrapper is `GenericDialect` with that +//! one switch flipped (issue #466); `lower` parses through +//! `DFParser::parse_sql_with_dialect` with it and plans the statement itself, +//! exactly as the ClickHouse path already does. + +use std::any::TypeId; + +use datafusion::sql::sqlparser::dialect::{Dialect, GenericDialect}; + +#[derive(Debug, Default)] +pub(crate) struct GenericWithAggregateFilter; + +/// Forward every boolean switch `GenericDialect` overrides, so the only +/// behavioural difference is `supports_filter_during_aggregation`. +macro_rules! forward_to_generic { + ($($method:ident),* $(,)?) => { + $(fn $method(&self) -> bool { + GenericDialect.$method() + })* + }; +} + +impl Dialect for GenericWithAggregateFilter { + /// The parser's own `dialect_of!(… is GenericDialect)` checks keep + /// matching, so generic-only syntax paths stay enabled. + fn dialect(&self) -> TypeId { + GenericDialect.dialect() + } + + fn is_delimited_identifier_start(&self, ch: char) -> bool { + GenericDialect.is_delimited_identifier_start(ch) + } + + fn is_identifier_start(&self, ch: char) -> bool { + GenericDialect.is_identifier_start(ch) + } + + fn is_identifier_part(&self, ch: char) -> bool { + GenericDialect.is_identifier_part(ch) + } + + fn supports_filter_during_aggregation(&self) -> bool { + true + } + + forward_to_generic!( + supports_unicode_string_literal, + supports_group_by_expr, + supports_connect_by, + supports_match_recognize, + supports_start_transaction_modifier, + supports_window_function_null_treatment_arg, + supports_dictionary_syntax, + supports_window_clause_named_window_reference, + supports_parenthesized_set_variables, + supports_select_wildcard_except, + support_map_literal_syntax, + allow_extract_custom, + allow_extract_single_quotes, + supports_create_index_with_clause, + ); +} + +#[cfg(test)] +mod tests { + use super::*; + use datafusion::sql::parser::DFParser; + + // Every switch `GenericDialect` sets is mirrored, and only the aggregate + // FILTER switch differs. + #[test] + fn mirrors_generic_except_for_aggregate_filter() { + let ours = GenericWithAggregateFilter; + let generic = GenericDialect; + assert_eq!(ours.dialect(), generic.dialect()); + for ch in ['"', '`', '_', '#', '@', '$', 'a', '1', ' '] { + assert_eq!( + ours.is_delimited_identifier_start(ch), + generic.is_delimited_identifier_start(ch) + ); + assert_eq!( + ours.is_identifier_start(ch), + generic.is_identifier_start(ch) + ); + assert_eq!(ours.is_identifier_part(ch), generic.is_identifier_part(ch)); + } + assert_eq!( + ours.supports_group_by_expr(), + generic.supports_group_by_expr() + ); + assert!(!generic.supports_filter_during_aggregation()); + assert!(ours.supports_filter_during_aggregation()); + } + + // The generic dialect rejects an aggregate FILTER clause; ours parses it. + #[test] + fn parses_aggregate_filter_clause() { + let sql = "SELECT count(*) FILTER (WHERE a > 1) FROM t"; + assert!(DFParser::parse_sql_with_dialect(sql, &GenericDialect).is_err()); + assert_eq!( + DFParser::parse_sql_with_dialect(sql, &GenericWithAggregateFilter) + .unwrap() + .len(), + 1 + ); + } +} diff --git a/crates/frontend-sql/src/sql/mod.rs b/crates/frontend-sql/src/sql/mod.rs index 9ebf7f4a..d6e41eb7 100644 --- a/crates/frontend-sql/src/sql/mod.rs +++ b/crates/frontend-sql/src/sql/mod.rs @@ -49,6 +49,7 @@ use datafusion::logical_expr::{ use datafusion::optimizer::analyzer::function_rewrite::ApplyFunctionRewrites; use datafusion::optimizer::{AnalyzerRule, OptimizerConfig}; use datafusion::prelude::{SessionConfig, SessionContext}; +use datafusion::sql::parser::DFParser; use asap_sql_function_catalog::{AggSemantic, Arity, RewriteKind}; use asap_types::pre_asap::agg_intent::AggIntent; @@ -69,11 +70,13 @@ use crate::error::SqlError as LoweringError; mod clickhouse_ast; mod collection_planning; +mod dialect; mod expr; mod types; pub use types::SqlCatalog; +use self::dialect::GenericWithAggregateFilter; use self::expr::df_expr_to_unresolved; use self::types::{arrow_to_dtype, scalar_value_to_asap, schema_to_arrow}; @@ -172,16 +175,26 @@ impl<'a> SqlLowerer<'a> { accuracy: &AccuracyTarget, ) -> Result { let ctx = self.build_context()?; - let plan = if matches!(self.dialect, SqlDialect::ClickhouseSQL) { - let state = ctx.state(); + let state = ctx.state(); + let statement = if matches!(self.dialect, SqlDialect::ClickhouseSQL) { let mut statement = state.sql_to_statement(sql, "ClickHouse")?; if let datafusion::sql::parser::Statement::Statement(ast) = &mut statement { clickhouse_ast::normalize(ast); } - state.statement_to_plan(statement).await? + statement } else { - ctx.sql(sql).await?.into_unoptimized_plan() + // Not `ctx.sql(sql)`: that parses under the by-name `generic` + // dialect, which cannot see an aggregate `FILTER (WHERE …)`. + let mut statements = DFParser::parse_sql_with_dialect(sql, &GenericWithAggregateFilter) + .map_err(|e| datafusion::error::DataFusionError::SQL(e, None))?; + let (Some(statement), true) = (statements.pop_front(), statements.is_empty()) else { + return Err(LoweringError::UnsupportedFeature( + "exactly one SQL statement per query".into(), + )); + }; + statement }; + let plan = state.statement_to_plan(statement).await?; let rewriter = ApplyFunctionRewrites::new(vec![Arc::new(ClickHouseBuiltinRewrite)]); let plan = rewriter.analyze(plan, ctx.state().options())?; // Output schemas omit predicate and nested-expression types. Check the @@ -705,6 +718,7 @@ impl<'a> SqlLowerer<'a> { reduction: Reduction::Reduce(GroupKeys::none()), measures: vec![AggIntent::HistogramQuantile { q }], output_names: vec!["value".into()], + filters: vec![], having: None, child: Rc::new(input), }, @@ -750,31 +764,21 @@ impl<'a> SqlLowerer<'a> { fn lower_aggregate(&self, agg: &logical_expr::Aggregate) -> Result { let input = self.lower_plan(&agg.input)?; - // Canonical Count counts rows and has no nullable argument or FILTER. - // Check the original typed expression before derived-column rewriting - // erases the argument's nullability. - for expression in &agg.aggr_expr { - if let Expr::AggregateFunction(function) = unalias(expression) { - if function.func.name().eq_ignore_ascii_case("count") && !function.distinct { - let nullable = function - .args - .iter() - .try_fold(false, |nullable, argument| { - argument - .nullable(agg.input.schema().as_ref()) - .map(|next| nullable || next) - }) - .map_err(|error| LoweringError::UnsupportedFeature(error.to_string()))?; - if nullable || function.filter.is_some() { - return Err(LoweringError::UnsupportedFeature( - "COUNT of a nullable expression or with FILTER requires explicit per-aggregate null/filter semantics".into(), - )); - } - } - } - } + // Each measure's row predicate (`FILTER (WHERE …)`, or the NULL-skip + // a `count(expr)` implies), read off the original typed expression + // before derived-column rewriting erases the argument's nullability. + let measure_filters = agg + .aggr_expr + .iter() + .map(|e| measure_filter(e, agg.input.schema())) + .collect::, LoweringError>>()?; if agg.aggr_expr.iter().any(is_temporal_aggregate) { + if measure_filters.iter().any(Option::is_some) { + return Err(LoweringError::UnsupportedFeature( + "FILTER on an ASAP temporal aggregate".into(), + )); + } return self.lower_temporal_aggregate(agg, input); } @@ -782,6 +786,11 @@ impl<'a> SqlLowerer<'a> { // scan. `Aggregate.by` is a single key set, so each level becomes its own // `Aggregate` and they are merged (issue #118). if let Some(gs) = agg.group_expr.iter().find_map(as_grouping_set) { + if measure_filters.iter().any(Option::is_some) { + return Err(LoweringError::UnsupportedFeature( + "FILTER on a measure inside a multi-level grouping".into(), + )); + } return self.lower_grouping_sets(agg, gs, input); } @@ -827,6 +836,11 @@ impl<'a> SqlLowerer<'a> { .iter() .map(|e| derived.rewrite_agg(e)) .collect::, LoweringError>>()?; + // A measure filter reads the aggregate's input rows, so the columns + // it names must survive any derived-column `Project` inserted below. + for column in measure_filters.iter().flatten().flat_map(Expr::column_refs) { + derived.passthrough(&Expr::Column(column.clone()))?; + } let child = Rc::new(derived.wrap(input)?); // DataFusion names the aggregate outputs in its own schema (e.g. @@ -845,6 +859,19 @@ impl<'a> SqlLowerer<'a> { .iter() .map(lower_agg_intent) .collect::, LoweringError>>()?; + // Empty when nothing is filtered — the one canonical unfiltered shape. + let filters = if measure_filters.iter().any(Option::is_some) { + measure_filters + .iter() + .map(|f| { + f.as_ref() + .map(|f| Ok(Predicate(Rc::new(df_expr_to_unresolved(f)?)))) + .transpose() + }) + .collect::, LoweringError>>()? + } else { + Vec::new() + }; Ok(Unresolved::Aggregate { // SQL `GROUP BY` is always an inclusion list, never PromQL's // `without(...)` exclusion form — and always a genuine reduction, @@ -853,6 +880,7 @@ impl<'a> SqlLowerer<'a> { reduction: Reduction::Reduce(GroupKeys::by(keys)), measures, output_names, + filters, having: None, child, }) @@ -1001,6 +1029,7 @@ impl<'a> SqlLowerer<'a> { reduction: Reduction::PerEntity, measures: vec![intent], output_names: vec![], + filters: vec![], having: None, child: Rc::new(child), }) @@ -1093,6 +1122,7 @@ impl<'a> SqlLowerer<'a> { reduction: Reduction::Reduce(GroupKeys::by(level_keys)), measures: measures.clone(), output_names: output_names.clone(), + filters: vec![], having: None, child: Rc::new(input.clone()), }; @@ -1572,9 +1602,8 @@ impl FunctionRewrite for ClickHouseBuiltinRewrite { f.null_treatment, ), // `f(cond)` -> `sum(CASE WHEN cond THEN 1 ELSE 0 END)` — see - // `RewriteKind::CountIfToSum`'s doc for why a plain `count(...) - // FILTER (WHERE cond)` doesn't work here (`AggIntent::Count` - // never consults its argument). + // `RewriteKind::CountIfToSum`'s doc; moving the `-If` family onto + // `Aggregate.filters` (issue #466) is a follow-up. RewriteKind::CountIfToSum => { let cond = f.args.into_iter().next().expect( "countif's stub signature fixes its arity at 1 -- the planner \ @@ -1601,6 +1630,67 @@ impl FunctionRewrite for ClickHouseBuiltinRewrite { // ── Aggregate / group-key helpers ─────────────────────────────────────────────── +/// The row predicate one aggregate call carries (issue #466): its explicit +/// `FILTER (WHERE p)`, plus — for a plain `count(expr)`, which canonical +/// `AggIntent::Count` lowers to a row count that never looks at `expr` — the +/// NULL-skipping SQL gives it. `count(CASE WHEN p THEN x END)` is the +/// conditional-count idiom, so it becomes `p [AND x IS NOT NULL]` rather +/// than the opaque `CASE … IS NOT NULL`; any other nullable argument becomes +/// `expr IS NOT NULL`. `None` when the call updates on every row. +fn measure_filter(expr: &Expr, input: &DFSchema) -> Result, LoweringError> { + let Expr::AggregateFunction(agg_fn) = unalias(expr) else { + return Ok(None); + }; + let mut conjuncts: Vec = agg_fn.filter.iter().map(|f| (**f).clone()).collect(); + let counts_rows = agg_fn.func.name().eq_ignore_ascii_case("count") && !agg_fn.distinct; + if counts_rows { + for argument in &agg_fn.args { + let nullable = argument + .nullable(input) + .map_err(|error| LoweringError::UnsupportedFeature(error.to_string()))?; + if !nullable { + continue; + } + match conditional_count_arm(argument) { + Some((when, then)) => { + conjuncts.push(when.clone()); + if then + .nullable(input) + .map_err(|error| LoweringError::UnsupportedFeature(error.to_string()))? + { + conjuncts.push(then.clone().is_not_null()); + } + } + None => conjuncts.push(argument.clone().is_not_null()), + } + } + } + Ok(conjuncts.into_iter().reduce(Expr::and)) +} + +/// `CASE WHEN p THEN x END` (searched, one arm, no `ELSE` or `ELSE NULL`) +/// as `(p, x)`. +fn conditional_count_arm(expr: &Expr) -> Option<(&Expr, &Expr)> { + let Expr::Case(case) = unalias(expr) else { + return None; + }; + if case.expr.is_some() { + return None; + } + let else_is_null = match case.else_expr.as_deref() { + None => true, + Some(Expr::Literal(value)) => value.is_null(), + Some(_) => false, + }; + if !else_is_null { + return None; + } + let [(when, then)] = case.when_then_expr.as_slice() else { + return None; + }; + Some((when, then)) +} + /// Map a DataFusion aggregate expression directly to the canonical /// [`AggIntent`] — issue #179's "dedicated function → canonical /// intent directly" front-end construction, no `AggFunc` intermediate. The @@ -1648,12 +1738,9 @@ fn lower_agg_intent(expr: &Expr) -> Result, LoweringError> }; Ok(match semantic { AggSemantic::Correlation => { - if agg_fn.filter.is_some() - || agg_fn.order_by.is_some() - || agg_fn.null_treatment.is_some() - { + if agg_fn.order_by.is_some() || agg_fn.null_treatment.is_some() { return Err(LoweringError::UnsupportedAggregate( - "corr with FILTER, ORDER BY, or explicit null treatment".into(), + "corr with ORDER BY or explicit null treatment".into(), )); } let [left, right] = agg_fn.args.as_slice() else { diff --git a/crates/frontend-sql/tests/pearson_corr.rs b/crates/frontend-sql/tests/pearson_corr.rs index 82e9d60d..a2301d20 100644 --- a/crates/frontend-sql/tests/pearson_corr.rs +++ b/crates/frontend-sql/tests/pearson_corr.rs @@ -108,7 +108,6 @@ async fn corr_repeated_input_and_serialization() { async fn corr_rejects_unrepresented_forms() { for sql in [ "SELECT corr(DISTINCT x, y) FROM a", - "SELECT corr(x, y) FILTER (WHERE g > 0) FROM a", "SELECT corr(x, y ORDER BY g) FROM a", "SELECT corr(x, y) OVER () FROM a", "SELECT corr(x) FROM a", @@ -122,6 +121,27 @@ async fn corr_rejects_unrepresented_forms() { } } +// A `FILTER` clause becomes the measure's own predicate (#466), leaving the +// two inputs untouched. +#[tokio::test] +async fn corr_filter_is_a_measure_filter() { + let query = lower("SELECT corr(x, y) FILTER (WHERE g > 0) FROM a").await; + let (measures, _) = aggregate(&query); + assert_eq!(measures[0].input_cols(), vec![0, 1]); + fn filters(query: &QueryExpr) -> &[Option] { + match query { + QueryExpr::Aggregate { filters, .. } => filters, + QueryExpr::Project { child, .. } | QueryExpr::Filter { child, .. } => filters(child), + other => panic!("expected aggregate, got {other:?}"), + } + } + assert!( + matches!(filters(&query), [Some(_)]), + "{:?}", + filters(&query) + ); +} + // Exact fallback retains the complete typed query and compiles to an executable DAG. #[tokio::test] async fn corr_survives_exact_plan_compilation() { diff --git a/crates/frontend-sql/tests/sql_lowering.rs b/crates/frontend-sql/tests/sql_lowering.rs index 5482e6f6..eeda467b 100644 --- a/crates/frontend-sql/tests/sql_lowering.rs +++ b/crates/frontend-sql/tests/sql_lowering.rs @@ -8,8 +8,8 @@ use asap_frontend_sql::{lower_sql, lower_sql_dialect, SqlCatalog, SqlError as LoweringError}; use asap_types::pre_asap::schema::{Column, DataType, Schema}; use asap_types::pre_asap::{ - AggIntent, CompareOpKind, GroupKeys, JoinKind, QueryExpr, Reduction, ScalarValue, Source, - WindowFrameBound, WindowFrameOffset, WindowFrameUnits, WindowFuncKind, + AggIntent, CompareOpKind, GroupKeys, JoinKind, Predicate, QueryExpr, Reduction, ScalarValue, + Source, WindowFrameBound, WindowFrameOffset, WindowFrameUnits, WindowFuncKind, }; use asap_types::types::AccuracyTarget; use asap_types::workload::SqlDialect; @@ -1850,11 +1850,9 @@ async fn count_if_lowers_to_a_sum_over_a_derived_indicator_column() { // ClickHouse's `countIf(cond)` has no DataFusion equivalent at all, so it // goes through the same stub-UDAF + catalog-driven `FunctionRewrite` // mechanism `uniqExact` (#221) does — rewritten, before `lower_agg_intent` - // ever runs, to `sum(CASE WHEN cond THEN 1 ELSE 0 END)`. Not a plain - // `count(...) FILTER (WHERE cond)`: `AggIntent::Count` never consults its - // argument (always a row count), so the filter would be silently dropped; - // summing a 0/1 indicator keeps `cond` observable through the ordinary - // `Sum` path instead. + // ever runs, to `sum(CASE WHEN cond THEN 1 ELSE 0 END)`. A per-measure + // filter (#466) could express it as a filtered `Count` now; that move is + // a follow-up, so the indicator sum is still the shape to expect. let qe = lower_clickhouse("SELECT countIf(bytes > 100) AS big FROM metrics").await; let (by, measures) = find_aggregate(&qe).expect("expected an Aggregate"); assert!(by.is_empty()); @@ -2075,8 +2073,11 @@ async fn current_timestamp_lowers_to_typed_current_timestamp_leaf() { assert_eq!(schema.columns[0].dtype, DataType::Timestamp); } +// A `count` over a non-null input is a plain row count; over a nullable +// input it keeps SQL's NULL-skipping as the measure's own filter (#466), and +// only the multi-level grouping path, which cannot carry one, still rejects it. #[tokio::test] -async fn count_preserves_non_null_inputs_and_rejects_erased_null_semantics() { +async fn count_null_semantics_become_a_measure_filter() { let catalog = SqlCatalog::new().with_table( "samples", Schema::new(vec![ @@ -2090,25 +2091,53 @@ async fn count_preserves_non_null_inputs_and_rejects_erased_null_semantics() { "SELECT count(value) FROM samples", "SELECT count(value + 1) FROM samples", ] { - lower_sql(sql, &catalog, AccuracyTarget::Exact) + let qe = lower_sql(sql, &catalog, AccuracyTarget::Exact) .await .unwrap_or_else(|error| panic!("{sql}: {error}")); + assert!( + aggregate_filters(&qe).is_empty(), + "{sql}: unfiltered row count" + ); } for sql in [ "SELECT count(nullable_value) FROM samples", "SELECT count(NULL) FROM samples", "SELECT count(nullable_value + 1) FROM samples", - "SELECT count(*), count(nullable_value) FROM samples", - "SELECT count(nullable_value) FROM samples GROUP BY ROLLUP(value)", ] { - let error = lower_sql(sql, &catalog, AccuracyTarget::Exact) + let qe = lower_sql(sql, &catalog, AccuracyTarget::Exact) .await - .unwrap_err(); + .unwrap_or_else(|error| panic!("{sql}: {error}")); + let [Some(Predicate(cond))] = aggregate_filters(&qe) else { + panic!( + "{sql}: expected one filtered Count, got {:?}", + aggregate_filters(&qe) + ); + }; assert!( - error.to_string().contains("explicit per-aggregate"), - "{sql}: {error}" + matches!(cond.as_ref(), QueryExpr::IsNotNull(_)), + "{sql}: {cond:?}" ); } + // Only the second measure is filtered. + let qe = lower_sql( + "SELECT count(*), count(nullable_value) FROM samples", + &catalog, + AccuracyTarget::Exact, + ) + .await + .unwrap(); + assert!(matches!(aggregate_filters(&qe), [None, Some(_)])); + let error = lower_sql( + "SELECT count(nullable_value) FROM samples GROUP BY ROLLUP(value)", + &catalog, + AccuracyTarget::Exact, + ) + .await + .unwrap_err(); + assert!( + matches!(error, LoweringError::UnsupportedFeature(_)), + "{error}" + ); } /// A native SQL map grouping key retains its typed key/value schema. @@ -2555,3 +2584,123 @@ async fn distinct_with_derived_sibling() { assert!(result.is_ok(), "{sql}: {result:?}"); } } + +// ── Issue #466: per-measure FILTER predicates ───────────────────────────────── + +/// The first `Aggregate`'s `filters`, positional against its child. +fn aggregate_filters(qe: &QueryExpr) -> &[Option] { + let Some(QueryExpr::Aggregate { filters, .. }) = find_aggregate_node(qe) else { + panic!("expected an Aggregate, got {qe:?}"); + }; + filters +} + +// The motivating query: one scan, one grouping, one conditional count next to +// a plain sum — a single `Aggregate` whose Count carries the condition, with no +// `Join` and no derived column for the `CASE`. +#[tokio::test] +async fn conditional_count_lowers_to_a_filtered_measure() { + let qe = lower( + "SELECT service, count(CASE WHEN latency > 1.0 THEN 1 END), sum(bytes) \ + FROM metrics GROUP BY service", + ) + .await; + assert!(find_join(&qe).is_none(), "no join: {qe:?}"); + let (by, measures) = find_aggregate(&qe).unwrap(); + assert_eq!(by.keys(), &[1]); + assert!( + matches!( + measures.as_slice(), + [AggIntent::Count { .. }, AggIntent::Sum { col: Some(3) }] + ), + "{measures:?}" + ); + let [Some(Predicate(cond)), None] = aggregate_filters(&qe) else { + panic!("expected [Some, None], got {:?}", aggregate_filters(&qe)); + }; + assert!( + matches!(cond.as_ref(), QueryExpr::Compare { left, op: CompareOpKind::Gt, .. } + if matches!(left.as_ref(), QueryExpr::Column(2))), + "latency > 1.0 against the scan, got {cond:?}" + ); + let Some(QueryExpr::Aggregate { child, .. }) = find_aggregate_node(&qe) else { + unreachable!() + }; + assert!( + matches!(child.as_ref(), QueryExpr::Scan { .. }), + "{child:?}" + ); +} + +// `FILTER (WHERE …)` parses under the DataFusion dialect and lands on exactly +// the measure it annotates. +#[tokio::test] +async fn filter_clause_lowers_to_a_measure_filter() { + let qe = lower("SELECT sum(bytes) FILTER (WHERE service = 'a'), count(*) FROM metrics").await; + let [Some(Predicate(cond)), None] = aggregate_filters(&qe) else { + panic!("expected [Some, None], got {:?}", aggregate_filters(&qe)); + }; + assert!( + matches!(cond.as_ref(), QueryExpr::Compare { left, op: CompareOpKind::Eq, right } + if matches!(left.as_ref(), QueryExpr::Column(1)) + && matches!(right.as_ref(), QueryExpr::Literal(ScalarValue::Utf8(s)) if s == "a")), + "{cond:?}" + ); +} + +// SQL `count(expr)` skips NULLs; canonical `Count` counts rows and never sees +// `expr`, so a nullable argument becomes the measure filter `expr IS NOT NULL` +// instead of being rejected (the pre-#466 behavior) or silently over-counted. +#[tokio::test] +async fn count_of_a_nullable_expression_filters_nulls() { + let qe = lower("SELECT count(nullif(bytes, 0)) FROM metrics").await; + let [Some(Predicate(cond))] = aggregate_filters(&qe) else { + panic!("expected [Some], got {:?}", aggregate_filters(&qe)); + }; + assert!(matches!(cond.as_ref(), QueryExpr::IsNotNull(_)), "{cond:?}"); + assert!( + matches!( + find_aggregate(&qe).unwrap().1.as_slice(), + [AggIntent::Count { .. }] + ), + "still a row count" + ); +} + +// The columns a measure filter reads must survive the derived-column +// `Project` a reducer expression inserts beneath the aggregate. +#[tokio::test] +async fn measure_filter_columns_survive_a_derived_column_projection() { + let qe = lower("SELECT sum(bytes * 2) FILTER (WHERE latency > 1.0) FROM metrics").await; + let Some(QueryExpr::Aggregate { child, .. }) = find_aggregate_node(&qe) else { + unreachable!() + }; + assert!( + matches!(child.as_ref(), QueryExpr::Project { .. }), + "{child:?}" + ); + let [Some(Predicate(cond))] = aggregate_filters(&qe) else { + panic!("expected [Some], got {:?}", aggregate_filters(&qe)); + }; + let QueryExpr::Compare { left, .. } = cond.as_ref() else { + panic!("{cond:?}"); + }; + let QueryExpr::Column(id) = left.as_ref() else { + panic!("{left:?}"); + }; + assert_eq!(child.output_schema().unwrap().columns[*id].name, "latency"); +} + +// `GROUP BY ROLLUP` fans one measure list out into one `Aggregate` per level; +// a filtered measure there is rejected rather than silently unfiltered. +#[tokio::test] +async fn measure_filter_inside_a_rollup_is_rejected() { + let err = lower_sql( + "SELECT service, count(*) FILTER (WHERE latency > 1.0) FROM metrics GROUP BY ROLLUP(service)", + &catalog(), + AccuracyTarget::Exact, + ) + .await + .unwrap_err(); + assert!(matches!(err, LoweringError::UnsupportedFeature(_)), "{err}"); +} diff --git a/crates/integration-tests/tests/aggregate.rs b/crates/integration-tests/tests/aggregate.rs index dd0bef4e..051eeb6b 100644 --- a/crates/integration-tests/tests/aggregate.rs +++ b/crates/integration-tests/tests/aggregate.rs @@ -35,6 +35,7 @@ fn agg(by: Vec, intent: AggIntent, child: QueryExpr) -> QueryExpr { reduction: Reduction::by(by), measures: vec![intent], output_names: vec!["".into()], + filters: vec![], having: None, child: Rc::new(QueryExpr::TimeRange { range: Duration::from_secs(1), diff --git a/crates/integration-tests/tests/binary_op.rs b/crates/integration-tests/tests/binary_op.rs index 63ba458f..35f1c632 100644 --- a/crates/integration-tests/tests/binary_op.rs +++ b/crates/integration-tests/tests/binary_op.rs @@ -42,6 +42,7 @@ fn rate_agg(metric: &str) -> QueryExpr { reduction: Reduction::PerEntity, measures: vec![AggIntent::Rate], output_names: vec!["".into()], + filters: vec![], having: None, child: Rc::new(QueryExpr::TimeRange { range: Duration::from_secs(300), @@ -55,6 +56,7 @@ fn sum_by_job(metric: &str) -> QueryExpr { reduction: Reduction::by(vec![2]), measures: vec![AggIntent::Sum { col: None }], output_names: vec!["".into()], + filters: vec![], having: None, child: Rc::new(scan(metric, &["job"])), } @@ -256,6 +258,7 @@ fn q36_sum_of_negation_nests() { reduction: Reduction::by(vec![]), measures: vec![AggIntent::Sum { col: None }], output_names: vec!["".into()], + filters: vec![], having: None, child: Rc::new(QueryExpr::BinaryOp { op: BinaryOpKind::Arithmetic(ArithmeticOpKind::Mul), diff --git a/crates/integration-tests/tests/exact_composition.rs b/crates/integration-tests/tests/exact_composition.rs index 40b593a1..e3a64690 100644 --- a/crates/integration-tests/tests/exact_composition.rs +++ b/crates/integration-tests/tests/exact_composition.rs @@ -56,6 +56,7 @@ fn agg(by: Vec, intent: AggIntent, child: Rc) -> Rc reduction: Reduction::by(by), measures: vec![intent], output_names: vec![], + filters: vec![], having: None, child, }) @@ -66,6 +67,7 @@ fn per_entity(intent: AggIntent, child: Rc) -> Rc { reduction: Reduction::PerEntity, measures: vec![intent], output_names: vec![], + filters: vec![], having: None, child, }) @@ -689,6 +691,7 @@ fn summary_construction_follows_its_value_input_phase() { input: SummaryUpdate::column(asap_types::pre_asap::ColumnRef::SampleValue), reduction: Reduction::by(vec![]), grouping: Default::default(), + filter: None, }, schema: asap_types::post_asap::SummarySchema { fields: vec![], diff --git a/crates/integration-tests/tests/nested.rs b/crates/integration-tests/tests/nested.rs index 2d937082..af25af3b 100644 --- a/crates/integration-tests/tests/nested.rs +++ b/crates/integration-tests/tests/nested.rs @@ -29,6 +29,7 @@ fn agg(by: Vec, intent: AggIntent, child: QueryExpr) -> QueryExpr { reduction: Reduction::by(by), measures: vec![intent], output_names: vec!["".into()], + filters: vec![], having: None, child: Rc::new(child), } @@ -39,6 +40,7 @@ fn agg_per_entity(intent: AggIntent, child: QueryExpr) -> QueryExpr { reduction: Reduction::PerEntity, measures: vec![intent], output_names: vec!["".into()], + filters: vec![], having: None, child: Rc::new(child), } @@ -292,6 +294,7 @@ fn q39_sum_without_instance_over_rate() { reduction: Reduction::Reduce(GroupKeys::without(vec![2])), // exclude `instance` measures: vec![AggIntent::Sum { col: None }], output_names: vec!["".into()], + filters: vec![], having: None, child: Rc::new(inner_rate), }; diff --git a/crates/integration-tests/tests/time_range.rs b/crates/integration-tests/tests/time_range.rs index 940facc3..d3ab732f 100644 --- a/crates/integration-tests/tests/time_range.rs +++ b/crates/integration-tests/tests/time_range.rs @@ -34,6 +34,7 @@ fn range_agg(range_secs: u64, intent: AggIntent, metric: &str) -> QueryExpr { reduction: Reduction::PerEntity, measures: vec![intent], output_names: vec!["".into()], + filters: vec![], having: None, child: Rc::new(QueryExpr::TimeRange { range: Duration::from_secs(range_secs), diff --git a/crates/sql-function-catalog/src/lib.rs b/crates/sql-function-catalog/src/lib.rs index bf7dd143..a95a2a63 100644 --- a/crates/sql-function-catalog/src/lib.rs +++ b/crates/sql-function-catalog/src/lib.rs @@ -238,15 +238,12 @@ pub enum RewriteKind { /// is needed once the call wears DataFusion's own name. CountDistinct, /// `f(cond)` -> `sum(CASE WHEN cond THEN 1 ELSE 0 END)` -- ClickHouse's - /// conditional-count family. Not a plain `count(...) FILTER (WHERE - /// cond)`: `AggIntent::Count` never consults its argument (it always - /// means "row count"), so a per-call *filtered* count needs a shape - /// whose value actually depends on `cond` to survive `lower_agg_intent` - /// unchanged. Summing a 0/1 indicator does, and lands on the existing - /// `Sum` path -- including the general non-column-argument + /// conditional-count family. Predates per-measure filters (issue #466, + /// `Aggregate.filters`); the `-If` combinators' move onto that field is + /// left for a follow-up, so the indicator sum stays: it lands on the + /// existing `Sum` path -- including the general non-column-argument /// materialization `asap-frontend-sql`'s `lower_aggregate` already does - /// for any reducer over an expression (issue #110) -- so no new - /// `AggIntent` variant or lowering path is needed either. + /// for any reducer over an expression (issue #110). CountIfToSum, /// No native DataFusion aggregate shape to rewrite to at all -- the call /// survives unchanged (`ClickHouseBuiltinRewrite` is a no-op for it) and diff --git a/crates/types/src/dag_export.rs b/crates/types/src/dag_export.rs index 70c04846..289baa25 100644 --- a/crates/types/src/dag_export.rs +++ b/crates/types/src/dag_export.rs @@ -1210,6 +1210,7 @@ fn build_no_recheck( reduction, measures, output_names, + filters, having, child, } => { @@ -1218,6 +1219,7 @@ fn build_no_recheck( "reduction": reduction, "measures": measures, "output_names": output_names, + "filters": filters, "having": having, }); push_node( @@ -1551,6 +1553,7 @@ mod tests { accuracy: AccuracyTarget::Exact, }], output_names: vec![], + filters: vec![], having: None, child: Rc::new(scan("metrics", value_col())), }), @@ -1668,6 +1671,7 @@ mod tests { accuracy: AccuracyTarget::Exact, }], output_names: vec![], + filters: vec![], having: None, child: Rc::new(scan("metrics", value_col())), }; @@ -1724,6 +1728,7 @@ mod tests { ), reduction: Reduction::by(vec![]), grouping: GroupingStrategy::default(), + filter: None, }, schema: SummarySchema { fields: vec![], diff --git a/crates/types/src/post_asap/cse.rs b/crates/types/src/post_asap/cse.rs index 6d1d0dc9..47b10701 100644 --- a/crates/types/src/post_asap/cse.rs +++ b/crates/types/src/post_asap/cse.rs @@ -81,6 +81,7 @@ fn same_node(left: &SummaryNode, right: &SummaryNode) -> bool { input: ai, reduction: ar, grouping: ag, + filter: afl, }, SummaryAgg { child: bc, @@ -88,8 +89,16 @@ fn same_node(left: &SummaryNode, right: &SummaryNode) -> bool { input: bi, reduction: br, grouping: bg, + filter: bfl, }, - ) => Rc::ptr_eq(ac, bc) && af == bf && same_value(ai, bi) && ar == br && ag == bg, + ) => { + Rc::ptr_eq(ac, bc) + && af == bf + && same_value(ai, bi) + && ar == br + && ag == bg + && same_value(afl, bfl) + } ( SummaryJoin { outer: ao, @@ -363,6 +372,7 @@ mod tests { input: SummaryUpdate::column(ColumnRef::SampleValue), reduction: Reduction::PerEntity, grouping: GroupingStrategy::default(), + filter: None, }, schema: SummarySchema { fields: vec![], diff --git a/crates/types/src/post_asap/executable_dag.rs b/crates/types/src/post_asap/executable_dag.rs index c156aa97..78dd6d73 100644 --- a/crates/types/src/post_asap/executable_dag.rs +++ b/crates/types/src/post_asap/executable_dag.rs @@ -14,7 +14,7 @@ use super::{ use crate::pre_asap::{ColumnRef, JoinKind, Predicate, QueryExpr, Reduction}; use thiserror::Error; -pub const POST_ASAP_DAG_WIRE_VERSION: u32 = 5; +pub const POST_ASAP_DAG_WIRE_VERSION: u32 = 6; #[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)] pub enum EdgeRole { @@ -69,6 +69,10 @@ pub enum ExecutableOperatorPayload { input: SummaryUpdate, reduction: Reduction, grouping: GroupingStrategy, + /// See `SummaryExpr::SummaryAgg::filter`. Wire version 6 added it; + /// a version-5 reader would otherwise take a filtered summary as + /// unfiltered. + filter: Option, }, SummaryJoin { key: ColumnRef, @@ -446,12 +450,14 @@ pub fn compile_executable_dag_with_node_ids( input, reduction, grouping, + filter, .. } => ExecutableOperatorPayload::SummaryAgg { family: family.clone(), input: input.clone(), reduction: reduction.clone(), grouping: grouping.clone(), + filter: filter.clone(), }, SummaryExpr::SummaryJoin { key, family, .. } => { ExecutableOperatorPayload::SummaryJoin { @@ -610,6 +616,7 @@ mod tests { input: SummaryUpdate::column(ColumnRef::SampleValue), reduction: Reduction::by(vec![]), grouping: GroupingStrategy::default(), + filter: None, }, ExecutableOperatorPayload::SummaryJoin { key: ColumnRef::SampleValue, @@ -756,6 +763,7 @@ mod tests { input: SummaryUpdate::column(ColumnRef::SampleValue), reduction: Reduction::by(vec![]), grouping: GroupingStrategy::default(), + filter: None, }, schema: SummarySchema { fields: vec![SummaryField { diff --git a/crates/types/src/post_asap/execution_data_state.rs b/crates/types/src/post_asap/execution_data_state.rs index e61cafe3..307efe88 100644 --- a/crates/types/src/post_asap/execution_data_state.rs +++ b/crates/types/src/post_asap/execution_data_state.rs @@ -49,7 +49,7 @@ use thiserror::Error; use super::expr::{ExactOperation, SummaryExpr, SummaryNode, ValueOperation}; use super::schema::{SummaryFamilyType, SummaryField, SummarySchema}; -use crate::pre_asap::query_expr::{aggregate_output_schema, QueryExprError}; +use crate::pre_asap::query_expr::{aggregate_output_schema, Predicate, QueryExprError}; use crate::pre_asap::schema::{Column, Schema}; /// When a post-ASAP value is produced. @@ -649,6 +649,7 @@ fn check_plain_operands( let ExactOperation::Aggregate { reduction, measures, + filters, .. } = op; let mut referenced: Vec = reduction @@ -658,6 +659,9 @@ fn check_plain_operands( for m in measures { referenced.extend(m.input_cols()); } + for Predicate(f) in filters.iter().flatten() { + referenced.extend(f.columns_referenced().into_iter().copied()); + } // With no explicit input column (the PromQL sample-value convention) // the operator reads every non-key column, so all must be plain. let implicit = measures.iter().any(|m| m.input_cols().is_empty()); @@ -839,6 +843,7 @@ mod tests { input: crate::post_asap::SummaryUpdate::column(ColumnRef::SampleValue), reduction: Reduction::by(vec![]), grouping: GroupingStrategy::default(), + filter: None, }, schema: SummarySchema { fields: vec![SummaryField { @@ -877,6 +882,7 @@ mod tests { measures: vec![AggIntent::Max { col: None }], output_names: vec![], having: None, + filters: vec![], } } @@ -1144,6 +1150,7 @@ mod tests { measures: vec![AggIntent::PearsonCorr { left: 0, right: 1 }], output_names: vec![], having: None, + filters: vec![], }); for operand in [0, 1] { let mut input = plain(&["x", "y", "unused"]); @@ -1166,6 +1173,7 @@ mod tests { measures: vec![AggIntent::Max { col: None }], output_names: vec![], having: None, + filters: vec![], }; let out = exact_operation_output_schema(&op, &child_schema).unwrap(); let names: Vec<_> = out.fields.iter().map(|f| f.name.as_str()).collect(); diff --git a/crates/types/src/post_asap/expr.rs b/crates/types/src/post_asap/expr.rs index 12e2c21c..6c48114a 100644 --- a/crates/types/src/post_asap/expr.rs +++ b/crates/types/src/post_asap/expr.rs @@ -18,6 +18,11 @@ pub enum ExactOperation { reduction: Reduction, measures: Vec, output_names: Vec, + /// Per-measure row predicates parallel to `measures`, positional + /// against the child's output rows — the same contract as + /// `QueryExpr::Aggregate.filters` (issue #466). + #[serde(default)] + filters: Vec>, having: Option, }, } @@ -199,6 +204,14 @@ pub enum SummaryExpr { /// `GroupingStrategy::PerSubpopulationInstance` (its `Default`), /// so no existing behavior changes. grouping: GroupingStrategy, + /// Row predicate gating this summary's updates (issue #466): only + /// rows where it is `TRUE` update the state; grouping keys are + /// still read from every row. Positional against `child`'s output. + /// A field rather than a `Filter` child so summaries that differ + /// only in predicate can still share one child. No binding rule + /// sets it yet — a filtered pre-ASAP measure stays `KeepPreAsap` — + /// so every producer today writes `None`. + filter: Option, }, /// Summary-aware join (KMV / theta for join-cardinality; join-sample for diff --git a/crates/types/src/pre_asap/canonicalize.rs b/crates/types/src/pre_asap/canonicalize.rs index aaf2c06d..3a66c5c4 100644 --- a/crates/types/src/pre_asap/canonicalize.rs +++ b/crates/types/src/pre_asap/canonicalize.rs @@ -213,6 +213,7 @@ fn try_promote_additive_top_ranking(expr: &QueryExpr) -> Option { let QueryExpr::Aggregate { reduction, measures, + filters, child: aggregate_child, .. } = agg_expr @@ -225,6 +226,11 @@ fn try_promote_additive_top_ranking(expr: &QueryExpr) -> Option { let [ranked_agg] = measures.as_slice() else { return None; }; + // A heavy-hitter sketch ranks the raw update stream; a filtered measure + // only counts part of it, and no binding rule applies the filter (#466). + if filters.iter().any(Option::is_some) { + return None; + } if ranked_col != by.len() { return None; } @@ -261,6 +267,7 @@ fn try_promote_additive_top_ranking(expr: &QueryExpr) -> Option { reduction: Reduction::by(partition_by.to_vec()), measures: vec![AggIntent::TopK { k: *k, accuracy }], output_names: Vec::new(), + filters: Vec::new(), having: None, child: Rc::new(agg_expr.clone()), }) @@ -373,6 +380,7 @@ mod tests { accuracy: AccuracyTarget::Exact, }], output_names: vec![], + filters: vec![], having: None, child: Rc::new(scan()), } @@ -417,6 +425,25 @@ mod tests { assert!(is_topk_over_count(&canonicalize(q))); } + // A heavy-hitter sketch ranks every row; a count that only counts some + // rows (#466) is not that, so the generic Sort + Limit stays. + #[test] + fn does_not_promote_a_filtered_count_ranking() { + let mut filtered = count_by_service(); + let QueryExpr::Aggregate { filters, .. } = &mut filtered else { + unreachable!() + }; + *filters = vec![Some(Predicate(Rc::new(QueryExpr::Compare { + left: Rc::new(QueryExpr::Column(2)), + op: CompareOpKind::Gt, + right: Rc::new(QueryExpr::Literal(ScalarValue::Float64(1.0))), + })))]; + let q = limit(5, 0, sort(desc(1), filtered)); + let canonical = canonicalize(q.clone()); + assert!(!is_topk_over_count(&canonical)); + assert_eq!(canonical, q); + } + #[test] fn promotes_through_a_passthrough_projection() { // …with a `SELECT service, count` projection between the Sort and the Agg. @@ -545,6 +572,7 @@ mod tests { reduction: Reduction::by(vec![1]), measures: vec![AggIntent::Sum { col: None }], output_names: vec![], + filters: vec![], having: None, child: Rc::new(scan()), }; @@ -573,6 +601,7 @@ mod tests { reduction: Reduction::PerEntity, measures: vec![counter], output_names: vec![], + filters: vec![], having: None, child: Rc::new(scan()), }; @@ -580,6 +609,7 @@ mod tests { reduction: Reduction::by(vec![1]), measures: vec![AggIntent::Sum { col: None }], output_names: vec![], + filters: vec![], having: None, child: Rc::new(derived), }; @@ -619,6 +649,7 @@ mod tests { reduction: Reduction::by(vec![1, 2]), measures: vec![agg], output_names: vec![], + filters: vec![], having: None, child: Rc::new(scan4()), } diff --git a/crates/types/src/pre_asap/column_resolution.rs b/crates/types/src/pre_asap/column_resolution.rs index f1db980a..9f819ad1 100644 --- a/crates/types/src/pre_asap/column_resolution.rs +++ b/crates/types/src/pre_asap/column_resolution.rs @@ -373,6 +373,7 @@ mod tests { reduction: Reduction::PerEntity, measures: vec![AggIntent::Rate], output_names: vec![], + filters: vec![], having: None, child: Rc::new(QueryExpr::TimeRange { range: Duration::from_secs(300), diff --git a/crates/types/src/pre_asap/cse.rs b/crates/types/src/pre_asap/cse.rs index 3f3a9cfe..9514451f 100644 --- a/crates/types/src/pre_asap/cse.rs +++ b/crates/types/src/pre_asap/cse.rs @@ -288,12 +288,20 @@ pub fn structural_hash(node: &QueryExpr, cache: &mut HashCache) -> u64 { reduction, measures, output_names, + filters, having, child, } => { hash_own_fields( &mut hasher, - &("Aggregate", reduction, measures, output_names, having), + &( + "Aggregate", + reduction, + measures, + output_names, + filters, + having, + ), ); child_hash(child, cache).hash(&mut hasher); } @@ -610,12 +618,14 @@ fn rebuild_children(table: &mut InternTable, expr: QueryExpr) -> QueryExpr { reduction, measures, output_names, + filters, having, child, } => Aggregate { reduction, measures, output_names, + filters, having, child: intern_child(table, child), }, @@ -762,7 +772,8 @@ pub fn share_common_subtrees(roots: Vec<(Id, QueryExpr)>) -> Vec<(Id, Rc(pub Rc>); +/// Whether any entry of an `Aggregate.filters` vector is set — the shape +/// no binding rule accepts yet (issue #466): a filtered measure stays +/// `KeepPreAsap`, and heavy-hitter promotion skips it. +pub fn any_measure_filtered(filters: &[Option>]) -> bool { + filters.iter().any(Option::is_some) +} + /// One item in a SELECT projection list. #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] #[serde(bound(serialize = "C: ColState", deserialize = "C: ColState"))] @@ -767,6 +774,15 @@ pub enum QueryExpr { /// (PromQL's convention). #[serde(default)] output_names: Vec, + /// Per-measure row predicates, parallel to `measures` — SQL + /// `FILTER (WHERE …)` semantics (issue #466): only rows where + /// `filters[i]` is `TRUE` update `measures[i]`; groups are still + /// formed from every row. Positional against `child`'s output + /// schema, like `Filter.pred` — not against this node's output like + /// `having`. `None` (or an entry past the end of a shorter vec) is + /// an unfiltered measure, so an empty vec is the pre-#466 shape. + #[serde(default)] + filters: Vec>>, #[serde(default)] having: Option>, child: Rc>, @@ -1791,6 +1807,7 @@ fn default_proj_name(expr: &QueryExpr, idx: usize, schema: &Schema) -> mod tests { use super::*; use crate::pre_asap::expr_ir::{ArithmeticOpKind, CompareOpKind}; + use crate::types::AccuracyTarget; fn col(name: &str, dtype: DataType, nullable: bool) -> Column { Column::new(name, dtype, nullable) @@ -2222,6 +2239,7 @@ mod tests { reduction: Reduction::Reduce(GroupKeys::without(vec![2])), // exclude `instance` measures: vec![AggIntent::Sum { col: None }], output_names: vec![], + filters: vec![], having: None, child: Rc::new(scan_node), }; @@ -2300,6 +2318,7 @@ mod tests { reduction: Reduction::PerEntity, measures: vec![AggIntent::Rate], output_names: vec![], + filters: vec![], having: None, child: Rc::new(QueryExpr::TimeRange { range: Duration::from_secs(300), @@ -2338,6 +2357,7 @@ mod tests { reduction: Reduction::PerEntity, measures: vec![AggIntent::Avg { col: None }], output_names: vec![], + filters: vec![], having: None, child: Rc::new(QueryExpr::TimeRange { range: Duration::from_secs(300), @@ -2387,6 +2407,7 @@ mod tests { reduction: Reduction::PerEntity, measures: vec![AggIntent::Rate], output_names: vec![], + filters: vec![], having: None, child: Rc::new(open_leaf), }; @@ -2399,6 +2420,7 @@ mod tests { reduction: Reduction::by(vec![2]), // `job` measures: vec![AggIntent::Sum { col: None }], output_names: vec![], + filters: vec![], having: None, child: Rc::new(rate), }; @@ -2571,6 +2593,57 @@ mod tests { /// modifier survives unchanged alongside it — the relational binary-op /// path (issue #220's Instance 2, left as follow-up) is untouched by the /// Instance-1 `PromqlScalar` → `PromqlScalarBridge` collapse. + // `filters` (#466) round-trips, and an `Aggregate` serialized before the + // field existed still deserializes as unfiltered. + #[test] + fn aggregate_filters_serde_round_trip_and_default() { + let child = Rc::new(scan( + vec![ + col("service", DataType::Utf8, false), + col("latency", DataType::Float64, false), + ], + None, + vec![], + )); + let filtered = QueryExpr::Aggregate { + reduction: Reduction::by(vec![0]), + measures: vec![ + AggIntent::Count { + accuracy: AccuracyTarget::Exact, + }, + AggIntent::Sum { col: Some(1) }, + ], + output_names: vec![], + filters: vec![ + Some(Predicate(Rc::new(QueryExpr::Compare { + left: Rc::new(QueryExpr::Column(1)), + op: CompareOpKind::Gt, + right: Rc::new(QueryExpr::Literal(ScalarValue::Float64(1.0))), + }))), + None, + ], + having: None, + child: Rc::clone(&child), + }; + let json = serde_json::to_value(&filtered).unwrap(); + assert_eq!( + serde_json::from_value::(json.clone()).unwrap(), + filtered + ); + + let mut legacy = json; + legacy["Aggregate"] + .as_object_mut() + .unwrap() + .remove("filters") + .expect("fixture sanity: filters was serialized"); + let decoded: QueryExpr = serde_json::from_value(legacy).unwrap(); + let QueryExpr::Aggregate { filters, .. } = &decoded else { + unreachable!() + }; + assert!(filters.is_empty()); + } + #[test] fn binary_op_schema_follows_the_vector_side_over_a_scalar_bridge_with_vector_match_intact() { let vector = scan( diff --git a/crates/types/src/pre_asap/resolve.rs b/crates/types/src/pre_asap/resolve.rs index bec0331d..25131942 100644 --- a/crates/types/src/pre_asap/resolve.rs +++ b/crates/types/src/pre_asap/resolve.rs @@ -58,8 +58,8 @@ use super::column_resolution::{ }; use super::expr_ir::ColumnRef; use super::query_expr::{ - aggregate_output_schema, ConcatDiscriminatorKey, GroupKeys, Predicate, ProjectItem, - QueryExprError, Reduction, ResolvedQueryExpr, SortKey, UnresolvedQueryExpr, + aggregate_output_schema, any_measure_filtered, ConcatDiscriminatorKey, GroupKeys, Predicate, + ProjectItem, QueryExprError, Reduction, ResolvedQueryExpr, SortKey, UnresolvedQueryExpr, }; use super::schema::{ColumnId, Schema}; use super::schema_resolver::SchemaResolver; @@ -202,6 +202,7 @@ fn resolve( reduction, measures, output_names, + filters, having, child, } => { @@ -212,6 +213,23 @@ fn resolve( .iter() .map(|m| resolve_agg_intent(m, &child_schema)) .collect::, ResolveError>>()?; + // A measure filter reads the rows being aggregated, so it binds + // against the child's schema, not the aggregate's output. + let filters = filters + .iter() + .map(|f| { + f.as_ref() + .map(|Predicate(p)| Ok(Predicate(Rc::new(resolve_expr(p, &child_schema)?)))) + .transpose() + }) + .collect::, ResolveError>>()?; + // One canonical spelling of "unfiltered" (empty), so structural + // equality and CSE never split on `[]` versus `[None, None]`. + let filters = if any_measure_filtered(&filters) { + filters + } else { + Vec::new() + }; let having = having .as_ref() .map(|Predicate(h)| -> Result { @@ -228,6 +246,7 @@ fn resolve( reduction, measures, output_names: output_names.clone(), + filters, having, child: Rc::new(child), } @@ -616,6 +635,68 @@ mod tests { BinaryOpKind, QueryExpr, Source, VectorMatch, VectorMatchKind, }; + // A measure filter (#466) binds positionally against the aggregate's + // input, and a vector with no set entry collapses to the empty spelling. + #[test] + fn resolve_measure_filters_against_the_child_schema() { + use crate::pre_asap::expr_ir::ScalarValue; + use crate::pre_asap::query_expr::Predicate; + use crate::pre_asap::{Column, DataType, GroupKeys}; + use crate::types::AccuracyTarget; + let scan = || UnresolvedQueryExpr::Scan { + source: Source::Table { + table_ref: "metrics".into(), + }, + predicates: vec![], + schema: Some(Schema::new(vec![ + Column::new("service", DataType::Utf8, false), + Column::new("latency", DataType::Float64, false), + Column::new("bytes", DataType::Int64, false), + ])), + }; + let aggregate = |filters| UnresolvedQueryExpr::Aggregate { + reduction: Reduction::Reduce(GroupKeys::by(vec![ColumnRef::Named("service".into())])), + measures: vec![ + AggIntent::Count { + accuracy: AccuracyTarget::Exact, + }, + AggIntent::Sum { + col: Some(ColumnRef::Named("bytes".into())), + }, + ], + output_names: vec![], + filters, + having: None, + child: Rc::new(scan()), + }; + let latency_gt_one = Predicate(Rc::new(UnresolvedQueryExpr::Compare { + left: Rc::new(UnresolvedQueryExpr::Column(ColumnRef::Named( + "latency".into(), + ))), + op: CompareOpKind::Gt, + right: Rc::new(UnresolvedQueryExpr::Literal(ScalarValue::Float64(1.0))), + })); + + let resolved = resolve_root(&aggregate(vec![Some(latency_gt_one), None])).unwrap(); + let QueryExpr::Aggregate { filters, .. } = &resolved else { + unreachable!() + }; + let [Some(Predicate(first)), None] = filters.as_slice() else { + panic!("expected one filtered and one unfiltered measure, got {filters:?}"); + }; + assert!( + matches!(first.as_ref(), QueryExpr::Compare { left, .. } + if matches!(left.as_ref(), QueryExpr::Column(1))), + "latency is input column 1, got {first:?}" + ); + + let resolved = resolve_root(&aggregate(vec![None, None])).unwrap(); + let QueryExpr::Aggregate { filters, .. } = &resolved else { + unreachable!() + }; + assert!(filters.is_empty()); + } + // Both sides resolve with qualifiers; an unknown right input is an error. #[test] fn resolve_pearson_corr_inputs() { diff --git a/crates/types/src/pre_asap/schema_resolver.rs b/crates/types/src/pre_asap/schema_resolver.rs index 8244d9be..b8d23b6a 100644 --- a/crates/types/src/pre_asap/schema_resolver.rs +++ b/crates/types/src/pre_asap/schema_resolver.rs @@ -227,6 +227,7 @@ pub(crate) fn collect_referenced_columns(tree: &UnresolvedQueryExpr) -> Vec Vec