1use crate::ohlcv::OhlcvBar;
17
18#[derive(Debug, Clone, Copy, PartialEq, Eq)]
22pub enum Direction {
23 Long,
25 Short,
27 Flat,
29}
30
31#[derive(Debug, Clone)]
35pub struct EngineSignal {
36 pub timestamp: u64,
38 pub symbol: String,
40 pub direction: Direction,
42 pub strength: f64,
44}
45
46#[derive(Debug, Clone)]
50pub struct EngineConfig {
51 pub initial_capital: f64,
53 pub commission: f64,
55 pub slippage_bps: f64,
57 pub data: Vec<OhlcvBar>,
59 pub capital_fraction: f64,
61}
62
63#[derive(Debug, Clone)]
67pub struct CompletedTrade {
68 pub entry_ts: u64,
70 pub exit_ts: u64,
72 pub direction: Direction,
74 pub entry_price: f64,
76 pub exit_price: f64,
78 pub pnl: f64,
80 pub pnl_pct: f64,
82}
83
84#[derive(Debug, Clone)]
86pub struct BacktestMetrics {
87 pub total_return: f64,
89 pub annualized_return: f64,
91 pub sharpe: f64,
93 pub sortino: f64,
95 pub max_drawdown: f64,
97 pub calmar: f64,
99 pub win_rate: f64,
101 pub profit_factor: f64,
103 pub avg_trade_return: f64,
105 pub num_trades: usize,
107}
108
109#[derive(Debug, Clone)]
111pub struct BacktestResult {
112 pub equity_curve: Vec<f64>,
114 pub trades: Vec<CompletedTrade>,
116 pub metrics: BacktestMetrics,
118}
119
120pub struct BacktestEngine;
124
125impl BacktestEngine {
126 pub fn run(signals: Vec<EngineSignal>, config: EngineConfig) -> BacktestResult {
135 let bars = &config.data;
136 if bars.is_empty() {
137 return BacktestEngine::empty_result(config.initial_capital);
138 }
139
140 let n = bars.len();
141 let mut equity = config.initial_capital;
142 let mut equity_curve: Vec<f64> = Vec::with_capacity(n);
143 let mut completed_trades: Vec<CompletedTrade> = Vec::new();
144
145 let mut open_direction: Option<Direction> = None;
147 let mut open_entry_price: f64 = 0.0;
148 let mut open_entry_ts: u64 = 0;
149 let mut open_size: f64 = 0.0; let mut open_notional: f64 = 0.0;
151
152 let mut sorted_signals = signals;
158 sorted_signals.sort_by(|a, b| a.timestamp.cmp(&b.timestamp));
159
160 let mut sig_idx = 0;
161
162 for bar_i in 0..n {
163 let bar = &bars[bar_i];
164 let bar_open_ms = bar.ts_open.nanos() as u64 / 1_000_000;
165
166 while sig_idx < sorted_signals.len()
169 && sorted_signals[sig_idx].timestamp < bar_open_ms
170 {
171 let sig = &sorted_signals[sig_idx];
172 let fill_price_raw = bar.open.value().to_f64_or(bar.open.value());
173 let slippage_mult = config.slippage_bps / 10_000.0;
174
175 match sig.direction {
176 Direction::Long | Direction::Short => {
177 if let Some(existing_dir) = open_direction {
179 if existing_dir != sig.direction {
180 let exit_price = apply_slippage(
181 fill_price_raw,
182 slippage_mult,
183 existing_dir,
184 true, );
186 let commission = exit_price * open_size * config.commission;
187 let pnl = compute_pnl(
188 existing_dir,
189 open_entry_price,
190 exit_price,
191 open_size,
192 ) - commission;
193 equity += pnl;
194 let pnl_pct = if open_notional != 0.0 {
195 pnl / open_notional
196 } else {
197 0.0
198 };
199 completed_trades.push(CompletedTrade {
200 entry_ts: open_entry_ts,
201 exit_ts: bar_open_ms,
202 direction: existing_dir,
203 entry_price: open_entry_price,
204 exit_price,
205 pnl,
206 pnl_pct,
207 });
208 open_direction = None;
209 }
210 }
211
212 if open_direction.is_none() {
214 let entry_price = apply_slippage(
215 fill_price_raw,
216 slippage_mult,
217 sig.direction,
218 false, );
220 let size_capital = equity * config.capital_fraction * sig.strength;
221 let size = if entry_price > 0.0 {
222 size_capital / entry_price
223 } else {
224 0.0
225 };
226 let commission = entry_price * size * config.commission;
227 equity -= commission;
228 open_direction = Some(sig.direction);
229 open_entry_price = entry_price;
230 open_entry_ts = bar_open_ms;
231 open_size = size;
232 open_notional = entry_price * size;
233 }
234 }
235 Direction::Flat => {
236 if let Some(existing_dir) = open_direction {
238 let exit_price = apply_slippage(
239 fill_price_raw,
240 slippage_mult,
241 existing_dir,
242 true,
243 );
244 let commission = exit_price * open_size * config.commission;
245 let pnl = compute_pnl(
246 existing_dir,
247 open_entry_price,
248 exit_price,
249 open_size,
250 ) - commission;
251 equity += pnl;
252 let pnl_pct = if open_notional != 0.0 {
253 pnl / open_notional
254 } else {
255 0.0
256 };
257 completed_trades.push(CompletedTrade {
258 entry_ts: open_entry_ts,
259 exit_ts: bar_open_ms,
260 direction: existing_dir,
261 entry_price: open_entry_price,
262 exit_price,
263 pnl,
264 pnl_pct,
265 });
266 open_direction = None;
267 }
268 }
269 }
270 sig_idx += 1;
271 }
272
273 let close_f = bar.close.value().to_f64_or(bar.close.value());
275 let mtm_equity = if let Some(dir) = open_direction {
276 let unrealized = compute_pnl(dir, open_entry_price, close_f, open_size);
277 equity + unrealized
278 } else {
279 equity
280 };
281 equity_curve.push(mtm_equity.max(0.0));
282 }
283
284 if let Some(dir) = open_direction {
286 let last_bar = &bars[n - 1];
287 let exit_price = last_bar.close.value().to_f64_or(last_bar.close.value());
288 let commission = exit_price * open_size * config.commission;
289 let pnl = compute_pnl(dir, open_entry_price, exit_price, open_size) - commission;
290 equity += pnl;
291 let pnl_pct = if open_notional != 0.0 { pnl / open_notional } else { 0.0 };
292 let bar_ts = last_bar.ts_close.nanos() as u64 / 1_000_000;
293 completed_trades.push(CompletedTrade {
294 entry_ts: open_entry_ts,
295 exit_ts: bar_ts,
296 direction: dir,
297 entry_price: open_entry_price,
298 exit_price,
299 pnl,
300 pnl_pct,
301 });
302 if let Some(last) = equity_curve.last_mut() {
304 *last = equity.max(0.0);
305 }
306 }
307
308 let metrics =
309 compute_metrics(&equity_curve, &completed_trades, config.initial_capital);
310
311 BacktestResult { equity_curve, trades: completed_trades, metrics }
312 }
313
314 fn empty_result(_initial_capital: f64) -> BacktestResult {
315 BacktestResult {
316 equity_curve: vec![],
317 trades: vec![],
318 metrics: BacktestMetrics {
319 total_return: 0.0,
320 annualized_return: 0.0,
321 sharpe: 0.0,
322 sortino: 0.0,
323 max_drawdown: 0.0,
324 calmar: 0.0,
325 win_rate: 0.0,
326 profit_factor: 0.0,
327 avg_trade_return: 0.0,
328 num_trades: 0,
329 },
330 }
331 }
332}
333
334fn apply_slippage(price: f64, slippage_mult: f64, dir: Direction, closing: bool) -> f64 {
338 let adverse = match (dir, closing) {
339 (Direction::Long, false) => 1.0 + slippage_mult, (Direction::Long, true) => 1.0 - slippage_mult, (Direction::Short, false) => 1.0 - slippage_mult, (Direction::Short, true) => 1.0 + slippage_mult, _ => 1.0,
344 };
345 price * adverse
346}
347
348fn compute_pnl(dir: Direction, entry: f64, exit: f64, size: f64) -> f64 {
350 match dir {
351 Direction::Long => (exit - entry) * size,
352 Direction::Short => (entry - exit) * size,
353 Direction::Flat => 0.0,
354 }
355}
356
357trait ToF64OrDefault {
359 fn to_f64_or(&self, _fallback: Self) -> f64
360 where
361 Self: Sized;
362}
363
364impl ToF64OrDefault for rust_decimal::Decimal {
365 fn to_f64_or(&self, _fallback: Self) -> f64 {
366 use rust_decimal::prelude::ToPrimitive;
367 self.to_f64().unwrap_or(0.0)
368 }
369}
370
371fn compute_metrics(
373 equity_curve: &[f64],
374 trades: &[CompletedTrade],
375 initial_capital: f64,
376) -> BacktestMetrics {
377 let n = equity_curve.len();
378
379 let final_equity = equity_curve.last().copied().unwrap_or(initial_capital);
381 let total_return = if initial_capital > 0.0 {
382 (final_equity - initial_capital) / initial_capital
383 } else {
384 0.0
385 };
386
387 let years = n as f64 / 252.0;
389 let annualized_return = if years > 0.0 {
390 (1.0 + total_return).powf(1.0 / years) - 1.0
391 } else {
392 0.0
393 };
394
395 let mut daily_returns: Vec<f64> = Vec::with_capacity(n.saturating_sub(1));
397 for i in 1..n {
398 if equity_curve[i - 1] > 0.0 {
399 daily_returns.push((equity_curve[i] - equity_curve[i - 1]) / equity_curve[i - 1]);
400 }
401 }
402
403 let sharpe = compute_sharpe_f64(&daily_returns);
404 let sortino = compute_sortino_f64(&daily_returns);
405
406 let max_drawdown = compute_max_drawdown(equity_curve);
408
409 let calmar = if max_drawdown > 0.0 {
411 annualized_return / max_drawdown
412 } else {
413 0.0
414 };
415
416 let num_trades = trades.len();
418 let (win_rate, profit_factor, avg_trade_return) = if num_trades == 0 {
419 (0.0, 0.0, 0.0)
420 } else {
421 let wins = trades.iter().filter(|t| t.pnl > 0.0).count();
422 let gross_profit: f64 = trades.iter().filter(|t| t.pnl > 0.0).map(|t| t.pnl).sum();
423 let gross_loss: f64 =
424 trades.iter().filter(|t| t.pnl < 0.0).map(|t| t.pnl.abs()).sum();
425 let pf = if gross_loss > 0.0 { gross_profit / gross_loss } else { f64::INFINITY };
426 let avg_ret: f64 = trades.iter().map(|t| t.pnl_pct).sum::<f64>() / num_trades as f64;
427 (wins as f64 / num_trades as f64, pf, avg_ret)
428 };
429
430 BacktestMetrics {
431 total_return,
432 annualized_return,
433 sharpe,
434 sortino,
435 max_drawdown,
436 calmar,
437 win_rate,
438 profit_factor,
439 avg_trade_return,
440 num_trades,
441 }
442}
443
444fn compute_sharpe_f64(returns: &[f64]) -> f64 {
445 let n = returns.len();
446 if n < 2 {
447 return 0.0;
448 }
449 let mean = returns.iter().sum::<f64>() / n as f64;
450 let var = returns.iter().map(|r| (r - mean).powi(2)).sum::<f64>() / (n as f64 - 1.0);
451 let std_dev = var.sqrt();
452 if std_dev == 0.0 {
453 return 0.0;
454 }
455 (mean / std_dev) * 252.0_f64.sqrt()
456}
457
458fn compute_sortino_f64(returns: &[f64]) -> f64 {
459 let n = returns.len();
460 if n < 2 {
461 return 0.0;
462 }
463 let mean = returns.iter().sum::<f64>() / n as f64;
464 let downside_var = returns
465 .iter()
466 .map(|r| if *r < 0.0 { r.powi(2) } else { 0.0 })
467 .sum::<f64>()
468 / (n as f64 - 1.0);
469 let downside_dev = downside_var.sqrt();
470 if downside_dev == 0.0 {
471 return 0.0;
472 }
473 (mean / downside_dev) * 252.0_f64.sqrt()
474}
475
476fn compute_max_drawdown(equity_curve: &[f64]) -> f64 {
477 let mut peak = f64::NEG_INFINITY;
478 let mut max_dd = 0.0_f64;
479 for &e in equity_curve {
480 if e > peak {
481 peak = e;
482 }
483 if peak > 0.0 {
484 let dd = (peak - e) / peak;
485 if dd > max_dd {
486 max_dd = dd;
487 }
488 }
489 }
490 max_dd
491}
492
493#[cfg(test)]
496mod tests {
497 use super::*;
498 use crate::ohlcv::OhlcvBar;
499 use crate::types::{NanoTimestamp, Price, Quantity, Symbol};
500 use rust_decimal_macros::dec;
501
502 fn bar(open: f64, high: f64, low: f64, close: f64, ts_ms: u64) -> OhlcvBar {
503 let sym = Symbol::new("TEST").unwrap();
504 let open_p = Price::new(rust_decimal::Decimal::try_from(open).unwrap()).unwrap();
505 let high_p = Price::new(rust_decimal::Decimal::try_from(high).unwrap()).unwrap();
506 let low_p = Price::new(rust_decimal::Decimal::try_from(low).unwrap()).unwrap();
507 let close_p = Price::new(rust_decimal::Decimal::try_from(close).unwrap()).unwrap();
508 OhlcvBar {
509 symbol: sym,
510 open: open_p,
511 high: high_p,
512 low: low_p,
513 close: close_p,
514 volume: Quantity::new(dec!(100)).unwrap(),
515 ts_open: NanoTimestamp::new((ts_ms * 1_000_000) as i64),
516 ts_close: NanoTimestamp::new((ts_ms * 1_000_000 + 1_000_000) as i64),
517 tick_count: 1,
518 }
519 }
520
521 fn long_signal(ts_ms: u64) -> EngineSignal {
522 EngineSignal {
523 timestamp: ts_ms,
524 symbol: "TEST".to_string(),
525 direction: Direction::Long,
526 strength: 1.0,
527 }
528 }
529
530 fn flat_signal(ts_ms: u64) -> EngineSignal {
531 EngineSignal {
532 timestamp: ts_ms,
533 symbol: "TEST".to_string(),
534 direction: Direction::Flat,
535 strength: 1.0,
536 }
537 }
538
539 fn short_signal(ts_ms: u64) -> EngineSignal {
540 EngineSignal {
541 timestamp: ts_ms,
542 symbol: "TEST".to_string(),
543 direction: Direction::Short,
544 strength: 1.0,
545 }
546 }
547
548 fn make_config(bars: Vec<OhlcvBar>) -> EngineConfig {
549 EngineConfig {
550 initial_capital: 10_000.0,
551 commission: 0.001,
552 slippage_bps: 5.0,
553 data: bars,
554 capital_fraction: 0.1,
555 }
556 }
557
558 #[test]
559 fn test_empty_bars_returns_empty_result() {
560 let result = BacktestEngine::run(vec![], make_config(vec![]));
561 assert!(result.equity_curve.is_empty());
562 assert_eq!(result.trades.len(), 0);
563 assert_eq!(result.metrics.num_trades, 0);
564 }
565
566 #[test]
567 fn test_no_signals_equity_equals_initial() {
568 let bars: Vec<OhlcvBar> = (0..5)
569 .map(|i| bar(100.0, 102.0, 99.0, 101.0, 1000 + i * 100))
570 .collect();
571 let config = make_config(bars);
572 let result = BacktestEngine::run(vec![], config);
573 for &eq in &result.equity_curve {
575 assert!((eq - 10_000.0).abs() < 1e-6, "eq={}", eq);
576 }
577 }
578
579 #[test]
580 fn test_long_trade_profitable() {
581 let bars = vec![
584 bar(100.0, 105.0, 99.0, 102.0, 1000),
585 bar(100.0, 125.0, 99.0, 120.0, 2000),
586 bar(120.0, 130.0, 118.0, 125.0, 3000),
587 ];
588 let signals = vec![long_signal(900)]; let config = make_config(bars);
590 let result = BacktestEngine::run(signals, config);
591 assert!(!result.equity_curve.is_empty());
593 }
594
595 #[test]
596 fn test_flat_signal_closes_position() {
597 let bars = vec![
598 bar(100.0, 105.0, 99.0, 102.0, 1000),
599 bar(102.0, 110.0, 100.0, 108.0, 2000),
600 bar(108.0, 112.0, 106.0, 110.0, 3000),
601 ];
602 let signals = vec![
603 long_signal(900), flat_signal(1500), ];
606 let config = make_config(bars);
607 let result = BacktestEngine::run(signals, config);
608 assert_eq!(result.trades.len(), 1);
609 assert_eq!(result.trades[0].direction, Direction::Long);
610 }
611
612 #[test]
613 fn test_short_trade_created() {
614 let bars = vec![
615 bar(100.0, 105.0, 99.0, 99.0, 1000),
616 bar(99.0, 100.0, 90.0, 90.0, 2000),
617 bar(90.0, 91.0, 80.0, 82.0, 3000),
618 ];
619 let signals = vec![short_signal(900)];
620 let config = make_config(bars);
621 let result = BacktestEngine::run(signals, config);
622 assert!(!result.equity_curve.is_empty());
623 }
624
625 #[test]
626 fn test_opposite_signal_closes_then_opens() {
627 let bars = vec![
628 bar(100.0, 105.0, 99.0, 102.0, 1000),
629 bar(102.0, 110.0, 100.0, 108.0, 2000),
630 bar(108.0, 112.0, 106.0, 110.0, 3000),
631 bar(110.0, 115.0, 108.0, 112.0, 4000),
632 ];
633 let signals = vec![
634 long_signal(900), short_signal(1500), ];
637 let config = make_config(bars);
638 let result = BacktestEngine::run(signals, config);
639 assert!(result.trades.len() >= 1);
641 assert_eq!(result.trades[0].direction, Direction::Long);
642 }
643
644 #[test]
645 fn test_commission_reduces_equity() {
646 let bars = vec![
647 bar(100.0, 100.0, 100.0, 100.0, 1000),
648 bar(100.0, 100.0, 100.0, 100.0, 2000),
649 ];
650 let signals = vec![long_signal(900)];
651 let mut config = make_config(bars);
652 config.commission = 0.01; config.slippage_bps = 0.0;
654 let result = BacktestEngine::run(signals, config);
655 let final_eq = result.equity_curve.last().copied().unwrap_or(10_000.0);
657 assert!(final_eq < 10_000.0, "Commission should reduce equity: {}", final_eq);
658 }
659
660 #[test]
661 fn test_slippage_applied_to_long_open() {
662 let bars = vec![
664 bar(100.0, 100.0, 100.0, 100.0, 1000),
665 bar(100.0, 100.0, 100.0, 100.0, 2000),
666 ];
667 let signals = vec![long_signal(900)];
668 let mut config = make_config(bars);
669 config.commission = 0.0;
670 config.slippage_bps = 100.0; let result = BacktestEngine::run(signals, config);
672 let final_eq = result.equity_curve.last().copied().unwrap_or(10_000.0);
674 assert!(final_eq <= 10_000.0, "Slippage should reduce equity: {}", final_eq);
675 }
676
677 #[test]
678 fn test_equity_curve_length_equals_bars() {
679 let bars: Vec<OhlcvBar> = (0..10)
680 .map(|i| bar(100.0, 105.0, 99.0, 102.0, 1000 + i * 100))
681 .collect();
682 let config = make_config(bars.clone());
683 let result = BacktestEngine::run(vec![], config);
684 assert_eq!(result.equity_curve.len(), bars.len());
685 }
686
687 #[test]
688 fn test_metrics_total_return_positive_for_winning_trade() {
689 let bars = vec![
691 bar(100.0, 100.0, 100.0, 100.0, 1000),
692 bar(200.0, 200.0, 200.0, 200.0, 2000),
693 bar(200.0, 200.0, 200.0, 200.0, 3000),
694 ];
695 let signals = vec![long_signal(900)];
696 let mut config = make_config(bars);
697 config.commission = 0.0;
698 config.slippage_bps = 0.0;
699 config.capital_fraction = 1.0;
700 let result = BacktestEngine::run(signals, config);
701 assert!(result.metrics.total_return > 0.0, "tr={}", result.metrics.total_return);
702 }
703
704 #[test]
705 fn test_metrics_win_rate_one_winner() {
706 let bars = vec![
707 bar(100.0, 100.0, 100.0, 100.0, 1000),
708 bar(200.0, 200.0, 200.0, 200.0, 2000),
709 bar(200.0, 200.0, 200.0, 200.0, 3000),
710 ];
711 let signals = vec![long_signal(900), flat_signal(1500)];
712 let mut config = make_config(bars);
713 config.commission = 0.0;
714 config.slippage_bps = 0.0;
715 config.capital_fraction = 1.0;
716 let result = BacktestEngine::run(signals, config);
717 assert_eq!(result.metrics.win_rate, 1.0);
718 }
719
720 #[test]
721 fn test_max_drawdown_computed() {
722 let bars = vec![
724 bar(100.0, 100.0, 100.0, 200.0, 1000),
725 bar(200.0, 200.0, 200.0, 50.0, 2000),
726 bar(50.0, 50.0, 50.0, 50.0, 3000),
727 ];
728 let result = BacktestEngine::run(vec![], make_config(bars));
729 assert!(result.metrics.max_drawdown >= 0.0);
733 }
734
735 #[test]
736 fn test_profit_factor_above_one_for_winning_trade() {
737 let bars = vec![
738 bar(100.0, 100.0, 100.0, 100.0, 1000),
739 bar(110.0, 110.0, 110.0, 110.0, 2000),
740 bar(110.0, 110.0, 110.0, 110.0, 3000),
741 ];
742 let signals = vec![long_signal(900), flat_signal(1500)];
743 let mut config = make_config(bars);
744 config.commission = 0.0;
745 config.slippage_bps = 0.0;
746 let result = BacktestEngine::run(signals, config);
747 assert!(result.metrics.profit_factor > 1.0 || result.metrics.profit_factor.is_infinite());
749 }
750
751 #[test]
752 fn test_num_trades_matches_completed_trades() {
753 let bars = vec![
754 bar(100.0, 100.0, 100.0, 100.0, 1000),
755 bar(110.0, 110.0, 110.0, 110.0, 2000),
756 bar(110.0, 110.0, 110.0, 115.0, 3000),
757 bar(115.0, 115.0, 115.0, 120.0, 4000),
758 ];
759 let signals = vec![
760 long_signal(900),
761 flat_signal(1500),
762 short_signal(2500),
763 flat_signal(3500),
764 ];
765 let config = make_config(bars);
766 let result = BacktestEngine::run(signals, config);
767 assert_eq!(result.metrics.num_trades, result.trades.len());
768 }
769
770 #[test]
771 fn test_strength_scales_position_size() {
772 let bars = vec![
773 bar(100.0, 100.0, 100.0, 100.0, 1000),
774 bar(200.0, 200.0, 200.0, 200.0, 2000),
775 ];
776 let sig_half = EngineSignal {
777 timestamp: 900,
778 symbol: "TEST".to_string(),
779 direction: Direction::Long,
780 strength: 0.5,
781 };
782 let sig_full = EngineSignal {
783 timestamp: 900,
784 symbol: "TEST".to_string(),
785 direction: Direction::Long,
786 strength: 1.0,
787 };
788 let mut config1 = make_config(bars.clone());
789 config1.commission = 0.0;
790 config1.slippage_bps = 0.0;
791 let mut config2 = make_config(bars);
792 config2.commission = 0.0;
793 config2.slippage_bps = 0.0;
794 let r1 = BacktestEngine::run(vec![sig_half], config1);
795 let r2 = BacktestEngine::run(vec![sig_full], config2);
796 let ret1 = r1.metrics.total_return;
797 let ret2 = r2.metrics.total_return;
798 assert!(ret2 > ret1, "ret2={} ret1={}", ret2, ret1);
800 }
801
802 #[test]
803 fn test_completed_trade_fields() {
804 let bars = vec![
805 bar(100.0, 100.0, 100.0, 100.0, 1000),
806 bar(110.0, 110.0, 110.0, 110.0, 2000),
807 bar(110.0, 110.0, 110.0, 115.0, 3000),
808 ];
809 let signals = vec![long_signal(900), flat_signal(1500)];
810 let mut config = make_config(bars);
811 config.commission = 0.0;
812 config.slippage_bps = 0.0;
813 let result = BacktestEngine::run(signals, config);
814 let t = &result.trades[0];
815 assert_eq!(t.direction, Direction::Long);
816 assert!(t.entry_price > 0.0);
817 assert!(t.exit_price > 0.0);
818 assert!(t.exit_ts > t.entry_ts || t.exit_ts == t.entry_ts);
819 }
820}