1use std::collections::VecDeque;
4
5use crate::error::{Error, Result};
6use crate::indicators::ema::Ema;
7use crate::traits::Indicator;
8
9#[derive(Debug, Clone)]
39pub struct Stc {
40 fast_period: usize,
41 slow_period: usize,
42 schaff_period: usize,
43 factor: f64,
44 fast_ema: Ema,
45 slow_ema: Ema,
46 macd_window: VecDeque<f64>,
47 d_window: VecDeque<f64>,
48 last_d: Option<f64>,
49 last_value: Option<f64>,
50}
51
52impl Stc {
53 pub fn new(fast: usize, slow: usize, schaff_period: usize, factor: f64) -> Result<Self> {
57 if fast == 0 || slow == 0 || schaff_period == 0 {
58 return Err(Error::PeriodZero);
59 }
60 if fast >= slow {
61 return Err(Error::InvalidPeriod {
62 message: "STC fast period must be strictly less than slow",
63 });
64 }
65 if !factor.is_finite() || factor <= 0.0 || factor > 1.0 {
66 return Err(Error::InvalidPeriod {
67 message: "STC factor must be a finite value in (0, 1]",
68 });
69 }
70 Ok(Self {
71 fast_period: fast,
72 slow_period: slow,
73 schaff_period,
74 factor,
75 fast_ema: Ema::new(fast)?,
76 slow_ema: Ema::new(slow)?,
77 macd_window: VecDeque::with_capacity(schaff_period),
78 d_window: VecDeque::with_capacity(schaff_period),
79 last_d: None,
80 last_value: None,
81 })
82 }
83
84 pub fn classic() -> Self {
86 Self::new(23, 50, 10, 0.5).expect("classic STC parameters are valid")
87 }
88
89 pub const fn params(&self) -> (usize, usize, usize, f64) {
91 (
92 self.fast_period,
93 self.slow_period,
94 self.schaff_period,
95 self.factor,
96 )
97 }
98}
99
100fn rolling_minmax(window: &VecDeque<f64>) -> (f64, f64) {
101 let mut lo = f64::INFINITY;
102 let mut hi = f64::NEG_INFINITY;
103 for &v in window {
104 if v < lo {
105 lo = v;
106 }
107 if v > hi {
108 hi = v;
109 }
110 }
111 (lo, hi)
112}
113
114impl Indicator for Stc {
115 type Input = f64;
116 type Output = f64;
117
118 fn update(&mut self, input: f64) -> Option<f64> {
119 let f = self.fast_ema.update(input);
120 let s = self.slow_ema.update(input);
121 let (f, s) = (f?, s?);
122 let macd = f - s;
123
124 if self.macd_window.len() == self.schaff_period {
125 self.macd_window.pop_front();
126 }
127 self.macd_window.push_back(macd);
128 if self.macd_window.len() < self.schaff_period {
129 return None;
130 }
131
132 let (lo, hi) = rolling_minmax(&self.macd_window);
133 let k = if hi > lo {
134 100.0 * (macd - lo) / (hi - lo)
135 } else {
136 0.0
137 };
138
139 let d = match self.last_d {
140 Some(prev) => prev + self.factor * (k - prev),
141 None => k,
142 };
143 self.last_d = Some(d);
144
145 if self.d_window.len() == self.schaff_period {
146 self.d_window.pop_front();
147 }
148 self.d_window.push_back(d);
149 if self.d_window.len() < self.schaff_period {
150 return None;
151 }
152
153 let (lo_d, hi_d) = rolling_minmax(&self.d_window);
154 let k2 = if hi_d > lo_d {
155 100.0 * (d - lo_d) / (hi_d - lo_d)
156 } else {
157 0.0
158 };
159
160 let stc = match self.last_value {
161 Some(prev) => prev + self.factor * (k2 - prev),
162 None => k2,
163 };
164 self.last_value = Some(stc);
165 Some(stc.clamp(0.0, 100.0))
166 }
167
168 fn reset(&mut self) {
169 self.fast_ema.reset();
170 self.slow_ema.reset();
171 self.macd_window.clear();
172 self.d_window.clear();
173 self.last_d = None;
174 self.last_value = None;
175 }
176
177 #[inline]
178 fn warmup_period(&self) -> usize {
179 self.slow_period + 2 * (self.schaff_period - 1)
183 }
184
185 #[inline]
186 fn is_ready(&self) -> bool {
187 self.last_value.is_some() && self.d_window.len() == self.schaff_period
188 }
189
190 #[inline]
191 fn name(&self) -> &'static str {
192 "STC"
193 }
194}
195
196#[cfg(test)]
197mod tests {
198 use super::*;
199 use crate::traits::BatchExt;
200
201 #[test]
202 fn rejects_zero_period() {
203 assert!(matches!(Stc::new(0, 50, 10, 0.5), Err(Error::PeriodZero)));
204 assert!(matches!(Stc::new(23, 0, 10, 0.5), Err(Error::PeriodZero)));
205 assert!(matches!(Stc::new(23, 50, 0, 0.5), Err(Error::PeriodZero)));
206 }
207
208 #[test]
209 fn rejects_invalid_params() {
210 assert!(matches!(
211 Stc::new(50, 23, 10, 0.5),
212 Err(Error::InvalidPeriod { .. })
213 ));
214 assert!(matches!(
215 Stc::new(23, 50, 10, 0.0),
216 Err(Error::InvalidPeriod { .. })
217 ));
218 assert!(matches!(
219 Stc::new(23, 50, 10, 1.5),
220 Err(Error::InvalidPeriod { .. })
221 ));
222 assert!(matches!(
223 Stc::new(23, 50, 10, f64::NAN),
224 Err(Error::InvalidPeriod { .. })
225 ));
226 }
227
228 #[test]
229 fn accessors_and_metadata() {
230 let stc = Stc::classic();
231 let (f, s, p, k) = stc.params();
232 assert_eq!((f, s, p), (23, 50, 10));
233 assert!((k - 0.5).abs() < 1e-12);
234 assert_eq!(stc.warmup_period(), 50 + 18);
235 assert_eq!(stc.name(), "STC");
236 }
237
238 #[test]
239 fn classic_factory() {
240 let (f, s, p, k) = Stc::classic().params();
241 assert_eq!((f, s, p), (23, 50, 10));
242 assert!((k - 0.5).abs() < 1e-12);
243 }
244
245 #[test]
246 fn constant_series_yields_zero() {
247 let mut stc = Stc::new(3, 5, 4, 0.5).unwrap();
250 let out = stc.batch(&[42.0_f64; 80]);
251 for v in out.iter().rev().take(5).flatten() {
252 assert_eq!(*v, 0.0);
253 }
254 }
255
256 #[test]
257 fn warmup_emits_first_value_at_warmup_period() {
258 let mut stc = Stc::new(2, 4, 3, 0.5).unwrap();
259 assert_eq!(stc.warmup_period(), 8);
261 let prices: Vec<f64> = (1..=10).map(f64::from).collect();
262 let out = stc.batch(&prices);
263 for v in out.iter().take(7) {
264 assert!(v.is_none());
265 }
266 assert!(out[7].is_some());
267 }
268
269 #[test]
270 fn output_is_bounded() {
271 let mut stc = Stc::classic();
272 let prices: Vec<f64> = (0..400)
273 .map(|i| 100.0 + (f64::from(i) * 0.2).sin() * 25.0)
274 .collect();
275 for v in stc.batch(&prices).iter().flatten() {
276 assert!((0.0..=100.0).contains(v), "STC out of [0, 100]: {v}");
277 }
278 }
279
280 #[test]
281 fn oscillating_series_visits_full_range() {
282 let mut stc = Stc::classic();
289 let prices: Vec<f64> = (0..400)
290 .map(|i| 100.0 + (f64::from(i) * 0.15).sin() * 30.0)
291 .collect();
292 let out = stc.batch(&prices);
293 let mut saw_high = false;
294 let mut saw_low = false;
295 for v in out.iter().flatten() {
296 if *v > 80.0 {
297 saw_high = true;
298 }
299 if *v < 20.0 {
300 saw_low = true;
301 }
302 }
303 assert!(
304 saw_high,
305 "STC should reach above 80 on a strong oscillation"
306 );
307 assert!(saw_low, "STC should reach below 20 on a strong oscillation");
308 }
309
310 #[test]
311 fn batch_equals_streaming() {
312 let prices: Vec<f64> = (1..=200)
313 .map(|i| 100.0 + (f64::from(i) * 0.2).sin() * 5.0)
314 .collect();
315 let mut a = Stc::classic();
316 let mut b = Stc::classic();
317 assert_eq!(
318 a.batch(&prices),
319 prices.iter().map(|p| b.update(*p)).collect::<Vec<_>>()
320 );
321 }
322
323 #[test]
324 fn reset_clears_state() {
325 let mut stc = Stc::classic();
326 stc.batch(&(1..=200).map(f64::from).collect::<Vec<_>>());
327 assert!(stc.is_ready());
328 stc.reset();
329 assert!(!stc.is_ready());
330 assert!(stc.last_value.is_none());
331 }
332}