1use std::collections::HashMap;
2
3#[cfg(feature = "serde")]
4use serde::{Deserialize, Serialize};
5
6use crate::model::Bar;
7
8use super::divergence::SlopeDivergence;
9use super::smoothing::{crossed_over, crossed_under, Ema, ExtremeWindow, Rma};
10use super::{Indicator, IndicatorAlert, IndicatorOutput};
11
12#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
18#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
19pub enum RsiSmoothing {
20 #[default]
23 Wilder,
24 Ema,
30}
31
32#[derive(Debug, Clone)]
39enum ChangeSmoother {
40 Wilder(Rma),
41 Ema { ema: Ema, len: usize, seen: usize },
42}
43
44impl ChangeSmoother {
45 fn new(method: RsiSmoothing, len: usize) -> Self {
46 match method {
47 RsiSmoothing::Wilder => Self::Wilder(Rma::new(len)),
48 RsiSmoothing::Ema => Self::Ema {
49 ema: Ema::new(len),
50 len,
51 seen: 0,
52 },
53 }
54 }
55
56 fn update(&mut self, change: f64) -> Option<f64> {
57 match self {
58 Self::Wilder(rma) => rma.update(change),
59 Self::Ema { ema, len, seen } => {
60 let value = ema.update(change)?;
61 *seen += 1;
62 (*seen >= *len).then_some(value)
63 }
64 }
65 }
66
67 fn reset(&mut self) {
68 match self {
69 Self::Wilder(rma) => rma.reset(),
70 Self::Ema { ema, seen, .. } => {
71 ema.reset();
72 *seen = 0;
73 }
74 }
75 }
76}
77
78#[derive(Debug, Clone)]
97pub struct Rsi {
98 mid_line: f64,
99 oversold: f64,
100 overbought: f64,
101 require_extreme_zone: bool,
102 rsi_len: usize,
103 ctx_len: usize,
104 smoothing: RsiSmoothing,
105
106 prev_close: Option<f64>,
107 avg_gain: ChangeSmoother,
108 avg_loss: ChangeSmoother,
109 rsi_avg: Ema,
110 signal_avg: Ema,
111 extreme_window: ExtremeWindow,
112 prev_rsi_line: Option<f64>,
113 prev_signal: Option<f64>,
114 bars_seen: usize,
115 warmup_period: usize,
116
117 ctx_avg_gain: ChangeSmoother,
118 ctx_avg_loss: ChangeSmoother,
119 ctx_avg: Ema,
120 divergence: SlopeDivergence,
121
122 alerts: RsiAlerts,
123}
124
125#[derive(Debug, Clone, Copy, PartialEq, Default)]
126pub struct RsiAlerts {
127 pub bull_extreme: bool,
128 pub bear_extreme: bool,
129 pub bull_mid_cross: bool,
130 pub bear_mid_cross: bool,
131 pub bull_divergence: bool,
132 pub bear_divergence: bool,
133 pub extreme_strength: f64,
134 pub divergence_strength: f64,
135}
136
137impl Rsi {
138 #[allow(clippy::too_many_arguments)]
139 pub fn new(
140 rsi_len: usize,
141 avg_len: usize,
142 sig_len: usize,
143 mid_line: f64,
144 overbought: f64,
145 oversold: f64,
146 lookback_extreme: usize,
147 require_extreme_zone: bool,
148 ctx_len: usize,
149 div_len: usize,
150 div_min: f64,
151 ) -> Self {
152 Self {
153 mid_line,
154 oversold,
155 overbought,
156 require_extreme_zone,
157 rsi_len,
158 ctx_len,
159 smoothing: RsiSmoothing::Wilder,
160 prev_close: None,
161 avg_gain: ChangeSmoother::new(RsiSmoothing::Wilder, rsi_len),
162 avg_loss: ChangeSmoother::new(RsiSmoothing::Wilder, rsi_len),
163 rsi_avg: Ema::new(avg_len),
164 signal_avg: Ema::new(sig_len),
165 extreme_window: ExtremeWindow::new(lookback_extreme),
166 prev_rsi_line: None,
167 prev_signal: None,
168 bars_seen: 0,
169 warmup_period: rsi_len + 1,
170 ctx_avg_gain: ChangeSmoother::new(RsiSmoothing::Wilder, ctx_len),
171 ctx_avg_loss: ChangeSmoother::new(RsiSmoothing::Wilder, ctx_len),
172 ctx_avg: Ema::new(avg_len),
173 divergence: SlopeDivergence::new(div_len, div_min),
174 alerts: RsiAlerts::default(),
175 }
176 }
177
178 pub fn with_defaults() -> Self {
179 Self::new(14, 3, 3, 50.0, 70.0, 30.0, 5, true, 100, 4, 10.0)
180 }
181
182 pub fn with_period(rsi_len: usize) -> Self {
183 Self::new(rsi_len, 3, 3, 50.0, 70.0, 30.0, 5, true, 100, 4, 10.0)
184 }
185
186 pub fn with_smoothing(mut self, method: RsiSmoothing) -> Self {
194 self.smoothing = method;
195 self.avg_gain = ChangeSmoother::new(method, self.rsi_len);
196 self.avg_loss = ChangeSmoother::new(method, self.rsi_len);
197 self.ctx_avg_gain = ChangeSmoother::new(method, self.ctx_len);
198 self.ctx_avg_loss = ChangeSmoother::new(method, self.ctx_len);
199 self
200 }
201
202 pub fn smoothing(&self) -> RsiSmoothing {
203 self.smoothing
204 }
205}
206
207impl Indicator for Rsi {
208 fn name(&self) -> &str {
209 "rsi"
210 }
211
212 fn warmup_period(&self) -> usize {
213 self.warmup_period.max(self.ctx_len + 1)
214 }
215
216 fn on_bar(&mut self, bar: &Bar) -> Option<IndicatorOutput> {
217 self.alerts = RsiAlerts::default();
218 self.bars_seen += 1;
219
220 let close = bar.close;
221 let prev_close = match self.prev_close {
222 None => {
223 self.prev_close = Some(close);
224 return None;
225 }
226 Some(p) => p,
227 };
228 self.prev_close = Some(close);
229
230 let change = close - prev_close;
231 let gain = change.max(0.0);
232 let loss = (-change).max(0.0);
233
234 let ctx_line = match (
235 self.ctx_avg_gain.update(gain),
236 self.ctx_avg_loss.update(loss),
237 ) {
238 (Some(ctx_avg_gain), Some(ctx_avg_loss)) => {
239 let ctx_raw = if ctx_avg_gain == 0.0 && ctx_avg_loss == 0.0 {
240 50.0
241 } else if ctx_avg_loss == 0.0 {
242 100.0
243 } else if ctx_avg_gain == 0.0 {
244 0.0
245 } else {
246 100.0 - 100.0 / (1.0 + ctx_avg_gain / ctx_avg_loss)
247 };
248 self.ctx_avg.update(ctx_raw)
249 }
250 _ => None,
251 };
252
253 let (avg_gain, avg_loss) = match (self.avg_gain.update(gain), self.avg_loss.update(loss)) {
254 (Some(g), Some(l)) => (g, l),
255 _ => return None,
256 };
257
258 let raw_rsi = if avg_gain == 0.0 && avg_loss == 0.0 {
259 50.0
260 } else if avg_loss == 0.0 {
261 100.0
262 } else if avg_gain == 0.0 {
263 0.0
264 } else {
265 (100.0 - 100.0 / (1.0 + avg_gain / avg_loss)).clamp(0.0, 100.0)
266 };
267
268 let rsi_line = self.rsi_avg.update(raw_rsi)?.clamp(0.0, 100.0);
269 let signal = self.signal_avg.update(rsi_line)?.clamp(0.0, 100.0);
270
271 let extreme = self.extreme_window.push(rsi_line);
272 let was_oversold = extreme
273 .map(|(low, _)| low <= self.oversold)
274 .unwrap_or(false);
275 let was_overbought = extreme
276 .map(|(_, high)| high >= self.overbought)
277 .unwrap_or(false);
278
279 if let (Some(prev_rsi), Some(prev_sig)) = (self.prev_rsi_line, self.prev_signal) {
280 let bull_cross = crossed_over(prev_rsi, prev_sig, rsi_line, signal);
281 let bear_cross = crossed_under(prev_rsi, prev_sig, rsi_line, signal);
282 self.alerts.bull_extreme = bull_cross && (!self.require_extreme_zone || was_oversold);
283 self.alerts.bear_extreme = bear_cross && (!self.require_extreme_zone || was_overbought);
284 self.alerts.bull_mid_cross =
285 crossed_over(prev_rsi, self.mid_line, rsi_line, self.mid_line);
286 self.alerts.bear_mid_cross =
287 crossed_under(prev_rsi, self.mid_line, rsi_line, self.mid_line);
288
289 self.alerts.extreme_strength = if self.alerts.bull_extreme {
290 extreme
291 .map(|(low, _)| ((self.oversold - low) / self.oversold.abs()).clamp(0.0, 1.0))
292 .unwrap_or(0.0)
293 } else if self.alerts.bear_extreme {
294 extreme
295 .map(|(_, high)| {
296 ((high - self.overbought) / self.overbought.abs()).clamp(0.0, 1.0)
297 })
298 .unwrap_or(0.0)
299 } else {
300 0.0
301 };
302 }
303 self.prev_rsi_line = Some(rsi_line);
304 self.prev_signal = Some(signal);
305
306 let mut extra = HashMap::new();
307 extra.insert("signal".to_string(), signal);
308 if let Some(ctx_line) = ctx_line {
309 let div = self.divergence.update(rsi_line, ctx_line);
310 self.alerts.bull_divergence = div.bull;
311 self.alerts.bear_divergence = div.bear;
312 self.alerts.divergence_strength = if div.bull || div.bear {
313 ((div.fast_dir.abs() - self.divergence.div_min()) / self.divergence.div_min())
314 .clamp(0.0, 1.0)
315 } else {
316 0.0
317 };
318 extra.insert("ctx".to_string(), ctx_line);
319 }
320
321 Some(IndicatorOutput::with_extra(rsi_line, extra))
322 }
323
324 fn reset(&mut self) {
325 self.prev_close = None;
326 self.avg_gain.reset();
327 self.avg_loss.reset();
328 self.rsi_avg.reset();
329 self.signal_avg.reset();
330 self.extreme_window.reset();
331 self.prev_rsi_line = None;
332 self.prev_signal = None;
333 self.bars_seen = 0;
334 self.ctx_avg_gain.reset();
335 self.ctx_avg_loss.reset();
336 self.ctx_avg.reset();
337 self.divergence.reset();
338 self.alerts = RsiAlerts::default();
339 }
340
341 fn alerts(&self) -> Vec<IndicatorAlert> {
342 let a = self.alerts;
343 let mut out = Vec::new();
344 if a.bull_extreme {
345 out.push(IndicatorAlert {
346 kind: "bull_extreme".to_string(),
347 note: "RSI · BULL CROSS OVERSOLD".to_string(),
348 strength: a.extreme_strength,
349 });
350 }
351 if a.bear_extreme {
352 out.push(IndicatorAlert {
353 kind: "bear_extreme".to_string(),
354 note: "RSI · BEAR CROSS OVERBOUGHT".to_string(),
355 strength: a.extreme_strength,
356 });
357 }
358 if a.bull_mid_cross {
359 out.push(IndicatorAlert {
360 kind: "bull_mid_cross".to_string(),
361 note: "RSI · CROSS ABOVE 50".to_string(),
362 strength: 1.0,
363 });
364 }
365 if a.bear_mid_cross {
366 out.push(IndicatorAlert {
367 kind: "bear_mid_cross".to_string(),
368 note: "RSI · CROSS BELOW 50".to_string(),
369 strength: 1.0,
370 });
371 }
372 if a.bull_divergence {
373 out.push(IndicatorAlert {
374 kind: "bull_divergence".to_string(),
375 note: "RSI · BULL DIVERGENCE".to_string(),
376 strength: a.divergence_strength,
377 });
378 }
379 if a.bear_divergence {
380 out.push(IndicatorAlert {
381 kind: "bear_divergence".to_string(),
382 note: "RSI · BEAR DIVERGENCE".to_string(),
383 strength: a.divergence_strength,
384 });
385 }
386 out
387 }
388}