Skip to main content

finance_solution/stocks/ta/
keltner.rs

1//! # Keltner Channels
2//!
3//! ```text
4//! middle = EMA(ema_period) of close
5//! atr    = Wilder ATR(atr_period) of high/low/close
6//! upper  = middle + atr_mult * atr
7//! lower  = middle − atr_mult * atr
8//! ```
9//!
10//! Default pack: **EMA 20, ATR 10, mult 2** ([`KeltnerParams::standard`]).
11//!
12//! ## Word problem
13//!
14//! > Compare Bollinger(20, 2) and Keltner(20, 10, 2) on the same closes. Which uses
15//! > volatility of closes only, and which uses true range of the bar?
16//!
17//! Bollinger → close stdev. Keltner → Wilder ATR (high/low/close). Teaching tables for both
18//! show warm-up `n/a` until their respective windows fill.
19//!
20//! ## Quant pattern
21//!
22//! ```
23//! use finance_solution::stocks::ta::{KeltnerParams, ValidatedKeltner, KeltnerState};
24//!
25//! const KC: KeltnerParams = KeltnerParams::standard();
26//! let eng = ValidatedKeltner::new(KC).unwrap();
27//! # let n = 40usize;
28//! # let high: Vec<_> = (0..n).map(|i| 101.0 + i as f64 * 0.1).collect();
29//! # let low: Vec<_> = (0..n).map(|i| 99.0 + i as f64 * 0.1).collect();
30//! # let close: Vec<_> = (0..n).map(|i| 100.0 + i as f64 * 0.1).collect();
31//! let s = eng.compute(&high, &low, &close).unwrap();
32//! let mut live = KeltnerState::new(KC).unwrap();
33//! let _ = live.push_bars(&high, &low, &close).unwrap();
34//! assert_eq!(s.middle.len(), n);
35//! ```
36//!
37//! ## Sample solution table
38//!
39//! ```text
40//! period   close  middle   upper   lower     atr
41//! ------  ------  ------  ------  ------  ------
42//!      9  100.90     n/a     n/a     n/a     n/a
43//!     19  101.90  101.20  103.00   99.40  0.9000
44//! ```
45
46use crate::stocks::ta::atr::{atr, AtrParams};
47use crate::stocks::ta::common::{opt_cell, require_hlc};
48use crate::stocks::ta::moving_average::ema;
49use crate::util::error::{require_finite, FinanceError, FinanceResult};
50use crate::util::primitives::PeriodLength;
51use crate::{columns_with_strings, print_table_locale_opt};
52
53/// Keltner parameter pack (Wilder ATR).
54#[derive(Clone, Copy, Debug, PartialEq)]
55pub struct KeltnerParams {
56    pub ema_period: usize,
57    pub atr_period: usize,
58    pub atr_mult: f64,
59}
60
61impl KeltnerParams {
62    /// `(20, 10, 2.0)`.
63    pub const fn standard() -> Self {
64        Self {
65            ema_period: 20,
66            atr_period: 10,
67            atr_mult: 2.0,
68        }
69    }
70
71    pub const fn new(ema_period: usize, atr_period: usize, atr_mult: f64) -> Self {
72        Self {
73            ema_period,
74            atr_period,
75            atr_mult,
76        }
77    }
78}
79
80/// Validated Keltner config.
81#[derive(Clone, Copy, Debug, PartialEq)]
82pub struct ValidatedKeltner {
83    params: KeltnerParams,
84}
85
86impl ValidatedKeltner {
87    pub fn new(params: KeltnerParams) -> FinanceResult<Self> {
88        PeriodLength::new(params.ema_period)?;
89        PeriodLength::new(params.atr_period)?;
90        require_finite("atr_mult", params.atr_mult)?;
91        if params.atr_mult < 0.0 {
92            return Err(FinanceError::Unsolvable {
93                message: "Keltner atr_mult must be non-negative",
94            });
95        }
96        Ok(Self { params })
97    }
98
99    pub fn params(self) -> KeltnerParams {
100        self.params
101    }
102
103    pub fn compute(self, high: &[f64], low: &[f64], close: &[f64]) -> FinanceResult<KeltnerSeries> {
104        keltner_validated(high, low, close, self)
105    }
106}
107
108#[derive(Clone, Debug, PartialEq)]
109pub struct KeltnerSeries {
110    pub middle: Vec<Option<f64>>,
111    pub upper: Vec<Option<f64>>,
112    pub lower: Vec<Option<f64>>,
113    pub atr: Vec<Option<f64>>,
114    pub params: KeltnerParams,
115}
116
117#[derive(Clone, Debug)]
118pub struct KeltnerSolution {
119    series: KeltnerSeries,
120    close: Vec<f64>,
121    formula: String,
122    symbolic_formula: String,
123}
124
125impl KeltnerSolution {
126    pub fn series(&self) -> &KeltnerSeries {
127        &self.series
128    }
129    pub fn formula(&self) -> &str {
130        &self.formula
131    }
132    pub fn symbolic_formula(&self) -> &str {
133        &self.symbolic_formula
134    }
135
136    /// # Sample output
137    /// ```text
138    /// period   close  middle   upper   lower     atr
139    /// ------  ------  ------  ------  ------  ------
140    ///     19  101.90  101.20  103.00   99.40  0.9000
141    /// ```
142    pub fn print_table(&self) {
143        self.print_table_locale_opt(None, None);
144    }
145
146    pub fn print_table_locale(&self, locale: &num_format::Locale, precision: usize) {
147        self.print_table_locale_opt(Some(locale), Some(precision));
148    }
149
150    fn print_table_locale_opt(
151        &self,
152        locale: Option<&num_format::Locale>,
153        precision: Option<usize>,
154    ) {
155        let columns = columns_with_strings(&[
156            ("period", "i", true),
157            ("close", "f", true),
158            ("middle", "f", true),
159            ("upper", "f", true),
160            ("lower", "f", true),
161            ("atr", "f", true),
162        ]);
163        let data = self
164            .close
165            .iter()
166            .enumerate()
167            .map(|(i, c)| {
168                vec![
169                    i.to_string(),
170                    c.to_string(),
171                    opt_cell(self.series.middle[i]),
172                    opt_cell(self.series.upper[i]),
173                    opt_cell(self.series.lower[i]),
174                    opt_cell(self.series.atr[i]),
175                ]
176            })
177            .collect();
178        print_table_locale_opt(&columns, data, locale, precision);
179    }
180}
181
182pub fn keltner(
183    high: &[f64],
184    low: &[f64],
185    close: &[f64],
186    params: KeltnerParams,
187) -> FinanceResult<KeltnerSeries> {
188    ValidatedKeltner::new(params)?.compute(high, low, close)
189}
190
191/// # Examples
192/// ```
193/// use finance_solution::stocks::ta::{keltner_solution, KeltnerParams};
194/// let n = 30usize;
195/// let high: Vec<_> = (0..n).map(|i| 11.0 + i as f64).collect();
196/// let low: Vec<_> = (0..n).map(|i| 9.0 + i as f64).collect();
197/// let close: Vec<_> = (0..n).map(|i| 10.0 + i as f64).collect();
198/// let sol = keltner_solution(&high, &low, &close, KeltnerParams::standard()).unwrap();
199/// assert!(sol.formula().contains("Wilder"));
200/// ```
201pub fn keltner_solution(
202    high: &[f64],
203    low: &[f64],
204    close: &[f64],
205    params: KeltnerParams,
206) -> FinanceResult<KeltnerSolution> {
207    let series = keltner(high, low, close, params)?;
208    let formula = format!(
209        "mid = EMA({})(close); atr = WilderATR({}); upper/lower = mid ± {} * atr",
210        params.ema_period, params.atr_period, params.atr_mult
211    );
212    let symbolic =
213        "mid = ema(close); atr = wilder_atr(high,low,close); bands = mid ± mult * atr".to_string();
214    Ok(KeltnerSolution {
215        series,
216        close: close.to_vec(),
217        formula,
218        symbolic_formula: symbolic,
219    })
220}
221
222fn keltner_validated(
223    high: &[f64],
224    low: &[f64],
225    close: &[f64],
226    v: ValidatedKeltner,
227) -> FinanceResult<KeltnerSeries> {
228    require_hlc(high, low, close)?;
229    let p = v.params;
230    let middle = ema(close, p.ema_period)?;
231    // Single Wilder ATR implementation (shared with `stocks::ta::atr`).
232    let atr_series = atr(high, low, close, AtrParams::new(p.atr_period))?;
233    let atr_vals = atr_series.atr;
234    let n = close.len();
235    let mut upper = vec![None; n];
236    let mut lower = vec![None; n];
237    for i in 0..n {
238        match (middle[i], atr_vals[i]) {
239            (Some(m), Some(a)) => {
240                upper[i] = Some(m + p.atr_mult * a);
241                lower[i] = Some(m - p.atr_mult * a);
242            }
243            _ => {}
244        }
245    }
246    Ok(KeltnerSeries {
247        middle,
248        upper,
249        lower,
250        atr: atr_vals,
251        params: p,
252    })
253}
254
255#[cfg(test)]
256mod tests {
257    use super::*;
258    use crate::stocks::ta::atr::{atr, AtrParams};
259
260    #[test]
261    fn runs() {
262        let n = 40;
263        let high: Vec<_> = (0..n).map(|i| 12.0 + i as f64 * 0.1).collect();
264        let low: Vec<_> = (0..n).map(|i| 10.0 + i as f64 * 0.1).collect();
265        let close: Vec<_> = (0..n).map(|i| 11.0 + i as f64 * 0.1).collect();
266        let s = keltner(&high, &low, &close, KeltnerParams::standard()).unwrap();
267        assert!(s.middle[19].is_some());
268        assert!(s.atr[9].is_some());
269    }
270
271    #[test]
272    fn atr_leg_matches_public_atr() {
273        let n = 50usize;
274        let high: Vec<_> = (0..n).map(|i| 12.0 + i as f64 * 0.05).collect();
275        let low: Vec<_> = (0..n).map(|i| 10.0 + i as f64 * 0.05).collect();
276        let close: Vec<_> = (0..n).map(|i| 11.0 + i as f64 * 0.05).collect();
277        let kc = keltner(&high, &low, &close, KeltnerParams::standard()).unwrap();
278        let a = atr(&high, &low, &close, AtrParams::period_10()).unwrap();
279        for i in 0..n {
280            match (kc.atr[i], a.atr[i]) {
281                (None, None) => {}
282                (Some(x), Some(y)) => assert!((x - y).abs() < 1e-12, "i={i}"),
283                other => panic!("mismatch at {i}: {other:?}"),
284            }
285        }
286    }
287}