#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum StiClass {
Sit,
Sti,
Ist,
Tsi,
Its,
Tis,
}
impl std::fmt::Display for StiClass {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Sit => write!(f, "SIT"),
Self::Sti => write!(f, "STI"),
Self::Ist => write!(f, "IST"),
Self::Tsi => write!(f, "TSI"),
Self::Its => write!(f, "ITS"),
Self::Tis => write!(f, "TIS"),
}
}
}
#[derive(Debug, Clone, Copy)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct StiClassResult {
pub class: StiClass,
pub sf_seas: f64,
pub sf_trnd: f64,
pub sif: f64,
pub tif: f64,
}
pub fn sti_class(series: &[f64], months_per_year: usize, years: usize) -> Option<StiClassResult> {
if months_per_year < 2 || years < 2 {
return None;
}
let needed = months_per_year.checked_mul(years)?;
if series.len() < needed {
return None;
}
let tail = &series[series.len() - needed..];
if tail.iter().any(|v| !v.is_finite()) {
return None;
}
let m = months_per_year;
let y = years;
let n = needed as f64;
let y_bar: f64 = tail.iter().sum::<f64>() / n;
let mut row_mean = vec![0.0_f64; y];
let mut col_mean = vec![0.0_f64; m];
for (idx, &v) in tail.iter().enumerate() {
let row = idx / m;
let col = idx % m;
row_mean[row] += v;
col_mean[col] += v;
}
for v in row_mean.iter_mut() {
*v /= m as f64;
}
for v in col_mean.iter_mut() {
*v /= y as f64;
}
let mut ss_total = 0.0_f64;
for &v in tail {
let d = v - y_bar;
ss_total += d * d;
}
let mut ss_row = 0.0_f64;
for &r in &row_mean {
let d = r - y_bar;
ss_row += d * d;
}
ss_row *= m as f64;
let mut ss_col = 0.0_f64;
for &c in &col_mean {
let d = c - y_bar;
ss_col += d * d;
}
ss_col *= y as f64;
let ss_err = ss_total - ss_row - ss_col;
if ss_total <= 0.0 || !ss_total.is_finite() {
return None;
}
if ss_err <= 0.0 || !ss_err.is_finite() {
return None;
}
let df_row = (y - 1) as f64;
let df_col = (m - 1) as f64;
let df_err = df_row * df_col;
let ms_row = ss_row / df_row;
let ms_col = ss_col / df_col;
let ms_err = ss_err / df_err;
let sf_seas = ss_col / (ss_col + ss_err);
let sf_trnd = ss_row / (ss_row + ss_err);
let sif = ms_col / ms_err;
let tif = ms_row / ms_err;
let class = rank_to_class(ms_col, ms_row, ms_err);
Some(StiClassResult {
class,
sf_seas,
sf_trnd,
sif,
tif,
})
}
fn rank_to_class(s: f64, t: f64, i: f64) -> StiClass {
if s >= t && s >= i {
if t >= i {
StiClass::Sti
} else {
StiClass::Sit
}
} else if t >= s && t >= i {
if s >= i {
StiClass::Tsi
} else {
StiClass::Tis
}
} else {
if s >= t {
StiClass::Ist
} else {
StiClass::Its
}
}
}
#[cfg(test)]
mod tests {
use super::*;
const M: usize = 12;
const Y: usize = 5;
#[test]
fn pure_seasonal_classifies_as_seasonal_dominant() {
let n = M * Y;
let series: Vec<f64> = (0..n)
.map(|i| {
(2.0 * std::f64::consts::PI * (i % M) as f64 / M as f64).sin()
+ 0.001 * ((i * 7) % 11) as f64 })
.collect();
let result = sti_class(&series, M, Y).expect("classify pure seasonal");
assert!(
matches!(result.class, StiClass::Sit | StiClass::Sti),
"expected seasonal-dominant class, got {}",
result.class,
);
assert!(result.sf_seas > 0.9, "SF_seas too low: {}", result.sf_seas);
}
#[test]
fn pure_trend_classifies_as_trend_dominant() {
let n = M * Y;
let series: Vec<f64> = (0..n)
.map(|i| i as f64 + 0.001 * ((i * 13) % 7) as f64)
.collect();
let result = sti_class(&series, M, Y).expect("classify pure trend");
assert!(
matches!(result.class, StiClass::Tsi | StiClass::Tis),
"expected trend-dominant class, got {}",
result.class,
);
assert!(result.sf_trnd > 0.9, "SF_trnd too low: {}", result.sf_trnd);
}
#[test]
fn noise_dominated_signal_has_small_strength_factors() {
let years = 20;
let n = M * years;
let mut state: u64 = 0xCAFEBABE_DEADBEEF;
let series: Vec<f64> = (0..n)
.map(|i| {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
let noise = ((state >> 33) as f64) / ((1u64 << 31) as f64) - 1.0;
let trend = 0.001 * i as f64;
let seasonal =
0.01 * (2.0 * std::f64::consts::PI * (i % M) as f64 / M as f64).sin();
100.0 * noise + trend + seasonal
})
.collect();
let result = sti_class(&series, M, years).expect("classify noise-dominated");
assert!(
result.sf_seas < 0.5,
"noise-dominated SF_seas = {} (expected < 0.5)",
result.sf_seas
);
assert!(
result.sf_trnd < 0.5,
"noise-dominated SF_trnd = {} (expected < 0.5)",
result.sf_trnd
);
}
#[test]
fn constant_series_returns_none() {
let series = vec![7.0; M * Y];
assert!(sti_class(&series, M, Y).is_none());
}
#[test]
fn too_short_returns_none() {
let series = vec![1.0; M * Y - 1];
assert!(sti_class(&series, M, Y).is_none());
}
#[test]
fn rejects_degenerate_grid_dims() {
let series = vec![1.0_f64; 100];
assert!(sti_class(&series, 1, 5).is_none());
assert!(sti_class(&series, 12, 1).is_none());
}
#[test]
fn rejects_non_finite_observations() {
let mut series: Vec<f64> = (0..M * Y).map(|i| i as f64).collect();
series[5] = f64::NAN;
assert!(sti_class(&series, M, Y).is_none());
}
#[test]
fn uses_only_last_window_for_classification() {
let n_head = M * Y;
let n_tail = M * Y;
let mut series: Vec<f64> = (0..n_head)
.map(|i| ((i * 17) % 23) as f64 - 11.0) .collect();
series.extend((0..n_tail).map(|i| {
(2.0 * std::f64::consts::PI * (i % M) as f64 / M as f64).sin()
+ 0.001 * ((i * 7) % 11) as f64
}));
let result = sti_class(&series, M, Y).expect("classify tail");
assert!(
matches!(result.class, StiClass::Sit | StiClass::Sti),
"tail should drive class, got {}",
result.class,
);
}
#[test]
fn strength_factors_in_unit_interval() {
let n = M * Y;
let series: Vec<f64> = (0..n)
.map(|i| {
(2.0 * std::f64::consts::PI * (i % M) as f64 / M as f64).sin()
+ 0.1 * i as f64
+ 0.05 * ((i * 7) % 11) as f64
})
.collect();
let result = sti_class(&series, M, Y).unwrap();
assert!((0.0..=1.0).contains(&result.sf_seas));
assert!((0.0..=1.0).contains(&result.sf_trnd));
}
}