#![allow(clippy::doc_markdown)]
use std::f64::consts::PI;
use crate::error::{Error, Result};
use crate::traits::Indicator;
#[derive(Debug, Clone)]
pub struct SuperSmoother {
period: usize,
c1: f64,
c2: f64,
c3: f64,
prev_input: Option<f64>,
prev_output_1: Option<f64>,
prev_output_2: Option<f64>,
count: usize,
}
impl SuperSmoother {
pub fn new(period: usize) -> Result<Self> {
if period == 0 {
return Err(Error::PeriodZero);
}
if period > crate::error::MAX_PERIOD {
return Err(Error::InvalidPeriod {
message: crate::error::PERIOD_ABOVE_MAX,
});
}
Ok(Self::with_critical_period(period, period as f64))
}
pub(crate) fn with_critical_period(period: usize, critical: f64) -> Self {
let arg = std::f64::consts::SQRT_2 * PI / critical;
let a1 = (-arg).exp();
let b1 = 2.0 * a1 * arg.cos();
let c2 = b1;
let c3 = -a1 * a1;
let c1 = 1.0 - c2 - c3;
Self {
period,
c1,
c2,
c3,
prev_input: None,
prev_output_1: None,
prev_output_2: None,
count: 0,
}
}
pub const fn period(&self) -> usize {
self.period
}
pub const fn coefficients(&self) -> (f64, f64, f64) {
(self.c1, self.c2, self.c3)
}
pub const fn value(&self) -> Option<f64> {
self.prev_output_1
}
}
impl Indicator for SuperSmoother {
type Input = f64;
type Output = f64;
#[inline]
fn update(&mut self, input: f64) -> Option<f64> {
if !input.is_finite() {
return None;
}
self.count += 1;
let output = match (self.prev_input, self.prev_output_1, self.prev_output_2) {
(Some(p_in), Some(y1), Some(y2)) => {
let avg = f64::midpoint(input, p_in);
self.c1 * avg + self.c2 * y1 + self.c3 * y2
}
_ => input,
};
self.prev_output_2 = self.prev_output_1;
self.prev_output_1 = Some(output);
self.prev_input = Some(input);
Some(output)
}
fn reset(&mut self) {
self.prev_input = None;
self.prev_output_1 = None;
self.prev_output_2 = None;
self.count = 0;
}
#[inline]
fn warmup_period(&self) -> usize {
1
}
#[inline]
fn is_ready(&self) -> bool {
self.prev_output_1.is_some()
}
#[inline]
fn name(&self) -> &'static str {
"SuperSmoother"
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::traits::BatchExt;
use approx::assert_relative_eq;
#[test]
fn new_rejects_zero_period() {
assert!(matches!(SuperSmoother::new(0), Err(Error::PeriodZero)));
}
#[test]
fn accessors_and_metadata() {
let mut ss = SuperSmoother::new(10).unwrap();
assert_eq!(ss.period(), 10);
assert_eq!(ss.name(), "SuperSmoother");
assert_eq!(ss.warmup_period(), 1);
let (c1, c2, c3) = ss.coefficients();
assert_relative_eq!(c1 + c2 + c3, 1.0, epsilon = 1e-12);
assert!(ss.value().is_none());
ss.update(42.0);
assert!(ss.value().is_some());
assert!(ss.is_ready());
}
#[test]
fn first_output_equals_input_then_filters() {
let mut ss = SuperSmoother::new(10).unwrap();
assert_eq!(ss.update(100.0), Some(100.0));
assert_eq!(ss.update(101.0), Some(101.0));
let third = ss.update(102.0).unwrap();
assert!((third - 102.0).abs() < 5.0);
}
#[test]
fn constant_series_converges_to_constant() {
let mut ss = SuperSmoother::new(20).unwrap();
let out = ss.batch(&[50.0_f64; 200]);
for x in out.iter().skip(50).flatten() {
assert_relative_eq!(*x, 50.0, epsilon = 1e-9);
}
}
#[test]
fn batch_equals_streaming() {
let prices: Vec<f64> = (0..120)
.map(|i| 100.0 + (f64::from(i) * 0.2).sin() * 5.0)
.collect();
let mut a = SuperSmoother::new(15).unwrap();
let mut b = SuperSmoother::new(15).unwrap();
let batch = a.batch(&prices);
let streamed: Vec<_> = prices.iter().map(|p| b.update(*p)).collect();
assert_eq!(batch, streamed);
}
#[test]
fn ignores_non_finite_input() {
let mut ss = SuperSmoother::new(10).unwrap();
ss.batch(&(1..=20).map(f64::from).collect::<Vec<_>>());
let before = ss.value();
assert!(before.is_some());
assert_eq!(ss.update(f64::NAN), None);
assert_eq!(ss.update(f64::INFINITY), None);
}
#[test]
fn reset_clears_state() {
let mut ss = SuperSmoother::new(10).unwrap();
ss.batch(&(1..=40).map(f64::from).collect::<Vec<_>>());
assert!(ss.is_ready());
ss.reset();
assert!(!ss.is_ready());
assert_eq!(ss.update(50.0), Some(50.0));
}
use crate::traits::BatchNanExt;
#[test]
fn new_rejects_period_above_max() {
assert!(matches!(
SuperSmoother::new(crate::error::MAX_PERIOD + 1),
Err(Error::InvalidPeriod { .. })
));
}
#[test]
fn first_value_lands_exactly_at_warmup() {
let mut ss = SuperSmoother::new(10).unwrap();
let out = ss.batch(&[5.0, 6.0, 7.0]);
assert_eq!(ss.warmup_period(), 1);
assert_eq!(out[0], Some(5.0));
}
#[test]
fn reset_replays_identically() {
let prices: Vec<f64> = (0..120)
.map(|i| 100.0 + (f64::from(i) * 0.2).sin() * 5.0)
.collect();
let fresh = SuperSmoother::new(12).unwrap().batch(&prices);
let mut ss = SuperSmoother::new(12).unwrap();
let first = ss.batch(&prices);
ss.reset();
let second = ss.batch(&prices);
assert_eq!(first, fresh);
assert_eq!(second, fresh);
}
#[test]
fn batch_nan_paths_match_streaming_bitwise() {
let prices: Vec<f64> = (0..120)
.map(|i| 100.0 + (f64::from(i) * 0.2).sin() * 5.0)
.collect();
let mut out = vec![0.0; prices.len()];
SuperSmoother::new(12)
.unwrap()
.batch_nan_into(&prices, &mut out);
let nan = SuperSmoother::new(12).unwrap().batch_nan(&prices);
let fast = SuperSmoother::new(12).unwrap().batch_fast(&prices);
let mut stream = SuperSmoother::new(12).unwrap();
let expected: Vec<u64> = prices
.iter()
.map(|&p| stream.update(p).unwrap_or(f64::NAN).to_bits())
.collect();
assert!(out.iter().zip(&expected).all(|(v, e)| v.to_bits() == *e));
assert!(nan.iter().zip(&expected).all(|(v, e)| v.to_bits() == *e));
assert!(fast.iter().zip(&expected).all(|(v, e)| v.to_bits() == *e));
}
#[test]
fn with_critical_period_hand_computed() {
let ss = SuperSmoother::with_critical_period(7, 2.0 * std::f64::consts::SQRT_2);
assert_eq!(ss.period(), 7);
let (c1, c2, c3) = ss.coefficients();
assert_relative_eq!(c2, 0.0, epsilon = 1e-12);
assert_relative_eq!(c3, -0.043_213_918_264, epsilon = 1e-12);
assert_relative_eq!(c1, 1.043_213_918_264, epsilon = 1e-12);
let a = SuperSmoother::new(9).unwrap().coefficients();
let b = SuperSmoother::with_critical_period(9, 9.0).coefficients();
assert_eq!(
(a.0.to_bits(), a.1.to_bits(), a.2.to_bits()),
(b.0.to_bits(), b.1.to_bits(), b.2.to_bits())
);
let half = SuperSmoother::with_critical_period(9, 4.5).coefficients();
assert!((half.0 - a.0).abs() > 1e-3);
}
#[test]
fn third_output_is_the_recursion_hand_computed() {
let mut ss = SuperSmoother::with_critical_period(7, 2.0 * std::f64::consts::SQRT_2);
let (c1, c2, c3) = ss.coefficients();
let out = ss.batch(&[100.0, 101.0, 102.0]);
assert_eq!(out[0], Some(100.0));
assert_eq!(out[1], Some(101.0));
let expected = c1 * 101.5 + c2 * 101.0 + c3 * 100.0;
assert_eq!(out[2].unwrap().to_bits(), expected.to_bits());
assert_relative_eq!(out[2].unwrap(), 101.564_820_877, epsilon = 1e-6);
}
}