Skip to main content

wickra_core/indicators/
super_smoother.rs

1//! Ehlers SuperSmoother filter.
2#![allow(clippy::doc_markdown)]
3
4use std::f64::consts::PI;
5
6use crate::error::{Error, Result};
7use crate::traits::Indicator;
8
9/// Ehlers' 2-pole Butterworth-style "SuperSmoother" lowpass filter.
10///
11/// From John Ehlers' *Cycle Analytics for Traders* (2013, ch. 3). For a given
12/// critical period `period`, the filter coefficients are:
13///
14/// ```text
15/// a1 = exp(-sqrt(2) * pi / period)
16/// b1 = 2 * a1 * cos(sqrt(2) * pi / period)
17/// c2 = b1
18/// c3 = -a1 * a1
19/// c1 = 1 - c2 - c3
20/// y[t] = c1 * (x[t] + x[t-1]) / 2 + c2 * y[t-1] + c3 * y[t-2]
21/// ```
22///
23/// The implementation needs two prior inputs and two prior outputs to begin
24/// running; until then it returns the input itself (a common Ehlers initial
25/// condition), which lets downstream filters warm up without long delays.
26///
27/// # Example
28///
29/// ```
30/// use wickra_core::{Indicator, SuperSmoother};
31///
32/// let mut ss = SuperSmoother::new(10).unwrap();
33/// let mut last = None;
34/// for i in 0..40 {
35///     last = ss.update(100.0 + f64::from(i));
36/// }
37/// assert!(last.is_some());
38/// ```
39#[derive(Debug, Clone)]
40pub struct SuperSmoother {
41    period: usize,
42    c1: f64,
43    c2: f64,
44    c3: f64,
45    prev_input: Option<f64>,
46    prev_output_1: Option<f64>,
47    prev_output_2: Option<f64>,
48    count: usize,
49}
50
51impl SuperSmoother {
52    /// Construct a new SuperSmoother with the given critical period.
53    ///
54    /// # Errors
55    ///
56    /// Returns [`Error::PeriodZero`] if `period == 0`.
57    pub fn new(period: usize) -> Result<Self> {
58        if period == 0 {
59            return Err(Error::PeriodZero);
60        }
61        if period > crate::error::MAX_PERIOD {
62            return Err(Error::InvalidPeriod {
63                message: crate::error::PERIOD_ABOVE_MAX,
64            });
65        }
66        Ok(Self::with_critical_period(period, period as f64))
67    }
68
69    /// Build a SuperSmoother whose coefficients use a fractional `critical`
70    /// period, reporting `period` from [`period`](Self::period). Ehlers' Reflex
71    /// and Trendflex smooth with half their lookback (`0.5 ยท Length`), which is
72    /// fractional for an odd length. The caller validates `period`.
73    pub(crate) fn with_critical_period(period: usize, critical: f64) -> Self {
74        let arg = std::f64::consts::SQRT_2 * PI / critical;
75        let a1 = (-arg).exp();
76        let b1 = 2.0 * a1 * arg.cos();
77        let c2 = b1;
78        let c3 = -a1 * a1;
79        let c1 = 1.0 - c2 - c3;
80        Self {
81            period,
82            c1,
83            c2,
84            c3,
85            prev_input: None,
86            prev_output_1: None,
87            prev_output_2: None,
88            count: 0,
89        }
90    }
91
92    /// Configured period.
93    pub const fn period(&self) -> usize {
94        self.period
95    }
96
97    /// Filter coefficients `(c1, c2, c3)`.
98    pub const fn coefficients(&self) -> (f64, f64, f64) {
99        (self.c1, self.c2, self.c3)
100    }
101
102    /// Current value if available.
103    pub const fn value(&self) -> Option<f64> {
104        self.prev_output_1
105    }
106}
107
108impl Indicator for SuperSmoother {
109    type Input = f64;
110    type Output = f64;
111
112    #[inline]
113    fn update(&mut self, input: f64) -> Option<f64> {
114        if !input.is_finite() {
115            return None;
116        }
117        self.count += 1;
118        let output = match (self.prev_input, self.prev_output_1, self.prev_output_2) {
119            (Some(p_in), Some(y1), Some(y2)) => {
120                let avg = f64::midpoint(input, p_in);
121                self.c1 * avg + self.c2 * y1 + self.c3 * y2
122            }
123            _ => input,
124        };
125        self.prev_output_2 = self.prev_output_1;
126        self.prev_output_1 = Some(output);
127        self.prev_input = Some(input);
128        Some(output)
129    }
130
131    fn reset(&mut self) {
132        self.prev_input = None;
133        self.prev_output_1 = None;
134        self.prev_output_2 = None;
135        self.count = 0;
136    }
137
138    #[inline]
139    fn warmup_period(&self) -> usize {
140        1
141    }
142
143    #[inline]
144    fn is_ready(&self) -> bool {
145        self.prev_output_1.is_some()
146    }
147
148    #[inline]
149    fn name(&self) -> &'static str {
150        "SuperSmoother"
151    }
152}
153
154#[cfg(test)]
155mod tests {
156    use super::*;
157    use crate::traits::BatchExt;
158    use approx::assert_relative_eq;
159
160    #[test]
161    fn new_rejects_zero_period() {
162        assert!(matches!(SuperSmoother::new(0), Err(Error::PeriodZero)));
163    }
164
165    #[test]
166    fn accessors_and_metadata() {
167        let mut ss = SuperSmoother::new(10).unwrap();
168        assert_eq!(ss.period(), 10);
169        assert_eq!(ss.name(), "SuperSmoother");
170        assert_eq!(ss.warmup_period(), 1);
171        let (c1, c2, c3) = ss.coefficients();
172        // Coefficients sum to 1 by construction (steady-state gain == 1).
173        assert_relative_eq!(c1 + c2 + c3, 1.0, epsilon = 1e-12);
174        assert!(ss.value().is_none());
175        ss.update(42.0);
176        assert!(ss.value().is_some());
177        assert!(ss.is_ready());
178    }
179
180    #[test]
181    fn first_output_equals_input_then_filters() {
182        let mut ss = SuperSmoother::new(10).unwrap();
183        // Initial condition: first two outputs equal their inputs.
184        assert_eq!(ss.update(100.0), Some(100.0));
185        assert_eq!(ss.update(101.0), Some(101.0));
186        let third = ss.update(102.0).unwrap();
187        // From step 3 onward, the recursive filter activates and the result
188        // is no longer the raw input.
189        assert!((third - 102.0).abs() < 5.0);
190    }
191
192    #[test]
193    fn constant_series_converges_to_constant() {
194        // Steady-state gain is 1 (c1 + c2 + c3 = 1), so a flat input yields a
195        // flat output after warmup.
196        let mut ss = SuperSmoother::new(20).unwrap();
197        let out = ss.batch(&[50.0_f64; 200]);
198        for x in out.iter().skip(50).flatten() {
199            assert_relative_eq!(*x, 50.0, epsilon = 1e-9);
200        }
201    }
202
203    #[test]
204    fn batch_equals_streaming() {
205        let prices: Vec<f64> = (0..120)
206            .map(|i| 100.0 + (f64::from(i) * 0.2).sin() * 5.0)
207            .collect();
208        let mut a = SuperSmoother::new(15).unwrap();
209        let mut b = SuperSmoother::new(15).unwrap();
210        let batch = a.batch(&prices);
211        let streamed: Vec<_> = prices.iter().map(|p| b.update(*p)).collect();
212        assert_eq!(batch, streamed);
213    }
214
215    #[test]
216    fn ignores_non_finite_input() {
217        let mut ss = SuperSmoother::new(10).unwrap();
218        ss.batch(&(1..=20).map(f64::from).collect::<Vec<_>>());
219        let before = ss.value();
220        assert!(before.is_some());
221        assert_eq!(ss.update(f64::NAN), None);
222        assert_eq!(ss.update(f64::INFINITY), None);
223    }
224
225    #[test]
226    fn reset_clears_state() {
227        let mut ss = SuperSmoother::new(10).unwrap();
228        ss.batch(&(1..=40).map(f64::from).collect::<Vec<_>>());
229        assert!(ss.is_ready());
230        ss.reset();
231        assert!(!ss.is_ready());
232        assert_eq!(ss.update(50.0), Some(50.0));
233    }
234
235    use crate::traits::BatchNanExt;
236
237    #[test]
238    fn new_rejects_period_above_max() {
239        assert!(matches!(
240            SuperSmoother::new(crate::error::MAX_PERIOD + 1),
241            Err(Error::InvalidPeriod { .. })
242        ));
243    }
244
245    #[test]
246    fn first_value_lands_exactly_at_warmup() {
247        let mut ss = SuperSmoother::new(10).unwrap();
248        let out = ss.batch(&[5.0, 6.0, 7.0]);
249        assert_eq!(ss.warmup_period(), 1);
250        assert_eq!(out[0], Some(5.0));
251    }
252
253    #[test]
254    fn reset_replays_identically() {
255        let prices: Vec<f64> = (0..120)
256            .map(|i| 100.0 + (f64::from(i) * 0.2).sin() * 5.0)
257            .collect();
258        let fresh = SuperSmoother::new(12).unwrap().batch(&prices);
259        let mut ss = SuperSmoother::new(12).unwrap();
260        let first = ss.batch(&prices);
261        ss.reset();
262        let second = ss.batch(&prices);
263        assert_eq!(first, fresh);
264        assert_eq!(second, fresh);
265    }
266
267    #[test]
268    fn batch_nan_paths_match_streaming_bitwise() {
269        let prices: Vec<f64> = (0..120)
270            .map(|i| 100.0 + (f64::from(i) * 0.2).sin() * 5.0)
271            .collect();
272        let mut out = vec![0.0; prices.len()];
273        SuperSmoother::new(12)
274            .unwrap()
275            .batch_nan_into(&prices, &mut out);
276        let nan = SuperSmoother::new(12).unwrap().batch_nan(&prices);
277        let fast = SuperSmoother::new(12).unwrap().batch_fast(&prices);
278        let mut stream = SuperSmoother::new(12).unwrap();
279        let expected: Vec<u64> = prices
280            .iter()
281            .map(|&p| stream.update(p).unwrap_or(f64::NAN).to_bits())
282            .collect();
283        assert!(out.iter().zip(&expected).all(|(v, e)| v.to_bits() == *e));
284        assert!(nan.iter().zip(&expected).all(|(v, e)| v.to_bits() == *e));
285        assert!(fast.iter().zip(&expected).all(|(v, e)| v.to_bits() == *e));
286    }
287
288    #[test]
289    fn with_critical_period_hand_computed() {
290        // critical = 2*sqrt(2) makes arg = sqrt(2)*pi / (2*sqrt(2)) = pi/2, so
291        // cos(arg) ~ 0 and b1 = c2 ~ 0, a1 = exp(-pi/2) = 0.207_879_576,
292        // c3 = -a1^2 = -0.043_213_918, c1 = 1 - c2 - c3 = 1.043_213_918.
293        let ss = SuperSmoother::with_critical_period(7, 2.0 * std::f64::consts::SQRT_2);
294        assert_eq!(ss.period(), 7);
295        let (c1, c2, c3) = ss.coefficients();
296        assert_relative_eq!(c2, 0.0, epsilon = 1e-12);
297        assert_relative_eq!(c3, -0.043_213_918_264, epsilon = 1e-12);
298        assert_relative_eq!(c1, 1.043_213_918_264, epsilon = 1e-12);
299        // `new(period)` is `with_critical_period(period, period)`.
300        let a = SuperSmoother::new(9).unwrap().coefficients();
301        let b = SuperSmoother::with_critical_period(9, 9.0).coefficients();
302        assert_eq!(
303            (a.0.to_bits(), a.1.to_bits(), a.2.to_bits()),
304            (b.0.to_bits(), b.1.to_bits(), b.2.to_bits())
305        );
306        // A fractional critical period yields different coefficients.
307        let half = SuperSmoother::with_critical_period(9, 4.5).coefficients();
308        assert!((half.0 - a.0).abs() > 1e-3);
309    }
310
311    #[test]
312    fn third_output_is_the_recursion_hand_computed() {
313        // Inputs 100, 101, 102: y0 = 100, y1 = 101 (pass-through seed), then
314        // y2 = c1 * (102 + 101)/2 + c2 * 101 + c3 * 100.
315        let mut ss = SuperSmoother::with_critical_period(7, 2.0 * std::f64::consts::SQRT_2);
316        let (c1, c2, c3) = ss.coefficients();
317        let out = ss.batch(&[100.0, 101.0, 102.0]);
318        assert_eq!(out[0], Some(100.0));
319        assert_eq!(out[1], Some(101.0));
320        let expected = c1 * 101.5 + c2 * 101.0 + c3 * 100.0;
321        assert_eq!(out[2].unwrap().to_bits(), expected.to_bits());
322        // Numerically: 1.043_213_918 * 101.5 - 0.043_213_918 * 100 = 101.564_820_88.
323        assert_relative_eq!(out[2].unwrap(), 101.564_820_877, epsilon = 1e-6);
324    }
325}