nautilus_indicators/momentum/
cmo.rs1use std::fmt::Display;
17
18use nautilus_model::data::{Bar, QuoteTick, TradeTick};
19
20use crate::{
21 average::{MovingAverageFactory, MovingAverageType},
22 indicator::{Indicator, MovingAverage},
23};
24
25#[repr(C)]
26#[derive(Debug)]
27#[cfg_attr(
28 feature = "python",
29 pyo3::pyclass(module = "nautilus_trader.indicators", unsendable)
30)]
31#[cfg_attr(
32 feature = "python",
33 pyo3_stub_gen::derive::gen_stub_pyclass(module = "nautilus_trader.indicators")
34)]
35pub struct ChandeMomentumOscillator {
36 pub period: usize,
37 pub ma_type: MovingAverageType,
38 pub value: f64,
39 pub count: usize,
40 pub initialized: bool,
41 previous_close: f64,
42 average_gain: Box<dyn MovingAverage + Send + 'static>,
43 average_loss: Box<dyn MovingAverage + Send + 'static>,
44 has_inputs: bool,
45}
46
47impl Display for ChandeMomentumOscillator {
48 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
49 write!(f, "{}({})", self.name(), self.period)
50 }
51}
52
53impl Indicator for ChandeMomentumOscillator {
54 fn name(&self) -> String {
55 stringify!(ChandeMomentumOscillator).to_string()
56 }
57
58 fn has_inputs(&self) -> bool {
59 self.has_inputs
60 }
61
62 fn initialized(&self) -> bool {
63 self.initialized
64 }
65
66 fn handle_quote(&mut self, _quote: &QuoteTick) -> anyhow::Result<()> {
67 Ok(())
68 }
69
70 fn handle_trade(&mut self, _trade: &TradeTick) {}
71
72 fn handle_bar(&mut self, bar: &Bar) {
73 self.update_raw((&bar.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 self.previous_close = 0.0;
82 self.average_gain.reset();
83 self.average_loss.reset();
84 }
85}
86
87impl ChandeMomentumOscillator {
88 #[must_use]
94 pub fn new(period: usize, ma_type: Option<MovingAverageType>) -> Self {
95 assert!(period > 0, "ChandeMomentumOscillator: period must be > 0");
96 let ma_type = ma_type.unwrap_or(MovingAverageType::Wilder);
97 Self {
98 period,
99 ma_type,
100 average_gain: MovingAverageFactory::create(ma_type, period),
101 average_loss: MovingAverageFactory::create(ma_type, period),
102 previous_close: 0.0,
103 value: 0.0,
104 count: 0,
105 initialized: false,
106 has_inputs: false,
107 }
108 }
109
110 pub fn update_raw(&mut self, close: f64) {
111 self.count += 1;
112
113 if !self.has_inputs {
114 self.previous_close = close;
115 self.has_inputs = true;
116 }
117
118 let gain: f64 = close - self.previous_close;
119 if gain > 0.0 {
120 self.average_gain.update_raw(gain);
121 self.average_loss.update_raw(0.0);
122 } else if gain < 0.0 {
123 self.average_gain.update_raw(0.0);
124 self.average_loss.update_raw(-gain);
125 } else {
126 self.average_gain.update_raw(0.0);
127 self.average_loss.update_raw(0.0);
128 }
129
130 if !self.initialized && self.average_gain.initialized() && self.average_loss.initialized() {
131 self.initialized = true;
132 }
133
134 if self.initialized {
135 let divisor = self.average_gain.value() + self.average_loss.value();
136 if divisor == 0.0 {
137 self.value = 0.0;
138 } else {
139 self.value =
142 100.0 * ((self.average_gain.value() - self.average_loss.value()) / divisor);
143 }
144 }
145 self.previous_close = close;
146 }
147}
148
149#[cfg(test)]
150mod tests {
151 use nautilus_model::data::{Bar, QuoteTick};
152 use rstest::rstest;
153
154 use crate::{
155 average::MovingAverageType, indicator::Indicator, momentum::cmo::ChandeMomentumOscillator,
156 stubs::*, testing::assert_approx_equal,
157 };
158
159 #[rstest]
160 fn test_cmo_initialized(cmo_10: ChandeMomentumOscillator) {
161 let display_str = format!("{cmo_10}");
162 assert_eq!(display_str, "ChandeMomentumOscillator(10)");
163 assert_eq!(cmo_10.period, 10);
164 assert!(!cmo_10.initialized);
165 }
166
167 #[rstest]
168 fn test_value_stays_within_bounds_when_gain_average_is_zero() {
169 let mut cmo = ChandeMomentumOscillator::new(14, None);
173 for price in std::iter::repeat_n(100.0, 21).chain(std::iter::repeat_n(50.0, 20)) {
174 cmo.update_raw(price);
175 assert!(
176 (-100.0..=100.0).contains(&cmo.value),
177 "value {} outside [-100, 100]",
178 cmo.value
179 );
180 }
181 assert_eq!(cmo.value, -100.0);
182 }
183
184 #[rstest]
185 fn test_initialized_with_required_inputs_returns_true(mut cmo_10: ChandeMomentumOscillator) {
186 for i in 0..12 {
187 cmo_10.update_raw(f64::from(i));
188 }
189 assert!(cmo_10.initialized);
190 }
191
192 #[rstest]
193 fn test_value_all_higher_inputs_returns_expected_value(mut cmo_10: ChandeMomentumOscillator) {
194 cmo_10.update_raw(109.93);
195 cmo_10.update_raw(110.0);
196 cmo_10.update_raw(109.77);
197 cmo_10.update_raw(109.96);
198 cmo_10.update_raw(110.29);
199 cmo_10.update_raw(110.53);
200 cmo_10.update_raw(110.27);
201 cmo_10.update_raw(110.21);
202 cmo_10.update_raw(110.06);
203 cmo_10.update_raw(110.19);
204 cmo_10.update_raw(109.83);
205 cmo_10.update_raw(109.9);
206 cmo_10.update_raw(110.0);
207 cmo_10.update_raw(110.03);
208 cmo_10.update_raw(110.13);
209 cmo_10.update_raw(109.95);
210 cmo_10.update_raw(109.75);
211 cmo_10.update_raw(110.15);
212 cmo_10.update_raw(109.9);
213 cmo_10.update_raw(110.04);
214 assert_approx_equal(cmo_10.value, 2.08962945624);
215 }
216
217 #[rstest]
218 fn test_value_with_one_input_returns_expected_value(mut cmo_10: ChandeMomentumOscillator) {
219 cmo_10.update_raw(1.00000);
220 assert_eq!(cmo_10.value, 0.0);
221 }
222
223 #[rstest]
224 fn test_reset(mut cmo_10: ChandeMomentumOscillator) {
225 cmo_10.update_raw(1.00020);
226 cmo_10.update_raw(1.00030);
227 cmo_10.update_raw(1.00050);
228 cmo_10.reset();
229 assert!(!cmo_10.initialized());
230 assert_eq!(cmo_10.count, 0);
231 assert_eq!(cmo_10.value, 0.0);
232 assert_eq!(cmo_10.previous_close, 0.0);
233 }
234
235 #[rstest]
236 fn test_handle_quote_tick(mut cmo_10: ChandeMomentumOscillator, stub_quote: QuoteTick) {
237 cmo_10.handle_quote(&stub_quote).unwrap();
238 assert_eq!(cmo_10.count, 0);
239 assert_eq!(cmo_10.value, 0.0);
240 }
241
242 #[rstest]
243 fn test_handle_bar(mut cmo_10: ChandeMomentumOscillator, bar_ethusdt_binance_minute_bid: Bar) {
244 cmo_10.handle_bar(&bar_ethusdt_binance_minute_bid);
245 assert_eq!(cmo_10.count, 1);
246 assert_eq!(cmo_10.value, 0.0);
247 }
248
249 #[rstest]
250 fn test_ma_type_affects_value() {
251 let mut cmo_sma = ChandeMomentumOscillator::new(3, Some(MovingAverageType::Simple));
252 let mut cmo_wilder = ChandeMomentumOscillator::new(3, Some(MovingAverageType::Wilder));
253 let prices = [1.0, 2.0, 3.0, 2.5, 3.5];
254 for price in prices {
255 cmo_sma.update_raw(price);
256 cmo_wilder.update_raw(price);
257 }
258 assert_ne!(cmo_sma.value, cmo_wilder.value);
259 }
260
261 #[rstest]
262 fn test_count_increments(mut cmo_10: ChandeMomentumOscillator) {
263 for i in 0..5 {
264 cmo_10.update_raw(f64::from(i));
265 }
266 assert_eq!(cmo_10.count, 5);
267 }
268
269 #[rstest]
270 fn test_reset_resets_inner_mas() {
271 let mut cmo = ChandeMomentumOscillator::new(3, None);
272 for price in [1.0, 2.0, 3.0] {
273 cmo.update_raw(price);
274 }
275 assert!(cmo.average_gain.initialized());
276 assert!(cmo.average_loss.initialized());
277 assert_ne!(cmo.average_gain.value(), 0.0);
278 cmo.reset();
279 assert!(!cmo.average_gain.initialized());
280 assert!(!cmo.average_loss.initialized());
281 assert_eq!(cmo.average_gain.value(), 0.0);
282 assert_eq!(cmo.average_loss.value(), 0.0);
283 }
284
285 #[rstest]
286 #[should_panic]
287 fn test_invalid_period_panics() {
288 let _ = ChandeMomentumOscillator::new(0, None);
289 }
290
291 #[rstest]
292 fn test_ma_type_propagation() {
293 let cmo = ChandeMomentumOscillator::new(5, Some(MovingAverageType::Simple));
294 assert_eq!(cmo.ma_type, MovingAverageType::Simple);
295 }
296
297 #[rstest]
298 fn test_zero_divisor_returns_zero() {
299 let mut cmo = ChandeMomentumOscillator::new(3, None);
300 for _ in 0..5 {
301 cmo.update_raw(100.0);
302 }
303 assert!(cmo.initialized);
304 assert_eq!(cmo.value, 0.0);
305 }
306
307 #[rstest]
308 fn test_random_walk_values_within_bounds() {
309 let prices = [
310 100.0, 100.5, 99.8, 100.3, 101.0, 100.7, 101.5, 101.2, 100.6, 101.1, 100.9, 101.4,
311 100.8, 101.2, 100.6, 100.9, 101.3, 101.0, 100.5, 101.1, 100.7, 101.4, 100.9, 100.8,
312 101.2, 100.6, 100.9, 101.3, 101.0, 100.5,
313 ];
314 let mut cmo = ChandeMomentumOscillator::new(10, None);
315 for price in prices {
316 cmo.update_raw(price);
317 }
318 assert!(cmo.initialized);
319 assert!(cmo.value <= 100.0 && cmo.value >= -100.0);
320 }
321}