use crate::error::FinError;
#[derive(Debug, Clone, Copy)]
pub struct AdfTest {
pub statistic: f64,
pub critical_values: [f64; 3],
}
impl AdfTest {
#[must_use]
pub fn is_stationary_at_1pct(&self) -> bool {
self.statistic < self.critical_values[0]
}
#[must_use]
pub fn is_stationary_at_5pct(&self) -> bool {
self.statistic < self.critical_values[1]
}
#[must_use]
pub fn is_stationary_at_10pct(&self) -> bool {
self.statistic < self.critical_values[2]
}
}
#[derive(Debug, Clone)]
pub struct CointegrationTest {
pub hedge_ratio: f64,
pub intercept: f64,
pub adf: AdfTest,
pub residuals: Vec<f64>,
}
impl CointegrationTest {
pub fn run(y: &[f64], x: &[f64]) -> Result<Self, FinError> {
if y.len() != x.len() {
return Err(FinError::InvalidInput(
"y and x must have the same length".into(),
));
}
if y.len() < 5 {
return Err(FinError::InvalidInput("need at least 5 observations".into()));
}
let n = y.len() as f64;
let sum_x: f64 = x.iter().sum();
let sum_y: f64 = y.iter().sum();
let sum_xx: f64 = x.iter().map(|v| v * v).sum();
let sum_xy: f64 = x.iter().zip(y.iter()).map(|(xi, yi)| xi * yi).sum();
let denom = n * sum_xx - sum_x * sum_x;
let (hedge_ratio, intercept) = if denom.abs() < f64::EPSILON {
(1.0, 0.0)
} else {
let beta = (n * sum_xy - sum_x * sum_y) / denom;
let alpha = (sum_y - beta * sum_x) / n;
(beta, alpha)
};
let residuals: Vec<f64> = x
.iter()
.zip(y.iter())
.map(|(xi, yi)| yi - hedge_ratio * xi - intercept)
.collect();
let adf = Self::adf_test(&residuals);
Ok(Self {
hedge_ratio,
intercept,
adf,
residuals,
})
}
fn adf_test(series: &[f64]) -> AdfTest {
let critical_values = [-3.43, -2.86, -2.57_f64];
let n = series.len();
if n < 3 {
return AdfTest {
statistic: 0.0,
critical_values,
};
}
let dy: Vec<f64> = (1..n).map(|i| series[i] - series[i - 1]).collect();
let lag: Vec<f64> = (0..n - 1).map(|i| series[i]).collect();
let m = dy.len() as f64;
let sum_lag: f64 = lag.iter().sum();
let sum_dy: f64 = dy.iter().sum();
let sum_ll: f64 = lag.iter().map(|v| v * v).sum();
let sum_ldy: f64 = lag.iter().zip(dy.iter()).map(|(l, d)| l * d).sum();
let denom = m * sum_ll - sum_lag * sum_lag;
if denom.abs() < f64::EPSILON {
return AdfTest {
statistic: 0.0,
critical_values,
};
}
let gamma = (m * sum_ldy - sum_lag * sum_dy) / denom;
let alpha_fd = (sum_dy - gamma * sum_lag) / m;
let residuals: Vec<f64> = lag
.iter()
.zip(dy.iter())
.map(|(l, d)| d - gamma * l - alpha_fd)
.collect();
let sse: f64 = residuals.iter().map(|r| r * r).sum();
let s2 = sse / (m - 2.0).max(1.0);
let se_gamma = if s2 <= 0.0 || denom.abs() < f64::EPSILON {
1.0
} else {
(s2 * m / denom).sqrt()
};
let statistic = if se_gamma.abs() < f64::EPSILON {
0.0
} else {
gamma / se_gamma
};
AdfTest {
statistic,
critical_values,
}
}
#[must_use]
pub fn is_cointegrated(&self) -> bool {
self.adf.is_stationary_at_5pct()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PairSignal {
EnterLongShort,
EnterShortLong,
Exit,
Hold,
}
#[derive(Debug, Clone)]
pub struct PairsStrategy {
pub hedge_ratio: f64,
pub spread_mean: f64,
pub spread_std: f64,
pub z_threshold: f64,
}
impl PairsStrategy {
pub fn new(
hedge_ratio: f64,
spread_mean: f64,
spread_std: f64,
z_threshold: f64,
) -> Result<Self, FinError> {
if spread_std <= 0.0 {
return Err(FinError::InvalidInput(
"spread_std must be positive".into(),
));
}
if z_threshold <= 0.0 {
return Err(FinError::InvalidInput(
"z_threshold must be positive".into(),
));
}
Ok(Self {
hedge_ratio,
spread_mean,
spread_std,
z_threshold,
})
}
#[must_use]
pub fn generate_signal(&self, spread: f64) -> PairSignal {
let z = (spread - self.spread_mean) / self.spread_std;
if z > self.z_threshold {
PairSignal::EnterShortLong
} else if z < -self.z_threshold {
PairSignal::EnterLongShort
} else if z.abs() < self.z_threshold * 0.5 {
PairSignal::Exit
} else {
PairSignal::Hold
}
}
#[must_use]
pub fn z_score(&self, spread: f64) -> f64 {
(spread - self.spread_mean) / self.spread_std
}
}
#[derive(Debug, Clone)]
pub struct SpreadTracker {
count: u64,
mean: f64,
m2: f64,
}
impl SpreadTracker {
#[must_use]
pub fn new() -> Self {
Self {
count: 0,
mean: 0.0,
m2: 0.0,
}
}
pub fn update(&mut self, value: f64) {
self.count += 1;
let delta = value - self.mean;
self.mean += delta / self.count as f64;
let delta2 = value - self.mean;
self.m2 += delta * delta2;
}
#[must_use]
pub fn mean(&self) -> Option<f64> {
if self.count == 0 {
None
} else {
Some(self.mean)
}
}
#[must_use]
pub fn variance(&self) -> Option<f64> {
if self.count < 2 {
None
} else {
Some(self.m2 / (self.count - 1) as f64)
}
}
#[must_use]
pub fn std_dev(&self) -> Option<f64> {
self.variance().map(f64::sqrt)
}
#[must_use]
pub fn count(&self) -> u64 {
self.count
}
}
impl Default for SpreadTracker {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn adf_stationary_at_5pct() {
let adf = AdfTest {
statistic: -3.5,
critical_values: [-3.43, -2.86, -2.57],
};
assert!(!adf.is_stationary_at_1pct()); assert!(adf.is_stationary_at_1pct());
assert!(adf.is_stationary_at_5pct());
assert!(adf.is_stationary_at_10pct());
}
#[test]
fn adf_non_stationary() {
let adf = AdfTest {
statistic: -1.0,
critical_values: [-3.43, -2.86, -2.57],
};
assert!(!adf.is_stationary_at_10pct());
}
#[test]
fn cointegration_perfect_linear() {
let x: Vec<f64> = (1..=20).map(|i| i as f64).collect();
let y: Vec<f64> = x.iter().map(|xi| 2.0 * xi + 0.001).collect();
let result = CointegrationTest::run(&y, &x).unwrap();
assert!((result.hedge_ratio - 2.0).abs() < 0.01);
}
#[test]
fn cointegration_too_short() {
let y = vec![1.0, 2.0];
let x = vec![1.0, 2.0];
assert!(CointegrationTest::run(&y, &x).is_err());
}
#[test]
fn cointegration_length_mismatch() {
let y = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let x = vec![1.0, 2.0, 3.0];
assert!(CointegrationTest::run(&y, &x).is_err());
}
#[test]
fn pairs_strategy_signals() {
let s = PairsStrategy::new(1.0, 0.0, 1.0, 2.0).unwrap();
assert_eq!(s.generate_signal(3.0), PairSignal::EnterShortLong);
assert_eq!(s.generate_signal(-3.0), PairSignal::EnterLongShort);
assert_eq!(s.generate_signal(0.1), PairSignal::Exit);
assert_eq!(s.generate_signal(1.5), PairSignal::Hold);
}
#[test]
fn pairs_strategy_invalid() {
assert!(PairsStrategy::new(1.0, 0.0, 0.0, 2.0).is_err());
assert!(PairsStrategy::new(1.0, 0.0, 1.0, 0.0).is_err());
}
#[test]
fn pairs_strategy_z_score() {
let s = PairsStrategy::new(1.0, 0.0, 2.0, 2.0).unwrap();
assert!((s.z_score(4.0) - 2.0).abs() < 1e-10);
}
#[test]
fn spread_tracker_mean_variance() {
let mut t = SpreadTracker::new();
assert!(t.mean().is_none());
assert!(t.variance().is_none());
for v in [2.0, 4.0, 4.0, 4.0, 5.0, 5.0, 7.0, 9.0] {
t.update(v);
}
let mean = t.mean().unwrap();
assert!((mean - 5.0).abs() < 1e-10);
let var = t.variance().unwrap();
assert!((var - 4.0).abs() < 0.01);
}
#[test]
fn spread_tracker_std_dev() {
let mut t = SpreadTracker::new();
for v in [1.0, 2.0, 3.0] {
t.update(v);
}
let std = t.std_dev().unwrap();
assert!((std - 1.0).abs() < 1e-10);
}
#[test]
fn spread_tracker_single_obs_no_variance() {
let mut t = SpreadTracker::new();
t.update(5.0);
assert!(t.mean().is_some());
assert!(t.variance().is_none());
}
#[test]
fn spread_tracker_welford_stability() {
let mut t = SpreadTracker::new();
let base = 1_000_000.0_f64;
for i in 0..100 {
t.update(base + i as f64);
}
let mean = t.mean().unwrap();
assert!((mean - (base + 49.5)).abs() < 1e-6);
}
}