fix: align ordered aggregate state type with the accumulator output (#9340)

* fix: align ordered aggregate state type with the accumulator output

Signed-off-by: Dennis Zhuang <killme2008@gmail.com>

* fix: reject mismatched aggregate states and keep hard-ordered aggregates unsplit

Signed-off-by: Dennis Zhuang <killme2008@gmail.com>

* fix: keep WITHIN GROUP aggregates splittable and relabel state fields by position

Signed-off-by: Dennis Zhuang <killme2008@gmail.com>

---------

Signed-off-by: Dennis Zhuang <killme2008@gmail.com>
This commit is contained in:
dennis zhuang
2026-09-24 07:02:14 +00:00
committed by GitHub
parent def0c2ec5c
commit affc0a1b1d
5 changed files with 348 additions and 73 deletions
+91 -72
View File
@@ -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<ArrayRef>) -> datafusion_common::Result<ArrayRef> {
let array_type = arrays
.iter()
.map(|array| array.data_type().clone())
.collect::<Vec<_>>();
let expected_type = self
.state_fields
.iter()
.map(|field| field.data_type().clone())
.collect::<Vec<_>>();
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::<Fields>();
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<ArrayRef>,
) -> datafusion_common::Result<StructArray> {
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::<datafusion_common::Result<Vec<_>>>()?;
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<ArrayData> {
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::<datafusion_common::Result<Vec<_>>>()?;
Ok(data
.into_builder()
.data_type(target.clone())
.child_data(children)
.build()?)
}
impl StateAccum {
pub fn new(
inner: Box<dyn Accumulator>,
@@ -636,40 +688,7 @@ impl Accumulator for StateAccum {
.iter()
.map(|s| s.to_array())
.collect::<Result<Vec<_>, _>>()?;
let array_type = array
.iter()
.map(|a| a.data_type().clone())
.collect::<Vec<_>>();
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::<Fields>();
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)))
}
@@ -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::<datafusion_common::Result<Vec<_>>>()?;
let distinct = aggregate_function.params.distinct;
@@ -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<ArrayRef> = 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::<arrow::datatypes::Int64Type>()
.value(0),
10
);
assert_eq!(
state
.column_by_name("a")
.unwrap()
.as_primitive::<arrow::datatypes::Int64Type>()
.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<AggregateUDF>, args: Vec<Expr>| {
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)]
)]));
}
@@ -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
@@ -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;