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.
38///
39/// Both smoothers run from the first close. `value` is `fast - slow`, `secondary` the slow line,
40/// `extra["fast"]`/`extra["slow"]` both lines, and `state` the relation. Through the registry both
41/// are EMAs with their first-sample seed, so output starts with the first bar; other smoother kinds
42/// start when their own seed is ready. [`Indicator::reset`] clears both smoothers.
43pub struct AdaptiveTrendRelationship {
44    fast: Box<dyn Smoother>,
45    slow: Box<dyn Smoother>,
46    prev_fast: Option<f64>,
47    prev_relation_above: Option<bool>,
48    alerts: Vec<IndicatorAlert>,
49}
50
51impl AdaptiveTrendRelationship {
52    pub fn new(
53        fast_kind: SmootherKind,
54        fast_len: usize,
55        slow_kind: SmootherKind,
56        slow_len: usize,
57    ) -> Self {
58        Self {
59            fast: fast_kind.build(fast_len),
60            slow: slow_kind.build(slow_len),
61            prev_fast: None,
62            prev_relation_above: None,
63            alerts: Vec::new(),
64        }
65    }
66
67    /// Builds the relationship from two arbitrary, already-constructed smoothers (e.g. a
68    /// multi-stage [`super::smoothing::SmootherChain`] wrapped to implement [`Smoother`]).
69    pub fn from_smoothers(fast: Box<dyn Smoother>, slow: Box<dyn Smoother>) -> Self {
70        Self {
71            fast,
72            slow,
73            prev_fast: None,
74            prev_relation_above: None,
75            alerts: Vec::new(),
76        }
77    }
78}
79
80impl Indicator for AdaptiveTrendRelationship {
81    fn name(&self) -> &str {
82        "trend_relationship"
83    }
84
85    fn reset(&mut self) {
86        self.fast.reset();
87        self.slow.reset();
88        self.prev_fast = None;
89        self.prev_relation_above = None;
90        self.alerts.clear();
91    }
92
93    fn on_bar(&mut self, bar: &Bar) -> Option<IndicatorOutput> {
94        self.alerts.clear();
95
96        let fast = self.fast.update(bar.close);
97        let slow = self.slow.update(bar.close);
98        let (fast, slow) = match (fast, slow) {
99            (Some(f), Some(s)) => (f, s),
100            _ => {
101                self.prev_fast = fast.or(self.prev_fast);
102                return None;
103            }
104        };
105
106        let above = fast > slow;
107        if let Some(prev_above) = self.prev_relation_above {
108            if prev_above != above {
109                let (kind, note) = if above {
110                    (
111                        "trend_relationship_cross_up",
112                        "Fast smoother crossed above slow",
113                    )
114                } else {
115                    (
116                        "trend_relationship_cross_down",
117                        "Fast smoother crossed below slow",
118                    )
119                };
120                self.alerts.push(IndicatorAlert::new(kind, note, 1.0));
121            }
122        }
123        self.prev_relation_above = Some(above);
124
125        let rising = self.prev_fast.map(|p| fast > p);
126        let relation = match (above, rising) {
127            (true, Some(true)) | (true, None) => TrendRelation::Bullish,
128            (false, Some(false)) | (false, None) => TrendRelation::Bearish,
129            _ => TrendRelation::Transition,
130        };
131        self.prev_fast = Some(fast);
132
133        let mut extra = HashMap::new();
134        extra.insert("fast".to_string(), fast);
135        extra.insert("slow".to_string(), slow);
136
137        Some(
138            IndicatorOutput::with_extra(fast - slow, extra)
139                .with_secondary(slow)
140                .with_state(relation.as_str()),
141        )
142    }
143
144    fn alerts(&self) -> Vec<IndicatorAlert> {
145        self.alerts.clone()
146    }
147}
148
149#[cfg(test)]
150mod tests {
151    use super::*;
152
153    fn bars_from_closes(closes: &[f64]) -> Vec<Bar> {
154        closes
155            .iter()
156            .enumerate()
157            .map(|(i, &c)| Bar::new(i as i64 * 60, c, c + 1.0, c - 1.0, c, 100.0))
158            .collect()
159    }
160
161    #[test]
162    fn test_classifies_bullish_when_fast_above_and_rising() {
163        let mut engine = AdaptiveTrendRelationship::new(SmootherKind::Sma, 2, SmootherKind::Sma, 4);
164        let closes: Vec<f64> = (0..15).map(|i| 100.0 + i as f64).collect();
165        let mut last_state = None;
166        for bar in bars_from_closes(&closes) {
167            if let Some(out) = engine.on_bar(&bar) {
168                last_state = out.state;
169            }
170        }
171        assert_eq!(last_state.as_deref(), Some("bullish"));
172    }
173
174    #[test]
175    fn test_classifies_bearish_when_fast_below_and_falling() {
176        let mut engine = AdaptiveTrendRelationship::new(SmootherKind::Sma, 2, SmootherKind::Sma, 4);
177        let closes: Vec<f64> = (0..15).map(|i| 200.0 - i as f64).collect();
178        let mut last_state = None;
179        for bar in bars_from_closes(&closes) {
180            if let Some(out) = engine.on_bar(&bar) {
181                last_state = out.state;
182            }
183        }
184        assert_eq!(last_state.as_deref(), Some("bearish"));
185    }
186
187    #[test]
188    fn test_crossover_alert_fires_once_per_cross() {
189        let mut engine = AdaptiveTrendRelationship::new(SmootherKind::Sma, 2, SmootherKind::Sma, 3);
190        // Falling then rising, forcing at least one crossover of fast vs. slow.
191        let closes = [100.0, 99.0, 98.0, 97.0, 100.0, 104.0, 108.0, 112.0];
192        let mut cross_count = 0;
193        for bar in bars_from_closes(&closes) {
194            engine.on_bar(&bar);
195            cross_count += engine.alerts().len();
196        }
197        assert!(cross_count >= 1);
198    }
199
200    #[test]
201    fn test_from_smoothers_accepts_chain() {
202        use super::super::smoothing::SmootherChain;
203
204        let fast = Box::new(SmootherChain::new(vec![
205            SmootherKind::Sma.build(2),
206            SmootherKind::Ema.build(2),
207        ]));
208        let slow = SmootherKind::Sma.build(5);
209        let mut engine = AdaptiveTrendRelationship::from_smoothers(fast, slow);
210
211        let closes: Vec<f64> = (0..10).map(|i| 100.0 + i as f64).collect();
212        let mut saw_output = false;
213        for bar in bars_from_closes(&closes) {
214            if engine.on_bar(&bar).is_some() {
215                saw_output = true;
216            }
217        }
218        assert!(saw_output);
219    }
220}