1use std::collections::HashSet;
23use std::fmt::Debug;
24use std::sync::Arc;
25
26use ahash::RandomState;
27use datafusion_common::cast::as_list_array;
28use datafusion_common::error::Result;
29use datafusion_common::hash_utils::create_hashes_with_hasher;
30use datafusion_common::utils::SingleRowListArrayBuilder;
31use datafusion_common::{ScalarValue, internal_err, not_impl_err};
32use datafusion_expr::function::{AccumulatorArgs, StateFieldsArgs};
33use datafusion_expr::utils::{AggregateOrderSensitivity, format_state_name};
34use datafusion_expr::{
35 Accumulator, AggregateUDF, AggregateUDFImpl, EmitTo, GroupsAccumulator, ReversedUDAF,
36 SetMonotonicity, Signature, TypeSignature, Volatility,
37};
38use datafusion_functions_aggregate_common::aggregate::groups_accumulator::nulls::filtered_null_mask;
39use datatypes::arrow;
40use datatypes::arrow::array::{
41 Array, ArrayRef, AsArray, BooleanArray, Int64Array, ListArray, UInt64Array,
42};
43use datatypes::arrow::buffer::{OffsetBuffer, ScalarBuffer};
44use datatypes::arrow::datatypes::{DataType, Field, FieldRef};
45
46use crate::function_registry::FunctionRegistry;
47
48type HashValueType = u64;
49
50const RANDOM_SEED_0: u64 = 0x4047821dc6144e4b;
52const RANDOM_SEED_1: u64 = 0x2abddf23ad417112;
53const RANDOM_SEED_2: u64 = 0x6a52eeecd26eff21;
54const RANDOM_SEED_3: u64 = 0x91cf673b965a7875;
55
56impl CountHash {
57 pub fn register(registry: &FunctionRegistry) {
58 registry.register_aggr(CountHash::udf_impl());
59 }
60
61 pub fn udf_impl() -> AggregateUDF {
62 AggregateUDF::new_from_impl(CountHash {
63 signature: Signature::one_of(
64 vec![TypeSignature::VariadicAny, TypeSignature::Nullary],
65 Volatility::Immutable,
66 ),
67 })
68 }
69}
70
71#[derive(Debug, Clone, Eq, PartialEq, Hash)]
72pub struct CountHash {
73 signature: Signature,
74}
75
76impl AggregateUDFImpl for CountHash {
77 fn name(&self) -> &str {
78 "count_hash"
79 }
80
81 fn signature(&self) -> &Signature {
82 &self.signature
83 }
84
85 fn return_type(&self, _arg_types: &[DataType]) -> Result<DataType> {
86 Ok(DataType::Int64)
87 }
88
89 fn is_nullable(&self) -> bool {
90 false
91 }
92
93 fn state_fields(&self, args: StateFieldsArgs) -> Result<Vec<FieldRef>> {
94 Ok(vec![Arc::new(Field::new_list(
95 format_state_name(args.name, "count_hash"),
96 Field::new_list_field(DataType::UInt64, true),
97 true,
100 ))])
101 }
102
103 fn accumulator(&self, acc_args: AccumulatorArgs) -> Result<Box<dyn Accumulator>> {
104 if acc_args.exprs.len() > 1 {
105 return not_impl_err!("count_hash with multiple arguments");
106 }
107
108 Ok(Box::new(CountHashAccumulator {
109 values: HashSet::default(),
110 random_state: RandomState::with_seeds(
111 RANDOM_SEED_0,
112 RANDOM_SEED_1,
113 RANDOM_SEED_2,
114 RANDOM_SEED_3,
115 ),
116 batch_hashes: vec![],
117 }))
118 }
119
120 fn aliases(&self) -> &[String] {
121 &[]
122 }
123
124 fn groups_accumulator_supported(&self, _args: AccumulatorArgs) -> bool {
125 true
126 }
127
128 fn create_groups_accumulator(
129 &self,
130 args: AccumulatorArgs,
131 ) -> Result<Box<dyn GroupsAccumulator>> {
132 if args.exprs.len() > 1 {
133 return not_impl_err!("count_hash with multiple arguments");
134 }
135
136 Ok(Box::new(CountHashGroupAccumulator::new()))
137 }
138
139 fn reverse_expr(&self) -> ReversedUDAF {
140 ReversedUDAF::Identical
141 }
142
143 fn order_sensitivity(&self) -> AggregateOrderSensitivity {
144 AggregateOrderSensitivity::Insensitive
145 }
146
147 fn default_value(&self, _data_type: &DataType) -> Result<ScalarValue> {
148 Ok(ScalarValue::Int64(Some(0)))
149 }
150
151 fn set_monotonicity(&self, _data_type: &DataType) -> SetMonotonicity {
152 SetMonotonicity::Increasing
153 }
154}
155
156#[derive(Debug)]
158pub struct CountHashGroupAccumulator {
159 distinct_sets: Vec<HashSet<HashValueType, RandomState>>,
161 random_state: RandomState,
162 batch_hashes: Vec<HashValueType>,
163}
164
165impl Default for CountHashGroupAccumulator {
166 fn default() -> Self {
167 Self::new()
168 }
169}
170
171impl CountHashGroupAccumulator {
172 pub fn new() -> Self {
173 Self {
174 distinct_sets: vec![],
175 random_state: RandomState::with_seeds(
176 RANDOM_SEED_0,
177 RANDOM_SEED_1,
178 RANDOM_SEED_2,
179 RANDOM_SEED_3,
180 ),
181 batch_hashes: vec![],
182 }
183 }
184
185 fn ensure_sets(&mut self, total_num_groups: usize) {
186 if self.distinct_sets.len() < total_num_groups {
187 self.distinct_sets
188 .resize_with(total_num_groups, HashSet::default);
189 }
190 }
191}
192
193impl GroupsAccumulator for CountHashGroupAccumulator {
194 fn update_batch(
195 &mut self,
196 values: &[ArrayRef],
197 group_indices: &[usize],
198 opt_filter: Option<&BooleanArray>,
199 total_num_groups: usize,
200 ) -> Result<()> {
201 assert_eq!(values.len(), 1, "count_hash expects a single argument");
202 self.ensure_sets(total_num_groups);
203
204 let array = &values[0];
205 self.batch_hashes.clear();
206 self.batch_hashes.resize(array.len(), 0);
207 let hashes = create_hashes_with_hasher(
208 &[ArrayRef::clone(array)],
209 &self.random_state,
210 &mut self.batch_hashes,
211 )?;
212
213 let nulls = array.logical_nulls();
216
217 match (nulls.as_ref(), opt_filter) {
218 (None, None) => {
219 for (row_idx, &group_idx) in group_indices.iter().enumerate() {
221 self.distinct_sets[group_idx].insert(hashes[row_idx]);
222 }
223 }
224 (Some(nulls), None) => {
225 for (row_idx, (&group_idx, is_valid)) in
227 group_indices.iter().zip(nulls.iter()).enumerate()
228 {
229 if is_valid {
230 self.distinct_sets[group_idx].insert(hashes[row_idx]);
231 }
232 }
233 }
234 (None, Some(filter)) => {
235 for (row_idx, (&group_idx, filter_value)) in
237 group_indices.iter().zip(filter.iter()).enumerate()
238 {
239 if let Some(true) = filter_value {
240 self.distinct_sets[group_idx].insert(hashes[row_idx]);
241 }
242 }
243 }
244 (Some(nulls), Some(filter)) => {
245 let iter = filter
247 .iter()
248 .zip(group_indices.iter())
249 .zip(nulls.iter())
250 .enumerate();
251
252 for (row_idx, ((filter_value, &group_idx), is_valid)) in iter {
253 if is_valid && filter_value == Some(true) {
254 self.distinct_sets[group_idx].insert(hashes[row_idx]);
255 }
256 }
257 }
258 }
259
260 Ok(())
261 }
262
263 fn evaluate(&mut self, emit_to: EmitTo) -> Result<ArrayRef> {
264 let distinct_sets: Vec<HashSet<u64, RandomState>> =
265 emit_to.take_needed(&mut self.distinct_sets);
266
267 let counts = distinct_sets
268 .iter()
269 .map(|set| set.len() as i64)
270 .collect::<Vec<_>>();
271 Ok(Arc::new(Int64Array::from(counts)))
272 }
273
274 fn merge_batch(
275 &mut self,
276 values: &[ArrayRef],
277 group_indices: &[usize],
278 total_num_groups: usize,
279 ) -> Result<()> {
280 assert_eq!(
281 values.len(),
282 1,
283 "count_hash merge expects a single state array"
284 );
285 self.ensure_sets(total_num_groups);
286
287 let list_array = as_list_array(&values[0])?;
288
289 for (i, &group_idx) in group_indices.iter().enumerate() {
291 if i < list_array.len() {
292 let inner_array = list_array.value(i);
293 let inner_array = inner_array.as_any().downcast_ref::<UInt64Array>().unwrap();
294 for j in 0..inner_array.len() {
296 if !inner_array.is_null(j) {
297 self.distinct_sets[group_idx].insert(inner_array.value(j));
298 }
299 }
300 }
301 }
302
303 Ok(())
304 }
305
306 fn state(&mut self, emit_to: EmitTo) -> Result<Vec<ArrayRef>> {
307 let distinct_sets: Vec<HashSet<u64, RandomState>> =
308 emit_to.take_needed(&mut self.distinct_sets);
309
310 let mut offsets = Vec::with_capacity(distinct_sets.len() + 1);
311 offsets.push(0);
312 let mut curr_len = 0i32;
313
314 let mut value_iter = distinct_sets
315 .into_iter()
316 .flat_map(|set| {
317 curr_len += set.len() as i32;
319 offsets.push(curr_len);
320 set.into_iter()
322 })
323 .peekable();
324 let data_array: ArrayRef = if value_iter.peek().is_none() {
325 arrow::array::new_empty_array(&DataType::UInt64) as _
326 } else {
327 Arc::new(UInt64Array::from_iter_values(value_iter))
328 };
329 let offset_buffer = OffsetBuffer::new(ScalarBuffer::from(offsets));
330
331 let list_array = ListArray::new(
332 Arc::new(Field::new_list_field(DataType::UInt64, true)),
333 offset_buffer,
334 data_array,
335 None,
336 );
337
338 Ok(vec![Arc::new(list_array) as _])
339 }
340
341 fn convert_to_state(
342 &self,
343 values: &[ArrayRef],
344 opt_filter: Option<&BooleanArray>,
345 ) -> Result<Vec<ArrayRef>> {
346 assert_eq!(values.len(), 1, "count_hash expects a single argument");
348 let values = ArrayRef::clone(&values[0]);
349
350 let offsets = OffsetBuffer::new(ScalarBuffer::from_iter(0..values.len() as i32 + 1));
351 let nulls = filtered_null_mask(opt_filter, &values);
352 let list_array = ListArray::new(
353 Arc::new(Field::new_list_field(DataType::UInt64, true)),
354 offsets,
355 values,
356 nulls,
357 );
358
359 Ok(vec![Arc::new(list_array)])
360 }
361
362 fn size(&self) -> usize {
363 let mut size = size_of::<Self>();
365
366 size += size_of::<Vec<HashSet<HashValueType, RandomState>>>()
368 + self.distinct_sets.capacity() * size_of::<HashSet<HashValueType, RandomState>>();
369
370 for set in &self.distinct_sets {
373 size += set.capacity() * size_of::<HashValueType>();
375 }
376
377 size
378 }
379}
380
381#[derive(Debug)]
382struct CountHashAccumulator {
383 values: HashSet<HashValueType, RandomState>,
384 random_state: RandomState,
385 batch_hashes: Vec<HashValueType>,
386}
387
388impl CountHashAccumulator {
389 fn fixed_size(&self) -> usize {
392 size_of_val(self) + (size_of::<HashValueType>() * self.values.capacity())
393 }
394}
395
396impl Accumulator for CountHashAccumulator {
397 fn state(&mut self) -> Result<Vec<ScalarValue>> {
399 let values = self.values.iter().cloned().collect::<Vec<_>>();
400 let arr = Arc::new(UInt64Array::from(values)) as _;
401 let list_scalar = SingleRowListArrayBuilder::new(arr).build_list_scalar();
402 Ok(vec![list_scalar])
403 }
404
405 fn update_batch(&mut self, values: &[ArrayRef]) -> Result<()> {
406 if values.is_empty() {
407 return Ok(());
408 }
409
410 let arr = &values[0];
411 if arr.data_type() == &DataType::Null {
412 return Ok(());
413 }
414
415 self.batch_hashes.clear();
416 self.batch_hashes.resize(arr.len(), 0);
417 let hashes = create_hashes_with_hasher(
418 &[ArrayRef::clone(arr)],
419 &self.random_state,
420 &mut self.batch_hashes,
421 )?;
422 for hash in hashes {
423 self.values.insert(*hash);
424 }
425 Ok(())
426 }
427
428 fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> {
434 if states.is_empty() {
435 return Ok(());
436 }
437 assert_eq!(states.len(), 1, "array_agg states must be singleton!");
438 let array = &states[0];
439 let list_array = array.as_list::<i32>();
440 for inner_array in list_array.iter() {
441 let Some(inner_array) = inner_array else {
442 return internal_err!(
443 "Intermediate results of count_hash should always be non null"
444 );
445 };
446 let hash_array = inner_array.as_any().downcast_ref::<UInt64Array>().unwrap();
447 for &hash in hash_array.values().iter().take(hash_array.len()) {
448 self.values.insert(hash);
449 }
450 }
451 Ok(())
452 }
453
454 fn evaluate(&mut self) -> Result<ScalarValue> {
455 Ok(ScalarValue::Int64(Some(self.values.len() as i64)))
456 }
457
458 fn size(&self) -> usize {
459 self.fixed_size()
460 }
461}
462
463#[cfg(test)]
464mod tests {
465 use datatypes::arrow::array::{Array, BooleanArray, Int32Array, Int64Array};
466
467 use super::*;
468
469 fn create_test_accumulator() -> CountHashAccumulator {
470 CountHashAccumulator {
471 values: HashSet::default(),
472 random_state: RandomState::with_seeds(
473 RANDOM_SEED_0,
474 RANDOM_SEED_1,
475 RANDOM_SEED_2,
476 RANDOM_SEED_3,
477 ),
478 batch_hashes: vec![],
479 }
480 }
481
482 #[test]
483 fn test_count_hash_accumulator() -> Result<()> {
484 let mut acc = create_test_accumulator();
485
486 let array = Arc::new(Int32Array::from(vec![
488 Some(1),
489 Some(2),
490 Some(3),
491 Some(1),
492 Some(2),
493 None,
494 ])) as ArrayRef;
495 acc.update_batch(&[array])?;
496 let result = acc.evaluate()?;
497 assert_eq!(result, ScalarValue::Int64(Some(4)));
498
499 let mut acc = create_test_accumulator();
501 let array = Arc::new(Int32Array::from(vec![] as Vec<Option<i32>>)) as ArrayRef;
502 acc.update_batch(&[array])?;
503 let result = acc.evaluate()?;
504 assert_eq!(result, ScalarValue::Int64(Some(0)));
505
506 let mut acc = create_test_accumulator();
508 let array = Arc::new(Int32Array::from(vec![None, None, None])) as ArrayRef;
509 acc.update_batch(&[array])?;
510 let result = acc.evaluate()?;
511 assert_eq!(result, ScalarValue::Int64(Some(1)));
512
513 Ok(())
514 }
515
516 #[test]
517 fn test_count_hash_accumulator_typed_null_state_merge() -> Result<()> {
518 let typed_nulls = Arc::new(Int32Array::from(vec![None, None])) as ArrayRef;
519
520 let mut fresh = create_test_accumulator();
521 fresh.update_batch(&[typed_nulls])?;
522 let fresh_state = fresh.state()?;
523 assert_eq!(fresh.evaluate()?, ScalarValue::Int64(Some(1)));
524
525 let persisted_state = Arc::new(
526 SingleRowListArrayBuilder::new(Arc::new(UInt64Array::from(vec![0])) as ArrayRef)
527 .build_list_array(),
528 ) as ArrayRef;
529 let mut restored = create_test_accumulator();
530 restored.merge_batch(&[persisted_state])?;
531
532 assert_eq!(restored.evaluate()?, fresh.evaluate()?);
533 assert_eq!(restored.state()?, fresh_state);
534
535 Ok(())
536 }
537
538 #[test]
539 fn test_count_hash_accumulator_merge() -> Result<()> {
540 let mut acc1 = create_test_accumulator();
542 let array1 = Arc::new(Int32Array::from(vec![Some(1), Some(2), Some(3)])) as ArrayRef;
543 acc1.update_batch(&[array1])?;
544 let state1 = acc1.state()?;
545
546 let mut acc2 = create_test_accumulator();
548 let array2 = Arc::new(Int32Array::from(vec![Some(3), Some(4), Some(5)])) as ArrayRef;
549 acc2.update_batch(&[array2])?;
550 let state2 = acc2.state()?;
551
552 let mut acc_merged = create_test_accumulator();
554 let state_array1 = state1[0].to_array()?;
555 let state_array2 = state2[0].to_array()?;
556
557 acc_merged.merge_batch(&[state_array1])?;
558 acc_merged.merge_batch(&[state_array2])?;
559
560 let result = acc_merged.evaluate()?;
561 assert_eq!(result, ScalarValue::Int64(Some(5)));
563
564 Ok(())
565 }
566
567 fn create_test_group_accumulator() -> CountHashGroupAccumulator {
568 CountHashGroupAccumulator::new()
569 }
570
571 #[test]
572 fn test_count_hash_group_accumulator() -> Result<()> {
573 let mut acc = create_test_group_accumulator();
574 let values = Arc::new(Int32Array::from(vec![1, 2, 1, 3, 2, 4, 5])) as ArrayRef;
575 let group_indices = vec![0, 1, 0, 0, 1, 2, 0];
576 let total_num_groups = 3;
577
578 acc.update_batch(&[values], &group_indices, None, total_num_groups)?;
579
580 let result_array = acc.evaluate(EmitTo::All)?;
581 let result = result_array.as_any().downcast_ref::<Int64Array>().unwrap();
582
583 assert_eq!(result.value(0), 3);
587 assert_eq!(result.value(1), 1);
588 assert_eq!(result.value(2), 1);
589
590 Ok(())
591 }
592
593 #[test]
594 fn test_count_hash_group_accumulator_with_filter() -> Result<()> {
595 let mut acc = create_test_group_accumulator();
596 let values = Arc::new(Int32Array::from(vec![1, 2, 3, 4, 5, 6])) as ArrayRef;
597 let group_indices = vec![0, 0, 1, 1, 2, 2];
598 let filter = BooleanArray::from(vec![true, false, true, true, false, true]);
599 let total_num_groups = 3;
600
601 acc.update_batch(&[values], &group_indices, Some(&filter), total_num_groups)?;
602
603 let result_array = acc.evaluate(EmitTo::All)?;
604 let result = result_array.as_any().downcast_ref::<Int64Array>().unwrap();
605
606 assert_eq!(result.value(0), 1);
610 assert_eq!(result.value(1), 2);
611 assert_eq!(result.value(2), 1);
612
613 Ok(())
614 }
615
616 #[test]
617 fn test_count_hash_group_accumulator_merge() -> Result<()> {
618 let mut acc1 = create_test_group_accumulator();
620 let values1 = Arc::new(Int32Array::from(vec![1, 2, 3, 4])) as ArrayRef;
621 let group_indices1 = vec![0, 0, 1, 1];
622 acc1.update_batch(&[values1], &group_indices1, None, 2)?;
623 let state1 = acc1.state(EmitTo::All)?;
625
626 let mut acc2 = create_test_group_accumulator();
628 let values2 = Arc::new(Int32Array::from(vec![5, 6, 1, 3])) as ArrayRef;
629 let group_indices2 = vec![2, 2, 0, 1];
631 acc2.update_batch(&[values2], &group_indices2, None, 3)?;
632 let merge_group_indices = vec![0, 2];
638 acc2.merge_batch(&state1, &merge_group_indices, 3)?;
639
640 let result_array = acc2.evaluate(EmitTo::All)?;
641 let result = result_array.as_any().downcast_ref::<Int64Array>().unwrap();
642
643 assert_eq!(result.value(0), 2);
648 assert_eq!(result.value(1), 1);
649 assert_eq!(result.value(2), 4);
650
651 Ok(())
652 }
653
654 #[test]
655 fn test_size() {
656 let acc = create_test_group_accumulator();
657 assert!(acc.size() > 0);
659 }
660}