1use crate::error::{Error, Result};
4use crate::ohlcv::Candle;
5use crate::traits::Indicator;
6
7#[derive(Debug, Clone, Copy, PartialEq)]
9pub struct AdxOutput {
10 pub plus_di: f64,
12 pub minus_di: f64,
14 pub adx: f64,
16}
17
18#[allow(clippy::struct_field_names)] #[derive(Debug, Clone)]
42pub struct Adx {
43 period: usize,
44 prev: Option<Candle>,
45
46 tr_seed: f64,
48 plus_dm_seed: f64,
49 minus_dm_seed: f64,
50 seed_count: usize,
51
52 tr_smooth: Option<f64>,
54 plus_dm_smooth: Option<f64>,
55 minus_dm_smooth: Option<f64>,
56
57 dx_buf: Vec<f64>,
59 adx_value: Option<f64>,
60 last_plus_di: f64,
61 last_minus_di: f64,
62}
63
64impl Adx {
65 pub fn new(period: usize) -> Result<Self> {
68 if period == 0 {
69 return Err(Error::PeriodZero);
70 }
71 if period > crate::error::MAX_PERIOD {
72 return Err(Error::InvalidPeriod {
73 message: crate::error::PERIOD_ABOVE_MAX,
74 });
75 }
76 Ok(Self {
77 period,
78 prev: None,
79 tr_seed: 0.0,
80 plus_dm_seed: 0.0,
81 minus_dm_seed: 0.0,
82 seed_count: 0,
83 tr_smooth: None,
84 plus_dm_smooth: None,
85 minus_dm_smooth: None,
86 dx_buf: Vec::with_capacity(period),
87 adx_value: None,
88 last_plus_di: 0.0,
89 last_minus_di: 0.0,
90 })
91 }
92
93 pub const fn period(&self) -> usize {
95 self.period
96 }
97}
98
99pub(crate) fn directional_movement(prev: &Candle, current: &Candle) -> (f64, f64) {
100 let up = current.high - prev.high;
101 let down = prev.low - current.low;
102 let plus_dm = if up > down && up > 0.0 { up } else { 0.0 };
103 let minus_dm = if down > up && down > 0.0 { down } else { 0.0 };
104 (plus_dm, minus_dm)
105}
106
107impl Indicator for Adx {
108 type Input = Candle;
109 type Output = AdxOutput;
110
111 fn update(&mut self, candle: Candle) -> Option<AdxOutput> {
112 let Some(prev) = self.prev else {
113 self.prev = Some(candle);
114 return None;
115 };
116 self.prev = Some(candle);
117
118 let tr = candle.true_range(Some(prev.close));
119 let (plus_dm, minus_dm) = directional_movement(&prev, &candle);
120 let n = self.period as f64;
121
122 let (tr_v, plus_v, minus_v) = if let (Some(t), Some(p), Some(m)) =
123 (self.tr_smooth, self.plus_dm_smooth, self.minus_dm_smooth)
124 {
125 let t_new = t - t / n + tr;
126 let p_new = p - p / n + plus_dm;
127 let m_new = m - m / n + minus_dm;
128 self.tr_smooth = Some(t_new);
129 self.plus_dm_smooth = Some(p_new);
130 self.minus_dm_smooth = Some(m_new);
131 (t_new, p_new, m_new)
132 } else {
133 self.tr_seed += tr;
134 self.plus_dm_seed += plus_dm;
135 self.minus_dm_seed += minus_dm;
136 self.seed_count += 1;
137 if self.seed_count < self.period {
138 return None;
139 }
140 self.tr_smooth = Some(self.tr_seed);
141 self.plus_dm_smooth = Some(self.plus_dm_seed);
142 self.minus_dm_smooth = Some(self.minus_dm_seed);
143 (self.tr_seed, self.plus_dm_seed, self.minus_dm_seed)
144 };
145
146 let plus_di = if tr_v == 0.0 {
147 0.0
148 } else {
149 100.0 * plus_v / tr_v
150 };
151 let minus_di = if tr_v == 0.0 {
152 0.0
153 } else {
154 100.0 * minus_v / tr_v
155 };
156 self.last_plus_di = plus_di;
157 self.last_minus_di = minus_di;
158
159 let dx_den = plus_di + minus_di;
160 let dx = if dx_den == 0.0 {
161 0.0
162 } else {
163 100.0 * (plus_di - minus_di).abs() / dx_den
164 };
165
166 if let Some(prev_adx) = self.adx_value {
167 let new_adx = (prev_adx * (n - 1.0) + dx) / n;
168 self.adx_value = Some(new_adx);
169 return Some(AdxOutput {
170 plus_di,
171 minus_di,
172 adx: new_adx,
173 });
174 }
175
176 self.dx_buf.push(dx);
177 if self.dx_buf.len() == self.period {
178 let seed = self.dx_buf.iter().sum::<f64>() / n;
179 self.adx_value = Some(seed);
180 return Some(AdxOutput {
181 plus_di,
182 minus_di,
183 adx: seed,
184 });
185 }
186 None
187 }
188
189 fn reset(&mut self) {
190 self.prev = None;
191 self.tr_seed = 0.0;
192 self.plus_dm_seed = 0.0;
193 self.minus_dm_seed = 0.0;
194 self.seed_count = 0;
195 self.tr_smooth = None;
196 self.plus_dm_smooth = None;
197 self.minus_dm_smooth = None;
198 self.dx_buf.clear();
199 self.adx_value = None;
200 self.last_plus_di = 0.0;
201 self.last_minus_di = 0.0;
202 }
203
204 #[inline]
205 fn warmup_period(&self) -> usize {
206 2 * self.period
207 }
208
209 #[inline]
210 fn is_ready(&self) -> bool {
211 self.adx_value.is_some()
212 }
213
214 #[inline]
215 fn name(&self) -> &'static str {
216 "ADX"
217 }
218}
219
220#[cfg(test)]
221mod tests {
222 use super::*;
223 use crate::traits::BatchExt;
224 use approx::assert_relative_eq;
225
226 fn c(h: f64, l: f64, cl: f64) -> Candle {
227 Candle::new(cl, h, l, cl, 1.0, 0).unwrap()
228 }
229
230 #[test]
231 fn pure_uptrend_yields_plus_di_dominant() {
232 let candles: Vec<Candle> = (0..50)
235 .map(|i| {
236 let base = 100.0 + f64::from(i) * 2.0;
237 c(base + 1.0, base - 0.5, base + 0.5)
238 })
239 .collect();
240 let mut adx = Adx::new(14).unwrap();
241 let last = adx
242 .batch(&candles)
243 .into_iter()
244 .flatten()
245 .last()
246 .expect("emits");
247 assert!(
248 last.plus_di > last.minus_di,
249 "+DI {} should exceed -DI {}",
250 last.plus_di,
251 last.minus_di
252 );
253 assert!(last.adx > 0.0);
254 }
255
256 #[test]
257 fn pure_downtrend_yields_minus_di_dominant() {
258 let candles: Vec<Candle> = (0..50)
259 .rev()
260 .map(|i| {
261 let base = 100.0 + f64::from(i) * 2.0;
262 c(base + 1.0, base - 0.5, base + 0.5)
263 })
264 .collect();
265 let mut adx = Adx::new(14).unwrap();
266 let last = adx
267 .batch(&candles)
268 .into_iter()
269 .flatten()
270 .last()
271 .expect("emits");
272 assert!(last.minus_di > last.plus_di);
273 }
274
275 #[test]
276 fn rejects_zero_period() {
277 assert!(Adx::new(0).is_err());
278 }
279
280 #[test]
284 fn accessors_and_metadata() {
285 let adx = Adx::new(14).unwrap();
286 assert_eq!(adx.period(), 14);
287 assert_eq!(adx.warmup_period(), 28);
288 assert_eq!(adx.name(), "ADX");
289 }
290
291 #[test]
297 fn zero_true_range_yields_zero_di_and_zero_adx() {
298 let candles: Vec<Candle> = (0..30).map(|_| c(10.0, 10.0, 10.0)).collect();
299 let mut adx = Adx::new(5).unwrap();
300 let last = adx
301 .batch(&candles)
302 .into_iter()
303 .flatten()
304 .last()
305 .expect("ADX emits after 2 * period candles");
306 assert_eq!(last.plus_di, 0.0);
307 assert_eq!(last.minus_di, 0.0);
308 assert_eq!(last.adx, 0.0);
309 }
310
311 #[test]
312 fn batch_equals_streaming() {
313 let candles: Vec<Candle> = (0..60)
314 .map(|i| {
315 let base = 100.0 + (f64::from(i) * 0.3).sin() * 5.0;
316 c(base + 1.0, base - 1.0, base)
317 })
318 .collect();
319 let mut a = Adx::new(14).unwrap();
320 let mut b = Adx::new(14).unwrap();
321 assert_eq!(
322 a.batch(&candles),
323 candles.iter().map(|x| b.update(*x)).collect::<Vec<_>>()
324 );
325 }
326
327 #[test]
328 fn reset_clears_state() {
329 let candles: Vec<Candle> = (0..40).map(|_| c(11.0, 9.0, 10.0)).collect();
330 let mut adx = Adx::new(14).unwrap();
331 adx.batch(&candles);
332 adx.reset();
333 assert!(!adx.is_ready());
334 }
335
336 #[test]
337 fn outputs_remain_finite() {
338 let candles: Vec<Candle> = (0..200)
339 .map(|i| {
340 let m = 100.0 + (f64::from(i) * 0.2).sin() * 5.0;
341 c(m + 1.0, m - 1.0, m)
342 })
343 .collect();
344 let mut adx = Adx::new(14).unwrap();
345 for v in adx.batch(&candles).into_iter().flatten() {
346 assert!(v.plus_di.is_finite() && v.minus_di.is_finite() && v.adx.is_finite());
347 }
348 let last = adx.batch(&candles).into_iter().flatten().last().unwrap();
350 assert!(last.adx <= 100.0 + 1e-6);
351 assert_relative_eq!(0.0_f64.max(last.adx), last.adx, epsilon = 1e-9);
352 }
353}