nautilus_indicators/momentum/
rsi.rs1use std::fmt::{Debug, Display};
17
18use nautilus_model::{
19 data::{Bar, QuoteTick, TradeTick},
20 enums::PriceType,
21};
22
23use crate::{
24 average::{MovingAverageFactory, MovingAverageType},
25 indicator::{Indicator, MovingAverage},
26};
27
28#[repr(C)]
30#[derive(Debug)]
31#[cfg_attr(
32 feature = "python",
33 pyo3::pyclass(module = "nautilus_trader.indicators", unsendable)
34)]
35#[cfg_attr(
36 feature = "python",
37 pyo3_stub_gen::derive::gen_stub_pyclass(module = "nautilus_trader.indicators")
38)]
39pub struct RelativeStrengthIndex {
40 pub period: usize,
41 pub ma_type: MovingAverageType,
42 pub value: f64,
43 pub count: usize,
44 pub initialized: bool,
45 has_inputs: bool,
46 last_value: f64,
47 average_gain: Box<dyn MovingAverage + Send + 'static>,
48 average_loss: Box<dyn MovingAverage + Send + 'static>,
49 rsi_max: f64,
50}
51
52impl Display for RelativeStrengthIndex {
53 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
54 write!(f, "{}({},{})", self.name(), self.period, self.ma_type)
55 }
56}
57
58impl Indicator for RelativeStrengthIndex {
59 fn name(&self) -> String {
60 stringify!(RelativeStrengthIndex).to_string()
61 }
62
63 fn has_inputs(&self) -> bool {
64 self.has_inputs
65 }
66
67 fn initialized(&self) -> bool {
68 self.initialized
69 }
70
71 fn handle_quote(&mut self, quote: &QuoteTick) -> anyhow::Result<()> {
72 self.update_raw(quote.extract_price(PriceType::Mid)?.into());
73 Ok(())
74 }
75
76 fn handle_trade(&mut self, trade: &TradeTick) {
77 self.update_raw((trade.price).into());
78 }
79
80 fn handle_bar(&mut self, bar: &Bar) {
81 self.update_raw((&bar.close).into());
82 }
83
84 fn reset(&mut self) {
85 self.value = 0.0;
86 self.last_value = 0.0;
87 self.count = 0;
88 self.has_inputs = false;
89 self.initialized = false;
90 self.average_gain.reset();
91 self.average_loss.reset();
92 }
93}
94
95impl RelativeStrengthIndex {
96 #[must_use]
98 pub fn new(period: usize, ma_type: Option<MovingAverageType>) -> Self {
99 let ma_type = ma_type.unwrap_or(MovingAverageType::Exponential);
100 Self {
101 period,
102 ma_type,
103 value: 0.0,
104 last_value: 0.0,
105 count: 0,
106 has_inputs: false,
107 average_gain: MovingAverageFactory::create(ma_type, period),
108 average_loss: MovingAverageFactory::create(ma_type, period),
109 rsi_max: 1.0,
110 initialized: false,
111 }
112 }
113
114 pub fn update_raw(&mut self, value: f64) {
115 if !self.has_inputs {
116 self.last_value = value;
117 self.has_inputs = true;
118 }
119 let gain = value - self.last_value;
120 if gain > 0.0 {
121 self.average_gain.update_raw(gain);
122 self.average_loss.update_raw(0.0);
123 } else if gain < 0.0 {
124 self.average_loss.update_raw(-gain);
125 self.average_gain.update_raw(0.0);
126 } else {
127 self.average_loss.update_raw(0.0);
128 self.average_gain.update_raw(0.0);
129 }
130 self.count = self.average_gain.count();
131 if !self.initialized && self.average_loss.initialized() && self.average_gain.initialized() {
132 self.initialized = true;
133 }
134
135 if self.average_loss.value() == 0.0 {
136 self.value = self.rsi_max;
137 self.last_value = value;
138 return;
139 }
140
141 let rs = self.average_gain.value() / self.average_loss.value();
142 self.value = self.rsi_max - (self.rsi_max / (1.0 + rs));
143 self.last_value = value;
144
145 if !self.initialized && self.count >= self.period {
146 self.initialized = true;
147 }
148 }
149}
150
151#[cfg(test)]
152mod tests {
153 use nautilus_model::data::{Bar, QuoteTick, TradeTick};
154 use rstest::rstest;
155
156 use crate::{
157 average::MovingAverageType, indicator::Indicator, momentum::rsi::RelativeStrengthIndex,
158 stubs::*, testing::assert_approx_equal,
159 };
160
161 #[rstest]
162 fn test_rsi_initialized(rsi_10: RelativeStrengthIndex) {
163 let display_str = format!("{rsi_10}");
164 assert_eq!(display_str, "RelativeStrengthIndex(10,EXPONENTIAL)");
165 assert_eq!(rsi_10.period, 10);
166 assert!(!rsi_10.initialized);
167 }
168
169 #[rstest]
170 fn test_initialized_with_required_inputs_returns_true(mut rsi_10: RelativeStrengthIndex) {
171 for i in 0..12 {
172 rsi_10.update_raw(f64::from(i));
173 }
174 assert!(rsi_10.initialized);
175 }
176
177 #[rstest]
178 fn test_value_with_one_input_returns_expected_value(mut rsi_10: RelativeStrengthIndex) {
179 rsi_10.update_raw(1.0);
180 assert_eq!(rsi_10.value, 1.0);
181 }
182
183 #[rstest]
184 fn test_value_all_higher_inputs_returns_expected_value(mut rsi_10: RelativeStrengthIndex) {
185 for i in 1..4 {
186 rsi_10.update_raw(f64::from(i));
187 }
188 assert_eq!(rsi_10.value, 1.0);
189 }
190
191 #[rstest]
192 fn test_value_with_all_lower_inputs_returns_expected_value(mut rsi_10: RelativeStrengthIndex) {
193 for i in (1..4).rev() {
194 rsi_10.update_raw(f64::from(i));
195 }
196 assert_eq!(rsi_10.value, 0.0);
197 }
198
199 #[rstest]
200 fn test_value_with_various_input_returns_expected_value(mut rsi_10: RelativeStrengthIndex) {
201 rsi_10.update_raw(3.0);
202 rsi_10.update_raw(2.0);
203 rsi_10.update_raw(5.0);
204 rsi_10.update_raw(6.0);
205 rsi_10.update_raw(7.0);
206 rsi_10.update_raw(6.0);
207
208 assert_approx_equal(rsi_10.value, 0.683736332583);
209 }
210
211 #[rstest]
212 fn test_value_at_returns_expected_value(mut rsi_10: RelativeStrengthIndex) {
213 rsi_10.update_raw(3.0);
214 rsi_10.update_raw(2.0);
215 rsi_10.update_raw(5.0);
216 rsi_10.update_raw(6.0);
217 rsi_10.update_raw(7.0);
218 rsi_10.update_raw(6.0);
219 rsi_10.update_raw(6.0);
220 rsi_10.update_raw(7.0);
221
222 assert_approx_equal(rsi_10.value, 0.761534466766);
223 }
224
225 #[rstest]
226 fn test_reset(mut rsi_10: RelativeStrengthIndex) {
227 rsi_10.update_raw(1.0);
228 rsi_10.update_raw(2.0);
229 rsi_10.reset();
230 assert!(!rsi_10.initialized());
231 assert_eq!(rsi_10.count, 0);
232 }
233
234 #[rstest]
235 fn test_reset_resets_inner_mas(mut rsi_10: RelativeStrengthIndex) {
236 rsi_10.update_raw(1.0);
237 rsi_10.update_raw(2.0);
238 rsi_10.reset();
239 assert_eq!(rsi_10.average_gain.count(), 0);
240 assert_eq!(rsi_10.average_loss.count(), 0);
241 }
242
243 #[rstest]
244 fn test_handle_quote_tick(mut rsi_10: RelativeStrengthIndex, stub_quote: QuoteTick) {
245 rsi_10.handle_quote(&stub_quote).unwrap();
246 assert_eq!(rsi_10.count, 1);
247 assert_eq!(rsi_10.value, 1.0);
248 }
249
250 #[rstest]
251 fn test_handle_trade_tick(mut rsi_10: RelativeStrengthIndex, stub_trade: TradeTick) {
252 rsi_10.handle_trade(&stub_trade);
253 assert_eq!(rsi_10.count, 1);
254 assert_eq!(rsi_10.value, 1.0);
255 }
256
257 #[rstest]
258 fn test_handle_bar(mut rsi_10: RelativeStrengthIndex, bar_ethusdt_binance_minute_bid: Bar) {
259 rsi_10.handle_bar(&bar_ethusdt_binance_minute_bid);
260 assert_eq!(rsi_10.count, 1);
261 assert_eq!(rsi_10.value, 1.0);
262 }
263
264 #[rstest]
265 fn test_constant_inputs_initializes_and_value_max(mut rsi_10: RelativeStrengthIndex) {
266 for _ in 0..12 {
267 rsi_10.update_raw(5.0);
268 }
269 assert!(rsi_10.initialized);
270 assert_eq!(rsi_10.value, 1.0);
271 }
272
273 #[rstest]
274 fn test_reset_resets_has_inputs_and_value(mut rsi_10: RelativeStrengthIndex) {
275 rsi_10.update_raw(1.0);
276 rsi_10.reset();
277 assert!(!rsi_10.has_inputs());
278 assert_eq!(rsi_10.value, 0.0);
279 }
280
281 fn run_rsi(values: &[f64], period: usize, ma_type: MovingAverageType) -> f64 {
283 let mut rsi = RelativeStrengthIndex::new(period, Some(ma_type));
284 for &v in values {
285 rsi.update_raw(v);
286 }
287 rsi.value
288 }
289
290 #[rstest]
291 fn test_ma_type_is_plumbed_into_inner_averages() {
292 let series = [
296 44.34, 44.09, 44.15, 43.61, 44.33, 44.83, 45.10, 45.42, 45.84, 46.08, 45.89, 46.03,
297 45.61, 46.28, 46.28,
298 ];
299
300 let wilder = run_rsi(&series, 14, MovingAverageType::Wilder);
301 let simple = run_rsi(&series, 14, MovingAverageType::Simple);
302 let exponential = run_rsi(&series, 14, MovingAverageType::Exponential);
303
304 assert_ne!(wilder, simple);
305 assert_ne!(wilder, exponential);
306 assert_ne!(simple, exponential);
307 }
308
309 #[rstest]
310 fn test_recovers_below_max_after_losses() {
311 let mut values: Vec<f64> = (1..=15).map(f64::from).collect();
315 values.extend([14.0, 12.0, 9.0, 5.0, 2.0]);
316
317 let value = run_rsi(&values, 14, MovingAverageType::Wilder);
318 assert!(
319 value < 1.0,
320 "RSI should drop below rsi_max after losses, was {value}"
321 );
322 }
323
324 #[rstest]
325 fn test_wilder_golden_series() {
326 let base: Vec<f64> = (1..=15).map(f64::from).collect();
329 let downs = [14.0, 12.0, 9.0, 5.0, 2.0];
330 let expected = [0.8935, 0.7269, 0.5586, 0.4192, 0.3489];
331
332 let mut rsi = RelativeStrengthIndex::new(14, Some(MovingAverageType::Wilder));
333 for &v in &base {
334 rsi.update_raw(v);
335 }
336
337 for (i, &v) in downs.iter().enumerate() {
338 rsi.update_raw(v);
339 assert!(
340 (rsi.value - expected[i]).abs() < 1e-4,
341 "step {i}: expected {}, was {}",
342 expected[i],
343 rsi.value
344 );
345 }
346 }
347}