1use crate::error::{Error, Result};
4use crate::ohlcv::Candle;
5use crate::traits::Indicator;
6
7#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9enum Trend {
10 Up,
11 Down,
12}
13
14#[derive(Debug, Clone)]
36pub struct Psar {
37 af_start: f64,
38 af_step: f64,
39 af_max: f64,
40
41 initialised: bool,
46 has_emitted: bool,
53 prev_high: f64,
54 prev_low: f64,
55 prev2_high: f64,
56 prev2_low: f64,
57 trend: Trend,
58 sar: f64,
59 ep: f64,
60 af: f64,
61}
62
63impl Psar {
64 pub fn new(af_start: f64, af_step: f64, af_max: f64) -> Result<Self> {
69 if !af_start.is_finite() || !af_step.is_finite() || !af_max.is_finite() {
70 return Err(Error::NonPositiveMultiplier);
71 }
72 if af_start <= 0.0 || af_step <= 0.0 || af_max <= 0.0 {
73 return Err(Error::NonPositiveMultiplier);
74 }
75 if af_start > af_max {
76 return Err(Error::InvalidPeriod {
77 message: "af_start must be <= af_max",
78 });
79 }
80 Ok(Self {
81 af_start,
82 af_step,
83 af_max,
84 initialised: false,
85 has_emitted: false,
86 prev_high: f64::NAN,
92 prev_low: f64::NAN,
93 prev2_high: f64::NAN,
94 prev2_low: f64::NAN,
95 trend: Trend::Up,
96 sar: f64::NAN,
97 ep: f64::NAN,
98 af: af_start,
99 })
100 }
101
102 pub fn classic() -> Self {
104 Self::new(0.02, 0.02, 0.20).expect("classic PSAR params are valid")
105 }
106}
107
108impl Indicator for Psar {
109 type Input = Candle;
110 type Output = f64;
111
112 fn update(&mut self, candle: Candle) -> Option<f64> {
113 if !self.initialised {
114 self.prev_high = candle.high;
117 self.prev_low = candle.low;
118 self.initialised = true;
119 return None;
120 }
121
122 let new_sar = if self.has_emitted {
123 let predicted = self.sar + self.af * (self.ep - self.sar);
128 match self.trend {
129 Trend::Up => predicted.min(self.prev_low).min(self.prev2_low),
130 Trend::Down => predicted.max(self.prev_high).max(self.prev2_high),
131 }
132 } else {
133 let up_move = candle.high - self.prev_high;
141 let down_move = self.prev_low - candle.low;
142 if down_move > 0.0 && down_move > up_move {
143 self.trend = Trend::Down;
144 self.sar = self.prev_high;
145 self.ep = candle.low;
146 } else {
147 self.trend = Trend::Up;
148 self.sar = self.prev_low;
149 self.ep = candle.high;
150 }
151 self.prev_high = candle.high;
152 self.prev_low = candle.low;
153 self.sar
154 };
155 let prev_h = self.prev_high;
156 let prev_l = self.prev_low;
157
158 let mut output_sar = new_sar;
159
160 let reversed = match self.trend {
162 Trend::Up => candle.low <= new_sar,
163 Trend::Down => candle.high >= new_sar,
164 };
165
166 if reversed {
167 output_sar = match self.trend {
171 Trend::Up => self.ep.max(prev_h).max(candle.high),
172 Trend::Down => self.ep.min(prev_l).min(candle.low),
173 };
174 self.trend = match self.trend {
175 Trend::Up => Trend::Down,
176 Trend::Down => Trend::Up,
177 };
178 self.ep = match self.trend {
179 Trend::Up => candle.high,
180 Trend::Down => candle.low,
181 };
182 self.af = self.af_start;
183 } else {
184 match self.trend {
186 Trend::Up => {
187 if candle.high > self.ep {
188 self.ep = candle.high;
189 self.af = (self.af + self.af_step).min(self.af_max);
190 }
191 }
192 Trend::Down => {
193 if candle.low < self.ep {
194 self.ep = candle.low;
195 self.af = (self.af + self.af_step).min(self.af_max);
196 }
197 }
198 }
199 }
200
201 self.sar = output_sar;
202 self.prev2_high = self.prev_high;
203 self.prev2_low = self.prev_low;
204 self.prev_high = candle.high;
205 self.prev_low = candle.low;
206 self.has_emitted = true;
207 Some(output_sar)
208 }
209
210 fn reset(&mut self) {
211 self.initialised = false;
215 self.has_emitted = false;
216 self.prev_high = f64::NAN;
217 self.prev_low = f64::NAN;
218 self.prev2_high = f64::NAN;
219 self.prev2_low = f64::NAN;
220 self.trend = Trend::Up;
221 self.sar = f64::NAN;
222 self.ep = f64::NAN;
223 self.af = self.af_start;
224 }
225
226 #[inline]
227 fn warmup_period(&self) -> usize {
228 2
229 }
230
231 #[inline]
232 fn is_ready(&self) -> bool {
233 self.has_emitted
240 }
241
242 #[inline]
243 fn name(&self) -> &'static str {
244 "PSAR"
245 }
246}
247
248#[cfg(test)]
249mod tests {
250 use super::*;
251 use crate::traits::BatchExt;
252
253 fn c(h: f64, l: f64, cl: f64) -> Candle {
254 Candle::new(cl, h, l, cl, 1.0, 0).unwrap()
255 }
256
257 #[test]
258 fn first_candle_returns_none() {
259 let mut psar = Psar::classic();
260 assert_eq!(psar.update(c(11.0, 9.0, 10.0)), None);
261 }
262
263 #[test]
264 fn pure_uptrend_sar_below_lows() {
265 let candles: Vec<Candle> = (0..40)
266 .map(|i| {
267 let base = 100.0 + f64::from(i);
268 c(base + 0.5, base - 0.5, base)
269 })
270 .collect();
271 let mut psar = Psar::classic();
272 let ok = psar
277 .batch(&candles)
278 .iter()
279 .enumerate()
280 .all(|(i, sar)| sar.is_none_or(|s| s <= candles[i].low + 1e-9));
281 assert!(ok, "SAR sat above a candle's low on a pure uptrend");
282 }
283
284 #[test]
285 fn pure_downtrend_sar_above_highs() {
286 let candles: Vec<Candle> = (0..40)
287 .rev()
288 .map(|i| {
289 let base = 100.0 + f64::from(i);
290 c(base + 0.5, base - 0.5, base)
291 })
292 .collect();
293 let mut psar = Psar::classic();
294 let ok = psar
298 .batch(&candles)
299 .iter()
300 .enumerate()
301 .skip(5)
302 .all(|(i, sar)| sar.is_none_or(|s| s >= candles[i].high - 1e-9));
303 assert!(ok, "SAR sat below a candle's high on a pure downtrend");
304 }
305
306 #[test]
307 fn batch_equals_streaming() {
308 let candles: Vec<Candle> = (0..60)
309 .map(|i| {
310 let m = 100.0 + (f64::from(i) * 0.3).sin() * 8.0;
311 c(m + 1.0, m - 1.0, m)
312 })
313 .collect();
314 let mut a = Psar::classic();
315 let mut b = Psar::classic();
316 assert_eq!(
317 a.batch(&candles),
318 candles.iter().map(|x| b.update(*x)).collect::<Vec<_>>()
319 );
320 }
321
322 #[test]
326 fn accessors_and_metadata() {
327 let psar = Psar::classic();
328 assert_eq!(psar.warmup_period(), 2);
329 assert_eq!(psar.name(), "PSAR");
330 }
331
332 #[test]
333 fn rejects_invalid_params() {
334 assert!(Psar::new(0.0, 0.02, 0.20).is_err());
335 assert!(Psar::new(0.02, 0.0, 0.20).is_err());
336 assert!(Psar::new(0.30, 0.02, 0.20).is_err());
337 assert!(Psar::new(f64::NAN, 0.02, 0.20).is_err());
338 }
339
340 #[test]
341 fn is_ready_only_after_first_some_value() {
342 let mut psar = Psar::classic();
347 assert!(!psar.is_ready(), "fresh PSAR must not be ready");
348 let first = psar.update(c(11.0, 9.0, 10.0));
349 assert!(first.is_none(), "seed candle returns None by design");
350 assert!(
351 !psar.is_ready(),
352 "is_ready must stay false until a Some value is produced"
353 );
354 let second = psar.update(c(12.0, 10.0, 11.0));
355 assert!(second.is_some(), "second candle must emit");
356 assert!(
357 psar.is_ready(),
358 "is_ready must flip to true once a real value has been returned"
359 );
360 }
361
362 #[test]
363 fn reset_allows_clean_reuse() {
364 let candles: Vec<Candle> = (0..40)
365 .map(|i| {
366 let base = 100.0 + f64::from(i);
367 c(base + 0.5, base - 0.5, base)
368 })
369 .collect();
370 let mut psar = Psar::classic();
371 let first = psar.batch(&candles);
372 assert!(psar.is_ready());
373 psar.reset();
374 assert!(!psar.is_ready());
375 let second = psar.batch(&candles);
377 assert_eq!(first, second);
378 }
379
380 fn hl(high: f64, low: f64) -> Candle {
381 c(high, low, f64::midpoint(high, low))
382 }
383
384 fn run(bars: &[(f64, f64)]) -> Vec<Option<f64>> {
385 let candles: Vec<Candle> = bars.iter().map(|&(h, l)| hl(h, l)).collect();
386 Psar::classic().batch(&candles)
387 }
388
389 fn assert_series(got: &[Option<f64>], expected: &[Option<f64>]) {
390 assert_eq!(got.len(), expected.len());
391 for (g, e) in got.iter().zip(expected) {
392 assert_eq!(g.is_some(), e.is_some());
393 if let (Some(g), Some(e)) = (g, e) {
394 approx::assert_relative_eq!(*g, *e, epsilon = 1e-12);
395 }
396 }
397 }
398
399 #[test]
400 fn rejects_every_invalid_parameter() {
401 assert!(matches!(
402 Psar::new(f64::NAN, 0.02, 0.2),
403 Err(Error::NonPositiveMultiplier)
404 ));
405 assert!(matches!(
406 Psar::new(0.02, f64::INFINITY, 0.2),
407 Err(Error::NonPositiveMultiplier)
408 ));
409 assert!(matches!(
410 Psar::new(0.02, 0.02, f64::NAN),
411 Err(Error::NonPositiveMultiplier)
412 ));
413 assert!(matches!(
414 Psar::new(-0.02, 0.02, 0.2),
415 Err(Error::NonPositiveMultiplier)
416 ));
417 assert!(matches!(
418 Psar::new(0.02, -0.02, 0.2),
419 Err(Error::NonPositiveMultiplier)
420 ));
421 assert!(matches!(
422 Psar::new(0.02, 0.02, 0.0),
423 Err(Error::NonPositiveMultiplier)
424 ));
425 assert!(matches!(
426 Psar::new(0.3, 0.02, 0.2),
427 Err(Error::InvalidPeriod { .. })
428 ));
429 assert!(Psar::new(0.2, 0.02, 0.2).is_ok());
430 }
431
432 #[test]
433 fn first_value_lands_at_index_one() {
434 let out = run(&[(10.0, 8.0), (11.0, 9.0), (12.0, 10.0)]);
435 assert_eq!(Psar::classic().warmup_period(), 2);
436 assert!(out[0].is_none());
437 assert!(out[1..].iter().all(Option::is_some));
438 }
439
440 #[test]
441 fn hand_computed_long_seed() {
442 let out = run(&[
448 (10.0, 8.0),
449 (11.0, 9.0),
450 (12.0, 10.0),
451 (13.0, 11.0),
452 (14.0, 12.0),
453 ]);
454 assert_series(
455 &out,
456 &[None, Some(8.0), Some(8.06), Some(8.2176), Some(8.504_544)],
457 );
458 }
459
460 #[test]
461 fn hand_computed_short_seed() {
462 let out = run(&[(12.0, 10.0), (11.0, 9.0), (10.0, 8.0), (9.0, 7.0)]);
467 assert_series(&out, &[None, Some(12.0), Some(11.94), Some(11.7824)]);
468 }
469
470 #[test]
471 fn equal_moves_seed_long() {
472 let out = run(&[(10.0, 8.0), (11.0, 7.0)]);
476 assert_series(&out, &[None, Some(11.0)]);
477 }
478
479 #[test]
480 fn hand_computed_immediate_reversal_long_to_short_then_back() {
481 let out = run(&[(10.0, 8.0), (9.0, 8.0), (9.5, 7.5), (10.5, 9.0)]);
488 assert_series(&out, &[None, Some(9.0), Some(7.5), Some(7.5)]);
489 }
490
491 #[test]
492 fn hand_computed_immediate_reversal_short_to_long() {
493 let out = run(&[(10.0, 8.0), (10.0, 7.0), (10.5, 8.0), (11.0, 9.0)]);
498 assert_series(&out, &[None, Some(7.0), Some(7.0), Some(7.0)]);
499 }
500
501 #[test]
502 fn hand_computed_reversal_clamped_above_the_extreme_point() {
503 let out = run(&[
510 (10.0, 8.0),
511 (11.0, 9.0),
512 (12.0, 10.0),
513 (13.0, 8.0),
514 (12.0, 7.0),
515 (11.5, 6.5),
516 ]);
517 assert_series(
518 &out,
519 &[
520 None,
521 Some(8.0),
522 Some(8.06),
523 Some(13.0),
524 Some(13.0),
525 Some(13.0),
526 ],
527 );
528 }
529
530 #[test]
531 fn acceleration_factor_caps_at_max() {
532 let candles: Vec<Candle> = [
539 (10.0, 8.0),
540 (11.0, 9.0),
541 (12.0, 10.0),
542 (13.0, 11.0),
543 (14.0, 12.0),
544 ]
545 .iter()
546 .map(|&(h, l)| hl(h, l))
547 .collect();
548 let out = Psar::new(0.1, 0.1, 0.2).unwrap().batch(&candles);
549 assert_series(&out, &[None, Some(8.0), Some(8.3), Some(9.0), Some(9.8)]);
550 }
551
552 #[test]
553 fn reset_matches_a_fresh_instance_and_batch_nan_into() {
554 let candles: Vec<Candle> = (0..60)
555 .map(|i| {
556 let m = 100.0 + (f64::from(i) * 0.3).sin() * 8.0;
557 c(m + 1.0, m - 1.0, m)
558 })
559 .collect();
560 let mut psar = Psar::classic();
561 let _ = psar.batch(&candles);
562 psar.reset();
563 let after_reset = psar.batch(&candles);
564 assert_eq!(after_reset, Psar::classic().batch(&candles));
565 let expected: Vec<u64> = after_reset
566 .iter()
567 .map(|v| v.unwrap_or(f64::NAN).to_bits())
568 .collect();
569 let mut out = vec![0.0; candles.len()];
570 Psar::classic().batch_nan_into(&candles, &mut out);
571 let got: Vec<u64> = out.iter().map(|v| v.to_bits()).collect();
572 assert_eq!(got, expected);
573 }
574}