use std::fmt;
use std::path::{Path, PathBuf};
use super::Report;
pub const UPDATE_BASELINES: &str = "TYPEDLM_UPDATE_BASELINES";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Verdict {
Regression,
Improvement,
NoSignificantChange,
}
#[derive(Debug, Clone)]
pub struct Comparison {
pub metric: String,
pub examples: usize,
pub baseline_score: f64,
pub current_score: f64,
pub difference: f64,
pub interval: (f64, f64),
pub worse: Vec<usize>,
pub better: Vec<usize>,
pub verdict: Verdict,
}
#[derive(Debug)]
#[non_exhaustive]
pub enum RegressionError {
Regression(Box<Comparison>),
BelowMinimum {
score: f64,
interval: (f64, f64),
minimum: f64,
},
NotComparable { reason: String },
Io { path: PathBuf, message: String },
InvalidBaseline { path: PathBuf, message: String },
}
pub fn compare(
baseline: &Report,
current: &Report,
tolerance: f64,
) -> Result<Comparison, RegressionError> {
let not_comparable = |reason: String| Err(RegressionError::NotComparable { reason });
if baseline.program != current.program {
return not_comparable(format!(
"program {} differs from the baseline's {}",
current.program, baseline.program
));
}
if baseline.metric != current.metric {
return not_comparable(format!(
"metric {} differs from the baseline's {}",
current.metric, baseline.metric
));
}
if baseline.dataset != current.dataset || baseline.scores.len() != current.scores.len() {
return not_comparable(
"the dataset changed since the baseline was recorded; update the baseline".into(),
);
}
let differences: Vec<f64> = current
.scores
.iter()
.zip(&baseline.scores)
.map(|(now, before)| now - before)
.collect();
let indices = |keep: fn(f64) -> bool| -> Vec<usize> {
differences
.iter()
.enumerate()
.filter(|(_, d)| keep(**d))
.map(|(i, _)| i)
.collect()
};
let difference = current.score - baseline.score;
let interval = paired_interval(&differences);
let verdict = if interval.1 < -tolerance {
Verdict::Regression
} else if interval.0 > 0.0 {
Verdict::Improvement
} else {
Verdict::NoSignificantChange
};
Ok(Comparison {
metric: current.metric.clone(),
examples: differences.len(),
baseline_score: baseline.score,
current_score: current.score,
difference,
interval,
worse: indices(|d| d < 0.0),
better: indices(|d| d > 0.0),
verdict,
})
}
impl Report {
pub fn require_score(&self, minimum: f64) -> Result<(), RegressionError> {
if self.interval.1 < minimum {
return Err(RegressionError::BelowMinimum {
score: self.score,
interval: self.interval,
minimum,
});
}
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct Baseline {
path: PathBuf,
tolerance: f64,
}
#[derive(Debug, Clone)]
pub enum BaselineOutcome {
Created,
Updated,
Compared(Comparison),
}
impl Baseline {
pub fn at(path: impl Into<PathBuf>) -> Self {
Self {
path: path.into(),
tolerance: 0.0,
}
}
pub fn tolerance(mut self, tolerance: f64) -> Self {
self.tolerance = tolerance;
self
}
pub fn path(&self) -> &Path {
&self.path
}
pub fn load(&self) -> Result<Report, RegressionError> {
let text = std::fs::read_to_string(&self.path).map_err(|e| self.io(e))?;
serde_json::from_str(&text).map_err(|e| RegressionError::InvalidBaseline {
path: self.path.clone(),
message: e.to_string(),
})
}
pub fn store(&self, report: &Report) -> Result<(), RegressionError> {
if let Some(parent) = self.path.parent().filter(|p| !p.as_os_str().is_empty()) {
std::fs::create_dir_all(parent).map_err(|e| self.io(e))?;
}
let mut json = serde_json::to_string_pretty(report).expect("reports serialize to JSON");
json.push('\n');
std::fs::write(&self.path, json).map_err(|e| self.io(e))
}
pub fn check(&self, report: &Report) -> Result<BaselineOutcome, RegressionError> {
let update = std::env::var(UPDATE_BASELINES).is_ok_and(|v| !v.is_empty() && v != "0");
if update {
self.store(report)?;
return Ok(BaselineOutcome::Updated);
}
if !self.path.exists() {
self.store(report)?;
return Ok(BaselineOutcome::Created);
}
let comparison = compare(&self.load()?, report, self.tolerance)?;
match comparison.verdict {
Verdict::Regression => Err(RegressionError::Regression(Box::new(comparison))),
_ => Ok(BaselineOutcome::Compared(comparison)),
}
}
fn io(&self, error: std::io::Error) -> RegressionError {
RegressionError::Io {
path: self.path.clone(),
message: error.to_string(),
}
}
}
fn paired_interval(differences: &[f64]) -> (f64, f64) {
let n = differences.len();
if n == 0 {
return (0.0, 0.0);
}
let mean = differences.iter().sum::<f64>() / n as f64;
if n == 1 {
return (mean, mean);
}
let variance = differences.iter().map(|d| (d - mean).powi(2)).sum::<f64>() / (n - 1) as f64;
let half = t_975(n - 1) * (variance / n as f64).sqrt();
(mean - half, mean + half)
}
fn t_975(degrees_of_freedom: usize) -> f64 {
const TABLE: [f64; 30] = [
12.706, 4.303, 3.182, 2.776, 2.571, 2.447, 2.365, 2.306, 2.262, 2.228, 2.201, 2.179, 2.160,
2.145, 2.131, 2.120, 2.110, 2.101, 2.093, 2.086, 2.080, 2.074, 2.069, 2.064, 2.060, 2.056,
2.052, 2.048, 2.045, 2.042,
];
match degrees_of_freedom {
0 => f64::INFINITY,
df @ 1..=30 => TABLE[df - 1],
df => {
const Z: f64 = 1.959_964;
Z + (Z.powi(3) + Z) / (4.0 * df as f64)
}
}
}
impl fmt::Display for Comparison {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let pct = |x: f64| format!("{:.1} %", x * 100.0);
let points = |x: f64| format!("{:+.1}", x * 100.0);
let verdict = match self.verdict {
Verdict::Regression => "regression",
Verdict::Improvement => "improvement",
Verdict::NoSignificantChange => "no significant change",
};
write!(
f,
"{} {} → {} ({} points, 95 % CI {} – {}) on {} examples, {} worse, {} better: {verdict}",
self.metric,
pct(self.baseline_score),
pct(self.current_score),
points(self.difference),
points(self.interval.0),
points(self.interval.1),
self.examples,
self.worse.len(),
self.better.len(),
)
}
}
impl fmt::Display for RegressionError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Regression(comparison) => {
write!(f, "{comparison}")?;
if !comparison.worse.is_empty() {
let list: Vec<String> = comparison.worse.iter().map(usize::to_string).collect();
write!(f, "; worse examples: {}", list.join(", "))?;
}
Ok(())
}
Self::BelowMinimum {
score,
interval,
minimum,
} => write!(
f,
"score {:.1} % (95 % CI {:.1} % – {:.1} %) is below the required {:.1} %",
score * 100.0,
interval.0 * 100.0,
interval.1 * 100.0,
minimum * 100.0
),
Self::NotComparable { reason } => {
write!(f, "cannot compare with the baseline: {reason}")
}
Self::Io { path, message } => write!(f, "{}: {message}", path.display()),
Self::InvalidBaseline { path, message } => {
write!(f, "{} is not a valid report: {message}", path.display())
}
}
}
}
impl std::error::Error for RegressionError {}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn t_quantiles() {
assert_eq!(t_975(7), 2.365);
assert!((t_975(40) - 2.021).abs() < 0.002);
assert!((t_975(1000) - 1.962).abs() < 0.002);
}
#[test]
fn paired_interval_of_constant_differences_is_a_point() {
assert_eq!(paired_interval(&[-1.0, -1.0, -1.0]), (-1.0, -1.0));
assert_eq!(paired_interval(&[0.0; 5]), (0.0, 0.0));
}
#[test]
fn paired_interval_matches_reference() {
let d = [-1.0, 0.0, 0.0, 0.0, 0.0, -1.0, 0.0, 0.0];
let (low, high) = paired_interval(&d);
assert!((low - (-0.6371)).abs() < 1e-3, "{low}");
assert!((high - 0.1371).abs() < 1e-3, "{high}");
}
}