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