Skip to main content

common_function/aggrs/approximate/
avg.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use 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/// The portable state used by the Float64 average aggregate functions.
34#[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    /// Returns the exact AVG1 representation of this state.
48    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    /// Decodes and validates an AVG1 state.
57    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    /// Returns the number of non-null input values in this state.
70    pub(crate) fn count(&self) -> u64 {
71        self.count
72    }
73
74    /// Returns the average, or `None` for the canonical empty state.
75    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/// The `avg_state` / `avg_merge` aggregate UDF.
95///
96/// Declares a canonical AVG1 empty state as `default_value` so window frames
97/// without rows observe the same contract as the accumulator's `evaluate`.
98#[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/// Accumulates and merges AVG1 states.
147#[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}