common_function/scalars/
avg_calc.rs1use 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#[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}