1use crate::error::{RillError, checked_increment, ensure_finite};
9#[cfg(feature = "serde")]
10use crate::persistence::ValidateState;
11use crate::traits::OnlineStatistic;
12
13#[derive(Debug, Clone)]
31#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
32pub struct ExponentiallyWeightedMean {
33 alpha: f64,
34 mean: f64,
35 count: u64,
36}
37
38impl ExponentiallyWeightedMean {
39 pub fn new(alpha: f64) -> Result<Self, RillError> {
43 ensure_finite("alpha", alpha)?;
44 if alpha <= 0.0 || alpha > 1.0 {
45 return Err(RillError::InvalidParameter {
46 name: "alpha",
47 value: alpha,
48 });
49 }
50 Ok(Self {
51 alpha,
52 mean: 0.0,
53 count: 0,
54 })
55 }
56
57 pub const fn alpha(&self) -> f64 {
59 self.alpha
60 }
61
62 pub const fn value(&self) -> f64 {
64 self.mean
65 }
66
67 pub const fn count(&self) -> u64 {
69 self.count
70 }
71}
72
73#[cfg(feature = "serde")]
74impl ValidateState for ExponentiallyWeightedMean {
75 fn validate_state(&self) -> Result<(), RillError> {
76 ensure_finite("alpha", self.alpha)?;
77 ensure_finite("ew mean", self.mean)?;
78 if self.alpha <= 0.0 || self.alpha > 1.0 {
79 return Err(RillError::InvalidParameter {
80 name: "alpha",
81 value: self.alpha,
82 });
83 }
84 Ok(())
85 }
86}
87
88impl OnlineStatistic for ExponentiallyWeightedMean {
89 fn update(&mut self, value: f64) -> Result<(), RillError> {
90 ensure_finite("value", value)?;
91 let next_count = checked_increment(self.count, "EW mean sample")?;
92 let next_mean = if self.count == 0 {
93 value
94 } else {
95 self.alpha * value + (1.0 - self.alpha) * self.mean
96 };
97 ensure_finite("EW mean", next_mean)?;
98 self.mean = next_mean;
99 self.count = next_count;
100 Ok(())
101 }
102
103 fn samples_seen(&self) -> u64 {
104 self.count
105 }
106
107 fn reset(&mut self) {
108 self.mean = 0.0;
109 self.count = 0;
110 }
111}
112
113#[cfg(test)]
114mod tests {
115 use super::*;
116
117 #[test]
118 fn first_sample_seeds_mean() {
119 let mut ew = ExponentiallyWeightedMean::new(0.3).unwrap();
120 ew.update(10.0).unwrap();
121 assert!((ew.value() - 10.0).abs() < 1e-12);
122 }
123
124 #[test]
125 fn weighted_update_matches_formula() {
126 let mut ew = ExponentiallyWeightedMean::new(0.5).unwrap();
127 ew.update(10.0).unwrap();
128 ew.update(20.0).unwrap();
129 assert!((ew.value() - 15.0).abs() < 1e-12);
130 }
131
132 #[test]
133 fn alpha_one_tracks_last_value() {
134 let mut ew = ExponentiallyWeightedMean::new(1.0).unwrap();
135 ew.update(3.0).unwrap();
136 ew.update(7.0).unwrap();
137 assert!((ew.value() - 7.0).abs() < 1e-12);
138 }
139
140 #[test]
141 fn invalid_alpha_rejected() {
142 assert!(ExponentiallyWeightedMean::new(0.0).is_err());
143 assert!(ExponentiallyWeightedMean::new(-0.1).is_err());
144 assert!(ExponentiallyWeightedMean::new(1.5).is_err());
145 assert!(ExponentiallyWeightedMean::new(f64::NAN).is_err());
146 }
147
148 #[test]
149 fn reset_clears_state() {
150 let mut ew = ExponentiallyWeightedMean::new(0.5).unwrap();
151 ew.update(10.0).unwrap();
152 ew.reset();
153 assert_eq!(ew.count(), 0);
154 assert_eq!(ew.value(), 0.0);
155 }
156}