Skip to main content

mantis_ta/indicators/volatility/
keltner.rs

1use super::ATR;
2use crate::indicators::{EMA, Indicator};
3use crate::types::{Candle, KeltnerOutput};
4
5/// Keltner Channels — EMA-centered volatility bands using ATR for width.
6///
7/// The middle band is an EMA of closing price; the upper/lower bands
8/// offset the middle band by `multiplier * ATR`. Unlike Bollinger Bands
9/// (stddev of price), Keltner Channels use true-range-based volatility,
10/// so bands react less to single-bar price gaps.
11///
12/// # Examples
13/// ```rust
14/// use mantis_ta::indicators::{Indicator, KeltnerChannels};
15/// use mantis_ta::types::Candle;
16///
17/// let candles: Vec<Candle> = (0..10)
18///     .map(|i| {
19///         let price = 100.0 + i as f64;
20///         Candle {
21///             timestamp: i as i64,
22///             open: price,
23///             high: price + 1.0,
24///             low: price - 1.0,
25///             close: price,
26///             volume: 1_000.0,
27///         }
28///     })
29///     .collect();
30///
31/// // EMA period 5, ATR period 3, multiplier 2.0 -> warmup at max(5, 3) - 1 = 4
32/// let out = KeltnerChannels::new(5, 3, 2.0).calculate(&candles);
33/// assert!(out.iter().take(4).all(|v| v.is_none()));
34/// let k = out[4].unwrap();
35/// assert!(k.upper > k.middle && k.middle > k.lower);
36/// ```
37#[derive(Debug, Clone)]
38pub struct KeltnerChannels {
39    ema: EMA,
40    atr: ATR,
41    ema_period: usize,
42    atr_period: usize,
43    multiplier: f64,
44}
45
46impl KeltnerChannels {
47    pub fn new(ema_period: usize, atr_period: usize, multiplier: f64) -> Self {
48        assert!(ema_period > 0 && atr_period > 0, "periods must be > 0");
49        assert!(multiplier > 0.0, "multiplier must be > 0");
50        Self {
51            ema: EMA::new(ema_period),
52            atr: ATR::new(atr_period),
53            ema_period,
54            atr_period,
55            multiplier,
56        }
57    }
58}
59
60impl Indicator for KeltnerChannels {
61    type Output = KeltnerOutput;
62
63    fn next(&mut self, candle: &Candle) -> Option<Self::Output> {
64        let middle = self.ema.next(candle);
65        let atr = self.atr.next(candle);
66        let (middle, atr) = match (middle, atr) {
67            (Some(m), Some(a)) => (m, a),
68            _ => return None,
69        };
70        let offset = self.multiplier * atr;
71        Some(KeltnerOutput {
72            upper: middle + offset,
73            middle,
74            lower: middle - offset,
75        })
76    }
77
78    fn reset(&mut self) {
79        self.ema.reset();
80        self.atr.reset();
81    }
82
83    fn warmup_period(&self) -> usize {
84        self.ema_period.max(self.atr_period)
85    }
86
87    fn clone_boxed(&self) -> Box<dyn Indicator<Output = Self::Output>> {
88        Box::new(self.clone())
89    }
90}
91
92#[cfg(test)]
93mod tests {
94    use super::*;
95
96    fn candle(price: f64) -> Candle {
97        Candle {
98            timestamp: 0,
99            open: price,
100            high: price + 1.0,
101            low: price - 1.0,
102            close: price,
103            volume: 1_000.0,
104        }
105    }
106
107    #[test]
108    fn keltner_emits_after_warmup() {
109        let mut kc = KeltnerChannels::new(5, 3, 2.0);
110        let candles: Vec<Candle> = (0..8).map(|i| candle(100.0 + i as f64)).collect();
111
112        let outputs: Vec<_> = candles.iter().map(|c| kc.next(c)).collect();
113        let wp = kc.warmup_period();
114        assert!(outputs.iter().take(wp - 1).all(|o| o.is_none()));
115        assert!(outputs[wp - 1].is_some());
116    }
117
118    #[test]
119    fn keltner_bands_bracket_middle() {
120        let mut kc = KeltnerChannels::new(5, 3, 2.0);
121        let candles: Vec<Candle> = (0..8).map(|i| candle(100.0 + i as f64)).collect();
122
123        for c in &candles {
124            if let Some(out) = kc.next(c) {
125                assert!(out.upper > out.middle);
126                assert!(out.middle > out.lower);
127                assert!((out.upper - out.middle - (out.middle - out.lower)).abs() < 1e-9);
128            }
129        }
130    }
131
132    #[test]
133    fn keltner_flat_prices_zero_width_bands() {
134        // Constant close, but nonzero high/low means ATR still nonzero.
135        let mut kc = KeltnerChannels::new(3, 3, 2.0);
136        let candles: Vec<Candle> = (0..5)
137            .map(|_| Candle {
138                timestamp: 0,
139                open: 50.0,
140                high: 50.0,
141                low: 50.0,
142                close: 50.0,
143                volume: 1_000.0,
144            })
145            .collect();
146
147        for c in &candles {
148            if let Some(out) = kc.next(c) {
149                assert!((out.middle - 50.0).abs() < 1e-9);
150                assert!((out.upper - 50.0).abs() < 1e-9);
151                assert!((out.lower - 50.0).abs() < 1e-9);
152            }
153        }
154    }
155
156    #[test]
157    fn keltner_streaming_matches_batch() {
158        let candles: Vec<Candle> = (0..15).map(|i| candle(90.0 + i as f64 * 0.7)).collect();
159
160        let batch = KeltnerChannels::new(5, 4, 1.5).calculate(&candles);
161
162        let mut streamed_kc = KeltnerChannels::new(5, 4, 1.5);
163        let streamed: Vec<_> = candles.iter().map(|c| streamed_kc.next(c)).collect();
164
165        assert_eq!(streamed, batch);
166    }
167
168    #[test]
169    fn keltner_reset_clears_state() {
170        let mut kc = KeltnerChannels::new(4, 3, 2.0);
171        let candles: Vec<Candle> = (0..6).map(|i| candle(100.0 + i as f64)).collect();
172        for c in &candles {
173            kc.next(c);
174        }
175        kc.reset();
176
177        let mut fresh = KeltnerChannels::new(4, 3, 2.0);
178        for c in &candles {
179            assert_eq!(kc.next(c), fresh.next(c));
180        }
181    }
182
183    #[test]
184    #[should_panic(expected = "periods must be > 0")]
185    fn keltner_rejects_zero_period() {
186        KeltnerChannels::new(0, 3, 2.0);
187    }
188
189    #[test]
190    #[should_panic(expected = "multiplier must be > 0")]
191    fn keltner_rejects_nonpositive_multiplier() {
192        KeltnerChannels::new(5, 3, 0.0);
193    }
194}