kestrel_chartkit/engine/
vwap_regime.rs1use std::collections::VecDeque;
2
3#[cfg(feature = "serde")]
4use serde::{Deserialize, Serialize};
5
6#[derive(Debug, Clone, Copy, PartialEq, Eq)]
7#[cfg_attr(
8 feature = "serde",
9 derive(Serialize, Deserialize),
10 serde(rename_all = "snake_case")
11)]
12pub enum SlopeState {
13 StronglyFalling,
14 ModeratelyFalling,
15 Flat,
16 ModeratelyRising,
17 StronglyRising,
18}
19
20#[derive(Debug, Clone, Copy, PartialEq)]
24#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
25pub struct VwapRegimeOutput {
26 pub slope_atr: f64,
27 pub slope_state: SlopeState,
28 pub distance_atr: f64,
29 pub z_score: f64,
30 pub cross_frequency: f64,
32 pub price_persistence: f64,
34}
35
36pub fn classify_slope(slope_atr: f64, flat_threshold: f64, strong_threshold: f64) -> SlopeState {
37 if slope_atr >= strong_threshold {
38 SlopeState::StronglyRising
39 } else if slope_atr >= flat_threshold {
40 SlopeState::ModeratelyRising
41 } else if slope_atr <= -strong_threshold {
42 SlopeState::StronglyFalling
43 } else if slope_atr <= -flat_threshold {
44 SlopeState::ModeratelyFalling
45 } else {
46 SlopeState::Flat
47 }
48}
49
50pub struct VwapRegimeTracker {
53 window: usize,
54 diffs: VecDeque<f64>,
55}
56
57impl VwapRegimeTracker {
58 pub fn new(window: usize) -> Self {
59 Self {
60 window,
61 diffs: VecDeque::new(),
62 }
63 }
64
65 pub fn reset(&mut self) {
66 self.diffs.clear();
67 }
68
69 #[allow(clippy::too_many_arguments)]
70 pub fn update(
71 &mut self,
72 price_minus_vwap: f64,
73 atr: f64,
74 sigma: f64,
75 slope_atr: f64,
76 flat_threshold: f64,
77 strong_threshold: f64,
78 ) -> VwapRegimeOutput {
79 self.diffs.push_back(price_minus_vwap);
80 if self.diffs.len() > self.window {
81 self.diffs.pop_front();
82 }
83
84 let mut crosses = 0usize;
85 for pair in self.diffs.iter().collect::<Vec<_>>().windows(2) {
86 if (*pair[0] >= 0.0) != (*pair[1] >= 0.0) {
87 crosses += 1;
88 }
89 }
90 let above = self.diffs.iter().filter(|d| **d >= 0.0).count();
91 let n = self.diffs.len().max(1);
92 let cross_frequency = crosses as f64 / n as f64;
93 let side_count = above.max(n - above);
94 let price_persistence = side_count as f64 / n as f64;
95
96 VwapRegimeOutput {
97 slope_atr,
98 slope_state: classify_slope(slope_atr, flat_threshold, strong_threshold),
99 distance_atr: if atr > 0.0 {
100 price_minus_vwap / atr
101 } else {
102 0.0
103 },
104 z_score: if sigma > 0.0 {
105 price_minus_vwap / sigma
106 } else {
107 0.0
108 },
109 cross_frequency,
110 price_persistence,
111 }
112 }
113}
114
115#[cfg(test)]
116mod tests {
117 use super::*;
118
119 #[test]
120 fn slope_classification_boundaries() {
121 assert_eq!(classify_slope(0.5, 0.1, 0.3), SlopeState::StronglyRising);
122 assert_eq!(classify_slope(0.2, 0.1, 0.3), SlopeState::ModeratelyRising);
123 assert_eq!(classify_slope(0.0, 0.1, 0.3), SlopeState::Flat);
124 assert_eq!(
125 classify_slope(-0.2, 0.1, 0.3),
126 SlopeState::ModeratelyFalling
127 );
128 assert_eq!(classify_slope(-0.5, 0.1, 0.3), SlopeState::StronglyFalling);
129 }
130
131 #[test]
132 fn persistent_one_sided_series_has_low_cross_frequency_high_persistence() {
133 let mut tracker = VwapRegimeTracker::new(20);
134 let mut out = None;
135 for _ in 0..20 {
136 out = Some(tracker.update(1.0, 2.0, 0.5, 0.2, 0.1, 0.3));
137 }
138 let out = out.unwrap();
139 assert!((out.cross_frequency).abs() < 1e-9);
140 assert!((out.price_persistence - 1.0).abs() < 1e-9);
141 assert!((out.distance_atr - 0.5).abs() < 1e-9);
142 assert!((out.z_score - 2.0).abs() < 1e-9);
143 }
144
145 #[test]
146 fn alternating_series_has_high_cross_frequency() {
147 let mut tracker = VwapRegimeTracker::new(20);
148 let mut out = None;
149 for i in 0..20 {
150 let diff = if i % 2 == 0 { 1.0 } else { -1.0 };
151 out = Some(tracker.update(diff, 2.0, 0.5, 0.0, 0.1, 0.3));
152 }
153 let out = out.unwrap();
154 assert!(out.cross_frequency > 0.8);
155 }
156}