nautilus_indicators/average/
rma.rs1use std::fmt::Display;
17
18use nautilus_model::{
19 data::{Bar, QuoteTick, TradeTick},
20 enums::PriceType,
21};
22
23use crate::indicator::{Indicator, MovingAverage};
24
25#[repr(C)]
26#[derive(Debug)]
27#[cfg_attr(
28 feature = "python",
29 pyo3::pyclass(module = "nautilus_trader.indicators")
30)]
31#[cfg_attr(
32 feature = "python",
33 pyo3_stub_gen::derive::gen_stub_pyclass(module = "nautilus_trader.indicators")
34)]
35pub struct WilderMovingAverage {
36 pub period: usize,
37 pub price_type: PriceType,
38 pub alpha: f64,
39 pub value: f64,
40 pub count: usize,
41 pub initialized: bool,
42 has_inputs: bool,
43}
44
45impl Display for WilderMovingAverage {
46 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
47 write!(f, "{}({})", self.name(), self.period)
48 }
49}
50
51impl Indicator for WilderMovingAverage {
52 fn name(&self) -> String {
53 stringify!(WilderMovingAverage).to_string()
54 }
55
56 fn has_inputs(&self) -> bool {
57 self.has_inputs
58 }
59 fn initialized(&self) -> bool {
60 self.initialized
61 }
62
63 fn handle_quote(&mut self, quote: &QuoteTick) -> anyhow::Result<()> {
64 self.update_raw(quote.extract_price(self.price_type)?.into());
65 Ok(())
66 }
67
68 fn handle_trade(&mut self, t: &TradeTick) {
69 self.update_raw((&t.price).into());
70 }
71
72 fn handle_bar(&mut self, b: &Bar) {
73 self.update_raw((&b.close).into());
74 }
75
76 fn reset(&mut self) {
77 self.value = 0.0;
78 self.count = 0;
79 self.has_inputs = false;
80 self.initialized = false;
81 }
82}
83
84impl WilderMovingAverage {
85 #[must_use]
91 pub fn new(period: usize, price_type: Option<PriceType>) -> Self {
92 assert!(
96 period > 0,
97 "WilderMovingAverage: period must be > 0 (received {period})"
98 );
99 Self {
100 period,
101 price_type: price_type.unwrap_or(PriceType::Last),
102 alpha: 1.0 / period as f64,
103 value: 0.0,
104 count: 0,
105 initialized: false,
106 has_inputs: false,
107 }
108 }
109}
110
111impl MovingAverage for WilderMovingAverage {
112 fn value(&self) -> f64 {
113 self.value
114 }
115 fn count(&self) -> usize {
116 self.count
117 }
118
119 fn update_raw(&mut self, price: f64) {
120 if !self.has_inputs {
121 self.has_inputs = true;
122 self.value = price;
123 self.count = 1;
124 self.initialized = self.count >= self.period;
125 return;
126 }
127
128 self.value = self.alpha.mul_add(price, (1.0 - self.alpha) * self.value);
129 self.count += 1;
130 if !self.initialized && self.count >= self.period {
131 self.initialized = true;
132 }
133 }
134}
135
136#[cfg(test)]
137mod tests {
138 use nautilus_model::{
139 data::{Bar, QuoteTick, TradeTick},
140 enums::PriceType,
141 };
142 use rstest::rstest;
143
144 use crate::{
145 average::rma::WilderMovingAverage,
146 indicator::{Indicator, MovingAverage},
147 stubs::*,
148 testing::assert_approx_equal,
149 };
150
151 #[rstest]
152 fn test_rma_initialized(indicator_rma_10: WilderMovingAverage) {
153 let rma = indicator_rma_10;
154 let display_str = format!("{rma}");
155 assert_eq!(display_str, "WilderMovingAverage(10)");
156 assert_eq!(rma.period, 10);
157 assert_eq!(rma.price_type, PriceType::Mid);
158 assert_eq!(rma.alpha, 0.1);
159 assert!(!rma.initialized);
160 }
161
162 #[rstest]
163 #[should_panic(expected = "WilderMovingAverage: period must be > 0")]
164 fn test_new_with_zero_period_panics() {
165 let _ = WilderMovingAverage::new(0, None);
166 }
167
168 #[rstest]
169 fn test_one_value_input(indicator_rma_10: WilderMovingAverage) {
170 let mut rma = indicator_rma_10;
171 rma.update_raw(1.0);
172 assert_eq!(rma.count, 1);
173 assert_eq!(rma.value, 1.0);
174 }
175
176 #[rstest]
177 fn test_rma_update_raw(indicator_rma_10: WilderMovingAverage) {
178 let mut rma = indicator_rma_10;
179 rma.update_raw(1.0);
180 rma.update_raw(2.0);
181 rma.update_raw(3.0);
182 rma.update_raw(4.0);
183 rma.update_raw(5.0);
184 rma.update_raw(6.0);
185 rma.update_raw(7.0);
186 rma.update_raw(8.0);
187 rma.update_raw(9.0);
188 rma.update_raw(10.0);
189
190 assert!(rma.has_inputs());
191 assert!(rma.initialized());
192 assert_eq!(rma.count, 10);
193 assert_approx_equal(rma.value, 4.486_784_401);
194 }
195
196 #[rstest]
197 fn test_reset(indicator_rma_10: WilderMovingAverage) {
198 let mut rma = indicator_rma_10;
199 rma.update_raw(1.0);
200 assert_eq!(rma.count, 1);
201 rma.reset();
202 assert_eq!(rma.count, 0);
203 assert_eq!(rma.value, 0.0);
204 assert!(!rma.initialized);
205 }
206
207 #[rstest]
208 fn test_handle_quote_tick_single(indicator_rma_10: WilderMovingAverage, stub_quote: QuoteTick) {
209 let mut rma = indicator_rma_10;
210 rma.handle_quote(&stub_quote).unwrap();
211 assert!(rma.has_inputs());
212 assert_eq!(rma.value, 1501.0);
213 }
214
215 #[rstest]
216 fn test_handle_quote_tick_multi(mut indicator_rma_10: WilderMovingAverage) {
217 let tick1 = stub_quote("1500.0", "1502.0");
218 let tick2 = stub_quote("1502.0", "1504.0");
219
220 indicator_rma_10.handle_quote(&tick1).unwrap();
221 indicator_rma_10.handle_quote(&tick2).unwrap();
222 assert_eq!(indicator_rma_10.count, 2);
223 assert_eq!(indicator_rma_10.value, 1_501.2);
224 }
225
226 #[rstest]
227 fn test_handle_trade_tick(indicator_rma_10: WilderMovingAverage, stub_trade: TradeTick) {
228 let mut rma = indicator_rma_10;
229 rma.handle_trade(&stub_trade);
230 assert!(rma.has_inputs());
231 assert_eq!(rma.value, 1500.0);
232 }
233
234 #[rstest]
235 fn handle_handle_bar(
236 mut indicator_rma_10: WilderMovingAverage,
237 bar_ethusdt_binance_minute_bid: Bar,
238 ) {
239 indicator_rma_10.handle_bar(&bar_ethusdt_binance_minute_bid);
240 assert!(indicator_rma_10.has_inputs);
241 assert!(!indicator_rma_10.initialized);
242 assert_eq!(indicator_rma_10.value, 1522.0);
243 }
244
245 #[rstest]
246 #[should_panic(expected = "WilderMovingAverage: period must be > 0")]
247 fn invalid_period_panics() {
248 let _ = WilderMovingAverage::new(0, None);
249 }
250
251 #[rstest]
252 #[case(1.0)]
253 #[case(123.456)]
254 #[case(9_876.543_21)]
255 fn first_tick_seeding_parity(#[case] seed_price: f64) {
256 let mut rma = WilderMovingAverage::new(10, None);
257
258 rma.update_raw(seed_price);
259
260 assert_eq!(rma.count(), 1);
261 assert_eq!(rma.value(), seed_price);
262 assert!(!rma.initialized());
263 }
264
265 #[rstest]
266 fn numeric_parity_with_reference_series() {
267 let mut rma = WilderMovingAverage::new(10, None);
268
269 for price in 1_u32..=10 {
270 rma.update_raw(f64::from(price));
271 }
272
273 assert!(rma.initialized());
274 assert_eq!(rma.count(), 10);
275 assert_approx_equal(rma.value(), 4.486_784_401);
276 }
277
278 #[rstest]
280 fn test_rma_period_one_behaviour() {
281 let mut rma = WilderMovingAverage::new(1, None);
282
283 rma.update_raw(42.0);
285 assert!(rma.initialized());
286 assert_eq!(rma.count(), 1);
287 assert!((rma.value() - 42.0).abs() < 1e-12);
288
289 rma.update_raw(100.0);
291 assert_eq!(rma.count(), 2);
292 assert!((rma.value() - 100.0).abs() < 1e-12);
293 }
294
295 #[rstest]
297 fn test_rma_large_period_not_initialized() {
298 let mut rma = WilderMovingAverage::new(1_000, None);
299
300 for p in 1_u32..=999 {
301 rma.update_raw(f64::from(p));
302 }
303
304 assert_eq!(rma.count(), 999);
305 assert!(!rma.initialized());
306 }
307
308 #[rstest]
309 fn test_reset_reseeds_properly() {
310 let mut rma = WilderMovingAverage::new(10, None);
311
312 rma.update_raw(10.0);
313 assert!(rma.has_inputs());
314 assert_eq!(rma.count(), 1);
315
316 rma.reset();
317 assert_eq!(rma.count(), 0);
318 assert!(!rma.has_inputs());
319 assert!(!rma.initialized());
320
321 rma.update_raw(20.0);
322 assert_eq!(rma.count(), 1);
323 assert!((rma.value() - 20.0).abs() < 1e-12);
324 }
325
326 #[rstest]
327 fn test_default_price_type_is_last() {
328 let rma = WilderMovingAverage::new(5, None);
329 assert_eq!(rma.price_type, PriceType::Last);
330 }
331
332 #[rstest]
333 fn test_update_with_nan_propagates() {
334 let mut rma = WilderMovingAverage::new(10, None);
335 rma.update_raw(f64::NAN);
336
337 assert!(rma.value().is_nan());
338 assert!(rma.has_inputs());
339 assert_eq!(rma.count(), 1);
340 }
341}