Skip to main content

kestrel_chartkit/indicator/
stoch_rsi.rs

1use std::collections::HashMap;
2
3use crate::model::Bar;
4
5use super::divergence::SlopeDivergence;
6use super::smoothing::{crossed_over, crossed_under, ExtremeWindow, Rma, Sma};
7use super::{Indicator, IndicatorAlert, IndicatorOutput};
8
9/// Stochastic RSI with a smoothed %K and a %D signal.
10///
11/// The raw Wilder RSI over `rsi_len` changes — not the smoothed line of [`super::rsi::Rsi`] — is
12/// placed in the range of its last `stoch_len` values, `100 * (rsi - min) / (max - min)`, `50` for a
13/// flat range. `value` is `%K`, an `SMA(k_len)` of that; `extra["signal"]` is `%D`, an
14/// `SMA(d_len)` of `%K`; both clamped to `0..=100`. A context line over longer windows serves the
15/// divergence alerts.
16///
17/// First output: with the `rsi_len + stoch_len + k_len - 1`-th bar. [`Indicator::reset`] clears all
18/// state.
19pub struct StochRsi {
20    mid_line: f64,
21    oversold: f64,
22    overbought: f64,
23    require_extreme_zone: bool,
24
25    prev_close: Option<f64>,
26    avg_gain: Rma,
27    avg_loss: Rma,
28    stoch_window: ExtremeWindow,
29    k_avg: Sma,
30    signal_avg: Sma,
31    extreme_window: ExtremeWindow,
32    prev_k_line: Option<f64>,
33    prev_signal: Option<f64>,
34    bars_seen: usize,
35    warmup_period: usize,
36
37    ctx_avg_gain: Rma,
38    ctx_avg_loss: Rma,
39    ctx_stoch_window: ExtremeWindow,
40    ctx_k_avg: Sma,
41    ctx_bars_seen: usize,
42    ctx_warmup_period: usize,
43    divergence: SlopeDivergence,
44
45    alerts: StochRsiAlerts,
46}
47
48#[derive(Debug, Clone, Copy, PartialEq, Default)]
49pub struct StochRsiAlerts {
50    pub bull_extreme: bool,
51    pub bear_extreme: bool,
52    pub bull_mid_cross: bool,
53    pub bear_mid_cross: bool,
54    pub bull_divergence: bool,
55    pub bear_divergence: bool,
56    pub extreme_strength: f64,
57    pub divergence_strength: f64,
58}
59
60impl StochRsi {
61    #[allow(clippy::too_many_arguments)]
62    pub fn new(
63        rsi_len: usize,
64        stoch_len: usize,
65        k_len: usize,
66        d_len: usize,
67        mid_line: f64,
68        overbought: f64,
69        oversold: f64,
70        lookback_extreme: usize,
71        require_extreme_zone: bool,
72        ctx_rsi_len: usize,
73        ctx_stoch_len: usize,
74        div_len: usize,
75        div_min: f64,
76    ) -> Self {
77        Self {
78            mid_line,
79            oversold,
80            overbought,
81            require_extreme_zone,
82            prev_close: None,
83            avg_gain: Rma::new(rsi_len),
84            avg_loss: Rma::new(rsi_len),
85            stoch_window: ExtremeWindow::new(stoch_len),
86            k_avg: Sma::new(k_len),
87            signal_avg: Sma::new(d_len),
88            extreme_window: ExtremeWindow::new(lookback_extreme),
89            prev_k_line: None,
90            prev_signal: None,
91            bars_seen: 0,
92            warmup_period: rsi_len + 1 + stoch_len.saturating_sub(1) + k_len.saturating_sub(1),
93            ctx_avg_gain: Rma::new(ctx_rsi_len),
94            ctx_avg_loss: Rma::new(ctx_rsi_len),
95            ctx_stoch_window: ExtremeWindow::new(ctx_stoch_len),
96            ctx_k_avg: Sma::new(k_len),
97            ctx_bars_seen: 0,
98            ctx_warmup_period: ctx_rsi_len
99                + 1
100                + ctx_stoch_len.saturating_sub(1)
101                + k_len.saturating_sub(1),
102            divergence: SlopeDivergence::new(div_len, div_min),
103            alerts: StochRsiAlerts::default(),
104        }
105    }
106
107    pub fn with_defaults() -> Self {
108        Self::new(14, 14, 3, 3, 50.0, 80.0, 20.0, 5, true, 50, 50, 4, 10.0)
109    }
110}
111
112impl Indicator for StochRsi {
113    fn name(&self) -> &str {
114        "stoch_rsi"
115    }
116
117    fn warmup_period(&self) -> usize {
118        self.warmup_period.max(self.ctx_warmup_period)
119    }
120
121    fn on_bar(&mut self, bar: &Bar) -> Option<IndicatorOutput> {
122        self.alerts = StochRsiAlerts::default();
123        self.bars_seen += 1;
124
125        let close = bar.close;
126        let prev_close = match self.prev_close {
127            None => {
128                self.prev_close = Some(close);
129                return None;
130            }
131            Some(p) => p,
132        };
133        self.prev_close = Some(close);
134
135        let change = close - prev_close;
136        let gain = change.max(0.0);
137        let loss = (-change).max(0.0);
138
139        self.ctx_bars_seen += 1;
140        let ctx_line = match (
141            self.ctx_avg_gain.update(gain),
142            self.ctx_avg_loss.update(loss),
143        ) {
144            (Some(ctx_avg_gain), Some(ctx_avg_loss)) => {
145                let ctx_rsi = if ctx_avg_gain == 0.0 && ctx_avg_loss == 0.0 {
146                    50.0
147                } else if ctx_avg_loss == 0.0 {
148                    100.0
149                } else if ctx_avg_gain == 0.0 {
150                    0.0
151                } else {
152                    100.0 - 100.0 / (1.0 + ctx_avg_gain / ctx_avg_loss)
153                };
154                self.ctx_stoch_window.push(ctx_rsi).and_then(|(low, high)| {
155                    let ctx_stoch_raw = if high > low {
156                        100.0 * (ctx_rsi - low) / (high - low)
157                    } else {
158                        50.0
159                    };
160                    self.ctx_k_avg.update(ctx_stoch_raw)
161                })
162            }
163            _ => None,
164        };
165
166        let (avg_gain, avg_loss) = match (self.avg_gain.update(gain), self.avg_loss.update(loss)) {
167            (Some(g), Some(l)) => (g, l),
168            _ => return None,
169        };
170
171        let rsi = if avg_gain == 0.0 && avg_loss == 0.0 {
172            50.0
173        } else if avg_loss == 0.0 {
174            100.0
175        } else if avg_gain == 0.0 {
176            0.0
177        } else {
178            100.0 - 100.0 / (1.0 + avg_gain / avg_loss)
179        };
180
181        let (low, high) = self.stoch_window.push(rsi)?;
182        let stoch_raw = if high > low {
183            (100.0 * (rsi - low) / (high - low)).clamp(0.0, 100.0)
184        } else {
185            50.0
186        };
187
188        let k_line = self.k_avg.update(stoch_raw)?.clamp(0.0, 100.0);
189        let signal = self.signal_avg.update(k_line)?.clamp(0.0, 100.0);
190
191        let extreme = self.extreme_window.push(k_line);
192        let was_oversold = extreme
193            .map(|(low, _)| low <= self.oversold)
194            .unwrap_or(false);
195        let was_overbought = extreme
196            .map(|(_, high)| high >= self.overbought)
197            .unwrap_or(false);
198
199        if let (Some(prev_k), Some(prev_sig)) = (self.prev_k_line, self.prev_signal) {
200            let bull_cross = crossed_over(prev_k, prev_sig, k_line, signal);
201            let bear_cross = crossed_under(prev_k, prev_sig, k_line, signal);
202            self.alerts.bull_extreme = bull_cross && (!self.require_extreme_zone || was_oversold);
203            self.alerts.bear_extreme = bear_cross && (!self.require_extreme_zone || was_overbought);
204            self.alerts.bull_mid_cross = crossed_over(prev_k, self.mid_line, k_line, self.mid_line);
205            self.alerts.bear_mid_cross =
206                crossed_under(prev_k, self.mid_line, k_line, self.mid_line);
207
208            self.alerts.extreme_strength = if self.alerts.bull_extreme {
209                let lowest = extreme.map(|(low, _)| low).unwrap_or(k_line);
210                ((self.oversold - lowest) / self.oversold.abs()).clamp(0.0, 1.0)
211            } else if self.alerts.bear_extreme {
212                let highest = extreme.map(|(_, high)| high).unwrap_or(k_line);
213                ((highest - self.overbought) / self.overbought.abs()).clamp(0.0, 1.0)
214            } else {
215                0.0
216            };
217        }
218        self.prev_k_line = Some(k_line);
219        self.prev_signal = Some(signal);
220
221        let mut extra = HashMap::new();
222        extra.insert("signal".to_string(), signal);
223        if let Some(ctx_line) = ctx_line {
224            let div = self.divergence.update(k_line, ctx_line);
225            self.alerts.bull_divergence = div.bull;
226            self.alerts.bear_divergence = div.bear;
227            self.alerts.divergence_strength = if div.bull || div.bear {
228                ((div.fast_dir.abs() - self.divergence.div_min()) / self.divergence.div_min())
229                    .clamp(0.0, 1.0)
230            } else {
231                0.0
232            };
233            extra.insert("ctx".to_string(), ctx_line);
234        }
235
236        Some(IndicatorOutput::with_extra(k_line, extra))
237    }
238
239    fn reset(&mut self) {
240        self.prev_close = None;
241        self.avg_gain.reset();
242        self.avg_loss.reset();
243        self.stoch_window.reset();
244        self.k_avg.reset();
245        self.signal_avg.reset();
246        self.extreme_window.reset();
247        self.prev_k_line = None;
248        self.prev_signal = None;
249        self.bars_seen = 0;
250        self.ctx_avg_gain.reset();
251        self.ctx_avg_loss.reset();
252        self.ctx_stoch_window.reset();
253        self.ctx_k_avg.reset();
254        self.ctx_bars_seen = 0;
255        self.divergence.reset();
256        self.alerts = StochRsiAlerts::default();
257    }
258
259    fn alerts(&self) -> Vec<IndicatorAlert> {
260        let a = self.alerts;
261        let mut out = Vec::new();
262        if a.bull_extreme {
263            out.push(IndicatorAlert {
264                kind: "bull_extreme".to_string(),
265                note: "SRSI · BULL CROSS OVERSOLD".to_string(),
266                strength: a.extreme_strength,
267            });
268        }
269        if a.bear_extreme {
270            out.push(IndicatorAlert {
271                kind: "bear_extreme".to_string(),
272                note: "SRSI · BEAR CROSS OVERBOUGHT".to_string(),
273                strength: a.extreme_strength,
274            });
275        }
276        if a.bull_mid_cross {
277            out.push(IndicatorAlert {
278                kind: "bull_mid_cross".to_string(),
279                note: "SRSI · CROSS ABOVE 50".to_string(),
280                strength: 1.0,
281            });
282        }
283        if a.bear_mid_cross {
284            out.push(IndicatorAlert {
285                kind: "bear_mid_cross".to_string(),
286                note: "SRSI · CROSS BELOW 50".to_string(),
287                strength: 1.0,
288            });
289        }
290        if a.bull_divergence {
291            out.push(IndicatorAlert {
292                kind: "bull_divergence".to_string(),
293                note: "STOCH RSI · BULL DIVERGENCE".to_string(),
294                strength: a.divergence_strength,
295            });
296        }
297        if a.bear_divergence {
298            out.push(IndicatorAlert {
299                kind: "bear_divergence".to_string(),
300                note: "STOCH RSI · BEAR DIVERGENCE".to_string(),
301                strength: a.divergence_strength,
302            });
303        }
304        out
305    }
306}