diff --git a/src/common/function/src/aggrs/aggr_wrapper.rs b/src/common/function/src/aggrs/aggr_wrapper.rs index 70651eee41d..710c383a834 100644 --- a/src/common/function/src/aggrs/aggr_wrapper.rs +++ b/src/common/function/src/aggrs/aggr_wrapper.rs @@ -25,7 +25,7 @@ use std::hash::{Hash, Hasher}; use std::sync::Arc; -use arrow::array::{ArrayRef, BooleanArray, StructArray}; +use arrow::array::{ArrayData, ArrayRef, BooleanArray, StructArray, make_array}; use arrow_schema::{FieldRef, Fields}; use common_telemetry::debug; use datafusion::functions_aggregate::all_default_aggregate_functions; @@ -84,6 +84,19 @@ pub fn is_all_aggr_exprs_steppable(aggr_exprs: &[Expr]) -> bool { return false; } + // DataFusion only sorts the input of an aggregate with a hard ordering requirement + // when the requirement is already satisfied or the aggregate has a reverse + // expression (apache/datafusion#25676). The state wrapper has none, so e.g. + // `nth_value(.. ORDER BY ..)` would read unsorted input on datanodes. Ordered-set + // aggregates like `approx_percentile_cont(..) WITHIN GROUP (ORDER BY ..)` are + // exempt: their ORDER BY names the value, and they don't need sorted input. + if !aggr_func.params.order_by.is_empty() + && aggr_func.func.order_sensitivity().hard_requires() + && !aggr_func.func.supports_within_group_clause() + { + return false; + } + // whether the corresponding state function exists in the registry FUNCTION_REGISTRY.is_aggr_func_exist(&aggr_state_func_name(aggr_func.func.name())) } else { @@ -526,43 +539,7 @@ impl StateGroupsAccum { } fn wrap_state_arrays(&self, arrays: Vec) -> datafusion_common::Result { - let array_type = arrays - .iter() - .map(|array| array.data_type().clone()) - .collect::>(); - let expected_type = self - .state_fields - .iter() - .map(|field| field.data_type().clone()) - .collect::>(); - if array_type != expected_type { - debug!( - "State mismatch, expected: {}, got: {} for expected fields: {:?} and given array types: {:?}", - self.state_fields.len(), - arrays.len(), - self.state_fields, - array_type, - ); - let guess_schema = arrays - .iter() - .enumerate() - .map(|(index, array)| { - Field::new( - format!("col_{index}[mismatch_state]").as_str(), - array.data_type().clone(), - true, - ) - }) - .collect::(); - let array = StructArray::try_new(guess_schema, arrays, None)?; - return Ok(Arc::new(array)); - } - - Ok(Arc::new(StructArray::try_new( - self.state_fields.clone(), - arrays, - None, - )?)) + Ok(Arc::new(state_struct_array(&self.state_fields, arrays)?)) } } @@ -610,6 +587,81 @@ impl GroupsAccumulator for StateGroupsAccum { } } +/// Wraps the state arrays of an accumulator into a struct of the declared state fields. +/// +/// The declared state type is derived from logical expressions, while the accumulator names +/// nested fields after physical expressions. For example, `array_agg(v ORDER BY ts)` declares +/// its orderings as `List(Struct("ts": ..))` but produces `List(Struct("ts@0": ..))`. Arrays that +/// differ only in nested field names are cast to the declared type; any other difference is an +/// error. +fn state_struct_array( + state_fields: &Fields, + arrays: Vec, +) -> datafusion_common::Result { + if arrays.len() != state_fields.len() { + return Err(datafusion_common::DataFusionError::Internal(format!( + "Expected {} state arrays for fields {:?}, got {}", + state_fields.len(), + state_fields, + arrays.len() + ))); + } + let arrays = arrays + .into_iter() + .zip(state_fields.iter()) + .map(|(array, field)| { + let expected = field.data_type(); + if array.data_type() == expected { + Ok(array) + } else if array.data_type().equals_datatype(expected) { + Ok(make_array(relabel_nested_fields( + array.to_data(), + expected, + )?)) + } else { + Err(datafusion_common::DataFusionError::Internal(format!( + "State field `{}` expects type {expected}, but the accumulator produced {}", + field.name(), + array.data_type() + ))) + } + }) + .collect::>>()?; + Ok(StructArray::try_new(state_fields.clone(), arrays, None)?) +} + +/// Rebuilds `data` with the type `target`, which must match it position by position apart +/// from nested field names and metadata. Unlike a cast, children are never matched by name. +fn relabel_nested_fields( + data: ArrayData, + target: &DataType, +) -> datafusion_common::Result { + let child_types = match target { + DataType::List(field) + | DataType::LargeList(field) + | DataType::FixedSizeList(field, _) + | DataType::Map(field, _) => vec![field.data_type()], + DataType::Struct(fields) => fields.iter().map(|f| f.data_type()).collect(), + _ if data.child_data().is_empty() => vec![], + _ => { + return Err(datafusion_common::DataFusionError::NotImplemented(format!( + "Relabeling nested fields of {target}" + ))); + } + }; + let children = data + .child_data() + .iter() + .zip(child_types) + .map(|(child, child_type)| relabel_nested_fields(child.clone(), child_type)) + .collect::>>()?; + Ok(data + .into_builder() + .data_type(target.clone()) + .child_data(children) + .build()?) +} + impl StateAccum { pub fn new( inner: Box, @@ -636,40 +688,7 @@ impl Accumulator for StateAccum { .iter() .map(|s| s.to_array()) .collect::, _>>()?; - let array_type = array - .iter() - .map(|a| a.data_type().clone()) - .collect::>(); - let expected_type: Vec<_> = self - .state_fields - .iter() - .map(|f| f.data_type().clone()) - .collect(); - if array_type != expected_type { - debug!( - "State mismatch, expected: {}, got: {} for expected fields: {:?} and given array types: {:?}", - self.state_fields.len(), - array.len(), - self.state_fields, - array_type, - ); - let guess_schema = array - .iter() - .enumerate() - .map(|(index, array)| { - Field::new( - format!("col_{index}[mismatch_state]").as_str(), - array.data_type().clone(), - true, - ) - }) - .collect::(); - let arr = StructArray::try_new(guess_schema, array, None)?; - - return Ok(ScalarValue::Struct(Arc::new(arr))); - } - - let struct_array = StructArray::try_new(self.state_fields.clone(), array, None)?; + let struct_array = state_struct_array(&self.state_fields, array)?; Ok(ScalarValue::Struct(Arc::new(struct_array))) } diff --git a/src/common/function/src/aggrs/aggr_wrapper/fix_order.rs b/src/common/function/src/aggrs/aggr_wrapper/fix_order.rs index 480b0b5a490..715b051407d 100644 --- a/src/common/function/src/aggrs/aggr_wrapper/fix_order.rs +++ b/src/common/function/src/aggrs/aggr_wrapper/fix_order.rs @@ -153,13 +153,17 @@ fn rewrite_expr( if is_fix { // then always fix the ordering field&distinct flag and more let order_by = aggregate_function.params.order_by.clone(); + // DataFusion always makes ordering fields nullable when building the physical + // aggregate (`ordering_fields` in `datafusion-functions-aggregate-common`), and some + // accumulators (e.g. `array_agg`) nest these fields in their state type. Keep the + // declared state type in line with what the accumulator produces. let ordering_fields: Vec<_> = order_by .iter() .map(|sort_expr| { sort_expr .expr .to_field(&aggregate_input.schema()) - .map(|(_, f)| f) + .map(|(_, f)| Arc::new(f.as_ref().clone().with_nullable(true))) }) .collect::>>()?; let distinct = aggregate_function.params.distinct; diff --git a/src/common/function/src/aggrs/aggr_wrapper/tests.rs b/src/common/function/src/aggrs/aggr_wrapper/tests.rs index 4274a03cd89..b68a735f182 100644 --- a/src/common/function/src/aggrs/aggr_wrapper/tests.rs +++ b/src/common/function/src/aggrs/aggr_wrapper/tests.rs @@ -27,8 +27,11 @@ use datafusion::catalog::{Session, TableProvider}; use datafusion::common::tree_node::TreeNodeRecursion; use datafusion::datasource::DefaultTableSource; use datafusion::execution::{RecordBatchStream, SendableRecordBatchStream, TaskContext}; +use datafusion::functions_aggregate::approx_percentile_cont::approx_percentile_cont_udaf; +use datafusion::functions_aggregate::array_agg::array_agg_udaf; use datafusion::functions_aggregate::average::avg_udaf; use datafusion::functions_aggregate::count::count_udaf; +use datafusion::functions_aggregate::nth_value::nth_value_udaf; use datafusion::functions_aggregate::sum::sum_udaf; use datafusion::optimizer::AnalyzerRule; use datafusion::optimizer::analyzer::type_coercion::TypeCoercion; @@ -1289,6 +1292,40 @@ async fn test_udaf_correct_eval_result() { order_by: vec![], null_treatment: None, }, + // The ordering state of `array_agg` nests the ORDER BY fields in `List(Struct(..))`. + TestCase { + func: array_agg_udaf(), + input_schema: Arc::new(arrow_schema::Schema::new(vec![ + Field::new("number", DataType::Float64, true), + Field::new( + "ts", + DataType::Timestamp(arrow_schema::TimeUnit::Millisecond, None), + false, + ), + ])), + args: vec![Expr::Column(Column::new_unqualified("number"))], + input: vec![ + Arc::new(Float64Array::from(vec![Some(3.), Some(1.), Some(2.)])), + Arc::new(TimestampMillisecondArray::from(vec![3000, 1000, 2000])), + ], + expected_output: Some(ScalarValue::List(ScalarValue::new_list_nullable( + &[ + ScalarValue::Float64(Some(1.)), + ScalarValue::Float64(Some(2.)), + ScalarValue::Float64(Some(3.)), + ], + &DataType::Float64, + ))), + expected_fn: None, + distinct: false, + filter: None, + order_by: vec![SortExpr::new( + Expr::Column(Column::new_unqualified("ts")), + true, + true, + )], + null_treatment: None, + }, // TODO(discord9): udd_merge/hll_merge/geo_path/quantile_aggr tests ]; let test_table_ref = TableReference::bare("TestTable"); @@ -1396,3 +1433,89 @@ async fn execute_phy_plan( } Ok(batches) } + +#[test] +fn test_state_struct_array_rejects_mismatched_state() { + let fields = Fields::from(vec![Field::new("sum", DataType::Int64, true)]); + let arrays: Vec = vec![Arc::new(Float64Array::from(vec![1.0]))]; + let err = state_struct_array(&fields, arrays).unwrap_err(); + assert!( + err.to_string() + .contains("State field `sum` expects type Int64, but the accumulator produced Float64"), + "{err}" + ); +} + +#[test] +fn test_state_struct_array_keeps_child_order() { + let int_field = |name: &str| Field::new(name, DataType::Int64, true); + // Same child types, names crossed: children must stay in place, not be matched by name. + let produced = StructArray::from(vec![ + ( + Arc::new(int_field("a")), + Arc::new(Int64Array::from(vec![10])) as ArrayRef, + ), + ( + Arc::new(int_field("b")), + Arc::new(Int64Array::from(vec![20])) as ArrayRef, + ), + ]); + let declared = Fields::from(vec![Field::new( + "state", + DataType::Struct(Fields::from(vec![int_field("b"), int_field("a")])), + true, + )]); + + let state = state_struct_array(&declared, vec![Arc::new(produced)]).unwrap(); + let state = state.column(0).as_struct(); + assert_eq!( + state + .column_by_name("b") + .unwrap() + .as_primitive::() + .value(0), + 10 + ); + assert_eq!( + state + .column_by_name("a") + .unwrap() + .as_primitive::() + .value(0), + 20 + ); +} + +#[test] +fn test_hard_ordered_aggr_not_steppable() { + let order_by = vec![SortExpr::new( + Expr::Column(Column::new_unqualified("ts")), + true, + true, + )]; + let aggr = |func: Arc, args: Vec| { + Expr::AggregateFunction(AggregateFunction::new_udf( + func, + args, + false, + None, + order_by.clone(), + None, + )) + }; + let number = Expr::Column(Column::new_unqualified("number")); + + assert!(!is_all_aggr_exprs_steppable(&[aggr( + nth_value_udaf(), + vec![number.clone(), lit(2i64)], + )])); + assert!(is_all_aggr_exprs_steppable(&[aggr( + array_agg_udaf(), + vec![number.clone()] + )])); + // WITHIN GROUP (ORDER BY number) + assert!(is_all_aggr_exprs_steppable(&[aggr( + approx_percentile_cont_udaf(), + vec![number, lit(0.5f64)] + )])); +} diff --git a/tests/cases/standalone/common/aggregate/array_agg.result b/tests/cases/standalone/common/aggregate/array_agg.result index 5a73c289899..236a4582f14 100644 --- a/tests/cases/standalone/common/aggregate/array_agg.result +++ b/tests/cases/standalone/common/aggregate/array_agg.result @@ -144,3 +144,94 @@ DROP TABLE doubles; Affected Rows: 0 +-- On partitioned tables the aggregate is split into partial state and merge +CREATE TABLE array_agg_partitioned ( + ts TIMESTAMP TIME INDEX, + k INT, + lat DOUBLE, + PRIMARY KEY(k) +) +PARTITION ON COLUMNS (k) (k < 10, k >= 10 AND k < 20, k >= 20); + +Affected Rows: 0 + +INSERT INTO array_agg_partitioned VALUES + (1000, 1, 1), + (2000, 11, 2), + (3000, 21, 3), + (4000, 2, 4), + (5000, 12, 5), + (6000, 22, 6), + (7000, 3, 7); + +Affected Rows: 7 + +SELECT array_agg(lat ORDER BY ts) FROM array_agg_partitioned; + ++-----------------------------------------------------------------------------------------+ +| array_agg(array_agg_partitioned.lat) ORDER BY [array_agg_partitioned.ts ASC NULLS LAST] | ++-----------------------------------------------------------------------------------------+ +| [1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0] | ++-----------------------------------------------------------------------------------------+ + +SELECT lat > 3 AS g, array_agg(lat ORDER BY ts DESC) FROM array_agg_partitioned GROUP BY g ORDER BY g; + ++-------+-------------------------------------------------------------------------------------------+ +| g | array_agg(array_agg_partitioned.lat) ORDER BY [array_agg_partitioned.ts DESC NULLS FIRST] | ++-------+-------------------------------------------------------------------------------------------+ +| false | [3.0, 2.0, 1.0] | +| true | [7.0, 6.0, 5.0, 4.0] | ++-------+-------------------------------------------------------------------------------------------+ + +SELECT array_agg(k ORDER BY k % 10, ts DESC) FROM array_agg_partitioned; + ++---------------------------------------------------------------------------------------------------------------------------------------------+ +| array_agg(array_agg_partitioned.k) ORDER BY [array_agg_partitioned.k % Int64(10) ASC NULLS LAST, array_agg_partitioned.ts DESC NULLS FIRST] | ++---------------------------------------------------------------------------------------------------------------------------------------------+ +| [21, 11, 1, 22, 12, 2, 3] | ++---------------------------------------------------------------------------------------------------------------------------------------------+ + +-- nth_value needs sorted input, so it isn't split into partial state and merge +SELECT nth_value(lat, 2 ORDER BY ts), nth_value(lat, 3 ORDER BY ts DESC) FROM array_agg_partitioned; + ++--------------------------------------------------------------------------------------------------+----------------------------------------------------------------------------------------------------+ +| nth_value(array_agg_partitioned.lat,Int64(2)) ORDER BY [array_agg_partitioned.ts ASC NULLS LAST] | nth_value(array_agg_partitioned.lat,Int64(3)) ORDER BY [array_agg_partitioned.ts DESC NULLS FIRST] | ++--------------------------------------------------------------------------------------------------+----------------------------------------------------------------------------------------------------+ +| 2.0 | 5.0 | ++--------------------------------------------------------------------------------------------------+----------------------------------------------------------------------------------------------------+ + +-- WITHIN GROUP aggregates don't need sorted input and are still split +-- SQLNESS REPLACE (peers.*) REDACTED +-- SQLNESS REPLACE (RoundRobinBatch.*) REDACTED +-- SQLNESS REPLACE (-+) - +-- SQLNESS REPLACE (\s\s+) _ +EXPLAIN SELECT sum(lat), approx_percentile_cont(0.75) WITHIN GROUP (ORDER BY lat DESC) FROM array_agg_partitioned; + ++-+-+ +| plan_type_| plan_| ++-+-+ +| logical_plan_| Aggregate: groupBy=[[]], aggr=[[__sum_merge(__sum_state(array_agg_partitioned.lat)) AS sum(array_agg_partitioned.lat), __approx_percentile_cont_merge(__approx_percentile_cont_state(array_agg_partitioned.lat,Float64(0.75)) ORDER BY [array_agg_partitioned.lat DESC NULLS FIRST]) AS approx_percentile_cont(Float64(0.75)) WITHIN GROUP [array_agg_partitioned.lat DESC NULLS FIRST]]]_| +|_|_MergeScan [is_placeholder=false, remote_input=[_| +|_| Aggregate: groupBy=[[]], aggr=[[__sum_state(array_agg_partitioned.lat), __approx_percentile_cont_state(array_agg_partitioned.lat, Float64(0.75)) ORDER BY [array_agg_partitioned.lat DESC NULLS FIRST]]]_| +|_|_TableScan: array_agg_partitioned_| +|_| ]]_| +| physical_plan | AggregateExec: mode=Final, gby=[], aggr=[__sum_merge(__sum_state(array_agg_partitioned.lat)) as sum(array_agg_partitioned.lat), __approx_percentile_cont_merge(__approx_percentile_cont_state(array_agg_partitioned.lat,Float64(0.75)) ORDER BY [array_agg_partitioned.lat DESC NULLS FIRST]) as approx_percentile_cont(Float64(0.75)) WITHIN GROUP [array_agg_partitioned.lat DESC NULLS FIRST]]_| +|_|_CoalescePartitionsExec_| +|_|_AggregateExec: mode=Partial, gby=[], aggr=[__sum_merge(__sum_state(array_agg_partitioned.lat)) as sum(array_agg_partitioned.lat), __approx_percentile_cont_merge(__approx_percentile_cont_state(array_agg_partitioned.lat,Float64(0.75)) ORDER BY [array_agg_partitioned.lat DESC NULLS FIRST]) as approx_percentile_cont(Float64(0.75)) WITHIN GROUP [array_agg_partitioned.lat DESC NULLS FIRST]] | +|_|_RepartitionExec: partitioning=REDACTED +|_|_MergeScanExec: REDACTED +|_|_| ++-+-+ + +SELECT sum(lat), approx_percentile_cont(0.75) WITHIN GROUP (ORDER BY lat DESC) FROM array_agg_partitioned; + ++--------------------------------+-------------------------------------------------------------------------------------------------+ +| sum(array_agg_partitioned.lat) | approx_percentile_cont(Float64(0.75)) WITHIN GROUP [array_agg_partitioned.lat DESC NULLS FIRST] | ++--------------------------------+-------------------------------------------------------------------------------------------------+ +| 28.0 | 2.25 | ++--------------------------------+-------------------------------------------------------------------------------------------------+ + +DROP TABLE array_agg_partitioned; + +Affected Rows: 0 + diff --git a/tests/cases/standalone/common/aggregate/array_agg.sql b/tests/cases/standalone/common/aggregate/array_agg.sql index dedabf1a18f..d0d7caa574f 100644 --- a/tests/cases/standalone/common/aggregate/array_agg.sql +++ b/tests/cases/standalone/common/aggregate/array_agg.sql @@ -54,3 +54,41 @@ DROP TABLE integers; DROP TABLE strings; DROP TABLE doubles; + +-- On partitioned tables the aggregate is split into partial state and merge +CREATE TABLE array_agg_partitioned ( + ts TIMESTAMP TIME INDEX, + k INT, + lat DOUBLE, + PRIMARY KEY(k) +) +PARTITION ON COLUMNS (k) (k < 10, k >= 10 AND k < 20, k >= 20); + +INSERT INTO array_agg_partitioned VALUES + (1000, 1, 1), + (2000, 11, 2), + (3000, 21, 3), + (4000, 2, 4), + (5000, 12, 5), + (6000, 22, 6), + (7000, 3, 7); + +SELECT array_agg(lat ORDER BY ts) FROM array_agg_partitioned; + +SELECT lat > 3 AS g, array_agg(lat ORDER BY ts DESC) FROM array_agg_partitioned GROUP BY g ORDER BY g; + +SELECT array_agg(k ORDER BY k % 10, ts DESC) FROM array_agg_partitioned; + +-- nth_value needs sorted input, so it isn't split into partial state and merge +SELECT nth_value(lat, 2 ORDER BY ts), nth_value(lat, 3 ORDER BY ts DESC) FROM array_agg_partitioned; + +-- WITHIN GROUP aggregates don't need sorted input and are still split +-- SQLNESS REPLACE (peers.*) REDACTED +-- SQLNESS REPLACE (RoundRobinBatch.*) REDACTED +-- SQLNESS REPLACE (-+) - +-- SQLNESS REPLACE (\s\s+) _ +EXPLAIN SELECT sum(lat), approx_percentile_cont(0.75) WITHIN GROUP (ORDER BY lat DESC) FROM array_agg_partitioned; + +SELECT sum(lat), approx_percentile_cont(0.75) WITHIN GROUP (ORDER BY lat DESC) FROM array_agg_partitioned; + +DROP TABLE array_agg_partitioned;