1use crate::error::{Error, Result};
4use crate::ohlcv::Candle;
5use crate::traits::Indicator;
6
7#[derive(Debug, Clone)]
29pub struct Atr {
30 period: usize,
31 n_minus_1: f64,
33 inv_period: f64,
36 prev_close: Option<f64>,
37 seed_sum: f64,
42 seed_count: usize,
44 avg: f64,
47 seeded: bool,
48}
49
50impl Atr {
51 pub fn new(period: usize) -> Result<Self> {
57 if period == 0 {
58 return Err(Error::PeriodZero);
59 }
60 if period > crate::error::MAX_PERIOD {
61 return Err(Error::InvalidPeriod {
62 message: crate::error::PERIOD_ABOVE_MAX,
63 });
64 }
65 Ok(Self {
66 period,
67 n_minus_1: (period - 1) as f64,
68 inv_period: 1.0 / period as f64,
69 prev_close: None,
70 seed_sum: -0.0,
71 seed_count: 0,
72 avg: 0.0,
73 seeded: false,
74 })
75 }
76
77 pub const fn period(&self) -> usize {
79 self.period
80 }
81
82 pub const fn value(&self) -> Option<f64> {
84 if self.seeded {
85 Some(self.avg)
86 } else {
87 None
88 }
89 }
90
91 pub fn batch_atr(&mut self, high: &[f64], low: &[f64], close: &[f64]) -> Vec<f64> {
99 let mut out = vec![0.0; high.len()];
100 self.batch_atr_into(high, low, close, &mut out);
101 out
102 }
103
104 pub fn batch_atr_into(&mut self, high: &[f64], low: &[f64], close: &[f64], out: &mut [f64]) {
119 let n = high.len();
120 assert!(
121 low.len() == n && close.len() == n && out.len() == n,
122 "high, low, close and the output must be equal length"
123 );
124 let p = self.period;
125 if self.seeded || self.seed_count != 0 || self.prev_close.is_some() || n < p {
126 for (i, slot) in out.iter_mut().enumerate() {
127 let candle = Candle::new_unchecked(close[i], high[i], low[i], close[i], 0.0, 0);
128 *slot = self.update(candle).unwrap_or(f64::NAN);
129 }
130 return;
131 }
132
133 out[..p - 1].fill(f64::NAN);
135 let mut prev_close = close[0];
137 let mut sum_tr = -0.0 + (high[0] - low[0]);
138 for i in 1..p {
139 let (h, l) = (high[i], low[i]);
140 let tr = (h - l)
141 .max((h - prev_close).abs())
142 .max((l - prev_close).abs());
143 prev_close = close[i];
144 sum_tr += tr;
145 }
146 let avg = sum_tr / p as f64;
147 out[p - 1] = avg;
148 let (prev_close, avg) = wickra_simd::dispatch(AtrTail {
150 high: &high[p..],
151 low: &low[p..],
152 close: &close[p..],
153 out: &mut out[p..],
154 state: (prev_close, avg),
155 n_minus_1: self.n_minus_1,
156 inv_period: self.inv_period,
157 });
158
159 self.prev_close = Some(prev_close);
161 self.seed_sum = sum_tr;
162 self.seed_count = p;
163 self.avg = avg;
164 self.seeded = true;
165 }
166
167 pub fn batch_atr_fast_into(
180 &mut self,
181 high: &[f64],
182 low: &[f64],
183 close: &[f64],
184 out: &mut [f64],
185 ) {
186 let n = high.len();
187 assert!(
188 low.len() == n && close.len() == n && out.len() == n,
189 "high, low, close and the output must be equal length"
190 );
191 let p = self.period;
192 if self.seeded
193 || self.seed_count != 0
194 || self.prev_close.is_some()
195 || n < p
196 || !crate::fast::in_range(high)
197 || !crate::fast::in_range(low)
198 || !crate::fast::in_range(close)
199 {
200 self.batch_atr_into(high, low, close, out);
201 return;
202 }
203 out[..p - 1].fill(f64::NAN);
204 let mut prev_close = close[0];
205 let mut sum_tr = -0.0 + (high[0] - low[0]);
206 for i in 1..p {
207 let (h, l) = (high[i], low[i]);
208 let tr = (h - l)
209 .max((h - prev_close).abs())
210 .max((l - prev_close).abs());
211 prev_close = close[i];
212 sum_tr += tr;
213 }
214 let seed = sum_tr / p as f64;
215 out[p - 1] = seed;
216 let avg = wickra_simd::dispatch(crate::fast::AtrFast {
217 high: &high[p..],
218 low: &low[p..],
219 prev_close: &close[p - 1..n - 1],
220 seed,
221 n_minus_1: self.n_minus_1,
222 inv_period: self.inv_period,
223 out: &mut out[p..],
224 _borrow: std::marker::PhantomData,
225 });
226 self.prev_close = Some(close[n - 1]);
227 self.seed_sum = sum_tr;
228 self.seed_count = p;
229 self.avg = avg;
230 self.seeded = true;
231 }
232
233 pub fn batch_atr_fast(&mut self, high: &[f64], low: &[f64], close: &[f64]) -> Vec<f64> {
235 let mut out = vec![0.0; high.len()];
236 self.batch_atr_fast_into(high, low, close, &mut out);
237 out
238 }
239}
240
241struct AtrTail<'a> {
245 high: &'a [f64],
246 low: &'a [f64],
247 close: &'a [f64],
248 out: &'a mut [f64],
249 state: (f64, f64),
250 n_minus_1: f64,
251 inv_period: f64,
252}
253
254#[allow(clippy::inline_always)]
257impl wickra_simd::Kernel for AtrTail<'_> {
258 type Output = (f64, f64);
259
260 #[inline(always)]
261 fn run<S: wickra_simd::Simd>(self, _simd: S) -> (f64, f64) {
262 let (mut prev_close, mut avg) = self.state;
263 let (n_minus_1, inv_period) = (self.n_minus_1, self.inv_period);
264 for (((slot, &h), &l), &c) in self
265 .out
266 .iter_mut()
267 .zip(self.high)
268 .zip(self.low)
269 .zip(self.close)
270 {
271 let tr = (h - l)
272 .max((h - prev_close).abs())
273 .max((l - prev_close).abs());
274 prev_close = c;
275 avg = avg.mul_add(n_minus_1, tr) * inv_period;
276 *slot = avg;
277 }
278 (prev_close, avg)
279 }
280}
281
282impl Indicator for Atr {
283 type Input = Candle;
284 type Output = f64;
285
286 #[inline]
287 fn update(&mut self, candle: Candle) -> Option<f64> {
288 let tr = candle.true_range(self.prev_close);
289 self.prev_close = Some(candle.close);
290
291 if self.seeded {
292 let new_avg = self.avg.mul_add(self.n_minus_1, tr) * self.inv_period;
294 self.avg = new_avg;
295 return Some(new_avg);
296 }
297
298 self.seed_sum += tr;
299 self.seed_count += 1;
300 if self.seed_count == self.period {
301 let seed = self.seed_sum / self.period as f64;
302 self.avg = seed;
303 self.seeded = true;
304 return Some(seed);
305 }
306 None
307 }
308
309 fn reset(&mut self) {
310 self.prev_close = None;
311 self.seed_sum = -0.0;
312 self.seed_count = 0;
313 self.avg = 0.0;
314 self.seeded = false;
315 }
316
317 #[inline]
318 fn warmup_period(&self) -> usize {
319 self.period
320 }
321
322 #[inline]
323 fn is_ready(&self) -> bool {
324 self.seeded
325 }
326
327 #[inline]
328 fn name(&self) -> &'static str {
329 "ATR"
330 }
331}
332
333#[cfg(test)]
334mod tests {
335 use super::*;
336 use crate::traits::BatchExt;
337 use approx::assert_relative_eq;
338
339 fn c(h: f64, l: f64, cl: f64) -> Candle {
340 Candle::new(cl, h, l, cl, 1.0, 0).unwrap()
342 }
343
344 fn atr_naive(hlc: &[(f64, f64, f64)], period: usize) -> Vec<Option<f64>> {
346 let n = period as f64;
347 let mut out = Vec::with_capacity(hlc.len());
348 let mut trs: Vec<f64> = Vec::new();
349 let mut avg: Option<f64> = None;
350 let mut prev_close: Option<f64> = None;
351 for &(h, l, cl) in hlc {
352 let tr = match prev_close {
353 None => h - l,
354 Some(pc) => (h - l).max((h - pc).abs()).max((l - pc).abs()),
355 };
356 prev_close = Some(cl);
357 if let Some(a) = avg {
358 let na = (a * (n - 1.0) + tr) / n;
359 avg = Some(na);
360 out.push(Some(na));
361 } else {
362 trs.push(tr);
363 if trs.len() == period {
364 avg = Some(trs.iter().sum::<f64>() / n);
365 out.push(avg);
366 } else {
367 out.push(None);
368 }
369 }
370 }
371 out
372 }
373
374 #[test]
375 fn rejects_zero_period() {
376 assert!(matches!(Atr::new(0), Err(Error::PeriodZero)));
377 }
378
379 #[test]
383 fn accessors_and_metadata() {
384 let mut atr = Atr::new(14).unwrap();
385 assert_eq!(atr.period(), 14);
386 assert_eq!(atr.name(), "ATR");
387 assert_eq!(atr.value(), None);
388 for _ in 0..14 {
389 atr.update(c(11.0, 9.0, 10.0));
390 }
391 assert!(atr.value().is_some());
392 }
393
394 #[test]
395 fn warmup_emits_on_period_th_candle() {
396 let candles = vec![
397 c(2.0, 1.0, 1.5),
398 c(3.0, 2.0, 2.5),
399 c(4.0, 3.0, 3.5),
400 c(5.0, 4.0, 4.5),
401 c(6.0, 5.0, 5.5),
402 ];
403 let mut atr = Atr::new(3).unwrap();
404 let out = atr.batch(&candles);
405 assert!(out[0].is_none());
406 assert!(out[1].is_none());
407 assert!(out[2].is_some());
408 assert!(out[3].is_some());
409 }
410
411 #[test]
412 fn constant_range_yields_constant_atr() {
413 let candles: Vec<Candle> = (0..30).map(|_| c(11.0, 9.0, 10.0)).collect();
415 let mut atr = Atr::new(14).unwrap();
416 let out = atr.batch(&candles);
417 for v in out.iter().skip(13).flatten() {
418 assert_relative_eq!(*v, 2.0, epsilon = 1e-12);
419 }
420 }
421
422 #[test]
423 fn gap_up_uses_high_minus_prev_close() {
424 let candles = vec![
426 c(6.0, 4.0, 5.0), c(10.0, 9.0, 9.5), ];
429 let mut atr = Atr::new(2).unwrap();
430 let out = atr.batch(&candles);
431 assert_relative_eq!(out[1].unwrap(), 3.5, epsilon = 1e-12);
434 }
435
436 #[test]
437 fn batch_equals_streaming() {
438 let candles: Vec<Candle> = (0..40)
439 .map(|i| {
440 let mid = f64::from(i) + 10.0;
441 c(mid + 0.5, mid - 0.5, mid)
442 })
443 .collect();
444 let mut a = Atr::new(14).unwrap();
445 let mut b = Atr::new(14).unwrap();
446 assert_eq!(
447 a.batch(&candles),
448 candles.iter().map(|x| b.update(*x)).collect::<Vec<_>>()
449 );
450 }
451
452 #[test]
453 fn reset_clears_state() {
454 let candles: Vec<Candle> = (0..20).map(|_| c(11.0, 9.0, 10.0)).collect();
455 let mut atr = Atr::new(5).unwrap();
456 atr.batch(&candles);
457 assert!(atr.is_ready());
458 atr.reset();
459 assert!(!atr.is_ready());
460 assert_eq!(atr.update(candles[0]), None);
461 }
462
463 #[test]
464 fn never_negative() {
465 let candles: Vec<Candle> = (0..200)
466 .map(|i| {
467 let base = 100.0 + (f64::from(i) * 0.3).sin() * 5.0;
468 c(base + 1.0, base - 1.0, base)
469 })
470 .collect();
471 let mut atr = Atr::new(14).unwrap();
472 for v in atr.batch(&candles).into_iter().flatten() {
473 assert!(v >= 0.0, "ATR must be non-negative: {v}");
474 }
475 }
476
477 fn bits_eq(a: &[f64], b: &[f64]) -> bool {
478 a.len() == b.len()
479 && a.iter()
480 .zip(b)
481 .all(|(x, y)| x == y || (x.is_nan() && y.is_nan()))
482 }
483
484 fn atr_replay(period: usize, high: &[f64], low: &[f64], close: &[f64]) -> Vec<f64> {
485 let mut a = Atr::new(period).unwrap();
486 (0..high.len())
487 .map(|i| {
488 let candle = Candle::new_unchecked(close[i], high[i], low[i], close[i], 0.0, 0);
489 a.update(candle).unwrap_or(f64::NAN)
490 })
491 .collect()
492 }
493
494 fn columns(n: usize) -> (Vec<f64>, Vec<f64>, Vec<f64>) {
496 let base: Vec<f64> = (0..n)
497 .map(|i| (f64::from(u32::try_from(i).unwrap()) * 0.3).sin() * 5.0 + 100.0)
498 .collect();
499 let high = base.iter().map(|b| b + 1.0).collect();
500 let low = base.iter().map(|b| b - 1.0).collect();
501 (high, low, base)
502 }
503
504 fn to_bits(v: &[f64]) -> Vec<u64> {
505 v.iter().map(|x| x.to_bits()).collect()
506 }
507
508 #[test]
511 fn batch_atr_into_overwrites_a_dirty_buffer() {
512 let (high, low, close) = columns(260);
513 let mut out = vec![3.0; high.len()];
514 Atr::new(14)
515 .unwrap()
516 .batch_atr_into(&high, &low, &close, &mut out);
517 assert_eq!(to_bits(&out), to_bits(&atr_replay(14, &high, &low, &close)));
518 }
519
520 #[test]
521 #[should_panic(expected = "high, low, close and the output must be equal length")]
522 fn batch_atr_into_rejects_mismatched_lengths() {
523 let (high, low, close) = columns(20);
524 let mut out = vec![0.0; 19];
525 Atr::new(5)
526 .unwrap()
527 .batch_atr_into(&high, &low, &close, &mut out);
528 }
529
530 #[test]
534 fn seed_matches_a_buffered_sum_bit_for_bit() {
535 let flat = [-0.0_f64; 6];
536 let mut atr = Atr::new(6).unwrap();
537 let seed = flat
538 .iter()
539 .filter_map(|&c| atr.update(Candle::new_unchecked(c, c, c, c, 0.0, 0)))
540 .last()
541 .unwrap();
542 let trs: Vec<f64> = std::iter::once(-0.0 - -0.0)
543 .chain(std::iter::repeat_n(0.0_f64, 5))
544 .collect();
545 let buffered = trs.iter().copied().sum::<f64>() / 6.0;
546 assert_eq!(seed.to_bits(), buffered.to_bits());
547 let (high, low, close) = columns(40);
548 let want = atr_replay(9, &high, &low, &close);
549 let got = Atr::new(9).unwrap().batch_atr(&high, &low, &close);
550 assert_eq!(to_bits(&got), to_bits(&want));
551 }
552
553 #[test]
556 fn atr_tail_is_identical_on_every_dispatch_path() {
557 let (high, low, close) = columns(3000);
558 let atr = Atr::new(14).unwrap();
559 let (mut a, mut b) = (vec![0.0; 2990], vec![0.0; 2990]);
560 let make = |out: &mut [f64]| -> (f64, f64) {
561 wickra_simd::run_baseline(AtrTail {
562 high: &high[10..],
563 low: &low[10..],
564 close: &close[10..],
565 out,
566 state: (close[9], 1.7),
567 n_minus_1: atr.n_minus_1,
568 inv_period: atr.inv_period,
569 })
570 };
571 let rb = make(&mut b);
572 let ra = wickra_simd::dispatch(AtrTail {
573 high: &high[10..],
574 low: &low[10..],
575 close: &close[10..],
576 out: &mut a,
577 state: (close[9], 1.7),
578 n_minus_1: atr.n_minus_1,
579 inv_period: atr.inv_period,
580 });
581 assert_eq!(to_bits(&a), to_bits(&b));
582 assert_eq!(to_bits(&[ra.0, ra.1]), to_bits(&[rb.0, rb.1]));
583 }
584
585 #[test]
586 fn batch_atr_fast_path_is_bit_identical() {
587 let (high, low, close) = columns(300);
588 let mut atr = Atr::new(14).unwrap();
589 let got = atr.batch_atr(&high, &low, &close);
590 assert!(bits_eq(&got, &atr_replay(14, &high, &low, &close)));
591 let mut ref_atr = Atr::new(14).unwrap();
592 for i in 0..high.len() {
593 ref_atr.update(Candle::new_unchecked(
594 close[i], high[i], low[i], close[i], 0.0, 0,
595 ));
596 }
597 let next = Candle::new_unchecked(101.0, 102.0, 100.0, 101.0, 0.0, 0);
598 assert_eq!(atr.update(next), ref_atr.update(next));
599 }
600
601 #[test]
602 fn batch_atr_falls_back_when_not_fresh() {
603 let (high, low, close) = columns(40);
604 let mut atr = Atr::new(14).unwrap();
605 atr.update(Candle::new_unchecked(
606 close[0], high[0], low[0], close[0], 0.0, 0,
607 ));
608 let mut ref_atr = Atr::new(14).unwrap();
609 ref_atr.update(Candle::new_unchecked(
610 close[0], high[0], low[0], close[0], 0.0, 0,
611 ));
612 let want: Vec<f64> = (0..high.len())
613 .map(|i| {
614 ref_atr
615 .update(Candle::new_unchecked(
616 close[i], high[i], low[i], close[i], 0.0, 0,
617 ))
618 .unwrap_or(f64::NAN)
619 })
620 .collect();
621 assert!(bits_eq(&atr.batch_atr(&high, &low, &close), &want));
622 }
623
624 #[test]
625 fn batch_atr_sub_period_slice_falls_back() {
626 let (high, low, close) = columns(5);
627 let mut atr = Atr::new(14).unwrap();
628 let got = atr.batch_atr(&high, &low, &close);
629 assert!(bits_eq(&got, &atr_replay(14, &high, &low, &close)));
630 assert!(got.iter().all(|x| x.is_nan()));
631 }
632
633 proptest::proptest! {
634 #![proptest_config(proptest::test_runner::Config::with_cases(48))]
635 #[test]
636 fn atr_matches_naive(
637 period in 1usize..15,
638 bars in proptest::collection::vec(
639 (10.0_f64..1000.0, 0.0_f64..50.0, 0.0_f64..1.0),
640 0..120,
641 ),
642 ) {
643 let hlc: Vec<(f64, f64, f64)> = bars
645 .iter()
646 .map(|&(low, range, frac)| (low + range, low, low + range * frac))
647 .collect();
648 let candles: Vec<Candle> = hlc.iter().map(|&(h, l, cl)| c(h, l, cl)).collect();
649 let mut atr = Atr::new(period).unwrap();
650 let got = atr.batch(&candles);
651 let want = atr_naive(&hlc, period);
652 proptest::prop_assert_eq!(got.len(), want.len());
653 for (g, w) in got.iter().zip(want.iter()) {
654 match (g, w) {
655 (None, None) => {}
656 (Some(a), Some(b)) => proptest::prop_assert!(
657 (a - b).abs() <= 1e-9 * a.abs().max(1.0),
658 "got={a} want={b}"
659 ),
660 _ => proptest::prop_assert!(false, "warmup mismatch"),
661 }
662 }
663 }
664 }
665}