1use std::sync::Arc;
16
17use datafusion::arrow::array::{ArrayRef, AsArray};
18use datafusion::common::cast::{as_list_array, as_primitive_array, as_struct_array};
19use datafusion::error::{DataFusionError, Result as DfResult};
20use datafusion::logical_expr::{Accumulator as DfAccumulator, AggregateUDF, Volatility};
21use datafusion::physical_plan::expressions::Literal;
22use datafusion::prelude::create_udaf;
23use datafusion_common::ScalarValue;
24use datafusion_expr::function::AccumulatorArgs;
25use datatypes::arrow::array::{ListArray, StructArray};
26use datatypes::arrow::datatypes::{DataType, Field, Float64Type};
27
28use crate::functions::quantile::quantile_impl;
29
30pub const QUANTILE_NAME: &str = "quantile";
31
32const VALUES_FIELD_NAME: &str = "values";
33const DEFAULT_LIST_FIELD_NAME: &str = "item";
34
35#[derive(Debug, Default)]
36pub struct QuantileAccumulator {
37 q: f64,
38 values: Vec<Option<f64>>,
39}
40
41pub fn quantile_udaf() -> Arc<AggregateUDF> {
44 Arc::new(create_udaf(
45 QUANTILE_NAME,
46 vec![DataType::Float64, DataType::Float64],
48 Arc::new(DataType::Float64),
50 Volatility::Volatile,
51 Arc::new(QuantileAccumulator::from_args),
53 Arc::new(vec![DataType::Struct(
55 vec![Field::new(
56 VALUES_FIELD_NAME,
57 DataType::List(Arc::new(Field::new(
58 DEFAULT_LIST_FIELD_NAME,
59 DataType::Float64,
60 true,
61 ))),
62 false,
63 )]
64 .into(),
65 )]),
66 ))
67}
68
69impl QuantileAccumulator {
70 fn new(q: f64) -> Self {
71 Self {
72 q,
73 ..Default::default()
74 }
75 }
76
77 pub fn from_args(args: AccumulatorArgs) -> DfResult<Box<dyn DfAccumulator>> {
78 if args.exprs.len() != 2 {
79 return Err(DataFusionError::Plan(
80 "Quantile function should have 2 inputs".to_string(),
81 ));
82 }
83
84 let q = match &args.exprs[0]
85 .downcast_ref::<Literal>()
86 .map(|lit| lit.value())
87 {
88 Some(ScalarValue::Float64(Some(q))) => *q,
89 _ => {
90 return Err(DataFusionError::Internal(
91 "Invalid quantile value".to_string(),
92 ));
93 }
94 };
95
96 Ok(Box::new(Self::new(q)))
97 }
98}
99
100impl DfAccumulator for QuantileAccumulator {
101 fn update_batch(&mut self, values: &[ArrayRef]) -> DfResult<()> {
102 let f64_array = values[1].as_primitive::<Float64Type>();
103
104 self.values.extend(f64_array);
105
106 Ok(())
107 }
108
109 fn evaluate(&mut self) -> DfResult<ScalarValue> {
110 let values: Vec<_> = self.values.iter().map(|v| v.unwrap_or(0.0)).collect();
111
112 let result = quantile_impl(&values, self.q);
113
114 ScalarValue::new_primitive::<Float64Type>(result, &DataType::Float64)
115 }
116
117 fn size(&self) -> usize {
118 std::mem::size_of::<Self>() + self.values.capacity() * std::mem::size_of::<Option<f64>>()
119 }
120
121 fn state(&mut self) -> DfResult<Vec<ScalarValue>> {
122 let values_array = Arc::new(ListArray::from_iter_primitive::<Float64Type, _, _>(vec![
123 Some(self.values.clone()),
124 ]));
125
126 let state_struct = StructArray::new(
127 vec![Field::new(
128 VALUES_FIELD_NAME,
129 DataType::List(Arc::new(Field::new(
130 DEFAULT_LIST_FIELD_NAME,
131 DataType::Float64,
132 true,
133 ))),
134 false,
135 )]
136 .into(),
137 vec![values_array],
138 None,
139 );
140
141 Ok(vec![ScalarValue::Struct(Arc::new(state_struct))])
142 }
143
144 fn merge_batch(&mut self, states: &[ArrayRef]) -> DfResult<()> {
145 if states.is_empty() {
146 return Ok(());
147 }
148
149 for state in states {
150 let state = as_struct_array(state)?;
151
152 for list in as_list_array(state.column(0))?.iter().flatten() {
153 let f64_array = as_primitive_array::<Float64Type>(&list)?.clone();
154 self.values.extend(&f64_array);
155 }
156 }
157
158 Ok(())
159 }
160}
161#[cfg(test)]
162mod tests {
163 use std::sync::Arc;
164
165 use datafusion::arrow::array::{ArrayRef, Float64Array};
166 use datafusion_common::ScalarValue;
167
168 use super::*;
169
170 fn create_f64_array(values: Vec<Option<f64>>) -> ArrayRef {
171 Arc::new(Float64Array::from(values)) as ArrayRef
172 }
173
174 #[test]
175 fn test_quantile_accumulator_empty() {
176 let mut accumulator = QuantileAccumulator::new(0.5);
177
178 let result = accumulator.evaluate().unwrap();
179
180 match result {
181 ScalarValue::Float64(_) => (),
182 _ => panic!("Expected Float64 scalar value"),
183 }
184 }
185
186 #[test]
187 fn test_quantile_accumulator_single_value() {
188 let mut accumulator = QuantileAccumulator::new(0.5);
189 let q = create_f64_array(vec![Some(0.5)]);
190 let input = create_f64_array(vec![Some(10.0)]);
191
192 accumulator.update_batch(&[q, input]).unwrap();
193 let result = accumulator.evaluate().unwrap();
194
195 assert_eq!(result, ScalarValue::Float64(Some(10.0)));
196 }
197
198 #[test]
199 fn test_quantile_accumulator_multiple_values() {
200 let mut accumulator = QuantileAccumulator::new(0.5);
201 let q = create_f64_array(vec![Some(0.5)]);
202 let input = create_f64_array(vec![Some(1.0), Some(2.0), Some(3.0), Some(4.0), Some(5.0)]);
203
204 accumulator.update_batch(&[q, input]).unwrap();
205 let result = accumulator.evaluate().unwrap();
206
207 assert_eq!(result, ScalarValue::Float64(Some(3.0)));
208 }
209
210 #[test]
211 fn test_quantile_accumulator_with_nulls() {
212 let mut accumulator = QuantileAccumulator::new(0.5);
213 let q = create_f64_array(vec![Some(0.5)]);
214 let input = create_f64_array(vec![Some(1.0), None, Some(3.0), Some(4.0), Some(5.0)]);
215
216 accumulator.update_batch(&[q, input]).unwrap();
217
218 let result = accumulator.evaluate().unwrap();
219 assert_eq!(result, ScalarValue::Float64(Some(3.0)));
220 }
221
222 #[test]
223 fn test_quantile_accumulator_multiple_batches() {
224 let mut accumulator = QuantileAccumulator::new(0.5);
225 let q = create_f64_array(vec![Some(0.5)]);
226 let input1 = create_f64_array(vec![Some(1.0), Some(2.0)]);
227 let input2 = create_f64_array(vec![Some(3.0), Some(4.0), Some(5.0)]);
228
229 accumulator.update_batch(&[q.clone(), input1]).unwrap();
230 accumulator.update_batch(&[q, input2]).unwrap();
231
232 let result = accumulator.evaluate().unwrap();
233 assert_eq!(result, ScalarValue::Float64(Some(3.0)));
234 }
235
236 #[test]
237 fn test_quantile_accumulator_different_quantiles() {
238 let mut min_accumulator = QuantileAccumulator::new(0.0);
239 let q = create_f64_array(vec![Some(0.0)]);
240 let input = create_f64_array(vec![Some(1.0), Some(2.0), Some(3.0), Some(4.0), Some(5.0)]);
241 min_accumulator.update_batch(&[q, input.clone()]).unwrap();
242 assert_eq!(
243 min_accumulator.evaluate().unwrap(),
244 ScalarValue::Float64(Some(1.0))
245 );
246
247 let mut q1_accumulator = QuantileAccumulator::new(0.25);
248 let q = create_f64_array(vec![Some(0.25)]);
249 q1_accumulator.update_batch(&[q, input.clone()]).unwrap();
250 assert_eq!(
251 q1_accumulator.evaluate().unwrap(),
252 ScalarValue::Float64(Some(2.0))
253 );
254
255 let mut q3_accumulator = QuantileAccumulator::new(0.75);
256 let q = create_f64_array(vec![Some(0.75)]);
257 q3_accumulator.update_batch(&[q, input.clone()]).unwrap();
258 assert_eq!(
259 q3_accumulator.evaluate().unwrap(),
260 ScalarValue::Float64(Some(4.0))
261 );
262
263 let mut max_accumulator = QuantileAccumulator::new(1.0);
264 let q = create_f64_array(vec![Some(1.0)]);
265 max_accumulator.update_batch(&[q, input]).unwrap();
266 assert_eq!(
267 max_accumulator.evaluate().unwrap(),
268 ScalarValue::Float64(Some(5.0))
269 );
270 }
271
272 #[test]
273 fn test_quantile_accumulator_size() {
274 let mut accumulator = QuantileAccumulator::new(0.5);
275 let q = create_f64_array(vec![Some(0.5)]);
276 let input = create_f64_array(vec![Some(1.0), Some(2.0), Some(3.0)]);
277
278 let initial_size = accumulator.size();
279 accumulator.update_batch(&[q, input]).unwrap();
280 let after_update_size = accumulator.size();
281
282 assert!(after_update_size >= initial_size);
283 }
284
285 #[test]
286 fn test_quantile_accumulator_state_and_merge() -> DfResult<()> {
287 let mut acc1 = QuantileAccumulator::new(0.5);
288 let q = create_f64_array(vec![Some(0.5)]);
289 let input1 = create_f64_array(vec![Some(1.0), Some(2.0)]);
290 acc1.update_batch(&[q, input1])?;
291
292 let state1 = acc1.state()?;
293
294 let mut acc2 = QuantileAccumulator::new(0.5);
295 let q = create_f64_array(vec![Some(0.5)]);
296 let input2 = create_f64_array(vec![Some(3.0), Some(4.0), Some(5.0)]);
297 acc2.update_batch(&[q, input2])?;
298
299 let mut struct_builders = vec![];
300 for scalar in &state1 {
301 if let ScalarValue::Struct(struct_array) = scalar {
302 struct_builders.push(struct_array.clone() as ArrayRef);
303 }
304 }
305
306 acc2.merge_batch(&struct_builders)?;
307
308 let result = acc2.evaluate()?;
309
310 assert_eq!(result, ScalarValue::Float64(Some(3.0)));
311
312 Ok(())
313 }
314
315 #[test]
316 fn test_quantile_accumulator_with_extreme_values() {
317 let mut accumulator = QuantileAccumulator::new(0.5);
318 let q = create_f64_array(vec![Some(0.5)]);
319 let input = create_f64_array(vec![Some(f64::MAX), Some(f64::MIN), Some(0.0)]);
320
321 accumulator.update_batch(&[q, input]).unwrap();
322 let _result = accumulator.evaluate().unwrap();
323 }
324
325 #[test]
326 fn test_quantile_udaf_creation() {
327 let udaf = quantile_udaf();
328
329 assert_eq!(udaf.name(), QUANTILE_NAME);
330 assert_eq!(udaf.return_type(&[]).unwrap(), DataType::Float64);
331 }
332}