1use crate::error::{Error, Result};
4use crate::traits::Indicator;
5
6#[derive(Debug, Clone)]
32pub struct Sma {
33 period: usize,
34 buf: Box<[f64]>,
38 head: usize,
40 count: usize,
42 sum: f64,
43 updates_since_recompute: usize,
47}
48
49const RECOMPUTE_EVERY: usize = 16;
55
56impl Sma {
57 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 pub const fn period(&self) -> usize {
83 self.period
84 }
85
86 pub(crate) fn is_fresh(&self) -> bool {
88 self.count == 0 && self.updates_since_recompute == 0
89 }
90
91 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 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 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 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 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 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 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 *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 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 #[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 #[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 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 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 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 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 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 #[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 #[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 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}