1use std::collections::{HashMap, VecDeque};
7
8use crate::model::MarketRegime;
9use crate::stats::linear_regression;
10
11#[derive(Debug, Clone, Default)]
14pub struct RegimeMarkovModel {
15 counts: HashMap<(MarketRegime, MarketRegime), u32>,
16 totals: HashMap<MarketRegime, u32>,
17}
18
19impl RegimeMarkovModel {
20 pub fn new() -> Self {
21 Self::default()
22 }
23
24 pub fn observe_transition(&mut self, from: MarketRegime, to: MarketRegime) {
27 *self.counts.entry((from, to)).or_insert(0) += 1;
28 *self.totals.entry(from).or_insert(0) += 1;
29 }
30
31 pub fn transition_probability(&self, from: MarketRegime, to: MarketRegime) -> f64 {
33 let total = match self.totals.get(&from) {
34 Some(&t) if t > 0 => t,
35 _ => return 0.0,
36 };
37 let count = self.counts.get(&(from, to)).copied().unwrap_or(0);
38 count as f64 / total as f64
39 }
40
41 pub fn next_state_distribution(&self, from: MarketRegime) -> Vec<(MarketRegime, f64)> {
44 let states = [
45 MarketRegime::BullishExpansion,
46 MarketRegime::BearishExpansion,
47 MarketRegime::Consolidation,
48 MarketRegime::Transition,
49 ];
50 let mut dist: Vec<(MarketRegime, f64)> = states
51 .into_iter()
52 .map(|to| (to, self.transition_probability(from, to)))
53 .filter(|(_, p)| *p > 0.0)
54 .collect();
55 dist.sort_by(|a, b| b.1.total_cmp(&a.1));
56 dist
57 }
58}
59
60#[derive(Debug, Clone)]
62pub struct RegimePersistenceTracker {
63 current: Option<MarketRegime>,
64 bars_in_regime: u32,
65 completed_streaks: VecDeque<u32>,
66 max_history: usize,
67}
68
69#[derive(Debug, Clone, Copy, PartialEq)]
70pub struct RegimePersistenceOutput {
71 pub bars_in_regime: u32,
72 pub changed: bool,
74 pub average_streak_length: Option<f64>,
77}
78
79impl RegimePersistenceTracker {
80 pub fn new(max_history: usize) -> Self {
81 Self {
82 current: None,
83 bars_in_regime: 0,
84 completed_streaks: VecDeque::new(),
85 max_history: max_history.max(1),
86 }
87 }
88
89 pub fn reset(&mut self) {
90 self.current = None;
91 self.bars_in_regime = 0;
92 self.completed_streaks.clear();
93 }
94
95 pub fn update(&mut self, regime: MarketRegime) -> RegimePersistenceOutput {
96 let changed = match self.current {
97 Some(prev) if prev != regime => {
98 if self.completed_streaks.len() >= self.max_history {
99 self.completed_streaks.pop_front();
100 }
101 self.completed_streaks.push_back(self.bars_in_regime);
102 self.bars_in_regime = 0;
103 true
104 }
105 None => false,
106 _ => false,
107 };
108
109 self.current = Some(regime);
110 self.bars_in_regime += 1;
111
112 let average_streak_length = if self.completed_streaks.is_empty() {
113 None
114 } else {
115 Some(
116 self.completed_streaks.iter().copied().sum::<u32>() as f64
117 / self.completed_streaks.len() as f64,
118 )
119 };
120
121 RegimePersistenceOutput {
122 bars_in_regime: self.bars_in_regime,
123 changed,
124 average_streak_length,
125 }
126 }
127}
128
129#[derive(Debug, Clone)]
132pub struct PredictabilityTracker {
133 window_len: usize,
134 buffer: VecDeque<f64>,
135}
136
137impl PredictabilityTracker {
138 pub fn new(window_len: usize) -> Self {
139 let window_len = window_len.max(2);
140 Self {
141 window_len,
142 buffer: VecDeque::with_capacity(window_len),
143 }
144 }
145
146 pub fn reset(&mut self) {
147 self.buffer.clear();
148 }
149
150 pub fn update(&mut self, close: f64) -> Option<f64> {
152 if self.buffer.len() >= self.window_len {
153 self.buffer.pop_front();
154 }
155 self.buffer.push_back(close);
156 if self.buffer.len() < self.window_len {
157 return None;
158 }
159 let values: Vec<f64> = self.buffer.iter().copied().collect();
160 linear_regression(&values).map(|r| r.r2)
161 }
162}
163
164#[derive(Debug, Clone, Copy, PartialEq, Eq)]
173#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
174pub enum HysteresisLevel {
175 Low,
176 Neutral,
177 High,
178}
179
180#[derive(Debug, Clone, Copy)]
181pub struct HysteresisBand {
182 enter_low: f64,
183 exit_low: f64,
184 exit_high: f64,
185 enter_high: f64,
186 level: HysteresisLevel,
187}
188
189impl HysteresisBand {
190 pub fn new(enter_low: f64, exit_low: f64, exit_high: f64, enter_high: f64) -> Self {
191 Self {
192 enter_low,
193 exit_low,
194 exit_high,
195 enter_high,
196 level: HysteresisLevel::Neutral,
197 }
198 }
199
200 pub fn level(&self) -> HysteresisLevel {
201 self.level
202 }
203
204 pub fn reset(&mut self) {
205 self.level = HysteresisLevel::Neutral;
206 }
207
208 pub fn update(&mut self, value: f64) -> HysteresisLevel {
209 self.level = match self.level {
210 HysteresisLevel::High => {
211 if value <= self.exit_high {
212 HysteresisLevel::Neutral
213 } else {
214 HysteresisLevel::High
215 }
216 }
217 HysteresisLevel::Low => {
218 if value >= self.exit_low {
219 HysteresisLevel::Neutral
220 } else {
221 HysteresisLevel::Low
222 }
223 }
224 HysteresisLevel::Neutral => {
225 if value >= self.enter_high {
226 HysteresisLevel::High
227 } else if value <= self.enter_low {
228 HysteresisLevel::Low
229 } else {
230 HysteresisLevel::Neutral
231 }
232 }
233 };
234 self.level
235 }
236}
237
238#[derive(Debug, Clone)]
242pub struct AdaptiveCycleTracker {
243 prev_sign: Option<i8>,
244 bars_since_crossing: u32,
245 recent_swing_lengths: VecDeque<u32>,
246 max_history: usize,
247}
248
249#[derive(Debug, Clone, Copy, PartialEq)]
250pub struct AdaptiveCycleOutput {
251 pub bars_since_crossing: u32,
252 pub average_swing_length: Option<f64>,
255}
256
257impl AdaptiveCycleTracker {
258 pub fn new(max_history: usize) -> Self {
259 Self {
260 prev_sign: None,
261 bars_since_crossing: 0,
262 recent_swing_lengths: VecDeque::new(),
263 max_history: max_history.max(1),
264 }
265 }
266
267 pub fn reset(&mut self) {
268 self.prev_sign = None;
269 self.bars_since_crossing = 0;
270 self.recent_swing_lengths.clear();
271 }
272
273 pub fn update(&mut self, value: f64) -> AdaptiveCycleOutput {
274 let sign: i8 = if value > 0.0 {
275 1
276 } else if value < 0.0 {
277 -1
278 } else {
279 0
280 };
281
282 let crossed =
285 matches!(self.prev_sign, Some(prev) if prev != 0 && sign != 0 && prev != sign);
286 if crossed {
287 if self.recent_swing_lengths.len() >= self.max_history {
288 self.recent_swing_lengths.pop_front();
289 }
290 self.recent_swing_lengths
291 .push_back(self.bars_since_crossing);
292 self.bars_since_crossing = 0;
293 }
294
295 self.bars_since_crossing += 1;
296
297 if sign != 0 {
298 self.prev_sign = Some(sign);
299 }
300
301 let average_swing_length = if self.recent_swing_lengths.is_empty() {
302 None
303 } else {
304 Some(
305 self.recent_swing_lengths.iter().copied().sum::<u32>() as f64
306 / self.recent_swing_lengths.len() as f64,
307 )
308 };
309
310 AdaptiveCycleOutput {
311 bars_since_crossing: self.bars_since_crossing,
312 average_swing_length,
313 }
314 }
315}
316
317#[cfg(test)]
318mod tests {
319 use super::*;
320
321 #[test]
322 fn test_markov_model_learns_empirical_transitions() {
323 let mut model = RegimeMarkovModel::new();
324 let sequence = [
325 MarketRegime::BullishExpansion,
326 MarketRegime::BullishExpansion,
327 MarketRegime::Consolidation,
328 MarketRegime::BullishExpansion,
329 MarketRegime::BullishExpansion,
330 ];
331 for pair in sequence.windows(2) {
332 model.observe_transition(pair[0], pair[1]);
333 }
334
335 let p_stay = model.transition_probability(
337 MarketRegime::BullishExpansion,
338 MarketRegime::BullishExpansion,
339 );
340 assert!((p_stay - 2.0 / 3.0).abs() < 1e-9);
341
342 let dist = model.next_state_distribution(MarketRegime::BullishExpansion);
343 assert_eq!(dist[0].0, MarketRegime::BullishExpansion);
344
345 assert!(model
347 .next_state_distribution(MarketRegime::BearishExpansion)
348 .is_empty());
349 }
350
351 #[test]
352 fn test_persistence_tracker_counts_streaks_and_changes() {
353 let mut tracker = RegimePersistenceTracker::new(10);
354 let regimes = [
355 MarketRegime::BullishExpansion,
356 MarketRegime::BullishExpansion,
357 MarketRegime::BullishExpansion,
358 MarketRegime::Consolidation,
359 MarketRegime::Consolidation,
360 ];
361 let mut outputs = Vec::new();
362 for regime in regimes {
363 outputs.push(tracker.update(regime));
364 }
365
366 assert!(!outputs[0].changed);
367 assert_eq!(outputs[2].bars_in_regime, 3);
368 assert!(outputs[3].changed, "regime change must be flagged");
369 assert_eq!(outputs[3].bars_in_regime, 1);
370 assert_eq!(outputs[3].average_streak_length, Some(3.0));
371 }
372
373 #[test]
374 fn test_predictability_tracker_scores_clean_trend_high() {
375 let mut tracker = PredictabilityTracker::new(10);
376 let mut last = None;
377 for i in 0..10 {
378 last = tracker.update(100.0 + i as f64);
379 }
380 assert!(
381 last.unwrap() > 0.99,
382 "a perfectly linear trend must score near 1.0"
383 );
384 }
385
386 #[test]
387 fn test_predictability_tracker_scores_noise_low() {
388 let mut tracker = PredictabilityTracker::new(6);
389 let mut last = None;
390 for v in [100.0, 105.0, 98.0, 107.0, 96.0, 109.0] {
391 last = tracker.update(v);
392 }
393 assert!(
394 last.unwrap() < 0.3,
395 "zig-zagging noise must score low predictability"
396 );
397 }
398
399 #[test]
400 fn test_hysteresis_band_does_not_chatter_near_a_single_threshold() {
401 let mut band = HysteresisBand::new(-2.0, -1.0, 1.0, 2.0);
402 assert_eq!(band.update(2.5), HysteresisLevel::High);
403 for v in [1.8, 1.2, 1.9, 1.3, 1.7] {
406 assert_eq!(band.update(v), HysteresisLevel::High);
407 }
408 assert_eq!(band.update(0.5), HysteresisLevel::Neutral);
410 }
411
412 #[test]
413 fn test_hysteresis_band_extreme_to_extreme_passes_through_neutral() {
414 let mut band = HysteresisBand::new(-2.0, -1.0, 1.0, 2.0);
415 assert_eq!(band.update(2.5), HysteresisLevel::High);
416 assert_eq!(band.update(-2.5), HysteresisLevel::Neutral);
417 assert_eq!(band.update(-2.5), HysteresisLevel::Low);
418 }
419
420 #[test]
421 fn test_adaptive_cycle_tracker_measures_swing_length() {
422 let mut tracker = AdaptiveCycleTracker::new(10);
423 let values = [1.0, 1.0, 1.0, 1.0, -1.0, -1.0, -1.0, -1.0];
425 let mut last = AdaptiveCycleOutput {
426 bars_since_crossing: 0,
427 average_swing_length: None,
428 };
429 for v in values {
430 last = tracker.update(v);
431 }
432 assert_eq!(last.average_swing_length, Some(4.0));
433 assert_eq!(last.bars_since_crossing, 4);
434 }
435}