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