datafusion_functions_aggregate_common/aggregate/count_distinct/
groups.rs1use 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 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}