Skip to main content

common_function/scalars/
welford_stddev.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 `stddev_pop_calc`.
16
17use std::fmt;
18use std::fmt::Display;
19use std::sync::Arc;
20
21use datafusion_common::DataFusionError;
22use datafusion_common::arrow::array::{Array, AsArray, Float64Builder};
23use datafusion_expr::{ColumnarValue, ScalarFunctionArgs, Signature, Volatility};
24use datatypes::arrow::datatypes::DataType;
25
26use crate::aggrs::approximate::welford::WelfordState;
27use crate::function::{Function, extract_args};
28use crate::function_registry::FunctionRegistry;
29
30const NAME: &str = "stddev_pop_calc";
31
32/// Calculates population standard deviation from a serialized Welford state.
33#[derive(Debug)]
34pub(crate) struct WelfordStddevFunction {
35    signature: Signature,
36}
37
38impl WelfordStddevFunction {
39    pub fn register(registry: &FunctionRegistry) {
40        registry.register_scalar(Self::default());
41    }
42}
43
44impl Default for WelfordStddevFunction {
45    fn default() -> Self {
46        Self {
47            signature: Signature::exact(vec![DataType::Binary], Volatility::Immutable),
48        }
49    }
50}
51
52impl Display for WelfordStddevFunction {
53    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
54        write!(f, "{}", NAME.to_ascii_uppercase())
55    }
56}
57
58impl Function for WelfordStddevFunction {
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] = extract_args(self.name(), &args)?;
76        let Some(states) = arg.as_binary_opt::<i32>() else {
77            return Err(DataFusionError::Execution(format!(
78                "'{}' expects argument to be Binary datatype, got {}",
79                self.name(),
80                arg.data_type()
81            )));
82        };
83        let mut builder = Float64Builder::with_capacity(states.len());
84        for state in states.iter() {
85            match state.and_then(decode_population_stddev) {
86                Some(stddev) => builder.append_value(stddev),
87                None => builder.append_null(),
88            }
89        }
90
91        Ok(ColumnarValue::Array(Arc::new(builder.finish())))
92    }
93}
94
95fn decode_population_stddev(encoded: &[u8]) -> Option<f64> {
96    match WelfordState::decode(encoded) {
97        Ok(state) => state.population_stddev(),
98        Err(error) => {
99            common_telemetry::trace!("Failed to decode Welford state: {}", error);
100            None
101        }
102    }
103}
104
105#[cfg(test)]
106mod tests {
107    use std::sync::Arc;
108
109    use arrow_schema::Field;
110    use datafusion_common::arrow::array::{Array, AsArray, BinaryArray};
111    use datafusion_common::arrow::datatypes::Float64Type;
112    use datafusion_expr::{ColumnarValue, ScalarFunctionArgs};
113    use datatypes::arrow::datatypes::DataType;
114
115    use super::*;
116    use crate::aggrs::approximate::welford::WelfordState;
117    use crate::function::Function;
118
119    fn invoke(states: BinaryArray) -> ColumnarValue {
120        WelfordStddevFunction::default()
121            .invoke_with_args(ScalarFunctionArgs {
122                number_rows: states.len(),
123                args: vec![ColumnarValue::Array(Arc::new(states))],
124                arg_fields: vec![],
125                return_field: Arc::new(Field::new("x", DataType::Float64, true)),
126                config_options: Arc::new(Default::default()),
127            })
128            .unwrap()
129    }
130
131    #[test]
132    fn test_populated_welford_state_returns_population_stddev() {
133        let populated = WelfordState {
134            count: 4,
135            mean: 2.5,
136            m2: 5.0,
137        }
138        .encode();
139
140        let ColumnarValue::Array(output) =
141            invoke(BinaryArray::from(vec![Some(populated.as_slice())]))
142        else {
143            panic!("Expected array result");
144        };
145        let output = output.as_primitive::<Float64Type>();
146        assert!((output.value(0) - 1.25_f64.sqrt()).abs() < 1e-12);
147    }
148
149    #[test]
150    fn test_empty_malformed_and_null_states_return_null() {
151        let empty = WelfordState::default().encode();
152        let noncanonical_singleton = WelfordState {
153            count: 1,
154            mean: 0.0,
155            m2: 1.0,
156        }
157        .encode();
158
159        let ColumnarValue::Array(output) = invoke(BinaryArray::from(vec![
160            Some(empty.as_slice()),
161            Some(b"invalid".as_slice()),
162            Some(noncanonical_singleton.as_slice()),
163            None,
164        ])) else {
165            panic!("Expected array result");
166        };
167        assert_eq!(output.null_count(), 4);
168    }
169
170    #[test]
171    fn test_stddev_pop_calc_metadata() {
172        let function = WelfordStddevFunction::default();
173
174        assert_eq!(function.name(), "stddev_pop_calc");
175        assert_eq!(
176            function.return_type(&[DataType::Binary]).unwrap(),
177            DataType::Float64
178        );
179    }
180
181    #[test]
182    fn test_stddev_pop_calc_rejects_wrong_argument_count() {
183        let error = WelfordStddevFunction::default()
184            .invoke_with_args(ScalarFunctionArgs {
185                args: vec![],
186                arg_fields: vec![],
187                number_rows: 0,
188                return_field: Arc::new(Field::new("x", DataType::Float64, true)),
189                config_options: Arc::new(Default::default()),
190            })
191            .unwrap_err();
192
193        assert!(
194            error
195                .to_string()
196                .contains("stddev_pop_calc function requires 1 argument, got 0")
197        );
198    }
199}