Skip to main content

datafusion_extra_functions/common/mode/
native.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18use datafusion::common::utils::SingleRowListArrayBuilder;
19use datafusion::physical_expr::aggregate::utils::Hashable;
20use datafusion::{arrow, common, error, logical_expr, scalar};
21use std::{cmp, collections, fmt, hash, mem, sync};
22
23#[derive(fmt::Debug)]
24pub struct PrimitiveModeAccumulator<T>
25where
26    T: arrow::array::ArrowPrimitiveType + Send,
27    T::Native: Eq + hash::Hash,
28{
29    value_counts: collections::HashMap<T::Native, i64>,
30    data_type: arrow::datatypes::DataType,
31}
32
33impl<T> PrimitiveModeAccumulator<T>
34where
35    T: arrow::array::ArrowPrimitiveType + Send,
36    T::Native: Eq + hash::Hash + Clone,
37{
38    pub fn new(data_type: &arrow::datatypes::DataType) -> Self {
39        Self {
40            value_counts: collections::HashMap::default(),
41            data_type: data_type.clone(),
42        }
43    }
44}
45
46impl<T> logical_expr::Accumulator for PrimitiveModeAccumulator<T>
47where
48    T: arrow::array::ArrowPrimitiveType + Send + fmt::Debug,
49    T::Native: Eq + hash::Hash + Clone + PartialOrd + fmt::Debug,
50{
51    fn update_batch(&mut self, values: &[arrow::array::ArrayRef]) -> error::Result<()> {
52        if values.is_empty() {
53            return Ok(());
54        }
55        let arr = common::cast::as_primitive_array::<T>(&values[0])?;
56
57        for value in arr.iter().flatten() {
58            let counter = self.value_counts.entry(value).or_insert(0);
59            *counter += 1;
60        }
61
62        Ok(())
63    }
64
65    fn state(&mut self) -> error::Result<Vec<scalar::ScalarValue>> {
66        let values =
67            arrow::array::PrimitiveArray::<T>::from_iter_values(self.value_counts.keys().copied())
68                .with_data_type(self.data_type.clone());
69        let counts =
70            arrow::array::Int64Array::from_iter_values(self.value_counts.values().copied());
71
72        Ok(vec![
73            SingleRowListArrayBuilder::new(sync::Arc::new(values)).build_list_scalar(),
74            SingleRowListArrayBuilder::new(sync::Arc::new(counts)).build_list_scalar(),
75        ])
76    }
77
78    fn merge_batch(&mut self, states: &[arrow::array::ArrayRef]) -> error::Result<()> {
79        super::for_each_state_row(states, |values, counts| {
80            let values = common::cast::as_primitive_array::<T>(values)?;
81            for (value, count) in values.iter().zip(counts.values()) {
82                if let Some(value) = value {
83                    *self.value_counts.entry(value).or_insert(0) += *count;
84                }
85            }
86            Ok(())
87        })
88    }
89
90    fn evaluate(&mut self) -> error::Result<scalar::ScalarValue> {
91        let mut max_value: Option<T::Native> = None;
92        let mut max_count: i64 = 0;
93
94        self.value_counts.iter().for_each(|(value, &count)| {
95            match count.cmp(&max_count) {
96                cmp::Ordering::Greater => {
97                    max_value = Some(*value);
98                    max_count = count;
99                }
100                cmp::Ordering::Equal => {
101                    // On ties the smallest value wins, matching PostgreSQL's mode().
102                    max_value = match max_value {
103                        Some(ref current_max_value) if value < current_max_value => Some(*value),
104                        Some(ref current_max_value) => Some(*current_max_value),
105                        None => Some(*value),
106                    };
107                }
108                _ => {} // Do nothing if count is less than max_count
109            }
110        });
111
112        scalar::ScalarValue::new_primitive::<T>(max_value, &self.data_type)
113    }
114
115    fn size(&self) -> usize {
116        mem::size_of_val(&self.value_counts)
117            + self.value_counts.len() * mem::size_of::<(T::Native, i64)>()
118    }
119}
120
121#[derive(Debug)]
122pub struct FloatModeAccumulator<T>
123where
124    T: arrow::array::ArrowPrimitiveType,
125{
126    value_counts: collections::HashMap<Hashable<T::Native>, i64>,
127    data_type: arrow::datatypes::DataType,
128}
129
130impl<T> FloatModeAccumulator<T>
131where
132    T: arrow::array::ArrowPrimitiveType,
133{
134    pub fn new(data_type: &arrow::datatypes::DataType) -> Self {
135        Self {
136            value_counts: collections::HashMap::default(),
137            data_type: data_type.clone(),
138        }
139    }
140}
141
142impl<T> logical_expr::Accumulator for FloatModeAccumulator<T>
143where
144    T: arrow::array::ArrowPrimitiveType + Send + fmt::Debug,
145    T::Native: PartialOrd + fmt::Debug + Clone,
146{
147    fn update_batch(&mut self, values: &[arrow::array::ArrayRef]) -> error::Result<()> {
148        if values.is_empty() {
149            return Ok(());
150        }
151
152        let arr = common::cast::as_primitive_array::<T>(&values[0])?;
153
154        for value in arr.iter().flatten() {
155            let counter = self.value_counts.entry(Hashable(value)).or_insert(0);
156            *counter += 1;
157        }
158
159        Ok(())
160    }
161
162    fn state(&mut self) -> error::Result<Vec<scalar::ScalarValue>> {
163        let values = arrow::array::PrimitiveArray::<T>::from_iter_values(
164            self.value_counts.keys().map(|key| key.0),
165        )
166        .with_data_type(self.data_type.clone());
167        let counts =
168            arrow::array::Int64Array::from_iter_values(self.value_counts.values().copied());
169
170        Ok(vec![
171            SingleRowListArrayBuilder::new(sync::Arc::new(values)).build_list_scalar(),
172            SingleRowListArrayBuilder::new(sync::Arc::new(counts)).build_list_scalar(),
173        ])
174    }
175
176    fn merge_batch(&mut self, states: &[arrow::array::ArrayRef]) -> error::Result<()> {
177        super::for_each_state_row(states, |values, counts| {
178            let values = common::cast::as_primitive_array::<T>(values)?;
179            for (value, count) in values.iter().zip(counts.values()) {
180                if let Some(value) = value {
181                    *self.value_counts.entry(Hashable(value)).or_insert(0) += *count;
182                }
183            }
184            Ok(())
185        })
186    }
187
188    fn evaluate(&mut self) -> error::Result<scalar::ScalarValue> {
189        let mut max_value: Option<T::Native> = None;
190        let mut max_count: i64 = 0;
191
192        self.value_counts.iter().for_each(|(value, &count)| {
193            match count.cmp(&max_count) {
194                cmp::Ordering::Greater => {
195                    max_value = Some(value.0);
196                    max_count = count;
197                }
198                cmp::Ordering::Equal => {
199                    // On ties the smallest value wins, matching PostgreSQL's mode().
200                    max_value = match max_value {
201                        Some(current_max_value) if value.0 < current_max_value => Some(value.0),
202                        Some(current_max_value) => Some(current_max_value),
203                        None => Some(value.0),
204                    };
205                }
206                _ => {} // Do nothing if count is less than max_count
207            }
208        });
209
210        scalar::ScalarValue::new_primitive::<T>(max_value, &self.data_type)
211    }
212
213    fn size(&self) -> usize {
214        mem::size_of_val(&self.value_counts)
215            + self.value_counts.len() * mem::size_of::<(Hashable<T::Native>, i64)>()
216    }
217}
218
219#[cfg(test)]
220mod tests {
221
222    use super::*;
223
224    use datafusion::logical_expr::Accumulator;
225    use std::sync;
226
227    fn merge_from(dest: &mut impl Accumulator, src: &mut impl Accumulator) -> error::Result<()> {
228        let arrays = src
229            .state()?
230            .iter()
231            .map(|value| value.to_array())
232            .collect::<error::Result<Vec<_>>>()?;
233        dest.merge_batch(&arrays)
234    }
235
236    #[test]
237    fn test_mode_accumulator_single_mode_int64() -> error::Result<()> {
238        let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Int64Type>::new(
239            &arrow::datatypes::DataType::Int64,
240        );
241        let values: arrow::array::ArrayRef =
242            sync::Arc::new(arrow::array::Int64Array::from(vec![1, 2, 2, 3, 3, 3]));
243        acc.update_batch(&[values])?;
244        let result = acc.evaluate()?;
245        assert_eq!(
246            result,
247            scalar::ScalarValue::new_primitive::<arrow::datatypes::Int64Type>(
248                Some(3),
249                &arrow::datatypes::DataType::Int64
250            )?
251        );
252        Ok(())
253    }
254
255    #[test]
256    fn test_mode_accumulator_with_nulls_int64() -> error::Result<()> {
257        let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Int64Type>::new(
258            &arrow::datatypes::DataType::Int64,
259        );
260        let values: arrow::array::ArrayRef = sync::Arc::new(arrow::array::Int64Array::from(vec![
261            None,
262            Some(1),
263            Some(2),
264            Some(2),
265            Some(3),
266            Some(3),
267            Some(3),
268        ]));
269        acc.update_batch(&[values])?;
270        let result = acc.evaluate()?;
271        assert_eq!(
272            result,
273            scalar::ScalarValue::new_primitive::<arrow::datatypes::Int64Type>(
274                Some(3),
275                &arrow::datatypes::DataType::Int64
276            )?
277        );
278        Ok(())
279    }
280
281    #[test]
282    fn test_mode_accumulator_tie_case_int64() -> error::Result<()> {
283        let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Int64Type>::new(
284            &arrow::datatypes::DataType::Int64,
285        );
286        let values: arrow::array::ArrayRef =
287            sync::Arc::new(arrow::array::Int64Array::from(vec![1, 2, 2, 3, 3]));
288        acc.update_batch(&[values])?;
289        let result = acc.evaluate()?;
290        assert_eq!(
291            result,
292            scalar::ScalarValue::new_primitive::<arrow::datatypes::Int64Type>(
293                Some(2),
294                &arrow::datatypes::DataType::Int64
295            )?
296        );
297        Ok(())
298    }
299
300    #[test]
301    fn test_mode_accumulator_only_nulls_int64() -> error::Result<()> {
302        let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Int64Type>::new(
303            &arrow::datatypes::DataType::Int64,
304        );
305        let values: arrow::array::ArrayRef =
306            sync::Arc::new(arrow::array::Int64Array::from(vec![None, None, None, None]));
307        acc.update_batch(&[values])?;
308        let result = acc.evaluate()?;
309        assert_eq!(
310            result,
311            scalar::ScalarValue::new_primitive::<arrow::datatypes::Int64Type>(
312                None,
313                &arrow::datatypes::DataType::Int64
314            )?
315        );
316        Ok(())
317    }
318
319    #[test]
320    fn test_mode_accumulator_merge_overlapping_keys_int64() -> error::Result<()> {
321        let mut left = PrimitiveModeAccumulator::<arrow::datatypes::Int64Type>::new(
322            &arrow::datatypes::DataType::Int64,
323        );
324        let left_values: arrow::array::ArrayRef =
325            sync::Arc::new(arrow::array::Int64Array::from(vec![2, 2, 2]));
326        left.update_batch(&[left_values])?;
327
328        let mut right = PrimitiveModeAccumulator::<arrow::datatypes::Int64Type>::new(
329            &arrow::datatypes::DataType::Int64,
330        );
331        // Right-only or replace-instead-of-add both pick 1. Summed counts pick 2.
332        let right_values: arrow::array::ArrayRef =
333            sync::Arc::new(arrow::array::Int64Array::from(vec![1, 1, 1, 1, 2, 2]));
334        right.update_batch(&[right_values])?;
335
336        merge_from(&mut right, &mut left)?;
337        let result = right.evaluate()?;
338        assert_eq!(
339            result,
340            scalar::ScalarValue::new_primitive::<arrow::datatypes::Int64Type>(
341                Some(2),
342                &arrow::datatypes::DataType::Int64
343            )?
344        );
345        Ok(())
346    }
347
348    #[test]
349    fn test_mode_accumulator_single_mode_float64() -> error::Result<()> {
350        let mut acc = FloatModeAccumulator::<arrow::datatypes::Float64Type>::new(
351            &arrow::datatypes::DataType::Float64,
352        );
353        let values: arrow::array::ArrayRef =
354            sync::Arc::new(arrow::array::Float64Array::from(vec![
355                1.0, 2.0, 2.0, 3.0, 3.0, 3.0,
356            ]));
357        acc.update_batch(&[values])?;
358        let result = acc.evaluate()?;
359        assert_eq!(
360            result,
361            scalar::ScalarValue::new_primitive::<arrow::datatypes::Float64Type>(
362                Some(3.0),
363                &arrow::datatypes::DataType::Float64
364            )?
365        );
366        Ok(())
367    }
368
369    #[test]
370    fn test_mode_accumulator_with_nulls_float64() -> error::Result<()> {
371        let mut acc = FloatModeAccumulator::<arrow::datatypes::Float64Type>::new(
372            &arrow::datatypes::DataType::Float64,
373        );
374        let values: arrow::array::ArrayRef =
375            sync::Arc::new(arrow::array::Float64Array::from(vec![
376                None,
377                Some(1.0),
378                Some(2.0),
379                Some(2.0),
380                Some(3.0),
381                Some(3.0),
382                Some(3.0),
383            ]));
384        acc.update_batch(&[values])?;
385        let result = acc.evaluate()?;
386        assert_eq!(
387            result,
388            scalar::ScalarValue::new_primitive::<arrow::datatypes::Float64Type>(
389                Some(3.0),
390                &arrow::datatypes::DataType::Float64
391            )?
392        );
393        Ok(())
394    }
395
396    #[test]
397    fn test_mode_accumulator_tie_case_float64() -> error::Result<()> {
398        let mut acc = FloatModeAccumulator::<arrow::datatypes::Float64Type>::new(
399            &arrow::datatypes::DataType::Float64,
400        );
401        let values: arrow::array::ArrayRef =
402            sync::Arc::new(arrow::array::Float64Array::from(vec![
403                1.0, 2.0, 2.0, 3.0, 3.0,
404            ]));
405        acc.update_batch(&[values])?;
406        let result = acc.evaluate()?;
407        assert_eq!(
408            result,
409            scalar::ScalarValue::new_primitive::<arrow::datatypes::Float64Type>(
410                Some(2.0),
411                &arrow::datatypes::DataType::Float64
412            )?
413        );
414        Ok(())
415    }
416
417    #[test]
418    fn test_mode_accumulator_only_nulls_float64() -> error::Result<()> {
419        let mut acc = FloatModeAccumulator::<arrow::datatypes::Float64Type>::new(
420            &arrow::datatypes::DataType::Float64,
421        );
422        let values: arrow::array::ArrayRef =
423            sync::Arc::new(arrow::array::Float64Array::from(vec![
424                None, None, None, None,
425            ]));
426        acc.update_batch(&[values])?;
427        let result = acc.evaluate()?;
428        assert_eq!(
429            result,
430            scalar::ScalarValue::new_primitive::<arrow::datatypes::Float64Type>(
431                None,
432                &arrow::datatypes::DataType::Float64
433            )?
434        );
435        Ok(())
436    }
437
438    #[test]
439    fn test_mode_accumulator_merge_overlapping_keys_float64() -> error::Result<()> {
440        let mut left = FloatModeAccumulator::<arrow::datatypes::Float64Type>::new(
441            &arrow::datatypes::DataType::Float64,
442        );
443        let left_values: arrow::array::ArrayRef =
444            sync::Arc::new(arrow::array::Float64Array::from(vec![2.0, 2.0, 2.0]));
445        left.update_batch(&[left_values])?;
446
447        let mut right = FloatModeAccumulator::<arrow::datatypes::Float64Type>::new(
448            &arrow::datatypes::DataType::Float64,
449        );
450        let right_values: arrow::array::ArrayRef =
451            sync::Arc::new(arrow::array::Float64Array::from(vec![
452                1.0, 1.0, 1.0, 1.0, 2.0, 2.0,
453            ]));
454        right.update_batch(&[right_values])?;
455
456        merge_from(&mut right, &mut left)?;
457        let result = right.evaluate()?;
458        assert_eq!(
459            result,
460            scalar::ScalarValue::new_primitive::<arrow::datatypes::Float64Type>(
461                Some(2.0),
462                &arrow::datatypes::DataType::Float64
463            )?
464        );
465        Ok(())
466    }
467
468    #[test]
469    fn test_mode_accumulator_single_mode_date64() -> error::Result<()> {
470        let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Date64Type>::new(
471            &arrow::datatypes::DataType::Date64,
472        );
473        let values: arrow::array::ArrayRef = sync::Arc::new(arrow::array::Date64Array::from(vec![
474            1609459200000,
475            1609545600000,
476            1609545600000,
477            1609632000000,
478            1609632000000,
479            1609632000000,
480        ]));
481        acc.update_batch(&[values])?;
482        let result = acc.evaluate()?;
483        assert_eq!(
484            result,
485            scalar::ScalarValue::new_primitive::<arrow::datatypes::Date64Type>(
486                Some(1609632000000),
487                &arrow::datatypes::DataType::Date64
488            )?
489        );
490        Ok(())
491    }
492
493    #[test]
494    fn test_mode_accumulator_with_nulls_date64() -> error::Result<()> {
495        let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Date64Type>::new(
496            &arrow::datatypes::DataType::Date64,
497        );
498        let values: arrow::array::ArrayRef = sync::Arc::new(arrow::array::Date64Array::from(vec![
499            None,
500            Some(1609459200000),
501            Some(1609545600000),
502            Some(1609545600000),
503            Some(1609632000000),
504            Some(1609632000000),
505            Some(1609632000000),
506        ]));
507        acc.update_batch(&[values])?;
508        let result = acc.evaluate()?;
509        assert_eq!(
510            result,
511            scalar::ScalarValue::new_primitive::<arrow::datatypes::Date64Type>(
512                Some(1609632000000),
513                &arrow::datatypes::DataType::Date64
514            )?
515        );
516        Ok(())
517    }
518
519    #[test]
520    fn test_mode_accumulator_tie_case_date64() -> error::Result<()> {
521        let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Date64Type>::new(
522            &arrow::datatypes::DataType::Date64,
523        );
524        let values: arrow::array::ArrayRef = sync::Arc::new(arrow::array::Date64Array::from(vec![
525            1609459200000,
526            1609545600000,
527            1609545600000,
528            1609632000000,
529            1609632000000,
530        ]));
531        acc.update_batch(&[values])?;
532        let result = acc.evaluate()?;
533        assert_eq!(
534            result,
535            scalar::ScalarValue::new_primitive::<arrow::datatypes::Date64Type>(
536                Some(1609545600000),
537                &arrow::datatypes::DataType::Date64
538            )?
539        );
540        Ok(())
541    }
542
543    #[test]
544    fn test_mode_accumulator_only_nulls_date64() -> error::Result<()> {
545        let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Date64Type>::new(
546            &arrow::datatypes::DataType::Date64,
547        );
548        let values: arrow::array::ArrayRef = sync::Arc::new(arrow::array::Date64Array::from(vec![
549            None, None, None, None,
550        ]));
551        acc.update_batch(&[values])?;
552        let result = acc.evaluate()?;
553        assert_eq!(
554            result,
555            scalar::ScalarValue::new_primitive::<arrow::datatypes::Date64Type>(
556                None,
557                &arrow::datatypes::DataType::Date64
558            )?
559        );
560        Ok(())
561    }
562
563    #[test]
564    fn test_mode_accumulator_single_mode_time64() -> error::Result<()> {
565        let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Time64MicrosecondType>::new(
566            &arrow::datatypes::DataType::Time64(arrow::datatypes::TimeUnit::Microsecond),
567        );
568        let values: arrow::array::ArrayRef =
569            sync::Arc::new(arrow::array::Time64MicrosecondArray::from(vec![
570                3600000000,
571                7200000000,
572                7200000000,
573                10800000000,
574                10800000000,
575                10800000000,
576            ]));
577        acc.update_batch(&[values])?;
578        let result = acc.evaluate()?;
579        assert_eq!(
580            result,
581            scalar::ScalarValue::new_primitive::<arrow::datatypes::Time64MicrosecondType>(
582                Some(10800000000),
583                &arrow::datatypes::DataType::Time64(arrow::datatypes::TimeUnit::Microsecond)
584            )?
585        );
586        Ok(())
587    }
588
589    #[test]
590    fn test_mode_accumulator_with_nulls_time64() -> error::Result<()> {
591        let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Time64MicrosecondType>::new(
592            &arrow::datatypes::DataType::Time64(arrow::datatypes::TimeUnit::Microsecond),
593        );
594        let values: arrow::array::ArrayRef =
595            sync::Arc::new(arrow::array::Time64MicrosecondArray::from(vec![
596                None,
597                Some(3600000000),
598                Some(7200000000),
599                Some(7200000000),
600                Some(10800000000),
601                Some(10800000000),
602                Some(10800000000),
603            ]));
604        acc.update_batch(&[values])?;
605        let result = acc.evaluate()?;
606        assert_eq!(
607            result,
608            scalar::ScalarValue::new_primitive::<arrow::datatypes::Time64MicrosecondType>(
609                Some(10800000000),
610                &arrow::datatypes::DataType::Time64(arrow::datatypes::TimeUnit::Microsecond)
611            )?
612        );
613        Ok(())
614    }
615
616    #[test]
617    fn test_mode_accumulator_tie_case_time64() -> error::Result<()> {
618        let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Time64MicrosecondType>::new(
619            &arrow::datatypes::DataType::Time64(arrow::datatypes::TimeUnit::Microsecond),
620        );
621        let values: arrow::array::ArrayRef =
622            sync::Arc::new(arrow::array::Time64MicrosecondArray::from(vec![
623                3600000000,
624                7200000000,
625                7200000000,
626                10800000000,
627                10800000000,
628            ]));
629        acc.update_batch(&[values])?;
630        let result = acc.evaluate()?;
631        assert_eq!(
632            result,
633            scalar::ScalarValue::new_primitive::<arrow::datatypes::Time64MicrosecondType>(
634                Some(7200000000),
635                &arrow::datatypes::DataType::Time64(arrow::datatypes::TimeUnit::Microsecond)
636            )?
637        );
638        Ok(())
639    }
640
641    #[test]
642    fn test_mode_accumulator_only_nulls_time64() -> error::Result<()> {
643        let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Time64MicrosecondType>::new(
644            &arrow::datatypes::DataType::Time64(arrow::datatypes::TimeUnit::Microsecond),
645        );
646        let values: arrow::array::ArrayRef =
647            sync::Arc::new(arrow::array::Time64MicrosecondArray::from(vec![
648                None, None, None, None,
649            ]));
650        acc.update_batch(&[values])?;
651        let result = acc.evaluate()?;
652        assert_eq!(
653            result,
654            scalar::ScalarValue::new_primitive::<arrow::datatypes::Time64MicrosecondType>(
655                None,
656                &arrow::datatypes::DataType::Time64(arrow::datatypes::TimeUnit::Microsecond)
657            )?
658        );
659        Ok(())
660    }
661}