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