Skip to main content

promql/functions/
quantile_aggr.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
15use 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
41/// Create a quantile `AggregateUDF` for PromQL quantile operator,
42/// which calculates φ-quantile (0 ≤ φ ≤ 1) over dimensions
43pub fn quantile_udaf() -> Arc<AggregateUDF> {
44    Arc::new(create_udaf(
45        QUANTILE_NAME,
46        // Input type: (φ, values)
47        vec![DataType::Float64, DataType::Float64],
48        // Output type: the φ-quantile
49        Arc::new(DataType::Float64),
50        Volatility::Volatile,
51        // Create the accumulator
52        Arc::new(QuantileAccumulator::from_args),
53        // Intermediate state types
54        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}