use crate::error::{ForecastError, Result};
use crate::postprocess::PredictionIntervals;
#[derive(Debug, Clone)]
pub struct AciPredictor {
alpha_target: f64,
gamma: f64,
alpha_t: f64,
max_window: usize,
residuals: Vec<f64>,
n_observed: usize,
n_misses: usize,
}
impl AciPredictor {
pub fn new(target_coverage: f64, gamma: f64) -> Self {
assert!(
target_coverage > 0.0 && target_coverage < 1.0,
"target_coverage must be in (0, 1)"
);
assert!(gamma > 0.0, "gamma must be > 0");
let alpha = 1.0 - target_coverage;
Self {
alpha_target: alpha,
gamma,
alpha_t: alpha,
max_window: 500,
residuals: Vec::new(),
n_observed: 0,
n_misses: 0,
}
}
pub fn with_max_window(mut self, max_window: usize) -> Self {
assert!(max_window > 0, "max_window must be > 0");
self.max_window = max_window;
self
}
pub fn target_coverage(&self) -> f64 {
1.0 - self.alpha_target
}
pub fn alpha_t(&self) -> f64 {
self.alpha_t
}
pub fn gamma(&self) -> f64 {
self.gamma
}
pub fn n_observed(&self) -> usize {
self.n_observed
}
pub fn empirical_miscoverage(&self) -> f64 {
if self.n_observed == 0 {
0.0
} else {
self.n_misses as f64 / self.n_observed as f64
}
}
pub fn fit(&mut self, abs_residuals: &[f64]) -> Result<()> {
if abs_residuals.is_empty() {
return Err(ForecastError::EmptyData);
}
for &r in abs_residuals {
if !r.is_finite() {
return Err(ForecastError::InvalidParameter(
"residuals must be finite".to_string(),
));
}
}
let start = abs_residuals.len().saturating_sub(self.max_window);
self.residuals = abs_residuals[start..].to_vec();
Ok(())
}
pub fn current_radius(&self) -> f64 {
if self.residuals.is_empty() {
return 0.0;
}
let mut sorted = self.residuals.clone();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap());
let n = sorted.len();
let coverage = 1.0 - self.alpha_t;
let level = (coverage * (n as f64 + 1.0) / n as f64).min(1.0);
let idx = ((n as f64) * level).ceil() as usize;
let idx = idx.saturating_sub(1).min(n - 1);
sorted[idx]
}
pub fn predict_interval(&self, forecast: f64) -> (f64, f64) {
let r = self.current_radius();
(forecast - r, forecast + r)
}
pub fn predict_intervals(&self, forecasts: &[f64]) -> Result<PredictionIntervals> {
let r = self.current_radius();
let lower: Vec<f64> = forecasts.iter().map(|&f| f - r).collect();
let upper: Vec<f64> = forecasts.iter().map(|&f| f + r).collect();
PredictionIntervals::from_bounds(lower, upper, 1.0 - self.alpha_target)
}
pub fn observe(&mut self, forecast: f64, actual: f64) {
let radius = self.current_radius();
let lower = forecast - radius;
let upper = forecast + radius;
let inside = actual >= lower && actual <= upper;
let err: f64 = if inside { 0.0 } else { 1.0 };
self.alpha_t += self.gamma * (self.alpha_target - err);
self.alpha_t = self.alpha_t.clamp(0.001, 0.999);
let abs_err = (forecast - actual).abs();
self.residuals.push(abs_err);
if self.residuals.len() > self.max_window {
let drop = self.residuals.len() - self.max_window;
self.residuals.drain(..drop);
}
self.n_observed += 1;
if !inside {
self.n_misses += 1;
}
}
pub fn reset_alpha(&mut self) {
self.alpha_t = self.alpha_target;
self.n_observed = 0;
self.n_misses = 0;
}
}
#[cfg(test)]
mod tests {
use super::*;
use rand::rngs::StdRng;
use rand::{Rng, SeedableRng};
#[test]
fn fit_seeds_residual_buffer() {
let mut aci = AciPredictor::new(0.90, 0.01);
aci.fit(&[1.0, 2.0, 3.0, 4.0, 5.0]).unwrap();
let r = aci.current_radius();
assert!((4.0..=5.0).contains(&r));
}
#[test]
fn fit_empty_errors() {
let mut aci = AciPredictor::new(0.90, 0.01);
let err = aci.fit(&[]).unwrap_err();
assert!(matches!(err, ForecastError::EmptyData));
}
#[test]
fn predict_interval_uses_current_radius() {
let mut aci = AciPredictor::new(0.90, 0.01);
aci.fit(&[0.5; 50]).unwrap();
let (lo, hi) = aci.predict_interval(10.0);
assert!((lo - 9.5).abs() < 1e-9);
assert!((hi - 10.5).abs() < 1e-9);
}
#[test]
fn observe_inside_increases_alpha_t() {
let mut aci = AciPredictor::new(0.90, 0.05);
aci.fit(&[1.0; 30]).unwrap();
let alpha_before = aci.alpha_t();
aci.observe(0.0, 0.0);
assert!(
aci.alpha_t() > alpha_before,
"covered observation should push α_t UP, but {} <= {}",
aci.alpha_t(),
alpha_before
);
}
#[test]
fn observe_outside_decreases_alpha_t() {
let mut aci = AciPredictor::new(0.90, 0.05);
aci.fit(&[0.1; 30]).unwrap();
let alpha_before = aci.alpha_t();
aci.observe(0.0, 1000.0);
assert!(
aci.alpha_t() < alpha_before,
"missed observation should push α_t DOWN, but {} >= {}",
aci.alpha_t(),
alpha_before
);
}
#[test]
fn long_run_coverage_approaches_target() {
let mut aci = AciPredictor::new(0.90, 0.02).with_max_window(200);
let mut rng = StdRng::seed_from_u64(42);
let init: Vec<f64> = (0..50).map(|_| rng.gen::<f64>() * 2.0).collect();
aci.fit(&init).unwrap();
for _ in 0..2000 {
let u1: f64 = rng.gen::<f64>().max(1e-300);
let u2: f64 = rng.gen();
let z = (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos();
aci.observe(0.0, z);
}
let cov = 1.0 - aci.empirical_miscoverage();
assert!(
(0.83..0.97).contains(&cov),
"long-run coverage should approach target 0.90, got {:.3}",
cov
);
}
#[test]
fn responds_to_drift() {
let mut aci = AciPredictor::new(0.90, 0.05).with_max_window(100);
aci.fit(&[0.5; 50]).unwrap();
let radius_before_drift = aci.current_radius();
for i in 0..500 {
let actual = if i % 2 == 0 { 0.4 } else { -0.4 };
aci.observe(0.0, actual);
}
for i in 0..500 {
let actual = if i % 2 == 0 { 5.0 } else { -5.0 };
aci.observe(0.0, actual);
}
let radius_after_drift = aci.current_radius();
assert!(
radius_after_drift > radius_before_drift,
"ACI should widen intervals after drift: before={:.3}, after={:.3}",
radius_before_drift,
radius_after_drift
);
assert!(
radius_after_drift > 1.5,
"radius after drift should reflect new residual scale, got {:.3}",
radius_after_drift
);
}
#[test]
fn predict_intervals_returns_correct_coverage() {
let mut aci = AciPredictor::new(0.95, 0.01);
aci.fit(&[1.0; 100]).unwrap();
let intervals = aci.predict_intervals(&[10.0, 20.0, 30.0]).unwrap();
assert_eq!(intervals.coverage(), 0.95);
assert_eq!(intervals.len(), 3);
}
#[test]
fn reset_alpha_clears_counters_only() {
let mut aci = AciPredictor::new(0.90, 0.05);
aci.fit(&[1.0; 30]).unwrap();
aci.observe(0.0, 5.0); aci.observe(0.0, 0.5); assert_eq!(aci.n_observed(), 2);
aci.reset_alpha();
assert_eq!(aci.n_observed(), 0);
assert!((aci.alpha_t() - 0.10).abs() < 1e-12);
assert!(!aci.residuals.is_empty());
}
}