Skip to main content

common_function/aggrs/approximate/
uddsketch.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//! Implementation of the `uddsketch_state` UDAF that generate the state of
16//! UDDSketch for a given set of values.
17//!
18//! The generated state can be used to compute approximate quantiles using
19//! `uddsketch_calc` UDF.
20
21use std::sync::Arc;
22
23use common_query::prelude::*;
24use datafusion::common::cast::{as_binary_array, as_primitive_array};
25use datafusion::common::not_impl_err;
26use datafusion::error::{DataFusionError, Result as DfResult};
27use datafusion::logical_expr::function::AccumulatorArgs;
28use datafusion::logical_expr::{Accumulator as DfAccumulator, AggregateUDF, Volatility};
29use datafusion::physical_plan::expressions::Literal;
30use datafusion::prelude::create_udaf;
31use datatypes::arrow::array::{Array, ArrayRef};
32use datatypes::arrow::datatypes::{DataType, Float64Type};
33use uddsketch::{BatchWorkspace, UddSketch};
34
35use crate::uddsketch_compat;
36
37pub const UDDSKETCH_STATE_NAME: &str = "uddsketch_state";
38
39pub const UDDSKETCH_MERGE_NAME: &str = "uddsketch_merge";
40
41const MAX_BUCKETS: u32 = 1_000_000;
42
43#[derive(Debug)]
44pub struct UddSketchState {
45    uddsketch: UddSketch,
46    workspace: BatchWorkspace,
47    values: Vec<f64>,
48}
49
50impl UddSketchState {
51    pub fn new(bucket_size: u32, error_rate: f64) -> DfResult<Self> {
52        if bucket_size > MAX_BUCKETS {
53            return Err(DataFusionError::Plan(format!(
54                "UDDSketch bucket size exceeds the maximum of {}",
55                MAX_BUCKETS
56            )));
57        }
58        let uddsketch = UddSketch::new(bucket_size, error_rate)
59            .map_err(|e| DataFusionError::Plan(e.to_string()))?;
60        Ok(Self {
61            uddsketch,
62            workspace: BatchWorkspace::default(),
63            values: Vec::new(),
64        })
65    }
66
67    pub fn state_udf_impl() -> AggregateUDF {
68        create_udaf(
69            UDDSKETCH_STATE_NAME,
70            vec![DataType::Int64, DataType::Float64, DataType::Float64],
71            Arc::new(DataType::Binary),
72            Volatility::Immutable,
73            Arc::new(|args| {
74                let (bucket_size, error_rate) = downcast_accumulator_args(args)?;
75                Ok(Box::new(UddSketchState::new(bucket_size, error_rate)?))
76            }),
77            Arc::new(vec![DataType::Binary]),
78        )
79    }
80
81    /// Create a UDF for the `uddsketch_merge` function.
82    ///
83    /// `uddsketch_merge` accepts bucket size, error rate, and a binary column of states generated by `uddsketch_state`
84    /// and merges them into a single state.
85    ///
86    /// The bucket size and error rate must be the same as the original state.
87    pub fn merge_udf_impl() -> AggregateUDF {
88        create_udaf(
89            UDDSKETCH_MERGE_NAME,
90            vec![DataType::Int64, DataType::Float64, DataType::Binary],
91            Arc::new(DataType::Binary),
92            Volatility::Immutable,
93            Arc::new(|args| {
94                let (bucket_size, error_rate) = downcast_accumulator_args(args)?;
95                Ok(Box::new(UddSketchState::new(bucket_size, error_rate)?))
96            }),
97            Arc::new(vec![DataType::Binary]),
98        )
99    }
100
101    fn merge(&mut self, raw: &[u8]) -> DfResult<()> {
102        let uddsketch = uddsketch_compat::decode(raw).map_err(|e| {
103            common_telemetry::trace!("Failed to deserialize UDDSketch: {}", e);
104            DataFusionError::Plan("Failed to deserialize UDDSketch from binary".to_string())
105        })?;
106        if uddsketch.count() == 0 {
107            return Ok(());
108        }
109        if self.uddsketch.max_buckets() != uddsketch.max_buckets()
110            || self.uddsketch.initial_error().to_bits() != uddsketch.initial_error().to_bits()
111        {
112            return Err(DataFusionError::Plan(format!(
113                "Merging UDDSketch with different parameters: arguments={:?} vs actual input={:?}",
114                (self.uddsketch.max_buckets(), self.uddsketch.initial_error()),
115                (uddsketch.max_buckets(), uddsketch.initial_error())
116            )));
117        }
118        self.uddsketch
119            .merge(&uddsketch)
120            .map_err(|e| DataFusionError::Plan(e.to_string()))
121    }
122}
123
124fn downcast_accumulator_args(args: AccumulatorArgs) -> DfResult<(u32, f64)> {
125    let bucket_size = match args.exprs[0]
126        .downcast_ref::<Literal>()
127        .map(|lit| lit.value())
128    {
129        Some(ScalarValue::Int64(Some(value))) => u32::try_from(*value).map_err(|_| {
130            DataFusionError::Plan(format!("Invalid UDDSketch bucket size: {}", value))
131        })?,
132        _ => {
133            return not_impl_err!(
134                "{} not supported for bucket size: {}",
135                UDDSKETCH_STATE_NAME,
136                &args.exprs[0]
137            );
138        }
139    };
140
141    let error_rate = match args.exprs[1]
142        .downcast_ref::<Literal>()
143        .map(|lit| lit.value())
144    {
145        Some(ScalarValue::Float64(Some(value))) => *value,
146        _ => {
147            return not_impl_err!(
148                "{} not supported for error rate: {}",
149                UDDSKETCH_STATE_NAME,
150                &args.exprs[1]
151            );
152        }
153    };
154
155    Ok((bucket_size, error_rate))
156}
157
158impl DfAccumulator for UddSketchState {
159    fn update_batch(&mut self, values: &[ArrayRef]) -> DfResult<()> {
160        let array = &values[2]; // the third column is data value
161        match array.data_type() {
162            DataType::Float64 => {
163                let f64_array = as_primitive_array::<Float64Type>(array)?;
164                let values: &[f64] = if f64_array.null_count() == 0 {
165                    f64_array.values().as_ref()
166                } else {
167                    self.values.clear();
168                    self.values.extend(f64_array.iter().flatten());
169                    self.values.as_slice()
170                };
171                self.uddsketch
172                    .add_batch_with_workspace(values, &mut self.workspace)
173                    .map_err(|e| DataFusionError::Execution(e.to_string()))?;
174            }
175            // meaning instantiate as `uddsketch_merge`
176            DataType::Binary => self.merge_batch(std::slice::from_ref(array))?,
177            _ => {
178                return not_impl_err!(
179                    "UDDSketch functions do not support data type: {}",
180                    array.data_type()
181                );
182            }
183        }
184
185        Ok(())
186    }
187
188    fn evaluate(&mut self) -> DfResult<ScalarValue> {
189        Ok(ScalarValue::Binary(Some(self.uddsketch.encode().map_err(
190            |e| DataFusionError::Internal(format!("Failed to serialize UDDSketch: {}", e)),
191        )?)))
192    }
193
194    fn size(&self) -> usize {
195        std::mem::size_of::<Self>() - std::mem::size_of::<UddSketch>()
196            + self.uddsketch.allocated_size()
197            + self.workspace.allocated_size()
198            + self.values.capacity() * std::mem::size_of::<f64>()
199    }
200
201    fn state(&mut self) -> DfResult<Vec<ScalarValue>> {
202        Ok(vec![ScalarValue::Binary(Some(
203            self.uddsketch.encode().map_err(|e| {
204                DataFusionError::Internal(format!("Failed to serialize UDDSketch: {}", e))
205            })?,
206        ))])
207    }
208
209    fn merge_batch(&mut self, states: &[ArrayRef]) -> DfResult<()> {
210        let array = &states[0];
211        let binary_array = as_binary_array(array)?;
212        for v in binary_array.iter().flatten() {
213            self.merge(v)?;
214        }
215
216        Ok(())
217    }
218}
219
220#[cfg(test)]
221mod tests {
222    use datafusion::arrow::array::{BinaryArray, Float64Array};
223    use uddsketch::UddSketchRef;
224
225    use super::*;
226
227    #[test]
228    fn test_uddsketch_state_basic() {
229        let mut state = UddSketchState::new(10, 0.01).unwrap();
230        state.uddsketch.add(1.0).unwrap();
231        state.uddsketch.add(2.0).unwrap();
232        state.uddsketch.add(3.0).unwrap();
233
234        let result = state.evaluate().unwrap();
235        if let ScalarValue::Binary(Some(bytes)) = result {
236            let encoded = UddSketchRef::parse(&bytes).unwrap();
237            assert_eq!(encoded.count(), 3);
238        } else {
239            panic!("Expected binary scalar value");
240        }
241    }
242
243    #[test]
244    fn test_uddsketch_state_roundtrip() {
245        let mut state = UddSketchState::new(10, 0.01).unwrap();
246        state.uddsketch.add(1.0).unwrap();
247        state.uddsketch.add(2.0).unwrap();
248
249        // Serialize
250        let serialized = state.evaluate().unwrap();
251
252        // Create new state and merge the serialized data
253        let mut new_state = UddSketchState::new(10, 0.01).unwrap();
254        if let ScalarValue::Binary(Some(bytes)) = &serialized {
255            new_state.merge(bytes).unwrap();
256
257            let original_sketch = UddSketchRef::parse(bytes).unwrap();
258            let new_result = new_state.evaluate().unwrap();
259            if let ScalarValue::Binary(Some(new_bytes)) = new_result {
260                let new_sketch = UddSketchRef::parse(&new_bytes).unwrap();
261                assert_eq!(original_sketch.count(), new_sketch.count());
262                assert_eq!(original_sketch.sum(), new_sketch.sum());
263                assert_eq!(
264                    original_sketch.max_error().unwrap(),
265                    new_sketch.max_error().unwrap()
266                );
267                // Compare a few quantiles to ensure statistical equivalence
268                for q in [0.1, 0.5, 0.9].iter() {
269                    let original = original_sketch.quantile(*q).unwrap().unwrap();
270                    let merged = new_sketch.quantile(*q).unwrap().unwrap();
271                    assert!(
272                        (original - merged).abs() < 1e-10,
273                        "Quantile {} mismatch: original={}, new={}",
274                        q,
275                        original,
276                        merged
277                    );
278                }
279            } else {
280                panic!("Expected binary scalar value");
281            }
282        } else {
283            panic!("Expected binary scalar value");
284        }
285    }
286
287    #[test]
288    fn test_uddsketch_state_merges_legacy_state() {
289        let mut state = UddSketchState::new(128, 0.01).unwrap();
290
291        state.merge(uddsketch_compat::LEGACY_STATE).unwrap();
292
293        let ScalarValue::Binary(Some(encoded)) = state.evaluate().unwrap() else {
294            panic!("Expected binary scalar value");
295        };
296        let sketch = UddSketchRef::parse(&encoded).unwrap();
297        assert_eq!(sketch.count(), 4);
298        assert_eq!(sketch.sum(), 1.0);
299        assert_eq!(sketch.quantile(0.5).unwrap(), Some(0.9900000000000001));
300    }
301
302    #[test]
303    fn test_uddsketch_state_merges_compacted_legacy_state() {
304        let mut legacy_state = uddsketch_compat::COMPACTED_LEGACY_SKETCH.to_vec();
305        legacy_state.extend_from_slice(&0.01_f64.to_le_bytes());
306        let mut state = UddSketchState::new(7, 0.01).unwrap();
307
308        state.merge(&legacy_state).unwrap();
309
310        let ScalarValue::Binary(Some(encoded)) = state.evaluate().unwrap() else {
311            panic!("Expected binary scalar value");
312        };
313        let sketch = UddSketchRef::parse(&encoded).unwrap();
314        assert_eq!(sketch.count(), 201);
315        assert_eq!(sketch.times_compacted(), 12);
316    }
317
318    #[test]
319    fn test_uddsketch_state_batch_update() {
320        let mut state = UddSketchState::new(10, 0.01).unwrap();
321        let values = vec![Some(1.0f64), None, Some(2.0), Some(3.0)];
322        let array = Arc::new(Float64Array::from(values)) as ArrayRef;
323
324        state
325            .update_batch(&[array.clone(), array.clone(), array])
326            .unwrap();
327
328        let result = state.evaluate().unwrap();
329        if let ScalarValue::Binary(Some(bytes)) = result {
330            let encoded = UddSketchRef::parse(&bytes).unwrap();
331            assert_eq!(encoded.count(), 3);
332        } else {
333            panic!("Expected binary scalar value");
334        }
335    }
336
337    #[test]
338    fn test_uddsketch_state_non_null_batch_avoids_values_buffer() {
339        let mut state = UddSketchState::new(10, 0.01).unwrap();
340        let array = Arc::new(Float64Array::from(vec![1.0, 2.0, 3.0])) as ArrayRef;
341
342        state
343            .update_batch(&[array.clone(), array.clone(), array])
344            .unwrap();
345
346        assert_eq!(state.uddsketch.count(), 3);
347        assert_eq!(state.values.capacity(), 0);
348    }
349
350    #[test]
351    fn test_uddsketch_state_non_null_sliced_batch() {
352        let mut state = UddSketchState::new(10, 0.01).unwrap();
353        let array = Float64Array::from(vec![1.0, 2.0, 3.0, 4.0]);
354        let array = array.slice(1, 2);
355        let array = Arc::new(array) as ArrayRef;
356
357        state
358            .update_batch(&[array.clone(), array.clone(), array])
359            .unwrap();
360
361        assert_eq!(state.uddsketch.count(), 2);
362        assert_eq!(state.uddsketch.sum(), 5.0);
363    }
364
365    #[test]
366    fn test_uddsketch_state_merge_batch() {
367        let mut state1 = UddSketchState::new(10, 0.01).unwrap();
368        state1.uddsketch.add(1.0).unwrap();
369        let state1_binary = state1.evaluate().unwrap();
370
371        let mut state2 = UddSketchState::new(10, 0.01).unwrap();
372        state2.uddsketch.add(2.0).unwrap();
373        let state2_binary = state2.evaluate().unwrap();
374
375        let mut merged_state = UddSketchState::new(10, 0.01).unwrap();
376        if let (ScalarValue::Binary(Some(bytes1)), ScalarValue::Binary(Some(bytes2))) =
377            (&state1_binary, &state2_binary)
378        {
379            let binary_array = Arc::new(BinaryArray::from(vec![
380                bytes1.as_slice(),
381                bytes2.as_slice(),
382            ])) as ArrayRef;
383            merged_state.merge_batch(&[binary_array]).unwrap();
384
385            let result = merged_state.evaluate().unwrap();
386            if let ScalarValue::Binary(Some(bytes)) = result {
387                let encoded = UddSketchRef::parse(&bytes).unwrap();
388                assert_eq!(encoded.count(), 2);
389            } else {
390                panic!("Expected binary scalar value");
391            }
392        } else {
393            panic!("Expected binary scalar values");
394        }
395    }
396
397    #[test]
398    fn test_uddsketch_state_size() {
399        let mut state = UddSketchState::new(10, 0.01).unwrap();
400        let initial_size = state.size();
401
402        // Add some values to create buckets
403        let array = Arc::new(Float64Array::from_iter_values((0..64).map(f64::from))) as ArrayRef;
404        state
405            .update_batch(&[array.clone(), array.clone(), array])
406            .unwrap();
407
408        let size_with_values = state.size();
409        assert!(
410            size_with_values > initial_size,
411            "Size should increase after adding values: initial={}, with_values={}",
412            initial_size,
413            size_with_values
414        );
415    }
416
417    #[test]
418    fn test_uddsketch_state_rejects_invalid_config() {
419        assert!(UddSketchState::new(6, 0.01).is_err());
420        assert!(UddSketchState::new(10, 1.0).is_err());
421
422        let mut maximum = UddSketchState::new(1_000_000, 0.01).unwrap();
423        let ScalarValue::Binary(Some(encoded)) = maximum.evaluate().unwrap() else {
424            panic!("Expected binary scalar value");
425        };
426        UddSketchRef::parse(&encoded).unwrap();
427        assert!(UddSketchState::new(1_000_001, 0.01).is_err());
428    }
429
430    #[test]
431    fn test_uddsketch_state_rejects_nan_batch() {
432        let mut state = UddSketchState::new(10, 0.01).unwrap();
433        let array = Arc::new(Float64Array::from(vec![1.0, f64::NAN])) as ArrayRef;
434
435        let error = state
436            .update_batch(&[array.clone(), array.clone(), array])
437            .unwrap_err();
438        assert!(error.to_string().contains("NaN values are not supported"));
439    }
440}