diff --git a/src/flow/src/batching_mode/utils.rs b/src/flow/src/batching_mode/utils.rs index c09fc1ff40..fd22e42f22 100644 --- a/src/flow/src/batching_mode/utils.rs +++ b/src/flow/src/batching_mode/utils.rs @@ -88,7 +88,7 @@ impl IncrementalAggregateMergeColumn { } } -#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[derive(Debug, Clone, PartialEq, Eq)] pub enum IncrementalAggregateMergeOp { Sum, Min, @@ -99,6 +99,10 @@ pub enum IncrementalAggregateMergeOp { BitOr, BitXor, AvgDeltaMerge, + StateDeltaMerge { + function_name: &'static str, + params: Vec, + }, } /// Analysis result for an incremental aggregate plan. @@ -351,6 +355,14 @@ fn merge_op_for_aggregate_expr( return Err(format!("unsupported aggregate NULL treatment: {aggr_expr}")); } + let state_delta_merge = |function_name, params| { + Ok(IncrementalAggregateMergeOp::StateDeltaMerge { + function_name, + params, + }) + }; + let is_type = |expr: &Expr, data_type| expr.get_type(input_schema).ok() == Some(data_type); + match aggr_func.func.name().to_ascii_lowercase().as_str() { "sum" | "count" => Ok(IncrementalAggregateMergeOp::Sum), "min" => Ok(IncrementalAggregateMergeOp::Min), @@ -360,16 +372,36 @@ fn merge_op_for_aggregate_expr( "bit_and" => Ok(IncrementalAggregateMergeOp::BitAnd), "bit_or" => Ok(IncrementalAggregateMergeOp::BitOr), "bit_xor" => Ok(IncrementalAggregateMergeOp::BitXor), - "avg_state" => match aggr_func.params.args.as_slice() { - [_] => Ok(IncrementalAggregateMergeOp::AvgDeltaMerge), - _ => Err(aggr_expr.to_string()), - }, - "avg_merge" => match aggr_func.params.args.as_slice() { - [arg] if arg.get_type(input_schema).ok() == Some(ArrowDataType::Binary) => { - Ok(IncrementalAggregateMergeOp::AvgDeltaMerge) + // Preserve state-family parameters; value coercion is handled by the aggregate. + "avg_state" if aggr_func.params.args.len() == 1 => { + Ok(IncrementalAggregateMergeOp::AvgDeltaMerge) + } + "hll" if aggr_func.params.args.len() == 1 => state_delta_merge("__hll_delta_merge", vec![]), + "stddev_pop_state" if aggr_func.params.args.len() == 1 => { + state_delta_merge("__stddev_pop_state_delta_merge", vec![]) + } + "uddsketch_state" if aggr_func.params.args.len() == 3 => { + let [bucket_size, error_rate, _] = aggr_func.params.args.as_slice() else { + unreachable!(); + }; + if !matches!(bucket_size, Expr::Literal(ScalarValue::Int64(Some(_)), _)) + || !matches!(error_rate, Expr::Literal(ScalarValue::Float64(Some(_)), _)) + { + return Err(aggr_expr.to_string()); } - _ => Err(aggr_expr.to_string()), - }, + state_delta_merge( + "__uddsketch_state_delta_merge", + vec![bucket_size.clone(), error_rate.clone()], + ) + } + // AVG's binary merge form is admitted because its state argument is + // already the aggregate result stored by the sink. + "avg_merge" + if aggr_func.params.args.len() == 1 + && is_type(&aggr_func.params.args[0], ArrowDataType::Binary) => + { + Ok(IncrementalAggregateMergeOp::AvgDeltaMerge) + } _ => Err(aggr_expr.to_string()), } } @@ -547,7 +579,7 @@ pub fn analyze_incremental_aggregate_plan( merge_columns.push(IncrementalAggregateMergeColumn { input_field_name: input_field_name.clone(), output_field_name, - merge_op, + merge_op: merge_op.clone(), }); } } @@ -657,10 +689,13 @@ pub async fn rewrite_incremental_aggregate_with_sink_merge( let delta_alias = "__flow_delta"; let sink_alias = "__flow_sink"; - let state_merge = analysis - .merge_columns - .iter() - .any(|column| matches!(column.merge_op, IncrementalAggregateMergeOp::AvgDeltaMerge)); + let state_merge = analysis.merge_columns.iter().any(|column| { + matches!( + &column.merge_op, + IncrementalAggregateMergeOp::AvgDeltaMerge + | IncrementalAggregateMergeOp::StateDeltaMerge { .. } + ) + }); let mut selected_columns = analysis.group_key_names.clone(); selected_columns.extend( analysis @@ -804,8 +839,9 @@ pub async fn rewrite_incremental_aggregate_with_sink_merge( group_exprs.push(expr); } else if let Some(merge_col) = merge_columns.get(output_field_name) { if matches!( - merge_col.merge_op, + &merge_col.merge_op, IncrementalAggregateMergeOp::AvgDeltaMerge + | IncrementalAggregateMergeOp::StateDeltaMerge { .. } ) { state_aggr_exprs.push(build_state_delta_merge_expr(engine, merge_col)?); } else { @@ -866,29 +902,46 @@ fn build_state_delta_merge_expr( engine: &QueryEngineRef, merge_col: &IncrementalAggregateMergeColumn, ) -> Result { + let (function_name, params) = match &merge_col.merge_op { + IncrementalAggregateMergeOp::AvgDeltaMerge => ("__avg_state_delta_merge", vec![]), + IncrementalAggregateMergeOp::StateDeltaMerge { + function_name, + params, + } => (*function_name, params.clone()), + _ => { + return InvalidQuerySnafu { + reason: "non-state aggregate passed to state delta merge".to_string(), + } + .fail(); + } + }; let Some(udaf) = engine .engine_state() - .aggr_function("__avg_state_delta_merge") + .aggr_function(function_name) .or_else(|| { engine .engine_state() .session_state() .aggregate_functions() - .get("__avg_state_delta_merge") + .get(function_name) .map(|udaf| udaf.as_ref().clone()) }) else { return InvalidQuerySnafu { - reason: "Aggregate function __avg_state_delta_merge is not registered".to_string(), + reason: format!("Aggregate function {function_name} is not registered"), } .fail(); }; - Ok(udaf - .call(vec![ - qualified_col("__flow_delta", merge_col.input_field_name.clone()), - qualified_col("__flow_sink", merge_col.output_field_name.clone()), - ]) - .alias(merge_col.output_field_name.clone())) + let mut args = params; + args.push(qualified_col( + "__flow_delta", + merge_col.input_field_name.clone(), + )); + args.push(qualified_col( + "__flow_sink", + merge_col.output_field_name.clone(), + )); + Ok(udaf.call(args).alias(merge_col.output_field_name.clone())) } fn build_left_join_merge_expr( @@ -898,7 +951,7 @@ fn build_left_join_merge_expr( ) -> Result { let left = qualified_col(delta_alias, merge_col.input_field_name.clone()); let right = qualified_col(sink_alias, merge_col.output_field_name.clone()); - let merged = match merge_col.merge_op { + let merged = match merge_col.merge_op.clone() { IncrementalAggregateMergeOp::Sum => when(is_null(left.clone()), right.clone()) .when(is_null(right.clone()), left.clone()) .otherwise(binary_expr(left.clone(), Operator::Plus, right.clone())) @@ -947,7 +1000,8 @@ fn build_left_join_merge_expr( .with_context(|_| DatafusionSnafu { context: "Failed to build BIT_XOR merge expression".to_string(), })?, - IncrementalAggregateMergeOp::AvgDeltaMerge => { + IncrementalAggregateMergeOp::AvgDeltaMerge + | IncrementalAggregateMergeOp::StateDeltaMerge { .. } => { return InvalidQuerySnafu { reason: "state aggregate must be built with its delta UDAF".to_string(), } diff --git a/src/flow/src/batching_mode/utils/test.rs b/src/flow/src/batching_mode/utils/test.rs index 929b0f9da0..7bc9302986 100644 --- a/src/flow/src/batching_mode/utils/test.rs +++ b/src/flow/src/batching_mode/utils/test.rs @@ -16,10 +16,14 @@ use std::collections::BTreeMap; use std::sync::Arc; use catalog::RegisterTableRequest; -use common_recordbatch::RecordBatch; +use common_query::OutputData; +use common_recordbatch::recordbatch::merge_record_batches; +use common_recordbatch::{RecordBatch, util}; use common_time::Timestamp; use datafusion_common::tree_node::TreeNode as _; use datafusion_expr::GroupingSet; +use datatypes::arrow::array::{Array, AsArray}; +use datatypes::arrow::datatypes::{Float64Type, Int64Type, UInt64Type}; use datatypes::prelude::{ConcreteDataType, MutableVector, Scalar, ScalarVectorBuilder, VectorRef}; use datatypes::schema::{ColumnSchema, Schema}; use datatypes::timestamp::TimestampMillisecond; @@ -1878,6 +1882,281 @@ async fn test_analyze_incremental_aggregate_plan_supports_avg_with_native_aggreg })); } +#[tokio::test] +async fn test_analyze_incremental_aggregate_plan_supports_mixed_state_families() { + let analysis = analyze_test_sql( + "SELECT avg_state(number) AS avg_num, \ + hll(CAST(number AS VARCHAR)) AS hll_a, \ + hll(CAST(number AS VARCHAR)) AS hll_b, \ + uddsketch_state(128, 0.01, CAST(number AS DOUBLE)) AS percentile_a, \ + uddsketch_state(256, 0.02, number) AS percentile_b, \ + stddev_pop_state(number) AS stddev_state, \ + sum(number) AS total, ts FROM numbers_with_ts GROUP BY ts", + ) + .await; + + assert!( + analysis.unsupported_exprs.is_empty(), + "mixed state aggregate should be supported: {:?}", + analysis.unsupported_exprs + ); + assert_eq!(analysis.merge_columns.len(), 7); + assert!(analysis.merge_columns.iter().any(|column| { + column.output_field_name == "hll_a" + && column.merge_op + == (IncrementalAggregateMergeOp::StateDeltaMerge { + function_name: "__hll_delta_merge", + params: vec![], + }) + })); + assert!(analysis.merge_columns.iter().any(|column| { + column.output_field_name == "hll_b" + && column.input_field_name == "hll_a" + && column.merge_op + == (IncrementalAggregateMergeOp::StateDeltaMerge { + function_name: "__hll_delta_merge", + params: vec![], + }) + })); + for (output_field_name, function_name, param_count) in [ + ("percentile_a", "__uddsketch_state_delta_merge", 2), + ("percentile_b", "__uddsketch_state_delta_merge", 2), + ("stddev_state", "__stddev_pop_state_delta_merge", 0), + ] { + let column = analysis + .merge_columns + .iter() + .find(|column| column.output_field_name == output_field_name) + .unwrap(); + assert!(matches!( + &column.merge_op, + IncrementalAggregateMergeOp::StateDeltaMerge { + function_name: actual_name, + params, + } if *actual_name == function_name && params.len() == param_count + )); + } +} + +#[tokio::test] +async fn test_rewrite_incremental_aggregate_merges_populated_mixed_state_families() { + let query_engine = create_test_query_engine(); + let old_sql = "SELECT hll(CAST(number AS VARCHAR)) AS hll_a, hll(CAST(number AS VARCHAR)) AS hll_b, uddsketch_state(128, 0.000001, number) AS percentile_a, uddsketch_state(256, 0.02, number) AS percentile_b, stddev_pop_state(number) AS stddev_state, sum(number) AS total, CASE WHEN number <= 5 THEN 1 ELSE 2 END AS grp FROM numbers_with_ts WHERE number <= 3 OR number = 6 GROUP BY grp"; + let new_sql = "SELECT hll(CAST(number AS VARCHAR)) AS hll_a, hll(CAST(number AS VARCHAR)) AS hll_b, uddsketch_state(128, 0.000001, number) AS percentile_a, uddsketch_state(256, 0.02, number) AS percentile_b, stddev_pop_state(number) AS stddev_state, sum(number) AS total, CASE WHEN number <= 5 THEN 1 WHEN number <= 8 THEN 2 ELSE 3 END AS grp FROM numbers_with_ts WHERE number >= 4 AND number != 6 GROUP BY grp"; + let old_plan = sql_to_df_plan(QueryContext::arc(), query_engine.clone(), old_sql, false) + .await + .unwrap(); + let old_output = query_engine + .execute(old_plan, QueryContext::arc()) + .await + .unwrap(); + let OutputData::Stream(old_stream) = old_output.data else { + panic!("expected old aggregate execution to be a stream"); + }; + let old_batches = util::collect(old_stream).await.unwrap(); + let old_schema = old_batches.first().unwrap().schema.clone(); + let old_batch = merge_record_batches(old_schema, &old_batches).unwrap(); + assert_eq!( + old_batch.num_rows(), + 2, + "old state must have one row per group" + ); + let old_groups = old_batch + .column_by_name("grp") + .unwrap() + .as_primitive::(); + assert_eq!(old_groups.null_count(), 0); + let mut old_group_values = (0..old_groups.len()) + .map(|index| old_groups.value(index)) + .collect::>(); + old_group_values.sort_unstable(); + assert_eq!(old_group_values, [1, 2]); + let sink_table = MemTable::table("state_merge_sink", old_batch); + let sink_table_name = [ + "greptime".to_string(), + "public".to_string(), + "state_merge_sink".to_string(), + ]; + + let new_plan = sql_to_df_plan(QueryContext::arc(), query_engine.clone(), new_sql, false) + .await + .unwrap(); + let analysis = analyze_incremental_aggregate_plan(&new_plan) + .unwrap() + .unwrap(); + assert!(analysis.unsupported_exprs.is_empty()); + let rewritten = rewrite_incremental_aggregate_with_sink_merge( + &new_plan, + &analysis, + &query_engine, + sink_table, + &sink_table_name, + None, + ) + .await + .unwrap(); + let rendered = format!("{}", rewritten.display_indent()); + for function_name in [ + "__hll_delta_merge", + "__uddsketch_state_delta_merge", + "__stddev_pop_state_delta_merge", + ] { + assert!(rendered.contains(function_name), "{rendered}"); + } + assert_eq!( + analysis.output_field_names, + vec![ + "hll_a", + "hll_b", + "percentile_a", + "percentile_b", + "stddev_state", + "total", + "grp", + ], + "repeated HLL aliases must preserve output order" + ); + + let output = query_engine + .execute(rewritten, QueryContext::arc()) + .await + .unwrap(); + let OutputData::Stream(stream) = output.data else { + panic!("expected rewritten plan to execute as a stream"); + }; + let batches = util::collect(stream).await.unwrap(); + let merged_schema = batches.first().unwrap().schema.clone(); + let merged_batch = merge_record_batches(merged_schema, &batches).unwrap(); + assert_eq!( + merged_batch.num_rows(), + 3, + "rewrite must produce one row per group" + ); + let merged_groups = merged_batch + .column_by_name("grp") + .unwrap() + .as_primitive::(); + assert_eq!(merged_groups.null_count(), 0); + let mut merged_group_values = (0..merged_groups.len()) + .map(|index| merged_groups.value(index)) + .collect::>(); + merged_group_values.sort_unstable(); + assert_eq!(merged_group_values, [1, 2, 3]); + let merged_table = MemTable::table("merged_states", merged_batch); + query_engine + .engine_state() + .catalog_manager() + .as_any() + .downcast_ref::() + .unwrap() + .register_table_sync(RegisterTableRequest { + catalog: "greptime".to_string(), + schema: "public".to_string(), + table_name: "merged_states".to_string(), + table_id: 4097, + table: merged_table, + }) + .unwrap(); + let checks = "SELECT grp, sum(total) AS total, hll_count(hll_merge(hll_a)) AS hll_a, hll_count(hll_merge(hll_b)) AS hll_b, stddev_pop_calc(stddev_pop_merge(stddev_state)) AS stddev, uddsketch_calc(0.5, uddsketch_merge(128, 0.000001, percentile_a)) AS p50_a, uddsketch_calc(0.5, uddsketch_merge(256, 0.02, percentile_b)) AS p50_b FROM merged_states GROUP BY grp ORDER BY grp"; + let checks_plan = sql_to_df_plan(QueryContext::arc(), query_engine.clone(), checks, false) + .await + .unwrap(); + let checks_output = query_engine + .execute(checks_plan, QueryContext::arc()) + .await + .unwrap(); + let OutputData::Stream(checks_stream) = checks_output.data else { + panic!("expected state check execution to be a stream"); + }; + let checks_batches = util::collect(checks_stream).await.unwrap(); + let checks_schema = checks_batches.first().unwrap().schema.clone(); + let checks = merge_record_batches(checks_schema, &checks_batches).unwrap(); + assert_eq!(checks.num_rows(), 3); + assert_eq!( + checks + .schema + .column_schemas() + .iter() + .map(|column| column.name.as_str()) + .collect::>(), + vec!["grp", "total", "hll_a", "hll_b", "stddev", "p50_a", "p50_b"] + ); + let group = checks.column(0).as_primitive::(); + let total = checks.column(1).as_primitive::(); + let hll_a = checks.column(2).as_primitive::(); + let hll_b = checks.column(3).as_primitive::(); + let stddev = checks.column(4).as_primitive::(); + let p50_a = checks.column(5).as_primitive::(); + let p50_b = checks.column(6).as_primitive::(); + assert_eq!(group.null_count(), 0); + assert_eq!(total.null_count(), 0); + assert_eq!(hll_a.null_count(), 0); + assert_eq!(hll_b.null_count(), 0); + assert_eq!(stddev.null_count(), 0); + assert_eq!(p50_a.null_count(), 0); + assert_eq!(p50_b.null_count(), 0); + for expected in [ + (1_i64, 15_u64, 5_u64, 2_f64.sqrt(), 3_f64, 3_f64), + (2_i64, 21_u64, 3_u64, (2_f64 / 3.0).sqrt(), 7_f64, 7_f64), + (3_i64, 19_u64, 2_u64, 0.5_f64, 10_f64, 10_f64), + ] { + let index = (expected.0 - 1) as usize; + assert_eq!(group.value(index), expected.0); + assert_eq!(total.value(index), expected.1); + assert_eq!(hll_a.value(index), expected.2); + assert_eq!(hll_b.value(index), expected.2); + assert!((stddev.value(index) - expected.3).abs() < 1e-12); + // UDDSketch's relative-error bounds are 1e-6 and 2e-2 respectively. + assert!((p50_a.value(index) - expected.4).abs() <= expected.4 * 0.000001); + assert!((p50_b.value(index) - expected.5).abs() <= expected.5 * 0.02); + } +} + +#[tokio::test] +async fn test_analyze_incremental_aggregate_plan_state_producer_metadata_and_rejections() { + let analysis = analyze_test_sql( + "SELECT hll(CAST(number AS VARCHAR)) AS hll_state, \ + stddev_pop_state(number) AS stddev_state, \ + uddsketch_state(128, 0.000001, number) AS percentile_state, ts \ + FROM numbers_with_ts GROUP BY ts", + ) + .await; + assert!(analysis.unsupported_exprs.is_empty()); + let percentile = analysis + .merge_columns + .iter() + .find(|column| column.output_field_name == "percentile_state") + .unwrap(); + let IncrementalAggregateMergeOp::StateDeltaMerge { + function_name, + params, + } = &percentile.merge_op + else { + panic!("expected UDDSketch state delta merge"); + }; + assert_eq!(*function_name, "__uddsketch_state_delta_merge"); + assert_eq!(params.len(), 2); + assert!(matches!( + params[0], + Expr::Literal(ScalarValue::Int64(Some(128)), _) + )); + assert!(matches!( + params[1], + Expr::Literal(ScalarValue::Float64(Some(rate)), _) if rate == 0.000001 + )); + + for sql in [ + "SELECT uddsketch_state(CAST(number AS BIGINT), 0.01, number) AS state, ts FROM numbers_with_ts GROUP BY ts", + "SELECT uddsketch_state(128, NULL, number) AS state, ts FROM numbers_with_ts GROUP BY ts", + "SELECT uddsketch_state(NULL, 0.01, number) AS state, ts FROM numbers_with_ts GROUP BY ts", + "SELECT avg(number) AS state, ts FROM numbers_with_ts GROUP BY ts", + "SELECT stddev_pop(number) AS state, ts FROM numbers_with_ts GROUP BY ts", + ] { + let analysis = analyze_test_sql(sql).await; + assert!(!analysis.unsupported_exprs.is_empty(), "must reject {sql}"); + } +} + #[tokio::test] async fn test_analyze_incremental_aggregate_plan_rejects_distinct() { let query_engine = create_test_query_engine();