Skip to main content

fin_primitives/scenario/
mod.rs

1//! Risk scenario backtesting: replays historical bars through risk rules.
2//!
3//! ## Responsibility
4//! Risk scenario backtesting: replays a sequence of historical OHLCV bars through a
5//! user-provided risk rule and reports how many times the rule would have triggered,
6//! along with the maximum drawdown observed during the scenario.
7//!
8//! ## Guarantees
9//! - All arithmetic uses `rust_decimal::Decimal`
10//! - `ScenarioBacktester::run` never panics; all results are returned in a typed report
11//! - Equity is simulated as bar-close price by default (caller-supplied equity function)
12//!
13//! ## NOT Responsible For
14//! - Realistic fill simulation (see `position` module)
15//! - Multi-asset scenarios
16
17use crate::ohlcv::OhlcvBar;
18use crate::risk::{DrawdownTracker, RiskBreach, RiskRule};
19use rust_decimal::Decimal;
20
21/// Summary report produced by [`ScenarioBacktester::run`].
22#[derive(Debug, Clone)]
23pub struct ScenarioReport {
24    /// Total number of bars processed.
25    pub bars_processed: usize,
26    /// Number of bars on which at least one rule triggered.
27    pub trigger_count: usize,
28    /// All individual breach events, one entry per bar that triggered.
29    pub breaches: Vec<BarBreach>,
30    /// Maximum drawdown (%) observed at any point during the scenario.
31    pub max_drawdown_pct: Decimal,
32    /// Starting equity (first bar's simulated equity).
33    pub start_equity: Decimal,
34    /// Ending equity (last bar's simulated equity).
35    pub end_equity: Decimal,
36    /// Total equity return: `(end - start) / start * 100` (percent).
37    ///
38    /// Returns `None` when `start_equity == 0`.
39    pub total_return_pct: Option<Decimal>,
40}
41
42/// A breach event at a specific bar index.
43#[derive(Debug, Clone)]
44pub struct BarBreach {
45    /// Zero-based bar index.
46    pub bar_index: usize,
47    /// Risk breaches that fired on this bar.
48    pub breaches: Vec<RiskBreach>,
49    /// Simulated equity at this bar.
50    pub equity: Decimal,
51    /// Drawdown percentage at this bar.
52    pub drawdown_pct: Decimal,
53}
54
55/// Replays historical OHLCV bars through a set of `RiskRule`s.
56///
57/// The caller supplies:
58/// 1. A slice of [`OhlcvBar`] bars (historical data).
59/// 2. One or more [`RiskRule`] implementations.
60/// 3. An equity function `F: Fn(&OhlcvBar) -> Decimal` that maps each bar to a
61///    simulated equity value (e.g. close price, portfolio NAV).
62///
63/// # Example
64/// ```rust
65/// use fin_primitives::scenario::ScenarioBacktester;
66/// use fin_primitives::risk::{MaxDrawdownRule, DrawdownTracker};
67/// use fin_primitives::ohlcv::OhlcvBar;
68/// use fin_primitives::types::{Symbol, Price, Quantity, NanoTimestamp};
69/// use rust_decimal_macros::dec;
70///
71/// let sym = Symbol::new("SPY").unwrap();
72/// let ts = NanoTimestamp::new(0);
73/// let bars: Vec<OhlcvBar> = (0..10).map(|i| {
74///     let close_val = dec!(100) - rust_decimal::Decimal::from(i) * dec!(2);
75///     let close = Price::new(close_val).unwrap();
76///     let open = Price::new(dec!(102)).unwrap();
77///     let high = Price::new(dec!(103)).unwrap();
78///     let low = Price::new(close_val).unwrap();
79///     OhlcvBar::new(sym.clone(), open, high, low, close,
80///                   Quantity::new(dec!(1000)).unwrap(), ts, ts, 100).unwrap()
81/// }).collect();
82///
83/// let rule = MaxDrawdownRule { threshold_pct: dec!(10) };
84/// let report = ScenarioBacktester::new(bars)
85///     .add_rule(Box::new(rule))
86///     .run(|bar| bar.close.value());
87///
88/// assert_eq!(report.bars_processed, 10);
89/// ```
90pub struct ScenarioBacktester {
91    bars: Vec<OhlcvBar>,
92    rules: Vec<Box<dyn RiskRule>>,
93}
94
95impl ScenarioBacktester {
96    /// Creates a new `ScenarioBacktester` with the given historical bars.
97    pub fn new(bars: Vec<OhlcvBar>) -> Self {
98        Self { bars, rules: Vec::new() }
99    }
100
101    /// Adds a risk rule to the set evaluated at each bar.
102    ///
103    /// Rules are evaluated independently; all triggered rules produce breach events.
104    pub fn add_rule(mut self, rule: Box<dyn RiskRule>) -> Self {
105        self.rules.push(rule);
106        self
107    }
108
109    /// Runs the scenario, returning a [`ScenarioReport`].
110    ///
111    /// `equity_fn` maps each bar to a simulated equity value.
112    /// The most common choices are `|bar| bar.close.value()` (close-based equity)
113    /// or a portfolio NAV calculation that uses positions from the caller.
114    pub fn run<F>(&self, equity_fn: F) -> ScenarioReport
115    where
116        F: Fn(&OhlcvBar) -> Decimal,
117    {
118        if self.bars.is_empty() {
119            return ScenarioReport {
120                bars_processed: 0,
121                trigger_count: 0,
122                breaches: vec![],
123                max_drawdown_pct: Decimal::ZERO,
124                start_equity: Decimal::ZERO,
125                end_equity: Decimal::ZERO,
126                total_return_pct: None,
127            };
128        }
129
130        let first_equity = equity_fn(&self.bars[0]);
131        let mut tracker = DrawdownTracker::new(first_equity);
132        let mut all_breaches: Vec<BarBreach> = Vec::new();
133        let mut trigger_count = 0usize;
134        let mut last_equity = first_equity;
135
136        for (i, bar) in self.bars.iter().enumerate() {
137            let equity = equity_fn(bar);
138            tracker.update(equity);
139            let dd_pct = tracker.current_drawdown_pct();
140            last_equity = equity;
141
142            let bar_breaches: Vec<RiskBreach> = self
143                .rules
144                .iter()
145                .filter_map(|rule| rule.check(equity, dd_pct))
146                .collect();
147
148            if !bar_breaches.is_empty() {
149                trigger_count += 1;
150                all_breaches.push(BarBreach {
151                    bar_index: i,
152                    breaches: bar_breaches,
153                    equity,
154                    drawdown_pct: dd_pct,
155                });
156            }
157        }
158
159        let max_dd = tracker.worst_drawdown_pct();
160        let total_return_pct = if first_equity.is_zero() {
161            None
162        } else {
163            Some((last_equity - first_equity) / first_equity * Decimal::ONE_HUNDRED)
164        };
165
166        ScenarioReport {
167            bars_processed: self.bars.len(),
168            trigger_count,
169            breaches: all_breaches,
170            max_drawdown_pct: max_dd,
171            start_equity: first_equity,
172            end_equity: last_equity,
173            total_return_pct,
174        }
175    }
176}
177
178#[cfg(test)]
179mod tests {
180    use super::*;
181    use crate::risk::MaxDrawdownRule;
182    use crate::types::{NanoTimestamp, Price, Quantity, Symbol};
183    use rust_decimal_macros::dec;
184
185    fn sym() -> Symbol {
186        Symbol::new("SPY").unwrap()
187    }
188
189    fn ts() -> NanoTimestamp {
190        NanoTimestamp::new(0)
191    }
192
193    fn make_bar(close: rust_decimal::Decimal) -> OhlcvBar {
194        let p = Price::new(close).unwrap();
195        let high = Price::new(close + dec!(1)).unwrap();
196        OhlcvBar::new(
197            sym(),
198            p,
199            high,
200            p,
201            p,
202            Quantity::new(dec!(1000)).unwrap(),
203            ts(),
204            ts(),
205            10,
206        )
207        .unwrap()
208    }
209
210    #[test]
211    fn test_no_triggers_when_equity_rises() {
212        let bars: Vec<_> = (1..=10).map(|i| make_bar(dec!(100) + rust_decimal::Decimal::from(i))).collect();
213        let rule = MaxDrawdownRule { threshold_pct: dec!(5) };
214        let report = ScenarioBacktester::new(bars).add_rule(Box::new(rule)).run(|bar| bar.close.value());
215        assert_eq!(report.bars_processed, 10);
216        assert_eq!(report.trigger_count, 0);
217        assert_eq!(report.max_drawdown_pct, Decimal::ZERO);
218    }
219
220    #[test]
221    fn test_triggers_when_drawdown_exceeds_threshold() {
222        // Start at 100, drop to 80 (20% drawdown), threshold is 10%
223        let closes = [
224            dec!(100), dec!(99), dec!(95), dec!(90), dec!(85), dec!(80),
225        ];
226        let bars: Vec<_> = closes.iter().map(|&c| make_bar(c)).collect();
227        let rule = MaxDrawdownRule { threshold_pct: dec!(10) };
228        let report = ScenarioBacktester::new(bars).add_rule(Box::new(rule)).run(|bar| bar.close.value());
229        assert!(report.trigger_count > 0, "expected at least one trigger");
230        assert!(report.max_drawdown_pct > dec!(10));
231    }
232
233    #[test]
234    fn test_empty_bars_returns_zero_report() {
235        let report = ScenarioBacktester::new(vec![]).run(|bar| bar.close.value());
236        assert_eq!(report.bars_processed, 0);
237        assert_eq!(report.trigger_count, 0);
238        assert!(report.total_return_pct.is_none());
239    }
240
241    #[test]
242    fn test_total_return_pct_computed() {
243        let bars = vec![make_bar(dec!(100)), make_bar(dec!(110))];
244        let report = ScenarioBacktester::new(bars).run(|bar| bar.close.value());
245        // (110-100)/100*100 = 10%
246        assert_eq!(report.total_return_pct.unwrap(), dec!(10));
247    }
248
249    #[test]
250    fn test_multiple_rules_both_can_fire() {
251        let closes = [dec!(100), dec!(50)]; // 50% drawdown
252        let bars: Vec<_> = closes.iter().map(|&c| make_bar(c)).collect();
253        let rule1 = MaxDrawdownRule { threshold_pct: dec!(10) };
254        let rule2 = MaxDrawdownRule { threshold_pct: dec!(20) };
255        let report = ScenarioBacktester::new(bars)
256            .add_rule(Box::new(rule1))
257            .add_rule(Box::new(rule2))
258            .run(|bar| bar.close.value());
259        // Both rules should fire on bar index 1 (50% drawdown)
260        let bar1 = report.breaches.iter().find(|b| b.bar_index == 1).unwrap();
261        assert_eq!(bar1.breaches.len(), 2);
262    }
263
264    #[test]
265    fn test_max_drawdown_tracked() {
266        let closes = [dec!(200), dec!(180), dec!(160), dec!(190), dec!(210)];
267        let bars: Vec<_> = closes.iter().map(|&c| make_bar(c)).collect();
268        let report = ScenarioBacktester::new(bars).run(|bar| bar.close.value());
269        // Peak 200, trough 160: 20% drawdown
270        assert_eq!(report.max_drawdown_pct, dec!(20));
271    }
272
273    // ── ScenarioEngine ────────────────────────────────────────────────────
274
275    #[test]
276    fn test_apply_absolute_shift() {
277        let engine = ScenarioEngine;
278        let shocked = engine.apply_shock(100.0, &ShockType::AbsoluteShift(-30.0));
279        assert!((shocked - 70.0).abs() < 1e-9);
280    }
281
282    #[test]
283    fn test_apply_relative_shift() {
284        let engine = ScenarioEngine;
285        let shocked = engine.apply_shock(100.0, &ShockType::RelativeShift(-0.20));
286        assert!((shocked - 80.0).abs() < 1e-9);
287    }
288
289    #[test]
290    fn test_apply_volatility_scaling() {
291        let engine = ScenarioEngine;
292        let shocked = engine.apply_shock(100.0, &ShockType::VolatilityScaling(1.5));
293        // VolatilityScaling scales price by the factor
294        assert!((shocked - 150.0).abs() < 1e-9);
295    }
296
297    #[test]
298    fn test_apply_correlation_breakdown() {
299        let engine = ScenarioEngine;
300        // CorrelationBreakdown shifts by the factor as an absolute amount
301        let shocked = engine.apply_shock(100.0, &ShockType::CorrelationBreakdown(10.0));
302        assert!((shocked - 110.0).abs() < 1e-9);
303    }
304
305    #[test]
306    fn test_run_scenario_equity_crash() {
307        use std::collections::HashMap;
308        let mut portfolio: HashMap<String, f64> = HashMap::new();
309        portfolio.insert("equity".to_owned(), 100.0);
310        portfolio.insert("vol".to_owned(), 20.0);
311        let s = Scenario::equity_crash();
312        let engine = ScenarioEngine;
313        let shocked = engine.run_scenario(&portfolio, &s);
314        // equity should drop to ~70
315        let eq = shocked["equity"];
316        assert!(eq < 100.0, "equity should drop: {eq}");
317    }
318
319    #[test]
320    fn test_scenario_pnl_loss() {
321        use std::collections::HashMap;
322        let mut original: HashMap<String, f64> = HashMap::new();
323        original.insert("equity".to_owned(), 100.0);
324        let mut shocked: HashMap<String, f64> = HashMap::new();
325        shocked.insert("equity".to_owned(), 70.0);
326        let mut positions: HashMap<String, f64> = HashMap::new();
327        positions.insert("equity".to_owned(), 10.0);
328        let engine = ScenarioEngine;
329        let pnl = engine.scenario_pnl(&original, &shocked, &positions);
330        // 10 units * (70-100) = -300
331        assert!((pnl - (-300.0)).abs() < 1e-9, "pnl={pnl}");
332    }
333
334    #[test]
335    fn test_worst_case_scenario() {
336        use std::collections::HashMap;
337        let mut portfolio: HashMap<String, f64> = HashMap::new();
338        portfolio.insert("equity".to_owned(), 100.0);
339        let mut positions: HashMap<String, f64> = HashMap::new();
340        positions.insert("equity".to_owned(), 1.0);
341        let scenarios = vec![
342            Scenario::equity_crash(),
343            Scenario::rate_shock(),
344        ];
345        let engine = ScenarioEngine;
346        let (worst, pnl) = engine.worst_case(&portfolio, &scenarios, &positions);
347        assert!(pnl <= 0.0 || pnl.is_finite());
348        assert!(!worst.name.is_empty());
349    }
350
351    #[test]
352    fn test_built_in_scenarios_valid() {
353        assert!(!Scenario::equity_crash().shocks.is_empty());
354        assert!(!Scenario::credit_crisis().shocks.is_empty());
355        assert!(!Scenario::rate_shock().shocks.is_empty());
356        assert!(!Scenario::fx_devaluation().shocks.is_empty());
357    }
358}
359
360// ─────────────────────────────────────────
361//  Scenario analysis and stress-testing framework
362// ─────────────────────────────────────────
363
364use std::collections::HashMap;
365
366/// Type of shock applied to an asset's price.
367///
368/// - `AbsoluteShift(delta)`: adds `delta` directly to the price.
369/// - `RelativeShift(frac)`: multiplies price by `(1 + frac)`.
370/// - `VolatilityScaling(factor)`: scales price by `factor` (models vol regime shift).
371/// - `CorrelationBreakdown(delta)`: adds `delta` to price (models spread/basis blow-out).
372#[derive(Debug, Clone)]
373pub enum ShockType {
374    /// Add an absolute amount to the price (can be negative).
375    AbsoluteShift(f64),
376    /// Multiply price by `(1 + fraction)` (e.g. -0.30 = -30%).
377    RelativeShift(f64),
378    /// Scale price by `factor` (e.g. 1.5 = +50% vol-driven move).
379    VolatilityScaling(f64),
380    /// Add `delta` to price (models correlation breakdown / basis widening).
381    CorrelationBreakdown(f64),
382}
383
384/// A shock applied to a specific asset.
385#[derive(Debug, Clone)]
386pub struct AssetShock {
387    /// Identifier for the asset being shocked.
388    pub asset_id: String,
389    /// The type and magnitude of the shock.
390    pub shock: ShockType,
391}
392
393/// A named stress scenario containing a set of asset shocks.
394#[derive(Debug, Clone)]
395pub struct Scenario {
396    /// Short human-readable name (e.g. "equity_crash").
397    pub name: String,
398    /// Longer narrative description.
399    pub description: String,
400    /// Shocks applied to individual assets.
401    pub shocks: Vec<AssetShock>,
402    /// Subjective or historical probability of this scenario occurring.
403    pub probability: f64,
404}
405
406impl Scenario {
407    /// 2008-style equity crash: equities down 30%, volatility up 50%.
408    pub fn equity_crash() -> Self {
409        Self {
410            name: "equity_crash".to_owned(),
411            description: "2008-style equity market crash: equities -30%, implied vol +50%."
412                .to_owned(),
413            probability: 0.05,
414            shocks: vec![
415                AssetShock {
416                    asset_id: "equity".to_owned(),
417                    shock: ShockType::RelativeShift(-0.30),
418                },
419                AssetShock {
420                    asset_id: "vol".to_owned(),
421                    shock: ShockType::RelativeShift(0.50),
422                },
423            ],
424        }
425    }
426
427    /// Credit crisis: credit spreads blow out, IG credit down 20%.
428    pub fn credit_crisis() -> Self {
429        Self {
430            name: "credit_crisis".to_owned(),
431            description: "Credit crisis: IG credit -20%, HY spreads widen by +500 bps.".to_owned(),
432            probability: 0.03,
433            shocks: vec![
434                AssetShock {
435                    asset_id: "ig_credit".to_owned(),
436                    shock: ShockType::RelativeShift(-0.20),
437                },
438                AssetShock {
439                    asset_id: "hy_spread".to_owned(),
440                    shock: ShockType::AbsoluteShift(5.0),
441                },
442            ],
443        }
444    }
445
446    /// Rate shock: parallel upward shift of +200 bps across the yield curve.
447    pub fn rate_shock() -> Self {
448        Self {
449            name: "rate_shock".to_owned(),
450            description: "Sudden 200 bps rate hike across the yield curve.".to_owned(),
451            probability: 0.04,
452            shocks: vec![AssetShock {
453                asset_id: "rates".to_owned(),
454                shock: ShockType::AbsoluteShift(2.0),
455            }],
456        }
457    }
458
459    /// EM FX devaluation: emerging-market currencies depreciate 20%.
460    pub fn fx_devaluation() -> Self {
461        Self {
462            name: "fx_devaluation".to_owned(),
463            description: "EM FX devaluation: EM currencies -20% vs USD.".to_owned(),
464            probability: 0.06,
465            shocks: vec![AssetShock {
466                asset_id: "em_fx".to_owned(),
467                shock: ShockType::RelativeShift(-0.20),
468            }],
469        }
470    }
471}
472
473/// Engine that applies scenarios to portfolio prices and computes P&L impact.
474///
475/// # Example
476/// ```rust
477/// use std::collections::HashMap;
478/// use fin_primitives::scenario::{ScenarioEngine, ShockType};
479///
480/// let engine = ScenarioEngine;
481/// let shocked = engine.apply_shock(100.0, &ShockType::RelativeShift(-0.30));
482/// assert!((shocked - 70.0).abs() < 1e-9);
483/// ```
484pub struct ScenarioEngine;
485
486impl ScenarioEngine {
487    /// Apply a single shock to a base price and return the shocked price.
488    pub fn apply_shock(&self, price: f64, shock: &ShockType) -> f64 {
489        match shock {
490            ShockType::AbsoluteShift(delta) => price + delta,
491            ShockType::RelativeShift(frac) => price * (1.0 + frac),
492            ShockType::VolatilityScaling(factor) => price * factor,
493            ShockType::CorrelationBreakdown(delta) => price + delta,
494        }
495    }
496
497    /// Apply all shocks in a scenario to every matching asset in the portfolio.
498    ///
499    /// Returns a new `HashMap` with shocked prices. Assets not mentioned in the
500    /// scenario retain their original prices.
501    pub fn run_scenario(
502        &self,
503        portfolio: &HashMap<String, f64>,
504        scenario: &Scenario,
505    ) -> HashMap<String, f64> {
506        let mut result = portfolio.clone();
507        for asset_shock in &scenario.shocks {
508            if let Some(price) = result.get_mut(&asset_shock.asset_id) {
509                *price = self.apply_shock(*price, &asset_shock.shock);
510            }
511        }
512        result
513    }
514
515    /// Compute the P&L impact of a scenario on a set of positions.
516    ///
517    /// `P&L = Σ positions[asset] * (shocked_price[asset] - original_price[asset])`
518    ///
519    /// Assets present in `positions` but absent from either price map are skipped.
520    pub fn scenario_pnl(
521        &self,
522        original: &HashMap<String, f64>,
523        shocked: &HashMap<String, f64>,
524        positions: &HashMap<String, f64>,
525    ) -> f64 {
526        positions.iter().fold(0.0, |acc, (asset, &qty)| {
527            let orig = original.get(asset).copied().unwrap_or(0.0);
528            let shock = shocked.get(asset).copied().unwrap_or(orig);
529            acc + qty * (shock - orig)
530        })
531    }
532
533    /// Find the worst-case scenario (most negative P&L) from a list of scenarios.
534    ///
535    /// Returns a reference to the worst scenario and its P&L.
536    ///
537    /// # Panics
538    /// Panics if `scenarios` is empty.
539    pub fn worst_case<'a>(
540        &self,
541        portfolio: &HashMap<String, f64>,
542        scenarios: &'a [Scenario],
543        positions: &HashMap<String, f64>,
544    ) -> (&'a Scenario, f64) {
545        assert!(!scenarios.is_empty(), "scenarios must not be empty");
546        let mut worst_scenario = &scenarios[0];
547        let shocked = self.run_scenario(portfolio, worst_scenario);
548        let mut worst_pnl = self.scenario_pnl(portfolio, &shocked, positions);
549
550        for scenario in scenarios.iter().skip(1) {
551            let shocked = self.run_scenario(portfolio, scenario);
552            let pnl = self.scenario_pnl(portfolio, &shocked, positions);
553            if pnl < worst_pnl {
554                worst_pnl = pnl;
555                worst_scenario = scenario;
556            }
557        }
558        (worst_scenario, worst_pnl)
559    }
560}