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 {
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 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 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}