1pub mod engine;
20pub mod walk_forward;
21
22pub use engine::{
23 BacktestEngine, BacktestMetrics, BacktestResult as EngineBacktestResult, CompletedTrade,
24 Direction, EngineConfig, EngineSignal,
25};
26pub use walk_forward::{
27 ParamRange, WalkForwardConfig, WalkForwardOptimizer, WalkForwardResult, WfPeriod,
28};
29
30use crate::error::FinError;
31use crate::ohlcv::OhlcvBar;
32use crate::position::PositionLedger;
33use crate::types::{NanoTimestamp, Price, Quantity, Side};
34use rust_decimal::Decimal;
35use std::collections::HashMap;
36
37#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
41pub struct BacktestConfig {
42 pub initial_capital: Decimal,
44 pub commission_rate: Decimal,
46}
47
48impl BacktestConfig {
49 pub fn new(initial_capital: Decimal, commission_rate: Decimal) -> Result<Self, FinError> {
54 if initial_capital <= Decimal::ZERO {
55 return Err(FinError::InvalidInput(
56 "initial_capital must be positive".to_owned(),
57 ));
58 }
59 if commission_rate < Decimal::ZERO {
60 return Err(FinError::InvalidInput(
61 "commission_rate must be non-negative".to_owned(),
62 ));
63 }
64 Ok(Self { initial_capital, commission_rate })
65 }
66}
67
68#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
70pub struct BacktestResult {
71 pub total_return: Decimal,
73 pub sharpe_ratio: Decimal,
75 pub max_drawdown: Decimal,
77 pub win_rate: Decimal,
79 pub trade_count: u64,
81 pub final_equity: Decimal,
83 pub equity_curve: Vec<Decimal>,
85}
86
87#[derive(Debug, Clone, Copy, PartialEq, Eq)]
91pub enum SignalDirection {
92 Buy,
94 Sell,
96 Hold,
98}
99
100#[derive(Debug, Clone)]
102pub struct Signal {
103 pub direction: SignalDirection,
105 pub quantity: Decimal,
107}
108
109impl Signal {
110 pub fn new(direction: SignalDirection, quantity: Decimal) -> Self {
112 Self { direction, quantity }
113 }
114
115 pub fn hold() -> Self {
117 Self::new(SignalDirection::Hold, Decimal::ZERO)
118 }
119}
120
121pub trait Strategy: Send {
127 fn on_bar(&mut self, bar: &OhlcvBar) -> Option<Signal>;
131}
132
133pub struct Backtester {
140 config: BacktestConfig,
141}
142
143impl Backtester {
144 pub fn new(config: BacktestConfig) -> Self {
146 Self { config }
147 }
148
149 pub fn run(
155 &self,
156 bars: &[OhlcvBar],
157 strategy: &mut dyn Strategy,
158 ) -> Result<BacktestResult, FinError> {
159 if bars.is_empty() {
160 return Err(FinError::InvalidInput("bars slice must not be empty".to_owned()));
161 }
162
163 let mut ledger = PositionLedger::new(self.config.initial_capital);
164 let mut equity_curve: Vec<Decimal> = Vec::with_capacity(bars.len());
165 let mut trade_count: u64 = 0;
166 let mut daily_returns: Vec<Decimal> = Vec::with_capacity(bars.len());
167 let mut prev_equity = self.config.initial_capital;
168 let mut peak_equity = self.config.initial_capital;
169 let mut max_drawdown = Decimal::ZERO;
170
171 let mut winning_trades: u64 = 0;
173 let mut total_closed: u64 = 0;
174
175 for bar in bars {
176 if let Some(sig) = strategy.on_bar(bar) {
178 if sig.direction != SignalDirection::Hold && sig.quantity > Decimal::ZERO {
179 let side = match sig.direction {
180 SignalDirection::Buy => Side::Bid,
181 SignalDirection::Sell => Side::Ask,
182 SignalDirection::Hold => unreachable!(),
183 };
184
185 let price = Price::new(bar.close.value())?;
186 let qty = Quantity::new(sig.quantity)?;
187 let commission = bar.close.value() * sig.quantity * self.config.commission_rate;
188
189 let fill = crate::position::Fill::with_commission(
190 bar.symbol.clone(),
191 side,
192 qty,
193 price,
194 NanoTimestamp::new(bar.ts_close.nanos()),
195 commission,
196 );
197
198 let realized_before = ledger.realized_pnl_total();
200 if ledger.apply_fill(fill).is_ok() {
201 let realized_after = ledger.realized_pnl_total();
202 let pnl_delta = realized_after - realized_before;
203 if pnl_delta != Decimal::ZERO {
204 total_closed += 1;
205 if pnl_delta > Decimal::ZERO {
206 winning_trades += 1;
207 }
208 }
209 }
210 trade_count += 1;
211 }
212 }
213
214 let mut mark_prices: HashMap<String, Price> = HashMap::new();
216 mark_prices.insert(
217 bar.symbol.as_str().to_owned(),
218 Price::new(bar.close.value())?,
219 );
220 let equity = ledger.equity(&mark_prices).unwrap_or(prev_equity);
221
222 if equity > peak_equity {
224 peak_equity = equity;
225 }
226 if peak_equity > Decimal::ZERO {
227 let dd = (peak_equity - equity) / peak_equity;
228 if dd > max_drawdown {
229 max_drawdown = dd;
230 }
231 }
232
233 if prev_equity > Decimal::ZERO {
235 daily_returns.push((equity - prev_equity) / prev_equity);
236 }
237
238 equity_curve.push(equity);
239 prev_equity = equity;
240 }
241
242 let final_equity = equity_curve.last().copied().unwrap_or(self.config.initial_capital);
243
244 let total_return = if self.config.initial_capital > Decimal::ZERO {
245 (final_equity - self.config.initial_capital) / self.config.initial_capital
246 } else {
247 Decimal::ZERO
248 };
249
250 let sharpe_ratio = compute_sharpe(&daily_returns);
251
252 let win_rate = if total_closed > 0 {
253 Decimal::from(winning_trades) / Decimal::from(total_closed)
254 } else {
255 Decimal::ZERO
256 };
257
258 Ok(BacktestResult {
259 total_return,
260 sharpe_ratio,
261 max_drawdown,
262 win_rate,
263 trade_count,
264 final_equity,
265 equity_curve,
266 })
267 }
268}
269
270fn compute_sharpe(returns: &[Decimal]) -> Decimal {
275 use rust_decimal::prelude::ToPrimitive;
276
277 let n = returns.len();
278 if n < 2 {
279 return Decimal::ZERO;
280 }
281
282 let sum: Decimal = returns.iter().sum();
284 let mean_f = sum.to_f64().unwrap_or(0.0) / n as f64;
285
286 let var: f64 = returns
288 .iter()
289 .map(|r| {
290 let x = r.to_f64().unwrap_or(0.0) - mean_f;
291 x * x
292 })
293 .sum::<f64>()
294 / (n as f64 - 1.0);
295
296 let std_dev = var.sqrt();
297 if std_dev == 0.0 {
298 return Decimal::ZERO;
299 }
300
301 let sharpe_daily = mean_f / std_dev;
302 let sharpe_annual = sharpe_daily * 252.0_f64.sqrt();
303
304 Decimal::try_from(sharpe_annual).unwrap_or(Decimal::ZERO)
305}
306
307#[cfg(test)]
313mod tests {
314 use super::*;
315 use crate::types::{NanoTimestamp, Price, Quantity};
316 use rust_decimal_macros::dec;
317
318 fn make_bar(close: Decimal, ts: i64) -> OhlcvBar {
319 let sym = crate::types::Symbol::new("TEST").unwrap();
320 let p = Price::new(close).unwrap();
321 OhlcvBar {
322 symbol: sym,
323 open: p,
324 high: p,
325 low: p,
326 close: p,
327 volume: Quantity::new(dec!(1000)).unwrap(),
328 ts_open: NanoTimestamp::new(ts),
329 ts_close: NanoTimestamp::new(ts + 1),
330 tick_count: 1,
331 }
332 }
333
334 struct BuyAndHold {
336 bought: bool,
337 }
338
339 impl Strategy for BuyAndHold {
340 fn on_bar(&mut self, _bar: &OhlcvBar) -> Option<Signal> {
341 if !self.bought {
342 self.bought = true;
343 Some(Signal::new(SignalDirection::Buy, dec!(1)))
344 } else {
345 Some(Signal::hold())
346 }
347 }
348 }
349
350 #[test]
351 fn test_buy_and_hold_rising_market() {
352 let bars: Vec<OhlcvBar> = (0..10)
353 .map(|i| make_bar(dec!(100) + Decimal::from(i), i))
354 .collect();
355 let config = BacktestConfig::new(dec!(10_000), dec!(0)).unwrap();
356 let result = Backtester::new(config)
357 .run(&bars, &mut BuyAndHold { bought: false })
358 .unwrap();
359 assert!(result.final_equity > dec!(9_900), "final_equity={}", result.final_equity);
363 assert_eq!(result.trade_count, 1);
364 }
365
366 #[test]
367 fn test_empty_bars_errors() {
368 let config = BacktestConfig::new(dec!(10_000), dec!(0)).unwrap();
369 let result = Backtester::new(config).run(&[], &mut BuyAndHold { bought: false });
370 assert!(result.is_err());
371 }
372
373 #[test]
374 fn test_max_drawdown_flat_market_is_zero() {
375 struct HoldOnly;
377 impl Strategy for HoldOnly {
378 fn on_bar(&mut self, _bar: &OhlcvBar) -> Option<Signal> {
379 Some(Signal::hold())
380 }
381 }
382 let bars: Vec<OhlcvBar> = (0..5).map(|i| make_bar(dec!(100), i)).collect();
383 let config = BacktestConfig::new(dec!(10_000), dec!(0)).unwrap();
384 let result = Backtester::new(config).run(&bars, &mut HoldOnly).unwrap();
385 assert_eq!(result.max_drawdown, dec!(0));
386 }
387
388 #[test]
389 fn test_backtest_config_invalid_capital() {
390 assert!(BacktestConfig::new(dec!(-1), dec!(0)).is_err());
391 }
392
393 #[test]
394 fn test_walk_forward_basic() {
395 use crate::backtest::walk_forward::WalkForwardConfig;
396 use std::collections::HashMap;
397 let bars: Vec<OhlcvBar> = (0..30)
398 .map(|i| make_bar(dec!(100) + Decimal::from(i), i))
399 .collect();
400 let bt_config = BacktestConfig::new(dec!(10_000), dec!(0)).unwrap();
401 let wf_config = WalkForwardConfig {
402 train_window: 15,
403 test_window: 5,
404 step: 5,
405 param_space: vec![],
406 };
407 let wfo = WalkForwardOptimizer::new(wf_config, bt_config).unwrap();
408 let result = wfo
409 .run(&bars, |_train, _params: &HashMap<String, f64>| {
410 Box::new(BuyAndHold { bought: false })
411 })
412 .unwrap();
413 assert!(!result.periods.is_empty());
414 }
415
416 #[test]
417 fn test_walk_forward_insufficient_bars() {
418 use crate::backtest::walk_forward::WalkForwardConfig;
419 use std::collections::HashMap;
420 let bars: Vec<OhlcvBar> = (0..5).map(|i| make_bar(dec!(100), i)).collect();
421 let bt_config = BacktestConfig::new(dec!(10_000), dec!(0)).unwrap();
422 let wf_config = WalkForwardConfig {
423 train_window: 10,
424 test_window: 5,
425 step: 5,
426 param_space: vec![],
427 };
428 let wfo = WalkForwardOptimizer::new(wf_config, bt_config).unwrap();
429 let result = wfo.run(&bars, |_train, _params: &HashMap<String, f64>| {
430 Box::new(BuyAndHold { bought: false })
431 });
432 assert!(result.is_err());
433 }
434
435 #[test]
436 fn test_sharpe_constant_returns_zero() {
437 let returns = vec![dec!(0.01); 10];
439 let s = compute_sharpe(&returns);
441 assert_eq!(s, Decimal::ZERO);
443 }
444}