Skip to main content

fin_primitives/backtest/
mod.rs

1//! Bar-by-bar backtester, the `Strategy` trait, and a walk-forward optimizer.
2//!
3//! ## Responsibility
4//! Provides a bar-by-bar backtester, a `Strategy` trait for signal generation,
5//! equity curve tracking, and a full walk-forward optimizer with grid search.
6//!
7//! ## Sub-modules
8//! - [`walk_forward`]: grid-search walk-forward optimizer with per-period OOS evaluation
9//!
10//! ## Guarantees
11//! - Bars are processed in the order supplied; no look-ahead
12//! - `BacktestResult::max_drawdown` is always in `[0, 1]`
13//! - Commission is deducted from cash on every fill
14//!
15//! ## NOT Responsible For
16//! - Live order routing
17//! - Slippage models beyond the commission rate
18
19pub 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// ─── Config / Result ──────────────────────────────────────────────────────────
38
39/// Configuration for a single backtest run.
40#[derive(Debug, Clone)]
41#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
42pub struct BacktestConfig {
43    /// Starting cash balance.
44    pub initial_capital: Decimal,
45    /// Commission rate as a fraction of notional (e.g. `dec!(0.001)` = 0.1%).
46    pub commission_rate: Decimal,
47}
48
49impl BacktestConfig {
50    /// Creates a new `BacktestConfig`.
51    ///
52    /// # Errors
53    /// Returns [`FinError::InvalidInput`] if `initial_capital` or `commission_rate` are negative.
54    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/// Summary statistics produced after a completed backtest run.
70#[derive(Debug, Clone)]
71#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
72pub struct BacktestResult {
73    /// Total return over the period: `(final_equity - initial_capital) / initial_capital`.
74    pub total_return: Decimal,
75    /// Annualised Sharpe ratio (assuming 252 trading days; `NaN`-free — returns 0 if std == 0).
76    pub sharpe_ratio: Decimal,
77    /// Maximum peak-to-trough equity drawdown as a fraction, e.g. `dec!(0.15)` = 15%.
78    pub max_drawdown: Decimal,
79    /// Fraction of closed trades that were profitable.
80    pub win_rate: Decimal,
81    /// Total number of trades (fills) executed.
82    pub trade_count: u64,
83    /// Final equity value at the end of the period.
84    pub final_equity: Decimal,
85    /// Equity curve sampled once per bar.
86    pub equity_curve: Vec<Decimal>,
87}
88
89// ─── Signal ───────────────────────────────────────────────────────────────────
90
91/// Direction of a trading signal.
92#[derive(Debug, Clone, Copy, PartialEq, Eq)]
93pub enum SignalDirection {
94    /// Enter a long position.
95    Buy,
96    /// Enter a short position / exit a long position.
97    Sell,
98    /// Do nothing.
99    Hold,
100}
101
102/// A trading signal produced by a `Strategy` on each bar.
103#[derive(Debug, Clone)]
104pub struct Signal {
105    /// Desired direction.
106    pub direction: SignalDirection,
107    /// Number of units to trade (must be non-negative).
108    pub quantity: Decimal,
109}
110
111impl Signal {
112    /// Creates a new signal.
113    pub fn new(direction: SignalDirection, quantity: Decimal) -> Self {
114        Self { direction, quantity }
115    }
116
117    /// Convenience constructor for a hold signal with zero quantity.
118    pub fn hold() -> Self {
119        Self::new(SignalDirection::Hold, Decimal::ZERO)
120    }
121}
122
123// ─── Strategy trait ───────────────────────────────────────────────────────────
124
125/// User-supplied strategy that generates trade signals from bars.
126///
127/// Implement this trait and pass a `&mut dyn Strategy` to [`Backtester::run`].
128pub trait Strategy: Send {
129    /// Called once per bar in chronological order.
130    ///
131    /// Return `None` to skip trading on this bar.
132    fn on_bar(&mut self, bar: &OhlcvBar) -> Option<Signal>;
133}
134
135// ─── Backtester ───────────────────────────────────────────────────────────────
136
137/// Bar-by-bar backtester.
138///
139/// Processes OHLCV bars in the supplied order, routes signals from `Strategy::on_bar`
140/// into a `PositionLedger`, and records the equity curve.
141pub struct Backtester {
142    config: BacktestConfig,
143}
144
145impl Backtester {
146    /// Creates a new backtester with the given config.
147    pub fn new(config: BacktestConfig) -> Self {
148        Self { config }
149    }
150
151    /// Runs the backtest over `bars` using `strategy`.
152    ///
153    /// # Errors
154    /// - [`FinError::InvalidInput`] if `bars` is empty.
155    /// - Propagates [`FinError`] from position accounting.
156    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        // Track realized P&L by comparing ledger realized_pnl_total before/after each fill
174        let mut winning_trades: u64 = 0;
175        let mut total_closed: u64 = 0;
176
177        for bar in bars {
178            // Ask strategy for a signal
179            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                    // Capture realized P&L before and after to detect a profitable trade.
201                    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            // Mark-to-market equity
217            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            // Drawdown
225            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            // Daily return
236            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
272/// Computes the annualised Sharpe ratio from a slice of per-bar returns.
273///
274/// Assumes 252 trading days per year. Returns zero if the slice is empty or
275/// the standard deviation is zero (no variance in returns).
276fn 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    // Mean
285    let sum: Decimal = returns.iter().sum();
286    let mean_f = sum.to_f64().unwrap_or(0.0) / n as f64;
287
288    // Variance
289    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// Walk-forward types and optimizer are fully implemented in the walk_forward
310// sub-module and re-exported above via `pub use walk_forward::{...}`.
311
312// ─── tests ────────────────────────────────────────────────────────────────────
313
314#[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    /// Buy-and-hold: buys 1 unit on the first bar, holds.
337    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        // Bought 1 unit @ 100 from 10 000 cash; final close = 109.
362        // equity = (10 000 - 100) + unrealized_pnl(109) = 9 900 + 9 = 9 909.
363        // Assert equity grew relative to first bar (9 900) and trade count is 1.
364        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        // Hold strategy: no trades, cash never changes, drawdown must be zero.
378        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        // All returns identical → zero stddev → sharpe = 0
440        let returns = vec![dec!(0.01); 10];
441        // sharpe with zero variance returns Decimal::ZERO
442        let s = compute_sharpe(&returns);
443        // If all returns are equal the sample variance is zero, so sharpe = 0
444        assert_eq!(s, Decimal::ZERO);
445    }
446}