common_function/aggrs/approximate/
uddsketch.rs1use 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 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]; 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 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 let serialized = state.evaluate().unwrap();
253
254 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 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 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}