1use crate::error::{Error, Result};
4use crate::traits::Indicator;
5
6#[derive(Debug, Clone)]
26pub struct Rsi {
27 period: usize,
28 n_minus_1: f64,
30 inv_period: f64,
33 prev_close: f64,
36 has_prev: bool,
37 seed_buf_gains: Vec<f64>,
40 seed_buf_losses: Vec<f64>,
41 avg_gain: f64,
44 avg_loss: f64,
45 avgs_seeded: bool,
46 last_value: Option<f64>,
47}
48
49impl Rsi {
50 pub fn new(period: usize) -> Result<Self> {
56 if period == 0 {
57 return Err(Error::PeriodZero);
58 }
59 if period > crate::error::MAX_PERIOD {
60 return Err(Error::InvalidPeriod {
61 message: crate::error::PERIOD_ABOVE_MAX,
62 });
63 }
64 Ok(Self {
65 period,
66 n_minus_1: (period - 1) as f64,
67 inv_period: 1.0 / period as f64,
68 prev_close: 0.0,
69 has_prev: false,
70 seed_buf_gains: Vec::with_capacity(period),
71 seed_buf_losses: Vec::with_capacity(period),
72 avg_gain: 0.0,
73 avg_loss: 0.0,
74 avgs_seeded: false,
75 last_value: None,
76 })
77 }
78
79 pub const fn period(&self) -> usize {
81 self.period
82 }
83
84 pub const fn value(&self) -> Option<f64> {
86 self.last_value
87 }
88
89 pub fn batch_nan(&mut self, inputs: &[f64]) -> Vec<f64> {
95 crate::traits::BatchNanExt::batch_nan(self, inputs)
96 }
97
98 fn seed_batch(&mut self, inputs: &[f64], out: &mut [f64]) -> (f64, f64, f64) {
104 let p = self.period;
105 out[..p].fill(f64::NAN);
106 let mut prev = inputs[0];
108 let (mut sum_gain, mut sum_loss) = (0.0_f64, 0.0_f64);
109 for &x in &inputs[1..=p] {
110 let diff = x - prev;
111 prev = x;
112 let gain = if diff > 0.0 { diff } else { 0.0 };
113 let loss = if diff < 0.0 { -diff } else { 0.0 };
114 self.seed_buf_gains.push(gain);
115 self.seed_buf_losses.push(loss);
116 sum_gain += gain;
117 sum_loss += loss;
118 }
119 let p_f64 = p as f64;
120 let ag = sum_gain / p_f64;
121 let al = sum_loss / p_f64;
122 out[p] = Self::rsi_from_avgs(ag, al);
123 (prev, ag, al)
124 }
125
126 #[inline]
127 fn rsi_from_avgs(avg_gain: f64, avg_loss: f64) -> f64 {
128 let denom = avg_gain + avg_loss;
134 if denom == 0.0 {
135 50.0
136 } else {
137 100.0 * avg_gain / denom
138 }
139 }
140}
141
142struct WilderTail<'a> {
147 inputs: &'a [f64],
148 out: &'a mut [f64],
149 state: (f64, f64, f64),
150 n_minus_1: f64,
151 inv_period: f64,
152}
153
154#[allow(clippy::inline_always)]
157impl wickra_simd::Kernel for WilderTail<'_> {
158 type Output = (f64, f64, f64);
159
160 #[inline(always)]
161 fn run<S: wickra_simd::Simd>(self, _simd: S) -> (f64, f64, f64) {
162 let (mut prev, mut ag, mut al) = self.state;
163 let (n_minus_1, inv_period) = (self.n_minus_1, self.inv_period);
164 for (slot, &x) in self.out.iter_mut().zip(self.inputs) {
165 let diff = x - prev;
166 prev = x;
167 let gain = if diff > 0.0 { diff } else { 0.0 };
168 let loss = if diff < 0.0 { -diff } else { 0.0 };
169 ag = ag.mul_add(n_minus_1, gain) * inv_period;
170 al = al.mul_add(n_minus_1, loss) * inv_period;
171 *slot = Rsi::rsi_from_avgs(ag, al);
172 }
173 (prev, ag, al)
174 }
175}
176
177impl Indicator for Rsi {
178 type Input = f64;
179 type Output = f64;
180
181 fn update(&mut self, input: f64) -> Option<f64> {
182 if !input.is_finite() {
183 return None;
184 }
185
186 if !self.has_prev {
187 self.prev_close = input;
188 self.has_prev = true;
189 return None;
190 }
191 let prev = self.prev_close;
192 self.prev_close = input;
193
194 let diff = input - prev;
195 let gain = if diff > 0.0 { diff } else { 0.0 };
196 let loss = if diff < 0.0 { -diff } else { 0.0 };
197
198 if self.avgs_seeded {
199 let new_ag = self.avg_gain.mul_add(self.n_minus_1, gain) * self.inv_period;
202 let new_al = self.avg_loss.mul_add(self.n_minus_1, loss) * self.inv_period;
203 self.avg_gain = new_ag;
204 self.avg_loss = new_al;
205 let v = Self::rsi_from_avgs(new_ag, new_al);
206 self.last_value = Some(v);
207 return Some(v);
208 }
209
210 self.seed_buf_gains.push(gain);
211 self.seed_buf_losses.push(loss);
212 if self.seed_buf_gains.len() == self.period {
213 let ag = self.seed_buf_gains.iter().sum::<f64>() / self.period as f64;
214 let al = self.seed_buf_losses.iter().sum::<f64>() / self.period as f64;
215 self.avg_gain = ag;
216 self.avg_loss = al;
217 self.avgs_seeded = true;
218 let v = Self::rsi_from_avgs(ag, al);
219 self.last_value = Some(v);
220 return Some(v);
221 }
222 None
223 }
224
225 fn reset(&mut self) {
226 self.prev_close = 0.0;
227 self.has_prev = false;
228 self.seed_buf_gains.clear();
229 self.seed_buf_losses.clear();
230 self.avg_gain = 0.0;
231 self.avg_loss = 0.0;
232 self.avgs_seeded = false;
233 self.last_value = None;
234 }
235
236 #[inline]
237 fn warmup_period(&self) -> usize {
238 self.period + 1
239 }
240
241 #[inline]
242 fn is_ready(&self) -> bool {
243 self.last_value.is_some()
244 }
245
246 #[inline]
247 fn name(&self) -> &'static str {
248 "RSI"
249 }
250
251 fn batch_nan_into(&mut self, inputs: &[f64], out: &mut [f64]) {
261 assert_eq!(
262 inputs.len(),
263 out.len(),
264 "batch output length must equal input length"
265 );
266 let p = self.period;
267 let n = inputs.len();
268 if self.has_prev
269 || self.avgs_seeded
270 || !self.seed_buf_gains.is_empty()
271 || n <= p
272 || !inputs.iter().all(|x| x.is_finite())
273 {
274 for (slot, &x) in out.iter_mut().zip(inputs) {
275 *slot = self.update(x).unwrap_or(f64::NAN);
276 }
277 return;
278 }
279
280 let (mut prev, mut ag, mut al) = self.seed_batch(inputs, out);
281 (prev, ag, al) = wickra_simd::dispatch(WilderTail {
284 inputs: &inputs[p + 1..],
285 out: &mut out[p + 1..],
286 state: (prev, ag, al),
287 n_minus_1: self.n_minus_1,
288 inv_period: self.inv_period,
289 });
290
291 self.prev_close = prev;
293 self.has_prev = true;
294 self.avg_gain = ag;
295 self.avg_loss = al;
296 self.avgs_seeded = true;
297 self.last_value = Some(out[n - 1]);
298 }
299
300 fn batch_fast_into(&mut self, inputs: &[f64], out: &mut [f64]) {
307 assert_eq!(
308 inputs.len(),
309 out.len(),
310 "batch output length must equal input length"
311 );
312 let p = self.period;
313 let n = inputs.len();
314 if self.has_prev
315 || self.avgs_seeded
316 || !self.seed_buf_gains.is_empty()
317 || n <= p
318 || !crate::fast::in_range(inputs)
319 {
320 self.batch_nan_into(inputs, out);
321 return;
322 }
323 let (_, ag, al) = self.seed_batch(inputs, out);
324 let (ag, al) = wickra_simd::dispatch(crate::fast::RsiFast {
325 x: inputs,
326 period: p,
327 avg_gain: ag,
328 avg_loss: al,
329 n_minus_1: self.n_minus_1,
330 inv_period: self.inv_period,
331 out,
332 _borrow: std::marker::PhantomData,
333 });
334 self.prev_close = inputs[n - 1];
335 self.has_prev = true;
336 self.avg_gain = ag;
337 self.avg_loss = al;
338 self.avgs_seeded = true;
339 self.last_value = Some(out[n - 1]);
340 }
341}
342
343#[cfg(test)]
344mod tests {
345 use super::*;
346 use crate::traits::BatchExt;
347 use approx::assert_relative_eq;
348
349 fn rsi_naive(prices: &[f64], period: usize) -> Vec<Option<f64>> {
351 let n = period as f64;
352 let mut out = vec![None; prices.len()];
353 let mut gains: Vec<f64> = Vec::new();
354 let mut losses: Vec<f64> = Vec::new();
355 let mut avg_gain: Option<f64> = None;
356 let mut avg_loss: Option<f64> = None;
357 let rsi_val = |ag: f64, al: f64| -> f64 {
358 if al == 0.0 {
359 if ag == 0.0 {
360 50.0
361 } else {
362 100.0
363 }
364 } else {
365 100.0 - 100.0 / (1.0 + ag / al)
366 }
367 };
368 for i in 1..prices.len() {
369 let diff = prices[i] - prices[i - 1];
370 let gain = if diff > 0.0 { diff } else { 0.0 };
371 let loss = if diff < 0.0 { -diff } else { 0.0 };
372 if let (Some(ag), Some(al)) = (avg_gain, avg_loss) {
373 let nag = (ag * (n - 1.0) + gain) / n;
374 let nal = (al * (n - 1.0) + loss) / n;
375 avg_gain = Some(nag);
376 avg_loss = Some(nal);
377 out[i] = Some(rsi_val(nag, nal));
378 } else {
379 gains.push(gain);
380 losses.push(loss);
381 if gains.len() == period {
382 let ag = gains.iter().sum::<f64>() / n;
383 let al = losses.iter().sum::<f64>() / n;
384 avg_gain = Some(ag);
385 avg_loss = Some(al);
386 out[i] = Some(rsi_val(ag, al));
387 }
388 }
389 }
390 out
391 }
392
393 #[test]
394 fn new_rejects_zero_period() {
395 assert!(matches!(Rsi::new(0), Err(Error::PeriodZero)));
396 }
397
398 #[test]
402 fn accessors_and_metadata() {
403 let mut rsi = Rsi::new(14).unwrap();
404 assert_eq!(rsi.period(), 14);
405 assert_eq!(rsi.name(), "RSI");
406 assert_eq!(rsi.value(), None);
407 for i in 1..=15 {
408 rsi.update(100.0 + f64::from(i));
409 }
410 assert!(rsi.value().is_some());
411 }
412
413 #[test]
419 fn naive_helper_flat_series_yields_50() {
420 let ks = rsi_naive(&[42.0; 20], 5);
421 for r in ks.into_iter().skip(5) {
422 assert_eq!(r.expect("ready after period+1 inputs"), 50.0);
423 }
424 }
425
426 #[test]
432 fn naive_helper_monotone_up_yields_100() {
433 let prices: Vec<f64> = (1..=20).map(f64::from).collect();
434 let ks = rsi_naive(&prices, 5);
435 for r in ks.into_iter().skip(5) {
436 assert_eq!(r.expect("ready after period+1 inputs"), 100.0);
437 }
438 }
439
440 #[test]
441 fn warmup_period_is_period_plus_one() {
442 let rsi = Rsi::new(14).unwrap();
443 assert_eq!(rsi.warmup_period(), 15);
444 }
445
446 #[test]
447 fn first_emission_at_index_period() {
448 let prices: Vec<f64> = (1..=20).map(f64::from).collect();
450 let mut rsi = Rsi::new(14).unwrap();
451 let out = rsi.batch(&prices);
452 for x in &out[..14] {
454 assert!(x.is_none());
455 }
456 assert!(out[14].is_some());
457 }
458
459 #[test]
460 fn pure_uptrend_yields_rsi_100() {
461 let prices: Vec<f64> = (1..=20).map(f64::from).collect();
462 let mut rsi = Rsi::new(14).unwrap();
463 let out = rsi.batch(&prices);
464 for v in out.iter().filter_map(|x| x.as_ref()) {
466 assert_relative_eq!(*v, 100.0, epsilon = 1e-9);
467 }
468 }
469
470 #[test]
471 fn pure_downtrend_yields_rsi_0() {
472 let prices: Vec<f64> = (1..=20).rev().map(f64::from).collect();
473 let mut rsi = Rsi::new(14).unwrap();
474 let out = rsi.batch(&prices);
475 for v in out.iter().filter_map(|x| x.as_ref()) {
476 assert_relative_eq!(*v, 0.0, epsilon = 1e-9);
477 }
478 }
479
480 #[test]
481 fn flat_series_yields_rsi_50() {
482 let prices = [10.0_f64; 30];
483 let mut rsi = Rsi::new(14).unwrap();
484 let out = rsi.batch(&prices);
485 for v in out.iter().filter_map(|x| x.as_ref()) {
486 assert_relative_eq!(*v, 50.0, epsilon = 1e-12);
487 }
488 }
489
490 #[test]
491 fn classic_wilder_textbook_values() {
492 let prices = [
497 44.34, 44.09, 44.15, 43.61, 44.33, 44.83, 45.10, 45.42, 45.84, 46.08, 45.89, 46.03,
498 45.61, 46.28, 46.28,
499 ];
500 let mut rsi = Rsi::new(14).unwrap();
501 let out = rsi.batch(&prices);
502 let first = out[14].expect("first RSI emitted at index period");
503 assert_relative_eq!(first, 70.464, epsilon = 0.05);
504 }
505
506 #[test]
507 fn rsi_stays_in_0_100_range() {
508 let prices: Vec<f64> = (0..200)
509 .map(|i| 100.0 + (f64::from(i) * 0.7).sin() * 10.0)
510 .collect();
511 let mut rsi = Rsi::new(14).unwrap();
512 for x in rsi.batch(&prices).into_iter().flatten() {
513 assert!((0.0..=100.0).contains(&x), "RSI out of range: {x}");
514 }
515 }
516
517 #[test]
518 fn reset_clears_state() {
519 let mut rsi = Rsi::new(5).unwrap();
520 rsi.batch(&[1.0, 2.0, 3.0, 2.0, 4.0, 5.0, 6.0]);
521 assert!(rsi.is_ready());
522 rsi.reset();
523 assert!(!rsi.is_ready());
524 assert_eq!(rsi.update(1.0), None);
525 }
526
527 #[test]
528 fn batch_equals_streaming() {
529 let prices: Vec<f64> = (1..=40)
530 .map(|i| (f64::from(i) * 0.3).sin() * 5.0 + f64::from(i))
531 .collect();
532 let mut a = Rsi::new(7).unwrap();
533 let mut b = Rsi::new(7).unwrap();
534 assert_eq!(
535 a.batch(&prices),
536 prices.iter().map(|p| b.update(*p)).collect::<Vec<_>>()
537 );
538 }
539
540 #[test]
541 fn ignores_non_finite_input() {
542 let mut rsi = Rsi::new(3).unwrap();
543 rsi.batch(&[1.0, 2.0, 3.0, 4.0]);
544 let before = rsi.value();
545 assert!(before.is_some());
546 assert_eq!(rsi.update(f64::NAN), None);
547 assert_eq!(rsi.update(f64::INFINITY), None);
548 assert_eq!(rsi.value(), before);
549 }
550
551 fn bits_eq(a: &[f64], b: &[f64]) -> bool {
552 a.len() == b.len()
553 && a.iter()
554 .zip(b)
555 .all(|(x, y)| x == y || (x.is_nan() && y.is_nan()))
556 }
557
558 fn rsi_replay(period: usize, series: &[f64]) -> Vec<f64> {
559 let mut r = Rsi::new(period).unwrap();
560 series
561 .iter()
562 .map(|&x| r.update(x).unwrap_or(f64::NAN))
563 .collect()
564 }
565
566 #[test]
567 fn batch_nan_fast_path_is_bit_identical() {
568 let series: Vec<f64> = (0..300)
569 .map(|i| (f64::from(i) * 0.3).sin() * 5.0 + f64::from(i) * 0.1 + 100.0)
570 .collect();
571 let mut rsi = Rsi::new(14).unwrap();
572 let got = rsi.batch_nan(&series);
573 assert!(bits_eq(&got, &rsi_replay(14, &series)));
574 let mut ref_rsi = Rsi::new(14).unwrap();
575 for &x in &series {
576 ref_rsi.update(x);
577 }
578 assert_eq!(rsi.update(123.0), ref_rsi.update(123.0));
579 }
580
581 #[test]
584 fn wilder_tail_is_identical_on_every_dispatch_path() {
585 let series: Vec<f64> = (0..3000)
586 .map(|i| (f64::from(i) * 0.071).sin() * 9.0 + f64::from(i % 11) * 0.3 + 70.0)
587 .collect();
588 let rsi = Rsi::new(14).unwrap();
589 let mut a = vec![0.0; series.len() - 1];
590 let mut b = vec![0.0; series.len() - 1];
591 let state = (series[0], 0.4, 0.6);
592 let ra = wickra_simd::dispatch(WilderTail {
593 inputs: &series[1..],
594 out: &mut a,
595 state,
596 n_minus_1: rsi.n_minus_1,
597 inv_period: rsi.inv_period,
598 });
599 let rb = wickra_simd::run_baseline(WilderTail {
600 inputs: &series[1..],
601 out: &mut b,
602 state,
603 n_minus_1: rsi.n_minus_1,
604 inv_period: rsi.inv_period,
605 });
606 let bits = |v: &[f64]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
607 assert_eq!(bits(&a), bits(&b));
608 assert_eq!(bits(&[ra.0, ra.1, ra.2]), bits(&[rb.0, rb.1, rb.2]));
609 }
610
611 #[test]
614 fn batch_nan_into_overwrites_a_dirty_buffer() {
615 let series: Vec<f64> = (0..200)
616 .map(|i| (f64::from(i) * 0.3).sin() * 5.0 + 60.0)
617 .collect();
618 let mut out = vec![7.0; series.len()];
619 Rsi::new(14).unwrap().batch_nan_into(&series, &mut out);
620 assert!(bits_eq(&out, &rsi_replay(14, &series)));
621 }
622
623 #[test]
624 fn batch_nan_falls_back_on_non_finite() {
625 let series = [10.0, 11.0, 9.0, f64::NAN, 12.0, 13.0, 8.0];
626 let mut rsi = Rsi::new(3).unwrap();
627 assert!(bits_eq(&rsi.batch_nan(&series), &rsi_replay(3, &series)));
628 }
629
630 #[test]
631 fn batch_nan_falls_back_when_not_fresh() {
632 let mut rsi = Rsi::new(3).unwrap();
633 rsi.update(50.0);
634 let series = [51.0, 49.0, 52.0, 53.0, 50.0];
635 let mut ref_rsi = Rsi::new(3).unwrap();
636 ref_rsi.update(50.0);
637 let want: Vec<f64> = series
638 .iter()
639 .map(|&x| ref_rsi.update(x).unwrap_or(f64::NAN))
640 .collect();
641 assert!(bits_eq(&rsi.batch_nan(&series), &want));
642 }
643
644 #[test]
645 fn batch_nan_too_short_to_seed_falls_back() {
646 let series = [10.0, 11.0, 12.0];
648 let mut rsi = Rsi::new(3).unwrap();
649 assert!(bits_eq(&rsi.batch_nan(&series), &rsi_replay(3, &series)));
650 }
651
652 proptest::proptest! {
653 #![proptest_config(proptest::test_runner::Config::with_cases(48))]
654 #[test]
655 fn rsi_matches_naive(
656 period in 1usize..20,
657 prices in proptest::collection::vec(1.0_f64..1000.0, 0..150),
658 ) {
659 let mut rsi = Rsi::new(period).unwrap();
660 let got = rsi.batch(&prices);
661 let want = rsi_naive(&prices, period);
662 proptest::prop_assert_eq!(got.len(), want.len());
663 for (g, w) in got.iter().zip(want.iter()) {
664 match (g, w) {
665 (None, None) => {}
666 (Some(a), Some(b)) => proptest::prop_assert!(
667 (a - b).abs() < 1e-7,
668 "got={a} want={b}"
669 ),
670 _ => proptest::prop_assert!(false, "warmup mismatch"),
671 }
672 }
673 }
674 }
675}