datafusion_extra_functions/common/mode/
bytes.rs1use 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 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 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 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}