1use crate::error::{Error, Result};
4use crate::ohlcv::Candle;
5use crate::traits::Indicator;
6
7#[derive(Debug, Clone, Copy, PartialEq, Eq)]
8enum Trend {
9 Up,
10 Down,
11}
12
13#[derive(Debug, Clone, Copy)]
15struct Accel {
16 init: f64,
17 step: f64,
18 max: f64,
19}
20
21impl Accel {
22 fn validate(self) -> Result<Self> {
23 if !(self.init.is_finite() && self.step.is_finite() && self.max.is_finite()) {
24 return Err(Error::NonPositiveMultiplier);
25 }
26 if self.init <= 0.0 || self.step <= 0.0 || self.max <= 0.0 {
27 return Err(Error::NonPositiveMultiplier);
28 }
29 if self.init > self.max {
30 return Err(Error::InvalidPeriod {
31 message: "acceleration init must be <= max",
32 });
33 }
34 Ok(self)
35 }
36}
37
38#[derive(Debug, Clone)]
71pub struct SarExt {
72 start_value: f64,
73 offset_on_reverse: f64,
74 long: Accel,
75 short: Accel,
76
77 initialised: bool,
78 has_emitted: bool,
79 prev_high: f64,
80 prev_low: f64,
81 prev2_high: f64,
82 prev2_low: f64,
83 trend: Trend,
84 sar: f64,
85 ep: f64,
86 af: f64,
87}
88
89impl SarExt {
90 #[allow(clippy::too_many_arguments)]
102 pub fn new(
103 start_value: f64,
104 offset_on_reverse: f64,
105 accel_init_long: f64,
106 accel_long: f64,
107 accel_max_long: f64,
108 accel_init_short: f64,
109 accel_short: f64,
110 accel_max_short: f64,
111 ) -> Result<Self> {
112 if !start_value.is_finite() || !offset_on_reverse.is_finite() || offset_on_reverse < 0.0 {
113 return Err(Error::NonPositiveMultiplier);
114 }
115 let long = Accel {
116 init: accel_init_long,
117 step: accel_long,
118 max: accel_max_long,
119 }
120 .validate()?;
121 let short = Accel {
122 init: accel_init_short,
123 step: accel_short,
124 max: accel_max_short,
125 }
126 .validate()?;
127 Ok(Self {
128 start_value,
129 offset_on_reverse,
130 long,
131 short,
132 initialised: false,
133 has_emitted: false,
134 prev_high: f64::NAN,
135 prev_low: f64::NAN,
136 prev2_high: f64::NAN,
137 prev2_low: f64::NAN,
138 trend: Trend::Up,
139 sar: f64::NAN,
140 ep: f64::NAN,
141 af: long.init,
142 })
143 }
144
145 pub fn classic() -> Self {
148 Self::new(0.0, 0.0, 0.02, 0.02, 0.20, 0.02, 0.02, 0.20)
149 .expect("classic SAREXT params are valid")
150 }
151
152 fn signed(&self, sar: f64) -> f64 {
153 match self.trend {
154 Trend::Up => sar,
155 Trend::Down => -sar,
156 }
157 }
158}
159
160impl Indicator for SarExt {
161 type Input = Candle;
162 type Output = f64;
163
164 fn update(&mut self, candle: Candle) -> Option<f64> {
165 if !self.initialised {
166 self.prev_high = candle.high;
169 self.prev_low = candle.low;
170 self.initialised = true;
171 return None;
172 }
173
174 let new_sar = if self.has_emitted {
175 let predicted = self.sar + self.af * (self.ep - self.sar);
176 match self.trend {
177 Trend::Up => predicted.min(self.prev_low).min(self.prev2_low),
178 Trend::Down => predicted.max(self.prev_high).max(self.prev2_high),
179 }
180 } else {
181 let up_move = candle.high - self.prev_high;
189 let down_move = self.prev_low - candle.low;
190 let long = if self.start_value == 0.0 {
191 !(down_move > 0.0 && down_move > up_move)
192 } else {
193 self.start_value > 0.0
194 };
195 let auto_sar = if long { self.prev_low } else { self.prev_high };
196 self.sar = if self.start_value == 0.0 {
197 auto_sar
198 } else {
199 self.start_value.abs()
200 };
201 if long {
202 self.trend = Trend::Up;
203 self.ep = candle.high;
204 self.af = self.long.init;
205 } else {
206 self.trend = Trend::Down;
207 self.ep = candle.low;
208 self.af = self.short.init;
209 }
210 self.prev_high = candle.high;
211 self.prev_low = candle.low;
212 self.sar
213 };
214 let prev_h = self.prev_high;
215 let prev_l = self.prev_low;
216
217 let mut output_sar = new_sar;
218 let reversed = match self.trend {
219 Trend::Up => candle.low <= new_sar,
220 Trend::Down => candle.high >= new_sar,
221 };
222
223 if reversed {
224 output_sar = match self.trend {
228 Trend::Up => self.ep.max(prev_h).max(candle.high),
229 Trend::Down => self.ep.min(prev_l).min(candle.low),
230 };
231 self.trend = match self.trend {
232 Trend::Up => Trend::Down,
233 Trend::Down => Trend::Up,
234 };
235 match self.trend {
236 Trend::Up => {
237 output_sar -= output_sar.abs() * self.offset_on_reverse;
238 self.ep = candle.high;
239 self.af = self.long.init;
240 }
241 Trend::Down => {
242 output_sar += output_sar.abs() * self.offset_on_reverse;
243 self.ep = candle.low;
244 self.af = self.short.init;
245 }
246 }
247 } else {
248 match self.trend {
249 Trend::Up => {
250 if candle.high > self.ep {
251 self.ep = candle.high;
252 self.af = (self.af + self.long.step).min(self.long.max);
253 }
254 }
255 Trend::Down => {
256 if candle.low < self.ep {
257 self.ep = candle.low;
258 self.af = (self.af + self.short.step).min(self.short.max);
259 }
260 }
261 }
262 }
263
264 self.sar = output_sar;
265 self.prev2_high = self.prev_high;
266 self.prev2_low = self.prev_low;
267 self.prev_high = candle.high;
268 self.prev_low = candle.low;
269 self.has_emitted = true;
270 Some(self.signed(output_sar))
271 }
272
273 fn reset(&mut self) {
274 self.initialised = false;
275 self.has_emitted = false;
276 self.prev_high = f64::NAN;
277 self.prev_low = f64::NAN;
278 self.prev2_high = f64::NAN;
279 self.prev2_low = f64::NAN;
280 self.trend = Trend::Up;
281 self.sar = f64::NAN;
282 self.ep = f64::NAN;
283 self.af = self.long.init;
284 }
285
286 #[inline]
287 fn warmup_period(&self) -> usize {
288 2
289 }
290
291 #[inline]
292 fn is_ready(&self) -> bool {
293 self.has_emitted
294 }
295
296 #[inline]
297 fn name(&self) -> &'static str {
298 "SAREXT"
299 }
300}
301
302#[cfg(test)]
303mod tests {
304 use super::*;
305 use crate::traits::BatchExt;
306
307 fn c(h: f64, l: f64, cl: f64) -> Candle {
308 Candle::new(cl, h, l, cl, 1.0, 0).unwrap()
309 }
310
311 fn classic() -> SarExt {
312 SarExt::classic()
313 }
314
315 #[test]
316 fn rejects_invalid_params() {
317 assert!(SarExt::new(0.0, 0.0, 0.0, 0.02, 0.2, 0.02, 0.02, 0.2).is_err());
319 assert!(SarExt::new(0.0, 0.0, 0.02, 0.02, 0.2, 0.0, 0.02, 0.2).is_err());
320 assert!(SarExt::new(0.0, 0.0, 0.30, 0.02, 0.2, 0.02, 0.02, 0.2).is_err());
321 assert!(SarExt::new(0.0, 0.0, f64::NAN, 0.02, 0.2, 0.02, 0.02, 0.2).is_err());
324 assert!(SarExt::new(0.0, 0.0, 0.02, 0.02, 0.2, 0.02, f64::INFINITY, 0.2).is_err());
325 assert!(SarExt::new(f64::NAN, 0.0, 0.02, 0.02, 0.2, 0.02, 0.02, 0.2).is_err());
327 assert!(SarExt::new(0.0, -1.0, 0.02, 0.02, 0.2, 0.02, 0.02, 0.2).is_err());
328 }
329
330 #[test]
331 fn accessors_and_metadata() {
332 let s = classic();
333 assert_eq!(s.warmup_period(), 2);
334 assert_eq!(s.name(), "SAREXT");
335 assert!(!s.is_ready());
336 }
337
338 #[test]
339 fn seed_returns_none_then_emits() {
340 let mut s = classic();
341 assert_eq!(s.update(c(11.0, 9.0, 10.0)), None);
342 assert!(!s.is_ready());
343 assert!(s.update(c(12.0, 10.0, 11.0)).is_some());
344 assert!(s.is_ready());
345 }
346
347 #[test]
348 fn uptrend_is_positive_and_below_lows() {
349 let candles: Vec<Candle> = (0..40)
350 .map(|i| {
351 let base = 100.0 + f64::from(i);
352 c(base + 0.5, base - 0.5, base)
353 })
354 .collect();
355 let mut s = classic();
356 let ok = s
357 .batch(&candles)
358 .iter()
359 .enumerate()
360 .all(|(i, v)| v.is_none_or(|x| x > 0.0 && x <= candles[i].low + 1e-9));
361 assert!(ok, "long-phase SAREXT must be positive and below the low");
362 }
363
364 #[test]
365 fn downtrend_is_negative_and_above_highs() {
366 let candles: Vec<Candle> = (0..40)
367 .rev()
368 .map(|i| {
369 let base = 100.0 + f64::from(i);
370 c(base + 0.5, base - 0.5, base)
371 })
372 .collect();
373 let mut s = classic();
374 let ok = s
375 .batch(&candles)
376 .iter()
377 .enumerate()
378 .skip(5)
379 .all(|(i, v)| v.is_none_or(|x| x < 0.0 && -x >= candles[i].high - 1e-9));
380 assert!(ok, "short-phase SAREXT must be negative and above the high");
381 }
382
383 #[test]
384 fn positive_start_value_begins_long() {
385 let mut s = SarExt::new(95.0, 0.0, 0.02, 0.02, 0.2, 0.02, 0.02, 0.2).unwrap();
387 assert_eq!(s.update(c(101.0, 99.0, 100.0)), None);
388 let v = s.update(c(102.0, 100.0, 101.0)).unwrap();
389 assert!(v > 0.0);
390 }
391
392 #[test]
393 fn negative_start_value_begins_short() {
394 let mut s = SarExt::new(-105.0, 0.0, 0.02, 0.02, 0.2, 0.02, 0.02, 0.2).unwrap();
396 assert_eq!(s.update(c(101.0, 99.0, 100.0)), None);
397 let v = s.update(c(100.0, 98.0, 99.0)).unwrap();
398 assert!(v < 0.0);
399 }
400
401 #[test]
402 fn offset_on_reverse_pushes_sar_further() {
403 let candles: Vec<Candle> = (0..12)
406 .map(|i| {
407 let base = if i < 6 {
408 100.0 - f64::from(i) * 2.0
409 } else {
410 88.0 + f64::from(i - 6) * 2.0
411 };
412 c(base + 1.0, base - 1.0, base)
413 })
414 .collect();
415 let plain = SarExt::new(0.0, 0.0, 0.02, 0.02, 0.2, 0.02, 0.02, 0.2)
416 .unwrap()
417 .batch(&candles);
418 let offset = SarExt::new(0.0, 0.1, 0.02, 0.02, 0.2, 0.02, 0.02, 0.2)
419 .unwrap()
420 .batch(&candles);
421 assert_ne!(plain, offset);
423 }
424
425 #[test]
426 fn batch_equals_streaming() {
427 let candles: Vec<Candle> = (0..60)
428 .map(|i| {
429 let m = 100.0 + (f64::from(i) * 0.3).sin() * 8.0;
430 c(m + 1.0, m - 1.0, m)
431 })
432 .collect();
433 let mut a = classic();
434 let mut b = classic();
435 assert_eq!(
436 a.batch(&candles),
437 candles.iter().map(|x| b.update(*x)).collect::<Vec<_>>()
438 );
439 }
440
441 #[test]
442 fn reset_allows_clean_reuse() {
443 let candles: Vec<Candle> = (0..40)
444 .map(|i| {
445 let base = 100.0 + f64::from(i);
446 c(base + 0.5, base - 0.5, base)
447 })
448 .collect();
449 let mut s = classic();
450 let first = s.batch(&candles);
451 assert!(s.is_ready());
452 s.reset();
453 assert!(!s.is_ready());
454 assert_eq!(first, s.batch(&candles));
455 }
456
457 fn hl(high: f64, low: f64) -> Candle {
458 c(high, low, f64::midpoint(high, low))
459 }
460
461 fn run(sar: SarExt, bars: &[(f64, f64)]) -> Vec<Option<f64>> {
462 let mut sar = sar;
463 let candles: Vec<Candle> = bars.iter().map(|&(h, l)| hl(h, l)).collect();
464 sar.batch(&candles)
465 }
466
467 fn with(start: f64, offset: f64) -> SarExt {
468 SarExt::new(start, offset, 0.02, 0.02, 0.2, 0.02, 0.02, 0.2).unwrap()
469 }
470
471 fn assert_series(got: &[Option<f64>], expected: &[Option<f64>]) {
472 assert_eq!(got.len(), expected.len());
473 for (g, e) in got.iter().zip(expected) {
474 assert_eq!(g.is_some(), e.is_some());
475 if let (Some(g), Some(e)) = (g, e) {
476 approx::assert_relative_eq!(*g, *e, epsilon = 1e-12);
477 }
478 }
479 }
480
481 #[test]
482 fn rejects_every_invalid_parameter_with_its_variant() {
483 let nonpos = |r: Result<SarExt>| matches!(r, Err(Error::NonPositiveMultiplier));
484 assert!(nonpos(SarExt::new(
485 f64::INFINITY,
486 0.0,
487 0.02,
488 0.02,
489 0.2,
490 0.02,
491 0.02,
492 0.2
493 )));
494 assert!(nonpos(SarExt::new(
495 0.0,
496 f64::NAN,
497 0.02,
498 0.02,
499 0.2,
500 0.02,
501 0.02,
502 0.2
503 )));
504 assert!(nonpos(SarExt::new(
505 0.0, -0.1, 0.02, 0.02, 0.2, 0.02, 0.02, 0.2
506 )));
507 assert!(nonpos(SarExt::new(
508 0.0,
509 0.0,
510 0.02,
511 0.02,
512 f64::NAN,
513 0.02,
514 0.02,
515 0.2
516 )));
517 assert!(nonpos(SarExt::new(
518 0.0, 0.0, 0.02, 0.0, 0.2, 0.02, 0.02, 0.2
519 )));
520 assert!(nonpos(SarExt::new(
521 0.0, 0.0, 0.02, 0.02, 0.0, 0.02, 0.02, 0.2
522 )));
523 assert!(nonpos(SarExt::new(
524 0.0, 0.0, 0.02, 0.02, 0.2, 0.02, 0.0, 0.2
525 )));
526 assert!(nonpos(SarExt::new(
527 0.0, 0.0, 0.02, 0.02, 0.2, 0.02, 0.02, -0.2
528 )));
529 let invalid = |r: Result<SarExt>| matches!(r, Err(Error::InvalidPeriod { .. }));
530 assert!(invalid(SarExt::new(
531 0.0, 0.0, 0.3, 0.02, 0.2, 0.02, 0.02, 0.2
532 )));
533 assert!(invalid(SarExt::new(
534 0.0, 0.0, 0.02, 0.02, 0.2, 0.3, 0.02, 0.2
535 )));
536 }
537
538 #[test]
539 fn first_value_lands_at_index_one() {
540 let out = run(classic(), &[(10.0, 8.0), (11.0, 9.0), (12.0, 10.0)]);
541 assert!(out[0].is_none());
542 assert!(out[1..].iter().all(Option::is_some));
543 }
544
545 #[test]
546 fn hand_computed_auto_long_seed() {
547 let out = run(
551 classic(),
552 &[
553 (10.0, 8.0),
554 (11.0, 9.0),
555 (12.0, 10.0),
556 (13.0, 11.0),
557 (14.0, 12.0),
558 ],
559 );
560 assert_series(
561 &out,
562 &[None, Some(8.0), Some(8.06), Some(8.2176), Some(8.504_544)],
563 );
564 }
565
566 #[test]
567 fn hand_computed_auto_short_seed() {
568 let out = run(
572 classic(),
573 &[(12.0, 10.0), (11.0, 9.0), (10.0, 8.0), (9.0, 7.0)],
574 );
575 assert_series(&out, &[None, Some(-12.0), Some(-11.94), Some(-11.7824)]);
576 }
577
578 #[test]
579 fn hand_computed_immediate_reversals() {
580 let out = run(
585 classic(),
586 &[(10.0, 8.0), (9.0, 8.0), (9.5, 7.5), (10.5, 9.0)],
587 );
588 assert_series(&out, &[None, Some(-9.0), Some(7.5), Some(7.5)]);
589 let out = run(classic(), &[(10.0, 8.0), (10.0, 7.0), (10.5, 8.0)]);
592 assert_series(&out, &[None, Some(7.0), Some(7.0)]);
593 }
594
595 #[test]
596 fn hand_computed_reversal_clamped_above_the_extreme_point() {
597 let bars = [
601 (10.0, 8.0),
602 (11.0, 9.0),
603 (12.0, 10.0),
604 (13.0, 8.0),
605 (12.0, 7.0),
606 ];
607 let out = run(classic(), &bars);
608 assert_series(
609 &out,
610 &[None, Some(8.0), Some(8.06), Some(-13.0), Some(-13.0)],
611 );
612 }
613
614 #[test]
615 fn positive_start_value_forces_long_at_that_sar() {
616 let bars = [(10.0, 8.0), (9.0, 7.9), (9.5, 8.5)];
620 assert_series(&run(classic(), &bars[..2]), &[None, Some(-10.0)]);
621 assert_series(&run(with(7.5, 0.0), &bars), &[None, Some(7.5), Some(7.53)]);
622 }
623
624 #[test]
625 fn negative_start_value_forces_short_with_the_short_acceleration() {
626 let sar = SarExt::new(-12.0, 0.0, 0.02, 0.02, 0.2, 0.05, 0.05, 0.3).unwrap();
631 let out = run(sar, &[(10.0, 8.0), (11.0, 9.0), (10.5, 8.5), (10.0, 8.0)]);
632 assert_series(&out, &[None, Some(-12.0), Some(-11.85), Some(-11.515)]);
633 }
634
635 #[test]
636 fn forced_long_start_can_reverse_immediately_with_offset() {
637 let out = run(with(8.5, 0.01), &[(10.0, 8.0), (11.0, 8.2)]);
640 assert_series(&out, &[None, Some(-11.11)]);
641 }
642
643 #[test]
644 fn offset_on_reverse_hand_computed_both_directions() {
645 let out = run(
650 with(0.0, 0.1),
651 &[(10.0, 8.0), (9.0, 8.0), (10.0, 7.5), (11.0, 9.0)],
652 );
653 assert_series(&out, &[None, Some(-9.9), Some(6.75), Some(6.815)]);
654 }
655
656 #[test]
657 fn separate_long_acceleration_with_cap() {
658 let sar = SarExt::new(0.0, 0.0, 0.03, 0.01, 0.05, 0.02, 0.02, 0.2).unwrap();
663 let bars = [
664 (10.0, 8.0),
665 (11.0, 9.0),
666 (12.0, 10.0),
667 (13.0, 11.0),
668 (14.0, 12.0),
669 (15.0, 13.0),
670 ];
671 let out = run(sar, &bars);
672 assert_series(
673 &out,
674 &[
675 None,
676 Some(8.0),
677 Some(8.09),
678 Some(8.2464),
679 Some(8.484_08),
680 Some(8.759_876),
681 ],
682 );
683 }
684
685 #[test]
686 fn classic_matches_psar_magnitude() {
687 let candles: Vec<Candle> = (0..80)
688 .map(|i| {
689 let m = 100.0 + (f64::from(i) * 0.3).sin() * 8.0;
690 c(m + 1.0, m - 1.0, m)
691 })
692 .collect();
693 let ext = classic().batch(&candles);
694 let psar = crate::Psar::classic().batch(&candles);
695 let same = ext.iter().zip(&psar).all(|(e, p)| e.map(f64::abs) == *p);
696 assert!(same, "unsigned SAREXT must equal PSAR");
697 }
698
699 #[test]
700 fn reset_matches_a_fresh_instance_and_batch_nan_into() {
701 let candles: Vec<Candle> = (0..60)
702 .map(|i| {
703 let m = 100.0 + (f64::from(i) * 0.3).sin() * 8.0;
704 c(m + 1.0, m - 1.0, m)
705 })
706 .collect();
707 let mut s = with(-150.0, 0.05);
708 let _ = s.batch(&candles);
709 s.reset();
710 let after_reset = s.batch(&candles);
711 assert_eq!(after_reset, with(-150.0, 0.05).batch(&candles));
712 let expected: Vec<u64> = after_reset
713 .iter()
714 .map(|v| v.unwrap_or(f64::NAN).to_bits())
715 .collect();
716 let mut out = vec![0.0; candles.len()];
717 with(-150.0, 0.05).batch_nan_into(&candles, &mut out);
718 let got: Vec<u64> = out.iter().map(|v| v.to_bits()).collect();
719 assert_eq!(got, expected);
720 }
721}