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::common::{opt_cell, require_hlc, true_range};
47use crate::stocks::ta::moving_average::ema;
48use crate::util::error::{require_finite, FinanceError, FinanceResult};
49use crate::util::primitives::PeriodLength;
50use crate::{columns_with_strings, print_table_locale_opt};
51
52/// Keltner parameter pack (Wilder ATR).
53#[derive(Clone, Copy, Debug, PartialEq)]
54pub struct KeltnerParams {
55    pub ema_period: usize,
56    pub atr_period: usize,
57    pub atr_mult: f64,
58}
59
60impl KeltnerParams {
61    /// `(20, 10, 2.0)`.
62    pub const fn standard() -> Self {
63        Self {
64            ema_period: 20,
65            atr_period: 10,
66            atr_mult: 2.0,
67        }
68    }
69
70    pub const fn new(ema_period: usize, atr_period: usize, atr_mult: f64) -> Self {
71        Self {
72            ema_period,
73            atr_period,
74            atr_mult,
75        }
76    }
77}
78
79/// Validated Keltner config.
80#[derive(Clone, Copy, Debug, PartialEq)]
81pub struct ValidatedKeltner {
82    params: KeltnerParams,
83}
84
85impl ValidatedKeltner {
86    pub fn new(params: KeltnerParams) -> FinanceResult<Self> {
87        PeriodLength::new(params.ema_period)?;
88        PeriodLength::new(params.atr_period)?;
89        require_finite("atr_mult", params.atr_mult)?;
90        if params.atr_mult < 0.0 {
91            return Err(FinanceError::Unsolvable {
92                message: "Keltner atr_mult must be non-negative",
93            });
94        }
95        Ok(Self { params })
96    }
97
98    pub fn params(self) -> KeltnerParams {
99        self.params
100    }
101
102    pub fn compute(self, high: &[f64], low: &[f64], close: &[f64]) -> FinanceResult<KeltnerSeries> {
103        keltner_validated(high, low, close, self)
104    }
105}
106
107#[derive(Clone, Debug, PartialEq)]
108pub struct KeltnerSeries {
109    pub middle: Vec<Option<f64>>,
110    pub upper: Vec<Option<f64>>,
111    pub lower: Vec<Option<f64>>,
112    pub atr: Vec<Option<f64>>,
113    pub params: KeltnerParams,
114}
115
116#[derive(Clone, Debug)]
117pub struct KeltnerSolution {
118    series: KeltnerSeries,
119    close: Vec<f64>,
120    formula: String,
121    symbolic_formula: String,
122}
123
124impl KeltnerSolution {
125    pub fn series(&self) -> &KeltnerSeries {
126        &self.series
127    }
128    pub fn formula(&self) -> &str {
129        &self.formula
130    }
131    pub fn symbolic_formula(&self) -> &str {
132        &self.symbolic_formula
133    }
134
135    /// # Sample output
136    /// ```text
137    /// period   close  middle   upper   lower     atr
138    /// ------  ------  ------  ------  ------  ------
139    ///     19  101.90  101.20  103.00   99.40  0.9000
140    /// ```
141    pub fn print_table(&self) {
142        self.print_table_locale_opt(None, None);
143    }
144
145    pub fn print_table_locale(&self, locale: &num_format::Locale, precision: usize) {
146        self.print_table_locale_opt(Some(locale), Some(precision));
147    }
148
149    fn print_table_locale_opt(
150        &self,
151        locale: Option<&num_format::Locale>,
152        precision: Option<usize>,
153    ) {
154        let columns = columns_with_strings(&[
155            ("period", "i", true),
156            ("close", "f", true),
157            ("middle", "f", true),
158            ("upper", "f", true),
159            ("lower", "f", true),
160            ("atr", "f", true),
161        ]);
162        let data = self
163            .close
164            .iter()
165            .enumerate()
166            .map(|(i, c)| {
167                vec![
168                    i.to_string(),
169                    c.to_string(),
170                    opt_cell(self.series.middle[i]),
171                    opt_cell(self.series.upper[i]),
172                    opt_cell(self.series.lower[i]),
173                    opt_cell(self.series.atr[i]),
174                ]
175            })
176            .collect();
177        print_table_locale_opt(&columns, data, locale, precision);
178    }
179}
180
181pub fn keltner(
182    high: &[f64],
183    low: &[f64],
184    close: &[f64],
185    params: KeltnerParams,
186) -> FinanceResult<KeltnerSeries> {
187    ValidatedKeltner::new(params)?.compute(high, low, close)
188}
189
190/// # Examples
191/// ```
192/// use finance_solution::stocks::ta::{keltner_solution, KeltnerParams};
193/// let n = 30usize;
194/// let high: Vec<_> = (0..n).map(|i| 11.0 + i as f64).collect();
195/// let low: Vec<_> = (0..n).map(|i| 9.0 + i as f64).collect();
196/// let close: Vec<_> = (0..n).map(|i| 10.0 + i as f64).collect();
197/// let sol = keltner_solution(&high, &low, &close, KeltnerParams::standard()).unwrap();
198/// assert!(sol.formula().contains("Wilder"));
199/// ```
200pub fn keltner_solution(
201    high: &[f64],
202    low: &[f64],
203    close: &[f64],
204    params: KeltnerParams,
205) -> FinanceResult<KeltnerSolution> {
206    let series = keltner(high, low, close, params)?;
207    let formula = format!(
208        "mid = EMA({})(close); atr = WilderATR({}); upper/lower = mid ± {} * atr",
209        params.ema_period, params.atr_period, params.atr_mult
210    );
211    let symbolic =
212        "mid = ema(close); atr = wilder_atr(high,low,close); bands = mid ± mult * atr".to_string();
213    Ok(KeltnerSolution {
214        series,
215        close: close.to_vec(),
216        formula,
217        symbolic_formula: symbolic,
218    })
219}
220
221fn keltner_validated(
222    high: &[f64],
223    low: &[f64],
224    close: &[f64],
225    v: ValidatedKeltner,
226) -> FinanceResult<KeltnerSeries> {
227    require_hlc(high, low, close)?;
228    let p = v.params;
229    let middle = ema(close, p.ema_period)?;
230    let atr = wilder_atr(high, low, close, p.atr_period)?;
231    let n = close.len();
232    let mut upper = vec![None; n];
233    let mut lower = vec![None; n];
234    for i in 0..n {
235        match (middle[i], atr[i]) {
236            (Some(m), Some(a)) => {
237                upper[i] = Some(m + p.atr_mult * a);
238                lower[i] = Some(m - p.atr_mult * a);
239            }
240            _ => {}
241        }
242    }
243    Ok(KeltnerSeries {
244        middle,
245        upper,
246        lower,
247        atr,
248        params: p,
249    })
250}
251
252/// Wilder ATR series. First ATR at index `period-1` = SMA of first `period` true ranges.
253fn wilder_atr(
254    high: &[f64],
255    low: &[f64],
256    close: &[f64],
257    period: usize,
258) -> FinanceResult<Vec<Option<f64>>> {
259    let n = close.len();
260    let mut out = vec![None; n];
261    if n == 0 || period == 0 {
262        return Ok(out);
263    }
264    let mut trs = Vec::with_capacity(n);
265    for i in 0..n {
266        let prev = if i == 0 { None } else { Some(close[i - 1]) };
267        trs.push(true_range(high[i], low[i], prev));
268    }
269    if n < period {
270        return Ok(out);
271    }
272    let sum: f64 = trs[..period].iter().sum();
273    let mut prev_atr = sum / period as f64;
274    out[period - 1] = Some(prev_atr);
275    for i in period..n {
276        prev_atr = (prev_atr * (period as f64 - 1.0) + trs[i]) / period as f64;
277        out[i] = Some(prev_atr);
278    }
279    Ok(out)
280}
281
282#[cfg(test)]
283mod tests {
284    use super::*;
285
286    #[test]
287    fn runs() {
288        let n = 40;
289        let high: Vec<_> = (0..n).map(|i| 12.0 + i as f64 * 0.1).collect();
290        let low: Vec<_> = (0..n).map(|i| 10.0 + i as f64 * 0.1).collect();
291        let close: Vec<_> = (0..n).map(|i| 11.0 + i as f64 * 0.1).collect();
292        let s = keltner(&high, &low, &close, KeltnerParams::standard()).unwrap();
293        assert!(s.middle[19].is_some());
294        assert!(s.atr[9].is_some());
295    }
296}