1use crate::error::{RillError, checked_increment, ensure_finite};
6#[cfg(feature = "serde")]
7use crate::persistence::ValidateState;
8use crate::traits::OnlineStatistic;
9
10#[derive(Debug, Clone, Default)]
27#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
28pub struct Mean {
29 count: u64,
30 mean: f64,
31}
32
33impl Mean {
34 pub const fn new() -> Self {
36 Self {
37 count: 0,
38 mean: 0.0,
39 }
40 }
41
42 pub const fn value(&self) -> f64 {
44 self.mean
45 }
46
47 pub const fn count(&self) -> u64 {
49 self.count
50 }
51}
52
53#[cfg(feature = "serde")]
54impl ValidateState for Mean {
55 fn validate_state(&self) -> Result<(), RillError> {
56 ensure_finite("mean", self.mean)?;
57 Ok(())
58 }
59}
60
61impl OnlineStatistic for Mean {
62 fn update(&mut self, value: f64) -> Result<(), RillError> {
63 ensure_finite("value", value)?;
64 let next_count = checked_increment(self.count, "mean sample")?;
65 let delta = value - self.mean;
66 ensure_finite("mean delta", delta)?;
67 let next_mean = self.mean + delta / next_count as f64;
68 ensure_finite("mean", next_mean)?;
69
70 self.count = next_count;
71 self.mean = next_mean;
72 Ok(())
73 }
74
75 fn samples_seen(&self) -> u64 {
76 self.count
77 }
78
79 fn reset(&mut self) {
80 self.count = 0;
81 self.mean = 0.0;
82 }
83}
84
85#[cfg(test)]
86mod tests {
87 use super::*;
88 use rand::SeedableRng;
89
90 #[test]
91 fn mean_of_simple_sequence() {
92 let mut m = Mean::new();
93 for x in [1.0, 2.0, 3.0, 4.0, 5.0] {
94 m.update(x).unwrap();
95 }
96 assert_eq!(m.value(), 3.0);
97 assert_eq!(m.count(), 5);
98 }
99
100 #[test]
101 fn mean_empty_is_zero() {
102 let m = Mean::new();
103 assert_eq!(m.value(), 0.0);
104 assert_eq!(m.count(), 0);
105 }
106
107 #[test]
108 fn mean_rejects_non_finite() {
109 let mut m = Mean::new();
110 assert!(m.update(f64::NAN).is_err());
111 assert!(m.update(f64::INFINITY).is_err());
112 assert_eq!(m.count(), 0);
113 }
114
115 #[test]
116 fn mean_rejects_overflow_without_mutating_state() {
117 let mut m = Mean::new();
118 m.update(f64::MAX).unwrap();
119 let before = m.clone();
120 assert!(m.update(-f64::MAX).is_err());
121 assert_eq!(m.count(), before.count());
122 assert_eq!(m.value(), before.value());
123 }
124
125 #[test]
126 fn mean_reset() {
127 let mut m = Mean::new();
128 m.update(10.0).unwrap();
129 m.update(20.0).unwrap();
130 m.reset();
131 assert_eq!(m.count(), 0);
132 assert_eq!(m.value(), 0.0);
133 }
134
135 #[test]
136 fn mean_matches_batch_formula() {
137 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(42);
138 let mut m = Mean::new();
139 let mut data = Vec::new();
140 for _ in 0..1000 {
141 let x = rand::Rng::gen_range(&mut rng, -100.0..100.0);
142 m.update(x).unwrap();
143 data.push(x);
144 }
145 let batch: f64 = data.iter().sum::<f64>() / data.len() as f64;
146 assert!((m.value() - batch).abs() < 1e-9);
147 }
148}