wickra_core/indicators/
regime_label.rs1use std::collections::VecDeque;
4
5use crate::error::{Error, Result};
6use crate::indicators::rolling_moments::ShiftedMoments;
7use crate::indicators::rolling_quantile::quantile_sorted;
8use crate::indicators::sorted_window;
9use crate::traits::Indicator;
10
11#[derive(Debug, Clone)]
50pub struct RegimeLabel {
51 vol_period: usize,
52 lookback: usize,
53 prev_price: Option<f64>,
54 ret_window: VecDeque<f64>,
56 ret_moments: ShiftedMoments,
57 vol_window: VecDeque<f64>,
59 scratch: Vec<f64>,
62 last: Option<f64>,
63}
64
65impl RegimeLabel {
66 pub fn new(vol_period: usize, lookback: usize) -> Result<Self> {
76 if vol_period < 2 {
77 return Err(Error::InvalidPeriod {
78 message: "regime label needs vol_period >= 2",
79 });
80 }
81 if vol_period > crate::error::MAX_PERIOD {
82 return Err(Error::InvalidPeriod {
83 message: crate::error::PERIOD_ABOVE_MAX,
84 });
85 }
86 if lookback < 2 {
87 return Err(Error::InvalidPeriod {
88 message: "regime label needs lookback >= 2",
89 });
90 }
91 if lookback > crate::error::MAX_PERIOD {
92 return Err(Error::InvalidPeriod {
93 message: crate::error::PERIOD_ABOVE_MAX,
94 });
95 }
96 Ok(Self {
97 vol_period,
98 lookback,
99 prev_price: None,
100 ret_window: VecDeque::with_capacity(vol_period),
101 ret_moments: ShiftedMoments::new(),
102 vol_window: VecDeque::with_capacity(lookback),
103 scratch: Vec::with_capacity(lookback),
104 last: None,
105 })
106 }
107
108 pub const fn params(&self) -> (usize, usize) {
110 (self.vol_period, self.lookback)
111 }
112}
113
114impl Indicator for RegimeLabel {
115 type Input = f64;
116 type Output = f64;
117
118 fn update(&mut self, input: f64) -> Option<f64> {
119 if !input.is_finite() || input <= 0.0 {
120 return None;
121 }
122 let Some(prev) = self.prev_price else {
123 self.prev_price = Some(input);
124 return None;
125 };
126 self.prev_price = Some(input);
127 let r = (input / prev).ln();
128 if self.ret_window.len() == self.vol_period {
130 let old = self.ret_window.pop_front().expect("non-empty");
131 self.ret_moments.evict(old);
132 }
133 self.ret_window.push_back(r);
134 self.ret_moments.push(r);
135 if self.ret_moments.needs_reseed(self.vol_period) {
136 self.ret_moments.reseed(self.ret_window.iter().copied());
137 }
138 if self.ret_window.len() < self.vol_period {
139 return None;
140 }
141 let vol = self.ret_moments.sample_variance(self.vol_period).sqrt();
142 if self.vol_window.len() == self.lookback {
144 let oldest = self.vol_window.pop_front().expect("window is full");
145 sorted_window::remove(&mut self.scratch, oldest);
146 }
147 self.vol_window.push_back(vol);
148 sorted_window::insert(&mut self.scratch, vol);
149 if self.vol_window.len() < self.lookback {
150 return None;
151 }
152 let q1 = quantile_sorted(&self.scratch, 0.25);
154 let q3 = quantile_sorted(&self.scratch, 0.75);
155 let label = if vol < q1 {
156 -1.0
157 } else if vol > q3 {
158 1.0
159 } else {
160 0.0
161 };
162 self.last = Some(label);
163 Some(label)
164 }
165
166 fn reset(&mut self) {
167 self.prev_price = None;
168 self.ret_window.clear();
169 self.ret_moments.reset();
170 self.vol_window.clear();
171 self.scratch.clear();
172 self.last = None;
173 }
174
175 #[inline]
176 fn warmup_period(&self) -> usize {
177 self.vol_period + self.lookback
180 }
181
182 #[inline]
183 fn is_ready(&self) -> bool {
184 self.last.is_some()
185 }
186
187 #[inline]
188 fn name(&self) -> &'static str {
189 "RegimeLabel"
190 }
191}
192
193#[cfg(test)]
194mod tests {
195 use super::*;
196 use crate::traits::BatchExt;
197
198 #[test]
199 fn rejects_bad_periods() {
200 assert!(matches!(
201 RegimeLabel::new(1, 20),
202 Err(Error::InvalidPeriod { .. })
203 ));
204 assert!(matches!(
205 RegimeLabel::new(5, 1),
206 Err(Error::InvalidPeriod { .. })
207 ));
208 }
209
210 #[test]
211 fn accessors_and_metadata() {
212 let rl = RegimeLabel::new(5, 20).unwrap();
213 assert_eq!(rl.params(), (5, 20));
214 assert_eq!(rl.warmup_period(), 25);
215 assert_eq!(rl.name(), "RegimeLabel");
216 assert!(!rl.is_ready());
217 }
218
219 #[test]
220 fn detects_stressed_regime_on_volatility_spike() {
221 let mut rl = RegimeLabel::new(4, 8).unwrap();
224 let mut prices: Vec<f64> = (0..24)
225 .map(|i| 100.0 + (f64::from(i) * 0.7).sin() * 0.2)
226 .collect();
227 let mut base = *prices.last().unwrap();
228 for i in 0..8 {
229 base *= if i % 2 == 0 { 1.08 } else { 0.93 };
230 prices.push(base);
231 }
232 let out = rl.batch(&prices);
233 assert!(
234 out.iter().flatten().any(|&v| v == 1.0),
235 "expected a stressed (+1) regime label"
236 );
237 }
238
239 #[test]
240 fn detects_calm_regime_after_volatility_drop() {
241 let mut rl = RegimeLabel::new(4, 8).unwrap();
243 let mut prices: Vec<f64> = Vec::new();
244 let mut base = 100.0;
245 for i in 0..24 {
246 base *= if i % 2 == 0 { 1.05 } else { 0.96 };
247 prices.push(base);
248 }
249 for i in 0..12 {
250 prices.push(base + (f64::from(i) * 0.7).sin() * 0.05);
251 }
252 let out = rl.batch(&prices);
253 assert!(
254 out.iter().flatten().any(|&v| v == -1.0),
255 "expected a calm (-1) regime label"
256 );
257 }
258
259 #[test]
260 fn zero_volatility_is_neutral() {
261 let mut rl = RegimeLabel::new(4, 8).unwrap();
267 for v in rl.batch(&[100.0; 40]).into_iter().flatten() {
268 assert_eq!(v, 0.0);
269 }
270 }
271
272 #[test]
273 fn output_is_ternary() {
274 let mut rl = RegimeLabel::new(5, 20).unwrap();
275 let prices: Vec<f64> = (0..300)
276 .map(|i| 100.0 + (f64::from(i) * 0.3).sin() * (1.0 + (f64::from(i) * 0.05).sin() * 5.0))
277 .collect();
278 for v in rl.batch(&prices).into_iter().flatten() {
279 assert!(v == -1.0 || v == 0.0 || v == 1.0, "non-ternary label {v}");
280 }
281 }
282
283 #[test]
284 fn ignores_non_finite_and_non_positive() {
285 let mut rl = RegimeLabel::new(4, 6).unwrap();
286 let prices: Vec<f64> = (0..40)
287 .map(|i| 100.0 + (f64::from(i) * 0.5).sin() * 2.0)
288 .collect();
289 let out = rl.batch(&prices);
290 let last = *out.last().unwrap();
291 assert!(last.is_some());
292 assert_eq!(rl.update(f64::NAN), None);
293 assert_eq!(rl.update(-1.0), None);
294 assert_eq!(rl.update(0.0), None);
295 }
296
297 #[test]
298 fn reset_clears_state() {
299 let mut rl = RegimeLabel::new(4, 6).unwrap();
300 rl.batch(&(1..=40).map(f64::from).collect::<Vec<_>>());
301 assert!(rl.is_ready());
302 rl.reset();
303 assert!(!rl.is_ready());
304 assert_eq!(rl.update(1.0), None);
305 }
306
307 #[test]
308 fn batch_equals_streaming() {
309 let prices: Vec<f64> = (1..=160)
310 .map(|i| 100.0 + (f64::from(i) * 0.25).sin() * 4.0)
311 .collect();
312 let batch = RegimeLabel::new(5, 20).unwrap().batch(&prices);
313 let mut b = RegimeLabel::new(5, 20).unwrap();
314 let streamed: Vec<_> = prices.iter().map(|p| b.update(*p)).collect();
315 assert_eq!(batch, streamed);
316 }
317}