use std::collections::HashMap;
use rayon::prelude::*;
use serde::{Deserialize, Serialize};
use crate::models::chart::Candle;
use super::config::BacktestConfig;
use super::error::{BacktestError, Result};
use super::optimizer::{GridSearch, OptimizationReport, ParamValue};
use super::result::{BacktestResult, PerformanceMetrics};
use super::strategy::Strategy;
#[non_exhaustive]
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WindowResult {
pub window: usize,
pub optimized_params: HashMap<String, ParamValue>,
pub in_sample: BacktestResult,
pub out_of_sample: BacktestResult,
}
#[non_exhaustive]
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WalkForwardReport {
pub strategy_name: String,
pub windows: Vec<WindowResult>,
pub aggregate_metrics: PerformanceMetrics,
pub consistency_ratio: f64,
pub optimization_reports: Vec<OptimizationReport>,
}
#[non_exhaustive]
#[derive(Debug, Clone)]
pub struct WalkForwardConfig {
pub grid: GridSearch,
pub config: BacktestConfig,
pub in_sample_bars: usize,
pub out_of_sample_bars: usize,
pub step_bars: Option<usize>,
}
impl WalkForwardConfig {
pub fn new(grid: GridSearch, config: BacktestConfig) -> Self {
Self {
grid,
config,
in_sample_bars: 252,
out_of_sample_bars: 63,
step_bars: None,
}
}
pub fn in_sample_bars(mut self, bars: usize) -> Self {
self.in_sample_bars = bars;
self
}
pub fn out_of_sample_bars(mut self, bars: usize) -> Self {
self.out_of_sample_bars = bars;
self
}
pub fn step_bars(mut self, bars: usize) -> Self {
self.step_bars = Some(bars);
self
}
pub fn run<S, F>(
&self,
symbol: &str,
candles: &[Candle],
factory: F,
) -> Result<WalkForwardReport>
where
S: Strategy + Clone + Send,
F: Fn(&HashMap<String, ParamValue>) -> S,
F: Send + Sync,
{
self.validate(candles.len())?;
crate::backtesting::engine::validate_series_order(candles, &[])?;
let step = self.step_bars.unwrap_or(self.out_of_sample_bars);
let total_bars = self.in_sample_bars + self.out_of_sample_bars;
let starts: Vec<usize> = {
let mut v = Vec::new();
let mut start = 0usize;
while start + total_bars <= candles.len() {
v.push(start);
start += step;
}
v
};
let mut windows: Vec<WindowResult> = Vec::with_capacity(starts.len());
let mut opt_reports: Vec<OptimizationReport> = Vec::with_capacity(starts.len());
let results: Vec<Result<(WindowResult, OptimizationReport)>> = starts
.par_iter()
.enumerate()
.map(|(idx, &start)| self.run_one_window(idx, start, symbol, candles, &factory))
.collect();
for r in results {
let (w, o) = r?;
windows.push(w);
opt_reports.push(o);
}
let strategy_name = windows[0].in_sample.strategy_name.clone();
let consistency_ratio = calculate_consistency_ratio(&windows);
let aggregate_metrics = aggregate_oos_metrics(
&windows,
self.config.risk_free_rate,
self.config.bars_per_year,
);
Ok(WalkForwardReport {
strategy_name,
windows,
aggregate_metrics,
consistency_ratio,
optimization_reports: opt_reports,
})
}
fn run_one_window<S, F>(
&self,
window_idx: usize,
start: usize,
symbol: &str,
candles: &[Candle],
factory: &F,
) -> Result<(WindowResult, OptimizationReport)>
where
S: Strategy + Clone + Send,
F: Fn(&HashMap<String, ParamValue>) -> S,
F: Send + Sync,
{
let is_end = start + self.in_sample_bars;
let oos_end = is_end + self.out_of_sample_bars;
let is_candles = &candles[start..is_end];
let oos_candles = &candles[is_end..oos_end];
let opt_report = self
.grid
.run(symbol, is_candles, &self.config, factory)
.map_err(|e| {
BacktestError::invalid_param(
"walk_forward",
format!("window {window_idx} optimisation failed: {e}"),
)
})?;
let best_params = opt_report.best.params.clone();
let is_result = opt_report.best.result.clone();
let oos_strategy = factory(&best_params);
let oos_result = crate::backtesting::BacktestEngine::new(self.config.clone())
.simulate(symbol, oos_candles, oos_strategy, &[])
.map_err(|e| {
BacktestError::invalid_param(
"walk_forward",
format!("window {window_idx} OOS run failed: {e}"),
)
})?;
Ok((
WindowResult {
window: window_idx,
optimized_params: best_params,
in_sample: is_result,
out_of_sample: oos_result,
},
opt_report,
))
}
fn validate(&self, num_candles: usize) -> Result<()> {
if self.in_sample_bars == 0 {
return Err(BacktestError::invalid_param(
"in_sample_bars",
"must be greater than zero",
));
}
if self.out_of_sample_bars == 0 {
return Err(BacktestError::invalid_param(
"out_of_sample_bars",
"must be greater than zero",
));
}
if self.step_bars == Some(0) {
return Err(BacktestError::invalid_param(
"step_bars",
"must be greater than zero",
));
}
let total_bars = self.in_sample_bars + self.out_of_sample_bars;
if num_candles < total_bars {
return Err(BacktestError::insufficient_data(total_bars, num_candles));
}
Ok(())
}
}
fn calculate_consistency_ratio(windows: &[WindowResult]) -> f64 {
if windows.is_empty() {
return 0.0;
}
let profitable = windows
.iter()
.filter(|w| w.out_of_sample.is_profitable())
.count();
profitable as f64 / windows.len() as f64
}
fn aggregate_oos_metrics(
windows: &[WindowResult],
risk_free_rate: f64,
bars_per_year: f64,
) -> PerformanceMetrics {
use crate::backtesting::result::EquityPoint;
let all_trades: Vec<_> = windows
.iter()
.flat_map(|w| w.out_of_sample.trades.iter().cloned())
.collect();
let mut combined_equity: Vec<EquityPoint> = Vec::new();
let mut running_equity = windows[0].out_of_sample.initial_capital;
for (window_idx, window) in windows.iter().enumerate() {
let window_initial = window.out_of_sample.initial_capital;
if window_initial <= 0.0 {
continue;
}
for (point_idx, point) in window.out_of_sample.equity_curve.iter().enumerate() {
if window_idx > 0 && point_idx == 0 {
continue;
}
let scaled_equity = running_equity * (point.equity / window_initial);
combined_equity.push(EquityPoint {
timestamp: point.timestamp,
equity: scaled_equity,
drawdown_pct: 0.0,
});
}
if let Some(last) = combined_equity.last() {
running_equity = last.equity;
}
}
let mut peak = f64::NEG_INFINITY;
for point in &mut combined_equity {
peak = peak.max(point.equity);
point.drawdown_pct = if peak > 0.0 {
(peak - point.equity) / peak
} else {
0.0
};
}
let initial_capital = windows
.first()
.map(|w| w.out_of_sample.initial_capital)
.unwrap_or(10_000.0);
let total_signals: usize = windows.iter().map(|w| w.out_of_sample.signals.len()).sum();
let executed_signals: usize = windows
.iter()
.map(|w| {
w.out_of_sample
.signals
.iter()
.filter(|s| s.executed)
.count()
})
.sum();
PerformanceMetrics::calculate(
&all_trades,
&combined_equity,
initial_capital,
total_signals,
executed_signals,
risk_free_rate,
bars_per_year,
)
}
#[cfg(test)]
#[path = "walk_forward_tests.rs"]
mod tests;