Skip to main content

kestrel_chartkit/indicator/
wavetrend.rs

1use super::smoothing::{crossed_over, crossed_under, Ema, Sma};
2use super::{Indicator, IndicatorAlert, IndicatorOutput};
3use crate::model::Bar;
4use std::collections::HashMap;
5
6#[cfg(feature = "serde")]
7use serde::{Deserialize, Serialize};
8
9/// WaveTrend Alerts data structure.
10#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
11#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
12pub struct WaveTrendAlerts {
13    pub bull_cross: bool,
14    pub bear_cross: bool,
15    pub overbought_cross: bool,
16    pub oversold_cross: bool,
17}
18
19/// WaveTrend oscillator.
20///
21/// ```text
22/// ap  = (high + low + close) / 3
23/// esa = Ema(n1)(ap)          d = Ema(n1)(|ap - esa|)
24/// ci  = (ap - esa) / (0.015 * d)                        (0 for d <= 1e-8)
25/// wt1 = Ema(n2)(ci)          wt2 = SMA(4)(wt1)
26/// ```
27///
28/// All EMAs are the shared [`Ema`] with its first-sample seed, running from the first bar.
29/// `value` and `extra["wt1"]`: wt1; `extra["wt2"]`, `extra["hist"]` (`wt1 - wt2`) and the two
30/// levels. First output: with the fourth bar, once `wt2` exists. [`Indicator::reset`] clears all
31/// averages.
32#[derive(Debug, Clone)]
33pub struct WaveTrendEngine {
34    n1: usize,
35    n2: usize,
36    ob_level: f64,
37    os_level: f64,
38    ema_ap: Ema,
39    ema_d: Ema,
40    ema_wt1: Ema,
41    sma_wt2: Sma,
42    prev_wt1: Option<f64>,
43    prev_wt2: Option<f64>,
44    alerts: WaveTrendAlerts,
45}
46
47impl WaveTrendEngine {
48    pub fn new(n1: usize, n2: usize, ob_level: f64, os_level: f64) -> Self {
49        Self {
50            n1,
51            n2,
52            ob_level,
53            os_level,
54            ema_ap: Ema::new(n1),
55            ema_d: Ema::new(n1),
56            ema_wt1: Ema::new(n2),
57            sma_wt2: Sma::new(4),
58            prev_wt1: None,
59            prev_wt2: None,
60            alerts: WaveTrendAlerts::default(),
61        }
62    }
63
64    pub fn with_defaults() -> Self {
65        Self::new(10, 21, 60.0, -60.0)
66    }
67
68    pub fn alerts(&self) -> WaveTrendAlerts {
69        self.alerts
70    }
71}
72
73impl Indicator for WaveTrendEngine {
74    fn name(&self) -> &str {
75        "wavetrend"
76    }
77
78    fn warmup_period(&self) -> usize {
79        self.n1 + self.n2 + 4
80    }
81
82    fn on_bar(&mut self, bar: &Bar) -> Option<IndicatorOutput> {
83        let ap = bar.typical_price();
84        let esa = self.ema_ap.update(ap)?;
85        let d = self.ema_d.update((ap - esa).abs())?;
86
87        let ci = if d > 1e-8 {
88            (ap - esa) / (0.015 * d)
89        } else {
90            0.0
91        };
92
93        let wt1 = self.ema_wt1.update(ci)?;
94        let wt2 = self.sma_wt2.update(wt1)?;
95
96        self.alerts = WaveTrendAlerts::default();
97
98        if let (Some(p_wt1), Some(p_wt2)) = (self.prev_wt1, self.prev_wt2) {
99            self.alerts.bull_cross = crossed_over(p_wt1, p_wt2, wt1, wt2);
100            self.alerts.bear_cross = crossed_under(p_wt1, p_wt2, wt1, wt2);
101            self.alerts.overbought_cross = self.alerts.bear_cross && (wt1 >= self.ob_level);
102            self.alerts.oversold_cross = self.alerts.bull_cross && (wt1 <= self.os_level);
103        }
104
105        self.prev_wt1 = Some(wt1);
106        self.prev_wt2 = Some(wt2);
107
108        let mut extra = HashMap::new();
109        extra.insert("wt1".to_string(), wt1);
110        extra.insert("wt2".to_string(), wt2);
111        extra.insert("hist".to_string(), wt1 - wt2);
112        extra.insert("ob_level".to_string(), self.ob_level);
113        extra.insert("os_level".to_string(), self.os_level);
114
115        Some(IndicatorOutput::with_extra(wt1, extra))
116    }
117
118    fn reset(&mut self) {
119        self.ema_ap.reset();
120        self.ema_d.reset();
121        self.ema_wt1.reset();
122        self.sma_wt2.reset();
123        self.prev_wt1 = None;
124        self.prev_wt2 = None;
125        self.alerts = WaveTrendAlerts::default();
126    }
127
128    fn alerts(&self) -> Vec<IndicatorAlert> {
129        let mut res = Vec::new();
130        if self.alerts.bull_cross {
131            res.push(IndicatorAlert::new(
132                "wt_bull_cross",
133                "WaveTrend Bullish Cross",
134                0.8,
135            ));
136        }
137        if self.alerts.bear_cross {
138            res.push(IndicatorAlert::new(
139                "wt_bear_cross",
140                "WaveTrend Bearish Cross",
141                0.8,
142            ));
143        }
144        res
145    }
146}
147
148#[cfg(test)]
149mod tests {
150    use super::*;
151
152    #[test]
153    fn test_wavetrend_basic() {
154        let mut wt = WaveTrendEngine::with_defaults();
155        let mut outputs = Vec::new();
156        for i in 0..50 {
157            let b = Bar::new(i, 100.0, 105.0, 95.0, 100.0 + (i as f64 * 0.5), 1000.0);
158            if let Some(out) = wt.on_bar(&b) {
159                outputs.push(out);
160            }
161        }
162        assert!(!outputs.is_empty());
163        let last = outputs.last().unwrap();
164        assert!(last.extra.contains_key("wt1"));
165        assert!(last.extra.contains_key("wt2"));
166    }
167}