diff --git a/datafusion/functions-aggregate/src/approx_distinct.rs b/datafusion/functions-aggregate/src/approx_distinct.rs index 74f116cf48725..dc9270ebe41df 100644 --- a/datafusion/functions-aggregate/src/approx_distinct.rs +++ b/datafusion/functions-aggregate/src/approx_distinct.rs @@ -126,6 +126,65 @@ impl Accumulator for ApproxDistinctBitmapWrapper { } } +/// A validated, zero-copy view of a serialized `approx_distinct` partial state, +/// as produced by [`GroupHll::serialize`] or by the per-group [`Accumulator`]s. +enum SerializedHll<'a> { + /// The raw [`NUM_REGISTERS`] registers of a dense sketch. + Dense(&'a [u8; NUM_REGISTERS]), + /// Little-endian hashes of at most [`SPARSE_LIMIT`] distinct values. An + /// empty state decodes as an empty sparse state. + Sparse(&'a [[u8; size_of::()]]), +} + +impl<'a> SerializedHll<'a> { + fn decode(bytes: &'a [u8]) -> Result { + if let Ok(registers) = <&[u8; NUM_REGISTERS]>::try_from(bytes) { + return Ok(Self::Dense(registers)); + } + let (chunks, rest) = bytes.as_chunks::<{ size_of::() }>(); + if !rest.is_empty() { + return internal_err!( + "approx_distinct: malformed sparse state: length {} is not a multiple of {}", + bytes.len(), + size_of::() + ); + } + if chunks.len() > SPARSE_LIMIT { + return internal_err!( + "approx_distinct: malformed sparse state: length {} exceeds sparse limit {}", + bytes.len(), + SPARSE_LIMIT * size_of::() + ); + } + Ok(Self::Sparse(chunks)) + } +} + +/// Merge the serialized partial states in `states` into `hll`. +fn merge_states( + hll: &mut HyperLogLog, + states: &[ArrayRef], +) -> Result<()> { + assert_eq!(1, states.len(), "expect only 1 element in the states"); + let binary_array = downcast_value!(states[0], BinaryArray); + for v in binary_array.iter() { + let v = v.ok_or_else(|| { + internal_datafusion_err!("Impossibly got empty binary array from states") + })?; + match SerializedHll::decode(v)? { + SerializedHll::Dense(registers) => { + hll.merge(&HyperLogLog::new_with_registers(*registers)); + } + SerializedHll::Sparse(chunks) => { + for chunk in chunks { + hll.add_hashed(u64::from_le_bytes(*chunk)); + } + } + } + } + Ok(()) +} + #[derive(Debug)] struct HLLAccumulator { hll: HyperLogLog, @@ -166,16 +225,7 @@ impl Accumulator for HLLAccumulator { } fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> { - assert_eq!(1, states.len(), "expect only 1 element in the states"); - let binary_array = downcast_value!(states[0], BinaryArray); - for v in binary_array.iter() { - let v = v.ok_or_else(|| { - internal_datafusion_err!("Impossibly got empty binary array from states") - })?; - let other = v.try_into()?; - self.hll.merge(&other); - } - Ok(()) + merge_states(&mut self.hll, states) } fn state(&mut self) -> Result> { @@ -226,16 +276,7 @@ where } fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> { - assert_eq!(1, states.len(), "expect only 1 element in the states"); - let binary_array = downcast_value!(states[0], BinaryArray); - for v in binary_array.iter() { - let v = v.ok_or_else(|| { - internal_datafusion_err!("Impossibly got empty binary array from states") - })?; - let other = v.try_into()?; - self.hll.merge(&other); - } - Ok(()) + merge_states(&mut self.hll, states) } fn state(&mut self) -> Result> { @@ -335,36 +376,17 @@ impl GroupHll { } } - /// Merge a serialized state (produced by [`Self::serialize`] or by the - /// per-group [`Accumulator`]) into this sketch. + /// Merge a serialized state (see [`SerializedHll`]) into this sketch. fn merge_serialized(&mut self, bytes: &[u8]) -> Result { - if bytes.is_empty() { - return Ok(0); - } - if bytes.len() == NUM_REGISTERS { - let other: HyperLogLog = bytes.try_into()?; - Ok(self.merge_dense(&other)) - } else { - if !bytes.len().is_multiple_of(size_of::()) { - return internal_err!( - "approx_distinct: malformed sparse state: length {} is not a multiple of {}", - bytes.len(), - size_of::() - ); + Ok(match SerializedHll::decode(bytes)? { + SerializedHll::Dense(registers) => { + self.merge_dense(&HyperLogLog::new_with_registers(*registers)) } - if bytes.len() > SPARSE_LIMIT * size_of::() { - return internal_err!( - "approx_distinct: malformed sparse state: length {} exceeds sparse limit {}", - bytes.len(), - SPARSE_LIMIT * size_of::() - ); - } - let mut delta = 0; - for chunk in bytes.as_chunks::<{ size_of::() }>().0 { - delta += self.add_hash(u64::from_le_bytes(*chunk)); - } - Ok(delta) - } + SerializedHll::Sparse(chunks) => chunks + .iter() + .map(|chunk| self.add_hash(u64::from_le_bytes(*chunk))) + .sum(), + }) } /// Merge a dense sketch into this one, promoting to dense if necessary.