fix(flow): merge persisted sketch states in incremental batches

Signed-off-by: discord9 <discord9@163.com>
This commit is contained in:
discord9
2026-09-07 12:59:53 +08:00
parent dd845cfbc9
commit 2861ca6a75
2 changed files with 361 additions and 28 deletions
+81 -27
View File
@@ -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<Expr>,
},
}
/// 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<Expr, Error> {
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<Expr, Error> {
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(),
}
+280 -1
View File
@@ -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::<Int64Type>();
assert_eq!(old_groups.null_count(), 0);
let mut old_group_values = (0..old_groups.len())
.map(|index| old_groups.value(index))
.collect::<Vec<_>>();
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::<Int64Type>();
assert_eq!(merged_groups.null_count(), 0);
let mut merged_group_values = (0..merged_groups.len())
.map(|index| merged_groups.value(index))
.collect::<Vec<_>>();
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::<catalog::memory::MemoryCatalogManager>()
.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<_>>(),
vec!["grp", "total", "hll_a", "hll_b", "stddev", "p50_a", "p50_b"]
);
let group = checks.column(0).as_primitive::<Int64Type>();
let total = checks.column(1).as_primitive::<UInt64Type>();
let hll_a = checks.column(2).as_primitive::<UInt64Type>();
let hll_b = checks.column(3).as_primitive::<UInt64Type>();
let stddev = checks.column(4).as_primitive::<Float64Type>();
let p50_a = checks.column(5).as_primitive::<Float64Type>();
let p50_b = checks.column(6).as_primitive::<Float64Type>();
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();