Skip to main content

common_function/aggrs/
count_hash.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15//! `CountHash` / `count_hash` is a hash-based approximate distinct count function.
16//!
17//! It is a variant of `CountDistinct` that uses a hash function to approximate the
18//! distinct count.
19//! It is designed to be more efficient than `CountDistinct` for large datasets,
20//! but it is not as accurate, as the hash value may be collision.
21
22use 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
50// read from /dev/urandom 4047821dc6144e4b2abddf23ad4171126a52eeecd26eff2191cf673b965a7875
51const 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            // For count_hash accumulator, null list item stands for an
98            // empty value set (i.e., all NULL value so far for that group).
99            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/// GroupsAccumulator for `count_hash` aggregate function
157#[derive(Debug)]
158pub struct CountHashGroupAccumulator {
159    /// One HashSet per group to track distinct values
160    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        // Use a pattern similar to accumulate_indices to process rows
214        // that are not null and pass the filter
215        let nulls = array.logical_nulls();
216
217        match (nulls.as_ref(), opt_filter) {
218            (None, None) => {
219                // No nulls, no filter - process all rows
220                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                // Has nulls, no filter
226                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                // No nulls, has filter
236                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                // Has nulls and filter
246                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 each group in the incoming batch
290        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                // Add each value to our set for this group
295                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                // build offset
318                curr_len += set.len() as i32;
319                offsets.push(curr_len);
320                // convert into iter
321                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        // For a single hash value per row, create a list array with that value
347        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        // Base size of the struct
364        let mut size = size_of::<Self>();
365
366        // Size of the vector holding the HashSets
367        size += size_of::<Vec<HashSet<HashValueType, RandomState>>>()
368            + self.distinct_sets.capacity() * size_of::<HashSet<HashValueType, RandomState>>();
369
370        // Estimate HashSet contents size more efficiently
371        // Instead of iterating through all values which is expensive, use an approximation
372        for set in &self.distinct_sets {
373            // Base size of the HashSet
374            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    // calculating the size for fixed length values, taking first batch size *
390    // number of batches.
391    fn fixed_size(&self) -> usize {
392        size_of_val(self) + (size_of::<HashValueType>() * self.values.capacity())
393    }
394}
395
396impl Accumulator for CountHashAccumulator {
397    /// Returns the distinct values seen so far as (one element) ListArray.
398    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    /// Merges multiple sets of distinct values into the current set.
429    ///
430    /// The input to this function is a `ListArray` with **multiple** rows,
431    /// where each row contains the values from a partial aggregate's phase (e.g.
432    /// the result of calling `Self::state` on multiple accumulators).
433    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        // Test with some data
487        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        // Test with empty data
500        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        // Test with only nulls
507        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        // Accumulator 1
541        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        // Accumulator 2
547        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        // Merge state1 and state2 into a new accumulator
553        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        // Distinct values are {1, 2, 3, 4, 5}, so count is 5
562        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        // Group 0: {1, 3, 5} -> 3
584        // Group 1: {2} -> 1
585        // Group 2: {4} -> 1
586        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        // Group 0: {1} (2 is filtered out) -> 1
607        // Group 1: {3, 4} -> 2
608        // Group 2: {6} (5 is filtered out) -> 1
609        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        // Accumulator 1
619        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        // acc1 state: group 0 -> {1, 2}, group 1 -> {3, 4}
624        let state1 = acc1.state(EmitTo::All)?;
625
626        // Accumulator 2
627        let mut acc2 = create_test_group_accumulator();
628        let values2 = Arc::new(Int32Array::from(vec![5, 6, 1, 3])) as ArrayRef;
629        // Merge into different group indices
630        let group_indices2 = vec![2, 2, 0, 1];
631        acc2.update_batch(&[values2], &group_indices2, None, 3)?;
632        // acc2 state: group 0 -> {1}, group 1 -> {3}, group 2 -> {5, 6}
633
634        // Merge state from acc1 into acc2
635        // We will merge acc1's group 0 into acc2's group 0
636        // and acc1's group 1 into acc2's group 2
637        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        // Final state of acc2:
644        // Group 0: {1} U {1, 2} -> {1, 2}, count = 2
645        // Group 1: {3}, count = 1
646        // Group 2: {5, 6} U {3, 4} -> {3, 4, 5, 6}, count = 4
647        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        // Just test it doesn't crash and returns a value.
658        assert!(acc.size() > 0);
659    }
660}