Skip to main content

kestrel_chartkit/indicator/
trend_relationship.rs

1//! Adaptive trend-relationship engine: classifies trend state from the relative position and
2//! slope of two configurable [`super::smoothing::Smoother`] stages (any [`SmootherKind`], any
3//! length), instead of duplicating the same "fast MA vs. slow MA, is it rising" logic per script.
4//! Combine with [`super::smoothing::SmootherChain`] to compare cascades of more than one stage per
5//! side (e.g. an EMA-of-SMA fast leg against a plain slow EMA).
6
7use std::collections::HashMap;
8
9use crate::model::Bar;
10
11use super::smoothing::{Smoother, SmootherKind};
12use super::{Indicator, IndicatorAlert, IndicatorOutput};
13
14/// Relative trend classification of the fast smoother against the slow one.
15#[derive(Debug, Clone, Copy, PartialEq, Eq)]
16pub enum TrendRelation {
17    /// Fast above slow and still rising: a confirmed, still-strengthening uptrend.
18    Bullish,
19    /// Fast below slow and still falling: a confirmed, still-strengthening downtrend.
20    Bearish,
21    /// Fast/slow relationship disagrees with the fast leg's own slope (e.g. fast above slow but
22    /// now falling) — likely losing momentum or about to cross.
23    Transition,
24}
25
26impl TrendRelation {
27    fn as_str(self) -> &'static str {
28        match self {
29            TrendRelation::Bullish => "bullish",
30            TrendRelation::Bearish => "bearish",
31            TrendRelation::Transition => "transition",
32        }
33    }
34}
35
36/// Compares a fast and a slow [`Smoother`] fed the same source series, classifying the
37/// relationship as [`TrendRelation`] and alerting on fast/slow crossovers.
38pub struct AdaptiveTrendRelationship {
39    fast: Box<dyn Smoother>,
40    slow: Box<dyn Smoother>,
41    prev_fast: Option<f64>,
42    prev_relation_above: Option<bool>,
43    alerts: Vec<IndicatorAlert>,
44}
45
46impl AdaptiveTrendRelationship {
47    pub fn new(
48        fast_kind: SmootherKind,
49        fast_len: usize,
50        slow_kind: SmootherKind,
51        slow_len: usize,
52    ) -> Self {
53        Self {
54            fast: fast_kind.build(fast_len),
55            slow: slow_kind.build(slow_len),
56            prev_fast: None,
57            prev_relation_above: None,
58            alerts: Vec::new(),
59        }
60    }
61
62    /// Builds the relationship from two arbitrary, already-constructed smoothers (e.g. a
63    /// multi-stage [`super::smoothing::SmootherChain`] wrapped to implement [`Smoother`]).
64    pub fn from_smoothers(fast: Box<dyn Smoother>, slow: Box<dyn Smoother>) -> Self {
65        Self {
66            fast,
67            slow,
68            prev_fast: None,
69            prev_relation_above: None,
70            alerts: Vec::new(),
71        }
72    }
73}
74
75impl Indicator for AdaptiveTrendRelationship {
76    fn name(&self) -> &str {
77        "trend_relationship"
78    }
79
80    fn reset(&mut self) {
81        self.fast.reset();
82        self.slow.reset();
83        self.prev_fast = None;
84        self.prev_relation_above = None;
85        self.alerts.clear();
86    }
87
88    fn on_bar(&mut self, bar: &Bar) -> Option<IndicatorOutput> {
89        self.alerts.clear();
90
91        let fast = self.fast.update(bar.close);
92        let slow = self.slow.update(bar.close);
93        let (fast, slow) = match (fast, slow) {
94            (Some(f), Some(s)) => (f, s),
95            _ => {
96                self.prev_fast = fast.or(self.prev_fast);
97                return None;
98            }
99        };
100
101        let above = fast > slow;
102        if let Some(prev_above) = self.prev_relation_above {
103            if prev_above != above {
104                let (kind, note) = if above {
105                    (
106                        "trend_relationship_cross_up",
107                        "Fast smoother crossed above slow",
108                    )
109                } else {
110                    (
111                        "trend_relationship_cross_down",
112                        "Fast smoother crossed below slow",
113                    )
114                };
115                self.alerts.push(IndicatorAlert::new(kind, note, 1.0));
116            }
117        }
118        self.prev_relation_above = Some(above);
119
120        let rising = self.prev_fast.map(|p| fast > p);
121        let relation = match (above, rising) {
122            (true, Some(true)) | (true, None) => TrendRelation::Bullish,
123            (false, Some(false)) | (false, None) => TrendRelation::Bearish,
124            _ => TrendRelation::Transition,
125        };
126        self.prev_fast = Some(fast);
127
128        let mut extra = HashMap::new();
129        extra.insert("fast".to_string(), fast);
130        extra.insert("slow".to_string(), slow);
131
132        Some(
133            IndicatorOutput::with_extra(fast - slow, extra)
134                .with_secondary(slow)
135                .with_state(relation.as_str()),
136        )
137    }
138
139    fn alerts(&self) -> Vec<IndicatorAlert> {
140        self.alerts.clone()
141    }
142}
143
144#[cfg(test)]
145mod tests {
146    use super::*;
147
148    fn bars_from_closes(closes: &[f64]) -> Vec<Bar> {
149        closes
150            .iter()
151            .enumerate()
152            .map(|(i, &c)| Bar::new(i as i64 * 60, c, c + 1.0, c - 1.0, c, 100.0))
153            .collect()
154    }
155
156    #[test]
157    fn test_classifies_bullish_when_fast_above_and_rising() {
158        let mut engine = AdaptiveTrendRelationship::new(SmootherKind::Sma, 2, SmootherKind::Sma, 4);
159        let closes: Vec<f64> = (0..15).map(|i| 100.0 + i as f64).collect();
160        let mut last_state = None;
161        for bar in bars_from_closes(&closes) {
162            if let Some(out) = engine.on_bar(&bar) {
163                last_state = out.state;
164            }
165        }
166        assert_eq!(last_state.as_deref(), Some("bullish"));
167    }
168
169    #[test]
170    fn test_classifies_bearish_when_fast_below_and_falling() {
171        let mut engine = AdaptiveTrendRelationship::new(SmootherKind::Sma, 2, SmootherKind::Sma, 4);
172        let closes: Vec<f64> = (0..15).map(|i| 200.0 - i as f64).collect();
173        let mut last_state = None;
174        for bar in bars_from_closes(&closes) {
175            if let Some(out) = engine.on_bar(&bar) {
176                last_state = out.state;
177            }
178        }
179        assert_eq!(last_state.as_deref(), Some("bearish"));
180    }
181
182    #[test]
183    fn test_crossover_alert_fires_once_per_cross() {
184        let mut engine = AdaptiveTrendRelationship::new(SmootherKind::Sma, 2, SmootherKind::Sma, 3);
185        // Falling then rising, forcing at least one crossover of fast vs. slow.
186        let closes = [100.0, 99.0, 98.0, 97.0, 100.0, 104.0, 108.0, 112.0];
187        let mut cross_count = 0;
188        for bar in bars_from_closes(&closes) {
189            engine.on_bar(&bar);
190            cross_count += engine.alerts().len();
191        }
192        assert!(cross_count >= 1);
193    }
194
195    #[test]
196    fn test_from_smoothers_accepts_chain() {
197        use super::super::smoothing::SmootherChain;
198
199        let fast = Box::new(SmootherChain::new(vec![
200            SmootherKind::Sma.build(2),
201            SmootherKind::Ema.build(2),
202        ]));
203        let slow = SmootherKind::Sma.build(5);
204        let mut engine = AdaptiveTrendRelationship::from_smoothers(fast, slow);
205
206        let closes: Vec<f64> = (0..10).map(|i| 100.0 + i as f64).collect();
207        let mut saw_output = false;
208        for bar in bars_from_closes(&closes) {
209            if engine.on_bar(&bar).is_some() {
210                saw_output = true;
211            }
212        }
213        assert!(saw_output);
214    }
215}