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        .as_any()
127        .downcast_ref::<Literal>()
128        .map(|lit| lit.value())
129    {
130        Some(ScalarValue::Int64(Some(value))) => u32::try_from(*value).map_err(|_| {
131            DataFusionError::Plan(format!("Invalid UDDSketch bucket size: {}", value))
132        })?,
133        _ => {
134            return not_impl_err!(
135                "{} not supported for bucket size: {}",
136                UDDSKETCH_STATE_NAME,
137                &args.exprs[0]
138            );
139        }
140    };
141
142    let error_rate = match args.exprs[1]
143        .as_any()
144        .downcast_ref::<Literal>()
145        .map(|lit| lit.value())
146    {
147        Some(ScalarValue::Float64(Some(value))) => *value,
148        _ => {
149            return not_impl_err!(
150                "{} not supported for error rate: {}",
151                UDDSKETCH_STATE_NAME,
152                &args.exprs[1]
153            );
154        }
155    };
156
157    Ok((bucket_size, error_rate))
158}
159
160impl DfAccumulator for UddSketchState {
161    fn update_batch(&mut self, values: &[ArrayRef]) -> DfResult<()> {
162        let array = &values[2]; // the third column is data value
163        match array.data_type() {
164            DataType::Float64 => {
165                let f64_array = as_primitive_array::<Float64Type>(array)?;
166                let values: &[f64] = if f64_array.null_count() == 0 {
167                    f64_array.values().as_ref()
168                } else {
169                    self.values.clear();
170                    self.values.extend(f64_array.iter().flatten());
171                    self.values.as_slice()
172                };
173                self.uddsketch
174                    .add_batch_with_workspace(values, &mut self.workspace)
175                    .map_err(|e| DataFusionError::Execution(e.to_string()))?;
176            }
177            // meaning instantiate as `uddsketch_merge`
178            DataType::Binary => self.merge_batch(std::slice::from_ref(array))?,
179            _ => {
180                return not_impl_err!(
181                    "UDDSketch functions do not support data type: {}",
182                    array.data_type()
183                );
184            }
185        }
186
187        Ok(())
188    }
189
190    fn evaluate(&mut self) -> DfResult<ScalarValue> {
191        Ok(ScalarValue::Binary(Some(self.uddsketch.encode().map_err(
192            |e| DataFusionError::Internal(format!("Failed to serialize UDDSketch: {}", e)),
193        )?)))
194    }
195
196    fn size(&self) -> usize {
197        std::mem::size_of::<Self>() - std::mem::size_of::<UddSketch>()
198            + self.uddsketch.allocated_size()
199            + self.workspace.allocated_size()
200            + self.values.capacity() * std::mem::size_of::<f64>()
201    }
202
203    fn state(&mut self) -> DfResult<Vec<ScalarValue>> {
204        Ok(vec![ScalarValue::Binary(Some(
205            self.uddsketch.encode().map_err(|e| {
206                DataFusionError::Internal(format!("Failed to serialize UDDSketch: {}", e))
207            })?,
208        ))])
209    }
210
211    fn merge_batch(&mut self, states: &[ArrayRef]) -> DfResult<()> {
212        let array = &states[0];
213        let binary_array = as_binary_array(array)?;
214        for v in binary_array.iter().flatten() {
215            self.merge(v)?;
216        }
217
218        Ok(())
219    }
220}
221
222#[cfg(test)]
223mod tests {
224    use datafusion::arrow::array::{BinaryArray, Float64Array};
225    use uddsketch::UddSketchRef;
226
227    use super::*;
228
229    #[test]
230    fn test_uddsketch_state_basic() {
231        let mut state = UddSketchState::new(10, 0.01).unwrap();
232        state.uddsketch.add(1.0).unwrap();
233        state.uddsketch.add(2.0).unwrap();
234        state.uddsketch.add(3.0).unwrap();
235
236        let result = state.evaluate().unwrap();
237        if let ScalarValue::Binary(Some(bytes)) = result {
238            let encoded = UddSketchRef::parse(&bytes).unwrap();
239            assert_eq!(encoded.count(), 3);
240        } else {
241            panic!("Expected binary scalar value");
242        }
243    }
244
245    #[test]
246    fn test_uddsketch_state_roundtrip() {
247        let mut state = UddSketchState::new(10, 0.01).unwrap();
248        state.uddsketch.add(1.0).unwrap();
249        state.uddsketch.add(2.0).unwrap();
250
251        // Serialize
252        let serialized = state.evaluate().unwrap();
253
254        // Create new state and merge the serialized data
255        let mut new_state = UddSketchState::new(10, 0.01).unwrap();
256        if let ScalarValue::Binary(Some(bytes)) = &serialized {
257            new_state.merge(bytes).unwrap();
258
259            let original_sketch = UddSketchRef::parse(bytes).unwrap();
260            let new_result = new_state.evaluate().unwrap();
261            if let ScalarValue::Binary(Some(new_bytes)) = new_result {
262                let new_sketch = UddSketchRef::parse(&new_bytes).unwrap();
263                assert_eq!(original_sketch.count(), new_sketch.count());
264                assert_eq!(original_sketch.sum(), new_sketch.sum());
265                assert_eq!(
266                    original_sketch.max_error().unwrap(),
267                    new_sketch.max_error().unwrap()
268                );
269                // Compare a few quantiles to ensure statistical equivalence
270                for q in [0.1, 0.5, 0.9].iter() {
271                    let original = original_sketch.quantile(*q).unwrap().unwrap();
272                    let merged = new_sketch.quantile(*q).unwrap().unwrap();
273                    assert!(
274                        (original - merged).abs() < 1e-10,
275                        "Quantile {} mismatch: original={}, new={}",
276                        q,
277                        original,
278                        merged
279                    );
280                }
281            } else {
282                panic!("Expected binary scalar value");
283            }
284        } else {
285            panic!("Expected binary scalar value");
286        }
287    }
288
289    #[test]
290    fn test_uddsketch_state_merges_legacy_state() {
291        let mut state = UddSketchState::new(128, 0.01).unwrap();
292
293        state.merge(uddsketch_compat::LEGACY_STATE).unwrap();
294
295        let ScalarValue::Binary(Some(encoded)) = state.evaluate().unwrap() else {
296            panic!("Expected binary scalar value");
297        };
298        let sketch = UddSketchRef::parse(&encoded).unwrap();
299        assert_eq!(sketch.count(), 4);
300        assert_eq!(sketch.sum(), 1.0);
301        assert_eq!(sketch.quantile(0.5).unwrap(), Some(0.9900000000000001));
302    }
303
304    #[test]
305    fn test_uddsketch_state_merges_compacted_legacy_state() {
306        let mut legacy_state = uddsketch_compat::COMPACTED_LEGACY_SKETCH.to_vec();
307        legacy_state.extend_from_slice(&0.01_f64.to_le_bytes());
308        let mut state = UddSketchState::new(7, 0.01).unwrap();
309
310        state.merge(&legacy_state).unwrap();
311
312        let ScalarValue::Binary(Some(encoded)) = state.evaluate().unwrap() else {
313            panic!("Expected binary scalar value");
314        };
315        let sketch = UddSketchRef::parse(&encoded).unwrap();
316        assert_eq!(sketch.count(), 201);
317        assert_eq!(sketch.times_compacted(), 12);
318    }
319
320    #[test]
321    fn test_uddsketch_state_batch_update() {
322        let mut state = UddSketchState::new(10, 0.01).unwrap();
323        let values = vec![Some(1.0f64), None, Some(2.0), Some(3.0)];
324        let array = Arc::new(Float64Array::from(values)) as ArrayRef;
325
326        state
327            .update_batch(&[array.clone(), array.clone(), array])
328            .unwrap();
329
330        let result = state.evaluate().unwrap();
331        if let ScalarValue::Binary(Some(bytes)) = result {
332            let encoded = UddSketchRef::parse(&bytes).unwrap();
333            assert_eq!(encoded.count(), 3);
334        } else {
335            panic!("Expected binary scalar value");
336        }
337    }
338
339    #[test]
340    fn test_uddsketch_state_non_null_batch_avoids_values_buffer() {
341        let mut state = UddSketchState::new(10, 0.01).unwrap();
342        let array = Arc::new(Float64Array::from(vec![1.0, 2.0, 3.0])) as ArrayRef;
343
344        state
345            .update_batch(&[array.clone(), array.clone(), array])
346            .unwrap();
347
348        assert_eq!(state.uddsketch.count(), 3);
349        assert_eq!(state.values.capacity(), 0);
350    }
351
352    #[test]
353    fn test_uddsketch_state_non_null_sliced_batch() {
354        let mut state = UddSketchState::new(10, 0.01).unwrap();
355        let array = Float64Array::from(vec![1.0, 2.0, 3.0, 4.0]);
356        let array = array.slice(1, 2);
357        let array = Arc::new(array) as ArrayRef;
358
359        state
360            .update_batch(&[array.clone(), array.clone(), array])
361            .unwrap();
362
363        assert_eq!(state.uddsketch.count(), 2);
364        assert_eq!(state.uddsketch.sum(), 5.0);
365    }
366
367    #[test]
368    fn test_uddsketch_state_merge_batch() {
369        let mut state1 = UddSketchState::new(10, 0.01).unwrap();
370        state1.uddsketch.add(1.0).unwrap();
371        let state1_binary = state1.evaluate().unwrap();
372
373        let mut state2 = UddSketchState::new(10, 0.01).unwrap();
374        state2.uddsketch.add(2.0).unwrap();
375        let state2_binary = state2.evaluate().unwrap();
376
377        let mut merged_state = UddSketchState::new(10, 0.01).unwrap();
378        if let (ScalarValue::Binary(Some(bytes1)), ScalarValue::Binary(Some(bytes2))) =
379            (&state1_binary, &state2_binary)
380        {
381            let binary_array = Arc::new(BinaryArray::from(vec![
382                bytes1.as_slice(),
383                bytes2.as_slice(),
384            ])) as ArrayRef;
385            merged_state.merge_batch(&[binary_array]).unwrap();
386
387            let result = merged_state.evaluate().unwrap();
388            if let ScalarValue::Binary(Some(bytes)) = result {
389                let encoded = UddSketchRef::parse(&bytes).unwrap();
390                assert_eq!(encoded.count(), 2);
391            } else {
392                panic!("Expected binary scalar value");
393            }
394        } else {
395            panic!("Expected binary scalar values");
396        }
397    }
398
399    #[test]
400    fn test_uddsketch_state_size() {
401        let mut state = UddSketchState::new(10, 0.01).unwrap();
402        let initial_size = state.size();
403
404        // Add some values to create buckets
405        let array = Arc::new(Float64Array::from_iter_values((0..64).map(f64::from))) as ArrayRef;
406        state
407            .update_batch(&[array.clone(), array.clone(), array])
408            .unwrap();
409
410        let size_with_values = state.size();
411        assert!(
412            size_with_values > initial_size,
413            "Size should increase after adding values: initial={}, with_values={}",
414            initial_size,
415            size_with_values
416        );
417    }
418
419    #[test]
420    fn test_uddsketch_state_rejects_invalid_config() {
421        assert!(UddSketchState::new(6, 0.01).is_err());
422        assert!(UddSketchState::new(10, 1.0).is_err());
423
424        let mut maximum = UddSketchState::new(1_000_000, 0.01).unwrap();
425        let ScalarValue::Binary(Some(encoded)) = maximum.evaluate().unwrap() else {
426            panic!("Expected binary scalar value");
427        };
428        UddSketchRef::parse(&encoded).unwrap();
429        assert!(UddSketchState::new(1_000_001, 0.01).is_err());
430    }
431
432    #[test]
433    fn test_uddsketch_state_rejects_nan_batch() {
434        let mut state = UddSketchState::new(10, 0.01).unwrap();
435        let array = Arc::new(Float64Array::from(vec![1.0, f64::NAN])) as ArrayRef;
436
437        let error = state
438            .update_batch(&[array.clone(), array.clone(), array])
439            .unwrap_err();
440        assert!(error.to_string().contains("NaN values are not supported"));
441    }
442}