From 715f5066019d80637dfefcf8e31d6e8c397d58f1 Mon Sep 17 00:00:00 2001 From: Mason Hall Date: Tue, 22 Sep 2026 15:54:31 -0400 Subject: [PATCH 1/9] handle grouped accumulator state in ungrouped approx distinct --- .../src/approx_distinct.rs | 36 ++++++++++++++++--- 1 file changed, 32 insertions(+), 4 deletions(-) diff --git a/datafusion/functions-aggregate/src/approx_distinct.rs b/datafusion/functions-aggregate/src/approx_distinct.rs index 74f116cf48725..2b447f2b9b317 100644 --- a/datafusion/functions-aggregate/src/approx_distinct.rs +++ b/datafusion/functions-aggregate/src/approx_distinct.rs @@ -126,6 +126,36 @@ impl Accumulator for ApproxDistinctBitmapWrapper { } } +fn merge_serialized(hll: &mut HyperLogLog, bytes: &[u8]) -> Result<()> { + if bytes.is_empty() { + return Ok(()); + } + if bytes.len() == NUM_REGISTERS { + let other: HyperLogLog = bytes.try_into()?; + hll.merge(&other); + } else { + if !bytes.len().is_multiple_of(size_of::()) { + return internal_err!( + "approx_distinct: malformed sparse state: length {} is not a multiple of {}", + bytes.len(), + size_of::() + ); + } + if bytes.len() > SPARSE_LIMIT * size_of::() { + return internal_err!( + "approx_distinct: malformed sparse state: length {} exceeds sparse limit {}", + bytes.len(), + SPARSE_LIMIT * size_of::() + ); + } + for chunk in bytes.chunks_exact(size_of::()) { + let h = u64::from_le_bytes(chunk.try_into().unwrap()); + hll.add_hashed(h); + } + } + Ok(()) +} + #[derive(Debug)] struct HLLAccumulator { hll: HyperLogLog, @@ -172,8 +202,7 @@ impl Accumulator for HLLAccumulator { let v = v.ok_or_else(|| { internal_datafusion_err!("Impossibly got empty binary array from states") })?; - let other = v.try_into()?; - self.hll.merge(&other); + merge_serialized(&mut self.hll, v)?; } Ok(()) } @@ -232,8 +261,7 @@ where let v = v.ok_or_else(|| { internal_datafusion_err!("Impossibly got empty binary array from states") })?; - let other = v.try_into()?; - self.hll.merge(&other); + merge_serialized(&mut self.hll, v)?; } Ok(()) } From c9e145dc415f15160f6fa19877d3e2e32b0f6dd7 Mon Sep 17 00:00:00 2001 From: Mason Hall Date: Tue, 22 Sep 2026 16:07:57 -0400 Subject: [PATCH 2/9] add tests, fmt --- .../src/approx_distinct.rs | 62 ++++++++++++++++++- 1 file changed, 60 insertions(+), 2 deletions(-) diff --git a/datafusion/functions-aggregate/src/approx_distinct.rs b/datafusion/functions-aggregate/src/approx_distinct.rs index 2b447f2b9b317..de0af2dd452a1 100644 --- a/datafusion/functions-aggregate/src/approx_distinct.rs +++ b/datafusion/functions-aggregate/src/approx_distinct.rs @@ -126,7 +126,10 @@ impl Accumulator for ApproxDistinctBitmapWrapper { } } -fn merge_serialized(hll: &mut HyperLogLog, bytes: &[u8]) -> Result<()> { +fn merge_serialized( + hll: &mut HyperLogLog, + bytes: &[u8], +) -> Result<()> { if bytes.is_empty() { return Ok(()); } @@ -1039,7 +1042,7 @@ mod tests { use arrow::array::{ AsArray, Decimal32Array, Decimal64Array, Decimal128Array, Decimal256Array, Int64Array, IntervalDayTimeArray, IntervalMonthDayNanoArray, - IntervalYearMonthArray, StringViewArray, + IntervalYearMonthArray, StringArray, StringViewArray, }; use arrow::datatypes::{IntervalDayTime, IntervalMonthDayNano, i256}; use std::sync::Arc; @@ -1383,6 +1386,61 @@ mod tests { ); assert_eq!(distinct_count(&mut acc_single), 3); } + + /// Grouped state (empty, sparse and dense rows) must merge into the + /// ungrouped accumulator. + #[test] + fn hll_acc_merges_grouped_state() { + let sparse = 10; + let dense = SPARSE_LIMIT * 4; + let values: ArrayRef = Arc::new(StringArray::from_iter_values( + (0..sparse + dense).map(|i| format!("value-{i}")), + )); + let group_indices: Vec = (0..sparse + dense) + .map(|i| if i < sparse { 1 } else { 2 }) + .collect(); + + let mut grouped = HllGroupsAccumulator::new(); + grouped + .update_batch(std::slice::from_ref(&values), &group_indices, None, 3) + .unwrap(); + let state = grouped.state(EmitTo::All).unwrap(); + + let mut direct = HLLAccumulator::new(); + direct.update_batch(std::slice::from_ref(&values)).unwrap(); + + let mut merged = HLLAccumulator::new(); + merged.merge_batch(&state).unwrap(); + + assert_eq!(distinct_count(&mut merged), distinct_count(&mut direct)); + } + + /// Grouped state (empty, sparse and dense rows) must merge into the + /// ungrouped numeric accumulator. + #[test] + fn numeric_acc_merges_grouped_state() { + let sparse = 10; + let dense = SPARSE_LIMIT * 4; + let values: ArrayRef = + Arc::new(Int64Array::from_iter_values(0..(sparse + dense) as i64)); + let group_indices: Vec = (0..sparse + dense) + .map(|i| if i < sparse { 1 } else { 2 }) + .collect(); + + let mut grouped = HllGroupsAccumulator::new(); + grouped + .update_batch(std::slice::from_ref(&values), &group_indices, None, 3) + .unwrap(); + let state = grouped.state(EmitTo::All).unwrap(); + + let mut direct = NumericHLLAccumulator::::new(); + direct.update_batch(std::slice::from_ref(&values)).unwrap(); + + let mut merged = NumericHLLAccumulator::::new(); + merged.merge_batch(&state).unwrap(); + + assert_eq!(merged.evaluate().unwrap(), direct.evaluate().unwrap()); + } } fn h(v: u64) -> u64 { From 21b7c0cea2ffcdee294d66834d75299e48014c07 Mon Sep 17 00:00:00 2001 From: Mason Hall Date: Wed, 23 Sep 2026 15:42:56 -0400 Subject: [PATCH 3/9] fix clippy --- datafusion/functions-aggregate/src/approx_distinct.rs | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/datafusion/functions-aggregate/src/approx_distinct.rs b/datafusion/functions-aggregate/src/approx_distinct.rs index de0af2dd452a1..a47b46682cc71 100644 --- a/datafusion/functions-aggregate/src/approx_distinct.rs +++ b/datafusion/functions-aggregate/src/approx_distinct.rs @@ -151,8 +151,8 @@ fn merge_serialized( SPARSE_LIMIT * size_of::() ); } - for chunk in bytes.chunks_exact(size_of::()) { - let h = u64::from_le_bytes(chunk.try_into().unwrap()); + for chunk in bytes.as_chunks::<{ size_of::() }>().0 { + let h = u64::from_le_bytes(*chunk); hll.add_hashed(h); } } From abf16b73ed38df26704df2df665749d9434eb47b Mon Sep 17 00:00:00 2001 From: Mason Hall Date: Tue, 29 Sep 2026 13:50:29 -0400 Subject: [PATCH 4/9] factor out duplicate code --- .../src/approx_distinct.rs | 115 ++++++++---------- 1 file changed, 53 insertions(+), 62 deletions(-) diff --git a/datafusion/functions-aggregate/src/approx_distinct.rs b/datafusion/functions-aggregate/src/approx_distinct.rs index a47b46682cc71..3ce2e5005705b 100644 --- a/datafusion/functions-aggregate/src/approx_distinct.rs +++ b/datafusion/functions-aggregate/src/approx_distinct.rs @@ -126,34 +126,60 @@ impl Accumulator for ApproxDistinctBitmapWrapper { } } -fn merge_serialized( - hll: &mut HyperLogLog, - bytes: &[u8], -) -> Result<()> { - if bytes.is_empty() { - return Ok(()); - } - if bytes.len() == NUM_REGISTERS { - let other: HyperLogLog = bytes.try_into()?; - hll.merge(&other); - } else { - if !bytes.len().is_multiple_of(size_of::()) { +/// A validated, zero-copy view of a serialized `approx_distinct` partial state, +/// as produced by [`GroupHll::serialize`] or by the per-group [`Accumulator`]s. +enum SerializedHll<'a> { + /// The raw [`NUM_REGISTERS`] registers of a dense sketch. + Dense(&'a [u8; NUM_REGISTERS]), + /// Little-endian hashes of at most [`SPARSE_LIMIT`] distinct values. An + /// empty state decodes as an empty sparse state. + Sparse(&'a [[u8; size_of::()]]), +} + +impl<'a> SerializedHll<'a> { + fn decode(bytes: &'a [u8]) -> Result { + if let Ok(registers) = <&[u8; NUM_REGISTERS]>::try_from(bytes) { + return Ok(Self::Dense(registers)); + } + let (chunks, rest) = bytes.as_chunks::<{ size_of::() }>(); + if !rest.is_empty() { return internal_err!( "approx_distinct: malformed sparse state: length {} is not a multiple of {}", bytes.len(), size_of::() ); } - if bytes.len() > SPARSE_LIMIT * size_of::() { + if chunks.len() > SPARSE_LIMIT { return internal_err!( "approx_distinct: malformed sparse state: length {} exceeds sparse limit {}", bytes.len(), SPARSE_LIMIT * size_of::() ); } - for chunk in bytes.as_chunks::<{ size_of::() }>().0 { - let h = u64::from_le_bytes(*chunk); - hll.add_hashed(h); + Ok(Self::Sparse(chunks)) + } +} + +/// Merge the serialized partial states in `states` into `hll`. +fn merge_states( + hll: &mut HyperLogLog, + states: &[ArrayRef], +) -> Result<()> { + assert_eq!(1, states.len(), "expect only 1 element in the states"); + let binary_array = downcast_value!(states[0], BinaryArray); + for v in binary_array.iter() { + let v = v.ok_or_else(|| { + internal_datafusion_err!("Impossibly got empty binary array from states") + })?; + match SerializedHll::decode(v)? { + SerializedHll::Dense(registers) => { + hll.merge(&HyperLogLog::new_with_registers(*registers)); + } + SerializedHll::Sparse(chunks) => { + for chunk in chunks { + hll.add_hashed(u64::from_le_bytes(*chunk)); + } + } } } Ok(()) @@ -199,15 +225,7 @@ impl Accumulator for HLLAccumulator { } fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> { - assert_eq!(1, states.len(), "expect only 1 element in the states"); - let binary_array = downcast_value!(states[0], BinaryArray); - for v in binary_array.iter() { - let v = v.ok_or_else(|| { - internal_datafusion_err!("Impossibly got empty binary array from states") - })?; - merge_serialized(&mut self.hll, v)?; - } - Ok(()) + merge_states(&mut self.hll, states) } fn state(&mut self) -> Result> { @@ -258,15 +276,7 @@ where } fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> { - assert_eq!(1, states.len(), "expect only 1 element in the states"); - let binary_array = downcast_value!(states[0], BinaryArray); - for v in binary_array.iter() { - let v = v.ok_or_else(|| { - internal_datafusion_err!("Impossibly got empty binary array from states") - })?; - merge_serialized(&mut self.hll, v)?; - } - Ok(()) + merge_states(&mut self.hll, states) } fn state(&mut self) -> Result> { @@ -366,36 +376,17 @@ impl GroupHll { } } - /// Merge a serialized state (produced by [`Self::serialize`] or by the - /// per-group [`Accumulator`]) into this sketch. + /// Merge a serialized state (see [`SerializedHll`]) into this sketch. fn merge_serialized(&mut self, bytes: &[u8]) -> Result { - if bytes.is_empty() { - return Ok(0); - } - if bytes.len() == NUM_REGISTERS { - let other: HyperLogLog = bytes.try_into()?; - Ok(self.merge_dense(&other)) - } else { - if !bytes.len().is_multiple_of(size_of::()) { - return internal_err!( - "approx_distinct: malformed sparse state: length {} is not a multiple of {}", - bytes.len(), - size_of::() - ); - } - if bytes.len() > SPARSE_LIMIT * size_of::() { - return internal_err!( - "approx_distinct: malformed sparse state: length {} exceeds sparse limit {}", - bytes.len(), - SPARSE_LIMIT * size_of::() - ); - } - let mut delta = 0; - for chunk in bytes.as_chunks::<{ size_of::() }>().0 { - delta += self.add_hash(u64::from_le_bytes(*chunk)); + Ok(match SerializedHll::decode(bytes)? { + SerializedHll::Dense(registers) => { + self.merge_dense(&HyperLogLog::new_with_registers(*registers)) } - Ok(delta) - } + SerializedHll::Sparse(chunks) => chunks + .iter() + .map(|chunk| self.add_hash(u64::from_le_bytes(*chunk))) + .sum(), + }) } /// Merge a dense sketch into this one, promoting to dense if necessary. From c3496a35b2a0d509a5b3087b396a01fa935df8a8 Mon Sep 17 00:00:00 2001 From: Mason Hall Date: Thu, 24 Sep 2026 11:21:01 -0400 Subject: [PATCH 5/9] Add a test to enforce compatibility between an aggregate's accumulators --- .../functions-aggregate/tests/state_compat.rs | 828 ++++++++++++++++++ 1 file changed, 828 insertions(+) create mode 100644 datafusion/functions-aggregate/tests/state_compat.rs diff --git a/datafusion/functions-aggregate/tests/state_compat.rs b/datafusion/functions-aggregate/tests/state_compat.rs new file mode 100644 index 0000000000000..4cbc926122d96 --- /dev/null +++ b/datafusion/functions-aggregate/tests/state_compat.rs @@ -0,0 +1,828 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! Checks that the intermediate state produced by an aggregate's +//! [`Accumulator`] and its [`GroupsAccumulator`] are interchangeable. +//! +//! For every function in [`all_default_aggregate_functions`], every argument +//! shape that the function's signature accepts (from a fixed menu of candidate +//! types and literals), with and without `DISTINCT` and, for functions that are +//! not order insensitive, with and without `ORDER BY`, the test builds both +//! accumulator kinds and checks that state produced by one can be merged by the +//! other with the same result as the ungrouped two-phase path: +//! +//! * `Accumulator::state` -> `Accumulator::merge_batch` (the reference) +//! * `GroupsAccumulator::state` -> `Accumulator::merge_batch` +//! * `Accumulator::state` -> `GroupsAccumulator::merge_batch` +//! * `GroupsAccumulator::state` -> `GroupsAccumulator::merge_batch` +//! * `GroupsAccumulator::convert_to_state` -> both merges +//! +//! It also checks that every state matches the types declared by +//! `state_fields`. +//! +//! The input has groups of very different sizes (including an empty group), +//! since state encodings often depend on how much data a group has seen. + +use std::collections::BTreeMap; +use std::collections::BTreeSet; +use std::panic::{AssertUnwindSafe, catch_unwind}; +use std::sync::Arc; + +use arrow::array::{Array, ArrayRef, Int64Array, UInt32Array}; +use arrow::compute::{cast, concat, take}; +use arrow::datatypes::{DataType, Field, FieldRef, Schema, TimeUnit}; +use arrow::record_batch::{RecordBatch, RecordBatchOptions}; +use datafusion_common::{DataFusionError, Result, ScalarValue}; +use datafusion_expr::type_coercion::functions::fields_with_udf; +use datafusion_expr::{AggregateUDF, EmitTo}; +use datafusion_functions_aggregate::all_default_aggregate_functions; +use datafusion_physical_expr::PhysicalSortExpr; +use datafusion_physical_expr::aggregate::{AggregateExprBuilder, AggregateFunctionExpr}; +use datafusion_physical_expr::expressions::{Column, Literal}; +use datafusion_physical_expr_common::physical_expr::PhysicalExpr; + +/// Why a registered function is not exercised. +#[derive(Clone, Copy, Debug)] +enum Reason { + /// No native `GroupsAccumulator`: `groups_accumulator_supported` returns + /// false and `create_groups_accumulator` returns the trait's default error. + NoGroupsAccumulator, + /// Replaced during planning, so `accumulator` always fails. + NoAccumulator, +} + +/// Functions that are not expected to be exercised by this test, with the +/// reason. +const NOT_EXERCISED: &[(&str, Reason)] = &[ + ("any_value", Reason::NoGroupsAccumulator), + ("approx_median", Reason::NoGroupsAccumulator), + ("approx_percentile_cont", Reason::NoGroupsAccumulator), + ( + "approx_percentile_cont_with_weight", + Reason::NoGroupsAccumulator, + ), + ("covar_pop", Reason::NoGroupsAccumulator), + ("covar_samp", Reason::NoGroupsAccumulator), + ("grouping", Reason::NoAccumulator), + ("nth_value", Reason::NoGroupsAccumulator), + ("regr_avgx", Reason::NoGroupsAccumulator), + ("regr_avgy", Reason::NoGroupsAccumulator), + ("regr_count", Reason::NoGroupsAccumulator), + ("regr_intercept", Reason::NoGroupsAccumulator), + ("regr_r2", Reason::NoGroupsAccumulator), + ("regr_slope", Reason::NoGroupsAccumulator), + ("regr_sxx", Reason::NoGroupsAccumulator), + ("regr_sxy", Reason::NoGroupsAccumulator), + ("regr_syy", Reason::NoGroupsAccumulator), +]; + +/// What the candidate cases revealed about one function. +#[derive(Default)] +struct Coverage { + /// Cases that built an `AggregateFunctionExpr`. + built: usize, + /// Cases for which `create_accumulator` succeeded. + with_accumulator: usize, + /// Cases for which `groups_accumulator_supported` returned true or + /// `create_groups_accumulator` returned something other than the trait's + /// default error. + with_groups_accumulator: usize, + /// Cases that were checked for state compatibility. + exercised: usize, +} + +/// Whether `result` is the error returned by the default implementation of +/// `AggregateUDFImpl::create_groups_accumulator`. +fn is_default_groups_error(result: &Result) -> bool { + matches!( + result, + Err(DataFusionError::NotImplemented(msg)) + if msg.starts_with("GroupsAccumulator hasn't been implemented for") + ) +} + +fn candidate_types() -> Vec { + vec![ + DataType::Boolean, + DataType::Int32, + DataType::Int64, + DataType::UInt64, + DataType::Float64, + DataType::Decimal128(10, 2), + DataType::Utf8, + DataType::LargeUtf8, + DataType::Utf8View, + DataType::Binary, + DataType::Date32, + DataType::Timestamp(TimeUnit::Nanosecond, None), + ] +} + +fn candidate_literals() -> Vec { + vec![ + ScalarValue::Float64(Some(0.5)), + ScalarValue::Int64(Some(2)), + ScalarValue::Utf8(Some(",".to_string())), + ] +} + +/// Number of rows in each group. Group 0 is intentionally empty. +const GROUP_SIZES: &[usize] = &[0, 1, 5, 300, 3000]; + +/// Maximum number of arguments tried per function. +const MAX_ARITY: usize = 3; + +#[derive(Clone, Debug)] +enum ArgKind { + Column(DataType), + Literal(ScalarValue), +} + +/// A concrete way to call an aggregate function. +#[derive(Clone, Debug)] +struct Case { + args: Vec, + distinct: bool, + ordered: bool, +} + +impl Case { + fn describe(&self, name: &str) -> String { + let args = self + .args + .iter() + .map(|a| match a { + ArgKind::Column(dt) => format!("col {dt}"), + ArgKind::Literal(v) => format!("lit {v:?}"), + }) + .collect::>() + .join(", "); + let distinct = if self.distinct { "DISTINCT " } else { "" }; + let order_by = if self.ordered { " ORDER BY o" } else { "" }; + format!("{name}({distinct}{args}{order_by})") + } +} + +#[test] +fn accumulator_and_groups_accumulator_states_are_compatible() { + // Panics are caught and reported as failures; keep them from also being + // printed by the default hook. + std::panic::set_hook(Box::new(|_| {})); + + let mut failures: Vec = vec![]; + let mut coverage: BTreeMap = BTreeMap::new(); + + for udaf in all_default_aggregate_functions() { + let name = udaf.name().to_string(); + let cov = coverage.entry(name.clone()).or_default(); + + for case in candidate_cases(&udaf) { + let Some(expr) = build_expr(&udaf, &case) else { + continue; + }; + cov.built += 1; + + let has_accumulator = guard(|| expr.create_accumulator()).is_ok(); + let supported = + guard(|| Ok(expr.groups_accumulator_supported())).unwrap_or(false); + let groups_accumulator = guard(|| expr.create_groups_accumulator()); + if has_accumulator { + cov.with_accumulator += 1; + } + if supported || !is_default_groups_error(&groups_accumulator) { + cov.with_groups_accumulator += 1; + } + + // Only aggregates with a native GroupsAccumulator are interesting: + // otherwise `GroupsAccumulatorAdapter` wraps the `Accumulator` and + // the state formats agree by construction. + if !(supported && has_accumulator && groups_accumulator.is_ok()) { + continue; + } + + cov.exercised += 1; + let desc = case.describe(&name); + let errors = guard(|| Ok(check_case(&expr, &case))).unwrap_or_else(|e| { + let mut errors = Errors::default(); + errors.push(&format!("{e}")); + errors + }); + failures.extend(errors.lines().into_iter().map(|e| format!("{desc}: {e}"))); + } + } + + let not_exercised: BTreeMap<&str, Reason> = NOT_EXERCISED.iter().copied().collect(); + for (name, cov) in &coverage { + if let Some(msg) = check_coverage(cov, not_exercised.get(name.as_str())) { + failures.push(format!("{name}: {msg}")); + } + } + for name in not_exercised.keys() { + if !coverage.contains_key(*name) { + failures.push(format!( + "{name}: listed in NOT_EXERCISED but not registered" + )); + } + } + + let summary = coverage + .iter() + .filter(|(_, cov)| cov.exercised > 0) + .map(|(name, cov)| format!(" {name}: {} case(s)", cov.exercised)) + .collect::>() + .join("\n"); + println!("exercised:\n{summary}"); + + // Restore the default hook so the assertion below is reported. + let _ = std::panic::take_hook(); + assert!( + failures.is_empty(), + "{} state compatibility failure(s):\n{}", + failures.len(), + failures.join("\n") + ); +} + +/// Checks a function's coverage against its `NOT_EXERCISED` entry, returning a +/// failure message if they disagree. +fn check_coverage(cov: &Coverage, reason: Option<&Reason>) -> Option { + let Some(reason) = reason else { + if cov.exercised > 0 { + return None; + } + return Some(if cov.with_groups_accumulator > 0 { + format!( + "has a native GroupsAccumulator ({} case(s)) but no case exercised \ + it; extend the candidate types/literals", + cov.with_groups_accumulator + ) + } else { + "no case exercised a native GroupsAccumulator; extend the candidate \ + types/literals or add it to NOT_EXERCISED" + .to_string() + }); + }; + + if cov.exercised > 0 { + return Some(format!( + "listed in NOT_EXERCISED as {reason:?} but {} case(s) were exercised; \ + remove it from the list", + cov.exercised + )); + } + if cov.built == 0 { + return Some(format!( + "listed in NOT_EXERCISED as {reason:?} but no candidate case builds, \ + so the reason cannot be checked" + )); + } + match reason { + Reason::NoGroupsAccumulator if cov.with_groups_accumulator > 0 => Some(format!( + "listed in NOT_EXERCISED as {reason:?} but {} case(s) report or \ + create a native GroupsAccumulator; remove it from the list and \ + extend the candidate types/literals so it is exercised", + cov.with_groups_accumulator + )), + Reason::NoAccumulator if cov.with_accumulator > 0 => Some(format!( + "listed in NOT_EXERCISED as {reason:?} but {} case(s) create an \ + Accumulator", + cov.with_accumulator + )), + _ => None, + } +} + +/// Enumerates the argument shapes that the function's signature accepts, after +/// the same coercion the planner applies. +fn candidate_cases(udaf: &AggregateUDF) -> Vec { + let mut seen = BTreeSet::new(); + let mut cases = vec![]; + let orderings: &[bool] = if udaf.order_sensitivity().is_insensitive() { + &[false] + } else { + &[false, true] + }; + + for first in candidate_types() { + for arity in 1..=MAX_ARITY { + for rest in arg_shapes(&first, arity - 1) { + let mut args = vec![ArgKind::Column(first.clone())]; + args.extend(rest); + + let Some(args) = coerce(udaf, &args) else { + continue; + }; + if !seen.insert(format!("{args:?}")) { + continue; + } + for distinct in [false, true] { + for &ordered in orderings { + cases.push(Case { + args: args.clone(), + distinct, + ordered, + }); + } + } + } + } + } + cases +} + +/// All combinations of `n` trailing arguments, each either a column of the same +/// type as the first argument or one of the candidate literals. +fn arg_shapes(first: &DataType, n: usize) -> Vec> { + let mut options = vec![ArgKind::Column(first.clone())]; + options.extend(candidate_literals().into_iter().map(ArgKind::Literal)); + + let mut shapes: Vec> = vec![vec![]]; + for _ in 0..n { + shapes = shapes + .into_iter() + .flat_map(|prefix| { + options.iter().map(move |o| { + let mut next = prefix.clone(); + next.push(o.clone()); + next + }) + }) + .collect(); + } + shapes +} + +/// Applies signature coercion, returning the coerced arguments or `None` if the +/// function does not accept them. +fn coerce(udaf: &AggregateUDF, args: &[ArgKind]) -> Option> { + let fields: Vec = args + .iter() + .enumerate() + .map(|(i, a)| { + let dt = match a { + ArgKind::Column(dt) => dt.clone(), + ArgKind::Literal(v) => v.data_type(), + }; + Arc::new(Field::new(format!("c{i}"), dt, true)) + }) + .collect(); + let coerced = fields_with_udf(&fields, udaf).ok()?; + + args.iter() + .zip(coerced) + .map(|(a, f)| match a { + ArgKind::Column(_) => Some(ArgKind::Column(f.data_type().clone())), + ArgKind::Literal(v) => v.cast_to(f.data_type()).ok().map(ArgKind::Literal), + }) + .collect() +} + +/// Builds the physical aggregate expression. Column arguments refer to columns +/// `c0..cN` of a schema that contains only the column arguments. +fn build_expr(udaf: &Arc, case: &Case) -> Option { + let schema = Arc::new(input_schema(case)); + let mut col_idx = 0; + let exprs: Vec> = case + .args + .iter() + .map(|a| match a { + ArgKind::Column(_) => { + let e = Arc::new(Column::new(&format!("c{col_idx}"), col_idx)); + col_idx += 1; + e as Arc + } + ArgKind::Literal(v) => Arc::new(Literal::new(v.clone())) as _, + }) + .collect(); + + let mut builder = AggregateExprBuilder::new(Arc::clone(udaf), exprs) + .schema(Arc::clone(&schema)) + .alias("agg"); + if case.distinct { + builder = builder.distinct(); + } + if case.ordered { + let o = Column::new_with_schema(ORDER_COLUMN, &schema).ok()?; + builder = builder.order_by(vec![PhysicalSortExpr::new_default(Arc::new(o))]); + } + guard(|| builder.build()).ok() +} + +/// Name of the column used by `ORDER BY` cases. +const ORDER_COLUMN: &str = "o"; + +/// Column arguments `c0..cN`, followed by the ordering column if the case is +/// ordered. +fn input_schema(case: &Case) -> Schema { + let mut fields: Vec = case + .args + .iter() + .filter_map(|a| match a { + ArgKind::Column(dt) => Some(dt.clone()), + ArgKind::Literal(_) => None, + }) + .enumerate() + .map(|(i, dt)| Field::new(format!("c{i}"), dt, true)) + .collect(); + if case.ordered { + fields.push(Field::new(ORDER_COLUMN, DataType::Int64, false)); + } + Schema::new(fields) +} + +/// Input data: one array per aggregate argument (literals included, as +/// `AggregateExec` passes them) plus the group index of each row. Rows of +/// different groups are interleaved. +struct Input { + args: Vec, + group_indices: Vec, + num_groups: usize, +} + +impl Input { + fn new(expr: &AggregateFunctionExpr, case: &Case) -> Result { + // Round-robin over the groups that still need rows. + let mut remaining = GROUP_SIZES.to_vec(); + let mut group_indices = vec![]; + while remaining.iter().any(|r| *r > 0) { + for (g, r) in remaining.iter_mut().enumerate() { + if *r > 0 { + group_indices.push(g); + *r -= 1; + } + } + } + + // Every 7th row is null. Column k is scaled differently so that + // multi-column aggregates (covariance, correlation) see different inputs. + let num_rows = group_indices.len(); + let schema = Arc::new(input_schema(case)); + let columns = schema + .fields() + .iter() + .enumerate() + .map(|(k, field)| { + if field.name() == ORDER_COLUMN { + return Ok( + Arc::new(Int64Array::from_iter_values(0..num_rows as i64)) + as ArrayRef, + ); + } + let values: Int64Array = (0..num_rows as i64) + .map(|i| (i % 7 != 3).then_some(i * (k as i64 + 1) % 1009)) + .collect(); + cast_values(&(Arc::new(values) as ArrayRef), field.data_type()) + }) + .collect::>>()?; + let batch = RecordBatch::try_new_with_options( + schema, + columns, + &RecordBatchOptions::new().with_row_count(Some(num_rows)), + )?; + + // Evaluate the arguments the same way `AggregateExec` does: the + // argument expressions followed by the ordering expressions. + let args = expr + .expressions() + .iter() + .chain(expr.order_bys().iter().map(|s| &s.expr)) + .map(|e| e.evaluate(&batch)?.into_array(num_rows)) + .collect::>>()?; + + Ok(Self { + args, + group_indices, + num_groups: GROUP_SIZES.len(), + }) + } + + /// The indices of the rows of group `g`, in input order. + fn rows_of(&self, g: usize) -> UInt32Array { + self.group_indices + .iter() + .enumerate() + .filter(|(_, gi)| **gi == g) + .map(|(i, _)| i as u32) + .collect() + } +} + +/// Casts generated integers to the target type, going through `Utf8` for the +/// binary types, which have no direct cast from integers. +fn cast_values(values: &ArrayRef, dt: &DataType) -> Result { + Ok(match dt { + DataType::Binary | DataType::LargeBinary | DataType::BinaryView => { + cast(&cast(values, &DataType::Utf8)?, dt)? + } + _ => cast(values, dt)?, + }) +} + +fn take_all(arrays: &[ArrayRef], indices: &UInt32Array) -> Result> { + Ok(arrays + .iter() + .map(|a| take(a.as_ref(), indices, None)) + .collect::, _>>()?) +} + +/// Runs every state exchange for one case. +fn check_case(expr: &AggregateFunctionExpr, case: &Case) -> Errors { + let mut errors = Errors::default(); + let input = match Input::new(expr, case) { + Ok(input) => input, + Err(e) => { + errors.push(&format!("generating input: {e}")); + return errors; + } + }; + let state_types: Vec = match expr.state_fields() { + Ok(fields) => fields.iter().map(|f| f.data_type().clone()).collect(), + Err(e) => { + errors.push(&format!("state_fields: {e}")); + return errors; + } + }; + let all_groups: Vec = (0..input.num_groups).collect(); + + // Reference: per-group Accumulator -> state -> Accumulator::merge_batch. + let mut acc_states: Vec> = vec![]; + let mut expected: Vec = vec![]; + for g in 0..input.num_groups { + let reference = guard(|| { + let mut acc = expr.create_accumulator()?; + acc.update_batch(&take_all(&input.args, &input.rows_of(g))?)?; + let state = acc + .state()? + .iter() + .map(|s| s.to_array()) + .collect::>>()?; + let value = merge_into_accumulator(expr, &state)?; + Ok((state, value)) + }); + match reference { + Ok((state, value)) => { + errors.check_types("Accumulator::state", &state, &state_types); + acc_states.push(state); + expected.push(value); + } + Err(e) => { + errors.push_group(&format!("reference Accumulator path: {e}"), g); + return errors; + } + } + } + + // Accumulator::state -> GroupsAccumulator::merge_batch, all groups at once. + errors.compare_all( + "Accumulator::state -> GroupsAccumulator::merge_batch", + guard(|| { + let stacked = stack_states(&acc_states)?; + merge_into_groups(expr, &stacked, &all_groups, input.num_groups) + }), + &expected, + ); + + // GroupsAccumulator::state, emitted for all groups. + let groups_state = guard(|| { + let mut gacc = expr.create_groups_accumulator()?; + gacc.update_batch(&input.args, &input.group_indices, None, input.num_groups)?; + gacc.state(EmitTo::All) + }); + match groups_state { + Ok(groups_state) => { + errors.check_types("GroupsAccumulator::state", &groups_state, &state_types); + + for (g, want) in expected.iter().enumerate() { + errors.compare( + "GroupsAccumulator::state -> Accumulator::merge_batch", + g, + guard(|| { + let state: Vec = + groups_state.iter().map(|a| a.slice(g, 1)).collect(); + merge_into_accumulator(expr, &state) + }), + want, + ); + } + + errors.compare_all( + "GroupsAccumulator::state -> GroupsAccumulator::merge_batch", + guard(|| { + merge_into_groups(expr, &groups_state, &all_groups, input.num_groups) + }), + &expected, + ); + } + Err(e) => errors.push(&format!("GroupsAccumulator::state: {e}")), + } + + // convert_to_state produces one state row per input row. + let converted = guard(|| { + let gacc = expr.create_groups_accumulator()?; + gacc.convert_to_state(&input.args, None) + }); + match converted { + Ok(converted) => { + errors.check_types( + "GroupsAccumulator::convert_to_state", + &converted, + &state_types, + ); + + for (g, want) in expected.iter().enumerate() { + errors.compare( + "convert_to_state -> Accumulator::merge_batch", + g, + guard(|| { + let state = take_all(&converted, &input.rows_of(g))?; + merge_into_accumulator(expr, &state) + }), + want, + ); + } + + errors.compare_all( + "convert_to_state -> GroupsAccumulator::merge_batch", + guard(|| { + merge_into_groups( + expr, + &converted, + &input.group_indices, + input.num_groups, + ) + }), + &expected, + ); + } + Err(e) => errors.push(&format!("GroupsAccumulator::convert_to_state: {e}")), + } + + errors +} + +fn merge_into_accumulator( + expr: &AggregateFunctionExpr, + state: &[ArrayRef], +) -> Result { + let mut acc = expr.create_accumulator()?; + acc.merge_batch(state)?; + acc.evaluate() +} + +fn merge_into_groups( + expr: &AggregateFunctionExpr, + state: &[ArrayRef], + group_indices: &[usize], + num_groups: usize, +) -> Result { + let mut gacc = expr.create_groups_accumulator()?; + gacc.merge_batch(state, group_indices, num_groups)?; + gacc.evaluate(EmitTo::All) +} + +/// Concatenates one-row states into one array per state column. +fn stack_states(states: &[Vec]) -> Result> { + let num_cols = states.first().map(|s| s.len()).unwrap_or(0); + (0..num_cols) + .map(|c| { + let arrays: Vec<&dyn Array> = states.iter().map(|s| s[c].as_ref()).collect(); + Ok(concat(&arrays)?) + }) + .collect() +} + +/// Runs `f`, turning a panic into an error. +fn guard(f: impl FnOnce() -> Result) -> Result { + catch_unwind(AssertUnwindSafe(f)).unwrap_or_else(|panic| { + let msg = panic + .downcast_ref::() + .cloned() + .or_else(|| panic.downcast_ref::<&str>().map(|s| s.to_string())) + .unwrap_or_else(|| "".to_string()); + Err(DataFusionError::Execution(format!("panicked: {msg}"))) + }) +} + +/// Failures for one case. The same message for several groups is reported +/// once, listing the groups. +#[derive(Default)] +struct Errors(Vec<(String, Vec)>); + +impl Errors { + fn push(&mut self, msg: &str) { + self.add(msg, None); + } + + fn push_group(&mut self, msg: &str, group: usize) { + self.add(msg, Some(group)); + } + + fn add(&mut self, msg: &str, group: Option) { + // Only keep the first line, dropping e.g. the "please file a bug" text + // of internal errors. + let msg = msg.lines().next().unwrap_or_default().to_string(); + match self.0.iter_mut().find(|(m, _)| *m == msg) { + Some((_, groups)) => groups.extend(group), + None => self.0.push((msg, group.into_iter().collect())), + } + } + + fn lines(&self) -> Vec { + self.0 + .iter() + .map(|(msg, groups)| { + if groups.is_empty() { + return msg.clone(); + } + let sizes = groups + .iter() + .map(|g| GROUP_SIZES[*g].to_string()) + .collect::>() + .join(", "); + format!("{msg} [groups with {sizes} rows]") + }) + .collect() + } + + fn check_types(&mut self, what: &str, state: &[ArrayRef], expected: &[DataType]) { + let actual: Vec = state.iter().map(|a| a.data_type().clone()).collect(); + if actual != expected { + self.push(&format!( + "{what} produced types {actual:?} but state_fields declares {expected:?}" + )); + } + } + + fn compare_all( + &mut self, + what: &str, + got: Result, + expected: &[ScalarValue], + ) { + let got = match got { + Ok(arr) => arr, + Err(e) => return self.push(&format!("{what}: {e}")), + }; + if got.len() != expected.len() { + return self.push(&format!( + "{what}: expected {} groups, got {}", + expected.len(), + got.len() + )); + } + for (g, want) in expected.iter().enumerate() { + self.compare(what, g, ScalarValue::try_from_array(got.as_ref(), g), want); + } + } + + fn compare( + &mut self, + what: &str, + group: usize, + got: Result, + want: &ScalarValue, + ) { + match got { + Err(e) => self.push_group(&format!("{what}: {e}"), group), + Ok(got) if !scalars_match(&got, want) => { + self.push_group(&format!("{what}: got {got:?}, expected {want:?}"), group) + } + Ok(_) => {} + } + } +} + +/// Exact equality, except that floats are compared with a relative tolerance +/// because merging in a different order can change rounding. +fn scalars_match(a: &ScalarValue, b: &ScalarValue) -> bool { + match (a, b) { + (ScalarValue::Float64(Some(x)), ScalarValue::Float64(Some(y))) => { + floats_match(*x, *y) + } + (ScalarValue::Float32(Some(x)), ScalarValue::Float32(Some(y))) => { + floats_match(*x as f64, *y as f64) + } + _ => a == b, + } +} + +fn floats_match(x: f64, y: f64) -> bool { + if x.is_nan() || y.is_nan() { + return x.is_nan() && y.is_nan(); + } + (x - y).abs() <= 1e-9 * x.abs().max(y.abs()).max(1.0) +} From 560040a2226d2e3588f94a787980fe8871b5c8ae Mon Sep 17 00:00:00 2001 From: Mason Hall Date: Mon, 5 Oct 2026 13:18:55 -0400 Subject: [PATCH 6/9] remove tests --- .../src/approx_distinct.rs | 57 +------------------ 1 file changed, 1 insertion(+), 56 deletions(-) diff --git a/datafusion/functions-aggregate/src/approx_distinct.rs b/datafusion/functions-aggregate/src/approx_distinct.rs index b179b406ccf3a..7da68c47bcd2a 100644 --- a/datafusion/functions-aggregate/src/approx_distinct.rs +++ b/datafusion/functions-aggregate/src/approx_distinct.rs @@ -1045,7 +1045,7 @@ mod tests { use arrow::array::{ AsArray, Decimal32Array, Decimal64Array, Decimal128Array, Decimal256Array, Int64Array, IntervalDayTimeArray, IntervalMonthDayNanoArray, - IntervalYearMonthArray, StringArray, StringViewArray, + IntervalYearMonthArray, StringViewArray, }; use arrow::datatypes::{IntervalDayTime, IntervalMonthDayNano, i256}; use std::sync::Arc; @@ -1389,61 +1389,6 @@ mod tests { ); assert_eq!(distinct_count(&mut acc_single), 3); } - - /// Grouped state (empty, sparse and dense rows) must merge into the - /// ungrouped accumulator. - #[test] - fn hll_acc_merges_grouped_state() { - let sparse = 10; - let dense = SPARSE_LIMIT * 4; - let values: ArrayRef = Arc::new(StringArray::from_iter_values( - (0..sparse + dense).map(|i| format!("value-{i}")), - )); - let group_indices: Vec = (0..sparse + dense) - .map(|i| if i < sparse { 1 } else { 2 }) - .collect(); - - let mut grouped = HllGroupsAccumulator::new(); - grouped - .update_batch(std::slice::from_ref(&values), &group_indices, None, 3) - .unwrap(); - let state = grouped.state(EmitTo::All).unwrap(); - - let mut direct = HLLAccumulator::new(); - direct.update_batch(std::slice::from_ref(&values)).unwrap(); - - let mut merged = HLLAccumulator::new(); - merged.merge_batch(&state).unwrap(); - - assert_eq!(distinct_count(&mut merged), distinct_count(&mut direct)); - } - - /// Grouped state (empty, sparse and dense rows) must merge into the - /// ungrouped numeric accumulator. - #[test] - fn numeric_acc_merges_grouped_state() { - let sparse = 10; - let dense = SPARSE_LIMIT * 4; - let values: ArrayRef = - Arc::new(Int64Array::from_iter_values(0..(sparse + dense) as i64)); - let group_indices: Vec = (0..sparse + dense) - .map(|i| if i < sparse { 1 } else { 2 }) - .collect(); - - let mut grouped = HllGroupsAccumulator::new(); - grouped - .update_batch(std::slice::from_ref(&values), &group_indices, None, 3) - .unwrap(); - let state = grouped.state(EmitTo::All).unwrap(); - - let mut direct = NumericHLLAccumulator::::new(); - direct.update_batch(std::slice::from_ref(&values)).unwrap(); - - let mut merged = NumericHLLAccumulator::::new(); - merged.merge_batch(&state).unwrap(); - - assert_eq!(merged.evaluate().unwrap(), direct.evaluate().unwrap()); - } } fn h(v: u64) -> u64 { From 52043c8aa139702b63b7bb51e2b8b2ae9ea2a7e0 Mon Sep 17 00:00:00 2001 From: Mason Hall Date: Mon, 5 Oct 2026 13:40:38 -0400 Subject: [PATCH 7/9] Fix the `avg` accumulator --- datafusion/functions-aggregate/src/average.rs | 40 ++++++++++++++----- 1 file changed, 29 insertions(+), 11 deletions(-) diff --git a/datafusion/functions-aggregate/src/average.rs b/datafusion/functions-aggregate/src/average.rs index 85b63a9a17970..94e930ea2031d 100644 --- a/datafusion/functions-aggregate/src/average.rs +++ b/datafusion/functions-aggregate/src/average.rs @@ -653,10 +653,14 @@ impl Accumulator for AvgAccumulator { } fn state(&mut self) -> Result> { - Ok(vec![ - ScalarValue::from(self.count), - ScalarValue::Float64(self.sum), - ]) + // With no non-NULL values both fields are NULL, matching + // `AvgGroupsAccumulator::state` + let (count, sum) = if self.count == 0 { + (None, None) + } else { + (Some(self.count), self.sum) + }; + Ok(vec![ScalarValue::UInt64(count), ScalarValue::Float64(sum)]) } fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> { @@ -833,9 +837,16 @@ where } fn state(&mut self) -> Result> { + // With no non-NULL values both fields are NULL, matching + // `AvgGroupsAccumulator::state` + let (count, sum) = if self.count == 0 { + (None, None) + } else { + (Some(self.count), self.sum) + }; Ok(vec![ - ScalarValue::from(self.count), - ScalarValue::new_primitive::(self.sum, &self.sum_data_type)?, + ScalarValue::UInt64(count), + ScalarValue::new_primitive::(sum, &self.sum_data_type)?, ]) } @@ -916,14 +927,21 @@ impl Accumulator for DurationAvgAccumulator { } fn state(&mut self) -> Result> { + // With no non-NULL values both fields are NULL, matching + // `AvgGroupsAccumulator::state` + let (count, sum) = if self.count == 0 { + (None, None) + } else { + (Some(self.count), self.sum) + }; let duration_value = match self.time_unit { - TimeUnit::Second => ScalarValue::DurationSecond(self.sum), - TimeUnit::Millisecond => ScalarValue::DurationMillisecond(self.sum), - TimeUnit::Microsecond => ScalarValue::DurationMicrosecond(self.sum), - TimeUnit::Nanosecond => ScalarValue::DurationNanosecond(self.sum), + TimeUnit::Second => ScalarValue::DurationSecond(sum), + TimeUnit::Millisecond => ScalarValue::DurationMillisecond(sum), + TimeUnit::Microsecond => ScalarValue::DurationMicrosecond(sum), + TimeUnit::Nanosecond => ScalarValue::DurationNanosecond(sum), }; - Ok(vec![ScalarValue::from(self.count), duration_value]) + Ok(vec![ScalarValue::UInt64(count), duration_value]) } fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> { From 59fe8c5c5c6bc14668d5b373cb3c6af454fc02fd Mon Sep 17 00:00:00 2001 From: Mason Hall Date: Mon, 5 Oct 2026 15:24:34 -0400 Subject: [PATCH 8/9] document the requirement and expose the test publicly --- datafusion/expr-common/src/accumulator.rs | 24 +- .../expr-common/src/groups_accumulator.rs | 44 ++ datafusion/expr/src/udaf.rs | 6 + datafusion/functions-aggregate/Cargo.toml | 1 + datafusion/functions-aggregate/src/lib.rs | 2 + datafusion/functions-aggregate/src/testing.rs | 31 + .../{tests => src/testing}/state_compat.rs | 531 ++++++++++++------ docs/source/contributor-guide/howtos.md | 7 +- .../functions/adding-udfs.md | 22 + .../library-user-guide/upgrading/56.0.0.md | 47 ++ 10 files changed, 529 insertions(+), 186 deletions(-) create mode 100644 datafusion/functions-aggregate/src/testing.rs rename datafusion/functions-aggregate/{tests => src/testing}/state_compat.rs (63%) diff --git a/datafusion/expr-common/src/accumulator.rs b/datafusion/expr-common/src/accumulator.rs index 95cee870e2f0b..3d290799c2298 100644 --- a/datafusion/expr-common/src/accumulator.rs +++ b/datafusion/expr-common/src/accumulator.rs @@ -97,8 +97,10 @@ impl Drop for AggregateMetricTimer<'_> { /// `Accumulator`s are stateful objects that implement a single group. They /// aggregate values from multiple rows together into a final output aggregate. /// -/// [`GroupsAccumulator]` is an additional more performant (but also complex) API -/// that manages state for multiple groups at once. +/// [`GroupsAccumulator`] is an additional more performant (but also complex) API +/// that manages state for multiple groups at once. If an aggregate implements +/// both, their intermediate states must be interchangeable (see +/// [`GroupsAccumulator`] for details). /// /// An accumulator knows how to: /// * update its state from inputs via [`update_batch`] @@ -118,6 +120,7 @@ impl Drop for AggregateMetricTimer<'_> { /// [`state`]: Self::state /// [`evaluate`]: Self::evaluate /// [`merge_batch`]: Self::merge_batch +/// [`GroupsAccumulator`]: crate::groups_accumulator::GroupsAccumulator /// [window function]: https://en.wikipedia.org/wiki/Window_function_(SQL) pub trait Accumulator: Send + Sync + Debug + std::any::Any { /// Supplies optional metrics owned by this aggregate expression. @@ -270,6 +273,14 @@ pub trait Accumulator: Send + Sync + Debug + std::any::Any { /// values if the number of intermediate values is not known at /// planning time (e.g. for `MEDIAN`) /// + /// If the aggregate also implements a [`GroupsAccumulator`], the state + /// returned here must be accepted by [`GroupsAccumulator::merge_batch`], + /// and [`Self::merge_batch`] must accept state produced by the + /// [`GroupsAccumulator`]. See [`GroupsAccumulator`] for details. + /// + /// [`GroupsAccumulator`]: crate::groups_accumulator::GroupsAccumulator + /// [`GroupsAccumulator::merge_batch`]: crate::groups_accumulator::GroupsAccumulator::merge_batch + /// /// # Multi-phase repartitioned Grouping /// /// Many multi-phase grouping plans contain a Repartition operation @@ -383,6 +394,15 @@ pub trait Accumulator: Send + Sync + Debug + std::any::Any { /// The `states` array passed was formed by concatenating the /// results of calling [`Self::state`] on zero or more other /// `Accumulator` instances. + /// + /// If the aggregate also implements a [`GroupsAccumulator`], `states` may + /// instead contain state produced by [`GroupsAccumulator::state`] or + /// [`GroupsAccumulator::convert_to_state`]. See [`GroupsAccumulator`] for + /// details. + /// + /// [`GroupsAccumulator`]: crate::groups_accumulator::GroupsAccumulator + /// [`GroupsAccumulator::state`]: crate::groups_accumulator::GroupsAccumulator::state + /// [`GroupsAccumulator::convert_to_state`]: crate::groups_accumulator::GroupsAccumulator::convert_to_state fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()>; /// Retracts (removed) an update (caused by the given inputs) to diff --git a/datafusion/expr-common/src/groups_accumulator.rs b/datafusion/expr-common/src/groups_accumulator.rs index 1c004cb70b931..0b3b516af0667 100644 --- a/datafusion/expr-common/src/groups_accumulator.rs +++ b/datafusion/expr-common/src/groups_accumulator.rs @@ -188,7 +188,33 @@ impl<'a> GroupSelection<'a> { /// expected that each `GroupsAccumulator` will use something like `Vec<..>` /// to store the group states. /// +/// # State Compatibility with `Accumulator` +/// +/// The intermediate state of a `GroupsAccumulator` and of the [`Accumulator`] +/// for the same aggregate (created by the same `AggregateUDFImpl` with the same +/// arguments) must be interchangeable: +/// +/// * Each row of the state returned by [`Self::state`] or +/// [`Self::convert_to_state`] must be accepted by +/// [`Accumulator::merge_batch`]. +/// * [`Self::merge_batch`] must accept state returned by +/// [`Accumulator::state`]. +/// +/// Merging state from the other kind of accumulator must produce the same +/// result as merging the equivalent state from the same kind. +/// +/// Aggregates that do not implement a `GroupsAccumulator` meet this +/// requirement automatically, as they are run with a +/// [`GroupsAccumulatorAdapter`] that wraps their [`Accumulator`]. +/// +/// Use [`check_state_compatibility`] to check that an aggregate meets this +/// requirement. +/// /// [`Accumulator`]: crate::accumulator::Accumulator +/// [`Accumulator::state`]: crate::accumulator::Accumulator::state +/// [`Accumulator::merge_batch`]: crate::accumulator::Accumulator::merge_batch +/// [`GroupsAccumulatorAdapter`]: https://docs.rs/datafusion/latest/datafusion/physical_expr/struct.GroupsAccumulatorAdapter.html +/// [`check_state_compatibility`]: https://docs.rs/datafusion-functions-aggregate/latest/datafusion_functions_aggregate/testing/fn.check_state_compatibility.html /// [Aggregating Millions of Groups Fast blog]: https://arrow.apache.org/blog/2023/08/05/datafusion_fast_grouping/ pub trait GroupsAccumulator: Send + std::any::Any { /// Supplies optional metrics owned by this aggregate expression. @@ -285,7 +311,13 @@ pub trait GroupsAccumulator: Send + std::any::Any { /// See [`Self::evaluate`] for details on the required output /// order and `emit_to`. /// + /// Each row of the returned state must also be accepted by + /// [`Accumulator::merge_batch`]. See the [State Compatibility] section + /// for details. + /// /// [`Accumulator::state`]: crate::accumulator::Accumulator::state + /// [`Accumulator::merge_batch`]: crate::accumulator::Accumulator::merge_batch + /// [State Compatibility]: GroupsAccumulator#state-compatibility-with-accumulator fn state(&mut self, emit_to: EmitTo) -> Result>; /// Returns intermediate aggregate state without changing the logical state @@ -325,6 +357,12 @@ pub trait GroupsAccumulator: Send + std::any::Any { /// there is no `opt_filter` — aggregate filters are applied during the /// partial (update) phase, so by the time intermediate states are merged /// no per-row filtering is needed. + /// + /// `values` may also contain state produced by [`Accumulator::state`]. See + /// the [State Compatibility] section for details. + /// + /// [`Accumulator::state`]: crate::accumulator::Accumulator::state + /// [State Compatibility]: GroupsAccumulator#state-compatibility-with-accumulator fn merge_batch( &mut self, values: &[ArrayRef], @@ -366,7 +404,13 @@ pub trait GroupsAccumulator: Send + std::any::Any { /// state directly to the next aggregation phase with minimal processing /// using this method. /// + /// As with [`Self::state`], each row of the returned state must also be + /// accepted by [`Accumulator::merge_batch`]. See the + /// [State Compatibility] section for details. + /// /// [`Accumulator::state`]: crate::accumulator::Accumulator::state + /// [`Accumulator::merge_batch`]: crate::accumulator::Accumulator::merge_batch + /// [State Compatibility]: GroupsAccumulator#state-compatibility-with-accumulator fn convert_to_state( &self, values: &[ArrayRef], diff --git a/datafusion/expr/src/udaf.rs b/datafusion/expr/src/udaf.rs index 7358568f0afe1..9952855479ba3 100644 --- a/datafusion/expr/src/udaf.rs +++ b/datafusion/expr/src/udaf.rs @@ -662,6 +662,12 @@ pub trait AggregateUDFImpl: Debug + DynEq + DynHash + Send + Sync + Any { /// /// For maximum performance, a [`GroupsAccumulator`] should be /// implemented in addition to [`Accumulator`]. + /// + /// The intermediate state of the returned [`GroupsAccumulator`] must be + /// interchangeable with the state of the [`Accumulator`] returned by + /// [`Self::accumulator`] for the same arguments: each must be able to merge + /// state produced by the other, and both must match + /// [`Self::state_fields`]. See [`GroupsAccumulator`] for details. fn create_groups_accumulator( &self, _args: AccumulatorArgs, diff --git a/datafusion/functions-aggregate/Cargo.toml b/datafusion/functions-aggregate/Cargo.toml index ead59720d216e..b53cb4ea9253d 100644 --- a/datafusion/functions-aggregate/Cargo.toml +++ b/datafusion/functions-aggregate/Cargo.toml @@ -112,3 +112,4 @@ harness = false [features] force_hash_collisions = ["datafusion-common/force_hash_collisions"] +testing = [] diff --git a/datafusion/functions-aggregate/src/lib.rs b/datafusion/functions-aggregate/src/lib.rs index e3f2714abbf25..ae67c41456b31 100644 --- a/datafusion/functions-aggregate/src/lib.rs +++ b/datafusion/functions-aggregate/src/lib.rs @@ -91,6 +91,8 @@ pub mod sum; pub mod variance; pub mod planner; +#[cfg(any(test, feature = "testing"))] +pub mod testing; mod utils; use crate::approx_percentile_cont::approx_percentile_cont_udaf; diff --git a/datafusion/functions-aggregate/src/testing.rs b/datafusion/functions-aggregate/src/testing.rs new file mode 100644 index 0000000000000..7b01629a318bf --- /dev/null +++ b/datafusion/functions-aggregate/src/testing.rs @@ -0,0 +1,31 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! Utilities for testing aggregate function implementations, including user +//! defined aggregate functions. +//! +//! This module requires the `testing` feature, which is intended to be enabled +//! only in `[dev-dependencies]`: +//! +//! ```toml +//! [dev-dependencies] +//! datafusion-functions-aggregate = { version = "...", features = ["testing"] } +//! ``` + +mod state_compat; + +pub use state_compat::check_state_compatibility; diff --git a/datafusion/functions-aggregate/tests/state_compat.rs b/datafusion/functions-aggregate/src/testing/state_compat.rs similarity index 63% rename from datafusion/functions-aggregate/tests/state_compat.rs rename to datafusion/functions-aggregate/src/testing/state_compat.rs index 4cbc926122d96..1f8cf0476fc9d 100644 --- a/datafusion/functions-aggregate/tests/state_compat.rs +++ b/datafusion/functions-aggregate/src/testing/state_compat.rs @@ -18,26 +18,9 @@ //! Checks that the intermediate state produced by an aggregate's //! [`Accumulator`] and its [`GroupsAccumulator`] are interchangeable. //! -//! For every function in [`all_default_aggregate_functions`], every argument -//! shape that the function's signature accepts (from a fixed menu of candidate -//! types and literals), with and without `DISTINCT` and, for functions that are -//! not order insensitive, with and without `ORDER BY`, the test builds both -//! accumulator kinds and checks that state produced by one can be merged by the -//! other with the same result as the ungrouped two-phase path: -//! -//! * `Accumulator::state` -> `Accumulator::merge_batch` (the reference) -//! * `GroupsAccumulator::state` -> `Accumulator::merge_batch` -//! * `Accumulator::state` -> `GroupsAccumulator::merge_batch` -//! * `GroupsAccumulator::state` -> `GroupsAccumulator::merge_batch` -//! * `GroupsAccumulator::convert_to_state` -> both merges -//! -//! It also checks that every state matches the types declared by -//! `state_fields`. -//! -//! The input has groups of very different sizes (including an empty group), -//! since state encodings often depend on how much data a group has seen. +//! [`Accumulator`]: datafusion_expr::Accumulator +//! [`GroupsAccumulator`]: datafusion_expr::GroupsAccumulator -use std::collections::BTreeMap; use std::collections::BTreeSet; use std::panic::{AssertUnwindSafe, catch_unwind}; use std::sync::Arc; @@ -46,50 +29,94 @@ use arrow::array::{Array, ArrayRef, Int64Array, UInt32Array}; use arrow::compute::{cast, concat, take}; use arrow::datatypes::{DataType, Field, FieldRef, Schema, TimeUnit}; use arrow::record_batch::{RecordBatch, RecordBatchOptions}; -use datafusion_common::{DataFusionError, Result, ScalarValue}; +use datafusion_common::{DataFusionError, Result, ScalarValue, exec_err}; use datafusion_expr::type_coercion::functions::fields_with_udf; use datafusion_expr::{AggregateUDF, EmitTo}; -use datafusion_functions_aggregate::all_default_aggregate_functions; use datafusion_physical_expr::PhysicalSortExpr; use datafusion_physical_expr::aggregate::{AggregateExprBuilder, AggregateFunctionExpr}; use datafusion_physical_expr::expressions::{Column, Literal}; use datafusion_physical_expr_common::physical_expr::PhysicalExpr; -/// Why a registered function is not exercised. -#[derive(Clone, Copy, Debug)] -enum Reason { - /// No native `GroupsAccumulator`: `groups_accumulator_supported` returns - /// false and `create_groups_accumulator` returns the trait's default error. - NoGroupsAccumulator, - /// Replaced during planning, so `accumulator` always fails. - NoAccumulator, +/// Checks that `udaf`'s [`Accumulator`] and [`GroupsAccumulator`] can each +/// merge the intermediate state produced by the other. +/// +/// See the [State Compatibility] section of [`GroupsAccumulator`] for the +/// requirement this checks. +/// +/// # What is checked +/// +/// The function is called with every argument shape its signature accepts +/// (after the same coercion the planner applies), built from a fixed set of +/// candidate argument types and literals, with up to three arguments. Each +/// shape is tried with and without `DISTINCT` and, unless the function is +/// order insensitive, with and without `ORDER BY`. +/// +/// Cases for which the function provides a native [`GroupsAccumulator`] +/// (`groups_accumulator_supported` returns true) are run over input split into +/// groups of different sizes. For each case, the result of the ungrouped +/// two-phase path (`Accumulator::state` -> `Accumulator::merge_batch`) is the +/// reference, and the following must produce the same result: +/// +/// * `GroupsAccumulator::state` -> `Accumulator::merge_batch` +/// * `Accumulator::state` -> `GroupsAccumulator::merge_batch` +/// * `GroupsAccumulator::state` -> `GroupsAccumulator::merge_batch` +/// * `GroupsAccumulator::convert_to_state` -> both merges +/// +/// Results must be equal, except that floating point results are compared +/// with a small relative tolerance. Every state must also match the types +/// declared by `state_fields`. Panics are caught and reported as failures. +/// +/// # Returns +/// +/// `Ok(())` if every check passes. This includes the case where the function +/// has no native [`GroupsAccumulator`], since its [`Accumulator`] is then +/// wrapped in a `GroupsAccumulatorAdapter` and the states agree by +/// construction. +/// +/// Otherwise returns an error listing every failure. It also returns an error +/// if the signature accepts none of the candidate argument types, or if no +/// case could exercise a native [`GroupsAccumulator`] that the function +/// reports, since then nothing could be checked. +/// +/// # Example +/// +/// ``` +/// use datafusion_functions_aggregate::average::avg_udaf; +/// use datafusion_functions_aggregate::testing::check_state_compatibility; +/// +/// check_state_compatibility(&avg_udaf()).unwrap(); +/// ``` +/// +/// [`Accumulator`]: datafusion_expr::Accumulator +/// [`GroupsAccumulator`]: datafusion_expr::GroupsAccumulator +/// [State Compatibility]: datafusion_expr::GroupsAccumulator#state-compatibility-with-accumulator +pub fn check_state_compatibility(udaf: &Arc) -> Result<()> { + let name = udaf.name(); + let (coverage, failures) = check_udaf(udaf); + + if coverage.built == 0 { + return exec_err!( + "{name}: the signature accepts none of the candidate argument types, \ + so state compatibility could not be checked" + ); + } + if coverage.exercised == 0 && coverage.with_groups_accumulator > 0 { + return exec_err!( + "{name}: has a native GroupsAccumulator ({} case(s)) but no case could \ + exercise it", + coverage.with_groups_accumulator + ); + } + if !failures.is_empty() { + return exec_err!( + "{} state compatibility failure(s):\n{}", + failures.len(), + failures.join("\n") + ); + } + Ok(()) } -/// Functions that are not expected to be exercised by this test, with the -/// reason. -const NOT_EXERCISED: &[(&str, Reason)] = &[ - ("any_value", Reason::NoGroupsAccumulator), - ("approx_median", Reason::NoGroupsAccumulator), - ("approx_percentile_cont", Reason::NoGroupsAccumulator), - ( - "approx_percentile_cont_with_weight", - Reason::NoGroupsAccumulator, - ), - ("covar_pop", Reason::NoGroupsAccumulator), - ("covar_samp", Reason::NoGroupsAccumulator), - ("grouping", Reason::NoAccumulator), - ("nth_value", Reason::NoGroupsAccumulator), - ("regr_avgx", Reason::NoGroupsAccumulator), - ("regr_avgy", Reason::NoGroupsAccumulator), - ("regr_count", Reason::NoGroupsAccumulator), - ("regr_intercept", Reason::NoGroupsAccumulator), - ("regr_r2", Reason::NoGroupsAccumulator), - ("regr_slope", Reason::NoGroupsAccumulator), - ("regr_sxx", Reason::NoGroupsAccumulator), - ("regr_sxy", Reason::NoGroupsAccumulator), - ("regr_syy", Reason::NoGroupsAccumulator), -]; - /// What the candidate cases revealed about one function. #[derive(Default)] struct Coverage { @@ -105,6 +132,50 @@ struct Coverage { exercised: usize, } +/// Runs every candidate case for `udaf`, returning what was covered and a +/// description of each failure. +fn check_udaf(udaf: &Arc) -> (Coverage, Vec) { + let name = udaf.name(); + let mut cov = Coverage::default(); + let mut failures = vec![]; + + for case in candidate_cases(udaf) { + let Some(expr) = build_expr(udaf, &case) else { + continue; + }; + cov.built += 1; + + let has_accumulator = guard(|| expr.create_accumulator()).is_ok(); + let supported = + guard(|| Ok(expr.groups_accumulator_supported())).unwrap_or(false); + let groups_accumulator = guard(|| expr.create_groups_accumulator()); + if has_accumulator { + cov.with_accumulator += 1; + } + if supported || !is_default_groups_error(&groups_accumulator) { + cov.with_groups_accumulator += 1; + } + + // Only aggregates with a native GroupsAccumulator are interesting: + // otherwise `GroupsAccumulatorAdapter` wraps the `Accumulator` and the + // state formats agree by construction. + if !(supported && has_accumulator && groups_accumulator.is_ok()) { + continue; + } + + cov.exercised += 1; + let desc = case.describe(name); + let errors = guard(|| Ok(check_case(&expr, &case))).unwrap_or_else(|e| { + let mut errors = Errors::default(); + errors.push(&format!("{e}")); + errors + }); + failures.extend(errors.lines().into_iter().map(|e| format!("{desc}: {e}"))); + } + + (cov, failures) +} + /// Whether `result` is the error returned by the default implementation of /// `AggregateUDFImpl::create_groups_accumulator`. fn is_default_groups_error(result: &Result) -> bool { @@ -177,135 +248,6 @@ impl Case { } } -#[test] -fn accumulator_and_groups_accumulator_states_are_compatible() { - // Panics are caught and reported as failures; keep them from also being - // printed by the default hook. - std::panic::set_hook(Box::new(|_| {})); - - let mut failures: Vec = vec![]; - let mut coverage: BTreeMap = BTreeMap::new(); - - for udaf in all_default_aggregate_functions() { - let name = udaf.name().to_string(); - let cov = coverage.entry(name.clone()).or_default(); - - for case in candidate_cases(&udaf) { - let Some(expr) = build_expr(&udaf, &case) else { - continue; - }; - cov.built += 1; - - let has_accumulator = guard(|| expr.create_accumulator()).is_ok(); - let supported = - guard(|| Ok(expr.groups_accumulator_supported())).unwrap_or(false); - let groups_accumulator = guard(|| expr.create_groups_accumulator()); - if has_accumulator { - cov.with_accumulator += 1; - } - if supported || !is_default_groups_error(&groups_accumulator) { - cov.with_groups_accumulator += 1; - } - - // Only aggregates with a native GroupsAccumulator are interesting: - // otherwise `GroupsAccumulatorAdapter` wraps the `Accumulator` and - // the state formats agree by construction. - if !(supported && has_accumulator && groups_accumulator.is_ok()) { - continue; - } - - cov.exercised += 1; - let desc = case.describe(&name); - let errors = guard(|| Ok(check_case(&expr, &case))).unwrap_or_else(|e| { - let mut errors = Errors::default(); - errors.push(&format!("{e}")); - errors - }); - failures.extend(errors.lines().into_iter().map(|e| format!("{desc}: {e}"))); - } - } - - let not_exercised: BTreeMap<&str, Reason> = NOT_EXERCISED.iter().copied().collect(); - for (name, cov) in &coverage { - if let Some(msg) = check_coverage(cov, not_exercised.get(name.as_str())) { - failures.push(format!("{name}: {msg}")); - } - } - for name in not_exercised.keys() { - if !coverage.contains_key(*name) { - failures.push(format!( - "{name}: listed in NOT_EXERCISED but not registered" - )); - } - } - - let summary = coverage - .iter() - .filter(|(_, cov)| cov.exercised > 0) - .map(|(name, cov)| format!(" {name}: {} case(s)", cov.exercised)) - .collect::>() - .join("\n"); - println!("exercised:\n{summary}"); - - // Restore the default hook so the assertion below is reported. - let _ = std::panic::take_hook(); - assert!( - failures.is_empty(), - "{} state compatibility failure(s):\n{}", - failures.len(), - failures.join("\n") - ); -} - -/// Checks a function's coverage against its `NOT_EXERCISED` entry, returning a -/// failure message if they disagree. -fn check_coverage(cov: &Coverage, reason: Option<&Reason>) -> Option { - let Some(reason) = reason else { - if cov.exercised > 0 { - return None; - } - return Some(if cov.with_groups_accumulator > 0 { - format!( - "has a native GroupsAccumulator ({} case(s)) but no case exercised \ - it; extend the candidate types/literals", - cov.with_groups_accumulator - ) - } else { - "no case exercised a native GroupsAccumulator; extend the candidate \ - types/literals or add it to NOT_EXERCISED" - .to_string() - }); - }; - - if cov.exercised > 0 { - return Some(format!( - "listed in NOT_EXERCISED as {reason:?} but {} case(s) were exercised; \ - remove it from the list", - cov.exercised - )); - } - if cov.built == 0 { - return Some(format!( - "listed in NOT_EXERCISED as {reason:?} but no candidate case builds, \ - so the reason cannot be checked" - )); - } - match reason { - Reason::NoGroupsAccumulator if cov.with_groups_accumulator > 0 => Some(format!( - "listed in NOT_EXERCISED as {reason:?} but {} case(s) report or \ - create a native GroupsAccumulator; remove it from the list and \ - extend the candidate types/literals so it is exercised", - cov.with_groups_accumulator - )), - Reason::NoAccumulator if cov.with_accumulator > 0 => Some(format!( - "listed in NOT_EXERCISED as {reason:?} but {} case(s) create an \ - Accumulator", - cov.with_accumulator - )), - _ => None, - } -} - /// Enumerates the argument shapes that the function's signature accepts, after /// the same coercion the planner applies. fn candidate_cases(udaf: &AggregateUDF) -> Vec { @@ -826,3 +768,226 @@ fn floats_match(x: f64, y: f64) -> bool { } (x - y).abs() <= 1e-9 * x.abs().max(y.abs()).max(1.0) } + +#[cfg(test)] +mod tests { + use std::collections::BTreeMap; + + use arrow::datatypes::FieldRef; + use datafusion_expr::function::{AccumulatorArgs, StateFieldsArgs}; + use datafusion_expr::{ + Accumulator, AggregateUDFImpl, GroupsAccumulator, Signature, Volatility, + }; + + use super::*; + use crate::all_default_aggregate_functions; + use crate::count::count_udaf; + use crate::sum::sum_udaf; + + /// Why a built-in function is not exercised. + #[derive(Clone, Copy, Debug)] + enum Reason { + /// No native `GroupsAccumulator`: `groups_accumulator_supported` + /// returns false and `create_groups_accumulator` returns the trait's + /// default error. + NoGroupsAccumulator, + /// Replaced during planning, so `accumulator` always fails. + NoAccumulator, + } + + /// Built-in functions that are not expected to be exercised, with the + /// reason. + const NOT_EXERCISED: &[(&str, Reason)] = &[ + ("any_value", Reason::NoGroupsAccumulator), + ("approx_median", Reason::NoGroupsAccumulator), + ("approx_percentile_cont", Reason::NoGroupsAccumulator), + ( + "approx_percentile_cont_with_weight", + Reason::NoGroupsAccumulator, + ), + ("covar_pop", Reason::NoGroupsAccumulator), + ("covar_samp", Reason::NoGroupsAccumulator), + ("grouping", Reason::NoAccumulator), + ("nth_value", Reason::NoGroupsAccumulator), + ("regr_avgx", Reason::NoGroupsAccumulator), + ("regr_avgy", Reason::NoGroupsAccumulator), + ("regr_count", Reason::NoGroupsAccumulator), + ("regr_intercept", Reason::NoGroupsAccumulator), + ("regr_r2", Reason::NoGroupsAccumulator), + ("regr_slope", Reason::NoGroupsAccumulator), + ("regr_sxx", Reason::NoGroupsAccumulator), + ("regr_sxy", Reason::NoGroupsAccumulator), + ("regr_syy", Reason::NoGroupsAccumulator), + ]; + + /// Checks every built-in aggregate function, and that each one is either + /// exercised or listed in `NOT_EXERCISED` for the right reason. + #[test] + fn builtin_accumulator_and_groups_accumulator_states_are_compatible() { + let mut failures: Vec = vec![]; + let mut coverage: BTreeMap = BTreeMap::new(); + for udaf in all_default_aggregate_functions() { + let (cov, udaf_failures) = check_udaf(&udaf); + failures.extend(udaf_failures); + coverage.insert(udaf.name().to_string(), cov); + } + + let not_exercised: BTreeMap<&str, Reason> = + NOT_EXERCISED.iter().copied().collect(); + for (name, cov) in &coverage { + if let Some(msg) = check_coverage(cov, not_exercised.get(name.as_str())) { + failures.push(format!("{name}: {msg}")); + } + } + for name in not_exercised.keys() { + if !coverage.contains_key(*name) { + failures.push(format!( + "{name}: listed in NOT_EXERCISED but not registered" + )); + } + } + + let summary = coverage + .iter() + .filter(|(_, cov)| cov.exercised > 0) + .map(|(name, cov)| format!(" {name}: {} case(s)", cov.exercised)) + .collect::>() + .join("\n"); + println!("exercised:\n{summary}"); + + assert!( + failures.is_empty(), + "{} state compatibility failure(s):\n{}", + failures.len(), + failures.join("\n") + ); + } + + /// Checks a function's coverage against its `NOT_EXERCISED` entry, + /// returning a failure message if they disagree. + fn check_coverage(cov: &Coverage, reason: Option<&Reason>) -> Option { + let Some(reason) = reason else { + if cov.exercised > 0 { + return None; + } + return Some(if cov.with_groups_accumulator > 0 { + format!( + "has a native GroupsAccumulator ({} case(s)) but no case \ + exercised it; extend the candidate types/literals", + cov.with_groups_accumulator + ) + } else { + "no case exercised a native GroupsAccumulator; extend the \ + candidate types/literals or add it to NOT_EXERCISED" + .to_string() + }); + }; + + if cov.exercised > 0 { + return Some(format!( + "listed in NOT_EXERCISED as {reason:?} but {} case(s) were \ + exercised; remove it from the list", + cov.exercised + )); + } + if cov.built == 0 { + return Some(format!( + "listed in NOT_EXERCISED as {reason:?} but no candidate case \ + builds, so the reason cannot be checked" + )); + } + match reason { + Reason::NoGroupsAccumulator if cov.with_groups_accumulator > 0 => { + Some(format!( + "listed in NOT_EXERCISED as {reason:?} but {} case(s) report \ + or create a native GroupsAccumulator; remove it from the \ + list and extend the candidate types/literals so it is \ + exercised", + cov.with_groups_accumulator + )) + } + Reason::NoAccumulator if cov.with_accumulator > 0 => Some(format!( + "listed in NOT_EXERCISED as {reason:?} but {} case(s) create an \ + Accumulator", + cov.with_accumulator + )), + _ => None, + } + } + + /// An aggregate whose `Accumulator` is `count`'s but whose + /// `GroupsAccumulator` is `sum`'s, so their states are not interchangeable. + #[derive(Debug, PartialEq, Eq, Hash)] + struct Mismatched { + signature: Signature, + } + + impl Mismatched { + fn udaf(arg_type: DataType) -> Arc { + Arc::new(AggregateUDF::from(Self { + signature: Signature::exact(vec![arg_type], Volatility::Immutable), + })) + } + } + + impl AggregateUDFImpl for Mismatched { + fn name(&self) -> &str { + "mismatched" + } + + fn signature(&self) -> &Signature { + &self.signature + } + + fn return_type(&self, _arg_types: &[DataType]) -> Result { + Ok(DataType::Int64) + } + + fn accumulator(&self, args: AccumulatorArgs) -> Result> { + count_udaf().accumulator(args) + } + + fn state_fields(&self, args: StateFieldsArgs) -> Result> { + count_udaf().state_fields(args) + } + + fn groups_accumulator_supported(&self, _args: AccumulatorArgs) -> bool { + true + } + + fn create_groups_accumulator( + &self, + args: AccumulatorArgs, + ) -> Result> { + sum_udaf().create_groups_accumulator(args) + } + } + + #[test] + fn compatible_udaf_passes() { + check_state_compatibility(&crate::average::avg_udaf()).unwrap(); + } + + #[test] + fn incompatible_udaf_fails() { + let err = check_state_compatibility(&Mismatched::udaf(DataType::Int64)) + .unwrap_err() + .to_string(); + assert!( + err.contains("GroupsAccumulator::state -> Accumulator::merge_batch"), + "{err}" + ); + } + + #[test] + fn unsupported_signature_fails() { + let err = + check_state_compatibility(&Mismatched::udaf(DataType::FixedSizeBinary(3))) + .unwrap_err() + .to_string(); + assert!( + err.contains("accepts none of the candidate argument types"), + "{err}" + ); + } +} diff --git a/docs/source/contributor-guide/howtos.md b/docs/source/contributor-guide/howtos.md index 18d9391d24bbe..e76f0e7cc0810 100644 --- a/docs/source/contributor-guide/howtos.md +++ b/docs/source/contributor-guide/howtos.md @@ -45,7 +45,11 @@ Make a PR to update the [rust-toolchain] file in the root of the repository. - Scalar functions are further grouped into modules for families of functions (e.g. string, math, datetime). Functions should be added to the relevant module; if a new module needs to be created then a new [Rust feature] should also be added to allow DataFusion users to conditionally compile the modules as needed -- Aggregate functions can optionally implement a [`GroupsAccumulator`] for better performance +- Aggregate functions can optionally implement a [`GroupsAccumulator`] for better performance. Its intermediate + state must be interchangeable with the state of the function's [`Accumulator`]: each must be able to merge state + produced by the other (see the [`GroupsAccumulator`] docs). This is checked for all built-in aggregate functions by + a test in [`state_compat.rs`]; if a new function cannot be exercised by that test, add it to `NOT_EXERCISED` with the + reason Spark compatible functions are [located in separate crate][df-spark] but otherwise follow the same steps, though all function types (e.g. scalar, nested, aggregate) are grouped together in the single location. @@ -68,6 +72,7 @@ function types (e.g. scalar, nested, aggregate) are grouped together in the sing [`advanced_udaf.rs`]: https://github.com/apache/datafusion/blob/main/datafusion-examples/examples/udf/advanced_udaf.rs [`advanced_udwf.rs`]: https://github.com/apache/datafusion/blob/main/datafusion-examples/examples/udf/advanced_udwf.rs [`simple_udtf.rs`]: https://github.com/apache/datafusion/blob/main/datafusion-examples/examples/udf/simple_udtf.rs +[`state_compat.rs`]: https://github.com/apache/datafusion/blob/main/datafusion/functions-aggregate/src/testing/state_compat.rs [rust feature]: https://doc.rust-lang.org/cargo/reference/features.html **Testing** diff --git a/docs/source/library-user-guide/functions/adding-udfs.md b/docs/source/library-user-guide/functions/adding-udfs.md index 78e90dfa6b8d4..a41b6b20ec4c7 100644 --- a/docs/source/library-user-guide/functions/adding-udfs.md +++ b/docs/source/library-user-guide/functions/adding-udfs.md @@ -1386,7 +1386,29 @@ async fn main() -> Result<()> { ``` +### Implementing a `GroupsAccumulator` + +For better performance with many groups, an aggregate UDF can also implement a [`GroupsAccumulator`] by overriding +`groups_accumulator_supported` and `create_groups_accumulator` (see [`advanced_udaf.rs`]). The intermediate state of +the `GroupsAccumulator` must be interchangeable with the state of the `Accumulator`: each must be able to merge state +produced by the other. + +To check this, enable the `testing` feature of `datafusion-functions-aggregate` in your `[dev-dependencies]` and call +[`check_state_compatibility`] from a test: + +```rust,ignore +use datafusion_functions_aggregate::testing::check_state_compatibility; + +#[test] +fn state_compatibility() { + let udaf = Arc::new(AggregateUDF::from(MyUdaf::new())); + check_state_compatibility(&udaf).unwrap(); +} +``` + [`aggregateudf`]: https://docs.rs/datafusion/latest/datafusion/logical_expr/struct.AggregateUDF.html +[`groupsaccumulator`]: https://docs.rs/datafusion/latest/datafusion/logical_expr/trait.GroupsAccumulator.html +[`check_state_compatibility`]: https://docs.rs/datafusion-functions-aggregate/latest/datafusion_functions_aggregate/testing/fn.check_state_compatibility.html [`create_udaf`]: https://docs.rs/datafusion/latest/datafusion/logical_expr/fn.create_udaf.html [`aggregateudfimpl::distinct_handling`]: https://docs.rs/datafusion/latest/datafusion/logical_expr/trait.AggregateUDFImpl.html#method.distinct_handling [`advanced_udaf.rs`]: https://github.com/apache/datafusion/blob/main/datafusion-examples/examples/udf/advanced_udaf.rs diff --git a/docs/source/library-user-guide/upgrading/56.0.0.md b/docs/source/library-user-guide/upgrading/56.0.0.md index 1ea0f4671009d..336f9f49c0885 100644 --- a/docs/source/library-user-guide/upgrading/56.0.0.md +++ b/docs/source/library-user-guide/upgrading/56.0.0.md @@ -623,6 +623,53 @@ This guarantees unique aggregate state field names and allows Users or integrations that inspect aggregate state field names directly, including custom UDAFs and FFI integrations. +### `Accumulator` and `GroupsAccumulator` state must be interchangeable + +The intermediate state produced by an aggregate's `Accumulator` and its +`GroupsAccumulator` is now documented as interchangeable: each must be able to +merge state produced by the other (including state from +`GroupsAccumulator::convert_to_state`), with the same result as merging state +from the same kind of accumulator. + +The new `testing` feature of `datafusion-functions-aggregate` provides +`testing::check_state_compatibility`, which checks that an aggregate meets +this requirement. All built-in aggregate functions are checked with it. + +To meet this requirement, two built-in aggregates changed: + +- `approx_distinct`: the `Accumulator` now accepts the compact (hash list) + state produced by the `GroupsAccumulator` for groups with few distinct + values. +- `avg`: the `Accumulator` now returns `NULL` for both the count and the sum + when it has seen no non-`NULL` values, matching the `GroupsAccumulator`. + Previously it returned a count of `0`. + +**Who is affected:** + +Users who implement `AggregateUDFImpl` with both an `Accumulator` and a +`GroupsAccumulator` should check that each can merge the other's state, for +example by adding a test that calls `check_state_compatibility`: + +```toml +[dev-dependencies] +datafusion-functions-aggregate = { version = "56", features = ["testing"] } +``` + +```rust,ignore +use std::sync::Arc; +use datafusion_expr::AggregateUDF; +use datafusion_functions_aggregate::testing::check_state_compatibility; + +#[test] +fn state_compatibility() { + let udaf = Arc::new(AggregateUDF::from(MyUdaf::new())); + check_state_compatibility(&udaf).unwrap(); +} +``` + +Users who inspect the intermediate state of `avg` directly may see `NULL` +instead of a count of `0` for empty groups. + ### Physical filter pushdown resolves columns by position `datafusion_physical_plan::filter_pushdown::ChildFilterDescription::from_child` From 16e727fe9bad9004cec996e6204f8887d19ba071 Mon Sep 17 00:00:00 2001 From: Mason Hall Date: Wed, 7 Oct 2026 12:08:26 -0400 Subject: [PATCH 9/9] remove this requirement from public docs --- datafusion/expr-common/src/accumulator.rs | 21 +-------- .../expr-common/src/groups_accumulator.rs | 44 ------------------- datafusion/expr/src/udaf.rs | 6 --- .../src/testing/state_compat.rs | 30 +++++++++++-- docs/source/contributor-guide/howtos.md | 9 ++-- .../functions/adding-udfs.md | 22 ---------- .../library-user-guide/upgrading/56.0.0.md | 44 ++----------------- 7 files changed, 36 insertions(+), 140 deletions(-) diff --git a/datafusion/expr-common/src/accumulator.rs b/datafusion/expr-common/src/accumulator.rs index 3d290799c2298..b23fb6c10340e 100644 --- a/datafusion/expr-common/src/accumulator.rs +++ b/datafusion/expr-common/src/accumulator.rs @@ -98,9 +98,7 @@ impl Drop for AggregateMetricTimer<'_> { /// aggregate values from multiple rows together into a final output aggregate. /// /// [`GroupsAccumulator`] is an additional more performant (but also complex) API -/// that manages state for multiple groups at once. If an aggregate implements -/// both, their intermediate states must be interchangeable (see -/// [`GroupsAccumulator`] for details). +/// that manages state for multiple groups at once. /// /// An accumulator knows how to: /// * update its state from inputs via [`update_batch`] @@ -273,14 +271,6 @@ pub trait Accumulator: Send + Sync + Debug + std::any::Any { /// values if the number of intermediate values is not known at /// planning time (e.g. for `MEDIAN`) /// - /// If the aggregate also implements a [`GroupsAccumulator`], the state - /// returned here must be accepted by [`GroupsAccumulator::merge_batch`], - /// and [`Self::merge_batch`] must accept state produced by the - /// [`GroupsAccumulator`]. See [`GroupsAccumulator`] for details. - /// - /// [`GroupsAccumulator`]: crate::groups_accumulator::GroupsAccumulator - /// [`GroupsAccumulator::merge_batch`]: crate::groups_accumulator::GroupsAccumulator::merge_batch - /// /// # Multi-phase repartitioned Grouping /// /// Many multi-phase grouping plans contain a Repartition operation @@ -394,15 +384,6 @@ pub trait Accumulator: Send + Sync + Debug + std::any::Any { /// The `states` array passed was formed by concatenating the /// results of calling [`Self::state`] on zero or more other /// `Accumulator` instances. - /// - /// If the aggregate also implements a [`GroupsAccumulator`], `states` may - /// instead contain state produced by [`GroupsAccumulator::state`] or - /// [`GroupsAccumulator::convert_to_state`]. See [`GroupsAccumulator`] for - /// details. - /// - /// [`GroupsAccumulator`]: crate::groups_accumulator::GroupsAccumulator - /// [`GroupsAccumulator::state`]: crate::groups_accumulator::GroupsAccumulator::state - /// [`GroupsAccumulator::convert_to_state`]: crate::groups_accumulator::GroupsAccumulator::convert_to_state fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()>; /// Retracts (removed) an update (caused by the given inputs) to diff --git a/datafusion/expr-common/src/groups_accumulator.rs b/datafusion/expr-common/src/groups_accumulator.rs index 0b3b516af0667..1c004cb70b931 100644 --- a/datafusion/expr-common/src/groups_accumulator.rs +++ b/datafusion/expr-common/src/groups_accumulator.rs @@ -188,33 +188,7 @@ impl<'a> GroupSelection<'a> { /// expected that each `GroupsAccumulator` will use something like `Vec<..>` /// to store the group states. /// -/// # State Compatibility with `Accumulator` -/// -/// The intermediate state of a `GroupsAccumulator` and of the [`Accumulator`] -/// for the same aggregate (created by the same `AggregateUDFImpl` with the same -/// arguments) must be interchangeable: -/// -/// * Each row of the state returned by [`Self::state`] or -/// [`Self::convert_to_state`] must be accepted by -/// [`Accumulator::merge_batch`]. -/// * [`Self::merge_batch`] must accept state returned by -/// [`Accumulator::state`]. -/// -/// Merging state from the other kind of accumulator must produce the same -/// result as merging the equivalent state from the same kind. -/// -/// Aggregates that do not implement a `GroupsAccumulator` meet this -/// requirement automatically, as they are run with a -/// [`GroupsAccumulatorAdapter`] that wraps their [`Accumulator`]. -/// -/// Use [`check_state_compatibility`] to check that an aggregate meets this -/// requirement. -/// /// [`Accumulator`]: crate::accumulator::Accumulator -/// [`Accumulator::state`]: crate::accumulator::Accumulator::state -/// [`Accumulator::merge_batch`]: crate::accumulator::Accumulator::merge_batch -/// [`GroupsAccumulatorAdapter`]: https://docs.rs/datafusion/latest/datafusion/physical_expr/struct.GroupsAccumulatorAdapter.html -/// [`check_state_compatibility`]: https://docs.rs/datafusion-functions-aggregate/latest/datafusion_functions_aggregate/testing/fn.check_state_compatibility.html /// [Aggregating Millions of Groups Fast blog]: https://arrow.apache.org/blog/2023/08/05/datafusion_fast_grouping/ pub trait GroupsAccumulator: Send + std::any::Any { /// Supplies optional metrics owned by this aggregate expression. @@ -311,13 +285,7 @@ pub trait GroupsAccumulator: Send + std::any::Any { /// See [`Self::evaluate`] for details on the required output /// order and `emit_to`. /// - /// Each row of the returned state must also be accepted by - /// [`Accumulator::merge_batch`]. See the [State Compatibility] section - /// for details. - /// /// [`Accumulator::state`]: crate::accumulator::Accumulator::state - /// [`Accumulator::merge_batch`]: crate::accumulator::Accumulator::merge_batch - /// [State Compatibility]: GroupsAccumulator#state-compatibility-with-accumulator fn state(&mut self, emit_to: EmitTo) -> Result>; /// Returns intermediate aggregate state without changing the logical state @@ -357,12 +325,6 @@ pub trait GroupsAccumulator: Send + std::any::Any { /// there is no `opt_filter` — aggregate filters are applied during the /// partial (update) phase, so by the time intermediate states are merged /// no per-row filtering is needed. - /// - /// `values` may also contain state produced by [`Accumulator::state`]. See - /// the [State Compatibility] section for details. - /// - /// [`Accumulator::state`]: crate::accumulator::Accumulator::state - /// [State Compatibility]: GroupsAccumulator#state-compatibility-with-accumulator fn merge_batch( &mut self, values: &[ArrayRef], @@ -404,13 +366,7 @@ pub trait GroupsAccumulator: Send + std::any::Any { /// state directly to the next aggregation phase with minimal processing /// using this method. /// - /// As with [`Self::state`], each row of the returned state must also be - /// accepted by [`Accumulator::merge_batch`]. See the - /// [State Compatibility] section for details. - /// /// [`Accumulator::state`]: crate::accumulator::Accumulator::state - /// [`Accumulator::merge_batch`]: crate::accumulator::Accumulator::merge_batch - /// [State Compatibility]: GroupsAccumulator#state-compatibility-with-accumulator fn convert_to_state( &self, values: &[ArrayRef], diff --git a/datafusion/expr/src/udaf.rs b/datafusion/expr/src/udaf.rs index 9952855479ba3..7358568f0afe1 100644 --- a/datafusion/expr/src/udaf.rs +++ b/datafusion/expr/src/udaf.rs @@ -662,12 +662,6 @@ pub trait AggregateUDFImpl: Debug + DynEq + DynHash + Send + Sync + Any { /// /// For maximum performance, a [`GroupsAccumulator`] should be /// implemented in addition to [`Accumulator`]. - /// - /// The intermediate state of the returned [`GroupsAccumulator`] must be - /// interchangeable with the state of the [`Accumulator`] returned by - /// [`Self::accumulator`] for the same arguments: each must be able to merge - /// state produced by the other, and both must match - /// [`Self::state_fields`]. See [`GroupsAccumulator`] for details. fn create_groups_accumulator( &self, _args: AccumulatorArgs, diff --git a/datafusion/functions-aggregate/src/testing/state_compat.rs b/datafusion/functions-aggregate/src/testing/state_compat.rs index 1f8cf0476fc9d..9dbacd561facd 100644 --- a/datafusion/functions-aggregate/src/testing/state_compat.rs +++ b/datafusion/functions-aggregate/src/testing/state_compat.rs @@ -18,6 +18,28 @@ //! Checks that the intermediate state produced by an aggregate's //! [`Accumulator`] and its [`GroupsAccumulator`] are interchangeable. //! +//! # Built-in aggregate invariant +//! +//! DataFusion's own execution never mixes the two kinds of accumulator when +//! merging state, so this is not a requirement of the `AggregateUDFImpl` API. +//! It is, however, an invariant that every built-in aggregate function with a +//! native [`GroupsAccumulator`] maintains, so that systems built on DataFusion +//! can merge state produced by either kind with the other: +//! +//! * Each row of the state returned by `GroupsAccumulator::state` or +//! `GroupsAccumulator::convert_to_state` is accepted by +//! `Accumulator::merge_batch`. +//! * `GroupsAccumulator::merge_batch` accepts state returned by +//! `Accumulator::state`. +//! * Merging state from the other kind produces the same result as merging +//! the equivalent state from the same kind. +//! +//! The `builtin_accumulator_and_groups_accumulator_states_are_compatible` test +//! below enforces this for every function in +//! [`all_default_aggregate_functions`]. A new built-in function that the test +//! cannot exercise must be added to its `NOT_EXERCISED` list with the reason. +//! +//! [`all_default_aggregate_functions`]: crate::all_default_aggregate_functions //! [`Accumulator`]: datafusion_expr::Accumulator //! [`GroupsAccumulator`]: datafusion_expr::GroupsAccumulator @@ -40,8 +62,11 @@ use datafusion_physical_expr_common::physical_expr::PhysicalExpr; /// Checks that `udaf`'s [`Accumulator`] and [`GroupsAccumulator`] can each /// merge the intermediate state produced by the other. /// -/// See the [State Compatibility] section of [`GroupsAccumulator`] for the -/// requirement this checks. +/// DataFusion itself does not require this of user defined aggregates, since +/// it never merges state produced by one kind of accumulator with the other. +/// All built-in aggregate functions do maintain it, and this function can be +/// used to check the same of an aggregate that is used in a system that +/// relies on it. /// /// # What is checked /// @@ -89,7 +114,6 @@ use datafusion_physical_expr_common::physical_expr::PhysicalExpr; /// /// [`Accumulator`]: datafusion_expr::Accumulator /// [`GroupsAccumulator`]: datafusion_expr::GroupsAccumulator -/// [State Compatibility]: datafusion_expr::GroupsAccumulator#state-compatibility-with-accumulator pub fn check_state_compatibility(udaf: &Arc) -> Result<()> { let name = udaf.name(); let (coverage, failures) = check_udaf(udaf); diff --git a/docs/source/contributor-guide/howtos.md b/docs/source/contributor-guide/howtos.md index e76f0e7cc0810..8b026d2ba40a7 100644 --- a/docs/source/contributor-guide/howtos.md +++ b/docs/source/contributor-guide/howtos.md @@ -45,11 +45,10 @@ Make a PR to update the [rust-toolchain] file in the root of the repository. - Scalar functions are further grouped into modules for families of functions (e.g. string, math, datetime). Functions should be added to the relevant module; if a new module needs to be created then a new [Rust feature] should also be added to allow DataFusion users to conditionally compile the modules as needed -- Aggregate functions can optionally implement a [`GroupsAccumulator`] for better performance. Its intermediate - state must be interchangeable with the state of the function's [`Accumulator`]: each must be able to merge state - produced by the other (see the [`GroupsAccumulator`] docs). This is checked for all built-in aggregate functions by - a test in [`state_compat.rs`]; if a new function cannot be exercised by that test, add it to `NOT_EXERCISED` with the - reason +- Aggregate functions can optionally implement a [`GroupsAccumulator`] for better performance. Built-in aggregate + functions keep its intermediate state interchangeable with the state of the function's [`Accumulator`]: each must be + able to merge state produced by the other. This is checked for all built-in aggregate functions by a test in + [`state_compat.rs`]; if a new function cannot be exercised by that test, add it to `NOT_EXERCISED` with the reason Spark compatible functions are [located in separate crate][df-spark] but otherwise follow the same steps, though all function types (e.g. scalar, nested, aggregate) are grouped together in the single location. diff --git a/docs/source/library-user-guide/functions/adding-udfs.md b/docs/source/library-user-guide/functions/adding-udfs.md index a41b6b20ec4c7..78e90dfa6b8d4 100644 --- a/docs/source/library-user-guide/functions/adding-udfs.md +++ b/docs/source/library-user-guide/functions/adding-udfs.md @@ -1386,29 +1386,7 @@ async fn main() -> Result<()> { ``` -### Implementing a `GroupsAccumulator` - -For better performance with many groups, an aggregate UDF can also implement a [`GroupsAccumulator`] by overriding -`groups_accumulator_supported` and `create_groups_accumulator` (see [`advanced_udaf.rs`]). The intermediate state of -the `GroupsAccumulator` must be interchangeable with the state of the `Accumulator`: each must be able to merge state -produced by the other. - -To check this, enable the `testing` feature of `datafusion-functions-aggregate` in your `[dev-dependencies]` and call -[`check_state_compatibility`] from a test: - -```rust,ignore -use datafusion_functions_aggregate::testing::check_state_compatibility; - -#[test] -fn state_compatibility() { - let udaf = Arc::new(AggregateUDF::from(MyUdaf::new())); - check_state_compatibility(&udaf).unwrap(); -} -``` - [`aggregateudf`]: https://docs.rs/datafusion/latest/datafusion/logical_expr/struct.AggregateUDF.html -[`groupsaccumulator`]: https://docs.rs/datafusion/latest/datafusion/logical_expr/trait.GroupsAccumulator.html -[`check_state_compatibility`]: https://docs.rs/datafusion-functions-aggregate/latest/datafusion_functions_aggregate/testing/fn.check_state_compatibility.html [`create_udaf`]: https://docs.rs/datafusion/latest/datafusion/logical_expr/fn.create_udaf.html [`aggregateudfimpl::distinct_handling`]: https://docs.rs/datafusion/latest/datafusion/logical_expr/trait.AggregateUDFImpl.html#method.distinct_handling [`advanced_udaf.rs`]: https://github.com/apache/datafusion/blob/main/datafusion-examples/examples/udf/advanced_udaf.rs diff --git a/docs/source/library-user-guide/upgrading/56.0.0.md b/docs/source/library-user-guide/upgrading/56.0.0.md index 336f9f49c0885..4e75921a2f891 100644 --- a/docs/source/library-user-guide/upgrading/56.0.0.md +++ b/docs/source/library-user-guide/upgrading/56.0.0.md @@ -623,50 +623,14 @@ This guarantees unique aggregate state field names and allows Users or integrations that inspect aggregate state field names directly, including custom UDAFs and FFI integrations. -### `Accumulator` and `GroupsAccumulator` state must be interchangeable +### `avg` intermediate state for groups with no values -The intermediate state produced by an aggregate's `Accumulator` and its -`GroupsAccumulator` is now documented as interchangeable: each must be able to -merge state produced by the other (including state from -`GroupsAccumulator::convert_to_state`), with the same result as merging state -from the same kind of accumulator. - -The new `testing` feature of `datafusion-functions-aggregate` provides -`testing::check_state_compatibility`, which checks that an aggregate meets -this requirement. All built-in aggregate functions are checked with it. - -To meet this requirement, two built-in aggregates changed: - -- `approx_distinct`: the `Accumulator` now accepts the compact (hash list) - state produced by the `GroupsAccumulator` for groups with few distinct - values. -- `avg`: the `Accumulator` now returns `NULL` for both the count and the sum - when it has seen no non-`NULL` values, matching the `GroupsAccumulator`. - Previously it returned a count of `0`. +The `Accumulator` for `avg` now returns `NULL` for both the count and the sum +in its intermediate state when it has seen no non-`NULL` values, matching its +`GroupsAccumulator`. Previously it returned a count of `0`. **Who is affected:** -Users who implement `AggregateUDFImpl` with both an `Accumulator` and a -`GroupsAccumulator` should check that each can merge the other's state, for -example by adding a test that calls `check_state_compatibility`: - -```toml -[dev-dependencies] -datafusion-functions-aggregate = { version = "56", features = ["testing"] } -``` - -```rust,ignore -use std::sync::Arc; -use datafusion_expr::AggregateUDF; -use datafusion_functions_aggregate::testing::check_state_compatibility; - -#[test] -fn state_compatibility() { - let udaf = Arc::new(AggregateUDF::from(MyUdaf::new())); - check_state_compatibility(&udaf).unwrap(); -} -``` - Users who inspect the intermediate state of `avg` directly may see `NULL` instead of a count of `0` for empty groups.