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 Engine (LazyBear / Pine classic formulation).
20#[derive(Debug, Clone)]
21pub struct WaveTrendEngine {
22    n1: usize,
23    n2: usize,
24    ob_level: f64,
25    os_level: f64,
26    ema_ap: Ema,
27    ema_d: Ema,
28    ema_wt1: Ema,
29    sma_wt2: Sma,
30    prev_wt1: Option<f64>,
31    prev_wt2: Option<f64>,
32    alerts: WaveTrendAlerts,
33}
34
35impl WaveTrendEngine {
36    pub fn new(n1: usize, n2: usize, ob_level: f64, os_level: f64) -> Self {
37        Self {
38            n1,
39            n2,
40            ob_level,
41            os_level,
42            ema_ap: Ema::new(n1),
43            ema_d: Ema::new(n1),
44            ema_wt1: Ema::new(n2),
45            sma_wt2: Sma::new(4),
46            prev_wt1: None,
47            prev_wt2: None,
48            alerts: WaveTrendAlerts::default(),
49        }
50    }
51
52    pub fn with_defaults() -> Self {
53        Self::new(10, 21, 60.0, -60.0)
54    }
55
56    pub fn alerts(&self) -> WaveTrendAlerts {
57        self.alerts
58    }
59}
60
61impl Indicator for WaveTrendEngine {
62    fn name(&self) -> &str {
63        "wavetrend"
64    }
65
66    fn warmup_period(&self) -> usize {
67        self.n1 + self.n2 + 4
68    }
69
70    fn on_bar(&mut self, bar: &Bar) -> Option<IndicatorOutput> {
71        let ap = bar.typical_price();
72        let esa = self.ema_ap.update(ap);
73        let d = self.ema_d.update((ap - esa).abs());
74
75        let ci = if d > 1e-8 {
76            (ap - esa) / (0.015 * d)
77        } else {
78            0.0
79        };
80
81        let wt1 = self.ema_wt1.update(ci);
82        let wt2 = self.sma_wt2.update(wt1)?;
83
84        self.alerts = WaveTrendAlerts::default();
85
86        if let (Some(p_wt1), Some(p_wt2)) = (self.prev_wt1, self.prev_wt2) {
87            self.alerts.bull_cross = crossed_over(p_wt1, p_wt2, wt1, wt2);
88            self.alerts.bear_cross = crossed_under(p_wt1, p_wt2, wt1, wt2);
89            self.alerts.overbought_cross = self.alerts.bear_cross && (wt1 >= self.ob_level);
90            self.alerts.oversold_cross = self.alerts.bull_cross && (wt1 <= self.os_level);
91        }
92
93        self.prev_wt1 = Some(wt1);
94        self.prev_wt2 = Some(wt2);
95
96        let mut extra = HashMap::new();
97        extra.insert("wt1".to_string(), wt1);
98        extra.insert("wt2".to_string(), wt2);
99        extra.insert("hist".to_string(), wt1 - wt2);
100        extra.insert("ob_level".to_string(), self.ob_level);
101        extra.insert("os_level".to_string(), self.os_level);
102
103        Some(IndicatorOutput::with_extra(wt1, extra))
104    }
105
106    fn reset(&mut self) {
107        self.ema_ap.reset();
108        self.ema_d.reset();
109        self.ema_wt1.reset();
110        self.sma_wt2.reset();
111        self.prev_wt1 = None;
112        self.prev_wt2 = None;
113        self.alerts = WaveTrendAlerts::default();
114    }
115
116    fn alerts(&self) -> Vec<IndicatorAlert> {
117        let mut res = Vec::new();
118        if self.alerts.bull_cross {
119            res.push(IndicatorAlert::new(
120                "wt_bull_cross",
121                "WaveTrend Bullish Cross",
122                0.8,
123            ));
124        }
125        if self.alerts.bear_cross {
126            res.push(IndicatorAlert::new(
127                "wt_bear_cross",
128                "WaveTrend Bearish Cross",
129                0.8,
130            ));
131        }
132        res
133    }
134}
135
136#[cfg(test)]
137mod tests {
138    use super::*;
139
140    #[test]
141    fn test_wavetrend_basic() {
142        let mut wt = WaveTrendEngine::with_defaults();
143        let mut outputs = Vec::new();
144        for i in 0..50 {
145            let b = Bar::new(i, 100.0, 105.0, 95.0, 100.0 + (i as f64 * 0.5), 1000.0);
146            if let Some(out) = wt.on_bar(&b) {
147                outputs.push(out);
148            }
149        }
150        assert!(!outputs.is_empty());
151        let last = outputs.last().unwrap();
152        assert!(last.extra.contains_key("wt1"));
153        assert!(last.extra.contains_key("wt2"));
154    }
155}