diff --git a/datafusion/functions-aggregate/src/median.rs b/datafusion/functions-aggregate/src/median.rs index 7a399f73ec8e2..81a3c076dffbe 100644 --- a/datafusion/functions-aggregate/src/median.rs +++ b/datafusion/functions-aggregate/src/median.rs @@ -39,10 +39,11 @@ use arrow::datatypes::{ ArrowNativeType, ArrowPrimitiveType, Decimal32Type, Decimal64Type, FieldRef, }; +use datafusion_common::hash_utils::RandomState; use datafusion_common::types::{NativeType, logical_float64}; use datafusion_common::{ DataFusionError, Result, ScalarValue, assert_eq_or_internal_err, exec_datafusion_err, - internal_datafusion_err, + internal_datafusion_err, internal_err, }; use datafusion_expr::function::StateFieldsArgs; use datafusion_expr::{ @@ -288,7 +289,12 @@ impl Accumulator for MedianAccumulator { "failed to reserve {additional} values for median accumulator: {e}" ) })?; - self.all_values.extend(values.iter().flatten()); + if values.null_count() > 0 { + self.all_values.extend(values.iter().flatten()); + } else { + // Fast path: no nulls, so the values buffer can be appended wholesale. + self.all_values.extend_from_slice(values.values()); + } Ok(()) } @@ -310,11 +316,19 @@ impl Accumulator for MedianAccumulator { } fn retract_batch(&mut self, values: &[ArrayRef]) -> Result<()> { - let mut to_remove: HashMap, usize> = HashMap::new(); + let mut to_remove: HashMap, usize, RandomState> = + HashMap::default(); let arr = values[0].as_primitive::(); - for value in arr.iter().flatten() { - *to_remove.entry(Hashable(value)).or_default() += 1; + if arr.null_count() > 0 { + for value in arr.iter().flatten() { + *to_remove.entry(Hashable(value)).or_default() += 1; + } + } else { + // Fast path: no nulls, so skip the per-element validity check. + for value in arr.values().iter() { + *to_remove.entry(Hashable(*value)).or_default() += 1; + } } let mut i = 0; @@ -335,6 +349,15 @@ impl Accumulator for MedianAccumulator { i += 1; } } + + // Retracting values that are not tracked means the accumulator state + // has diverged from the window frame; continuing would silently + // produce wrong results, so surface it as an error. + if !to_remove.is_empty() { + return internal_err!( + "median retract_batch: retracted value(s) not present in the window" + ); + } Ok(()) } @@ -627,3 +650,58 @@ fn calculate_median(values: &mut [T::Native]) -> Option MedianAccumulator { + MedianAccumulator { + data_type: DataType::Float64, + all_values: vec![], + } + } + + #[test] + fn retract_batch_errors_on_untracked_value() { + let mut acc = median_accumulator(); + let values: ArrayRef = Arc::new(Float64Array::from(vec![1.0, 2.0])); + acc.update_batch(std::slice::from_ref(&values)).unwrap(); + + let retract: ArrayRef = Arc::new(Float64Array::from(vec![3.0])); + let err = acc + .retract_batch(std::slice::from_ref(&retract)) + .unwrap_err() + .to_string(); + assert!( + err.contains("not present in the window"), + "unexpected error: {err}" + ); + } + + #[test] + fn update_batch_with_and_without_nulls_agree() { + // The null-free fast path must accumulate the same values as the + // general path. + let dense: ArrayRef = Arc::new(Float64Array::from(vec![1.0, 2.0, 3.0])); + let sparse: ArrayRef = Arc::new(Float64Array::from(vec![ + Some(1.0), + None, + Some(2.0), + None, + Some(3.0), + ])); + + let mut dense_acc = median_accumulator(); + dense_acc + .update_batch(std::slice::from_ref(&dense)) + .unwrap(); + let mut sparse_acc = median_accumulator(); + sparse_acc + .update_batch(std::slice::from_ref(&sparse)) + .unwrap(); + + assert_eq!(dense_acc.all_values, sparse_acc.all_values); + } +} diff --git a/datafusion/functions-aggregate/src/percentile_cont.rs b/datafusion/functions-aggregate/src/percentile_cont.rs index 4b5e892fb5cdf..53eda31f6e512 100644 --- a/datafusion/functions-aggregate/src/percentile_cont.rs +++ b/datafusion/functions-aggregate/src/percentile_cont.rs @@ -32,6 +32,7 @@ use arrow::{ use num_traits::AsPrimitive; use arrow::array::ArrowNativeTypeOp; +use datafusion_common::hash_utils::RandomState; use datafusion_common::internal_err; use datafusion_common::types::{NativeType, logical_float64}; use datafusion_functions_aggregate_common::noop_accumulator::NoopAccumulator; @@ -427,7 +428,12 @@ where "failed to reserve {additional} values for percentile_cont accumulator: {e}" ) })?; - self.all_values.extend(values.iter().flatten()); + if values.null_count() > 0 { + self.all_values.extend(values.iter().flatten()); + } else { + // Fast path: no nulls, so the values buffer can be appended wholesale. + self.all_values.extend_from_slice(values.values()); + } Ok(()) } @@ -447,11 +453,19 @@ where } fn retract_batch(&mut self, values: &[ArrayRef]) -> Result<()> { - let mut to_remove: HashMap, usize> = HashMap::new(); + let mut to_remove: HashMap, usize, RandomState> = + HashMap::default(); let arr = values[0].as_primitive::(); - for value in arr.iter().flatten() { - *to_remove.entry(Hashable(value)).or_default() += 1; + if arr.null_count() > 0 { + for value in arr.iter().flatten() { + *to_remove.entry(Hashable(value)).or_default() += 1; + } + } else { + // Fast path: no nulls, so skip the per-element validity check. + for value in arr.values().iter() { + *to_remove.entry(Hashable(*value)).or_default() += 1; + } } let mut i = 0; @@ -472,6 +486,15 @@ where i += 1; } } + + // Retracting values that are not tracked means the accumulator state + // has diverged from the window frame; continuing would silently + // produce wrong results, so surface it as an error. + if !to_remove.is_empty() { + return internal_err!( + "percentile_cont retract_batch: retracted value(s) not present in the window" + ); + } Ok(()) } @@ -842,18 +865,60 @@ where #[cfg(test)] mod tests { - use super::calculate_percentile; + use super::*; + use arrow::array::Float64Array; use half::f16; + #[test] + fn retract_batch_errors_on_untracked_value() { + let mut acc = PercentileContAccumulator::::new(0.5); + let values: ArrayRef = Arc::new(Float64Array::from(vec![1.0, 2.0])); + acc.update_batch(std::slice::from_ref(&values)).unwrap(); + + let retract: ArrayRef = Arc::new(Float64Array::from(vec![3.0])); + let err = acc + .retract_batch(std::slice::from_ref(&retract)) + .unwrap_err() + .to_string(); + assert!( + err.contains("not present in the window"), + "unexpected error: {err}" + ); + } + + #[test] + fn update_batch_with_and_without_nulls_agree() { + // The null-free fast path must accumulate the same values as the + // general path. + let dense: ArrayRef = Arc::new(Float64Array::from(vec![1.0, 2.0, 3.0])); + let sparse: ArrayRef = Arc::new(Float64Array::from(vec![ + Some(1.0), + None, + Some(2.0), + None, + Some(3.0), + ])); + + let mut dense_acc = PercentileContAccumulator::::new(0.5); + dense_acc + .update_batch(std::slice::from_ref(&dense)) + .unwrap(); + let mut sparse_acc = PercentileContAccumulator::::new(0.5); + sparse_acc + .update_batch(std::slice::from_ref(&sparse)) + .unwrap(); + + assert_eq!(dense_acc.all_values, sparse_acc.all_values); + } + #[test] fn f16_interpolation_does_not_overflow_to_nan() { // Regression test for https://github.com/apache/datafusion/issues/18945 // Interpolating between 0 and the max finite f16 value previously overflowed // intermediate f16 computations and produced NaN. let mut values = vec![f16::from_f32(0.0), f16::from_f32(65504.0)]; - let result = - calculate_percentile::(&mut values, 0.5) - .expect("non-empty input"); + let result = calculate_percentile::(&mut values, 0.5) + .expect("non-empty input"); let result_f = result.to_f32(); assert!( !result_f.is_nan(),