1use crate::ohlcv::OhlcvBar;
18use crate::risk::{DrawdownTracker, RiskBreach, RiskRule};
19use rust_decimal::Decimal;
20
21#[derive(Debug, Clone)]
23pub struct ScenarioReport {
24 pub bars_processed: usize,
26 pub trigger_count: usize,
28 pub breaches: Vec<BarBreach>,
30 pub max_drawdown_pct: Decimal,
32 pub start_equity: Decimal,
34 pub end_equity: Decimal,
36 pub total_return_pct: Option<Decimal>,
40}
41
42#[derive(Debug, Clone)]
44pub struct BarBreach {
45 pub bar_index: usize,
47 pub breaches: Vec<RiskBreach>,
49 pub equity: Decimal,
51 pub drawdown_pct: Decimal,
53}
54
55pub struct ScenarioBacktester {
91 bars: Vec<OhlcvBar>,
92 rules: Vec<Box<dyn RiskRule>>,
93}
94
95impl ScenarioBacktester {
96 pub fn new(bars: Vec<OhlcvBar>) -> Self {
98 Self { bars, rules: Vec::new() }
99 }
100
101 pub fn add_rule(mut self, rule: Box<dyn RiskRule>) -> Self {
105 self.rules.push(rule);
106 self
107 }
108
109 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 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 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)]; 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 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 assert_eq!(report.max_drawdown_pct, dec!(20));
271 }
272
273 #[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 assert!((shocked - 150.0).abs() < 1e-9);
295 }
296
297 #[test]
298 fn test_apply_correlation_breakdown() {
299 let engine = ScenarioEngine;
300 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 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 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
360use std::collections::HashMap;
365
366#[derive(Debug, Clone)]
373pub enum ShockType {
374 AbsoluteShift(f64),
376 RelativeShift(f64),
378 VolatilityScaling(f64),
380 CorrelationBreakdown(f64),
382}
383
384#[derive(Debug, Clone)]
386pub struct AssetShock {
387 pub asset_id: String,
389 pub shock: ShockType,
391}
392
393#[derive(Debug, Clone)]
395pub struct Scenario {
396 pub name: String,
398 pub description: String,
400 pub shocks: Vec<AssetShock>,
402 pub probability: f64,
404}
405
406impl Scenario {
407 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 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 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 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
473pub struct ScenarioEngine;
485
486impl ScenarioEngine {
487 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 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 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 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}