use serde::{Deserialize, Serialize};
use std::collections::VecDeque;
use crate::metrics::MetricSample;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Forecast {
pub predicted_value: f32,
pub slope: f32,
pub ewma: f32,
pub confidence: f32,
pub horizon: usize,
}
pub struct ForecastEngine {
ewma_alpha: f32,
}
impl ForecastEngine {
#[must_use]
pub const fn new() -> Self {
Self { ewma_alpha: 0.3 }
}
#[must_use]
pub const fn with_alpha(alpha: f32) -> Self {
Self {
ewma_alpha: alpha.clamp(0.0, 1.0),
}
}
#[must_use]
pub fn forecast(&self, history: &VecDeque<MetricSample>, horizon: usize) -> Forecast {
let n = history.len();
let raw_values: Vec<f32> = history.iter().map(|s| s.value).collect();
let values = clamp_outliers(&raw_values);
let (slope, intercept) = linear_regression(&values);
let predicted_linear = slope.mul_add((n + horizon - 1) as f32, intercept);
let ewma = compute_ewma(&values, self.ewma_alpha);
let r_squared = compute_r_squared(&values, slope, intercept);
let blended = r_squared.mul_add(predicted_linear, (1.0 - r_squared) * ewma);
Forecast {
predicted_value: blended,
slope,
ewma,
confidence: r_squared,
horizon,
}
}
}
impl Default for ForecastEngine {
fn default() -> Self {
Self::new()
}
}
fn linear_regression(values: &[f32]) -> (f32, f32) {
let n = values.len() as f32;
if n < 2.0 {
return (0.0, values.first().copied().unwrap_or(0.0));
}
let sum_x: f32 = (0..values.len()).map(|i| i as f32).sum();
let sum_y: f32 = values.iter().copied().sum();
let sum_xy: f32 = values.iter().enumerate().map(|(i, &v)| i as f32 * v).sum();
let sum_x_sq: f32 = (0..values.len()).map(|i| (i as f32).powi(2)).sum();
let denominator = n.mul_add(sum_x_sq, -(sum_x * sum_x));
if denominator.abs() < f32::EPSILON {
return (0.0, sum_y / n);
}
let slope = n.mul_add(sum_xy, -(sum_x * sum_y)) / denominator;
let intercept = slope.mul_add(-sum_x, sum_y) / n;
(slope, intercept)
}
fn compute_r_squared(values: &[f32], slope: f32, intercept: f32) -> f32 {
let n = values.len();
if n < 3 {
return 0.5;
}
let mean_y: f32 = values.iter().copied().sum::<f32>() / n as f32;
let mut ss_res = 0.0_f32; let mut ss_tot = 0.0_f32;
for (i, &y) in values.iter().enumerate() {
let predicted = slope.mul_add(i as f32, intercept);
ss_res += (y - predicted).powi(2);
ss_tot += (y - mean_y).powi(2);
}
if ss_tot < f32::EPSILON {
return 0.9;
}
let r_squared = 1.0 - ss_res / ss_tot;
r_squared.clamp(0.0, 1.0)
}
fn compute_ewma(values: &[f32], alpha: f32) -> f32 {
if values.is_empty() {
return 0.0;
}
let alpha = alpha.clamp(0.0, 1.0);
let mut ewma = values[0];
for &v in &values[1..] {
ewma = alpha.mul_add(v, (1.0 - alpha) * ewma);
}
ewma
}
fn clamp_outliers(values: &[f32]) -> Vec<f32> {
if values.len() < 4 {
return values.to_vec();
}
let mut sorted: Vec<f32> = values.to_vec();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let mid = sorted.len() / 2;
let median = if sorted.len() % 2 == 0 {
f32::midpoint(sorted[mid - 1], sorted[mid])
} else {
sorted[mid]
};
let abs_devs: Vec<f32> = values.iter().map(|&v| (v - median).abs()).collect();
let mut sorted_devs = abs_devs.clone();
sorted_devs.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let mad = if sorted_devs.len() % 2 == 0 {
f32::midpoint(sorted_devs[mid - 1], sorted_devs[mid])
} else {
sorted_devs[mid]
};
let scaled_mad = 1.4826 * mad;
if scaled_mad < f32::EPSILON {
let has_non_zero_dev = abs_devs.iter().any(|&d| d > f32::EPSILON);
if !has_non_zero_dev {
return values.to_vec();
}
let fallback_scale = median.abs().max(0.01);
let threshold = 3.0 * fallback_scale;
let lower = median - threshold;
let upper = median + threshold;
return values.iter().map(|&v| v.clamp(lower, upper)).collect();
}
let threshold = 3.5 * scaled_mad;
let lower = median - threshold;
let upper = median + threshold;
values.iter().map(|&v| v.clamp(lower, upper)).collect()
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::Utc;
fn make_history(values: &[f32]) -> VecDeque<MetricSample> {
values
.iter()
.map(|&v| MetricSample {
kind: crate::metrics::MetricKind::CpuLoad,
value: v,
timestamp: Utc::now(),
})
.collect()
}
#[test]
fn forecast_linear_trend() {
let history = make_history(&[0.1, 0.2, 0.3, 0.4, 0.5]);
let engine = ForecastEngine::new();
let f = engine.forecast(&history, 3);
assert!(f.predicted_value > 0.5);
assert!(f.slope > 0.0);
assert!((f.slope - 0.1).abs() < 0.01);
assert!(f.confidence > 0.9); }
#[test]
fn forecast_decreasing_trend() {
let history = make_history(&[0.5, 0.4, 0.3, 0.2, 0.1]);
let engine = ForecastEngine::new();
let f = engine.forecast(&history, 2);
assert!(f.predicted_value < 0.1);
assert!(f.slope < 0.0);
}
#[test]
fn forecast_noisy_data_lower_confidence() {
let history = make_history(&[0.3, 0.7, 0.2, 0.6, 0.3]);
let engine = ForecastEngine::new();
let f = engine.forecast(&history, 3);
assert!(f.confidence < 0.8); }
#[test]
fn forecast_constant_data() {
let history = make_history(&[0.5, 0.5, 0.5, 0.5]);
let engine = ForecastEngine::new();
let f = engine.forecast(&history, 5);
assert!((f.predicted_value - 0.5).abs() < 0.1);
assert!((f.slope - 0.0).abs() < 0.01);
assert!(f.confidence > 0.8);
}
#[test]
fn forecast_two_points() {
let history = make_history(&[0.3, 0.5]);
let engine = ForecastEngine::new();
let f = engine.forecast(&history, 2);
assert!(f.predicted_value > 0.5);
assert!(f.slope > 0.0);
}
#[test]
fn forecast_ewma_blends_with_linear() {
let history = make_history(&[0.1, 0.9, 0.1, 0.9, 0.1]);
let engine = ForecastEngine::new();
let f = engine.forecast(&history, 3);
assert!(f.predicted_value > 0.0 && f.predicted_value < 1.0);
}
#[test]
fn forecast_horizon_affects_prediction() {
let history = make_history(&[0.1, 0.2, 0.3, 0.4, 0.5]);
let engine = ForecastEngine::new();
let f1 = engine.forecast(&history, 1);
let f10 = engine.forecast(&history, 10);
assert!(f10.predicted_value > f1.predicted_value);
}
#[test]
fn forecast_engine_default() {
let engine = ForecastEngine::default();
let history = make_history(&[0.1, 0.2, 0.3]);
let f = engine.forecast(&history, 1);
assert!(f.predicted_value > 0.2);
}
#[test]
fn forecast_engine_custom_alpha() {
let engine = ForecastEngine::with_alpha(0.8);
let history = make_history(&[0.1, 0.5, 0.9]);
let f = engine.forecast(&history, 1);
assert!(f.ewma > 0.5);
}
#[test]
fn forecast_serialization() {
let f = Forecast {
predicted_value: 0.5,
slope: 0.1,
ewma: 0.4,
confidence: 0.9,
horizon: 5,
};
let json = serde_json::to_string(&f).unwrap();
let back: Forecast = serde_json::from_str(&json).unwrap();
assert!((back.predicted_value - 0.5).abs() < 0.001);
assert_eq!(back.horizon, 5);
}
#[test]
fn linear_regression_perfect_fit() {
let values = [1.0, 2.0, 3.0, 4.0, 5.0];
let (slope, intercept) = linear_regression(&values);
assert!((slope - 1.0).abs() < 0.001);
assert!((intercept - 1.0).abs() < 0.001);
}
#[test]
fn linear_regression_flat() {
let values = [0.5, 0.5, 0.5, 0.5];
let (slope, intercept) = linear_regression(&values);
assert!((slope - 0.0).abs() < 0.001);
assert!((intercept - 0.5).abs() < 0.001);
}
#[test]
fn r_squared_perfect_linear() {
let values = [1.0, 2.0, 3.0, 4.0, 5.0];
let r2 = compute_r_squared(&values, 1.0, 1.0);
assert!((r2 - 1.0).abs() < 0.001);
}
#[test]
fn r_squared_noisy() {
let values = [0.3, 0.7, 0.2, 0.6, 0.3];
let (slope, intercept) = linear_regression(&values);
let r2 = compute_r_squared(&values, slope, intercept);
assert!(r2 < 0.5);
}
#[test]
fn r_squared_constant() {
let values = [0.5, 0.5, 0.5, 0.5];
let r2 = compute_r_squared(&values, 0.0, 0.5);
assert!(r2 > 0.8);
}
#[test]
fn ewma_computation() {
let values = [0.1, 0.2, 0.3, 0.4, 0.5];
let ewma = compute_ewma(&values, 0.3);
assert!(ewma > 0.1 && ewma < 0.5);
}
#[test]
fn ewma_empty() {
let ewma = compute_ewma(&[], 0.3);
assert_eq!(ewma, 0.0);
}
#[test]
fn ewma_single_value() {
let ewma = compute_ewma(&[0.42], 0.3);
assert!((ewma - 0.42).abs() < 0.001);
}
#[test]
fn forecast_outlier_does_not_dominate() {
let history = make_history(&[0.3, 0.3, 0.3, 0.3, 100.0, 0.3, 0.3, 0.3, 0.3]);
let engine = ForecastEngine::new();
let f = engine.forecast(&history, 3);
assert!(
f.predicted_value < 10.0,
"outlier should not dominate forecast: got {}",
f.predicted_value
);
assert!(
f.slope.abs() < 5.0,
"slope should not be dominated by outlier: got {}",
f.slope
);
}
#[test]
fn forecast_extreme_outlier_clamped() {
let history = make_history(&[0.3, 0.3, 0.3, 0.3, f32::MAX, 0.3, 0.3, 0.3, 0.3]);
let engine = ForecastEngine::new();
let f = engine.forecast(&history, 3);
assert!(!f.predicted_value.is_nan(), "forecast should not be NaN");
assert!(
!f.predicted_value.is_infinite(),
"forecast should not be infinite"
);
assert!(
f.predicted_value < 10.0,
"extreme outlier should be clamped: got {}",
f.predicted_value
);
}
#[test]
fn forecast_negative_outlier_clamped() {
let history = make_history(&[0.5, 0.5, 0.5, 0.5, -1000.0, 0.5, 0.5, 0.5, 0.5]);
let engine = ForecastEngine::new();
let f = engine.forecast(&history, 3);
assert!(!f.predicted_value.is_nan(), "forecast should not be NaN");
assert!(
f.predicted_value > -10.0,
"negative outlier should be clamped: got {}",
f.predicted_value
);
}
#[test]
fn clamp_outliers_preserves_normal_values() {
let values = vec![0.1, 0.2, 0.3, 0.4, 0.5];
let clamped = clamp_outliers(&values);
for (orig, clamped_val) in values.iter().zip(clamped.iter()) {
assert!((orig - clamped_val).abs() < 0.001);
}
}
#[test]
fn clamp_outliers_clamps_extreme() {
let values = vec![0.3_f32, 0.3, 0.3, 0.3, 100.0, 0.3, 0.3, 0.3, 0.3];
let clamped = clamp_outliers(&values);
assert!(
clamped[4] < 10.0,
"outlier should be clamped: got {}",
clamped[4]
);
for i in [0, 1, 2, 3, 5, 6, 7, 8] {
assert!((clamped[i] - 0.3).abs() < 0.001);
}
}
#[test]
fn clamp_outliers_short_input_unchanged() {
let values = vec![0.1, 0.2, 0.3];
let clamped = clamp_outliers(&values);
assert_eq!(clamped, values);
}
}