1use std::sync::Arc;
19
20use datafusion::arrow::array::Float64Array;
21use datafusion::arrow::datatypes::TimeUnit;
22use datafusion::common::DataFusionError;
23use datafusion::logical_expr::{ScalarUDF, Volatility};
24use datafusion::physical_plan::ColumnarValue;
25use datafusion_common::ScalarValue;
26use datafusion_expr::create_udf;
27use datatypes::arrow::array::Array;
28use datatypes::arrow::datatypes::DataType;
29
30use crate::error;
31use crate::functions::extract_array;
32use crate::range_array::RangeArray;
33
34struct FactorIterator<'a> {
36 is_scalar: bool,
37 array: Option<&'a Float64Array>,
38 scalar_val: f64,
39 index: usize,
40 len: usize,
41}
42
43impl<'a> FactorIterator<'a> {
44 fn new(value: &'a ColumnarValue, len: usize) -> Self {
45 let (is_scalar, array, scalar_val) = match value {
46 ColumnarValue::Array(arr) => {
47 (false, arr.as_any().downcast_ref::<Float64Array>(), f64::NAN)
48 }
49 ColumnarValue::Scalar(ScalarValue::Float64(Some(val))) => (true, None, *val),
50 _ => (true, None, f64::NAN),
51 };
52
53 Self {
54 is_scalar,
55 array,
56 scalar_val,
57 index: 0,
58 len,
59 }
60 }
61}
62
63impl<'a> Iterator for FactorIterator<'a> {
64 type Item = f64;
65
66 fn next(&mut self) -> Option<Self::Item> {
67 if self.index >= self.len {
68 return None;
69 }
70 self.index += 1;
71
72 if self.is_scalar {
73 return Some(self.scalar_val);
74 }
75
76 if let Some(array) = self.array {
77 if array.is_null(self.index - 1) {
78 Some(f64::NAN)
79 } else {
80 Some(array.value(self.index - 1))
81 }
82 } else {
83 Some(f64::NAN)
84 }
85 }
86}
87
88pub struct DoubleExponentialSmoothing;
103
104impl DoubleExponentialSmoothing {
105 pub const fn name() -> &'static str {
106 "prom_double_exponential_smoothing"
107 }
108
109 fn input_type() -> Vec<DataType> {
111 vec![
112 RangeArray::convert_data_type(DataType::Timestamp(TimeUnit::Millisecond, None)),
113 RangeArray::convert_data_type(DataType::Float64),
114 DataType::Float64,
116 DataType::Float64,
118 ]
119 }
120
121 fn return_type() -> DataType {
122 DataType::Float64
123 }
124
125 pub fn scalar_udf() -> ScalarUDF {
126 create_udf(
127 Self::name(),
128 Self::input_type(),
129 Self::return_type(),
130 Volatility::Volatile,
131 Arc::new(Self::double_exponential_smoothing) as _,
132 )
133 }
134
135 fn double_exponential_smoothing(
136 input: &[ColumnarValue],
137 ) -> Result<ColumnarValue, DataFusionError> {
138 error::ensure(
139 input.len() == 4,
140 DataFusionError::Plan(
141 "prom_double_exponential_smoothing function should have 4 inputs".to_string(),
142 ),
143 )?;
144
145 let ts_array = extract_array(&input[0])?;
146 let value_array = extract_array(&input[1])?;
147 let sf_col = &input[2];
148 let tf_col = &input[3];
149
150 let ts_range: RangeArray = RangeArray::try_new(ts_array.to_data().into())?;
151 let value_range: RangeArray = RangeArray::try_new(value_array.to_data().into())?;
152 let num_rows = ts_range.len();
153
154 error::ensure(
155 num_rows == value_range.len(),
156 DataFusionError::Execution(format!(
157 "{}: input arrays should have the same length, found {} and {}",
158 Self::name(),
159 num_rows,
160 value_range.len()
161 )),
162 )?;
163 error::ensure(
164 ts_range.value_type() == DataType::Timestamp(TimeUnit::Millisecond, None),
165 DataFusionError::Execution(format!(
166 "{}: expect TimestampMillisecond as time index array's type, found {}",
167 Self::name(),
168 ts_range.value_type()
169 )),
170 )?;
171 error::ensure(
172 value_range.value_type() == DataType::Float64,
173 DataFusionError::Execution(format!(
174 "{}: expect Float64 as value array's type, found {}",
175 Self::name(),
176 value_range.value_type()
177 )),
178 )?;
179
180 let mut result_array = Vec::with_capacity(ts_range.len());
182
183 let sf_iter = FactorIterator::new(sf_col, num_rows);
184 let tf_iter = FactorIterator::new(tf_col, num_rows);
185
186 let iter = (0..num_rows)
187 .map(|i| (ts_range.get(i), value_range.get(i)))
188 .zip(sf_iter.zip(tf_iter));
189
190 for ((timestamps, values), (sf, tf)) in iter {
191 let timestamps = timestamps.unwrap();
192 let values = values.unwrap();
193 let values = values
194 .as_any()
195 .downcast_ref::<Float64Array>()
196 .unwrap()
197 .values();
198 error::ensure(
199 timestamps.len() == values.len(),
200 DataFusionError::Execution(format!(
201 "{}: input arrays should have the same length, found {} and {}",
202 Self::name(),
203 timestamps.len(),
204 values.len()
205 )),
206 )?;
207
208 result_array.push(double_exponential_smoothing_impl(values, sf, tf));
209 }
210
211 let result = ColumnarValue::Array(Arc::new(Float64Array::from_iter(result_array)));
212 Ok(result)
213 }
214}
215
216fn calc_trend_value(i: usize, tf: f64, s0: f64, s1: f64, b: f64) -> f64 {
217 if i == 0 {
218 return b;
219 }
220 let x = tf * (s1 - s0);
221 let y = (1.0 - tf) * b;
222 x + y
223}
224
225fn double_exponential_smoothing_impl(values: &[f64], sf: f64, tf: f64) -> Option<f64> {
227 if sf.is_nan() || tf.is_nan() || values.is_empty() {
228 return Some(f64::NAN);
229 }
230 if sf < 0.0 || tf < 0.0 {
231 return Some(f64::NEG_INFINITY);
232 }
233 if sf > 1.0 || tf > 1.0 {
234 return Some(f64::INFINITY);
235 }
236
237 let l = values.len();
238 if l <= 2 {
239 return Some(f64::NAN);
241 }
242
243 let mut s0 = 0.0;
244 let mut s1 = values[0];
245 let mut b = values[1] - values[0];
246
247 for (i, value) in values.iter().enumerate().skip(1) {
248 let x = sf * value;
250 b = calc_trend_value(i - 1, tf, s0, s1, b);
252 let y = (1.0 - sf) * (s1 + b);
253 s0 = s1;
254 s1 = x + y;
255 }
256 Some(s1)
257}
258
259#[cfg(test)]
260mod tests {
261 use datafusion::arrow::array::{Float64Array, TimestampMillisecondArray};
262
263 use super::*;
264 use crate::functions::test_util::simple_range_udf_runner;
265
266 #[test]
267 fn test_double_exponential_smoothing_impl_empty() {
268 let sf = 0.5;
269 let tf = 0.5;
270 let values = &[];
271 assert!(
272 double_exponential_smoothing_impl(values, sf, tf)
273 .unwrap()
274 .is_nan()
275 );
276
277 let values = &[1.0, 2.0];
278 assert!(
279 double_exponential_smoothing_impl(values, sf, tf)
280 .unwrap()
281 .is_nan()
282 );
283 }
284
285 #[test]
286 fn test_double_exponential_smoothing_impl_nan() {
287 let values = &[1.0, 2.0, 3.0];
288 let sf = f64::NAN;
289 let tf = 0.5;
290 assert!(
291 double_exponential_smoothing_impl(values, sf, tf)
292 .unwrap()
293 .is_nan()
294 );
295
296 let values = &[1.0, 2.0, 3.0];
297 let sf = 0.5;
298 let tf = f64::NAN;
299 assert!(
300 double_exponential_smoothing_impl(values, sf, tf)
301 .unwrap()
302 .is_nan()
303 );
304 }
305
306 #[test]
307 fn test_double_exponential_smoothing_impl_validation_rules() {
308 let values = &[1.0, 2.0, 3.0];
309 let sf = -0.5;
310 let tf = 0.5;
311 assert_eq!(
312 double_exponential_smoothing_impl(values, sf, tf).unwrap(),
313 f64::NEG_INFINITY
314 );
315
316 let values = &[1.0, 2.0, 3.0];
317 let sf = 0.5;
318 let tf = -0.5;
319 assert_eq!(
320 double_exponential_smoothing_impl(values, sf, tf).unwrap(),
321 f64::NEG_INFINITY
322 );
323
324 let values = &[1.0, 2.0, 3.0];
325 let sf = 1.5;
326 let tf = 0.5;
327 assert_eq!(
328 double_exponential_smoothing_impl(values, sf, tf).unwrap(),
329 f64::INFINITY
330 );
331
332 let values = &[1.0, 2.0, 3.0];
333 let sf = 0.5;
334 let tf = 1.5;
335 assert_eq!(
336 double_exponential_smoothing_impl(values, sf, tf).unwrap(),
337 f64::INFINITY
338 );
339 }
340
341 #[test]
342 fn test_double_exponential_smoothing_impl() {
343 let sf = 0.5;
344 let tf = 0.1;
345 let values = &[1.0, 2.0, 3.0, 4.0, 5.0];
346 assert_eq!(double_exponential_smoothing_impl(values, sf, tf), Some(5.0));
347 let values = &[50.0, 52.0, 95.0, 59.0, 52.0, 45.0, 38.0, 10.0, 47.0, 40.0];
348 assert_eq!(
349 double_exponential_smoothing_impl(values, sf, tf),
350 Some(38.18119566835938)
351 );
352 }
353
354 #[test]
355 fn test_double_exponential_smoothing_impl_copy_oracle() {
356 let normal_values = (0..240)
357 .map(|i| (i as f64 - 120.0) * 0.25)
358 .collect::<Vec<_>>();
359 let special_values = (0..240)
360 .map(|i| match i % 8 {
361 0 => 0.0,
362 1 => -0.0,
363 2 => f64::INFINITY,
364 3 => f64::NEG_INFINITY,
365 4 => f64::NAN,
366 5 => f64::from_bits(0x7ff8_0000_0000_0001),
367 6 => 42.5,
368 _ => -42.5,
369 })
370 .collect::<Vec<_>>();
371 let factors = [
372 (0.0, 0.0),
373 (-0.0, 1.0),
374 (0.5, 0.1),
375 (1.0, 1.0),
376 (-0.5, 0.5),
377 (0.5, -0.5),
378 (1.5, 0.5),
379 (0.5, 1.5),
380 (f64::NAN, 0.5),
381 (0.5, f64::NAN),
382 (f64::INFINITY, 0.5),
383 (0.5, f64::INFINITY),
384 (f64::NEG_INFINITY, 0.5),
385 (0.5, f64::NEG_INFINITY),
386 ];
387
388 for (values_name, values) in [
389 ("normal", normal_values.as_slice()),
390 ("special", special_values.as_slice()),
391 ] {
392 for len in [0, 1, 2, 3, 20, 240] {
393 let values = &values[..len];
394 for (sf, tf) in factors {
395 let old = double_exponential_smoothing_impl_with_copy(values, sf, tf).unwrap();
396 let new = double_exponential_smoothing_impl(values, sf, tf).unwrap();
397 let case = format!("values={values_name}, len={len}, sf={sf:?}, tf={tf:?}");
398
399 if old.is_nan() || new.is_nan() {
400 assert!(
401 old.is_nan() && new.is_nan(),
402 "NaN mismatch for {case}: old={old:?}, new={new:?}"
403 );
404 assert_eq!(
405 old.to_bits(),
406 new.to_bits(),
407 "NaN bit difference for {case}: old={:#018x}, new={:#018x}",
408 old.to_bits(),
409 new.to_bits(),
410 );
411 } else {
412 assert_eq!(
413 old.to_bits(),
414 new.to_bits(),
415 "non-NaN bit difference for {case}: old={old:?}, new={new:?}"
416 );
417 }
418 }
419 }
420 }
421 }
422
423 fn double_exponential_smoothing_impl_with_copy(
424 values: &[f64],
425 sf: f64,
426 tf: f64,
427 ) -> Option<f64> {
428 if sf.is_nan() || tf.is_nan() || values.is_empty() {
429 return Some(f64::NAN);
430 }
431 if sf < 0.0 || tf < 0.0 {
432 return Some(f64::NEG_INFINITY);
433 }
434 if sf > 1.0 || tf > 1.0 {
435 return Some(f64::INFINITY);
436 }
437
438 if values.len() <= 2 {
439 return Some(f64::NAN);
440 }
441
442 let values = values.to_vec();
443 let mut s0 = 0.0;
444 let mut s1 = values[0];
445 let mut b = values[1] - values[0];
446
447 for (i, value) in values.iter().enumerate().skip(1) {
448 let x = sf * value;
449 b = calc_trend_value(i - 1, tf, s0, s1, b);
450 let y = (1.0 - sf) * (s1 + b);
451 s0 = s1;
452 s1 = x + y;
453 }
454 Some(s1)
455 }
456
457 #[test]
458 fn test_prom_double_exponential_smoothing_monotonic() {
459 let ranges = [(0, 5)];
460 let ts_array = Arc::new(TimestampMillisecondArray::from_iter(
461 [1000i64, 3000, 5000, 7000, 9000, 11000, 13000, 15000, 17000]
462 .into_iter()
463 .map(Some),
464 ));
465 let values_array = Arc::new(Float64Array::from_iter([1.0, 2.0, 3.0, 4.0, 5.0]));
466 let ts_range_array = RangeArray::from_ranges(ts_array, ranges).unwrap();
467 let value_range_array = RangeArray::from_ranges(values_array, ranges).unwrap();
468 simple_range_udf_runner(
469 DoubleExponentialSmoothing::scalar_udf(),
470 ts_range_array,
471 value_range_array,
472 vec![
473 ScalarValue::Float64(Some(0.5)),
474 ScalarValue::Float64(Some(0.1)),
475 ],
476 vec![Some(5.0)],
477 );
478 }
479
480 #[test]
481 fn test_prom_double_exponential_smoothing_non_monotonic() {
482 let ranges = [(0, 10)];
483 let ts_array = Arc::new(TimestampMillisecondArray::from_iter(
484 [
485 1000i64, 3000, 5000, 7000, 9000, 11000, 13000, 15000, 17000, 19000,
486 ]
487 .into_iter()
488 .map(Some),
489 ));
490 let values_array = Arc::new(Float64Array::from_iter([
491 50.0, 52.0, 95.0, 59.0, 52.0, 45.0, 38.0, 10.0, 47.0, 40.0,
492 ]));
493 let ts_range_array = RangeArray::from_ranges(ts_array, ranges).unwrap();
494 let value_range_array = RangeArray::from_ranges(values_array, ranges).unwrap();
495 simple_range_udf_runner(
496 DoubleExponentialSmoothing::scalar_udf(),
497 ts_range_array,
498 value_range_array,
499 vec![
500 ScalarValue::Float64(Some(0.5)),
501 ScalarValue::Float64(Some(0.1)),
502 ],
503 vec![Some(38.18119566835938)],
504 );
505 }
506
507 #[test]
508 fn test_promql_trends() {
509 let ranges = vec![(0, 801)];
510
511 let trends = vec![
512 ("0+10x1000 100+30x1000", 8000.0),
514 ("0+20x1000 200+30x1000", 16000.0),
515 ("0+30x1000 300+80x1000", 24000.0),
516 ("0+40x2000", 32000.0),
517 ("8000-10x1000", 0.0),
519 ("0-20x1000", -16000.0),
520 ("0+30x1000 300-80x1000", 24000.0),
521 ("0-40x1000 0+40x1000", -32000.0),
522 ];
523
524 for (query, expected) in trends {
525 let (ts_range_array, value_range_array) =
526 create_ts_and_value_range_arrays(query, ranges.clone());
527 simple_range_udf_runner(
528 DoubleExponentialSmoothing::scalar_udf(),
529 ts_range_array,
530 value_range_array,
531 vec![
532 ScalarValue::Float64(Some(0.01)),
533 ScalarValue::Float64(Some(0.1)),
534 ],
535 vec![Some(expected)],
536 );
537 }
538 }
539
540 fn create_ts_and_value_range_arrays(
541 input: &str,
542 ranges: Vec<(u32, u32)>,
543 ) -> (RangeArray, RangeArray) {
544 let promql_range = create_test_range_from_promql_series(input);
545 let ts_array = Arc::new(TimestampMillisecondArray::from_iter(
546 (0..(promql_range.len() as i64)).map(Some),
547 ));
548 let values_array = Arc::new(Float64Array::from_iter(promql_range));
549 let ts_range_array = RangeArray::from_ranges(ts_array, ranges.clone()).unwrap();
550 let value_range_array = RangeArray::from_ranges(values_array, ranges).unwrap();
551 (ts_range_array, value_range_array)
552 }
553
554 fn create_test_range_from_promql_series(input: &str) -> Vec<f64> {
557 input.split(' ').map(parse_promql_series_entry).fold(
558 Vec::new(),
559 |mut acc, (start, end, step, operation)| {
560 if operation.eq("+") {
561 let iter = (start..=((step * end) + start))
562 .step_by(step as usize)
563 .map(|x| x as f64);
564 acc.extend(iter);
565 } else {
566 let iter = (((-step * end) + start)..=start)
567 .rev()
568 .step_by(step as usize)
569 .map(|x| x as f64);
570 acc.extend(iter);
571 };
572 acc
573 },
574 )
575 }
576
577 fn parse_promql_series_entry(input: &str) -> (i32, i32, i32, &str) {
580 let mut parts = input.split('x');
581 let start_operation_step = parts.next().unwrap();
582 let operation = start_operation_step
583 .split(char::is_numeric)
584 .find(|&x| !x.is_empty())
585 .unwrap();
586 let start_step = start_operation_step
587 .split(operation)
588 .map(|s| s.parse::<i32>().unwrap())
589 .collect::<Vec<_>>();
590 let start = *start_step.first().unwrap();
591 let step = *start_step.last().unwrap();
592 let end = parts.next().unwrap().parse::<i32>().unwrap();
593 (start, end, step, operation)
594 }
595}