Skip to main content

nautilus_indicators/momentum/
cmo.rs

1// -------------------------------------------------------------------------------------------------
2//  Copyright (C) 2015-2026 Nautech Systems Pty Ltd. All rights reserved.
3//  https://nautechsystems.io
4//
5//  Licensed under the GNU Lesser General Public License Version 3.0 (the "License");
6//  You may not use this file except in compliance with the License.
7//  You may obtain a copy of the License at https://www.gnu.org/licenses/lgpl-3.0.en.html
8//
9//  Unless required by applicable law or agreed to in writing, software
10//  distributed under the License is distributed on an "AS IS" BASIS,
11//  WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12//  See the License for the specific language governing permissions and
13//  limitations under the License.
14// -------------------------------------------------------------------------------------------------
15
16use 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    /// Creates a new [`ChandeMomentumOscillator`] instance.
89    ///
90    /// # Panics
91    ///
92    /// Panics if `period` is not positive (> 0).
93    #[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                // Divide before scaling, so a zero gain average gives exactly -100
140                // rather than a value just outside the oscillator's range.
141                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        // A drop followed by flat prices decays the gain average to exactly zero, so
170        // the oscillator sits on its lower bound. Scaling before dividing put it just
171        // outside, which `VariableIndexDynamicAverage` then read as a weight above one.
172        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}