mod benchmark;
mod exits;
#[cfg(test)]
mod fixtures;
mod indicators;
mod margin;
mod positions;
mod simulate;
mod sizing;
pub(crate) use exits::{check_sl_tp, update_position_extremes, update_trailing_hwm};
pub(crate) use indicators::compute_for_candles;
pub(crate) use sizing::SizingSeries;
use crate::models::chart::{Candle, Dividend};
use self::benchmark::compute_benchmark_metrics;
use super::config::BacktestConfig;
use super::error::{BacktestError, Result};
use super::result::BacktestResult;
use super::strategy::Strategy;
pub(crate) fn validate_series_order(candles: &[Candle], dividends: &[Dividend]) -> Result<()> {
if !candles.windows(2).all(|w| w[0].timestamp <= w[1].timestamp) {
return Err(BacktestError::invalid_param(
"candles",
"must be sorted by timestamp (ascending)",
));
}
if !dividends
.windows(2)
.all(|w| w[0].timestamp <= w[1].timestamp)
{
return Err(BacktestError::invalid_param(
"dividends",
"must be sorted by timestamp (ascending)",
));
}
Ok(())
}
pub struct BacktestEngine {
config: BacktestConfig,
}
impl BacktestEngine {
pub fn new(config: BacktestConfig) -> Self {
Self { config }
}
pub fn run<S: Strategy>(
&self,
symbol: &str,
candles: &[Candle],
strategy: S,
) -> Result<BacktestResult> {
validate_series_order(candles, &[])?;
self.simulate(symbol, candles, strategy, &[])
}
pub fn run_with_dividends<S: Strategy>(
&self,
symbol: &str,
candles: &[Candle],
strategy: S,
dividends: &[Dividend],
) -> Result<BacktestResult> {
validate_series_order(candles, dividends)?;
self.simulate(symbol, candles, strategy, dividends)
}
pub fn run_with_benchmark<S: Strategy>(
&self,
symbol: &str,
candles: &[Candle],
strategy: S,
dividends: &[Dividend],
benchmark_symbol: &str,
benchmark_candles: &[Candle],
) -> Result<BacktestResult> {
validate_series_order(candles, dividends)?;
let mut result = self.simulate(symbol, candles, strategy, dividends)?;
result.benchmark = Some(compute_benchmark_metrics(
benchmark_symbol,
candles,
benchmark_candles,
&result.equity_curve,
self.config.risk_free_rate,
self.config.bars_per_year,
));
Ok(result)
}
}
#[cfg(test)]
mod tests {
use super::fixtures::make_candles;
use super::*;
use crate::backtesting::strategy::SmaCrossover;
#[test]
fn test_engine_basic() {
let mut prices = vec![100.0; 30];
for (i, price) in prices.iter_mut().enumerate().take(25).skip(15) {
*price = 100.0 + (i - 15) as f64 * 2.0;
}
for (i, price) in prices.iter_mut().enumerate().take(30).skip(25) {
*price = 118.0 - (i - 25) as f64 * 3.0;
}
let candles = make_candles(&prices);
let config = BacktestConfig::builder()
.initial_capital(10_000.0)
.commission_pct(0.0)
.slippage_pct(0.0)
.build()
.unwrap();
let engine = BacktestEngine::new(config);
let strategy = SmaCrossover::new(5, 10);
let result = engine.run("TEST", &candles, strategy).unwrap();
assert_eq!(result.symbol, "TEST");
assert_eq!(result.strategy_name, "SMA Crossover");
assert!(!result.equity_curve.is_empty());
}
#[test]
fn test_stop_loss() {
let mut prices = vec![100.0; 20];
for (i, price) in prices.iter_mut().enumerate().take(15).skip(10) {
*price = 100.0 + (i - 10) as f64 * 2.0;
}
for (i, price) in prices.iter_mut().enumerate().take(20).skip(15) {
*price = 108.0 - (i - 15) as f64 * 10.0;
}
let candles = make_candles(&prices);
let config = BacktestConfig::builder()
.initial_capital(10_000.0)
.stop_loss_pct(0.05) .commission_pct(0.0)
.slippage_pct(0.0)
.build()
.unwrap();
let engine = BacktestEngine::new(config);
let strategy = SmaCrossover::new(3, 6);
let result = engine.run("TEST", &candles, strategy).unwrap();
let _sl_signals: Vec<_> = result
.signals
.iter()
.filter(|s| {
s.reason
.as_ref()
.map(|r| r.contains("Stop-loss"))
.unwrap_or(false)
})
.collect();
assert!(!result.equity_curve.is_empty());
}
#[test]
fn test_trailing_stop() {
let mut prices: Vec<f64> = (0..20).map(|i| 100.0 + i as f64).collect();
prices.extend_from_slice(&[105.0, 103.0, 101.0]);
let candles = make_candles(&prices);
let config = BacktestConfig::builder()
.initial_capital(10_000.0)
.trailing_stop_pct(0.10)
.commission_pct(0.0)
.slippage_pct(0.0)
.build()
.unwrap();
let engine = BacktestEngine::new(config);
let strategy = SmaCrossover::new(3, 6);
let result = engine.run("TEST", &candles, strategy).unwrap();
let trail_exits: Vec<_> = result
.signals
.iter()
.filter(|s| {
s.reason
.as_ref()
.map(|r| r.contains("Trailing stop"))
.unwrap_or(false)
})
.collect();
let _ = trail_exits;
assert!(!result.equity_curve.is_empty());
}
}