Skip to main content

wickra_core/indicators/
trendflex.rs

1//! Ehlers Trendflex — a trend-sensitive sibling of Reflex.
2#![allow(clippy::doc_markdown)]
3
4use std::collections::VecDeque;
5
6use crate::error::{Error, Result};
7use crate::indicators::super_smoother::SuperSmoother;
8use crate::traits::Indicator;
9
10/// Ehlers' **Trendflex** — the trend-sensitive companion to
11/// [`Reflex`](crate::Reflex): it averages how far the SuperSmoothed price sits
12/// above or below its values over the lookback, then self-normalises.
13///
14/// From John Ehlers, "Reflex: A New Zero-Lag Indicator" (*Stocks & Commodities*,
15/// Feb 2020):
16///
17/// ```text
18/// Filt      = SuperSmoother(price, 0.5 · period)
19/// sum       = mean over i=1..period of ( Filt[0] − Filt[i] )
20/// ms        = 0.04·sum² + 0.96·ms[−1]                (adaptive normaliser)
21/// Trendflex = sum / sqrt(ms)                         (0 if ms == 0)
22/// ```
23///
24/// Where Reflex measures deviation from the straight *line* across the window
25/// (cycle sensitive, near zero lag), Trendflex measures deviation from the
26/// window's *values* (trend sensitive). It stays pinned to one side of zero
27/// during a trend and oscillates through zero in a range, so it doubles as a
28/// trend/range gauge. The adaptive mean-square normaliser keeps the output near a
29/// `±3` band on any instrument.
30///
31/// The first value lands after `period + 1` SuperSmoothed samples. Each `update`
32/// is O(`period`).
33///
34/// # Example
35///
36/// ```
37/// use wickra_core::{Indicator, Trendflex};
38///
39/// let mut indicator = Trendflex::new(20).unwrap();
40/// let mut last = None;
41/// for i in 0..120 {
42///     last = indicator.update(100.0 + f64::from(i));
43/// }
44/// assert!(last.is_some());
45/// ```
46#[derive(Debug, Clone)]
47pub struct Trendflex {
48    period: usize,
49    smoother: SuperSmoother,
50    filt: VecDeque<f64>,
51    ms: f64,
52    last: Option<f64>,
53}
54
55impl Trendflex {
56    /// Construct a Trendflex with the given lookback `period`.
57    ///
58    /// # Errors
59    ///
60    /// Returns [`Error::PeriodZero`] if `period == 0`.
61    pub fn new(period: usize) -> Result<Self> {
62        if period == 0 {
63            return Err(Error::PeriodZero);
64        }
65        if period > crate::error::MAX_PERIOD {
66            return Err(Error::InvalidPeriod {
67                message: crate::error::PERIOD_ABOVE_MAX,
68            });
69        }
70        Ok(Self {
71            period,
72            // Ehlers smooths with half the cycle length (`a1 = exp(-1.414·π / (0.5·Length))`).
73            smoother: SuperSmoother::with_critical_period(period, 0.5 * period as f64),
74            filt: VecDeque::with_capacity(period + 1),
75            ms: 0.0,
76            last: None,
77        })
78    }
79
80    /// Configured lookback period.
81    pub const fn period(&self) -> usize {
82        self.period
83    }
84
85    /// Current value if available.
86    pub const fn value(&self) -> Option<f64> {
87        self.last
88    }
89}
90
91impl Indicator for Trendflex {
92    type Input = f64;
93    type Output = f64;
94
95    #[inline]
96    fn update(&mut self, price: f64) -> Option<f64> {
97        if !price.is_finite() {
98            return None;
99        }
100        let filt = self.smoother.update(price)?;
101        if self.filt.len() == self.period + 1 {
102            self.filt.pop_front();
103        }
104        self.filt.push_back(filt);
105        if self.filt.len() < self.period + 1 {
106            return None;
107        }
108        let newest = self.filt[self.period];
109        let mut sum = 0.0;
110        for i in 1..=self.period {
111            sum += newest - self.filt[self.period - i];
112        }
113        sum /= self.period as f64;
114        self.ms = 0.04 * sum * sum + 0.96 * self.ms;
115        let trendflex = if self.ms > 0.0 {
116            sum / self.ms.sqrt()
117        } else {
118            0.0
119        };
120        self.last = Some(trendflex);
121        Some(trendflex)
122    }
123
124    fn reset(&mut self) {
125        self.smoother.reset();
126        self.filt.clear();
127        self.ms = 0.0;
128        self.last = None;
129    }
130
131    #[inline]
132    fn warmup_period(&self) -> usize {
133        self.period + 1
134    }
135
136    #[inline]
137    fn is_ready(&self) -> bool {
138        self.last.is_some()
139    }
140
141    #[inline]
142    fn name(&self) -> &'static str {
143        "Trendflex"
144    }
145}
146
147#[cfg(test)]
148mod tests {
149    use super::*;
150    use crate::traits::BatchExt;
151    use approx::assert_relative_eq;
152
153    #[test]
154    fn rejects_zero_period() {
155        assert!(matches!(Trendflex::new(0), Err(Error::PeriodZero)));
156    }
157
158    #[test]
159    fn accessors_and_metadata() {
160        let t = Trendflex::new(20).unwrap();
161        assert_eq!(t.period(), 20);
162        assert_eq!(t.warmup_period(), 21);
163        assert_eq!(t.name(), "Trendflex");
164        assert!(!t.is_ready());
165        assert_eq!(t.value(), None);
166    }
167
168    #[test]
169    fn first_emission_at_warmup_period() {
170        let mut t = Trendflex::new(5).unwrap();
171        let xs: Vec<f64> = (0..12).map(f64::from).collect();
172        let out = t.batch(&xs);
173        for v in out.iter().take(5) {
174            assert!(v.is_none());
175        }
176        assert!(out[5].is_some());
177    }
178
179    #[test]
180    fn constant_input_is_zero() {
181        let mut t = Trendflex::new(10).unwrap();
182        for v in t.batch(&[50.0; 100]).into_iter().flatten() {
183            assert_relative_eq!(v, 0.0, epsilon = 1e-9);
184        }
185    }
186
187    #[test]
188    fn uptrend_is_positive() {
189        // A steady rise keeps the current filtered value above its past values.
190        let mut t = Trendflex::new(10).unwrap();
191        let out: Vec<f64> = t
192            .batch(&(0..200).map(f64::from).collect::<Vec<_>>())
193            .into_iter()
194            .flatten()
195            .skip(100)
196            .collect();
197        for v in out {
198            assert!(v > 0.0, "uptrend should be positive, got {v}");
199        }
200    }
201
202    #[test]
203    fn downtrend_is_negative() {
204        let mut t = Trendflex::new(10).unwrap();
205        let out: Vec<f64> = t
206            .batch(&(0..200).map(|i| 200.0 - f64::from(i)).collect::<Vec<_>>())
207            .into_iter()
208            .flatten()
209            .skip(100)
210            .collect();
211        for v in out {
212            assert!(v < 0.0, "downtrend should be negative, got {v}");
213        }
214    }
215
216    #[test]
217    fn ignores_non_finite() {
218        let mut t = Trendflex::new(10).unwrap();
219        t.batch(&(0..40).map(f64::from).collect::<Vec<_>>());
220        let before = t.value();
221        assert_eq!(t.update(f64::NAN), None);
222        // The rejected input must not have disturbed the state.
223        assert_eq!(t.value(), before);
224    }
225
226    #[test]
227    fn reset_clears_state() {
228        let mut t = Trendflex::new(10).unwrap();
229        t.batch(&(0..40).map(f64::from).collect::<Vec<_>>());
230        assert!(t.is_ready());
231        t.reset();
232        assert!(!t.is_ready());
233        assert_eq!(t.value(), None);
234    }
235
236    #[test]
237    fn batch_equals_streaming() {
238        let xs: Vec<f64> = (0..120)
239            .map(|i| 100.0 + (f64::from(i) * 0.25).sin() * 9.0)
240            .collect();
241        let batch = Trendflex::new(20).unwrap().batch(&xs);
242        let mut b = Trendflex::new(20).unwrap();
243        let streamed: Vec<_> = xs.iter().map(|x| b.update(*x)).collect();
244        assert_eq!(batch, streamed);
245    }
246
247    use crate::traits::BatchNanExt;
248
249    #[test]
250    fn rejects_period_above_max() {
251        assert!(matches!(
252            Trendflex::new(crate::error::MAX_PERIOD + 1),
253            Err(Error::InvalidPeriod { .. })
254        ));
255    }
256
257    #[test]
258    fn first_value_lands_exactly_at_warmup_for_several_periods() {
259        for period in [1_usize, 2, 7] {
260            let mut r = Trendflex::new(period).unwrap();
261            let xs: Vec<f64> = (0..20)
262                .map(|i| 100.0 + (f64::from(i) * 0.4).sin() * 3.0)
263                .collect();
264            let out = r.batch(&xs);
265            let warmup = r.warmup_period();
266            assert!(out[..warmup - 1].iter().all(Option::is_none));
267            assert!(out[warmup - 1].is_some());
268        }
269    }
270
271    #[test]
272    fn reset_replays_identically() {
273        let xs: Vec<f64> = (0..120)
274            .map(|i| 100.0 + (f64::from(i) * 0.25).sin() * 9.0)
275            .collect();
276        let fresh = Trendflex::new(13).unwrap().batch(&xs);
277        let mut r = Trendflex::new(13).unwrap();
278        let first = r.batch(&xs);
279        r.reset();
280        let second = r.batch(&xs);
281        assert_eq!(first, fresh);
282        assert_eq!(second, fresh);
283    }
284
285    #[test]
286    fn batch_nan_paths_match_streaming_bitwise() {
287        let xs: Vec<f64> = (0..120)
288            .map(|i| 100.0 + (f64::from(i) * 0.25).sin() * 9.0)
289            .collect();
290        let mut out = vec![0.0; xs.len()];
291        Trendflex::new(13).unwrap().batch_nan_into(&xs, &mut out);
292        let nan = Trendflex::new(13).unwrap().batch_nan(&xs);
293        let fast = Trendflex::new(13).unwrap().batch_fast(&xs);
294        let mut stream = Trendflex::new(13).unwrap();
295        let expected: Vec<u64> = xs
296            .iter()
297            .map(|&p| stream.update(p).unwrap_or(f64::NAN).to_bits())
298            .collect();
299        assert!(out.iter().zip(&expected).all(|(v, e)| v.to_bits() == *e));
300        assert!(nan.iter().zip(&expected).all(|(v, e)| v.to_bits() == *e));
301        assert!(fast.iter().zip(&expected).all(|(v, e)| v.to_bits() == *e));
302    }
303
304    #[test]
305    fn smoother_uses_half_period_critical() {
306        // Ehlers: a1 = exp(-1.414 * pi / (0.5 * Length)); for an odd length the
307        // critical period is fractional (7 -> 3.5).
308        let r = Trendflex::new(7).unwrap();
309        let got = r.smoother.coefficients();
310        let want = SuperSmoother::with_critical_period(7, 3.5).coefficients();
311        assert_eq!(
312            (got.0.to_bits(), got.1.to_bits(), got.2.to_bits()),
313            (want.0.to_bits(), want.1.to_bits(), want.2.to_bits())
314        );
315        assert_eq!(r.smoother.period(), 7);
316        let full = SuperSmoother::new(7).unwrap().coefficients();
317        assert!((got.0 - full.0).abs() > 1e-3);
318    }
319
320    #[test]
321    fn first_value_hand_computed() {
322        // period = 2, inputs 0, 0, 6. SuperSmoother(critical 1.0) outputs
323        // 0, 0 (seed), then c1 * (6 + 0)/2 = 3*c1. Window filt = [0, 0, 3c1].
324        // sum = (3c1 - 0) + (3c1 - 0) = 6c1, divided by 2 -> 3c1;
325        // ms = 0.04 * (3c1)^2; trendflex = 3c1 / (0.2 * 3c1) = 5.
326        let mut t = Trendflex::new(2).unwrap();
327        let (c1, _, _) = t.smoother.coefficients();
328        assert!(c1 > 0.0);
329        let out = t.batch(&[0.0, 0.0, 6.0]);
330        assert_eq!(out[1], None);
331        assert_relative_eq!(out[2].unwrap(), 5.0, epsilon = 1e-12);
332        assert_relative_eq!(t.ms, 0.04 * (3.0 * c1) * (3.0 * c1), epsilon = 1e-12);
333        // Fourth input 6 again: filt = [0, 3c1, y3] with
334        // y3 = c1 * 6 + c2 * 3c1 + c3 * 0, sum = (2*y3 - 3c1) / 2.
335        let (_, c2, _) = t.smoother.coefficients();
336        let y3 = c1 * 6.0 + c2 * 3.0 * c1;
337        let sum = (2.0 * y3 - 3.0 * c1) / 2.0;
338        let ms = 0.04 * sum * sum + 0.96 * t.ms;
339        let next = t.update(6.0).unwrap();
340        assert_relative_eq!(next, sum / ms.sqrt(), epsilon = 1e-12);
341    }
342
343    #[test]
344    fn zero_series_takes_zero_normaliser_branch() {
345        // All-zero input keeps every filt value at exactly 0, so ms stays 0 and
346        // the guarded division returns 0.
347        let mut r = Trendflex::new(4).unwrap();
348        let out = r.batch(&[0.0; 30]);
349        assert!(out
350            .iter()
351            .flatten()
352            .all(|v| v.to_bits() == 0.0_f64.to_bits()));
353        assert_eq!(out.iter().flatten().count(), 30 - 4);
354    }
355}