Skip to main content

datafusion_extra_functions/common/mode/
bytes.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::arrow::array::AsArray;
19use datafusion::common::utils::SingleRowListArrayBuilder;
20use datafusion::{arrow, common, error, logical_expr, scalar};
21use std::{collections, mem, sync};
22
23#[derive(Debug)]
24pub struct BytesModeAccumulator {
25    value_counts: collections::HashMap<String, i64>,
26    data_type: arrow::datatypes::DataType,
27}
28
29impl BytesModeAccumulator {
30    pub fn new(data_type: &arrow::datatypes::DataType) -> Self {
31        Self {
32            value_counts: collections::HashMap::new(),
33            data_type: data_type.clone(),
34        }
35    }
36
37    fn update_counts<'a, V>(&mut self, array: V)
38    where
39        V: arrow::array::ArrayAccessor<Item = &'a str>,
40    {
41        for value in arrow::array::ArrayIter::new(array).flatten() {
42            // get_mut before insert avoids allocating a String for keys
43            // that are already present.
44            if let Some(count) = self.value_counts.get_mut(value) {
45                *count += 1;
46            } else {
47                self.value_counts.insert(value.to_string(), 1);
48            }
49        }
50    }
51}
52
53impl logical_expr::Accumulator for BytesModeAccumulator {
54    fn update_batch(&mut self, values: &[arrow::array::ArrayRef]) -> error::Result<()> {
55        if values.is_empty() {
56            return Ok(());
57        }
58
59        match &self.data_type {
60            arrow::datatypes::DataType::Utf8View => {
61                let array = values[0].as_string_view();
62                self.update_counts(array);
63            }
64            _ => {
65                let array = values[0].as_string::<i32>();
66                self.update_counts(array);
67            }
68        };
69
70        Ok(())
71    }
72
73    fn state(&mut self) -> error::Result<Vec<scalar::ScalarValue>> {
74        let values = arrow::array::StringArray::from_iter_values(self.value_counts.keys());
75        let counts =
76            arrow::array::Int64Array::from_iter_values(self.value_counts.values().copied());
77
78        Ok(vec![
79            SingleRowListArrayBuilder::new(sync::Arc::new(values)).build_list_scalar(),
80            SingleRowListArrayBuilder::new(sync::Arc::new(counts)).build_list_scalar(),
81        ])
82    }
83
84    fn merge_batch(&mut self, states: &[arrow::array::ArrayRef]) -> error::Result<()> {
85        super::for_each_state_row(states, |values, counts| {
86            let values = common::cast::as_string_array(values)?;
87            for (value, count) in values.iter().zip(counts.values()) {
88                if let Some(value) = value {
89                    *self.value_counts.entry(value.to_string()).or_insert(0) += *count;
90                }
91            }
92            Ok(())
93        })
94    }
95
96    fn evaluate(&mut self) -> error::Result<scalar::ScalarValue> {
97        let mode = self
98            .value_counts
99            .iter()
100            .max_by(|a, b| {
101                // Highest count wins; on ties the smallest value wins,
102                // matching PostgreSQL's mode() ordered-set aggregate.
103                a.1.cmp(b.1).then_with(|| b.0.cmp(a.0))
104            })
105            .map(|(value, _)| value.to_string());
106
107        match &self.data_type {
108            arrow::datatypes::DataType::Utf8View => Ok(scalar::ScalarValue::Utf8View(mode)),
109            _ => Ok(scalar::ScalarValue::Utf8(mode)),
110        }
111    }
112
113    fn size(&self) -> usize {
114        self.value_counts.capacity() * mem::size_of::<(String, i64)>()
115            + mem::size_of_val(&self.data_type)
116    }
117}
118
119#[cfg(test)]
120mod tests {
121
122    use super::*;
123
124    use datafusion::logical_expr::Accumulator;
125    use std::sync;
126
127    fn merge_from(dest: &mut impl Accumulator, src: &mut impl Accumulator) -> error::Result<()> {
128        let arrays = src
129            .state()?
130            .iter()
131            .map(|value| value.to_array())
132            .collect::<error::Result<Vec<_>>>()?;
133        dest.merge_batch(&arrays)
134    }
135
136    #[test]
137    fn test_mode_accumulator_single_mode_utf8() -> error::Result<()> {
138        let mut acc = BytesModeAccumulator::new(&arrow::datatypes::DataType::Utf8);
139        let values: arrow::array::ArrayRef = sync::Arc::new(arrow::array::StringArray::from(vec![
140            Some("apple"),
141            Some("banana"),
142            Some("apple"),
143            Some("orange"),
144            Some("banana"),
145            Some("apple"),
146        ]));
147
148        acc.update_batch(&[values])?;
149        let result = acc.evaluate()?;
150
151        assert_eq!(result, scalar::ScalarValue::Utf8(Some("apple".to_string())));
152        Ok(())
153    }
154
155    #[test]
156    fn test_mode_accumulator_tie_utf8() -> error::Result<()> {
157        let mut acc = BytesModeAccumulator::new(&arrow::datatypes::DataType::Utf8);
158        let values: arrow::array::ArrayRef = sync::Arc::new(arrow::array::StringArray::from(vec![
159            Some("apple"),
160            Some("banana"),
161            Some("apple"),
162            Some("orange"),
163            Some("banana"),
164        ]));
165
166        acc.update_batch(&[values])?;
167        let result = acc.evaluate()?;
168
169        assert_eq!(result, scalar::ScalarValue::Utf8(Some("apple".to_string())));
170        Ok(())
171    }
172
173    #[test]
174    fn test_mode_accumulator_all_nulls_utf8() -> error::Result<()> {
175        let mut acc = BytesModeAccumulator::new(&arrow::datatypes::DataType::Utf8);
176        let values: arrow::array::ArrayRef = sync::Arc::new(arrow::array::StringArray::from(vec![
177            None as Option<&str>,
178            None,
179            None,
180        ]));
181
182        acc.update_batch(&[values])?;
183        let result = acc.evaluate()?;
184
185        assert_eq!(result, scalar::ScalarValue::Utf8(None));
186        Ok(())
187    }
188
189    #[test]
190    fn test_mode_accumulator_with_nulls_utf8() -> error::Result<()> {
191        let mut acc = BytesModeAccumulator::new(&arrow::datatypes::DataType::Utf8);
192        let values: arrow::array::ArrayRef = sync::Arc::new(arrow::array::StringArray::from(vec![
193            Some("apple"),
194            None,
195            Some("banana"),
196            Some("apple"),
197            None,
198            None,
199            None,
200            Some("banana"),
201        ]));
202
203        acc.update_batch(&[values])?;
204        let result = acc.evaluate()?;
205
206        assert_eq!(result, scalar::ScalarValue::Utf8(Some("apple".to_string())));
207        Ok(())
208    }
209
210    #[test]
211    fn test_mode_accumulator_single_mode_utf8view() -> error::Result<()> {
212        let mut acc = BytesModeAccumulator::new(&arrow::datatypes::DataType::Utf8View);
213        let values: arrow::array::ArrayRef =
214            sync::Arc::new(arrow::array::GenericByteViewArray::from(vec![
215                Some("apple"),
216                Some("banana"),
217                Some("apple"),
218                Some("orange"),
219                Some("banana"),
220                Some("apple"),
221            ]));
222
223        acc.update_batch(&[values])?;
224        let result = acc.evaluate()?;
225
226        assert_eq!(
227            result,
228            scalar::ScalarValue::Utf8View(Some("apple".to_string()))
229        );
230        Ok(())
231    }
232
233    #[test]
234    fn test_mode_accumulator_tie_utf8view() -> error::Result<()> {
235        let mut acc = BytesModeAccumulator::new(&arrow::datatypes::DataType::Utf8View);
236        let values: arrow::array::ArrayRef =
237            sync::Arc::new(arrow::array::GenericByteViewArray::from(vec![
238                Some("apple"),
239                Some("banana"),
240                Some("apple"),
241                Some("orange"),
242                Some("banana"),
243            ]));
244
245        acc.update_batch(&[values])?;
246        let result = acc.evaluate()?;
247
248        assert_eq!(
249            result,
250            scalar::ScalarValue::Utf8View(Some("apple".to_string()))
251        );
252        Ok(())
253    }
254
255    #[test]
256    fn test_mode_accumulator_all_nulls_utf8view() -> error::Result<()> {
257        let mut acc = BytesModeAccumulator::new(&arrow::datatypes::DataType::Utf8View);
258        let values: arrow::array::ArrayRef =
259            sync::Arc::new(arrow::array::GenericByteViewArray::from(vec![
260                None as Option<&str>,
261                None,
262                None,
263            ]));
264
265        acc.update_batch(&[values])?;
266        let result = acc.evaluate()?;
267
268        assert_eq!(result, scalar::ScalarValue::Utf8View(None));
269        Ok(())
270    }
271
272    #[test]
273    fn test_mode_accumulator_with_nulls_utf8view() -> error::Result<()> {
274        let mut acc = BytesModeAccumulator::new(&arrow::datatypes::DataType::Utf8View);
275        let values: arrow::array::ArrayRef =
276            sync::Arc::new(arrow::array::GenericByteViewArray::from(vec![
277                Some("apple"),
278                None,
279                Some("banana"),
280                Some("apple"),
281                None,
282                None,
283                None,
284                Some("banana"),
285            ]));
286
287        acc.update_batch(&[values])?;
288        let result = acc.evaluate()?;
289
290        assert_eq!(
291            result,
292            scalar::ScalarValue::Utf8View(Some("apple".to_string()))
293        );
294        Ok(())
295    }
296
297    #[test]
298    fn test_mode_accumulator_merge_overlapping_keys_utf8() -> error::Result<()> {
299        let mut left = BytesModeAccumulator::new(&arrow::datatypes::DataType::Utf8);
300        let left_values: arrow::array::ArrayRef =
301            sync::Arc::new(arrow::array::StringArray::from(vec![
302                Some("banana"),
303                Some("banana"),
304                Some("banana"),
305            ]));
306        left.update_batch(&[left_values])?;
307
308        let mut right = BytesModeAccumulator::new(&arrow::datatypes::DataType::Utf8);
309        // Right-only or replace-instead-of-add both pick apple. Summed counts pick banana.
310        let right_values: arrow::array::ArrayRef =
311            sync::Arc::new(arrow::array::StringArray::from(vec![
312                Some("apple"),
313                Some("apple"),
314                Some("apple"),
315                Some("apple"),
316                Some("banana"),
317                Some("banana"),
318            ]));
319        right.update_batch(&[right_values])?;
320
321        merge_from(&mut right, &mut left)?;
322        let result = right.evaluate()?;
323        assert_eq!(
324            result,
325            scalar::ScalarValue::Utf8(Some("banana".to_string()))
326        );
327        Ok(())
328    }
329}