Skip to main content

kestrel_chartkit/indicator/
rsi.rs

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/// How the average up/down moves feeding the RSI ratio are smoothed.
13///
14/// This is a property of the RSI core itself and independent of `avg_len` (which smooths the
15/// finished RSI line) and `sig_len` (the signal line): those two keep their own smoothers in
16/// either mode.
17#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
18#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
19pub enum RsiSmoothing {
20    /// Wilder's own smoothing, `alpha = 1/N`, seeded with the SMA of the first `N` changes.
21    /// The default, and the historical behaviour of this indicator.
22    #[default]
23    Wilder,
24    /// Exponential smoothing, `alpha = 2/(N+1)`, seeded with the first actual change.
25    ///
26    /// A different formula, not a parametrisation of Wilder: `Ema(N)` reacts like
27    /// `Wilder(2N - 1)`. Matching a period against another implementation therefore does not by
28    /// itself produce matching values, and the seeds differ as well.
29    Ema,
30}
31
32/// Smooths the up/down moves for one RSI line, in whichever mode was selected.
33///
34/// Both modes report readiness the same way: `None` until `len` changes have been seen, so the
35/// RSI line starts at the same bar regardless of the mode. In `Ema` mode the internal EMA is
36/// already running before that — it is seeded with the first real change, never with an invented
37/// starting value — but its early, seed-dominated values are not published.
38#[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/// Relative Strength Index with a smoothed line, a signal line and a slower context line.
79///
80/// The raw RSI averages the up and down moves of the close over `rsi_len` changes — Wilder's way
81/// by default, exponentially under [`RsiSmoothing::Ema`] — and is `100 - 100 / (1 + up / down)`,
82/// with `50` when nothing moved, `100` without down moves and `0` without up moves.
83///
84/// **`value` is not that raw RSI.** It is the line: an `Ema(avg_len)` over the raw RSI with the
85/// first-sample seed of the shared [`Ema`], clamped to `0..=100`. With the registry default
86/// `avg_len = 3` the published RSI lags and deviates from the raw figure; `avg_len = 1` gives the
87/// raw RSI itself.
88///
89/// `extra["signal"]`: `Ema(sig_len)` over the line. `extra["ctx"]`: the same construction over
90/// `ctx_len` changes (default 100), present once that many changes exist; divergences are judged
91/// between line and context. Alerts fire on line/signal crosses (inside the extreme zones when
92/// `require_extreme_zone` is set), on crosses of `mid_line` and on divergences.
93///
94/// First output: once `rsi_len` changes exist, i.e. with the `rsi_len + 1`-th bar.
95/// [`Indicator::reset`] clears all averages.
96#[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    /// Selects how the up/down moves are smoothed; see [`RsiSmoothing`].
187    ///
188    /// Additive to the existing constructors, which keep the Wilder default. The choice applies
189    /// to the main line *and* the context line — each with its own period (`rsi_len`/`ctx_len`) —
190    /// so the divergence comparison is never between two differently smoothed series.
191    ///
192    /// Resets the smoothing state, so this belongs before the first bar, not mid-series.
193    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}