finance_solution/stocks/ta/
keltner.rs1use 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#[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 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#[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 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
190pub 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
252fn 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}