use serde::{Deserialize, Serialize};
use crate::indicator::{
config::RsiWindow,
streaming::{StreamingIndicator, moving_averages::StreamingEwm},
};
#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
pub struct StreamingRsi {
prev_price: Option<f64>,
avg_gain: StreamingEwm,
avg_loss: StreamingEwm,
}
impl StreamingRsi {
#[must_use]
pub fn new(window_size: RsiWindow) -> Self {
let size = window_size.0 as usize;
let size_u32 = u32::from(window_size.0);
let alpha = 1.0 / f64::from(size_u32);
let win = size;
Self {
prev_price: None,
avg_gain: StreamingEwm::new(alpha, win),
avg_loss: StreamingEwm::new(alpha, win),
}
}
}
impl StreamingIndicator for StreamingRsi {
type Input = f64;
type Output<'a> = Option<f64>;
fn update(&mut self, value: Self::Input) -> Self::Output<'_> {
let Some(prev) = self.prev_price else {
self.prev_price = Some(value);
return None;
};
let delta = value - prev;
self.prev_price = Some(value);
let (gain, loss) = if delta > 0.0 {
(delta, 0.0)
} else {
(0.0, delta.abs())
};
let g_val = self.avg_gain.update(gain);
let l_val = self.avg_loss.update(loss);
match (g_val, l_val) {
(Some(avg_gain), Some(avg_loss)) => {
if avg_loss == 0.0 {
if avg_gain == 0.0 {
Some(50.0)
} else {
Some(100.0)
}
} else {
let rs = avg_gain / avg_loss;
Some(100.0 - (100.0 / (1.0 + rs)))
}
}
_ => None,
}
}
fn reset(&mut self) {
self.prev_price = None;
self.avg_gain.reset();
self.avg_loss.reset();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn first_value_only_seeds_and_returns_none() {
let mut rsi = StreamingRsi::new(RsiWindow(3));
assert_eq!(rsi.update(100.0), None);
}
#[test]
fn warmup_takes_window_plus_one_prices() {
let mut rsi = StreamingRsi::new(RsiWindow(3));
assert_eq!(rsi.update(10.0), None); assert_eq!(rsi.update(11.0), None); assert_eq!(rsi.update(12.0), None); assert!(rsi.update(13.0).is_some()); }
#[test]
fn pure_uptrend_gives_100() {
let mut rsi = StreamingRsi::new(RsiWindow(3));
let mut last = None;
for p in [1.0, 2.0, 3.0, 4.0, 5.0] {
last = rsi.update(p);
}
assert_eq!(last, Some(100.0));
}
#[test]
fn pure_downtrend_gives_0() {
let mut rsi = StreamingRsi::new(RsiWindow(3));
let mut last = None;
for p in [5.0, 4.0, 3.0, 2.0, 1.0] {
last = rsi.update(p);
}
assert_eq!(last, Some(0.0));
}
#[test]
fn flat_line_gives_50() {
let mut rsi = StreamingRsi::new(RsiWindow(3));
let mut last = None;
for _ in 0..5 {
last = rsi.update(100.0);
}
assert_eq!(last, Some(50.0));
}
#[test]
fn output_stays_within_bounds() {
let mut rsi = StreamingRsi::new(RsiWindow(4));
let prices = [10.0, 12.0, 11.0, 15.0, 9.0, 20.0, 8.0, 13.0, 14.0, 7.0];
for p in prices {
if let Some(v) = rsi.update(p) {
assert!((0.0..=100.0).contains(&v), "RSI {v} out of bounds");
}
}
}
#[test]
fn reset_clears_state() {
let mut rsi = StreamingRsi::new(RsiWindow(2));
for p in [1.0, 2.0, 3.0, 4.0] {
rsi.update(p);
}
rsi.reset();
assert_eq!(rsi.update(50.0), None); assert_eq!(rsi.update(49.0), None); assert!(rsi.update(48.0).is_some()); }
}