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 .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]; 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 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 let serialized = state.evaluate().unwrap();
251
252 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 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 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}