Skip to main content

wickra_core/indicators/
keltner.rs

1//! Keltner Channels.
2
3use crate::error::{Error, Result};
4use crate::indicators::atr::Atr;
5use crate::indicators::ema::Ema;
6use crate::ohlcv::Candle;
7use crate::traits::Indicator;
8
9/// Keltner Channels output.
10#[derive(Debug, Clone, Copy, PartialEq)]
11pub struct KeltnerOutput {
12    /// Upper band = middle + multiplier * ATR.
13    pub upper: f64,
14    /// Middle band = EMA of the close.
15    pub middle: f64,
16    /// Lower band = middle - multiplier * ATR.
17    pub lower: f64,
18}
19
20/// Keltner Channels: an EMA centerline with bands sized by ATR.
21///
22/// This is the modern (Linda Raschke) form used by `TradingView`, `StockCharts` and
23/// most libraries: `middle = EMA(close, ema_period)`,
24/// `upper / lower = middle ± multiplier · ATR(atr_period)`.
25///
26/// # Example
27///
28/// ```
29/// use wickra_core::{Candle, Indicator, Keltner};
30///
31/// let mut indicator = Keltner::new(5, 5, 2.0).unwrap();
32/// let mut last = None;
33/// for i in 0..80 {
34///     let base = 100.0 + f64::from(i);
35///     let candle =
36///         Candle::new(base, base + 2.0, base - 2.0, base + 1.0, 10.0, i64::from(i)).unwrap();
37///     last = indicator.update(candle);
38/// }
39/// assert!(last.is_some());
40/// ```
41#[derive(Debug, Clone)]
42pub struct Keltner {
43    ema: Ema,
44    atr: Atr,
45    multiplier: f64,
46    ema_period: usize,
47    atr_period: usize,
48}
49
50impl Keltner {
51    /// # Errors
52    /// Returns [`Error::PeriodZero`] / [`Error::NonPositiveMultiplier`] on invalid inputs.
53    pub fn new(ema_period: usize, atr_period: usize, multiplier: f64) -> Result<Self> {
54        if !multiplier.is_finite() || multiplier <= 0.0 {
55            return Err(Error::NonPositiveMultiplier);
56        }
57        Ok(Self {
58            ema: Ema::new(ema_period)?,
59            atr: Atr::new(atr_period)?,
60            multiplier,
61            ema_period,
62            atr_period,
63        })
64    }
65
66    /// Classic configuration: EMA(20), ATR(10), 2.0x multiplier.
67    pub fn classic() -> Self {
68        Self::new(20, 10, 2.0).expect("classic Keltner parameters are valid")
69    }
70
71    /// Configured `(ema_period, atr_period, multiplier)`.
72    pub const fn periods(&self) -> (usize, usize, f64) {
73        (self.ema_period, self.atr_period, self.multiplier)
74    }
75}
76
77impl Indicator for Keltner {
78    type Input = Candle;
79    type Output = KeltnerOutput;
80
81    #[inline]
82    fn update(&mut self, candle: Candle) -> Option<KeltnerOutput> {
83        // Feed both sub-indicators on every candle so they warm up in parallel.
84        // Gating `atr.update` behind `ema.update(...)?` would starve the ATR of
85        // every candle consumed during the EMA's warmup, delaying the first
86        // emission past `warmup_period()` and seeding the ATR over the wrong
87        // window.
88        let mid = self.ema.update(candle.close);
89        let atr = self.atr.update(candle);
90        let (mid, atr) = (mid?, atr?);
91        Some(KeltnerOutput {
92            upper: mid + self.multiplier * atr,
93            middle: mid,
94            lower: mid - self.multiplier * atr,
95        })
96    }
97
98    fn reset(&mut self) {
99        self.ema.reset();
100        self.atr.reset();
101    }
102
103    #[inline]
104    fn warmup_period(&self) -> usize {
105        self.ema_period.max(self.atr_period)
106    }
107
108    #[inline]
109    fn is_ready(&self) -> bool {
110        self.ema.is_ready() && self.atr.is_ready()
111    }
112
113    #[inline]
114    fn name(&self) -> &'static str {
115        "KeltnerChannels"
116    }
117}
118
119#[cfg(test)]
120mod tests {
121    use super::*;
122    use crate::traits::BatchExt;
123    use approx::assert_relative_eq;
124
125    fn c(h: f64, l: f64, cl: f64) -> Candle {
126        Candle::new(cl, h, l, cl, 1.0, 0).unwrap()
127    }
128
129    #[test]
130    fn flat_market_collapses_bands() {
131        let candles: Vec<Candle> = (0..50).map(|_| c(10.0, 10.0, 10.0)).collect();
132        let mut k = Keltner::new(20, 10, 2.0).unwrap();
133        let last = k.batch(&candles).into_iter().flatten().last().unwrap();
134        assert_relative_eq!(last.upper, last.middle, epsilon = 1e-9);
135        assert_relative_eq!(last.lower, last.middle, epsilon = 1e-9);
136    }
137
138    #[test]
139    fn upper_above_middle_above_lower() {
140        let candles: Vec<Candle> = (0..100)
141            .map(|i| {
142                let m = 100.0 + (f64::from(i) * 0.2).sin() * 5.0;
143                c(m + 1.0, m - 1.0, m)
144            })
145            .collect();
146        let mut k = Keltner::classic();
147        for o in k.batch(&candles).into_iter().flatten() {
148            assert!(o.upper >= o.middle);
149            assert!(o.middle >= o.lower);
150        }
151    }
152
153    #[test]
154    fn batch_equals_streaming() {
155        let candles: Vec<Candle> = (0..50)
156            .map(|i| c(f64::from(i) + 1.0, f64::from(i) - 1.0, f64::from(i)))
157            .collect();
158        let mut a = Keltner::classic();
159        let mut b = Keltner::classic();
160        assert_eq!(
161            a.batch(&candles),
162            candles.iter().map(|x| b.update(*x)).collect::<Vec<_>>()
163        );
164    }
165
166    #[test]
167    fn rejects_invalid_input() {
168        assert!(Keltner::new(0, 10, 2.0).is_err());
169        assert!(Keltner::new(20, 10, 0.0).is_err());
170        assert!(Keltner::new(20, 10, -1.0).is_err());
171    }
172
173    /// Cover the const accessor `periods` (68-70) and the Indicator-impl
174    /// `name` body (106-108). Existing tests inspect band output but
175    /// never query the metadata.
176    #[test]
177    fn accessors_and_metadata() {
178        let k = Keltner::new(20, 10, 2.0).unwrap();
179        let (ema, atr, mult) = k.periods();
180        assert_eq!(ema, 20);
181        assert_eq!(atr, 10);
182        assert!((mult - 2.0).abs() < 1e-12);
183        assert_eq!(k.name(), "KeltnerChannels");
184    }
185
186    #[test]
187    fn reset_clears_state() {
188        let candles: Vec<Candle> = (0..50)
189            .map(|i| c(f64::from(i) + 1.0, f64::from(i) - 1.0, f64::from(i)))
190            .collect();
191        let mut k = Keltner::classic();
192        k.batch(&candles);
193        assert!(k.is_ready());
194        k.reset();
195        assert!(!k.is_ready());
196        assert_eq!(k.update(candles[0]), None);
197    }
198
199    #[test]
200    fn first_emission_matches_warmup_period() {
201        let candles: Vec<Candle> = (0..60)
202            .map(|i| {
203                let base = 100.0 + f64::from(i);
204                c(base + 1.0, base - 1.0, base)
205            })
206            .collect();
207        let mut k = Keltner::classic();
208        let out = k.batch(&candles);
209        let warmup = k.warmup_period();
210        assert_eq!(warmup, 20);
211        for (i, v) in out.iter().enumerate().take(warmup - 1) {
212            assert!(v.is_none(), "index {i} must be None during warmup");
213        }
214        assert!(
215            out[warmup - 1].is_some(),
216            "first KeltnerOutput must land at warmup_period - 1"
217        );
218    }
219
220    #[test]
221    fn matches_independent_ema_and_atr() {
222        // The EMA (on the close) and the ATR (on the candle) run as
223        // independent siblings; Keltner must equal feeding two standalone
224        // instances and combining them once both are ready.
225        let candles: Vec<Candle> = (0..60)
226            .map(|i| {
227                let m = 100.0 + (f64::from(i) * 0.2).sin() * 5.0;
228                c(m + 1.5, m - 1.5, m)
229            })
230            .collect();
231        let mut k = Keltner::classic();
232        let mut ema = Ema::new(20).unwrap();
233        let mut atr = Atr::new(10).unwrap();
234        for candle in &candles {
235            let got = k.update(*candle);
236            let mid = ema.update(candle.close);
237            let a = atr.update(*candle);
238            assert_eq!(got.is_some(), mid.is_some() && a.is_some());
239            if let (Some(o), Some(m), Some(av)) = (got, mid, a) {
240                assert_relative_eq!(o.middle, m, epsilon = 1e-9);
241                assert_relative_eq!(o.upper, m + 2.0 * av, epsilon = 1e-9);
242                assert_relative_eq!(o.lower, m - 2.0 * av, epsilon = 1e-9);
243            }
244        }
245    }
246
247    #[test]
248    fn rejects_every_invalid_parameter() {
249        assert!(matches!(Keltner::new(0, 10, 2.0), Err(Error::PeriodZero)));
250        assert!(matches!(Keltner::new(20, 0, 2.0), Err(Error::PeriodZero)));
251        assert!(matches!(
252            Keltner::new(20, 10, f64::NAN),
253            Err(Error::NonPositiveMultiplier)
254        ));
255        assert!(matches!(
256            Keltner::new(20, 10, f64::INFINITY),
257            Err(Error::NonPositiveMultiplier)
258        ));
259        assert!(matches!(
260            Keltner::new(20, 10, 0.0),
261            Err(Error::NonPositiveMultiplier)
262        ));
263        let too_big = crate::error::MAX_PERIOD + 1;
264        assert!(matches!(
265            Keltner::new(too_big, 10, 2.0),
266            Err(Error::InvalidPeriod { .. })
267        ));
268        assert!(matches!(
269            Keltner::new(20, too_big, 2.0),
270            Err(Error::InvalidPeriod { .. })
271        ));
272    }
273
274    #[test]
275    fn warmup_follows_the_longer_atr_period() {
276        // ATR(12) is slower than EMA(5): the first value lands at index 11.
277        let candles: Vec<Candle> = (0..30)
278            .map(|i| c(f64::from(i) + 1.0, f64::from(i) - 1.0, f64::from(i)))
279            .collect();
280        let mut k = Keltner::new(5, 12, 2.0).unwrap();
281        assert_eq!(k.warmup_period(), 12);
282        let out = k.batch(&candles);
283        assert!(out[..11].iter().all(Option::is_none));
284        assert!(out[11..].iter().all(Option::is_some));
285    }
286
287    #[test]
288    fn hand_computed_reference() {
289        // EMA(2) (alpha = 2/3, seeded with the mean), ATR(2) (Wilder, seeded
290        // with the mean true range; the first bar's TR is H − L), mult 1.5.
291        //   b0 H 11 L 9  C 10    TR 2
292        //   b1 H 12 L 10 C 11    TR 2   EMA 10.5   ATR 2
293        //   b2 H 14 L 11 C 13    TR 3   EMA 2/3·13 + 1/3·10.5 = 73/6   ATR (2 + 3)/2 = 2.5
294        //   b3 H 10 L 9  C 9.5   TR max(1, |10 − 13|, |9 − 13|) = 4
295        //                              EMA 2/3·9.5 + 1/3·73/6 = 187/18  ATR (2.5 + 4)/2 = 3.25
296        let candles = [
297            c(11.0, 9.0, 10.0),
298            c(12.0, 10.0, 11.0),
299            c(14.0, 11.0, 13.0),
300            c(10.0, 9.0, 9.5),
301        ];
302        let out = Keltner::new(2, 2, 1.5).unwrap().batch(&candles);
303        assert_eq!(out[0], None);
304        let b1 = out[1].unwrap();
305        assert_relative_eq!(b1.middle, 10.5, epsilon = 1e-12);
306        assert_relative_eq!(b1.upper, 13.5, epsilon = 1e-12);
307        assert_relative_eq!(b1.lower, 7.5, epsilon = 1e-12);
308        let b2 = out[2].unwrap();
309        assert_relative_eq!(b2.middle, 73.0 / 6.0, epsilon = 1e-12);
310        assert_relative_eq!(b2.upper, 73.0 / 6.0 + 3.75, epsilon = 1e-12);
311        assert_relative_eq!(b2.lower, 73.0 / 6.0 - 3.75, epsilon = 1e-12);
312        let b3 = out[3].unwrap();
313        assert_relative_eq!(b3.middle, 187.0 / 18.0, epsilon = 1e-12);
314        assert_relative_eq!(b3.upper, 187.0 / 18.0 + 4.875, epsilon = 1e-12);
315        assert_relative_eq!(b3.lower, 187.0 / 18.0 - 4.875, epsilon = 1e-12);
316    }
317
318    #[test]
319    fn reset_reproduces_a_fresh_run() {
320        let candles: Vec<Candle> = (0..60)
321            .map(|i| {
322                let m = 100.0 + (f64::from(i) * 0.3).sin() * 4.0;
323                c(m + 1.2, m - 0.8, m)
324            })
325            .collect();
326        let mut k = Keltner::new(7, 4, 1.5).unwrap();
327        let first = k.batch(&candles);
328        k.reset();
329        let second = k.batch(&candles);
330        assert_eq!(first, second);
331        assert_eq!(second, Keltner::new(7, 4, 1.5).unwrap().batch(&candles));
332    }
333}