1use crate::drift::detector::{DriftDetector, DriftLevel};
26use crate::error::{RillError, checked_finite_add, checked_increment, ensure_finite};
27use crate::persistence::ValidateState;
28
29pub const ADWIN_PORTABLE_STATE_VERSION: u32 = 1;
31
32#[derive(Debug, Clone)]
34#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
35#[non_exhaustive]
36pub struct AdwinConfig {
37 pub delta: f64,
40
41 pub warning_delta: f64,
44
45 pub max_window: usize,
49
50 pub min_samples: u64,
53}
54
55#[derive(Debug, Clone, PartialEq)]
60#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
61#[cfg_attr(feature = "serde", serde(deny_unknown_fields))]
62pub struct AdwinPortableStateV1 {
63 pub version: u32,
65 pub delta: f64,
67 pub warning_delta: f64,
69 pub max_window: usize,
71 pub min_samples: u64,
73 pub window: Vec<f64>,
75 pub total: f64,
77 pub samples: u64,
79 pub current_level: DriftLevel,
81}
82
83impl ValidateState for AdwinPortableStateV1 {
84 fn validate_state(&self) -> Result<(), RillError> {
85 if self.version != ADWIN_PORTABLE_STATE_VERSION {
86 return Err(RillError::IncompatibleStateVersion {
87 expected: ADWIN_PORTABLE_STATE_VERSION,
88 actual: self.version,
89 });
90 }
91 Adwin::new(AdwinConfig {
92 delta: self.delta,
93 warning_delta: self.warning_delta,
94 max_window: self.max_window,
95 min_samples: self.min_samples,
96 })?;
97 if self.window.len() > self.max_window {
98 return Err(RillError::InvalidState(
99 "ADWIN portable window exceeds max_window".to_owned(),
100 ));
101 }
102 if self.samples < self.window.len() as u64 {
103 return Err(RillError::InvalidState(
104 "ADWIN samples is smaller than the retained window".to_owned(),
105 ));
106 }
107 ensure_finite("portable ADWIN total", self.total)?;
108 let mut recomputed = 0.0;
109 for &value in &self.window {
110 ensure_finite("portable ADWIN window value", value)?;
111 recomputed = checked_finite_add(recomputed, value, "portable ADWIN window sum")?;
112 }
113 let tolerance = 1e-10 * recomputed.abs().max(self.total.abs()).max(1.0);
114 if (recomputed - self.total).abs() > tolerance {
115 return Err(RillError::InvalidState(
116 "ADWIN portable total does not match the retained window".to_owned(),
117 ));
118 }
119 if self.samples < self.min_samples && self.current_level != DriftLevel::None {
120 return Err(RillError::InvalidState(
121 "ADWIN state reports a level before min_samples".to_owned(),
122 ));
123 }
124 Ok(())
125 }
126}
127
128impl Default for AdwinConfig {
129 fn default() -> Self {
130 Self {
131 delta: 0.002,
132 warning_delta: 0.01,
133 max_window: 1000,
134 min_samples: 10,
135 }
136 }
137}
138
139#[derive(Debug, Clone)]
172#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
173pub struct Adwin {
174 config: AdwinConfig,
175 window: std::collections::VecDeque<f64>,
176 total: f64,
177 samples: u64,
178 current_level: DriftLevel,
179}
180
181impl Adwin {
182 pub fn new(config: AdwinConfig) -> Result<Self, RillError> {
190 ensure_finite("delta", config.delta)?;
191 if config.delta <= 0.0 || config.delta >= 1.0 {
192 return Err(RillError::InvalidSignificanceLevel(config.delta));
193 }
194 ensure_finite("warning_delta", config.warning_delta)?;
195 if config.warning_delta <= 0.0 || config.warning_delta >= 1.0 {
196 return Err(RillError::InvalidSignificanceLevel(config.warning_delta));
197 }
198 if config.warning_delta < config.delta {
199 return Err(RillError::InvalidParameter {
200 name: "warning_delta",
201 value: config.warning_delta,
202 });
203 }
204 if config.max_window == 0 {
205 return Err(RillError::InvalidCapacity(config.max_window));
206 }
207 if config.min_samples == 0 {
208 return Err(RillError::InvalidParameter {
209 name: "min_samples",
210 value: 0.0,
211 });
212 }
213 Ok(Self {
214 window: std::collections::VecDeque::with_capacity(config.max_window),
215 config,
216 total: 0.0,
217 samples: 0,
218 current_level: DriftLevel::None,
219 })
220 }
221
222 pub fn window_size(&self) -> usize {
224 self.window.len()
225 }
226
227 pub fn window_mean(&self) -> f64 {
229 if self.window.is_empty() {
230 0.0
231 } else {
232 self.total / self.window.len() as f64
233 }
234 }
235
236 pub const fn config(&self) -> &AdwinConfig {
238 &self.config
239 }
240
241 pub fn export_state_v1(&self) -> AdwinPortableStateV1 {
243 AdwinPortableStateV1 {
244 version: ADWIN_PORTABLE_STATE_VERSION,
245 delta: self.config.delta,
246 warning_delta: self.config.warning_delta,
247 max_window: self.config.max_window,
248 min_samples: self.config.min_samples,
249 window: self.window.iter().copied().collect(),
250 total: self.total,
251 samples: self.samples,
252 current_level: self.current_level,
253 }
254 }
255
256 pub fn restore_state_v1(
258 config: AdwinConfig,
259 state: AdwinPortableStateV1,
260 ) -> Result<Self, RillError> {
261 state.validate_state()?;
262 if config.delta != state.delta
263 || config.warning_delta != state.warning_delta
264 || config.max_window != state.max_window
265 || config.min_samples != state.min_samples
266 {
267 return Err(RillError::InvalidState(
268 "ADWIN portable state configuration mismatch".to_owned(),
269 ));
270 }
271 Adwin::new(config.clone())?;
272 Ok(Self {
273 config,
274 window: state.window.into(),
275 total: state.total,
276 samples: state.samples,
277 current_level: state.current_level,
278 })
279 }
280
281 fn hoeffding_bound(n0: f64, n1: f64, n: u64, delta: f64) -> f64 {
284 let m = n0 * n1 / (n0 + n1);
285 let ln_n = (n as f64).ln().max(1.0);
286 let delta_eff = delta / ln_n;
287 (1.0 / (2.0 * m) * (4.0 / delta_eff).ln()).sqrt()
288 }
289
290 fn check_splits(&self) -> Option<(usize, DriftLevel, f64)> {
294 let n = self.window.len();
295 if n < 2 {
296 return None;
297 }
298 let mut prefix = Vec::with_capacity(n + 1);
300 prefix.push(0.0_f64);
301 let mut acc = 0.0;
302 for &v in &self.window {
303 acc += v;
304 prefix.push(acc);
305 }
306 let total = prefix[n];
307 let n_total = n as u64;
308
309 let mut best_split: Option<(usize, DriftLevel, f64)> = None;
310 for (k, &sum0) in prefix.iter().enumerate().take(n).skip(1) {
312 let n0 = k as f64;
313 let n1 = (n - k) as f64;
314 let sum1 = total - sum0;
315 let mean0 = sum0 / n0;
316 let mean1 = sum1 / n1;
317 let diff = (mean0 - mean1).abs();
318
319 let eps_drift = Self::hoeffding_bound(n0, n1, n_total, self.config.delta);
321 if diff > eps_drift {
322 return Some((k, DriftLevel::Drift, diff));
323 }
324 let eps_warn = Self::hoeffding_bound(n0, n1, n_total, self.config.warning_delta);
326 if diff > eps_warn && best_split.is_none() {
327 best_split = Some((k, DriftLevel::Warning, diff));
328 }
329 }
330 best_split
331 }
332
333 fn trim_front(&mut self, count: usize) {
335 for _ in 0..count {
336 if let Some(v) = self.window.pop_front() {
337 self.total -= v;
338 }
339 }
340 }
341}
342
343impl Default for Adwin {
344 fn default() -> Self {
345 Self::new(AdwinConfig::default()).expect("default config is valid")
346 }
347}
348
349impl DriftDetector for Adwin {
350 fn update(&mut self, value: f64) -> Result<DriftLevel, RillError> {
351 ensure_finite("value", value)?;
352 let next_samples = checked_increment(self.samples, "ADWIN samples")?;
353 let mut next_total = checked_finite_add(self.total, value, "ADWIN window total")?;
354 let evicted = if self.window.len() == self.config.max_window {
355 self.window.front().copied()
356 } else {
357 None
358 };
359 if let Some(oldest) = evicted {
360 next_total = checked_finite_add(next_total, -oldest, "ADWIN window total")?;
361 }
362 self.samples = next_samples;
363 if evicted.is_some() {
365 self.window.pop_front();
366 }
367 self.window.push_back(value);
368 self.total = next_total;
369 if self.samples < self.config.min_samples || self.window.len() < 2 {
371 self.current_level = DriftLevel::None;
372 return Ok(DriftLevel::None);
373 }
374 if let Some((split, level, _diff)) = self.check_splits() {
376 if level == DriftLevel::Drift {
379 self.trim_front(split);
380 }
381 self.current_level = level;
382 } else {
383 self.current_level = DriftLevel::None;
384 }
385 Ok(self.current_level)
386 }
387
388 fn detected(&self) -> bool {
389 self.current_level == DriftLevel::Drift
390 }
391
392 fn warning(&self) -> bool {
393 self.current_level == DriftLevel::Warning
394 }
395
396 fn level(&self) -> DriftLevel {
397 self.current_level
398 }
399
400 fn samples_seen(&self) -> u64 {
401 self.samples
402 }
403
404 fn reset(&mut self) {
405 self.window.clear();
406 self.total = 0.0;
407 self.samples = 0;
408 self.current_level = DriftLevel::None;
409 }
410}
411
412#[cfg(test)]
413mod tests {
414 use super::*;
415
416 fn next_unit(seed: &mut u64) -> f64 {
418 *seed = seed
419 .wrapping_mul(6364136223846793005)
420 .wrapping_add(1442695040888963407);
421 ((*seed >> 11) as f64) / ((1u64 << 53) as f64)
422 }
423
424 #[test]
425 fn default_config_is_valid() {
426 let adwin = Adwin::default();
427 assert_eq!(adwin.samples_seen(), 0);
428 assert_eq!(adwin.level(), DriftLevel::None);
429 assert_eq!(adwin.window_size(), 0);
430 }
431
432 #[test]
433 fn detects_sudden_mean_shift() {
434 let mut adwin = Adwin::new(AdwinConfig {
435 delta: 0.05,
436 warning_delta: 0.1,
437 max_window: 500,
438 min_samples: 5,
439 })
440 .unwrap();
441 let mut seed = 42u64;
443 for _ in 0..100 {
444 let noise = 0.1 * (next_unit(&mut seed) - 0.5);
445 adwin.update(noise).unwrap();
446 }
447 assert_eq!(adwin.level(), DriftLevel::None);
448 let mut detected = false;
450 for _ in 0..200 {
451 let noise = 0.1 * (next_unit(&mut seed) - 0.5);
452 let level = adwin.update(5.0 + noise).unwrap();
453 if level == DriftLevel::Drift {
454 detected = true;
455 break;
456 }
457 }
458 assert!(detected, "ADWIN should detect the sudden mean shift");
459 }
460
461 #[test]
462 fn no_false_positive_on_stable_stream() {
463 let mut adwin = Adwin::new(AdwinConfig {
464 delta: 0.002,
465 warning_delta: 0.01,
466 max_window: 500,
467 min_samples: 10,
468 })
469 .unwrap();
470 let mut seed = 7u64;
471 for _ in 0..2000 {
472 let noise = 0.5 * (next_unit(&mut seed) - 0.5);
473 adwin.update(noise).unwrap();
474 }
475 assert!(
476 !adwin.detected(),
477 "false positive: drift reported on stable stream"
478 );
479 }
480
481 #[test]
482 fn detects_gradual_drift() {
483 let mut adwin = Adwin::new(AdwinConfig {
484 delta: 0.05,
485 warning_delta: 0.1,
486 max_window: 300,
487 min_samples: 5,
488 })
489 .unwrap();
490 let mut seed = 99u64;
492 let mut detected = false;
493 for i in 0..500 {
494 let mean = (i as f64 / 100.0).min(5.0);
495 let noise = 0.1 * (next_unit(&mut seed) - 0.5);
496 let level = adwin.update(mean + noise).unwrap();
497 if level == DriftLevel::Drift {
498 detected = true;
499 break;
500 }
501 }
502 assert!(detected, "ADWIN should detect gradual drift");
503 }
504
505 #[test]
506 fn window_trims_after_drift() {
507 let mut adwin = Adwin::new(AdwinConfig {
508 delta: 0.05,
509 warning_delta: 0.1,
510 max_window: 500,
511 min_samples: 5,
512 })
513 .unwrap();
514 for _ in 0..100 {
516 adwin.update(0.0).unwrap();
517 }
518 let size_before = adwin.window_size();
519 assert!(size_before > 0);
520 let mut trimmed = false;
522 for _ in 0..200 {
523 adwin.update(10.0).unwrap();
524 if adwin.detected() {
525 if adwin.window_size() < size_before + 200 {
529 trimmed = true;
530 break;
531 }
532 }
533 }
534 assert!(trimmed, "window should be trimmed after drift");
535 }
536
537 #[test]
538 fn max_window_enforced() {
539 let mut adwin = Adwin::new(AdwinConfig {
540 max_window: 50,
541 ..Default::default()
542 })
543 .unwrap();
544 for i in 0..200u64 {
545 adwin.update(i as f64).unwrap();
546 }
547 assert!(
550 adwin.window_size() <= 50,
551 "window should not exceed max_window, got {}",
552 adwin.window_size()
553 );
554 }
555
556 #[test]
557 fn min_samples_gates_detection() {
558 let mut adwin = Adwin::new(AdwinConfig {
559 delta: 0.5,
560 warning_delta: 0.5,
561 max_window: 100,
562 min_samples: 50,
563 })
564 .unwrap();
565 for _ in 0..48 {
567 adwin.update(0.0).unwrap();
568 }
569 adwin.update(100.0).unwrap();
571 assert_eq!(adwin.level(), DriftLevel::None);
572 let mut detected = false;
576 for _ in 0..50 {
577 let level = adwin.update(100.0).unwrap();
578 if level.is_change() {
579 detected = true;
580 }
581 }
582 assert!(detected, "should have detected drift after min_samples");
583 }
584
585 #[test]
586 fn reset_clears_state() {
587 let mut adwin = Adwin::default();
588 for _ in 0..50 {
589 adwin.update(1.0).unwrap();
590 }
591 assert!(adwin.window_size() > 0);
592 adwin.reset();
593 assert_eq!(adwin.window_size(), 0);
594 assert_eq!(adwin.samples_seen(), 0);
595 assert_eq!(adwin.level(), DriftLevel::None);
596 assert_eq!(adwin.window_mean(), 0.0);
597 }
598
599 #[test]
600 fn rejects_non_finite_input() {
601 let mut adwin = Adwin::default();
602 assert!(adwin.update(f64::NAN).is_err());
603 assert!(adwin.update(f64::INFINITY).is_err());
604 assert!(adwin.update(f64::NEG_INFINITY).is_err());
605 assert_eq!(adwin.samples_seen(), 0);
606 assert_eq!(adwin.window_size(), 0);
607 }
608
609 #[test]
610 fn rejects_invalid_config() {
611 assert!(
613 Adwin::new(AdwinConfig {
614 delta: 0.0,
615 ..Default::default()
616 })
617 .is_err()
618 );
619 assert!(
620 Adwin::new(AdwinConfig {
621 delta: 1.0,
622 ..Default::default()
623 })
624 .is_err()
625 );
626 assert!(
628 Adwin::new(AdwinConfig {
629 delta: 0.05,
630 warning_delta: 0.01,
631 ..Default::default()
632 })
633 .is_err()
634 );
635 assert!(
637 Adwin::new(AdwinConfig {
638 max_window: 0,
639 ..Default::default()
640 })
641 .is_err()
642 );
643 assert!(
645 Adwin::new(AdwinConfig {
646 min_samples: 0,
647 ..Default::default()
648 })
649 .is_err()
650 );
651 }
652
653 #[test]
654 fn window_mean_correct() {
655 let mut adwin = Adwin::new(AdwinConfig {
656 max_window: 100,
657 min_samples: 11, ..Default::default()
659 })
660 .unwrap();
661 for i in 1..=10 {
662 adwin.update(i as f64).unwrap();
663 }
664 assert!((adwin.window_mean() - 5.5).abs() < 1e-9);
666 }
667
668 #[test]
669 fn hoeffding_bound_decreases_with_more_data() {
670 let b1 = Adwin::hoeffding_bound(5.0, 5.0, 10, 0.01);
672 let b2 = Adwin::hoeffding_bound(50.0, 50.0, 100, 0.01);
673 assert!(
674 b2 < b1,
675 "bound should decrease with more data: {} vs {}",
676 b2,
677 b1
678 );
679 }
680
681 #[test]
682 fn portable_state_restore_preserves_future_results() {
683 let config = AdwinConfig {
684 delta: 0.05,
685 warning_delta: 0.1,
686 max_window: 80,
687 min_samples: 5,
688 };
689 let mut original = Adwin::new(config.clone()).unwrap();
690 for i in 0..60 {
691 original.update((i % 7) as f64 / 10.0).unwrap();
692 }
693 let state = original.export_state_v1();
694 state.validate_state().unwrap();
695 let mut restored = Adwin::restore_state_v1(config, state).unwrap();
696 for i in 0..100 {
697 let value = if i < 20 { 0.25 } else { 3.0 + i as f64 / 100.0 };
698 assert_eq!(
699 original.update(value).unwrap(),
700 restored.update(value).unwrap()
701 );
702 assert_eq!(original.export_state_v1(), restored.export_state_v1());
703 }
704 }
705
706 #[test]
707 fn portable_state_rejects_mismatch_and_corruption() {
708 let detector = Adwin::default();
709 let mut wrong_config = AdwinConfig::default();
710 wrong_config.max_window += 1;
711 assert!(Adwin::restore_state_v1(wrong_config, detector.export_state_v1()).is_err());
712
713 let mut corrupt = detector.export_state_v1();
714 corrupt.window.push(1.0);
715 corrupt.total = 2.0;
716 corrupt.samples = 1;
717 assert!(corrupt.validate_state().is_err());
718 let mut corrupt = detector.export_state_v1();
719 corrupt.version = 99;
720 assert!(corrupt.validate_state().is_err());
721 }
722
723 #[cfg(feature = "serde")]
724 #[test]
725 fn serde_roundtrip() {
726 let mut adwin = Adwin::new(AdwinConfig {
727 delta: 0.01,
728 warning_delta: 0.05,
729 max_window: 200,
730 min_samples: 5,
731 })
732 .unwrap();
733 for i in 0..50 {
734 adwin.update(i as f64 * 0.1).unwrap();
735 }
736 let json = serde_json::to_string(&adwin).unwrap();
737 let restored: Adwin = serde_json::from_str(&json).unwrap();
738 assert_eq!(restored.samples_seen(), 50);
739 assert_eq!(restored.window_size(), adwin.window_size());
740 assert!((restored.window_mean() - adwin.window_mean()).abs() < 1e-12);
741 assert_eq!(restored.level(), adwin.level());
742 }
743}