use crate::CoreError;
use crate::ml::features::PricePoint;
use chrono::{DateTime, Utc};
use rust_decimal::Decimal;
use rust_decimal::prelude::ToPrimitive;
use serde::{Deserialize, Serialize};
use std::collections::VecDeque;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RealizedVolatilityForecast {
pub timestamp: DateTime<Utc>,
pub horizon: usize,
pub predicted_volatility: f64,
pub confidence_interval: (f64, f64),
}
#[derive(Debug, Clone)]
pub struct RealizedVolatilityForecaster {
data: VecDeque<PricePoint>,
max_size: usize,
garch_params: (f64, f64, f64),
}
impl RealizedVolatilityForecaster {
pub fn new(max_size: usize) -> Self {
Self {
data: VecDeque::new(),
max_size,
garch_params: (0.1, 0.85, 0.00001), }
}
pub fn add_data(&mut self, point: PricePoint) {
self.data.push_back(point);
if self.data.len() > self.max_size {
self.data.pop_front();
}
}
pub fn fit(&mut self) -> anyhow::Result<()> {
if self.data.len() < 30 {
return Err(CoreError::Validation("Insufficient data for fitting".to_string()).into());
}
let data_vec: Vec<_> = self.data.iter().collect();
let returns: Vec<f64> = data_vec
.windows(2)
.map(|w| {
let prev = w[0].close.to_f64().unwrap_or(0.0);
let curr = w[1].close.to_f64().unwrap_or(0.0);
if prev > 0.0 {
((curr - prev) / prev).ln()
} else {
0.0
}
})
.collect();
let mean_return = returns.iter().sum::<f64>() / returns.len() as f64;
let squared_returns: Vec<f64> = returns.iter().map(|r| (r - mean_return).powi(2)).collect();
let unconditional_variance =
squared_returns.iter().sum::<f64>() / squared_returns.len() as f64;
let alpha = 0.1;
let beta = 0.85;
let omega = unconditional_variance * (1.0 - alpha - beta);
self.garch_params = (alpha, beta, omega.max(0.00001));
Ok(())
}
pub fn forecast(&self, horizon: usize) -> anyhow::Result<RealizedVolatilityForecast> {
if self.data.len() < 2 {
return Err(CoreError::Validation("Insufficient data".to_string()).into());
}
let returns: Vec<f64> = self
.data
.iter()
.rev()
.take(20.min(self.data.len()))
.collect::<Vec<_>>()
.windows(2)
.map(|w| {
let curr = w[0].close.to_f64().unwrap_or(0.0);
let prev = w[1].close.to_f64().unwrap_or(0.0);
if prev > 0.0 {
((curr - prev) / prev).ln()
} else {
0.0
}
})
.collect();
let last_squared_return = returns.last().copied().unwrap_or(0.0).powi(2);
let current_variance =
returns.iter().map(|r| r.powi(2)).sum::<f64>() / returns.len() as f64;
let (alpha, beta, omega) = self.garch_params;
let mut variance = current_variance;
let mut forecast_variance = variance;
for _ in 0..horizon {
forecast_variance = omega + alpha * last_squared_return + beta * variance;
variance = forecast_variance;
}
let predicted_volatility = (forecast_variance * 365.0).sqrt();
let std_error = predicted_volatility * 0.2; let lower = (predicted_volatility - 1.96 * std_error).max(0.0);
let upper = predicted_volatility + 1.96 * std_error;
Ok(RealizedVolatilityForecast {
timestamp: Utc::now(),
horizon,
predicted_volatility,
confidence_interval: (lower, upper),
})
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImpliedVolatility {
pub strike: Decimal,
pub time_to_maturity: f64,
pub iv: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct VolatilitySurfacePoint {
pub strike: Decimal,
pub maturity: f64,
pub iv: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct VolatilitySurface {
pub points: Vec<VolatilitySurfacePoint>,
pub spot_price: Decimal,
pub timestamp: DateTime<Utc>,
}
#[derive(Debug, Clone)]
pub struct VolatilitySurfaceModeler {
spot_price: Decimal,
}
impl VolatilitySurfaceModeler {
pub fn new(spot_price: Decimal) -> Self {
Self { spot_price }
}
pub fn build_surface(&self, observations: &[ImpliedVolatility]) -> VolatilitySurface {
let mut points = Vec::new();
for obs in observations {
points.push(VolatilitySurfacePoint {
strike: obs.strike,
maturity: obs.time_to_maturity,
iv: obs.iv,
});
}
VolatilitySurface {
points,
spot_price: self.spot_price,
timestamp: Utc::now(),
}
}
pub fn interpolate_iv(
&self,
surface: &VolatilitySurface,
strike: Decimal,
maturity: f64,
) -> f64 {
if surface.points.is_empty() {
return 0.3; }
let mut distances: Vec<(usize, f64)> = surface
.points
.iter()
.enumerate()
.map(|(i, p)| {
let strike_dist = (p.strike - strike).abs().to_f64().unwrap_or(0.0);
let maturity_dist = (p.maturity - maturity).abs();
let distance = (strike_dist.powi(2) + maturity_dist.powi(2)).sqrt();
(i, distance)
})
.collect();
distances.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap());
let k = 3.min(distances.len());
let mut weighted_iv = 0.0;
let mut total_weight = 0.0;
for &(idx, distance) in distances.iter().take(k) {
let weight = 1.0 / (distance + 0.001);
weighted_iv += surface.points[idx].iv * weight;
total_weight += weight;
}
if total_weight > 0.0 {
weighted_iv / total_weight
} else {
0.3
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct JumpEvent {
pub timestamp: DateTime<Utc>,
pub size: f64,
pub direction: JumpDirection,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum JumpDirection {
Up,
Down,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct JumpDetectionResult {
pub jumps: Vec<JumpEvent>,
pub intensity: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct JumpPrediction {
pub jump_probability: f64,
pub expected_jump_size: f64,
}
#[derive(Debug, Clone)]
pub struct JumpDetector {
data: VecDeque<PricePoint>,
max_size: usize,
threshold: f64,
}
impl JumpDetector {
pub fn new(max_size: usize, threshold: f64) -> Self {
Self {
data: VecDeque::new(),
max_size,
threshold,
}
}
pub fn add_data(&mut self, point: PricePoint) {
self.data.push_back(point);
if self.data.len() > self.max_size {
self.data.pop_front();
}
}
pub fn detect_jumps(&self) -> anyhow::Result<JumpDetectionResult> {
if self.data.len() < 10 {
return Err(CoreError::Validation("Insufficient data".to_string()).into());
}
let data_vec: Vec<_> = self.data.iter().collect();
let returns: Vec<(DateTime<Utc>, f64)> = data_vec
.windows(2)
.map(|w| {
let prev = w[0].close.to_f64().unwrap_or(0.0);
let curr = w[1].close.to_f64().unwrap_or(0.0);
let ret = if prev > 0.0 {
((curr - prev) / prev).ln()
} else {
0.0
};
(w[1].timestamp, ret)
})
.collect();
let mean_return = returns.iter().map(|(_, r)| r).sum::<f64>() / returns.len() as f64;
let variance = returns
.iter()
.map(|(_, r)| (r - mean_return).powi(2))
.sum::<f64>()
/ returns.len() as f64;
let std_dev = variance.sqrt();
let mut jumps = Vec::new();
for (timestamp, ret) in returns.iter() {
let z_score = (ret - mean_return) / std_dev;
if z_score.abs() > self.threshold {
jumps.push(JumpEvent {
timestamp: *timestamp,
size: ret.abs(),
direction: if *ret > 0.0 {
JumpDirection::Up
} else {
JumpDirection::Down
},
});
}
}
let intensity = jumps.len() as f64 / returns.len() as f64;
Ok(JumpDetectionResult { jumps, intensity })
}
pub fn predict_jump(&self) -> anyhow::Result<JumpPrediction> {
let detection = self.detect_jumps()?;
let lambda = detection.intensity;
let jump_probability = 1.0 - (-lambda).exp();
let expected_jump_size = if !detection.jumps.is_empty() {
detection.jumps.iter().map(|j| j.size).sum::<f64>() / detection.jumps.len() as f64
} else {
0.05 };
Ok(JumpPrediction {
jump_probability,
expected_jump_size,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::Duration;
use rust_decimal_macros::dec;
fn create_test_data() -> Vec<PricePoint> {
let mut data = Vec::new();
let mut timestamp = Utc::now();
for i in 0..100 {
let base_price = if i < 30 {
100.0
} else if i < 60 {
150.0 } else {
100.0 };
data.push(PricePoint {
timestamp,
open: Decimal::from_f64_retain(base_price).unwrap(),
high: Decimal::from_f64_retain(base_price + 1.0).unwrap(),
low: Decimal::from_f64_retain(base_price - 1.0).unwrap(),
close: Decimal::from_f64_retain(base_price).unwrap(),
volume: dec!(1000),
});
timestamp += Duration::hours(1);
}
data
}
#[test]
fn test_realized_volatility_forecaster() {
let data = create_test_data();
let mut forecaster = RealizedVolatilityForecaster::new(200);
for point in data {
forecaster.add_data(point);
}
forecaster.fit().unwrap();
let forecast = forecaster.forecast(5).unwrap();
assert_eq!(forecast.horizon, 5);
assert!(forecast.predicted_volatility > 0.0);
assert!(forecast.confidence_interval.0 >= 0.0);
assert!(forecast.confidence_interval.1 > forecast.confidence_interval.0);
}
#[test]
fn test_volatility_surface_modeler() {
let observations = vec![
ImpliedVolatility {
strike: dec!(100),
time_to_maturity: 0.25,
iv: 0.20,
},
ImpliedVolatility {
strike: dec!(110),
time_to_maturity: 0.25,
iv: 0.25,
},
ImpliedVolatility {
strike: dec!(100),
time_to_maturity: 0.50,
iv: 0.22,
},
];
let modeler = VolatilitySurfaceModeler::new(dec!(105));
let surface = modeler.build_surface(&observations);
assert_eq!(surface.points.len(), 3);
assert_eq!(surface.spot_price, dec!(105));
let iv = modeler.interpolate_iv(&surface, dec!(105), 0.30);
assert!(iv > 0.0);
}
#[test]
fn test_jump_detector() {
let data = create_test_data();
let mut detector = JumpDetector::new(200, 2.0);
for point in data {
detector.add_data(point);
}
let detection = detector.detect_jumps().unwrap();
assert!(detection.intensity >= 0.0);
let prediction = detector.predict_jump().unwrap();
assert!(prediction.jump_probability >= 0.0 && prediction.jump_probability <= 1.0);
if !detection.jumps.is_empty() {
assert!(prediction.expected_jump_size > 0.0);
}
}
}