1use crate::drift::detector::{DriftDetector, DriftLevel};
26use crate::error::{RillError, ensure_finite};
27
28#[derive(Debug, Clone)]
30#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
31#[non_exhaustive]
32pub struct AdwinConfig {
33 pub delta: f64,
36
37 pub warning_delta: f64,
40
41 pub max_window: usize,
45
46 pub min_samples: u64,
49}
50
51impl Default for AdwinConfig {
52 fn default() -> Self {
53 Self {
54 delta: 0.002,
55 warning_delta: 0.01,
56 max_window: 1000,
57 min_samples: 10,
58 }
59 }
60}
61
62#[derive(Debug, Clone)]
95#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
96pub struct Adwin {
97 config: AdwinConfig,
98 window: std::collections::VecDeque<f64>,
99 total: f64,
100 samples: u64,
101 current_level: DriftLevel,
102}
103
104impl Adwin {
105 pub fn new(config: AdwinConfig) -> Result<Self, RillError> {
113 ensure_finite("delta", config.delta)?;
114 if config.delta <= 0.0 || config.delta >= 1.0 {
115 return Err(RillError::InvalidSignificanceLevel(config.delta));
116 }
117 ensure_finite("warning_delta", config.warning_delta)?;
118 if config.warning_delta <= 0.0 || config.warning_delta >= 1.0 {
119 return Err(RillError::InvalidSignificanceLevel(config.warning_delta));
120 }
121 if config.warning_delta < config.delta {
122 return Err(RillError::InvalidParameter {
123 name: "warning_delta",
124 value: config.warning_delta,
125 });
126 }
127 if config.max_window == 0 {
128 return Err(RillError::InvalidCapacity(config.max_window));
129 }
130 if config.min_samples == 0 {
131 return Err(RillError::InvalidParameter {
132 name: "min_samples",
133 value: 0.0,
134 });
135 }
136 Ok(Self {
137 window: std::collections::VecDeque::with_capacity(config.max_window),
138 config,
139 total: 0.0,
140 samples: 0,
141 current_level: DriftLevel::None,
142 })
143 }
144
145 pub fn window_size(&self) -> usize {
147 self.window.len()
148 }
149
150 pub fn window_mean(&self) -> f64 {
152 if self.window.is_empty() {
153 0.0
154 } else {
155 self.total / self.window.len() as f64
156 }
157 }
158
159 pub const fn config(&self) -> &AdwinConfig {
161 &self.config
162 }
163
164 fn hoeffding_bound(n0: f64, n1: f64, n: u64, delta: f64) -> f64 {
167 let m = n0 * n1 / (n0 + n1);
168 let ln_n = (n as f64).ln().max(1.0);
169 let delta_eff = delta / ln_n;
170 (1.0 / (2.0 * m) * (4.0 / delta_eff).ln()).sqrt()
171 }
172
173 fn check_splits(&self) -> Option<(usize, DriftLevel, f64)> {
177 let n = self.window.len();
178 if n < 2 {
179 return None;
180 }
181 let mut prefix = Vec::with_capacity(n + 1);
183 prefix.push(0.0_f64);
184 let mut acc = 0.0;
185 for &v in &self.window {
186 acc += v;
187 prefix.push(acc);
188 }
189 let total = prefix[n];
190 let n_total = n as u64;
191
192 let mut best_split: Option<(usize, DriftLevel, f64)> = None;
193 for (k, &sum0) in prefix.iter().enumerate().take(n).skip(1) {
195 let n0 = k as f64;
196 let n1 = (n - k) as f64;
197 let sum1 = total - sum0;
198 let mean0 = sum0 / n0;
199 let mean1 = sum1 / n1;
200 let diff = (mean0 - mean1).abs();
201
202 let eps_drift = Self::hoeffding_bound(n0, n1, n_total, self.config.delta);
204 if diff > eps_drift {
205 return Some((k, DriftLevel::Drift, diff));
206 }
207 let eps_warn = Self::hoeffding_bound(n0, n1, n_total, self.config.warning_delta);
209 if diff > eps_warn && best_split.is_none() {
210 best_split = Some((k, DriftLevel::Warning, diff));
211 }
212 }
213 best_split
214 }
215
216 fn trim_front(&mut self, count: usize) {
218 for _ in 0..count {
219 if let Some(v) = self.window.pop_front() {
220 self.total -= v;
221 }
222 }
223 }
224}
225
226impl Default for Adwin {
227 fn default() -> Self {
228 Self::new(AdwinConfig::default()).expect("default config is valid")
229 }
230}
231
232impl DriftDetector for Adwin {
233 fn update(&mut self, value: f64) -> Result<DriftLevel, RillError> {
234 ensure_finite("value", value)?;
235 self.samples += 1;
236 self.window.push_back(value);
238 self.total += value;
239 if self.window.len() > self.config.max_window
241 && let Some(v) = self.window.pop_front()
242 {
243 self.total -= v;
244 }
245 if self.samples < self.config.min_samples || self.window.len() < 2 {
247 self.current_level = DriftLevel::None;
248 return Ok(DriftLevel::None);
249 }
250 if let Some((split, level, _diff)) = self.check_splits() {
252 if level == DriftLevel::Drift {
255 self.trim_front(split);
256 }
257 self.current_level = level;
258 } else {
259 self.current_level = DriftLevel::None;
260 }
261 Ok(self.current_level)
262 }
263
264 fn detected(&self) -> bool {
265 self.current_level == DriftLevel::Drift
266 }
267
268 fn warning(&self) -> bool {
269 self.current_level == DriftLevel::Warning
270 }
271
272 fn level(&self) -> DriftLevel {
273 self.current_level
274 }
275
276 fn samples_seen(&self) -> u64 {
277 self.samples
278 }
279
280 fn reset(&mut self) {
281 self.window.clear();
282 self.total = 0.0;
283 self.samples = 0;
284 self.current_level = DriftLevel::None;
285 }
286}
287
288#[cfg(test)]
289mod tests {
290 use super::*;
291
292 fn next_unit(seed: &mut u64) -> f64 {
294 *seed = seed
295 .wrapping_mul(6364136223846793005)
296 .wrapping_add(1442695040888963407);
297 ((*seed >> 11) as f64) / ((1u64 << 53) as f64)
298 }
299
300 #[test]
301 fn default_config_is_valid() {
302 let adwin = Adwin::default();
303 assert_eq!(adwin.samples_seen(), 0);
304 assert_eq!(adwin.level(), DriftLevel::None);
305 assert_eq!(adwin.window_size(), 0);
306 }
307
308 #[test]
309 fn detects_sudden_mean_shift() {
310 let mut adwin = Adwin::new(AdwinConfig {
311 delta: 0.05,
312 warning_delta: 0.1,
313 max_window: 500,
314 min_samples: 5,
315 })
316 .unwrap();
317 let mut seed = 42u64;
319 for _ in 0..100 {
320 let noise = 0.1 * (next_unit(&mut seed) - 0.5);
321 adwin.update(noise).unwrap();
322 }
323 assert_eq!(adwin.level(), DriftLevel::None);
324 let mut detected = false;
326 for _ in 0..200 {
327 let noise = 0.1 * (next_unit(&mut seed) - 0.5);
328 let level = adwin.update(5.0 + noise).unwrap();
329 if level == DriftLevel::Drift {
330 detected = true;
331 break;
332 }
333 }
334 assert!(detected, "ADWIN should detect the sudden mean shift");
335 }
336
337 #[test]
338 fn no_false_positive_on_stable_stream() {
339 let mut adwin = Adwin::new(AdwinConfig {
340 delta: 0.002,
341 warning_delta: 0.01,
342 max_window: 500,
343 min_samples: 10,
344 })
345 .unwrap();
346 let mut seed = 7u64;
347 for _ in 0..2000 {
348 let noise = 0.5 * (next_unit(&mut seed) - 0.5);
349 adwin.update(noise).unwrap();
350 }
351 assert!(
352 !adwin.detected(),
353 "false positive: drift reported on stable stream"
354 );
355 }
356
357 #[test]
358 fn detects_gradual_drift() {
359 let mut adwin = Adwin::new(AdwinConfig {
360 delta: 0.05,
361 warning_delta: 0.1,
362 max_window: 300,
363 min_samples: 5,
364 })
365 .unwrap();
366 let mut seed = 99u64;
368 let mut detected = false;
369 for i in 0..500 {
370 let mean = (i as f64 / 100.0).min(5.0);
371 let noise = 0.1 * (next_unit(&mut seed) - 0.5);
372 let level = adwin.update(mean + noise).unwrap();
373 if level == DriftLevel::Drift {
374 detected = true;
375 break;
376 }
377 }
378 assert!(detected, "ADWIN should detect gradual drift");
379 }
380
381 #[test]
382 fn window_trims_after_drift() {
383 let mut adwin = Adwin::new(AdwinConfig {
384 delta: 0.05,
385 warning_delta: 0.1,
386 max_window: 500,
387 min_samples: 5,
388 })
389 .unwrap();
390 for _ in 0..100 {
392 adwin.update(0.0).unwrap();
393 }
394 let size_before = adwin.window_size();
395 assert!(size_before > 0);
396 let mut trimmed = false;
398 for _ in 0..200 {
399 adwin.update(10.0).unwrap();
400 if adwin.detected() {
401 if adwin.window_size() < size_before + 200 {
405 trimmed = true;
406 break;
407 }
408 }
409 }
410 assert!(trimmed, "window should be trimmed after drift");
411 }
412
413 #[test]
414 fn max_window_enforced() {
415 let mut adwin = Adwin::new(AdwinConfig {
416 max_window: 50,
417 ..Default::default()
418 })
419 .unwrap();
420 for i in 0..200u64 {
421 adwin.update(i as f64).unwrap();
422 }
423 assert!(
426 adwin.window_size() <= 50,
427 "window should not exceed max_window, got {}",
428 adwin.window_size()
429 );
430 }
431
432 #[test]
433 fn min_samples_gates_detection() {
434 let mut adwin = Adwin::new(AdwinConfig {
435 delta: 0.5,
436 warning_delta: 0.5,
437 max_window: 100,
438 min_samples: 50,
439 })
440 .unwrap();
441 for _ in 0..48 {
443 adwin.update(0.0).unwrap();
444 }
445 adwin.update(100.0).unwrap();
447 assert_eq!(adwin.level(), DriftLevel::None);
448 let mut detected = false;
452 for _ in 0..50 {
453 let level = adwin.update(100.0).unwrap();
454 if level.is_change() {
455 detected = true;
456 }
457 }
458 assert!(detected, "should have detected drift after min_samples");
459 }
460
461 #[test]
462 fn reset_clears_state() {
463 let mut adwin = Adwin::default();
464 for _ in 0..50 {
465 adwin.update(1.0).unwrap();
466 }
467 assert!(adwin.window_size() > 0);
468 adwin.reset();
469 assert_eq!(adwin.window_size(), 0);
470 assert_eq!(adwin.samples_seen(), 0);
471 assert_eq!(adwin.level(), DriftLevel::None);
472 assert_eq!(adwin.window_mean(), 0.0);
473 }
474
475 #[test]
476 fn rejects_non_finite_input() {
477 let mut adwin = Adwin::default();
478 assert!(adwin.update(f64::NAN).is_err());
479 assert!(adwin.update(f64::INFINITY).is_err());
480 assert!(adwin.update(f64::NEG_INFINITY).is_err());
481 assert_eq!(adwin.samples_seen(), 0);
482 assert_eq!(adwin.window_size(), 0);
483 }
484
485 #[test]
486 fn rejects_invalid_config() {
487 assert!(
489 Adwin::new(AdwinConfig {
490 delta: 0.0,
491 ..Default::default()
492 })
493 .is_err()
494 );
495 assert!(
496 Adwin::new(AdwinConfig {
497 delta: 1.0,
498 ..Default::default()
499 })
500 .is_err()
501 );
502 assert!(
504 Adwin::new(AdwinConfig {
505 delta: 0.05,
506 warning_delta: 0.01,
507 ..Default::default()
508 })
509 .is_err()
510 );
511 assert!(
513 Adwin::new(AdwinConfig {
514 max_window: 0,
515 ..Default::default()
516 })
517 .is_err()
518 );
519 assert!(
521 Adwin::new(AdwinConfig {
522 min_samples: 0,
523 ..Default::default()
524 })
525 .is_err()
526 );
527 }
528
529 #[test]
530 fn window_mean_correct() {
531 let mut adwin = Adwin::new(AdwinConfig {
532 max_window: 100,
533 min_samples: 11, ..Default::default()
535 })
536 .unwrap();
537 for i in 1..=10 {
538 adwin.update(i as f64).unwrap();
539 }
540 assert!((adwin.window_mean() - 5.5).abs() < 1e-9);
542 }
543
544 #[test]
545 fn hoeffding_bound_decreases_with_more_data() {
546 let b1 = Adwin::hoeffding_bound(5.0, 5.0, 10, 0.01);
548 let b2 = Adwin::hoeffding_bound(50.0, 50.0, 100, 0.01);
549 assert!(
550 b2 < b1,
551 "bound should decrease with more data: {} vs {}",
552 b2,
553 b1
554 );
555 }
556
557 #[cfg(feature = "serde")]
558 #[test]
559 fn serde_roundtrip() {
560 let mut adwin = Adwin::new(AdwinConfig {
561 delta: 0.01,
562 warning_delta: 0.05,
563 max_window: 200,
564 min_samples: 5,
565 })
566 .unwrap();
567 for i in 0..50 {
568 adwin.update(i as f64 * 0.1).unwrap();
569 }
570 let json = serde_json::to_string(&adwin).unwrap();
571 let restored: Adwin = serde_json::from_str(&json).unwrap();
572 assert_eq!(restored.samples_seen(), 50);
573 assert_eq!(restored.window_size(), adwin.window_size());
574 assert!((restored.window_mean() - adwin.window_mean()).abs() < 1e-12);
575 assert_eq!(restored.level(), adwin.level());
576 }
577}