Skip to main content

wickra_core/indicators/
adx.rs

1//! Average Directional Index (ADX) with +DI / -DI components.
2
3use crate::error::{Error, Result};
4use crate::ohlcv::Candle;
5use crate::traits::Indicator;
6
7/// ADX output: the three Wilder lines.
8#[derive(Debug, Clone, Copy, PartialEq)]
9pub struct AdxOutput {
10    /// Plus Directional Indicator.
11    pub plus_di: f64,
12    /// Minus Directional Indicator.
13    pub minus_di: f64,
14    /// Average Directional Index (smoothed |DX|).
15    pub adx: f64,
16}
17
18/// Wilder's Average Directional Index.
19///
20/// Uses Wilder smoothing throughout. First `period` candles seed the directional
21/// movement / true range sums; the next `period` candles produce DX values that
22/// seed the ADX. The first complete `AdxOutput` is emitted after `2 * period`
23/// candles.
24///
25/// # Example
26///
27/// ```
28/// use wickra_core::{Candle, Indicator, Adx};
29///
30/// let mut indicator = Adx::new(5).unwrap();
31/// let mut last = None;
32/// for i in 0..80 {
33///     let base = 100.0 + f64::from(i);
34///     let candle =
35///         Candle::new(base, base + 2.0, base - 2.0, base + 1.0, 10.0, i64::from(i)).unwrap();
36///     last = indicator.update(candle);
37/// }
38/// assert!(last.is_some());
39/// ```
40#[allow(clippy::struct_field_names)] // adx_value pairs with adx (the output line) — renaming hurts clarity
41#[derive(Debug, Clone)]
42pub struct Adx {
43    period: usize,
44    prev: Option<Candle>,
45
46    // Wilder-smoothed sums during seeding.
47    tr_seed: f64,
48    plus_dm_seed: f64,
49    minus_dm_seed: f64,
50    seed_count: usize,
51
52    // Smoothed running values after seeding.
53    tr_smooth: Option<f64>,
54    plus_dm_smooth: Option<f64>,
55    minus_dm_smooth: Option<f64>,
56
57    // ADX seeding.
58    dx_buf: Vec<f64>,
59    adx_value: Option<f64>,
60    last_plus_di: f64,
61    last_minus_di: f64,
62}
63
64impl Adx {
65    /// # Errors
66    /// Returns [`Error::PeriodZero`] if `period == 0`.
67    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    /// Configured period.
94    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        // Strict uptrend: highs increase, lows increase, ADX should trend up,
233        // +DI should dominate -DI.
234        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    /// Cover the const accessor `period` (lines 89-91) and the Indicator-impl
281    /// `warmup_period` (199-201) + `name` (207-209). None of the trend tests
282    /// inspect these metadata methods.
283    #[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    /// Cover the `tr_v == 0.0` defensive branches in `update` (lines 142,
292    /// 147) — feeding a stream of perfectly flat candles (H == L == close
293    /// every bar) gives true-range 0 each step, so the smoothed `tr_smooth`
294    /// stays at 0.0 and the `plus_di` / `minus_di` divisions would otherwise
295    /// blow up. The indicator must emit zeros (DX denominator is also 0).
296    #[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        // Sanity: ADX is bounded by 100.
349        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}