Skip to main content

kestrel_chartkit/indicator/
rsi.rs

1use std::collections::HashMap;
2
3use crate::model::Bar;
4
5use super::divergence::SlopeDivergence;
6use super::smoothing::{crossed_over, crossed_under, Ema, ExtremeWindow, Rma};
7use super::{Indicator, IndicatorAlert, IndicatorOutput};
8
9#[derive(Debug, Clone)]
10pub struct Rsi {
11    mid_line: f64,
12    oversold: f64,
13    overbought: f64,
14    require_extreme_zone: bool,
15    ctx_len: usize,
16
17    prev_close: Option<f64>,
18    avg_gain: Rma,
19    avg_loss: Rma,
20    rsi_avg: Ema,
21    signal_avg: Ema,
22    extreme_window: ExtremeWindow,
23    prev_rsi_line: Option<f64>,
24    prev_signal: Option<f64>,
25    bars_seen: usize,
26    warmup_period: usize,
27
28    ctx_avg_gain: Rma,
29    ctx_avg_loss: Rma,
30    ctx_avg: Ema,
31    divergence: SlopeDivergence,
32
33    alerts: RsiAlerts,
34}
35
36#[derive(Debug, Clone, Copy, PartialEq, Default)]
37pub struct RsiAlerts {
38    pub bull_extreme: bool,
39    pub bear_extreme: bool,
40    pub bull_mid_cross: bool,
41    pub bear_mid_cross: bool,
42    pub bull_divergence: bool,
43    pub bear_divergence: bool,
44    pub extreme_strength: f64,
45    pub divergence_strength: f64,
46}
47
48impl Rsi {
49    #[allow(clippy::too_many_arguments)]
50    pub fn new(
51        rsi_len: usize,
52        avg_len: usize,
53        sig_len: usize,
54        mid_line: f64,
55        overbought: f64,
56        oversold: f64,
57        lookback_extreme: usize,
58        require_extreme_zone: bool,
59        ctx_len: usize,
60        div_len: usize,
61        div_min: f64,
62    ) -> Self {
63        Self {
64            mid_line,
65            oversold,
66            overbought,
67            require_extreme_zone,
68            ctx_len,
69            prev_close: None,
70            avg_gain: Rma::new(rsi_len),
71            avg_loss: Rma::new(rsi_len),
72            rsi_avg: Ema::new(avg_len),
73            signal_avg: Ema::new(sig_len),
74            extreme_window: ExtremeWindow::new(lookback_extreme),
75            prev_rsi_line: None,
76            prev_signal: None,
77            bars_seen: 0,
78            warmup_period: rsi_len + 1,
79            ctx_avg_gain: Rma::new(ctx_len),
80            ctx_avg_loss: Rma::new(ctx_len),
81            ctx_avg: Ema::new(avg_len),
82            divergence: SlopeDivergence::new(div_len, div_min),
83            alerts: RsiAlerts::default(),
84        }
85    }
86
87    pub fn with_defaults() -> Self {
88        Self::new(14, 3, 3, 50.0, 70.0, 30.0, 5, true, 100, 4, 10.0)
89    }
90
91    pub fn with_period(rsi_len: usize) -> Self {
92        Self::new(rsi_len, 3, 3, 50.0, 70.0, 30.0, 5, true, 100, 4, 10.0)
93    }
94}
95
96impl Indicator for Rsi {
97    fn name(&self) -> &str {
98        "rsi"
99    }
100
101    fn warmup_period(&self) -> usize {
102        self.warmup_period.max(self.ctx_len + 1)
103    }
104
105    fn on_bar(&mut self, bar: &Bar) -> Option<IndicatorOutput> {
106        self.alerts = RsiAlerts::default();
107        self.bars_seen += 1;
108
109        let close = bar.close;
110        let prev_close = match self.prev_close {
111            None => {
112                self.prev_close = Some(close);
113                return None;
114            }
115            Some(p) => p,
116        };
117        self.prev_close = Some(close);
118
119        let change = close - prev_close;
120        let gain = change.max(0.0);
121        let loss = (-change).max(0.0);
122
123        let ctx_line = match (
124            self.ctx_avg_gain.update(gain),
125            self.ctx_avg_loss.update(loss),
126        ) {
127            (Some(ctx_avg_gain), Some(ctx_avg_loss)) => {
128                let ctx_raw = if ctx_avg_gain == 0.0 && ctx_avg_loss == 0.0 {
129                    50.0
130                } else if ctx_avg_loss == 0.0 {
131                    100.0
132                } else if ctx_avg_gain == 0.0 {
133                    0.0
134                } else {
135                    100.0 - 100.0 / (1.0 + ctx_avg_gain / ctx_avg_loss)
136                };
137                Some(self.ctx_avg.update(ctx_raw))
138            }
139            _ => None,
140        };
141
142        let (avg_gain, avg_loss) = match (self.avg_gain.update(gain), self.avg_loss.update(loss)) {
143            (Some(g), Some(l)) => (g, l),
144            _ => return None,
145        };
146
147        let raw_rsi = if avg_gain == 0.0 && avg_loss == 0.0 {
148            50.0
149        } else if avg_loss == 0.0 {
150            100.0
151        } else if avg_gain == 0.0 {
152            0.0
153        } else {
154            (100.0 - 100.0 / (1.0 + avg_gain / avg_loss)).clamp(0.0, 100.0)
155        };
156
157        let rsi_line = self.rsi_avg.update(raw_rsi).clamp(0.0, 100.0);
158        let signal = self.signal_avg.update(rsi_line).clamp(0.0, 100.0);
159
160        let extreme = self.extreme_window.push(rsi_line);
161        let was_oversold = extreme
162            .map(|(low, _)| low <= self.oversold)
163            .unwrap_or(false);
164        let was_overbought = extreme
165            .map(|(_, high)| high >= self.overbought)
166            .unwrap_or(false);
167
168        if let (Some(prev_rsi), Some(prev_sig)) = (self.prev_rsi_line, self.prev_signal) {
169            let bull_cross = crossed_over(prev_rsi, prev_sig, rsi_line, signal);
170            let bear_cross = crossed_under(prev_rsi, prev_sig, rsi_line, signal);
171            self.alerts.bull_extreme = bull_cross && (!self.require_extreme_zone || was_oversold);
172            self.alerts.bear_extreme = bear_cross && (!self.require_extreme_zone || was_overbought);
173            self.alerts.bull_mid_cross =
174                crossed_over(prev_rsi, self.mid_line, rsi_line, self.mid_line);
175            self.alerts.bear_mid_cross =
176                crossed_under(prev_rsi, self.mid_line, rsi_line, self.mid_line);
177
178            self.alerts.extreme_strength = if self.alerts.bull_extreme {
179                extreme
180                    .map(|(low, _)| ((self.oversold - low) / self.oversold.abs()).clamp(0.0, 1.0))
181                    .unwrap_or(0.0)
182            } else if self.alerts.bear_extreme {
183                extreme
184                    .map(|(_, high)| {
185                        ((high - self.overbought) / self.overbought.abs()).clamp(0.0, 1.0)
186                    })
187                    .unwrap_or(0.0)
188            } else {
189                0.0
190            };
191        }
192        self.prev_rsi_line = Some(rsi_line);
193        self.prev_signal = Some(signal);
194
195        let mut extra = HashMap::new();
196        extra.insert("signal".to_string(), signal);
197        if let Some(ctx_line) = ctx_line {
198            let div = self.divergence.update(rsi_line, ctx_line);
199            self.alerts.bull_divergence = div.bull;
200            self.alerts.bear_divergence = div.bear;
201            self.alerts.divergence_strength = if div.bull || div.bear {
202                ((div.fast_dir.abs() - self.divergence.div_min()) / self.divergence.div_min())
203                    .clamp(0.0, 1.0)
204            } else {
205                0.0
206            };
207            extra.insert("ctx".to_string(), ctx_line);
208        }
209
210        Some(IndicatorOutput::with_extra(rsi_line, extra))
211    }
212
213    fn reset(&mut self) {
214        self.prev_close = None;
215        self.avg_gain.reset();
216        self.avg_loss.reset();
217        self.rsi_avg.reset();
218        self.signal_avg.reset();
219        self.extreme_window.reset();
220        self.prev_rsi_line = None;
221        self.prev_signal = None;
222        self.bars_seen = 0;
223        self.ctx_avg_gain.reset();
224        self.ctx_avg_loss.reset();
225        self.ctx_avg.reset();
226        self.divergence.reset();
227        self.alerts = RsiAlerts::default();
228    }
229
230    fn alerts(&self) -> Vec<IndicatorAlert> {
231        let a = self.alerts;
232        let mut out = Vec::new();
233        if a.bull_extreme {
234            out.push(IndicatorAlert {
235                kind: "bull_extreme".to_string(),
236                note: "RSI · BULL CROSS OVERSOLD".to_string(),
237                strength: a.extreme_strength,
238            });
239        }
240        if a.bear_extreme {
241            out.push(IndicatorAlert {
242                kind: "bear_extreme".to_string(),
243                note: "RSI · BEAR CROSS OVERBOUGHT".to_string(),
244                strength: a.extreme_strength,
245            });
246        }
247        if a.bull_mid_cross {
248            out.push(IndicatorAlert {
249                kind: "bull_mid_cross".to_string(),
250                note: "RSI · CROSS ABOVE 50".to_string(),
251                strength: 1.0,
252            });
253        }
254        if a.bear_mid_cross {
255            out.push(IndicatorAlert {
256                kind: "bear_mid_cross".to_string(),
257                note: "RSI · CROSS BELOW 50".to_string(),
258                strength: 1.0,
259            });
260        }
261        if a.bull_divergence {
262            out.push(IndicatorAlert {
263                kind: "bull_divergence".to_string(),
264                note: "RSI · BULL DIVERGENCE".to_string(),
265                strength: a.divergence_strength,
266            });
267        }
268        if a.bear_divergence {
269            out.push(IndicatorAlert {
270                kind: "bear_divergence".to_string(),
271                note: "RSI · BEAR DIVERGENCE".to_string(),
272                strength: a.divergence_strength,
273            });
274        }
275        out
276    }
277}