1use crate::error::{Error, Result};
4use crate::indicators::linreg::LinearRegression;
5use crate::indicators::rvi_volatility::RviVolatility;
6use crate::ohlcv::Candle;
7use crate::traits::Indicator;
8
9#[derive(Debug, Clone)]
38pub struct Inertia {
39 rvi_period: usize,
40 linreg_period: usize,
41 rvi: RviVolatility,
42 linreg: LinearRegression,
43}
44
45impl Inertia {
46 pub fn new(rvi_period: usize, linreg_period: usize) -> Result<Self> {
49 if rvi_period == 0 || linreg_period == 0 {
50 return Err(Error::PeriodZero);
51 }
52 Ok(Self {
53 rvi_period,
54 linreg_period,
55 rvi: RviVolatility::new(rvi_period)?,
56 linreg: LinearRegression::new(linreg_period)?,
57 })
58 }
59
60 pub fn classic() -> Self {
62 Self::new(14, 20).expect("classic Inertia parameters are valid")
63 }
64
65 pub const fn periods(&self) -> (usize, usize) {
67 (self.rvi_period, self.linreg_period)
68 }
69}
70
71impl Indicator for Inertia {
72 type Input = Candle;
73 type Output = f64;
74
75 #[inline]
76 fn update(&mut self, candle: Candle) -> Option<f64> {
77 let rvi = self.rvi.update(candle.close)?;
78 self.linreg.update(rvi)
79 }
80
81 fn reset(&mut self) {
82 self.rvi.reset();
83 self.linreg.reset();
84 }
85
86 #[inline]
87 fn warmup_period(&self) -> usize {
88 self.rvi.warmup_period() + self.linreg_period - 1
91 }
92
93 #[inline]
94 fn is_ready(&self) -> bool {
95 self.linreg.is_ready()
96 }
97
98 #[inline]
99 fn name(&self) -> &'static str {
100 "Inertia"
101 }
102}
103
104#[cfg(test)]
105mod tests {
106 use super::*;
107 use crate::traits::BatchExt;
108 use approx::assert_relative_eq;
109
110 fn candle(open: f64, high: f64, low: f64, close: f64, ts: i64) -> Candle {
111 Candle::new(open, high, low, close, 1.0, ts).unwrap()
112 }
113
114 #[test]
115 fn rejects_zero_period() {
116 assert!(matches!(Inertia::new(0, 20), Err(Error::PeriodZero)));
117 assert!(matches!(Inertia::new(14, 0), Err(Error::PeriodZero)));
118 }
119
120 #[test]
121 fn accessors_and_metadata() {
122 let inertia = Inertia::classic();
123 assert_eq!(inertia.periods(), (14, 20));
124 assert_eq!(inertia.warmup_period(), 46);
125 assert_eq!(inertia.name(), "Inertia");
126 }
127
128 #[test]
129 fn classic_factory() {
130 assert_eq!(Inertia::classic().periods(), (14, 20));
131 }
132
133 #[test]
134 fn warmup_emits_first_value_at_warmup_period() {
135 let mut inertia = Inertia::new(3, 4).unwrap();
139 assert_eq!(inertia.warmup_period(), 8);
140 for i in 0..7 {
141 assert_eq!(inertia.update(candle(10.0, 11.0, 9.0, 10.5, i)), None);
142 }
143 assert!(inertia.update(candle(10.0, 11.0, 9.0, 10.5, 7)).is_some());
144 }
145
146 #[test]
147 fn constant_rvi_yields_constant_inertia() {
148 let mut inertia = Inertia::new(3, 4).unwrap();
152 let mut last = None;
153 for i in 0..40 {
154 last = inertia.update(candle(10.0, 11.0, 9.0, 10.5, i));
155 }
156 let v = last.unwrap();
157 assert_relative_eq!(v, 50.0, epsilon = 1e-12);
158 }
159
160 #[test]
161 fn batch_equals_streaming() {
162 let candles: Vec<Candle> = (0..80_i64)
163 .map(|i| {
164 let o = 100.0 + (i as f64 * 0.3).sin() * 5.0;
165 let c = o + (i as f64 * 0.1).cos();
166 candle(o, o.max(c) + 0.5, o.min(c) - 0.5, c, i)
167 })
168 .collect();
169 let batch = Inertia::classic().batch(&candles);
170 let mut b = Inertia::classic();
171 let streamed: Vec<_> = candles.iter().map(|c| b.update(*c)).collect();
172 assert_eq!(batch, streamed);
173 }
174
175 #[test]
176 fn reset_clears_state() {
177 let mut inertia = Inertia::classic();
178 for i in 0..50 {
179 inertia.update(candle(10.0, 11.0, 9.0, 10.5, i));
180 }
181 assert!(inertia.is_ready());
182 inertia.reset();
183 assert!(!inertia.is_ready());
184 assert_eq!(inertia.update(candle(10.0, 11.0, 9.0, 10.5, 0)), None);
185 }
186
187 fn wave(len: i64) -> Vec<Candle> {
188 (0..len)
189 .map(|i| {
190 let step = f64::from(i32::try_from(i).unwrap());
191 let o = 100.0 + (step * 0.37).sin() * 6.0;
192 let cl = o + (step * 0.11).cos() * 1.5;
193 candle(o, o.max(cl) + 0.4, o.min(cl) - 0.4, cl, i)
194 })
195 .collect()
196 }
197
198 fn close_only(close: f64, ts: i64) -> Candle {
199 candle(close, close, close, close, ts)
200 }
201
202 #[test]
203 fn rejects_invalid_sub_periods() {
204 assert!(matches!(
207 Inertia::new(1, 20),
208 Err(Error::InvalidPeriod { .. })
209 ));
210 assert!(matches!(
211 Inertia::new(14, 1),
212 Err(Error::InvalidPeriod { .. })
213 ));
214 let too_big = crate::error::MAX_PERIOD + 1;
215 assert!(matches!(
216 Inertia::new(too_big, 20),
217 Err(Error::InvalidPeriod { .. })
218 ));
219 assert!(matches!(
220 Inertia::new(14, too_big),
221 Err(Error::InvalidPeriod { .. })
222 ));
223 }
224
225 #[test]
226 fn classic_first_value_lands_at_index_45() {
227 let candles = wave(60);
230 let out = Inertia::classic().batch(&candles);
231 let warmup = Inertia::classic().warmup_period();
232 assert_eq!(warmup, 46);
233 assert!(out[..warmup - 1].iter().all(Option::is_none));
234 assert!(out[warmup - 1..].iter().all(Option::is_some));
235 }
236
237 #[test]
238 fn hand_computed_reference_rvi2_linreg3() {
239 let closes = [10.0, 12.0, 11.0, 15.0, 14.0, 18.0];
252 let mut inertia = Inertia::new(2, 3).unwrap();
253 assert_eq!(inertia.warmup_period(), 5);
254 let out: Vec<Option<f64>> = closes
255 .iter()
256 .zip(0_i64..)
257 .map(|(&cl, ts)| inertia.update(close_only(cl, ts)))
258 .collect();
259 assert!(out[..4].iter().all(Option::is_none));
260 assert_relative_eq!(out[4].unwrap(), 14800.0 / 198.0, epsilon = 1e-9);
261 let expected5 = (5.0 * 4200.0 / 47.0 + 2.0 * 200.0 / 3.0 - 1000.0 / 11.0) / 6.0;
262 assert_relative_eq!(out[5].unwrap(), expected5, epsilon = 1e-9);
263 }
264
265 #[test]
266 fn reset_reproduces_a_fresh_run() {
267 let candles = wave(90);
268 let mut inertia = Inertia::classic();
269 let first = inertia.batch(&candles);
270 inertia.reset();
271 let second = inertia.batch(&candles);
272 let fresh = Inertia::classic().batch(&candles);
273 assert_eq!(first, second);
274 assert_eq!(second, fresh);
275 }
276
277 #[test]
278 fn batch_nan_into_matches_streaming_bits() {
279 let candles = wave(90);
280 let mut streaming = Inertia::new(5, 7).unwrap();
281 let expected: Vec<u64> = candles
282 .iter()
283 .map(|c| streaming.update(*c).unwrap_or(f64::NAN).to_bits())
284 .collect();
285 let mut out = vec![0.0; candles.len()];
286 Inertia::new(5, 7)
287 .unwrap()
288 .batch_nan_into(&candles, &mut out);
289 let got: Vec<u64> = out.iter().map(|v| v.to_bits()).collect();
290 assert_eq!(got, expected);
291 }
292
293 #[test]
294 fn flat_closes_hold_neutral_fifty() {
295 let mut inertia = Inertia::new(2, 3).unwrap();
298 let out: Vec<Option<f64>> = (0..8)
299 .map(|ts| inertia.update(close_only(7.0, ts)))
300 .collect();
301 assert!(out[4..]
302 .iter()
303 .all(|v| v.is_some_and(|x| (x - 50.0).abs() < 1e-12)));
304 }
305}