feat(flow): admit mergeable average states in incremental plans

Signed-off-by: discord9 <discord9@163.com>
This commit is contained in:
discord9
2026-09-09 12:48:26 +08:00
parent 621412080e
commit 1d405efe8b
2 changed files with 186 additions and 6 deletions
+18 -2
View File
@@ -21,6 +21,7 @@ use catalog::CatalogManagerRef;
use common_error::ext::BoxedError;
use common_function::aggrs::aggr_wrapper::get_aggr_func;
use common_telemetry::debug;
use datafusion::arrow::datatypes::DataType as ArrowDataType;
use datafusion::datasource::DefaultTableSource;
use datafusion::error::Result as DfResult;
use datafusion::logical_expr::Expr;
@@ -326,7 +327,10 @@ fn is_literal_or_cast_literal(expr: &Expr) -> bool {
}
}
fn merge_op_for_aggregate_expr(aggr_expr: &Expr) -> Result<IncrementalAggregateMergeOp, String> {
fn merge_op_for_aggregate_expr(
aggr_expr: &Expr,
input_schema: &DFSchema,
) -> Result<IncrementalAggregateMergeOp, String> {
let Some(aggr_func) = get_aggr_func(aggr_expr) else {
return Err(aggr_expr.to_string());
};
@@ -356,6 +360,9 @@ fn merge_op_for_aggregate_expr(aggr_expr: &Expr) -> Result<IncrementalAggregateM
"bit_or" => Ok(IncrementalAggregateMergeOp::BitOr),
"bit_xor" => Ok(IncrementalAggregateMergeOp::BitXor),
// Preserve state-family parameters; value coercion is handled by the aggregate.
"avg_state" if aggr_func.params.args.len() == 1 => {
state_delta_merge("__avg_state_delta_merge", vec![])
}
"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![])
@@ -374,6 +381,15 @@ fn merge_op_for_aggregate_expr(aggr_expr: &Expr) -> Result<IncrementalAggregateM
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
&& aggr_func.params.args[0].get_type(input_schema).ok()
== Some(ArrowDataType::Binary) =>
{
state_delta_merge("__avg_state_delta_merge", vec![])
}
_ => Err(aggr_expr.to_string()),
}
}
@@ -520,7 +536,7 @@ pub fn analyze_incremental_aggregate_plan(
&group_key_names,
));
for aggr_expr in aggr_exprs {
let merge_op = match merge_op_for_aggregate_expr(&aggr_expr) {
let merge_op = match merge_op_for_aggregate_expr(&aggr_expr, aggregate.input.schema()) {
Ok(merge_op) => merge_op,
Err(reason) => {
unsupported_exprs.push(reason);
+168 -4
View File
@@ -16,6 +16,7 @@ use std::collections::BTreeMap;
use std::sync::Arc;
use catalog::RegisterTableRequest;
use common_function::aggrs::approximate::avg::AvgState;
use common_query::OutputData;
use common_recordbatch::recordbatch::merge_record_batches;
use common_recordbatch::{RecordBatch, util};
@@ -1766,10 +1767,154 @@ async fn test_analyze_incremental_aggregate_plan_rejects_avg() {
assert!(!analysis.unsupported_exprs.is_empty());
}
#[tokio::test]
async fn test_analyze_incremental_aggregate_plan_supports_avg_state() {
let query_engine = create_test_query_engine();
let ctx = QueryContext::arc();
let sql = "SELECT avg_state(number) AS avg_num, ts FROM numbers_with_ts GROUP BY ts";
let plan = sql_to_df_plan(ctx, query_engine, sql, false).await.unwrap();
let analysis = analyze_incremental_aggregate_plan(&plan).unwrap().unwrap();
assert!(
analysis.unsupported_exprs.is_empty(),
"avg_state should be supported: {:?}",
analysis.unsupported_exprs
);
assert_eq!(analysis.merge_columns.len(), 1);
assert_eq!(analysis.merge_columns[0].output_field_name, "avg_num");
assert!(matches!(
&analysis.merge_columns[0].merge_op,
IncrementalAggregateMergeOp::StateDeltaMerge {
function_name: "__avg_state_delta_merge",
params,
} if params.is_empty()
));
}
#[tokio::test]
async fn test_analyze_incremental_aggregate_plan_supports_avg_merge() {
let query_engine = create_test_query_engine();
let producer_sql = "SELECT avg_state(number) AS avg_state, CASE WHEN number <= 5 THEN 1 ELSE 2 END AS grp FROM numbers_with_ts GROUP BY grp";
let producer_plan = sql_to_df_plan(
QueryContext::arc(),
query_engine.clone(),
producer_sql,
false,
)
.await
.unwrap();
let producer_output = query_engine
.execute(producer_plan, QueryContext::arc())
.await
.unwrap();
let OutputData::Stream(producer_stream) = producer_output.data else {
panic!("expected AVG state producer execution to be a stream");
};
let producer_batches = util::collect(producer_stream).await.unwrap();
let producer_schema = producer_batches.first().unwrap().schema.clone();
let avg_states = merge_record_batches(producer_schema, &producer_batches).unwrap();
assert_eq!(avg_states.num_rows(), 2);
let avg_states_table = MemTable::table("avg_states", avg_states);
query_engine
.engine_state()
.catalog_manager()
.as_any()
.downcast_ref::<catalog::memory::MemoryCatalogManager>()
.unwrap()
.register_table_sync(RegisterTableRequest {
catalog: "greptime".to_string(),
schema: "public".to_string(),
table_name: "avg_states".to_string(),
table_id: 4096,
table: avg_states_table,
})
.unwrap();
let sql = "SELECT avg_merge(avg_state) AS avg_num, grp FROM avg_states GROUP BY grp";
let plan = sql_to_df_plan(QueryContext::arc(), query_engine.clone(), sql, false)
.await
.unwrap();
let analysis = analyze_incremental_aggregate_plan(&plan).unwrap().unwrap();
assert!(
analysis.unsupported_exprs.is_empty(),
"avg_merge should be supported: {:?}",
analysis.unsupported_exprs
);
assert_eq!(analysis.merge_columns.len(), 1);
assert!(matches!(
&analysis.merge_columns[0].merge_op,
IncrementalAggregateMergeOp::StateDeltaMerge {
function_name: "__avg_state_delta_merge",
params,
} if params.is_empty()
));
let output = query_engine
.execute(plan, QueryContext::arc())
.await
.unwrap();
let OutputData::Stream(stream) = output.data else {
panic!("expected AVG merge execution to be a stream");
};
let batches = util::collect(stream).await.unwrap();
let schema = batches.first().unwrap().schema.clone();
let batch = merge_record_batches(schema, &batches).unwrap();
let groups = batch
.column_by_name("grp")
.unwrap()
.as_primitive::<Int64Type>();
let averages = batch.column_by_name("avg_num").unwrap().as_binary::<i32>();
let mut values = (0..batch.num_rows())
.map(|index| {
(
groups.value(index),
AvgState::decode(averages.value(index)).unwrap().average(),
)
})
.collect::<Vec<_>>();
values.sort_unstable_by_key(|(group, _)| *group);
assert_eq!(values, [(1, Some(3.0)), (2, Some(8.0))]);
}
#[tokio::test]
async fn test_analyze_incremental_aggregate_plan_supports_duplicate_avg_projections() {
let analysis = analyze_test_sql(
"SELECT avg_state(number) AS avg_a, avg_state(number) AS avg_b, avg_state(number + 1) AS avg_num_plus, ts FROM numbers_with_ts GROUP BY ts",
)
.await;
assert!(analysis.unsupported_exprs.is_empty());
assert_eq!(analysis.merge_columns.len(), 3);
assert!(analysis.merge_columns.iter().all(|column| {
matches!(
&column.merge_op,
IncrementalAggregateMergeOp::StateDeltaMerge {
function_name: "__avg_state_delta_merge",
params,
} if params.is_empty()
)
}));
assert_eq!(analysis.merge_columns[0].output_field_name, "avg_a");
assert_eq!(analysis.merge_columns[0].input_field_name, "avg_a");
let avg_b = analysis
.merge_columns
.iter()
.find(|column| column.output_field_name == "avg_b")
.unwrap();
assert_eq!(avg_b.input_field_name, "avg_a");
assert!(
analysis
.merge_columns
.iter()
.any(|column| column.output_field_name == "avg_num_plus")
);
}
#[tokio::test]
async fn test_analyze_incremental_aggregate_plan_supports_mixed_state_families() {
let analysis = analyze_test_sql(
"SELECT hll(CAST(number AS VARCHAR)) AS hll_a, \
"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, \
@@ -1783,7 +1928,7 @@ async fn test_analyze_incremental_aggregate_plan_supports_mixed_state_families()
"mixed state aggregate should be supported: {:?}",
analysis.unsupported_exprs
);
assert_eq!(analysis.merge_columns.len(), 6);
assert_eq!(analysis.merge_columns.len(), 7);
assert!(analysis.merge_columns.iter().any(|column| {
column.output_field_name == "hll_a"
&& column.merge_op
@@ -1824,8 +1969,8 @@ async fn test_analyze_incremental_aggregate_plan_supports_mixed_state_families()
#[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 CAST(NULL AS BIGINT) 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 CAST(NULL AS BIGINT) 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_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.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 CAST(NULL AS BIGINT) ELSE 2 END AS grp FROM numbers_with_ts WHERE number <= 3 OR number = 6 GROUP BY grp";
let new_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.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 CAST(NULL AS BIGINT) 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();
@@ -1880,6 +2025,7 @@ async fn test_rewrite_incremental_aggregate_merges_populated_mixed_state_familie
.unwrap();
let rendered = format!("{}", rewritten.display_indent());
for function_name in [
"__avg_state_delta_merge",
"__hll_delta_merge",
"__uddsketch_state_delta_merge",
"__stddev_pop_state_delta_merge",
@@ -1889,6 +2035,7 @@ async fn test_rewrite_incremental_aggregate_merges_populated_mixed_state_familie
assert_eq!(
analysis.output_field_names,
vec![
"avg_num",
"hll_a",
"hll_b",
"percentile_a",
@@ -1925,6 +2072,23 @@ async fn test_rewrite_incremental_aggregate_merges_populated_mixed_state_familie
.collect::<Vec<_>>();
merged_group_values.sort_unstable();
assert_eq!(merged_group_values, [None, Some(2), Some(3)]);
let avg_states = merged_batch
.column_by_name("avg_num")
.unwrap()
.as_binary::<i32>();
assert_eq!(avg_states.null_count(), 0);
for index in 0..merged_batch.num_rows() {
let expected = match (!merged_groups.is_null(index)).then(|| merged_groups.value(index)) {
None => Some(3.0),
Some(2) => Some(7.0),
Some(3) => Some(9.5),
group => panic!("unexpected group: {group:?}"),
};
assert_eq!(
AvgState::decode(avg_states.value(index)).unwrap().average(),
expected
);
}
let merged_table = MemTable::table("merged_states", merged_batch);
query_engine
.engine_state()