Skip to main content

market_model/strategy/
ma_crossover.rs

1use crate::strategy::PriceStrategy;
2use crate::types::Signal;
3
4/// Moving average crossover strategy.
5///
6/// Long when the fast MA is above the slow MA, short otherwise.
7pub struct MaCrossover {
8    pub fast_window: usize,
9    pub slow_window: usize,
10}
11
12impl MaCrossover {
13    pub fn new(fast_window: usize, slow_window: usize) -> Self {
14        assert!(fast_window < slow_window);
15        assert!(fast_window > 0);
16        Self {
17            fast_window,
18            slow_window,
19        }
20    }
21}
22
23impl PriceStrategy for MaCrossover {
24    fn signal(&self, history: &[f64]) -> Signal {
25        if history.len() < self.slow_window {
26            return Signal::Flat;
27        }
28        let fast_ma: f64 = history[history.len() - self.fast_window..]
29            .iter()
30            .sum::<f64>()
31            / self.fast_window as f64;
32        let slow_ma: f64 = history[history.len() - self.slow_window..]
33            .iter()
34            .sum::<f64>()
35            / self.slow_window as f64;
36        if fast_ma > slow_ma {
37            Signal::Long
38        } else if fast_ma < slow_ma {
39            Signal::Short
40        } else {
41            Signal::Flat
42        }
43    }
44}
45
46#[cfg(test)]
47mod tests {
48    use super::*;
49
50    #[test]
51    fn test_ma_crossover_long() {
52        let s = MaCrossover::new(2, 4);
53        // Recent prices rising: slow trend flat, fast trend rising.
54        let history = vec![10.0, 10.0, 10.0, 10.0, 12.0, 14.0];
55        // fast MA = (12+14)/2 = 13, slow MA = (10+10+12+14)/4 = 11.5
56        assert_eq!(s.signal(&history), Signal::Long);
57    }
58
59    #[test]
60    fn test_ma_crossover_short() {
61        let s = MaCrossover::new(2, 4);
62        // Recent prices falling.
63        let history = vec![14.0, 14.0, 14.0, 12.0, 10.0, 8.0];
64        // fast MA = 9, slow MA = 11
65        assert_eq!(s.signal(&history), Signal::Short);
66    }
67
68    #[test]
69    fn test_ma_crossover_flat_insufficient_data() {
70        let s = MaCrossover::new(3, 5);
71        let history = vec![10.0, 11.0];
72        assert_eq!(s.signal(&history), Signal::Flat);
73    }
74}