Skip to main content

common_function/scalars/
avg_calc.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 scalar function `avg_calc`.
16
17use std::fmt;
18use std::fmt::Display;
19use std::sync::Arc;
20
21use datafusion_common::arrow::array::{Array, AsArray, Float64Builder};
22use datafusion_common::{DataFusionError, ScalarValue};
23use datafusion_expr::{ColumnarValue, ScalarFunctionArgs, Signature, Volatility};
24use datatypes::arrow::datatypes::DataType;
25
26use crate::aggrs::approximate::avg::AvgState;
27use crate::function::Function;
28use crate::function_registry::FunctionRegistry;
29
30const NAME: &str = "avg_calc";
31
32/// Calculates an average from a serialized AVG1 state.
33#[derive(Debug)]
34pub(crate) struct AvgCalcFunction {
35    signature: Signature,
36}
37
38impl AvgCalcFunction {
39    pub fn register(registry: &FunctionRegistry) {
40        registry.register_scalar(Self::default());
41    }
42}
43
44impl Default for AvgCalcFunction {
45    fn default() -> Self {
46        Self {
47            signature: Signature::exact(vec![DataType::Binary], Volatility::Immutable),
48        }
49    }
50}
51
52impl Display for AvgCalcFunction {
53    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
54        write!(f, "{}", NAME.to_ascii_uppercase())
55    }
56}
57
58impl Function for AvgCalcFunction {
59    fn name(&self) -> &str {
60        NAME
61    }
62
63    fn return_type(&self, _: &[DataType]) -> datafusion_common::Result<DataType> {
64        Ok(DataType::Float64)
65    }
66
67    fn signature(&self) -> &Signature {
68        &self.signature
69    }
70
71    fn invoke_with_args(
72        &self,
73        args: ScalarFunctionArgs,
74    ) -> datafusion_common::Result<ColumnarValue> {
75        let [arg] = datafusion_common::utils::take_function_args(self.name(), &args.args)?;
76        match arg {
77            ColumnarValue::Scalar(ScalarValue::Binary(state)) => {
78                Ok(ColumnarValue::Scalar(ScalarValue::Float64(
79                    state
80                        .as_deref()
81                        .map(AvgState::decode)
82                        .transpose()?
83                        .and_then(|state| state.average()),
84                )))
85            }
86            ColumnarValue::Scalar(ScalarValue::Null) => {
87                Ok(ColumnarValue::Scalar(ScalarValue::Float64(None)))
88            }
89            ColumnarValue::Array(states) => {
90                let Some(states) = states.as_binary_opt::<i32>() else {
91                    return Err(invalid_type(self.name(), states.data_type()));
92                };
93                let mut builder = Float64Builder::with_capacity(states.len());
94                for state in states.iter() {
95                    builder.append_option(match state {
96                        Some(state) => AvgState::decode(state)?.average(),
97                        None => None,
98                    });
99                }
100                Ok(ColumnarValue::Array(Arc::new(builder.finish())))
101            }
102            _ => Err(invalid_type(self.name(), &arg.data_type())),
103        }
104    }
105}
106
107fn invalid_type(name: &str, data_type: &DataType) -> DataFusionError {
108    DataFusionError::Execution(format!(
109        "'{name}' expects argument to be Binary datatype, got {data_type}"
110    ))
111}
112
113#[cfg(test)]
114mod tests {
115    use std::sync::Arc;
116
117    use arrow_schema::Field;
118    use datafusion::arrow::array::{Array, AsArray, BinaryArray, Float64Array};
119    use datafusion::logical_expr::Accumulator;
120    use datafusion::prelude::SessionContext;
121    use datafusion_common::arrow::datatypes::Float64Type;
122    use datafusion_expr::{ColumnarValue, ScalarFunctionArgs};
123
124    use super::*;
125    use crate::aggrs::approximate::avg::AvgAccumulator;
126    use crate::function::{Function, FunctionContext};
127    use crate::function_registry::FUNCTION_REGISTRY;
128
129    fn produce_state(values: Vec<Option<f64>>) -> Vec<u8> {
130        let mut accumulator = AvgAccumulator::default();
131        accumulator
132            .update_batch(&[Arc::new(Float64Array::from(values))])
133            .unwrap();
134        let ScalarValue::Binary(Some(state)) = accumulator.evaluate().unwrap() else {
135            panic!("AVG state must be binary");
136        };
137        state
138    }
139
140    fn invoke(arg: ColumnarValue, number_rows: usize) -> datafusion_common::Result<ColumnarValue> {
141        AvgCalcFunction::default().invoke_with_args(ScalarFunctionArgs {
142            args: vec![arg],
143            arg_fields: vec![],
144            number_rows,
145            return_field: Arc::new(Field::new("x", DataType::Float64, true)),
146            config_options: Arc::new(Default::default()),
147        })
148    }
149
150    #[test]
151    fn scalar_and_array_states_decode_to_averages() {
152        let state = produce_state(vec![Some(1.0), Some(2.0), Some(6.0)]);
153        let ColumnarValue::Scalar(ScalarValue::Float64(Some(value))) =
154            invoke(ColumnarValue::Scalar(ScalarValue::Binary(Some(state))), 1).unwrap()
155        else {
156            panic!("Expected Float64 scalar");
157        };
158        assert_eq!(value, 3.0);
159
160        let ColumnarValue::Scalar(ScalarValue::Float64(None)) =
161            invoke(ColumnarValue::Scalar(ScalarValue::Binary(None)), 1).unwrap()
162        else {
163            panic!("Expected NULL Float64 scalar");
164        };
165        let empty = produce_state(vec![None]);
166        let ColumnarValue::Scalar(ScalarValue::Float64(None)) = invoke(
167            ColumnarValue::Scalar(ScalarValue::Binary(Some(empty.clone()))),
168            1,
169        )
170        .unwrap() else {
171            panic!("Expected NULL Float64 scalar");
172        };
173
174        let infinity = produce_state(vec![Some(f64::INFINITY)]);
175        let nan = produce_state(vec![Some(f64::NAN)]);
176        let ColumnarValue::Array(result) = invoke(
177            ColumnarValue::Array(Arc::new(BinaryArray::from(vec![
178                Some(empty.as_slice()),
179                None,
180                Some(infinity.as_slice()),
181                Some(nan.as_slice()),
182            ]))),
183            4,
184        )
185        .unwrap() else {
186            panic!("Expected Float64 array");
187        };
188        let result = result.as_primitive::<Float64Type>();
189        assert!(result.is_null(0));
190        assert!(result.is_null(1));
191        assert_eq!(result.value(2), f64::INFINITY);
192        assert!(result.value(3).is_nan());
193    }
194
195    #[test]
196    fn malformed_and_unknown_version_states_fail_the_whole_batch() {
197        let valid = produce_state(vec![Some(3.0)]);
198        let mut unknown_version = valid.clone();
199        unknown_version[..4].copy_from_slice(b"AVG2");
200
201        assert!(
202            invoke(
203                ColumnarValue::Scalar(ScalarValue::Binary(Some(unknown_version.clone()))),
204                1,
205            )
206            .is_err()
207        );
208        assert!(
209            invoke(
210                ColumnarValue::Array(Arc::new(BinaryArray::from(vec![
211                    Some(valid.as_slice()),
212                    Some(b"malformed".as_slice()),
213                    Some(unknown_version.as_slice()),
214                ]))),
215                3,
216            )
217            .is_err()
218        );
219    }
220
221    #[tokio::test]
222    async fn registry_query_decodes_avg_state_and_weighted_avg_merge() {
223        let ctx = SessionContext::new();
224        let avg_calc = FUNCTION_REGISTRY
225            .get_function(NAME)
226            .expect("avg_calc must be registered")
227            .provide(FunctionContext::default());
228        ctx.register_udf(avg_calc);
229        for name in ["avg_state", "avg_merge"] {
230            ctx.register_udaf(
231                FUNCTION_REGISTRY
232                    .get_aggr_func(name)
233                    .expect("AVG aggregate must be registered"),
234            );
235        }
236
237        let batches = ctx
238            .sql(
239                "SELECT avg_calc(avg_state(CAST(value AS DOUBLE))) FROM \
240                 (VALUES (1.0), (2.0), (6.0)) AS values_table(value)",
241            )
242            .await
243            .unwrap()
244            .collect()
245            .await
246            .unwrap();
247        let result = batches[0].column(0).as_primitive::<Float64Type>();
248        assert_eq!(result.value(0), 3.0);
249
250        let batches = ctx
251            .sql(
252                "WITH states AS (\
253                 SELECT avg_state(CAST(value AS DOUBLE)) AS state FROM (VALUES (1.0), (3.0)) AS left_values(value) \
254                 UNION ALL \
255                 SELECT avg_state(CAST(value AS DOUBLE)) AS state FROM (VALUES (6.0)) AS right_values(value)\
256                 ) SELECT avg_calc(avg_merge(state)) FROM states",
257            )
258            .await
259            .unwrap()
260            .collect()
261            .await
262            .unwrap();
263        let result = batches[0].column(0).as_primitive::<Float64Type>();
264        assert_eq!(result.value(0), 10.0 / 3.0);
265    }
266}