mantis_ta/indicators/volatility/
keltner.rs1use super::ATR;
2use crate::indicators::{EMA, Indicator};
3use crate::types::{Candle, KeltnerOutput};
4
5#[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 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}