diff --git a/datafusion/functions-aggregate/src/average.rs b/datafusion/functions-aggregate/src/average.rs index 8ce893c17dfd4..f59ac816a930a 100644 --- a/datafusion/functions-aggregate/src/average.rs +++ b/datafusion/functions-aggregate/src/average.rs @@ -1246,12 +1246,9 @@ where fn size(&self) -> usize { // Heap buffers self.counts.capacity() * size_of::() - + self.sums.capacity() * size_of::() - // Vec struct overhead (ptr, len, cap) for each field - + size_of::>() - + size_of::>() - // Null tracking buffers - + self.null_state.size() + + self.sums.capacity() * size_of::() + // Null tracking buffers + + self.null_state.size() } } @@ -1643,6 +1640,56 @@ mod tests { Ok(()) } + #[test] + fn avg_groups_size_excludes_accumulator_storage() -> Result<()> { + let mut accumulator = AvgGroupsAccumulator::::new( + &DataType::Float64, + &DataType::Float64, + |sum, count| Ok(sum / count as f64), + ); + + assert_eq!(accumulator.size(), 0); + + let values = Arc::new(Float64Array::from(vec![Some(2.0), None, Some(4.0)])); + accumulator.update_batch(&[values], &[0, 1, 2], None, 3)?; + + let expected_size = accumulator.counts.capacity() * size_of::() + + accumulator.sums.capacity() * size_of::() + + accumulator.null_state.size(); + assert_eq!(accumulator.size(), expected_size); + assert!(accumulator.size() > 0); + + Ok(()) + } + + #[test] + fn avg_groups_size_uses_sum_native_type() -> Result<()> { + let input_type = DataType::Decimal128(26, 0); + let sum_type = avg_sum_data_type(&input_type); + let return_type = Avg::new().return_type(std::slice::from_ref(&input_type))?; + assert_eq!(sum_type, DataType::Decimal256(76, 0)); + assert_eq!(return_type, DataType::Decimal128(30, 4)); + let mut accumulator = AvgGroupsAccumulator::< + Decimal128Type, + _, + Decimal256Type, + Decimal128Type, + >::new(&sum_type, &return_type, |_, _| Ok(0_i128)); + + let values = Arc::new( + Decimal128Array::from(vec![Some(2), None, Some(4)]) + .with_precision_and_scale(26, 0)?, + ); + accumulator.update_batch(&[values], &[0, 1, 2], None, 3)?; + + let expected_size = accumulator.counts.capacity() * size_of::() + + accumulator.sums.capacity() * size_of::() + + accumulator.null_state.size(); + assert_eq!(accumulator.size(), expected_size); + + Ok(()) + } + #[test] fn average_groups_preserving_reads() -> Result<()> { let mut accumulator = AvgGroupsAccumulator::::new(