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