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 trend: Trend,
82 sar: f64,
83 ep: f64,
84 af: f64,
85}
86
87impl SarExt {
88 #[allow(clippy::too_many_arguments)]
100 pub fn new(
101 start_value: f64,
102 offset_on_reverse: f64,
103 accel_init_long: f64,
104 accel_long: f64,
105 accel_max_long: f64,
106 accel_init_short: f64,
107 accel_short: f64,
108 accel_max_short: f64,
109 ) -> Result<Self> {
110 if !start_value.is_finite() || !offset_on_reverse.is_finite() || offset_on_reverse < 0.0 {
111 return Err(Error::NonPositiveMultiplier);
112 }
113 let long = Accel {
114 init: accel_init_long,
115 step: accel_long,
116 max: accel_max_long,
117 }
118 .validate()?;
119 let short = Accel {
120 init: accel_init_short,
121 step: accel_short,
122 max: accel_max_short,
123 }
124 .validate()?;
125 Ok(Self {
126 start_value,
127 offset_on_reverse,
128 long,
129 short,
130 initialised: false,
131 has_emitted: false,
132 prev_high: f64::NAN,
133 prev_low: f64::NAN,
134 trend: Trend::Up,
135 sar: f64::NAN,
136 ep: f64::NAN,
137 af: long.init,
138 })
139 }
140
141 pub fn classic() -> Self {
144 Self::new(0.0, 0.0, 0.02, 0.02, 0.20, 0.02, 0.02, 0.20)
145 .expect("classic SAREXT params are valid")
146 }
147
148 fn signed(&self, sar: f64) -> f64 {
149 match self.trend {
150 Trend::Up => sar,
151 Trend::Down => -sar,
152 }
153 }
154}
155
156impl Indicator for SarExt {
157 type Input = Candle;
158 type Output = f64;
159
160 fn update(&mut self, candle: Candle) -> Option<f64> {
161 if !self.initialised {
162 self.prev_high = candle.high;
163 self.prev_low = candle.low;
164 if self.start_value > 0.0 {
165 self.trend = Trend::Up;
166 self.sar = self.start_value;
167 self.ep = candle.high;
168 self.af = self.long.init;
169 } else if self.start_value < 0.0 {
170 self.trend = Trend::Down;
171 self.sar = -self.start_value;
172 self.ep = candle.low;
173 self.af = self.short.init;
174 } else {
175 self.trend = Trend::Up;
176 self.sar = candle.low;
177 self.ep = candle.high;
178 self.af = self.long.init;
179 }
180 self.initialised = true;
181 return None;
182 }
183
184 let mut new_sar = self.sar + self.af * (self.ep - self.sar);
185 let prev_h = self.prev_high;
186 let prev_l = self.prev_low;
187 new_sar = match self.trend {
188 Trend::Up => new_sar.min(prev_l).min(candle.low),
189 Trend::Down => new_sar.max(prev_h).max(candle.high),
190 };
191
192 let mut output_sar = new_sar;
193 let reversed = match self.trend {
194 Trend::Up => candle.low <= new_sar,
195 Trend::Down => candle.high >= new_sar,
196 };
197
198 if reversed {
199 output_sar = self.ep;
200 self.trend = match self.trend {
201 Trend::Up => Trend::Down,
202 Trend::Down => Trend::Up,
203 };
204 match self.trend {
205 Trend::Up => {
206 output_sar -= output_sar.abs() * self.offset_on_reverse;
207 self.ep = candle.high;
208 self.af = self.long.init;
209 }
210 Trend::Down => {
211 output_sar += output_sar.abs() * self.offset_on_reverse;
212 self.ep = candle.low;
213 self.af = self.short.init;
214 }
215 }
216 } else {
217 match self.trend {
218 Trend::Up => {
219 if candle.high > self.ep {
220 self.ep = candle.high;
221 self.af = (self.af + self.long.step).min(self.long.max);
222 }
223 }
224 Trend::Down => {
225 if candle.low < self.ep {
226 self.ep = candle.low;
227 self.af = (self.af + self.short.step).min(self.short.max);
228 }
229 }
230 }
231 }
232
233 self.sar = output_sar;
234 self.prev_high = candle.high;
235 self.prev_low = candle.low;
236 self.has_emitted = true;
237 Some(self.signed(output_sar))
238 }
239
240 fn reset(&mut self) {
241 self.initialised = false;
242 self.has_emitted = false;
243 self.prev_high = f64::NAN;
244 self.prev_low = f64::NAN;
245 self.trend = Trend::Up;
246 self.sar = f64::NAN;
247 self.ep = f64::NAN;
248 self.af = self.long.init;
249 }
250
251 #[inline]
252 fn warmup_period(&self) -> usize {
253 2
254 }
255
256 #[inline]
257 fn is_ready(&self) -> bool {
258 self.has_emitted
259 }
260
261 #[inline]
262 fn name(&self) -> &'static str {
263 "SAREXT"
264 }
265}
266
267#[cfg(test)]
268mod tests {
269 use super::*;
270 use crate::traits::BatchExt;
271
272 fn c(h: f64, l: f64, cl: f64) -> Candle {
273 Candle::new(cl, h, l, cl, 1.0, 0).unwrap()
274 }
275
276 fn classic() -> SarExt {
277 SarExt::classic()
278 }
279
280 #[test]
281 fn rejects_invalid_params() {
282 assert!(SarExt::new(0.0, 0.0, 0.0, 0.02, 0.2, 0.02, 0.02, 0.2).is_err());
284 assert!(SarExt::new(0.0, 0.0, 0.02, 0.02, 0.2, 0.0, 0.02, 0.2).is_err());
285 assert!(SarExt::new(0.0, 0.0, 0.30, 0.02, 0.2, 0.02, 0.02, 0.2).is_err());
286 assert!(SarExt::new(0.0, 0.0, f64::NAN, 0.02, 0.2, 0.02, 0.02, 0.2).is_err());
289 assert!(SarExt::new(0.0, 0.0, 0.02, 0.02, 0.2, 0.02, f64::INFINITY, 0.2).is_err());
290 assert!(SarExt::new(f64::NAN, 0.0, 0.02, 0.02, 0.2, 0.02, 0.02, 0.2).is_err());
292 assert!(SarExt::new(0.0, -1.0, 0.02, 0.02, 0.2, 0.02, 0.02, 0.2).is_err());
293 }
294
295 #[test]
296 fn accessors_and_metadata() {
297 let s = classic();
298 assert_eq!(s.warmup_period(), 2);
299 assert_eq!(s.name(), "SAREXT");
300 assert!(!s.is_ready());
301 }
302
303 #[test]
304 fn seed_returns_none_then_emits() {
305 let mut s = classic();
306 assert_eq!(s.update(c(11.0, 9.0, 10.0)), None);
307 assert!(!s.is_ready());
308 assert!(s.update(c(12.0, 10.0, 11.0)).is_some());
309 assert!(s.is_ready());
310 }
311
312 #[test]
313 fn uptrend_is_positive_and_below_lows() {
314 let candles: Vec<Candle> = (0..40)
315 .map(|i| {
316 let base = 100.0 + f64::from(i);
317 c(base + 0.5, base - 0.5, base)
318 })
319 .collect();
320 let mut s = classic();
321 let ok = s
322 .batch(&candles)
323 .iter()
324 .enumerate()
325 .all(|(i, v)| v.is_none_or(|x| x > 0.0 && x <= candles[i].low + 1e-9));
326 assert!(ok, "long-phase SAREXT must be positive and below the low");
327 }
328
329 #[test]
330 fn downtrend_is_negative_and_above_highs() {
331 let candles: Vec<Candle> = (0..40)
332 .rev()
333 .map(|i| {
334 let base = 100.0 + f64::from(i);
335 c(base + 0.5, base - 0.5, base)
336 })
337 .collect();
338 let mut s = classic();
339 let ok = s
340 .batch(&candles)
341 .iter()
342 .enumerate()
343 .skip(5)
344 .all(|(i, v)| v.is_none_or(|x| x < 0.0 && -x >= candles[i].high - 1e-9));
345 assert!(ok, "short-phase SAREXT must be negative and above the high");
346 }
347
348 #[test]
349 fn positive_start_value_begins_long() {
350 let mut s = SarExt::new(95.0, 0.0, 0.02, 0.02, 0.2, 0.02, 0.02, 0.2).unwrap();
352 assert_eq!(s.update(c(101.0, 99.0, 100.0)), None);
353 let v = s.update(c(102.0, 100.0, 101.0)).unwrap();
354 assert!(v > 0.0);
355 }
356
357 #[test]
358 fn negative_start_value_begins_short() {
359 let mut s = SarExt::new(-105.0, 0.0, 0.02, 0.02, 0.2, 0.02, 0.02, 0.2).unwrap();
361 assert_eq!(s.update(c(101.0, 99.0, 100.0)), None);
362 let v = s.update(c(100.0, 98.0, 99.0)).unwrap();
363 assert!(v < 0.0);
364 }
365
366 #[test]
367 fn offset_on_reverse_pushes_sar_further() {
368 let candles: Vec<Candle> = (0..12)
371 .map(|i| {
372 let base = if i < 6 {
373 100.0 - f64::from(i) * 2.0
374 } else {
375 88.0 + f64::from(i - 6) * 2.0
376 };
377 c(base + 1.0, base - 1.0, base)
378 })
379 .collect();
380 let plain = SarExt::new(0.0, 0.0, 0.02, 0.02, 0.2, 0.02, 0.02, 0.2)
381 .unwrap()
382 .batch(&candles);
383 let offset = SarExt::new(0.0, 0.1, 0.02, 0.02, 0.2, 0.02, 0.02, 0.2)
384 .unwrap()
385 .batch(&candles);
386 assert_ne!(plain, offset);
388 }
389
390 #[test]
391 fn batch_equals_streaming() {
392 let candles: Vec<Candle> = (0..60)
393 .map(|i| {
394 let m = 100.0 + (f64::from(i) * 0.3).sin() * 8.0;
395 c(m + 1.0, m - 1.0, m)
396 })
397 .collect();
398 let mut a = classic();
399 let mut b = classic();
400 assert_eq!(
401 a.batch(&candles),
402 candles.iter().map(|x| b.update(*x)).collect::<Vec<_>>()
403 );
404 }
405
406 #[test]
407 fn reset_allows_clean_reuse() {
408 let candles: Vec<Candle> = (0..40)
409 .map(|i| {
410 let base = 100.0 + f64::from(i);
411 c(base + 0.5, base - 0.5, base)
412 })
413 .collect();
414 let mut s = classic();
415 let first = s.batch(&candles);
416 assert!(s.is_ready());
417 s.reset();
418 assert!(!s.is_ready());
419 assert_eq!(first, s.batch(&candles));
420 }
421}