use crate::backtest::{BacktestConfig, Backtester, BacktestResult, Strategy};
use crate::error::FinError;
use crate::ohlcv::OhlcvBar;
use std::collections::HashMap;
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct ParamRange {
pub name: String,
pub min: f64,
pub max: f64,
pub step: f64,
}
impl ParamRange {
pub fn values(&self) -> Vec<f64> {
if self.step <= 0.0 || self.min > self.max {
return vec![self.min];
}
let mut vals = Vec::new();
let mut v = self.min;
while v <= self.max + f64::EPSILON {
vals.push(v);
v += self.step;
}
vals
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct WalkForwardConfig {
pub train_window: usize,
pub test_window: usize,
pub step: usize,
pub param_space: Vec<ParamRange>,
}
impl WalkForwardConfig {
pub fn validate(&self) -> Result<(), FinError> {
if self.train_window == 0 {
return Err(FinError::InvalidInput(
"train_window must be > 0".to_owned(),
));
}
if self.test_window == 0 {
return Err(FinError::InvalidInput(
"test_window must be > 0".to_owned(),
));
}
if self.step == 0 {
return Err(FinError::InvalidInput(
"step must be > 0".to_owned(),
));
}
Ok(())
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct WfPeriod {
pub train_start: usize,
pub train_end: usize,
pub test_start: usize,
pub test_end: usize,
pub best_params: HashMap<String, f64>,
pub in_sample_sharpe: f64,
pub out_of_sample_sharpe: f64,
pub oos_result: BacktestResult,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct WalkForwardResult {
pub periods: Vec<WfPeriod>,
pub aggregate_sharpe: f64,
pub stability_score: f64,
pub mean_oos_return: f64,
pub worst_oos_drawdown: f64,
}
impl WalkForwardResult {
pub fn is_robust(&self, min_stability: f64) -> bool {
self.aggregate_sharpe > 0.0 && self.stability_score >= min_stability
}
pub fn best_period(&self) -> Option<&WfPeriod> {
self.periods
.iter()
.max_by(|a, b| a.out_of_sample_sharpe.partial_cmp(&b.out_of_sample_sharpe).unwrap_or(std::cmp::Ordering::Equal))
}
pub fn worst_period(&self) -> Option<&WfPeriod> {
self.periods
.iter()
.min_by(|a, b| a.out_of_sample_sharpe.partial_cmp(&b.out_of_sample_sharpe).unwrap_or(std::cmp::Ordering::Equal))
}
}
pub struct WalkForwardOptimizer {
config: WalkForwardConfig,
bt_config: BacktestConfig,
}
impl WalkForwardOptimizer {
pub fn new(config: WalkForwardConfig, bt_config: BacktestConfig) -> Result<Self, FinError> {
config.validate()?;
Ok(Self { config, bt_config })
}
pub fn run<F>(
&self,
bars: &[OhlcvBar],
mut make_strategy: F,
) -> Result<WalkForwardResult, FinError>
where
F: FnMut(&[OhlcvBar], &HashMap<String, f64>) -> Box<dyn Strategy>,
{
let window = self.config.train_window + self.config.test_window;
if bars.len() < window {
return Err(FinError::InvalidInput(format!(
"need at least {} bars for one walk-forward window, got {}",
window,
bars.len()
)));
}
let backtester = Backtester::new(self.bt_config.clone());
let grid = build_grid(&self.config.param_space);
let mut periods: Vec<WfPeriod> = Vec::new();
let mut offset = 0usize;
while offset + window <= bars.len() {
let train_start = offset;
let train_end = offset + self.config.train_window;
let test_start = train_end;
let test_end = train_end + self.config.test_window;
let train_bars = &bars[train_start..train_end];
let test_bars = &bars[test_start..test_end];
let mut best_is_sharpe = f64::NEG_INFINITY;
let mut best_params: HashMap<String, f64> = HashMap::new();
let search_grid: &[HashMap<String, f64>] = &grid;
for param_set in search_grid {
let mut is_strategy = make_strategy(train_bars, param_set);
let is_result = match backtester.run(train_bars, is_strategy.as_mut()) {
Ok(r) => r,
Err(_) => continue,
};
let is_sharpe = is_result
.sharpe_ratio
.to_string()
.parse::<f64>()
.unwrap_or(f64::NEG_INFINITY);
if is_sharpe > best_is_sharpe {
best_is_sharpe = is_sharpe;
best_params = param_set.clone();
}
}
let mut oos_strategy = make_strategy(test_bars, &best_params);
let oos_result = backtester.run(test_bars, oos_strategy.as_mut())?;
let oos_sharpe = oos_result
.sharpe_ratio
.to_string()
.parse::<f64>()
.unwrap_or(0.0);
periods.push(WfPeriod {
train_start,
train_end,
test_start,
test_end,
best_params,
in_sample_sharpe: best_is_sharpe.max(0.0), out_of_sample_sharpe: oos_sharpe,
oos_result,
});
offset += self.config.step;
}
if periods.is_empty() {
return Err(FinError::InvalidInput(
"no walk-forward periods could be constructed".to_owned(),
));
}
let n = periods.len() as f64;
let aggregate_sharpe = periods.iter().map(|p| p.out_of_sample_sharpe).sum::<f64>() / n;
let positive_count = periods.iter().filter(|p| p.out_of_sample_sharpe > 0.0).count();
let stability_score = positive_count as f64 / n;
let mean_oos_return = periods
.iter()
.map(|p| {
p.oos_result
.total_return
.to_string()
.parse::<f64>()
.unwrap_or(0.0)
})
.sum::<f64>()
/ n;
let worst_oos_drawdown = periods
.iter()
.map(|p| {
p.oos_result
.max_drawdown
.to_string()
.parse::<f64>()
.unwrap_or(0.0)
})
.fold(0.0_f64, f64::max);
Ok(WalkForwardResult {
periods,
aggregate_sharpe,
stability_score,
mean_oos_return,
worst_oos_drawdown,
})
}
pub fn config(&self) -> &WalkForwardConfig {
&self.config
}
pub fn bt_config(&self) -> &BacktestConfig {
&self.bt_config
}
}
fn build_grid(param_space: &[ParamRange]) -> Vec<HashMap<String, f64>> {
if param_space.is_empty() {
return vec![HashMap::new()];
}
let mut grid: Vec<HashMap<String, f64>> = vec![HashMap::new()];
for param in param_space {
let vals = param.values();
let mut new_grid = Vec::with_capacity(grid.len() * vals.len());
for existing in &grid {
for &v in &vals {
let mut m = existing.clone();
m.insert(param.name.clone(), v);
new_grid.push(m);
}
}
grid = new_grid;
}
grid
}
#[cfg(test)]
mod tests {
use super::*;
use crate::backtest::{Signal, SignalDirection};
use crate::types::{NanoTimestamp, Price, Quantity, Symbol};
use rust_decimal::Decimal;
use rust_decimal_macros::dec;
fn make_bar(close: f64, ts: i64) -> OhlcvBar {
let sym = Symbol::new("TEST").unwrap();
let p = Price::new(Decimal::try_from(close).unwrap()).unwrap();
OhlcvBar {
symbol: sym,
open: p,
high: p,
low: p,
close: p,
volume: Quantity::new(dec!(1000)).unwrap(),
ts_open: NanoTimestamp::new(ts),
ts_close: NanoTimestamp::new(ts + 1),
tick_count: 1,
}
}
struct HoldAll;
impl Strategy for HoldAll {
fn on_bar(&mut self, _: &OhlcvBar) -> Option<Signal> {
None
}
}
struct BuyOnce {
bought: bool,
qty: f64,
}
impl BuyOnce {
fn from_params(params: &HashMap<String, f64>) -> Self {
Self {
bought: false,
qty: params.get("qty").copied().unwrap_or(1.0),
}
}
}
impl Strategy for BuyOnce {
fn on_bar(&mut self, _: &OhlcvBar) -> Option<Signal> {
if !self.bought {
self.bought = true;
let qty = Decimal::try_from(self.qty).unwrap_or(dec!(1));
return Some(Signal::new(SignalDirection::Buy, qty));
}
Some(Signal::hold())
}
}
#[test]
fn test_param_range_values_basic() {
let r = ParamRange { name: "x".to_owned(), min: 1.0, max: 3.0, step: 1.0 };
let vals = r.values();
assert_eq!(vals.len(), 3);
assert!((vals[0] - 1.0).abs() < 1e-10);
assert!((vals[2] - 3.0).abs() < 1e-10);
}
#[test]
fn test_param_range_single_value_when_min_equals_max() {
let r = ParamRange { name: "x".to_owned(), min: 5.0, max: 5.0, step: 1.0 };
let vals = r.values();
assert_eq!(vals.len(), 1);
assert!((vals[0] - 5.0).abs() < 1e-10);
}
#[test]
fn test_param_range_degenerate_step() {
let r = ParamRange { name: "x".to_owned(), min: 1.0, max: 5.0, step: 0.0 };
let vals = r.values();
assert_eq!(vals.len(), 1); }
#[test]
fn test_build_grid_empty_space() {
let grid = build_grid(&[]);
assert_eq!(grid.len(), 1);
assert!(grid[0].is_empty());
}
#[test]
fn test_build_grid_single_param() {
let params = vec![ParamRange { name: "p".to_owned(), min: 5.0, max: 15.0, step: 5.0 }];
let grid = build_grid(¶ms);
assert_eq!(grid.len(), 3); for m in &grid {
assert!(m.contains_key("p"));
}
}
#[test]
fn test_build_grid_two_params_cartesian() {
let params = vec![
ParamRange { name: "a".to_owned(), min: 1.0, max: 2.0, step: 1.0 },
ParamRange { name: "b".to_owned(), min: 10.0, max: 20.0, step: 10.0 },
];
let grid = build_grid(¶ms);
assert_eq!(grid.len(), 4); }
#[test]
fn test_config_validation_zero_train() {
let cfg = WalkForwardConfig {
train_window: 0,
test_window: 20,
step: 20,
param_space: vec![],
};
assert!(cfg.validate().is_err());
}
#[test]
fn test_config_validation_zero_test() {
let cfg = WalkForwardConfig {
train_window: 60,
test_window: 0,
step: 20,
param_space: vec![],
};
assert!(cfg.validate().is_err());
}
#[test]
fn test_config_validation_zero_step() {
let cfg = WalkForwardConfig {
train_window: 60,
test_window: 20,
step: 0,
param_space: vec![],
};
assert!(cfg.validate().is_err());
}
#[test]
fn test_optimizer_too_few_bars() {
let bars: Vec<OhlcvBar> = (0..5).map(|i| make_bar(100.0, i)).collect();
let cfg = WalkForwardConfig {
train_window: 60,
test_window: 20,
step: 20,
param_space: vec![],
};
let bt_cfg = BacktestConfig::new(dec!(10_000), dec!(0)).unwrap();
let opt = WalkForwardOptimizer::new(cfg, bt_cfg).unwrap();
let result = opt.run(&bars, |_, _| Box::new(HoldAll));
assert!(result.is_err());
}
#[test]
fn test_optimizer_hold_strategy_returns_result() {
let bars: Vec<OhlcvBar> = (0..100).map(|i| make_bar(100.0 + i as f64 * 0.1, i)).collect();
let cfg = WalkForwardConfig {
train_window: 40,
test_window: 20,
step: 20,
param_space: vec![],
};
let bt_cfg = BacktestConfig::new(dec!(10_000), dec!(0)).unwrap();
let opt = WalkForwardOptimizer::new(cfg, bt_cfg).unwrap();
let result = opt.run(&bars, |_, _| Box::new(HoldAll)).unwrap();
assert!(!result.periods.is_empty());
assert_eq!(result.aggregate_sharpe, 0.0);
assert_eq!(result.stability_score, 0.0);
}
#[test]
fn test_optimizer_with_param_grid() {
let bars: Vec<OhlcvBar> = (0..120).map(|i| make_bar(100.0 + i as f64 * 0.5, i)).collect();
let cfg = WalkForwardConfig {
train_window: 50,
test_window: 20,
step: 20,
param_space: vec![
ParamRange { name: "qty".to_owned(), min: 1.0, max: 3.0, step: 1.0 },
],
};
let bt_cfg = BacktestConfig::new(dec!(10_000), dec!(0)).unwrap();
let opt = WalkForwardOptimizer::new(cfg, bt_cfg).unwrap();
let result = opt.run(&bars, |_, params| Box::new(BuyOnce::from_params(params))).unwrap();
assert!(!result.periods.is_empty());
for p in &result.periods {
assert!(!p.best_params.is_empty());
assert!(p.best_params.contains_key("qty"));
}
}
#[test]
fn test_optimizer_stability_score_bounds() {
let bars: Vec<OhlcvBar> = (0..100).map(|i| make_bar(100.0, i)).collect();
let cfg = WalkForwardConfig {
train_window: 40,
test_window: 20,
step: 20,
param_space: vec![],
};
let bt_cfg = BacktestConfig::new(dec!(10_000), dec!(0)).unwrap();
let opt = WalkForwardOptimizer::new(cfg, bt_cfg).unwrap();
let result = opt.run(&bars, |_, _| Box::new(HoldAll)).unwrap();
assert!((0.0..=1.0).contains(&result.stability_score));
}
#[test]
fn test_wf_result_robustness_check() {
let result = WalkForwardResult {
periods: vec![],
aggregate_sharpe: 1.5,
stability_score: 0.8,
mean_oos_return: 0.05,
worst_oos_drawdown: 0.1,
};
assert!(result.is_robust(0.7));
assert!(!result.is_robust(0.9));
}
#[test]
fn test_wf_result_best_worst_period() {
let make_period = |oos_sharpe: f64| WfPeriod {
train_start: 0,
train_end: 50,
test_start: 50,
test_end: 70,
best_params: HashMap::new(),
in_sample_sharpe: 1.0,
out_of_sample_sharpe: oos_sharpe,
oos_result: crate::backtest::BacktestResult {
total_return: Decimal::ZERO,
sharpe_ratio: Decimal::ZERO,
max_drawdown: Decimal::ZERO,
win_rate: Decimal::ZERO,
trade_count: 0,
final_equity: dec!(10_000),
equity_curve: vec![],
},
};
let result = WalkForwardResult {
periods: vec![make_period(0.5), make_period(2.0), make_period(-0.3)],
aggregate_sharpe: 0.73,
stability_score: 0.67,
mean_oos_return: 0.0,
worst_oos_drawdown: 0.0,
};
assert!((result.best_period().unwrap().out_of_sample_sharpe - 2.0).abs() < 1e-10);
assert!((result.worst_period().unwrap().out_of_sample_sharpe + 0.3).abs() < 1e-10);
}
#[test]
fn test_optimizer_step_advances_window() {
let bars: Vec<OhlcvBar> = (0..150).map(|i| make_bar(100.0 + i as f64 * 0.1, i)).collect();
let cfg = WalkForwardConfig {
train_window: 50,
test_window: 30,
step: 10, param_space: vec![],
};
let bt_cfg = BacktestConfig::new(dec!(10_000), dec!(0)).unwrap();
let opt = WalkForwardOptimizer::new(cfg, bt_cfg).unwrap();
let result = opt.run(&bars, |_, _| Box::new(HoldAll)).unwrap();
assert!(result.periods.len() > 2);
}
}