use super::*;
use crate::backtesting::{
BacktestConfig, SmaCrossover,
optimizer::{OptimizeMetric, ParamRange},
};
use crate::models::chart::Candle;
fn make_candles(prices: &[f64]) -> Vec<Candle> {
prices
.iter()
.enumerate()
.map(|(i, &p)| Candle {
timestamp: i as i64,
open: p,
high: p * 1.01,
low: p * 0.99,
close: p,
volume: 1000,
adj_close: Some(p),
provider_id: None,
})
.collect()
}
fn trending_prices(n: usize) -> Vec<f64> {
(0..n).map(|i| 100.0 + i as f64 * 0.3).collect()
}
#[test]
fn test_walk_forward_basic() {
let prices = trending_prices(300);
let candles = make_candles(&prices);
let config = BacktestConfig::builder()
.commission_pct(0.0)
.slippage_pct(0.0)
.build()
.unwrap();
let grid = GridSearch::new()
.param("fast", ParamRange::int_range(3, 9, 3))
.param("slow", ParamRange::int_range(10, 20, 10))
.optimize_for(OptimizeMetric::TotalReturn);
let report = WalkForwardConfig::new(grid, config)
.in_sample_bars(200)
.out_of_sample_bars(100)
.run("TEST", &candles, |params| {
SmaCrossover::new(
params["fast"].as_int() as usize,
params["slow"].as_int() as usize,
)
})
.unwrap();
assert_eq!(report.windows.len(), 1);
assert_eq!(report.strategy_name, "SMA Crossover");
assert!(report.consistency_ratio >= 0.0);
assert!(report.consistency_ratio <= 1.0);
}
#[test]
fn test_walk_forward_multiple_windows() {
let prices = trending_prices(500);
let candles = make_candles(&prices);
let config = BacktestConfig::builder()
.commission_pct(0.0)
.slippage_pct(0.0)
.build()
.unwrap();
let grid = GridSearch::new()
.param("fast", ParamRange::int_range(3, 6, 3))
.param("slow", ParamRange::int_range(10, 10, 1))
.optimize_for(OptimizeMetric::TotalReturn);
let report = WalkForwardConfig::new(grid, config)
.in_sample_bars(200)
.out_of_sample_bars(100)
.step_bars(100)
.run("TEST", &candles, |params| {
SmaCrossover::new(
params["fast"].as_int() as usize,
params["slow"].as_int() as usize,
)
})
.unwrap();
assert!(report.windows.len() >= 2);
assert_eq!(report.optimization_reports.len(), report.windows.len());
}
#[test]
fn walk_forward_windows_are_order_stable() {
let candles = make_candles(&trending_prices(900));
let config = BacktestConfig::builder()
.commission_pct(0.0)
.slippage_pct(0.0)
.build()
.unwrap();
let run = || {
let grid = GridSearch::new()
.param("fast", ParamRange::int_range(3, 9, 3))
.param("slow", ParamRange::int_range(10, 20, 10))
.optimize_for(OptimizeMetric::TotalReturn);
WalkForwardConfig::new(grid, config.clone())
.in_sample_bars(200)
.out_of_sample_bars(100)
.run("TEST", &candles, |params| {
SmaCrossover::new(
params["fast"].as_int() as usize,
params["slow"].as_int() as usize,
)
})
.unwrap()
};
let a = run();
let b = run();
assert!(
a.windows.len() >= 3,
"need multiple windows to test ordering"
);
assert_eq!(a.windows.len(), b.windows.len());
for (i, (x, y)) in a.windows.iter().zip(b.windows.iter()).enumerate() {
assert_eq!(x.window, i, "window index out of order at position {i}");
assert_eq!(y.window, i, "window index out of order at position {i}");
assert_eq!(
x.optimized_params, y.optimized_params,
"window {i} diverged"
);
assert_eq!(
x.out_of_sample.metrics.total_return_pct, y.out_of_sample.metrics.total_return_pct,
"window {i} diverged"
);
assert_eq!(
x.out_of_sample.start_timestamp, y.out_of_sample.start_timestamp,
"window {i} diverged"
);
assert_eq!(
x.out_of_sample.end_timestamp, y.out_of_sample.end_timestamp,
"window {i} diverged"
);
assert_eq!(
x.in_sample.metrics.total_return_pct, y.in_sample.metrics.total_return_pct,
"window {i} diverged"
);
}
for pair in a.windows.windows(2) {
assert!(
pair[0].out_of_sample.start_timestamp < pair[1].out_of_sample.start_timestamp,
"windows not in chronological order: {} >= {}",
pair[0].out_of_sample.start_timestamp,
pair[1].out_of_sample.start_timestamp
);
}
assert_eq!(a.optimization_reports.len(), a.windows.len());
for (i, (x, y)) in a
.optimization_reports
.iter()
.zip(b.optimization_reports.iter())
.enumerate()
{
assert_eq!(x.best.params, y.best.params, "opt report {i} diverged");
}
assert_eq!(a.consistency_ratio, b.consistency_ratio);
}
#[test]
fn test_step_bars_zero_errors() {
let candles = make_candles(&trending_prices(300));
let config = BacktestConfig::default();
let grid = GridSearch::new()
.param("fast", ParamRange::int_range(3, 6, 3))
.param("slow", ParamRange::int_range(10, 10, 1));
let result = WalkForwardConfig::new(grid, config)
.in_sample_bars(200)
.out_of_sample_bars(100)
.step_bars(0)
.run("TEST", &candles, |params| {
SmaCrossover::new(
params["fast"].as_int() as usize,
params["slow"].as_int() as usize,
)
});
assert!(result.is_err());
}
#[test]
fn test_insufficient_data_errors() {
let candles = make_candles(&trending_prices(50));
let config = BacktestConfig::default();
let grid = GridSearch::new()
.param("fast", ParamRange::int_range(3, 6, 3))
.param("slow", ParamRange::int_range(10, 10, 1));
let result = WalkForwardConfig::new(grid, config)
.in_sample_bars(200) .out_of_sample_bars(100)
.run("TEST", &candles, |params| {
SmaCrossover::new(
params["fast"].as_int() as usize,
params["slow"].as_int() as usize,
)
});
assert!(result.is_err());
}
#[test]
fn test_consistency_ratio_all_profitable() {
let prices: Vec<f64> = (0..300).map(|i| 100.0 + i as f64).collect();
let candles = make_candles(&prices);
let config = BacktestConfig::builder()
.commission_pct(0.0)
.slippage_pct(0.0)
.build()
.unwrap();
let grid = GridSearch::new()
.param("fast", ParamRange::int_range(3, 3, 1))
.param("slow", ParamRange::int_range(10, 10, 1))
.optimize_for(OptimizeMetric::TotalReturn);
let report = WalkForwardConfig::new(grid, config)
.in_sample_bars(150)
.out_of_sample_bars(100)
.run("TEST", &candles, |params| {
SmaCrossover::new(
params["fast"].as_int() as usize,
params["slow"].as_int() as usize,
)
})
.unwrap();
assert!(report.consistency_ratio >= 0.0);
}
#[test]
fn test_aggregate_equity_timestamps_are_monotonic() {
let prices: Vec<f64> = (0..600).map(|i| 100.0 + (i as f64) * 0.5).collect();
let candles = make_candles(&prices);
let config = BacktestConfig::builder()
.commission_pct(0.0)
.slippage_pct(0.0)
.build()
.unwrap();
let grid = GridSearch::new()
.param("fast", ParamRange::int_range(3, 3, 1))
.param("slow", ParamRange::int_range(10, 10, 1))
.optimize_for(OptimizeMetric::TotalReturn);
let report = WalkForwardConfig::new(grid, config)
.in_sample_bars(100)
.out_of_sample_bars(50)
.run("TEST", &candles, |params| {
SmaCrossover::new(
params["fast"].as_int() as usize,
params["slow"].as_int() as usize,
)
})
.unwrap();
let curve = &report.aggregate_metrics;
assert!(
report.windows.len() >= 2,
"Expected multiple windows for timestamp test"
);
let timestamps: Vec<i64> = report
.windows
.iter()
.flat_map(|w| w.out_of_sample.equity_curve.iter().map(|ep| ep.timestamp))
.collect();
for window in &report.windows {
let ts: Vec<i64> = window
.out_of_sample
.equity_curve
.iter()
.map(|ep| ep.timestamp)
.collect();
for pair in ts.windows(2) {
assert!(
pair[0] < pair[1],
"Equity curve timestamps not strictly increasing within window: {} >= {}",
pair[0],
pair[1]
);
}
}
let _ = curve;
let _ = timestamps;
}
#[test]
fn test_aggregate_oos_equity_timestamps_are_gapless_across_windows() {
let prices: Vec<f64> = (0..600).map(|i| 100.0 + (i as f64) * 0.5).collect();
let candles = make_candles(&prices);
let config = BacktestConfig::builder()
.commission_pct(0.0)
.slippage_pct(0.0)
.build()
.unwrap();
let grid = GridSearch::new()
.param("fast", ParamRange::int_range(3, 3, 1))
.param("slow", ParamRange::int_range(10, 10, 1))
.optimize_for(OptimizeMetric::TotalReturn);
let report = WalkForwardConfig::new(grid, config)
.in_sample_bars(100)
.out_of_sample_bars(50)
.run("TEST", &candles, |params| {
SmaCrossover::new(
params["fast"].as_int() as usize,
params["slow"].as_int() as usize,
)
})
.unwrap();
assert!(
report.windows.len() >= 2,
"Need at least 2 OOS windows for this test"
);
let combined_ts: Vec<i64> = report
.windows
.iter()
.enumerate()
.flat_map(|(wi, w)| {
w.out_of_sample
.equity_curve
.iter()
.enumerate()
.filter(move |&(pi, _)| !(wi > 0 && pi == 0))
.map(|(_, ep)| ep.timestamp)
})
.collect();
for pair in combined_ts.windows(2) {
assert!(
pair[0] < pair[1],
"Combined equity curve timestamps not strictly increasing: {} >= {}",
pair[0],
pair[1]
);
}
let expected_first = report
.windows
.first()
.and_then(|w| w.out_of_sample.equity_curve.first())
.map(|ep| ep.timestamp)
.unwrap_or(0);
assert_eq!(
combined_ts.first().copied().unwrap_or(-1),
expected_first,
"First combined timestamp should equal the first OOS equity point timestamp"
);
}