kestrel_chartkit/indicator/
trend_relationship.rs1use std::collections::HashMap;
8
9use crate::model::Bar;
10
11use super::smoothing::{Smoother, SmootherKind};
12use super::{Indicator, IndicatorAlert, IndicatorOutput};
13
14#[derive(Debug, Clone, Copy, PartialEq, Eq)]
16pub enum TrendRelation {
17 Bullish,
19 Bearish,
21 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
36pub 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 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 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}