wickra_core/indicators/
sharpe_ratio.rs1use std::collections::VecDeque;
4
5use crate::error::{Error, Result};
6use crate::indicators::rolling_moments::ShiftedMoments;
7use crate::traits::Indicator;
8
9#[derive(Debug, Clone)]
45pub struct SharpeRatio {
46 period: usize,
47 risk_free: f64,
48 window: VecDeque<f64>,
49 moments: ShiftedMoments,
50}
51
52impl SharpeRatio {
53 pub fn new(period: usize, risk_free: f64) -> Result<Self> {
60 if period < 2 {
61 return Err(Error::InvalidPeriod {
62 message: "sharpe ratio needs period >= 2",
63 });
64 }
65 if period > crate::error::MAX_PERIOD {
66 return Err(Error::InvalidPeriod {
67 message: crate::error::PERIOD_ABOVE_MAX,
68 });
69 }
70 Ok(Self {
71 period,
72 risk_free,
73 window: VecDeque::with_capacity(period),
74 moments: ShiftedMoments::new(),
75 })
76 }
77
78 pub const fn period(&self) -> usize {
80 self.period
81 }
82
83 pub const fn risk_free(&self) -> f64 {
85 self.risk_free
86 }
87}
88
89impl Indicator for SharpeRatio {
90 type Input = f64;
91 type Output = f64;
92
93 #[inline]
94 fn update(&mut self, input: f64) -> Option<f64> {
95 if !input.is_finite() {
96 return None;
97 }
98 if self.window.len() == self.period {
99 let old = self.window.pop_front().expect("non-empty");
100 self.moments.evict(old);
101 }
102 self.window.push_back(input);
103 self.moments.push(input);
104 if self.moments.needs_reseed(self.period) {
105 self.moments.reseed(self.window.iter().copied());
106 }
107 if self.window.len() < self.period {
108 return None;
109 }
110 let mean = self.moments.mean(self.period);
111 let var = self.moments.sample_variance(self.period);
112 let sd = var.sqrt();
113 if sd == 0.0 {
114 return Some(0.0);
115 }
116 Some((mean - self.risk_free) / sd)
117 }
118
119 fn reset(&mut self) {
120 self.window.clear();
121 self.moments.reset();
122 }
123
124 #[inline]
125 fn warmup_period(&self) -> usize {
126 self.period
127 }
128
129 #[inline]
130 fn is_ready(&self) -> bool {
131 self.window.len() == self.period
132 }
133
134 #[inline]
135 fn name(&self) -> &'static str {
136 "SharpeRatio"
137 }
138}
139
140#[cfg(test)]
141mod tests {
142 use super::*;
143 use crate::traits::BatchExt;
144 use approx::assert_relative_eq;
145
146 #[test]
147 fn rejects_period_less_than_two() {
148 assert!(matches!(
149 SharpeRatio::new(1, 0.0),
150 Err(Error::InvalidPeriod { .. })
151 ));
152 assert!(matches!(
153 SharpeRatio::new(0, 0.0),
154 Err(Error::InvalidPeriod { .. })
155 ));
156 }
157
158 #[test]
159 fn accessors_and_metadata() {
160 let sr = SharpeRatio::new(20, 0.001).unwrap();
161 assert_eq!(sr.period(), 20);
162 assert_relative_eq!(sr.risk_free(), 0.001, epsilon = 1e-12);
163 assert_eq!(sr.name(), "SharpeRatio");
164 assert_eq!(sr.warmup_period(), 20);
165 }
166
167 #[test]
168 fn constant_returns_yield_zero() {
169 let mut sr = SharpeRatio::new(5, 0.0).unwrap();
170 let out = sr.batch(&[0.01; 10]);
171 for v in out.into_iter().flatten() {
172 assert_relative_eq!(v, 0.0, epsilon = 1e-12);
173 }
174 }
175
176 #[test]
177 fn reference_value() {
178 let mut sr = SharpeRatio::new(4, 0.0).unwrap();
183 let out = sr.batch(&[0.01, 0.02, 0.03, 0.04]);
184 let expected = 0.025_f64 / (0.000_166_666_666_666_666_67_f64).sqrt();
185 assert_relative_eq!(out[3].unwrap(), expected, epsilon = 1e-9);
186 }
187
188 #[test]
189 fn ignores_non_finite_input() {
190 let mut sr = SharpeRatio::new(3, 0.0).unwrap();
191 assert_eq!(sr.update(0.01), None);
192 assert_eq!(sr.update(f64::NAN), None);
193 assert_eq!(sr.update(0.02), None);
194 assert!(sr.update(0.03).is_some());
195 }
196
197 #[test]
198 fn warmup_returns_none() {
199 let mut sr = SharpeRatio::new(5, 0.0).unwrap();
200 for i in 0..4 {
201 assert_eq!(sr.update(f64::from(i) * 0.01), None);
202 }
203 assert!(sr.update(0.05).is_some());
204 }
205
206 #[test]
207 fn reset_clears_state() {
208 let mut sr = SharpeRatio::new(3, 0.0).unwrap();
209 sr.batch(&[0.01, 0.02, 0.03]);
210 assert!(sr.is_ready());
211 sr.reset();
212 assert!(!sr.is_ready());
213 assert_eq!(sr.update(0.01), None);
214 }
215
216 #[test]
217 fn batch_equals_streaming() {
218 let returns: Vec<f64> = (0..50)
219 .map(|i| 0.001 + (f64::from(i) * 0.2).sin() * 0.01)
220 .collect();
221 let batch = SharpeRatio::new(10, 0.0).unwrap().batch(&returns);
222 let mut s = SharpeRatio::new(10, 0.0).unwrap();
223 let streamed: Vec<_> = returns.iter().map(|p| s.update(*p)).collect();
224 assert_eq!(batch, streamed);
225 }
226}