1use datafusion::arrow::array::{ArrayRef, Float64Array};
16use datafusion::arrow::compute::sum;
17use datafusion::common::cast::{as_binary_array, as_primitive_array};
18use datafusion::common::not_impl_err;
19use datafusion::error::{DataFusionError, Result as DfResult};
20use datafusion::logical_expr::function::AccumulatorArgs;
21use datafusion::logical_expr::{
22 Accumulator as DfAccumulator, AggregateUDF, AggregateUDFImpl, Signature,
23};
24use datafusion_common::ScalarValue;
25use datatypes::arrow::datatypes::{DataType, Float64Type};
26
27pub const AVG_STATE_NAME: &str = "avg_state";
28pub const AVG_MERGE_NAME: &str = "avg_merge";
29
30const ENCODED_LEN: usize = 20;
31const MAGIC: &[u8; 4] = b"AVG1";
32
33#[derive(Debug, Clone, Copy, PartialEq)]
35pub struct AvgState {
36 count: u64,
37 sum: f64,
38}
39
40impl Default for AvgState {
41 fn default() -> Self {
42 Self { count: 0, sum: 0.0 }
43 }
44}
45
46impl AvgState {
47 pub(crate) fn encode(&self) -> [u8; ENCODED_LEN] {
49 let mut encoded = [0; ENCODED_LEN];
50 encoded[..4].copy_from_slice(MAGIC);
51 encoded[4..12].copy_from_slice(&self.count.to_le_bytes());
52 encoded[12..20].copy_from_slice(&self.sum.to_bits().to_le_bytes());
53 encoded
54 }
55
56 pub fn decode(encoded: &[u8]) -> DfResult<Self> {
58 if encoded.len() != ENCODED_LEN || &encoded[..4] != MAGIC {
59 return Err(invalid_state());
60 }
61 let count = decode_u64(encoded, 4);
62 let sum = f64::from_bits(decode_u64(encoded, 12));
63 if count == 0 && sum.to_bits() != 0 {
64 return Err(invalid_state());
65 }
66 Ok(Self { count, sum })
67 }
68
69 pub(crate) fn count(&self) -> u64 {
71 self.count
72 }
73
74 pub fn average(&self) -> Option<f64> {
76 (self.count() != 0).then(|| self.sum / self.count() as f64)
77 }
78}
79
80fn decode_u64(encoded: &[u8], offset: usize) -> u64 {
81 let mut bytes = [0; 8];
82 bytes.copy_from_slice(&encoded[offset..offset + 8]);
83 u64::from_le_bytes(bytes)
84}
85
86fn invalid_state() -> DataFusionError {
87 DataFusionError::Execution("Invalid AVG1 state".to_string())
88}
89
90fn count_overflow() -> DataFusionError {
91 DataFusionError::Execution("AVG count overflow".to_string())
92}
93
94#[derive(Debug, Clone, Eq, PartialEq, Hash)]
99struct AvgUdaf {
100 name: &'static str,
101 signature: Signature,
102 input: InputKind,
103}
104
105impl AggregateUDFImpl for AvgUdaf {
106 fn name(&self) -> &str {
107 self.name
108 }
109
110 fn signature(&self) -> &Signature {
111 &self.signature
112 }
113
114 fn return_type(&self, _arg_types: &[DataType]) -> DfResult<DataType> {
115 Ok(DataType::Binary)
116 }
117
118 fn accumulator(&self, acc_args: AccumulatorArgs) -> DfResult<Box<dyn DfAccumulator>> {
119 if acc_args.is_distinct {
120 return not_impl_err!("AVG DISTINCT aggregations are not available");
121 }
122 let input = match acc_args.exprs[0].data_type(acc_args.schema)? {
123 DataType::Float64 => InputKind::Float64,
124 DataType::Binary => InputKind::Binary,
125 data_type => return not_impl_err!("AVG functions do not support {data_type:?}"),
126 };
127 Ok(Box::new(AvgAccumulator {
128 state: AvgState::default(),
129 input,
130 }))
131 }
132
133 fn default_value(&self, _data_type: &DataType) -> DfResult<ScalarValue> {
134 Ok(ScalarValue::Binary(Some(
135 AvgState::default().encode().to_vec(),
136 )))
137 }
138}
139
140#[derive(Debug, Clone, Copy, Eq, PartialEq, Hash)]
141enum InputKind {
142 Float64,
143 Binary,
144}
145
146#[derive(Debug)]
148pub(crate) struct AvgAccumulator {
149 state: AvgState,
150 input: InputKind,
151}
152
153impl Default for AvgAccumulator {
154 fn default() -> Self {
155 Self {
156 state: AvgState::default(),
157 input: InputKind::Float64,
158 }
159 }
160}
161
162impl AvgAccumulator {
163 pub fn state_udf_impl() -> AggregateUDF {
164 AggregateUDF::new_from_impl(AvgUdaf {
165 name: AVG_STATE_NAME,
166 signature: Signature::exact(
167 vec![DataType::Float64],
168 datafusion::logical_expr::Volatility::Immutable,
169 ),
170 input: InputKind::Float64,
171 })
172 }
173
174 pub fn merge_udf_impl() -> AggregateUDF {
175 AggregateUDF::new_from_impl(AvgUdaf {
176 name: AVG_MERGE_NAME,
177 signature: Signature::exact(
178 vec![DataType::Binary],
179 datafusion::logical_expr::Volatility::Immutable,
180 ),
181 input: InputKind::Binary,
182 })
183 }
184
185 fn update_float64(&mut self, array: &ArrayRef) -> DfResult<()> {
186 let array = as_primitive_array::<Float64Type>(array)?;
187 let mut count = self.state.count;
188 for _ in array.iter().flatten() {
189 count = count.checked_add(1).ok_or_else(count_overflow)?;
190 }
191 let sum = sum(array)
192 .map(|batch_sum| self.state.sum + batch_sum)
193 .unwrap_or(self.state.sum);
194 self.state = AvgState { count, sum };
195 Ok(())
196 }
197
198 fn merge_states(&mut self, array: &ArrayRef) -> DfResult<()> {
199 let array = as_binary_array(array)?;
200 let states = array
201 .iter()
202 .flatten()
203 .map(AvgState::decode)
204 .collect::<DfResult<Vec<_>>>()?;
205 let count = states.iter().try_fold(self.state.count, |count, state| {
206 count.checked_add(state.count).ok_or_else(count_overflow)
207 })?;
208 let sums = states
209 .iter()
210 .filter(|state| state.count != 0)
211 .map(|state| Some(state.sum))
212 .collect::<Vec<_>>();
213 let sum = sum(&Float64Array::from(sums))
214 .map(|batch_sum| self.state.sum + batch_sum)
215 .unwrap_or(self.state.sum);
216 self.state = AvgState { count, sum };
217 Ok(())
218 }
219}
220
221impl DfAccumulator for AvgAccumulator {
222 fn update_batch(&mut self, values: &[ArrayRef]) -> DfResult<()> {
223 let array = &values[0];
224 match (self.input, array.data_type()) {
225 (InputKind::Float64, DataType::Float64) => self.update_float64(array),
226 (InputKind::Binary, DataType::Binary) => self.merge_states(array),
227 (_, data_type) => not_impl_err!("AVG input type does not match: {data_type:?}"),
228 }
229 }
230
231 fn evaluate(&mut self) -> DfResult<ScalarValue> {
232 Ok(ScalarValue::Binary(Some(self.state.encode().to_vec())))
233 }
234
235 fn size(&self) -> usize {
236 std::mem::size_of::<Self>()
237 }
238
239 fn state(&mut self) -> DfResult<Vec<ScalarValue>> {
240 Ok(vec![ScalarValue::Binary(Some(
241 self.state.encode().to_vec(),
242 ))])
243 }
244
245 fn merge_batch(&mut self, states: &[ArrayRef]) -> DfResult<()> {
246 self.merge_states(&states[0])
247 }
248}
249
250#[cfg(test)]
251mod tests {
252 use std::sync::Arc;
253
254 use arrow::array::{BinaryArray, Float64Array};
255 use datafusion_common::ScalarValue;
256 use datafusion_common::arrow::datatypes::DataType;
257 use datafusion_expr::TypeSignature;
258 use datafusion_physical_expr::aggregate::AggregateExprBuilder;
259 use datafusion_physical_expr::expressions::{Column, lit as physical_lit};
260
261 use super::*;
262 use crate::aggrs::aggr_wrapper::{aggr_delta_merge_func_name, aggr_state_func_name};
263 use crate::function_registry::FUNCTION_REGISTRY;
264
265 fn state(count: u64, sum: f64) -> Vec<u8> {
266 AvgState { count, sum }.encode().to_vec()
267 }
268
269 #[test]
270 fn codec_golden_and_roundtrip() {
271 let empty = AvgState::default().encode();
272 assert_eq!(empty.len(), ENCODED_LEN);
273 assert_eq!(empty.as_slice(), b"AVG1\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0");
274 let mut accumulator = AvgAccumulator::default();
275 accumulator
276 .update_batch(&[Arc::new(Float64Array::from(vec![Some(1.5)]))])
277 .unwrap();
278 let one = accumulator.state.encode();
279 assert_eq!(one.len(), ENCODED_LEN);
280 assert_eq!(
281 one.as_slice(),
282 &[
283 b'A', b'V', b'G', b'1', 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0xf8, 0x3f,
284 ]
285 );
286 assert_eq!(&one[12..20], &1.5f64.to_bits().to_le_bytes());
287 assert_eq!(AvgState::decode(&empty).unwrap().encode(), empty);
288 assert_eq!(AvgState::decode(&one).unwrap().encode(), one);
289 }
290
291 #[test]
292 fn codec_rejects_malformed_states() {
293 assert!(AvgState::decode(b"").is_err());
294 assert!(AvgState::decode(&[0; 19]).is_err());
295 assert!(AvgState::decode(&[0; 21]).is_err());
296 let mut avg2 = AvgState::default().encode();
297 avg2[..4].copy_from_slice(b"AVG2");
298 assert!(AvgState::decode(&avg2).is_err());
299 let mut wrong_magic = AvgState::default().encode();
300 wrong_magic[0] = b'X';
301 assert!(AvgState::decode(&wrong_magic).is_err());
302 for sum in [1.0, -0.0] {
303 assert!(AvgState::decode(&state(0, sum)).is_err());
304 }
305 let mut count = AvgState {
306 count: 0x0102_0304_0506_0708,
307 sum: 0.0,
308 }
309 .encode();
310 assert_eq!(&count[4..12], &0x0102_0304_0506_0708u64.to_le_bytes());
311 count[4..12].reverse();
312 assert_ne!(
313 AvgState::decode(&count).unwrap().count(),
314 0x0102_0304_0506_0708
315 );
316 let mut sum = AvgState { count: 1, sum: 1.5 }.encode();
317 assert_eq!(&sum[12..20], &1.5f64.to_bits().to_le_bytes());
318 sum[12..20].reverse();
319 assert_ne!(AvgState::decode(&sum).unwrap().average(), Some(1.5));
320 }
321
322 #[test]
323 fn codec_preserves_populated_float_bits() {
324 for bits in [
325 0.0f64.to_bits(),
326 (-0.0f64).to_bits(),
327 f64::INFINITY.to_bits(),
328 f64::NEG_INFINITY.to_bits(),
329 0x7ff8_0000_0000_0001,
330 0x7ff0_0000_0000_0001,
331 ] {
332 let encoded = state(1, f64::from_bits(bits));
333 assert_eq!(
334 AvgState::decode(&encoded).unwrap().encode().as_slice(),
335 encoded
336 );
337 }
338 }
339
340 #[test]
341 fn distinct_is_rejected() {
342 let udf = AvgAccumulator::state_udf_impl();
343 let schema = arrow_schema::Schema::empty();
344 let expr = physical_lit(1.0f64);
345 let field = Arc::new(arrow_schema::Field::new("in", DataType::Float64, true));
346 let args = AccumulatorArgs {
347 return_field: Arc::new(arrow_schema::Field::new("out", DataType::Binary, true)),
348 schema: &schema,
349 ignore_nulls: false,
350 order_bys: &[],
351 is_reversed: false,
352 name: AVG_STATE_NAME,
353 is_distinct: true,
354 exprs: std::slice::from_ref(&expr),
355 expr_fields: std::slice::from_ref(&field),
356 };
357 assert!(udf.accumulator(args).is_err());
358 }
359
360 #[test]
361 fn state_counts_nulls_and_empty_is_canonical() {
362 let mut accumulator = AvgAccumulator::default();
363 accumulator
364 .update_batch(&[Arc::new(Float64Array::from(vec![None, None]))])
365 .unwrap();
366 assert_eq!(accumulator.state.encode(), AvgState::default().encode());
367 accumulator
368 .update_batch(&[Arc::new(Float64Array::from(vec![
369 Some(1.0),
370 None,
371 Some(3.0),
372 Some(8.0),
373 ]))])
374 .unwrap();
375 assert_eq!(accumulator.state.count(), 3);
376 assert_eq!(accumulator.state.average(), Some(4.0));
377 }
378
379 #[test]
380 fn default_value_matches_empty_accumulator_evaluate() {
381 for udf in [
382 AvgAccumulator::state_udf_impl(),
383 AvgAccumulator::merge_udf_impl(),
384 ] {
385 let default = udf.default_value(&DataType::Binary).unwrap();
386 let expected = ScalarValue::Binary(Some(AvgState::default().encode().to_vec()));
387 assert_eq!(default, expected);
388 let mut empty = AvgAccumulator::default();
389 assert_eq!(
390 default,
391 empty.evaluate().unwrap(),
392 "{} default_value must equal the empty accumulator evaluate",
393 udf.name()
394 );
395 }
396 }
397
398 #[test]
399 fn merge_preserves_populated_negative_zero_for_empty_input() {
400 let mut accumulator = AvgAccumulator {
401 state: AvgState {
402 count: 1,
403 sum: -0.0,
404 },
405 input: InputKind::Binary,
406 };
407 let expected = accumulator.state.encode();
408 accumulator
409 .update_batch(&[Arc::new(BinaryArray::from(vec![
410 None,
411 Some(AvgState::default().encode().as_slice()),
412 ]))])
413 .unwrap();
414 assert_eq!(accumulator.state.encode(), expected);
415 }
416
417 #[test]
418 fn merge_ignores_nulls_and_merges_weighted_states() {
419 let mut accumulator = AvgAccumulator {
420 state: AvgState::default(),
421 input: InputKind::Binary,
422 };
423 accumulator
424 .update_batch(&[Arc::new(BinaryArray::from(vec![
425 Some(state(2, 4.0).as_slice()),
426 None,
427 Some(state(3, 15.0).as_slice()),
428 ]))])
429 .unwrap();
430 assert_eq!(accumulator.state.count(), 5);
431 assert_eq!(accumulator.state.average(), Some(19.0 / 5.0));
432 let before = accumulator.state;
433 assert!(
434 accumulator
435 .update_batch(&[Arc::new(BinaryArray::from(vec![Some(&[][..])]))])
436 .is_err()
437 );
438 assert_eq!(accumulator.state, before);
439 }
440
441 #[test]
442 fn overflow_does_not_mutate_update_or_merge() {
443 let mut update = AvgAccumulator {
444 state: AvgState {
445 count: u64::MAX,
446 sum: 1.0,
447 },
448 input: InputKind::Float64,
449 };
450 let before = update.state;
451 assert!(
452 update
453 .update_batch(&[Arc::new(Float64Array::from(vec![Some(2.0)]))])
454 .is_err()
455 );
456 assert_eq!(update.state, before);
457
458 let mut merge = AvgAccumulator {
459 state: AvgState {
460 count: u64::MAX,
461 sum: 1.0,
462 },
463 input: InputKind::Binary,
464 };
465 let before = merge.state;
466 assert!(
467 merge
468 .update_batch(&[Arc::new(BinaryArray::from(vec![Some(
469 state(1, 2.0).as_slice()
470 )]))])
471 .is_err()
472 );
473 assert_eq!(merge.state, before);
474 }
475
476 #[test]
477 fn registered_delta_merge_has_four_way_and_malformed_behavior() {
478 let udf = FUNCTION_REGISTRY
479 .get_aggr_func(&aggr_delta_merge_func_name(AVG_STATE_NAME))
480 .unwrap();
481 assert_eq!(udf.name(), "__avg_state_delta_merge");
482 assert_eq!(
483 udf.signature().type_signature,
484 TypeSignature::Exact(vec![DataType::Binary, DataType::Binary])
485 );
486 let schema = Arc::new(arrow_schema::Schema::new(vec![
487 arrow_schema::Field::new("delta", DataType::Binary, true),
488 arrow_schema::Field::new("persisted", DataType::Binary, true),
489 ]));
490 let expr = AggregateExprBuilder::new(
491 Arc::new(udf),
492 vec![
493 Arc::new(Column::new("delta", 0)),
494 Arc::new(Column::new("persisted", 1)),
495 ],
496 )
497 .schema(schema)
498 .alias("avg_delta_merge")
499 .build()
500 .unwrap();
501 let delta = state(2, 3.0);
502 let persisted = state(2, 7.0);
503 for (left, right, expected) in [
504 (
505 Some(delta.as_slice()),
506 None,
507 AvgState { count: 2, sum: 3.0 }.encode(),
508 ),
509 (
510 None,
511 Some(persisted.as_slice()),
512 AvgState { count: 2, sum: 7.0 }.encode(),
513 ),
514 (None, None, AvgState::default().encode()),
515 (
516 Some(delta.as_slice()),
517 Some(persisted.as_slice()),
518 AvgState {
519 count: 4,
520 sum: 10.0,
521 }
522 .encode(),
523 ),
524 ] {
525 let mut accumulator = expr.create_accumulator().unwrap();
526 accumulator
527 .update_batch(&[
528 Arc::new(BinaryArray::from(vec![left])),
529 Arc::new(BinaryArray::from(vec![right])),
530 ])
531 .unwrap();
532 let ScalarValue::Binary(Some(actual)) = accumulator.evaluate().unwrap() else {
533 panic!("AVG delta merge state must be binary");
534 };
535 assert_eq!(actual.as_slice(), expected.as_slice());
536 }
537 let mut accumulator = expr.create_accumulator().unwrap();
538 assert!(
539 accumulator
540 .update_batch(&[
541 Arc::new(BinaryArray::from(vec![Some(&[][..])])),
542 Arc::new(BinaryArray::from(vec![None])),
543 ])
544 .is_err()
545 );
546 let mut accumulator = expr.create_accumulator().unwrap();
547 assert!(
548 accumulator
549 .update_batch(&[
550 Arc::new(BinaryArray::from(vec![None])),
551 Arc::new(BinaryArray::from(vec![Some(&[][..])])),
552 ])
553 .is_err()
554 );
555 let mut accumulator = expr.create_accumulator().unwrap();
556 accumulator
557 .update_batch(&[
558 Arc::new(BinaryArray::from(vec![Some(delta.as_slice())])),
559 Arc::new(BinaryArray::from(vec![None])),
560 ])
561 .unwrap();
562 assert!(
563 accumulator
564 .update_batch(&[
565 Arc::new(BinaryArray::from(vec![Some(delta.as_slice())])),
566 Arc::new(BinaryArray::from(vec![Some(
567 AvgState {
568 count: u64::MAX,
569 sum: 1.0
570 }
571 .encode()
572 .as_slice()
573 )])),
574 ])
575 .is_err()
576 );
577 }
578
579 #[test]
580 fn avg_registry_does_not_replace_native_state_registry() {
581 let avg = FUNCTION_REGISTRY.get_aggr_func(AVG_STATE_NAME).unwrap();
582 let native = FUNCTION_REGISTRY
583 .get_aggr_func(&aggr_state_func_name("avg"))
584 .unwrap();
585 assert_eq!(
586 avg.return_type(&[DataType::Float64]).unwrap(),
587 DataType::Binary
588 );
589 assert!(matches!(
590 native.return_type(&[DataType::Float64]).unwrap(),
591 DataType::Struct(_)
592 ));
593 }
594}