Skip to main content

datafusion_functions_aggregate_common/aggregate/count_distinct/
groups.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 arrow::array::{
19    Array, ArrayRef, AsArray, BooleanArray, Int64Array, ListArray, ListBuilder,
20    PrimitiveArray, PrimitiveBuilder,
21};
22use arrow::buffer::{OffsetBuffer, ScalarBuffer};
23use arrow::datatypes::{ArrowPrimitiveType, Field};
24use datafusion_common::HashSet;
25use datafusion_common::hash_utils::RandomState;
26use datafusion_expr_common::groups_accumulator::{EmitTo, GroupsAccumulator};
27use std::hash::Hash;
28use std::mem::size_of;
29use std::sync::Arc;
30
31use crate::aggregate::groups_accumulator::accumulate::accumulate;
32
33pub struct PrimitiveDistinctCountGroupsAccumulator<T: ArrowPrimitiveType>
34where
35    T::Native: Eq + Hash,
36{
37    seen: HashSet<(usize, T::Native), RandomState>,
38    counts: Vec<i64>,
39}
40
41impl<T: ArrowPrimitiveType> PrimitiveDistinctCountGroupsAccumulator<T>
42where
43    T::Native: Eq + Hash,
44{
45    pub fn new() -> Self {
46        Self {
47            seen: HashSet::default(),
48            counts: Vec::new(),
49        }
50    }
51}
52
53impl<T: ArrowPrimitiveType> Default for PrimitiveDistinctCountGroupsAccumulator<T>
54where
55    T::Native: Eq + Hash,
56{
57    fn default() -> Self {
58        Self::new()
59    }
60}
61
62impl<T: ArrowPrimitiveType + Send + std::fmt::Debug> GroupsAccumulator
63    for PrimitiveDistinctCountGroupsAccumulator<T>
64where
65    T::Native: Eq + Hash,
66{
67    fn update_batch(
68        &mut self,
69        values: &[ArrayRef],
70        group_indices: &[usize],
71        opt_filter: Option<&BooleanArray>,
72        total_num_groups: usize,
73    ) -> datafusion_common::Result<()> {
74        debug_assert_eq!(values.len(), 1);
75        self.counts.resize(total_num_groups, 0);
76        let arr = values[0].as_primitive::<T>();
77        accumulate(group_indices, arr, opt_filter, |group_idx, value| {
78            if self.seen.insert((group_idx, value)) {
79                self.counts[group_idx] += 1;
80            }
81        });
82        Ok(())
83    }
84
85    fn evaluate(&mut self, emit_to: EmitTo) -> datafusion_common::Result<ArrayRef> {
86        let counts = emit_to.take_needed(&mut self.counts);
87
88        match emit_to {
89            EmitTo::All => {
90                self.seen.clear();
91            }
92            EmitTo::First(n) => {
93                let mut remaining = HashSet::default();
94                for (group_idx, value) in self.seen.drain() {
95                    if group_idx >= n {
96                        remaining.insert((group_idx - n, value));
97                    }
98                }
99                self.seen = remaining;
100            }
101        }
102
103        Ok(Arc::new(Int64Array::from(counts)))
104    }
105
106    fn state(&mut self, emit_to: EmitTo) -> datafusion_common::Result<Vec<ArrayRef>> {
107        let num_emitted = match emit_to {
108            EmitTo::All => self.counts.len(),
109            EmitTo::First(n) => n,
110        };
111
112        // Prefix-sum counts[..num_emitted] into offsets
113        let mut offsets = Vec::with_capacity(num_emitted + 1);
114        offsets.push(0i32);
115        let mut total = 0i32;
116        for &c in &self.counts[..num_emitted] {
117            total += c as i32;
118            offsets.push(total);
119        }
120
121        let mut all_values = vec![T::Native::default(); total as usize];
122        let mut cursors: Vec<i32> = offsets[..num_emitted].to_vec();
123
124        if matches!(emit_to, EmitTo::All) {
125            for (group_idx, value) in self.seen.drain() {
126                let pos = cursors[group_idx] as usize;
127                all_values[pos] = value;
128                cursors[group_idx] += 1;
129            }
130            self.counts.clear();
131        } else {
132            let mut remaining = HashSet::default();
133            for (group_idx, value) in self.seen.drain() {
134                if group_idx < num_emitted {
135                    let pos = cursors[group_idx] as usize;
136                    all_values[pos] = value;
137                    cursors[group_idx] += 1;
138                } else {
139                    remaining.insert((group_idx - num_emitted, value));
140                }
141            }
142            self.seen = remaining;
143            let _ = emit_to.take_needed(&mut self.counts);
144        }
145
146        let values_array = Arc::new(PrimitiveArray::<T>::new(
147            ScalarBuffer::from(all_values),
148            None,
149        ));
150        let list_array = ListArray::new(
151            Arc::new(Field::new_list_field(T::DATA_TYPE, true)),
152            OffsetBuffer::new(offsets.into()),
153            values_array,
154            None,
155        );
156
157        Ok(vec![Arc::new(list_array)])
158    }
159
160    fn merge_batch(
161        &mut self,
162        values: &[ArrayRef],
163        group_indices: &[usize],
164        total_num_groups: usize,
165    ) -> datafusion_common::Result<()> {
166        debug_assert_eq!(values.len(), 1);
167        self.counts.resize(total_num_groups, 0);
168        let list_array = values[0].as_list::<i32>();
169        let inner = list_array.values().as_primitive::<T>();
170        let inner_values = inner.values();
171        let offsets = list_array.offsets();
172
173        for (row_idx, &group_idx) in group_indices.iter().enumerate() {
174            let start = offsets[row_idx] as usize;
175            let end = offsets[row_idx + 1] as usize;
176            for &value in &inner_values[start..end] {
177                if self.seen.insert((group_idx, value)) {
178                    self.counts[group_idx] += 1;
179                }
180            }
181        }
182
183        Ok(())
184    }
185
186    fn convert_to_state(
187        &self,
188        values: &[ArrayRef],
189        opt_filter: Option<&BooleanArray>,
190    ) -> datafusion_common::Result<Vec<ArrayRef>> {
191        debug_assert_eq!(values.len(), 1);
192        let arr = values[0].as_primitive::<T>();
193
194        let values_builder = PrimitiveBuilder::<T>::with_capacity(arr.len());
195        let mut builder = ListBuilder::new(values_builder)
196            .with_field(Arc::new(Field::new_list_field(T::DATA_TYPE, true)));
197
198        for row in 0..arr.len() {
199            let included = arr.is_valid(row)
200                && opt_filter
201                    .is_none_or(|filter| filter.is_valid(row) && filter.value(row));
202            if included {
203                builder.values().append_value(arr.value(row));
204            }
205            builder.append(true);
206        }
207
208        Ok(vec![Arc::new(builder.finish())])
209    }
210    fn size(&self) -> usize {
211        size_of::<Self>()
212            + self.seen.capacity() * (size_of::<(usize, T::Native)>() + size_of::<u64>())
213            + self.counts.capacity() * size_of::<i64>()
214    }
215}
216
217#[cfg(test)]
218mod tests {
219    use super::*;
220    use arrow::array::Int32Array;
221    use arrow::datatypes::Int32Type;
222    use datafusion_common::Result;
223
224    #[test]
225    fn convert_to_state_roundtrips_through_merge() -> Result<()> {
226        let values = Arc::new(Int32Array::from(vec![
227            Some(1),
228            Some(2),
229            Some(2),
230            None,
231            Some(3),
232            Some(4),
233            Some(5),
234            Some(5),
235        ])) as ArrayRef;
236        let filter = BooleanArray::from(vec![
237            Some(true),
238            Some(true),
239            Some(true),
240            Some(true),
241            None,
242            Some(true),
243            Some(true),
244            Some(true),
245        ]);
246        let group_indices = vec![0usize, 1, 0, 1, 0, 0, 0, 0];
247
248        let mut direct = PrimitiveDistinctCountGroupsAccumulator::<Int32Type>::new();
249        direct.update_batch(
250            std::slice::from_ref(&values),
251            &group_indices,
252            Some(&filter),
253            2,
254        )?;
255        let direct = direct.evaluate(EmitTo::All)?;
256
257        let converter = PrimitiveDistinctCountGroupsAccumulator::<Int32Type>::new();
258        let state =
259            converter.convert_to_state(std::slice::from_ref(&values), Some(&filter))?;
260        assert_eq!(state[0].null_count(), 0);
261        let mut merged = PrimitiveDistinctCountGroupsAccumulator::<Int32Type>::new();
262        merged.merge_batch(&state, &group_indices, 2)?;
263        let merged = merged.evaluate(EmitTo::All)?;
264
265        assert_eq!(
266            direct.as_any().downcast_ref::<Int64Array>().unwrap(),
267            merged.as_any().downcast_ref::<Int64Array>().unwrap()
268        );
269        Ok(())
270    }
271
272    #[test]
273    fn convert_to_state_preserves_empty_and_filtered_rows() -> Result<()> {
274        let converter = PrimitiveDistinctCountGroupsAccumulator::<Int32Type>::new();
275        let empty_values =
276            Arc::new(Int32Array::from(Vec::<Option<i32>>::new())) as ArrayRef;
277        let state =
278            converter.convert_to_state(std::slice::from_ref(&empty_values), None)?;
279        assert_eq!(state[0].len(), 0);
280        assert_eq!(state[0].null_count(), 0);
281
282        let values = Arc::new(Int32Array::from(vec![Some(1), Some(2), None])) as ArrayRef;
283        let filter = BooleanArray::from(vec![Some(false), None, Some(false)]);
284        let group_indices = vec![0usize, 1, 0];
285
286        let state =
287            converter.convert_to_state(std::slice::from_ref(&values), Some(&filter))?;
288        assert_eq!(state[0].len(), values.len());
289        assert_eq!(state[0].null_count(), 0);
290        let list_state = state[0].as_list::<i32>();
291        for row in 0..list_state.len() {
292            assert_eq!(list_state.value_length(row), 0);
293        }
294
295        let mut merged = PrimitiveDistinctCountGroupsAccumulator::<Int32Type>::new();
296        merged.merge_batch(&state, &group_indices, 2)?;
297        let result = merged.evaluate(EmitTo::All)?;
298        assert_eq!(
299            result.as_any().downcast_ref::<Int64Array>().unwrap(),
300            &Int64Array::from(vec![0, 0])
301        );
302        Ok(())
303    }
304}