Skip to main content

wickra_core/indicators/
sma.rs

1//! Simple Moving Average.
2
3use crate::error::{Error, Result};
4use crate::traits::Indicator;
5
6/// Simple Moving Average over a fixed window.
7///
8/// Maintains a rolling sum so each update is O(1). Output equals
9/// `sum(last `period` prices) / period` once the window is full; `None` before.
10///
11/// On long-running streams a single-subtract incremental sum can accumulate
12/// rounding error (catastrophic cancellation when values of very different
13/// magnitudes are alternately added and removed). To keep drift bounded, the
14/// running sum is reseeded from the live window every `16 · period` updates —
15/// O(1) amortised cost (`O(period)` work amortised over `O(period)` updates),
16/// zero observable behaviour change on inputs that did not drift to begin
17/// with, and a strict cap on accumulated rounding for streams that did.
18///
19/// # Example
20///
21/// ```
22/// use wickra_core::{Indicator, Sma};
23///
24/// let mut indicator = Sma::new(3).unwrap();
25/// let mut last = None;
26/// for i in 0..80 {
27///     last = indicator.update(100.0 + f64::from(i));
28/// }
29/// assert!(last.is_some());
30/// ```
31#[derive(Debug, Clone)]
32pub struct Sma {
33    period: usize,
34    /// Fixed-capacity ring buffer of the last `period` finite inputs. A flat
35    /// `Box<[f64]>` with a manual write cursor beats `VecDeque` on this hot path:
36    /// sequential storage, branchless wraparound, no per-call bookkeeping.
37    buf: Box<[f64]>,
38    /// Index of the next slot to write — also the oldest element once full.
39    head: usize,
40    /// Number of slots filled, saturating at `period`.
41    count: usize,
42    sum: f64,
43    /// Number of finite updates since the running `sum` was last reseeded from
44    /// the live window. Caps accumulated floating-point drift on long streams.
45    /// See [`RECOMPUTE_EVERY`] below.
46    updates_since_recompute: usize,
47}
48
49/// How often (in finite updates) the incremental sum is reseeded from the live
50/// window. The multiplier `16` is the smallest power of two that keeps the
51/// amortised cost flat under any `period` while still bounding any drift to
52/// roughly `16 · period · ULP · max(|x|)` — sub-picodollar on real-world price
53/// scales.
54const RECOMPUTE_EVERY: usize = 16;
55
56impl Sma {
57    /// Construct a new SMA with the given window length.
58    ///
59    /// # Errors
60    ///
61    /// Returns [`Error::PeriodZero`] if `period == 0`.
62    pub fn new(period: usize) -> Result<Self> {
63        if period == 0 {
64            return Err(Error::PeriodZero);
65        }
66        if period > crate::error::MAX_PERIOD {
67            return Err(Error::InvalidPeriod {
68                message: crate::error::PERIOD_ABOVE_MAX,
69            });
70        }
71        Ok(Self {
72            period,
73            buf: vec![0.0; period].into_boxed_slice(),
74            head: 0,
75            count: 0,
76            sum: 0.0,
77            updates_since_recompute: 0,
78        })
79    }
80
81    /// Configured window length.
82    pub const fn period(&self) -> usize {
83        self.period
84    }
85
86    /// Whether the SMA has taken no input since construction or reset.
87    pub(crate) fn is_fresh(&self) -> bool {
88        self.count == 0 && self.updates_since_recompute == 0
89    }
90
91    /// Current value if available.
92    pub fn value(&self) -> Option<f64> {
93        if self.count == self.period {
94            Some(self.sum / self.period as f64)
95        } else {
96            None
97        }
98    }
99
100    /// Vectorized batch returning one `f64` per input (`NaN` during warmup).
101    ///
102    /// Kept as an inherent method so existing callers need no trait import; it
103    /// allocates the result and fills it through
104    /// [`batch_nan_into`](Indicator::batch_nan_into), which carries the fast path.
105    pub fn batch_nan(&mut self, inputs: &[f64]) -> Vec<f64> {
106        crate::traits::BatchNanExt::batch_nan(self, inputs)
107    }
108}
109
110impl Indicator for Sma {
111    type Input = f64;
112    type Output = f64;
113
114    #[inline]
115    fn update(&mut self, input: f64) -> Option<f64> {
116        if !input.is_finite() {
117            return None;
118        }
119        if self.count == self.period {
120            // Window full: overwrite the oldest slot (at `head`). Each step is a
121            // single f64 add/subtract — O(1) but introduces ~1 ULP of rounding
122            // noise. The periodic reseed below caps the accumulated drift.
123            self.sum -= self.buf[self.head];
124            self.buf[self.head] = input;
125            self.sum += input;
126        } else {
127            self.buf[self.head] = input;
128            self.sum += input;
129            self.count += 1;
130        }
131        // Branchless-ish wraparound, cheaper than `% period`.
132        self.head += 1;
133        if self.head == self.period {
134            self.head = 0;
135        }
136        self.updates_since_recompute += 1;
137        if self.updates_since_recompute >= RECOMPUTE_EVERY * self.period {
138            // Reseed in chronological order (oldest at `head`) so the running sum
139            // tracks a fresh from-scratch mean to the bit on stable inputs.
140            self.sum = self.buf[self.head..]
141                .iter()
142                .chain(&self.buf[..self.head])
143                .copied()
144                .sum();
145            self.updates_since_recompute = 0;
146        }
147        self.value()
148    }
149
150    fn reset(&mut self) {
151        self.head = 0;
152        self.count = 0;
153        self.sum = 0.0;
154        self.updates_since_recompute = 0;
155    }
156
157    #[inline]
158    fn warmup_period(&self) -> usize {
159        self.period
160    }
161
162    #[inline]
163    fn is_ready(&self) -> bool {
164        self.count == self.period
165    }
166
167    #[inline]
168    fn name(&self) -> &'static str {
169        "SMA"
170    }
171
172    /// For a fresh, all-finite slice this inlines `update`'s rolling sum and
173    /// drift-reseed, writing the mean as a bare `f64` (warmup → `NaN`) straight
174    /// into `out`. Same add/subtract order, same reseed cadence, same
175    /// `sum / period` division — so it is *bit-for-bit* equal to replaying
176    /// `update`, including the long-stream drift bound. Any other state, or a
177    /// non-finite element, defers to the exact `update` replay.
178    fn batch_nan_into(&mut self, inputs: &[f64], out: &mut [f64]) {
179        assert_eq!(
180            inputs.len(),
181            out.len(),
182            "batch output length must equal input length"
183        );
184        let p = self.period;
185        if self.count != 0
186            || self.updates_since_recompute != 0
187            || !inputs.iter().all(|x| x.is_finite())
188        {
189            for (slot, &x) in out.iter_mut().zip(inputs) {
190                *slot = self.update(x).unwrap_or(f64::NAN);
191            }
192            return;
193        }
194
195        let p_f64 = p as f64;
196        // Walk the ring one lap at a time and step through it with an iterator
197        // rather than indexing. Indexing put a bounds check in the hot loop,
198        // which under `panic = "unwind"` becomes an unwind edge carrying drop
199        // glue for `out` and blocks vectorisation; the same loop is roughly 40%
200        // faster without it.
201        //
202        // A lap is exactly `period` inputs, which is what makes this equivalent:
203        // the fast path only runs from a fresh state, so `head` is 0 at every
204        // lap boundary, and `RECOMPUTE_EVERY * period` is a whole multiple of
205        // `period`, so the drift reseed can only ever fall on one. At a reseed
206        // `head` is therefore 0 and the chronological order the reseed needs is
207        // simply the buffer in order. Only the final lap can be partial, since
208        // any shorter chunk means the input ran out.
209        let mut rest = inputs;
210        let mut written: &mut [f64] = out;
211        let mut lap = 0_usize;
212        while !rest.is_empty() {
213            let take = rest.len().min(p);
214            let (chunk, tail) = rest.split_at(take);
215            rest = tail;
216            let (lap_out, out_tail) = written.split_at_mut(take);
217            written = out_tail;
218            if lap == 0 {
219                for ((slot, &x), cell) in self.buf.iter_mut().zip(chunk).zip(lap_out.iter_mut()) {
220                    *slot = x;
221                    self.sum += x;
222                    self.count += 1;
223                    *cell = if self.count == p {
224                        self.sum / p_f64
225                    } else {
226                        f64::NAN
227                    };
228                }
229            } else {
230                for ((slot, &x), cell) in self.buf.iter_mut().zip(chunk).zip(lap_out.iter_mut()) {
231                    self.sum -= *slot;
232                    *slot = x;
233                    self.sum += x;
234                    *cell = self.sum / p_f64;
235                }
236            }
237            self.updates_since_recompute += take;
238            if self.updates_since_recompute >= RECOMPUTE_EVERY * p {
239                self.sum = self.buf.iter().copied().sum();
240                self.updates_since_recompute = 0;
241                // `update` reseeds *before* emitting the value for the input
242                // that tripped it, so this lap's last value has to come from
243                // the reseeded sum rather than the incremental one. The reseed
244                // cannot fire before `RECOMPUTE_EVERY` complete laps, so the
245                // window is full and this lap wrote a value for every input.
246                *lap_out
247                    .last_mut()
248                    .expect("a lap writes at least one value before it can reseed") =
249                    self.sum / p_f64;
250            }
251            lap += 1;
252        }
253        self.head = inputs.len() % p;
254    }
255
256    /// SIMD kernel: rolling sums as a prefix scan of `x[i] - x[i - period]`,
257    /// re-anchored on an exact window sum every `16 · period` values like the
258    /// exact path's reseed, each sum scaled by `1 / period`. Agrees with the
259    /// exact batch to within a few units in the last place (a multiply by the
260    /// reciprocal replaces the division, and the running sum is reassociated);
261    /// warmup `NaN`s and length are identical. Afterwards the window holds the
262    /// last `period` inputs and its sum is recomputed exactly, so streaming
263    /// continues from a freshly reseeded state.
264    fn batch_fast_into(&mut self, inputs: &[f64], out: &mut [f64]) {
265        assert_eq!(
266            inputs.len(),
267            out.len(),
268            "batch output length must equal input length"
269        );
270        let p = self.period;
271        let n = inputs.len();
272        if self.count != 0
273            || self.updates_since_recompute != 0
274            || n < p
275            || !crate::fast::in_range(inputs)
276        {
277            self.batch_nan_into(inputs, out);
278            return;
279        }
280        wickra_simd::dispatch(crate::fast::SmaFast {
281            x: inputs,
282            period: p,
283            out,
284            _borrow: std::marker::PhantomData,
285        });
286        for (idx, &x) in inputs.iter().enumerate().skip(n - p) {
287            self.buf[idx % p] = x;
288        }
289        self.head = n % p;
290        self.count = p;
291        self.sum = self.buf[self.head..]
292            .iter()
293            .chain(&self.buf[..self.head])
294            .copied()
295            .sum();
296        self.updates_since_recompute = 0;
297    }
298}
299
300#[cfg(test)]
301mod tests {
302    use super::*;
303
304    /// The same bound applies to every constructor that sizes a buffer from its
305    /// period; SMA allocates eagerly, so it is the sharpest case.
306    #[test]
307    fn rejects_a_period_above_the_maximum() {
308        assert!(matches!(
309            Sma::new(usize::MAX),
310            Err(Error::InvalidPeriod { .. })
311        ));
312        assert!(matches!(
313            Sma::new(crate::error::MAX_PERIOD + 1),
314            Err(Error::InvalidPeriod { .. })
315        ));
316        assert!(Sma::new(20).is_ok());
317    }
318    use crate::traits::BatchExt;
319    use approx::assert_relative_eq;
320    use std::collections::VecDeque;
321
322    #[test]
323    fn new_rejects_zero_period() {
324        assert!(matches!(Sma::new(0), Err(Error::PeriodZero)));
325    }
326
327    /// Cover the const accessor `period` (70-72) and the Indicator-impl
328    /// `warmup_period` (115-117) + `name` (123-125). Existing tests
329    /// inspect SMA output but never query the metadata.
330    #[test]
331    fn accessors_and_metadata() {
332        let sma = Sma::new(20).unwrap();
333        assert_eq!(sma.period(), 20);
334        assert_eq!(sma.warmup_period(), 20);
335        assert_eq!(sma.name(), "SMA");
336    }
337
338    #[test]
339    fn warmup_returns_none() {
340        let mut sma = Sma::new(3).unwrap();
341        assert_eq!(sma.update(1.0), None);
342        assert_eq!(sma.update(2.0), None);
343        assert_eq!(sma.update(3.0), Some(2.0));
344    }
345
346    #[test]
347    fn rolls_window_after_full() {
348        let mut sma = Sma::new(3).unwrap();
349        let out: Vec<_> = [1.0, 2.0, 3.0, 4.0, 5.0]
350            .iter()
351            .map(|p| sma.update(*p))
352            .collect();
353        assert_eq!(out, vec![None, None, Some(2.0), Some(3.0), Some(4.0)]);
354    }
355
356    #[test]
357    fn period_one_is_pass_through() {
358        let mut sma = Sma::new(1).unwrap();
359        assert_eq!(sma.update(5.0), Some(5.0));
360        assert_eq!(sma.update(10.0), Some(10.0));
361    }
362
363    #[test]
364    fn ignores_non_finite_input_but_keeps_state() {
365        let mut sma = Sma::new(3).unwrap();
366        sma.update(1.0);
367        sma.update(2.0);
368        sma.update(3.0);
369        assert_eq!(sma.update(f64::NAN), None);
370        assert_eq!(sma.update(f64::INFINITY), None);
371        // Non-finite inputs were not pushed; window still holds 1,2,3.
372        assert_eq!(sma.update(6.0), Some((2.0 + 3.0 + 6.0) / 3.0));
373    }
374
375    #[test]
376    fn reset_clears_state() {
377        let mut sma = Sma::new(3).unwrap();
378        sma.batch(&[1.0, 2.0, 3.0]);
379        assert!(sma.is_ready());
380        sma.reset();
381        assert!(!sma.is_ready());
382        assert_eq!(sma.update(10.0), None);
383    }
384
385    #[test]
386    fn batch_equals_streaming() {
387        let prices: Vec<f64> = (1..=20).map(f64::from).collect();
388        let mut a = Sma::new(5).unwrap();
389        let batch = a.batch(&prices);
390        let mut b = Sma::new(5).unwrap();
391        let streamed: Vec<_> = prices.iter().map(|p| b.update(*p)).collect();
392        assert_eq!(batch, streamed);
393    }
394
395    #[test]
396    fn known_reference_values() {
397        // SMA(3) of [2, 4, 6, 8, 10] -> [_, _, 4, 6, 8]
398        let mut sma = Sma::new(3).unwrap();
399        let out = sma.batch(&[2.0, 4.0, 6.0, 8.0, 10.0]);
400        assert_eq!(out[2], Some(4.0));
401        assert_eq!(out[3], Some(6.0));
402        assert_eq!(out[4], Some(8.0));
403    }
404
405    #[test]
406    fn constant_series_yields_constant_sma() {
407        let mut sma = Sma::new(5).unwrap();
408        let v = sma.batch(&[7.0; 10]);
409        for x in v.iter().skip(4) {
410            assert_relative_eq!(x.unwrap(), 7.0, epsilon = 1e-12);
411        }
412    }
413
414    /// NaN-aware bit-equality for the `f64`-with-NaN-warmup batch outputs.
415    fn bits_eq(a: &[f64], b: &[f64]) -> bool {
416        a.len() == b.len()
417            && a.iter()
418                .zip(b)
419                .all(|(x, y)| x == y || (x.is_nan() && y.is_nan()))
420    }
421
422    fn sma_replay(period: usize, series: &[f64]) -> Vec<f64> {
423        let mut s = Sma::new(period).unwrap();
424        series
425            .iter()
426            .map(|&x| s.update(x).unwrap_or(f64::NAN))
427            .collect()
428    }
429
430    #[test]
431    fn batch_nan_fast_path_is_bit_identical_with_reseed() {
432        // > 16*period inputs so the drift-reseed branch fires inside batch_nan.
433        let series: Vec<f64> = (0..500)
434            .map(|i| (f64::from(i) * 0.2).sin() * 10.0 + 50.0)
435            .collect();
436        let mut sma = Sma::new(14).unwrap();
437        let got = sma.batch_nan(&series);
438        assert!(bits_eq(&got, &sma_replay(14, &series)));
439        // State left where the replay would: continued updates agree.
440        let mut ref_sma = Sma::new(14).unwrap();
441        for &x in &series {
442            ref_sma.update(x);
443        }
444        assert_eq!(sma.update(42.0), ref_sma.update(42.0));
445    }
446
447    /// Into a caller buffer that already holds values, the fast path must write
448    /// every cell — warmup positions included — exactly as the replay would.
449    #[test]
450    fn batch_nan_into_overwrites_a_dirty_buffer() {
451        let series: Vec<f64> = (0..300).map(|i| f64::from(i % 17) * 1.5 + 3.0).collect();
452        let mut out = vec![123.0; series.len()];
453        Sma::new(9).unwrap().batch_nan_into(&series, &mut out);
454        assert!(bits_eq(&out, &sma_replay(9, &series)));
455    }
456
457    #[test]
458    fn batch_nan_falls_back_on_non_finite() {
459        let series = [1.0, 2.0, f64::NAN, 4.0, 5.0, 6.0];
460        let mut sma = Sma::new(3).unwrap();
461        assert!(bits_eq(&sma.batch_nan(&series), &sma_replay(3, &series)));
462    }
463
464    #[test]
465    fn batch_nan_falls_back_when_not_fresh() {
466        let mut sma = Sma::new(3).unwrap();
467        sma.update(99.0);
468        let series = [1.0, 2.0, 3.0, 4.0];
469        let mut ref_sma = Sma::new(3).unwrap();
470        ref_sma.update(99.0);
471        let want: Vec<f64> = series
472            .iter()
473            .map(|&x| ref_sma.update(x).unwrap_or(f64::NAN))
474            .collect();
475        assert!(bits_eq(&sma.batch_nan(&series), &want));
476    }
477
478    #[test]
479    fn batch_nan_sub_period_slice_is_all_nan() {
480        let series = [1.0, 2.0, 3.0];
481        let mut sma = Sma::new(10).unwrap();
482        let got = sma.batch_nan(&series);
483        assert!(bits_eq(&got, &sma_replay(10, &series)));
484        assert!(got.iter().all(|x| x.is_nan()));
485    }
486
487    proptest::proptest! {
488        #![proptest_config(proptest::test_runner::Config::with_cases(64))]
489        #[test]
490        fn sma_matches_naive_definition(
491            period in 1usize..20,
492            prices in proptest::collection::vec(-1000.0_f64..1000.0, 0..200),
493        ) {
494            let mut sma = Sma::new(period).unwrap();
495            let stream: Vec<_> = prices.iter().map(|p| sma.update(*p)).collect();
496            for (i, got) in stream.iter().enumerate() {
497                if i + 1 < period {
498                    proptest::prop_assert!(got.is_none());
499                } else {
500                    let window = &prices[i + 1 - period..=i];
501                    let expected = window.iter().sum::<f64>() / period as f64;
502                    let actual = got.expect("ready");
503                    proptest::prop_assert!(
504                        (actual - expected).abs() < 1e-9,
505                        "i={i} actual={actual} expected={expected}"
506                    );
507                }
508            }
509        }
510    }
511
512    /// Long-running stability check. Runs more updates than `RECOMPUTE_EVERY *
513    /// period` so the periodic reseed must fire several times, then asserts
514    /// that the reported SMA still equals a fresh from-scratch mean over the
515    /// live window to within tight floating-point tolerance. Inputs swing
516    /// between two magnitudes (`1e9` and `1.0`) — a pattern designed to
517    /// expose catastrophic cancellation in a naive single-subtract sum.
518    #[test]
519    fn long_stream_drift_stays_bounded() {
520        let period = 20;
521        let mut sma = Sma::new(period).unwrap();
522        let mut window: VecDeque<f64> = VecDeque::with_capacity(period);
523        // `RECOMPUTE_EVERY * period * 5` updates → recompute fires 5+ times.
524        let n_updates = 16 * period * 5;
525        for i in 0..n_updates {
526            let v = if i % 2 == 0 { 1e9 } else { 1.0 };
527            sma.update(v);
528            if window.len() == period {
529                window.pop_front();
530            }
531            window.push_back(v);
532        }
533        let from_scratch: f64 = window.iter().sum::<f64>() / period as f64;
534        let got = sma.value().expect("warmed up");
535        assert!(
536            (got - from_scratch).abs() < 1e-6,
537            "SMA drift exceeds 1e-6 over {n_updates} updates: got={got}, scratch={from_scratch}"
538        );
539    }
540}