use crate::ohlcv::OhlcvBar;
use crate::risk::{DrawdownTracker, RiskBreach, RiskRule};
use rust_decimal::Decimal;
#[derive(Debug, Clone)]
pub struct ScenarioReport {
pub bars_processed: usize,
pub trigger_count: usize,
pub breaches: Vec<BarBreach>,
pub max_drawdown_pct: Decimal,
pub start_equity: Decimal,
pub end_equity: Decimal,
pub total_return_pct: Option<Decimal>,
}
#[derive(Debug, Clone)]
pub struct BarBreach {
pub bar_index: usize,
pub breaches: Vec<RiskBreach>,
pub equity: Decimal,
pub drawdown_pct: Decimal,
}
pub struct ScenarioBacktester {
bars: Vec<OhlcvBar>,
rules: Vec<Box<dyn RiskRule>>,
}
impl ScenarioBacktester {
pub fn new(bars: Vec<OhlcvBar>) -> Self {
Self { bars, rules: Vec::new() }
}
pub fn add_rule(mut self, rule: Box<dyn RiskRule>) -> Self {
self.rules.push(rule);
self
}
pub fn run<F>(&self, equity_fn: F) -> ScenarioReport
where
F: Fn(&OhlcvBar) -> Decimal,
{
if self.bars.is_empty() {
return ScenarioReport {
bars_processed: 0,
trigger_count: 0,
breaches: vec![],
max_drawdown_pct: Decimal::ZERO,
start_equity: Decimal::ZERO,
end_equity: Decimal::ZERO,
total_return_pct: None,
};
}
let first_equity = equity_fn(&self.bars[0]);
let mut tracker = DrawdownTracker::new(first_equity);
let mut all_breaches: Vec<BarBreach> = Vec::new();
let mut trigger_count = 0usize;
let mut last_equity = first_equity;
for (i, bar) in self.bars.iter().enumerate() {
let equity = equity_fn(bar);
tracker.update(equity);
let dd_pct = tracker.current_drawdown_pct();
last_equity = equity;
let bar_breaches: Vec<RiskBreach> = self
.rules
.iter()
.filter_map(|rule| rule.check(equity, dd_pct))
.collect();
if !bar_breaches.is_empty() {
trigger_count += 1;
all_breaches.push(BarBreach {
bar_index: i,
breaches: bar_breaches,
equity,
drawdown_pct: dd_pct,
});
}
}
let max_dd = tracker.worst_drawdown_pct();
let total_return_pct = if first_equity.is_zero() {
None
} else {
Some((last_equity - first_equity) / first_equity * Decimal::ONE_HUNDRED)
};
ScenarioReport {
bars_processed: self.bars.len(),
trigger_count,
breaches: all_breaches,
max_drawdown_pct: max_dd,
start_equity: first_equity,
end_equity: last_equity,
total_return_pct,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::risk::MaxDrawdownRule;
use crate::types::{NanoTimestamp, Price, Quantity, Symbol};
use rust_decimal_macros::dec;
fn sym() -> Symbol {
Symbol::new("SPY").unwrap()
}
fn ts() -> NanoTimestamp {
NanoTimestamp::new(0)
}
fn make_bar(close: rust_decimal::Decimal) -> OhlcvBar {
let p = Price::new(close).unwrap();
let high = Price::new(close + dec!(1)).unwrap();
OhlcvBar::new(
sym(),
p,
high,
p,
p,
Quantity::new(dec!(1000)).unwrap(),
ts(),
ts(),
10,
)
.unwrap()
}
#[test]
fn test_no_triggers_when_equity_rises() {
let bars: Vec<_> = (1..=10).map(|i| make_bar(dec!(100) + rust_decimal::Decimal::from(i))).collect();
let rule = MaxDrawdownRule { threshold_pct: dec!(5) };
let report = ScenarioBacktester::new(bars).add_rule(Box::new(rule)).run(|bar| bar.close.value());
assert_eq!(report.bars_processed, 10);
assert_eq!(report.trigger_count, 0);
assert_eq!(report.max_drawdown_pct, Decimal::ZERO);
}
#[test]
fn test_triggers_when_drawdown_exceeds_threshold() {
let closes = [
dec!(100), dec!(99), dec!(95), dec!(90), dec!(85), dec!(80),
];
let bars: Vec<_> = closes.iter().map(|&c| make_bar(c)).collect();
let rule = MaxDrawdownRule { threshold_pct: dec!(10) };
let report = ScenarioBacktester::new(bars).add_rule(Box::new(rule)).run(|bar| bar.close.value());
assert!(report.trigger_count > 0, "expected at least one trigger");
assert!(report.max_drawdown_pct > dec!(10));
}
#[test]
fn test_empty_bars_returns_zero_report() {
let report = ScenarioBacktester::new(vec![]).run(|bar| bar.close.value());
assert_eq!(report.bars_processed, 0);
assert_eq!(report.trigger_count, 0);
assert!(report.total_return_pct.is_none());
}
#[test]
fn test_total_return_pct_computed() {
let bars = vec![make_bar(dec!(100)), make_bar(dec!(110))];
let report = ScenarioBacktester::new(bars).run(|bar| bar.close.value());
assert_eq!(report.total_return_pct.unwrap(), dec!(10));
}
#[test]
fn test_multiple_rules_both_can_fire() {
let closes = [dec!(100), dec!(50)]; let bars: Vec<_> = closes.iter().map(|&c| make_bar(c)).collect();
let rule1 = MaxDrawdownRule { threshold_pct: dec!(10) };
let rule2 = MaxDrawdownRule { threshold_pct: dec!(20) };
let report = ScenarioBacktester::new(bars)
.add_rule(Box::new(rule1))
.add_rule(Box::new(rule2))
.run(|bar| bar.close.value());
let bar1 = report.breaches.iter().find(|b| b.bar_index == 1).unwrap();
assert_eq!(bar1.breaches.len(), 2);
}
#[test]
fn test_max_drawdown_tracked() {
let closes = [dec!(200), dec!(180), dec!(160), dec!(190), dec!(210)];
let bars: Vec<_> = closes.iter().map(|&c| make_bar(c)).collect();
let report = ScenarioBacktester::new(bars).run(|bar| bar.close.value());
assert_eq!(report.max_drawdown_pct, dec!(20));
}
#[test]
fn test_apply_absolute_shift() {
let engine = ScenarioEngine;
let shocked = engine.apply_shock(100.0, &ShockType::AbsoluteShift(-30.0));
assert!((shocked - 70.0).abs() < 1e-9);
}
#[test]
fn test_apply_relative_shift() {
let engine = ScenarioEngine;
let shocked = engine.apply_shock(100.0, &ShockType::RelativeShift(-0.20));
assert!((shocked - 80.0).abs() < 1e-9);
}
#[test]
fn test_apply_volatility_scaling() {
let engine = ScenarioEngine;
let shocked = engine.apply_shock(100.0, &ShockType::VolatilityScaling(1.5));
assert!((shocked - 150.0).abs() < 1e-9);
}
#[test]
fn test_apply_correlation_breakdown() {
let engine = ScenarioEngine;
let shocked = engine.apply_shock(100.0, &ShockType::CorrelationBreakdown(10.0));
assert!((shocked - 110.0).abs() < 1e-9);
}
#[test]
fn test_run_scenario_equity_crash() {
use std::collections::HashMap;
let mut portfolio: HashMap<String, f64> = HashMap::new();
portfolio.insert("equity".to_owned(), 100.0);
portfolio.insert("vol".to_owned(), 20.0);
let s = Scenario::equity_crash();
let engine = ScenarioEngine;
let shocked = engine.run_scenario(&portfolio, &s);
let eq = shocked["equity"];
assert!(eq < 100.0, "equity should drop: {eq}");
}
#[test]
fn test_scenario_pnl_loss() {
use std::collections::HashMap;
let mut original: HashMap<String, f64> = HashMap::new();
original.insert("equity".to_owned(), 100.0);
let mut shocked: HashMap<String, f64> = HashMap::new();
shocked.insert("equity".to_owned(), 70.0);
let mut positions: HashMap<String, f64> = HashMap::new();
positions.insert("equity".to_owned(), 10.0);
let engine = ScenarioEngine;
let pnl = engine.scenario_pnl(&original, &shocked, &positions);
assert!((pnl - (-300.0)).abs() < 1e-9, "pnl={pnl}");
}
#[test]
fn test_worst_case_scenario() {
use std::collections::HashMap;
let mut portfolio: HashMap<String, f64> = HashMap::new();
portfolio.insert("equity".to_owned(), 100.0);
let mut positions: HashMap<String, f64> = HashMap::new();
positions.insert("equity".to_owned(), 1.0);
let scenarios = vec![
Scenario::equity_crash(),
Scenario::rate_shock(),
];
let engine = ScenarioEngine;
let (worst, pnl) = engine.worst_case(&portfolio, &scenarios, &positions);
assert!(pnl <= 0.0 || pnl.is_finite());
assert!(!worst.name.is_empty());
}
#[test]
fn test_built_in_scenarios_valid() {
assert!(!Scenario::equity_crash().shocks.is_empty());
assert!(!Scenario::credit_crisis().shocks.is_empty());
assert!(!Scenario::rate_shock().shocks.is_empty());
assert!(!Scenario::fx_devaluation().shocks.is_empty());
}
}
use std::collections::HashMap;
#[derive(Debug, Clone)]
pub enum ShockType {
AbsoluteShift(f64),
RelativeShift(f64),
VolatilityScaling(f64),
CorrelationBreakdown(f64),
}
#[derive(Debug, Clone)]
pub struct AssetShock {
pub asset_id: String,
pub shock: ShockType,
}
#[derive(Debug, Clone)]
pub struct Scenario {
pub name: String,
pub description: String,
pub shocks: Vec<AssetShock>,
pub probability: f64,
}
impl Scenario {
pub fn equity_crash() -> Self {
Self {
name: "equity_crash".to_owned(),
description: "2008-style equity market crash: equities -30%, implied vol +50%."
.to_owned(),
probability: 0.05,
shocks: vec![
AssetShock {
asset_id: "equity".to_owned(),
shock: ShockType::RelativeShift(-0.30),
},
AssetShock {
asset_id: "vol".to_owned(),
shock: ShockType::RelativeShift(0.50),
},
],
}
}
pub fn credit_crisis() -> Self {
Self {
name: "credit_crisis".to_owned(),
description: "Credit crisis: IG credit -20%, HY spreads widen by +500 bps.".to_owned(),
probability: 0.03,
shocks: vec![
AssetShock {
asset_id: "ig_credit".to_owned(),
shock: ShockType::RelativeShift(-0.20),
},
AssetShock {
asset_id: "hy_spread".to_owned(),
shock: ShockType::AbsoluteShift(5.0),
},
],
}
}
pub fn rate_shock() -> Self {
Self {
name: "rate_shock".to_owned(),
description: "Sudden 200 bps rate hike across the yield curve.".to_owned(),
probability: 0.04,
shocks: vec![AssetShock {
asset_id: "rates".to_owned(),
shock: ShockType::AbsoluteShift(2.0),
}],
}
}
pub fn fx_devaluation() -> Self {
Self {
name: "fx_devaluation".to_owned(),
description: "EM FX devaluation: EM currencies -20% vs USD.".to_owned(),
probability: 0.06,
shocks: vec![AssetShock {
asset_id: "em_fx".to_owned(),
shock: ShockType::RelativeShift(-0.20),
}],
}
}
}
pub struct ScenarioEngine;
impl ScenarioEngine {
pub fn apply_shock(&self, price: f64, shock: &ShockType) -> f64 {
match shock {
ShockType::AbsoluteShift(delta) => price + delta,
ShockType::RelativeShift(frac) => price * (1.0 + frac),
ShockType::VolatilityScaling(factor) => price * factor,
ShockType::CorrelationBreakdown(delta) => price + delta,
}
}
pub fn run_scenario(
&self,
portfolio: &HashMap<String, f64>,
scenario: &Scenario,
) -> HashMap<String, f64> {
let mut result = portfolio.clone();
for asset_shock in &scenario.shocks {
if let Some(price) = result.get_mut(&asset_shock.asset_id) {
*price = self.apply_shock(*price, &asset_shock.shock);
}
}
result
}
pub fn scenario_pnl(
&self,
original: &HashMap<String, f64>,
shocked: &HashMap<String, f64>,
positions: &HashMap<String, f64>,
) -> f64 {
positions.iter().fold(0.0, |acc, (asset, &qty)| {
let orig = original.get(asset).copied().unwrap_or(0.0);
let shock = shocked.get(asset).copied().unwrap_or(orig);
acc + qty * (shock - orig)
})
}
pub fn worst_case<'a>(
&self,
portfolio: &HashMap<String, f64>,
scenarios: &'a [Scenario],
positions: &HashMap<String, f64>,
) -> (&'a Scenario, f64) {
assert!(!scenarios.is_empty(), "scenarios must not be empty");
let mut worst_scenario = &scenarios[0];
let shocked = self.run_scenario(portfolio, worst_scenario);
let mut worst_pnl = self.scenario_pnl(portfolio, &shocked, positions);
for scenario in scenarios.iter().skip(1) {
let shocked = self.run_scenario(portfolio, scenario);
let pnl = self.scenario_pnl(portfolio, &shocked, positions);
if pnl < worst_pnl {
worst_pnl = pnl;
worst_scenario = scenario;
}
}
(worst_scenario, worst_pnl)
}
}