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