use super::{apply_policy, mean_kahan, quantile::quantile, MissingDataPolicy, TransformError};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum EdgePolicy {
#[default]
None,
Partial,
Pad,
}
pub fn rolling_mean(values: &[f64], window: usize, edge: EdgePolicy, missing: MissingDataPolicy) -> Result<Vec<Option<f64>>, TransformError> {
rolling_reduce(values, window, edge, missing, mean_kahan)
}
pub fn rolling_median(values: &[f64], window: usize, edge: EdgePolicy, missing: MissingDataPolicy) -> Result<Vec<Option<f64>>, TransformError> {
rolling_reduce(values, window, edge, missing, |w| quantile(w, 0.5, MissingDataPolicy::Propagate).unwrap_or(f64::NAN))
}
fn rolling_reduce(
values: &[f64],
window: usize,
edge: EdgePolicy,
missing: MissingDataPolicy,
reduce: impl Fn(&[f64]) -> f64,
) -> Result<Vec<Option<f64>>, TransformError> {
let working = apply_policy(values, missing)?;
let window = window.max(1);
let n = working.len();
let mut out = Vec::with_capacity(n);
for i in 0..n {
let avail = i + 1;
let value = if avail >= window {
Some(reduce(&working[i + 1 - window..=i]))
} else {
match edge {
EdgePolicy::None => None,
EdgePolicy::Partial => Some(reduce(&working[..avail])),
EdgePolicy::Pad => {
let pad_count = window - avail;
let first = working[0];
let mut padded: Vec<f64> = std::iter::repeat(first).take(pad_count).collect();
padded.extend_from_slice(&working[..avail]);
Some(reduce(&padded))
}
}
};
out.push(value);
}
Ok(out)
}
pub fn ema(values: &[f64], alpha: f64, edge: EdgePolicy, missing: MissingDataPolicy) -> Result<Vec<Option<f64>>, TransformError> {
let working = apply_policy(values, missing)?;
let n = working.len();
if n == 0 {
return Ok(Vec::new());
}
let alpha = if alpha.is_nan() { 1.0 } else { alpha.clamp(f64::EPSILON, 1.0) };
let implied_period = ((2.0 / alpha) - 1.0).round().max(1.0) as usize;
let seed = match edge {
EdgePolicy::Pad => mean_kahan(&working[..implied_period.min(n)]),
EdgePolicy::None | EdgePolicy::Partial => working[0],
};
let mut out = Vec::with_capacity(n);
let mut current = seed;
for (i, &v) in working.iter().enumerate() {
if i == 0 {
current = seed;
} else {
current = alpha * v + (1.0 - alpha) * current;
}
let emit = match edge {
EdgePolicy::None => i + 1 >= implied_period,
EdgePolicy::Partial | EdgePolicy::Pad => true,
};
out.push(if emit { Some(current) } else { None });
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rolling_mean_matches_hand_computed_windows() {
let values = [1.0, 2.0, 3.0, 4.0, 5.0];
let out = rolling_mean(&values, 3, EdgePolicy::None, MissingDataPolicy::Skip).unwrap();
assert_eq!(out, vec![None, None, Some(2.0), Some(3.0), Some(4.0)]);
}
#[test]
fn rolling_mean_edge_none_leading_gap_is_exactly_window_minus_one() {
let values = [1.0; 10];
let out = rolling_mean(&values, 4, EdgePolicy::None, MissingDataPolicy::Skip).unwrap();
assert_eq!(out[..3], [None, None, None]);
assert!(out[3..].iter().all(Option::is_some));
}
#[test]
fn rolling_mean_edge_partial_shrinking_window_every_index_defined() {
let values = [2.0, 4.0, 6.0, 8.0];
let out = rolling_mean(&values, 3, EdgePolicy::Partial, MissingDataPolicy::Skip).unwrap();
assert_eq!(out[0], Some(2.0), "avail=1: mean of [2.0]");
assert_eq!(out[1], Some(3.0), "avail=2: mean of [2.0, 4.0]");
assert_eq!(out[2], Some(4.0), "avail=3: full window [2,4,6]");
}
#[test]
fn rolling_mean_edge_pad_biases_toward_the_first_value_more_than_partial() {
let values = [10.0, 100.0];
let partial = rolling_mean(&values, 3, EdgePolicy::Partial, MissingDataPolicy::Skip).unwrap();
let pad = rolling_mean(&values, 3, EdgePolicy::Pad, MissingDataPolicy::Skip).unwrap();
assert_eq!(partial[1], Some(55.0));
assert_eq!(pad[1], Some(40.0));
}
#[test]
fn rolling_mean_window_floors_at_one() {
let values = [3.0, 5.0, 7.0];
let out = rolling_mean(&values, 0, EdgePolicy::None, MissingDataPolicy::Skip).unwrap();
assert_eq!(out, vec![Some(3.0), Some(5.0), Some(7.0)]);
}
#[test]
fn rolling_mean_empty_is_empty() {
assert!(rolling_mean(&[], 3, EdgePolicy::None, MissingDataPolicy::Skip).unwrap().is_empty());
}
#[test]
fn rolling_mean_all_equal_stays_constant() {
let out = rolling_mean(&[5.0; 8], 3, EdgePolicy::Partial, MissingDataPolicy::Skip).unwrap();
assert!(out.iter().all(|v| *v == Some(5.0)));
}
#[test]
fn rolling_mean_propagate_poisons_only_windows_touching_the_non_finite_value() {
let values = [1.0, 2.0, f64::NAN, 4.0, 5.0];
let out = rolling_mean(&values, 2, EdgePolicy::Partial, MissingDataPolicy::Propagate).unwrap();
assert_eq!(out[0], Some(1.0));
assert_eq!(out[1], Some(1.5));
assert!(out[2].unwrap().is_nan(), "window [2, NaN] must be poisoned");
assert!(out[3].unwrap().is_nan(), "window [NaN, 4] must be poisoned");
assert_eq!(out[4], Some(4.5), "window [4, 5] never touched the NaN, must be clean");
}
#[test]
fn rolling_mean_error_rejects_up_front() {
let err = rolling_mean(&[1.0, f64::NAN], 2, EdgePolicy::None, MissingDataPolicy::Error).unwrap_err();
assert_eq!(err, TransformError::NonFinite { index: 1 });
}
#[test]
fn rolling_median_matches_hand_computed_windows() {
let values = [1.0, 5.0, 2.0, 8.0, 3.0];
let out = rolling_median(&values, 3, EdgePolicy::None, MissingDataPolicy::Skip).unwrap();
assert_eq!(out[2], Some(2.0), "median of [1,5,2] sorted [1,2,5] -> 2");
assert_eq!(out[3], Some(5.0), "median of [5,2,8] sorted [2,5,8] -> 5");
}
#[test]
fn rolling_median_is_robust_to_a_single_outlier_unlike_rolling_mean() {
let values = [4.0, 5.0, 6.0, 1000.0, 5.0];
let mean_out = rolling_mean(&values, 3, EdgePolicy::None, MissingDataPolicy::Skip).unwrap();
let median_out = rolling_median(&values, 3, EdgePolicy::None, MissingDataPolicy::Skip).unwrap();
assert!(mean_out[3].unwrap() > 100.0, "mean gets dragged by the outlier");
assert!(median_out[3].unwrap() < 10.0, "median stays near the non-outlier values");
}
#[test]
fn ema_partial_seeds_from_the_first_value_and_recurses() {
let values = [10.0, 20.0, 30.0];
let alpha = 0.5;
let out = ema(&values, alpha, EdgePolicy::Partial, MissingDataPolicy::Skip).unwrap();
assert_eq!(out[0], Some(10.0));
assert_eq!(out[1], Some(15.0), "0.5*20 + 0.5*10");
assert_eq!(out[2], Some(22.5), "0.5*30 + 0.5*15");
}
#[test]
fn ema_none_delays_output_until_the_implied_period() {
let values = [10.0, 20.0, 30.0, 40.0];
let out = ema(&values, 0.5, EdgePolicy::None, MissingDataPolicy::Skip).unwrap();
assert_eq!(out[0], None);
assert_eq!(out[1], None);
assert!(out[2].is_some(), "index 2 is the (implied_period=3)rd sample, must emit");
assert!(out[3].is_some());
}
#[test]
fn ema_pad_seeds_from_the_sma_of_the_implied_period_not_the_raw_first_value() {
let values = [10.0, 100.0, 100.0];
let alpha = 1.0; let pad = ema(&values, alpha, EdgePolicy::Pad, MissingDataPolicy::Skip).unwrap();
let partial = ema(&values, alpha, EdgePolicy::Partial, MissingDataPolicy::Skip).unwrap();
assert_eq!(pad, partial, "an implied period of 1 makes Pad's own SMA seed degenerate to the same single raw value Partial uses");
let alpha2 = 2.0 / 4.0; let pad2 = ema(&values, alpha2, EdgePolicy::Pad, MissingDataPolicy::Skip).unwrap();
let partial2 = ema(&values, alpha2, EdgePolicy::Partial, MissingDataPolicy::Skip).unwrap();
assert_ne!(pad2[0], partial2[0], "a real (>1) implied period must make Pad's own smoothed seed differ from Partial's raw-first-value seed");
}
#[test]
fn ema_empty_is_empty() {
assert!(ema(&[], 0.5, EdgePolicy::Partial, MissingDataPolicy::Skip).unwrap().is_empty());
}
#[test]
fn ema_single_value_is_itself_under_partial_and_pad() {
assert_eq!(ema(&[9.0], 0.3, EdgePolicy::Partial, MissingDataPolicy::Skip).unwrap(), vec![Some(9.0)]);
assert_eq!(ema(&[9.0], 0.3, EdgePolicy::Pad, MissingDataPolicy::Skip).unwrap(), vec![Some(9.0)]);
}
#[test]
fn ema_alpha_is_clamped_to_a_sane_range_never_panics_and_never_produces_nan_output() {
let values = [1.0, 2.0, 3.0];
for bad_alpha in [-5.0, 0.0, 3.7, f64::NAN, f64::INFINITY] {
let out = ema(&values, bad_alpha, EdgePolicy::Partial, MissingDataPolicy::Skip).unwrap();
assert_eq!(out.len(), 3);
assert!(out.iter().all(|v| v.is_some_and(|x| x.is_finite())), "a degenerate alpha must clamp to a finite result, got {out:?}");
}
}
#[test]
fn ema_all_equal_stays_constant() {
let out = ema(&[4.0; 6], 0.4, EdgePolicy::Partial, MissingDataPolicy::Skip).unwrap();
assert!(out.iter().all(|v| *v == Some(4.0)));
}
}