1use crate::backtest::{BacktestConfig, Backtester, BacktestResult, Strategy};
83use crate::error::FinError;
84use crate::ohlcv::OhlcvBar;
85use std::collections::HashMap;
86
87#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
94pub struct ParamRange {
95 pub name: String,
97 pub min: f64,
99 pub max: f64,
101 pub step: f64,
103}
104
105impl ParamRange {
106 pub fn values(&self) -> Vec<f64> {
111 if self.step <= 0.0 || self.min > self.max {
112 return vec![self.min];
113 }
114 let mut vals = Vec::new();
115 let mut v = self.min;
116 while v <= self.max + f64::EPSILON {
117 vals.push(v);
118 v += self.step;
119 }
120 vals
121 }
122}
123
124#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
128pub struct WalkForwardConfig {
129 pub train_window: usize,
131 pub test_window: usize,
133 pub step: usize,
136 pub param_space: Vec<ParamRange>,
138}
139
140impl WalkForwardConfig {
141 pub fn validate(&self) -> Result<(), FinError> {
147 if self.train_window == 0 {
148 return Err(FinError::InvalidInput(
149 "train_window must be > 0".to_owned(),
150 ));
151 }
152 if self.test_window == 0 {
153 return Err(FinError::InvalidInput(
154 "test_window must be > 0".to_owned(),
155 ));
156 }
157 if self.step == 0 {
158 return Err(FinError::InvalidInput(
159 "step must be > 0".to_owned(),
160 ));
161 }
162 Ok(())
163 }
164}
165
166#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
173pub struct WfPeriod {
174 pub train_start: usize,
176 pub train_end: usize,
178 pub test_start: usize,
180 pub test_end: usize,
182 pub best_params: HashMap<String, f64>,
184 pub in_sample_sharpe: f64,
186 pub out_of_sample_sharpe: f64,
188 pub oos_result: BacktestResult,
190}
191
192#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
196pub struct WalkForwardResult {
197 pub periods: Vec<WfPeriod>,
199 pub aggregate_sharpe: f64,
201 pub stability_score: f64,
205 pub mean_oos_return: f64,
207 pub worst_oos_drawdown: f64,
209}
210
211impl WalkForwardResult {
212 pub fn is_robust(&self, min_stability: f64) -> bool {
215 self.aggregate_sharpe > 0.0 && self.stability_score >= min_stability
216 }
217
218 pub fn best_period(&self) -> Option<&WfPeriod> {
220 self.periods
221 .iter()
222 .max_by(|a, b| a.out_of_sample_sharpe.partial_cmp(&b.out_of_sample_sharpe).unwrap_or(std::cmp::Ordering::Equal))
223 }
224
225 pub fn worst_period(&self) -> Option<&WfPeriod> {
227 self.periods
228 .iter()
229 .min_by(|a, b| a.out_of_sample_sharpe.partial_cmp(&b.out_of_sample_sharpe).unwrap_or(std::cmp::Ordering::Equal))
230 }
231}
232
233pub struct WalkForwardOptimizer {
246 config: WalkForwardConfig,
247 bt_config: BacktestConfig,
248}
249
250impl WalkForwardOptimizer {
251 pub fn new(config: WalkForwardConfig, bt_config: BacktestConfig) -> Result<Self, FinError> {
256 config.validate()?;
257 Ok(Self { config, bt_config })
258 }
259
260 pub fn run<F>(
273 &self,
274 bars: &[OhlcvBar],
275 mut make_strategy: F,
276 ) -> Result<WalkForwardResult, FinError>
277 where
278 F: FnMut(&[OhlcvBar], &HashMap<String, f64>) -> Box<dyn Strategy>,
279 {
280 let window = self.config.train_window + self.config.test_window;
281 if bars.len() < window {
282 return Err(FinError::InvalidInput(format!(
283 "need at least {} bars for one walk-forward window, got {}",
284 window,
285 bars.len()
286 )));
287 }
288
289 let backtester = Backtester::new(self.bt_config.clone());
290 let grid = build_grid(&self.config.param_space);
291 let mut periods: Vec<WfPeriod> = Vec::new();
292 let mut offset = 0usize;
293
294 while offset + window <= bars.len() {
295 let train_start = offset;
296 let train_end = offset + self.config.train_window;
297 let test_start = train_end;
298 let test_end = train_end + self.config.test_window;
299
300 let train_bars = &bars[train_start..train_end];
301 let test_bars = &bars[test_start..test_end];
302
303 let mut best_is_sharpe = f64::NEG_INFINITY;
305 let mut best_params: HashMap<String, f64> = HashMap::new();
306
307 let search_grid: &[HashMap<String, f64>] = &grid;
308
309 for param_set in search_grid {
310 let mut is_strategy = make_strategy(train_bars, param_set);
312 let is_result = match backtester.run(train_bars, is_strategy.as_mut()) {
313 Ok(r) => r,
314 Err(_) => continue,
315 };
316 let is_sharpe = is_result
317 .sharpe_ratio
318 .to_string()
319 .parse::<f64>()
320 .unwrap_or(f64::NEG_INFINITY);
321
322 if is_sharpe > best_is_sharpe {
323 best_is_sharpe = is_sharpe;
324 best_params = param_set.clone();
325 }
326 }
327
328 let mut oos_strategy = make_strategy(test_bars, &best_params);
330 let oos_result = backtester.run(test_bars, oos_strategy.as_mut())?;
331
332 let oos_sharpe = oos_result
333 .sharpe_ratio
334 .to_string()
335 .parse::<f64>()
336 .unwrap_or(0.0);
337
338 periods.push(WfPeriod {
339 train_start,
340 train_end,
341 test_start,
342 test_end,
343 best_params,
344 in_sample_sharpe: best_is_sharpe.max(0.0), out_of_sample_sharpe: oos_sharpe,
346 oos_result,
347 });
348
349 offset += self.config.step;
350 }
351
352 if periods.is_empty() {
353 return Err(FinError::InvalidInput(
354 "no walk-forward periods could be constructed".to_owned(),
355 ));
356 }
357
358 let n = periods.len() as f64;
360 let aggregate_sharpe = periods.iter().map(|p| p.out_of_sample_sharpe).sum::<f64>() / n;
361 let positive_count = periods.iter().filter(|p| p.out_of_sample_sharpe > 0.0).count();
362 let stability_score = positive_count as f64 / n;
363
364 let mean_oos_return = periods
365 .iter()
366 .map(|p| {
367 p.oos_result
368 .total_return
369 .to_string()
370 .parse::<f64>()
371 .unwrap_or(0.0)
372 })
373 .sum::<f64>()
374 / n;
375
376 let worst_oos_drawdown = periods
377 .iter()
378 .map(|p| {
379 p.oos_result
380 .max_drawdown
381 .to_string()
382 .parse::<f64>()
383 .unwrap_or(0.0)
384 })
385 .fold(0.0_f64, f64::max);
386
387 Ok(WalkForwardResult {
388 periods,
389 aggregate_sharpe,
390 stability_score,
391 mean_oos_return,
392 worst_oos_drawdown,
393 })
394 }
395
396 pub fn config(&self) -> &WalkForwardConfig {
398 &self.config
399 }
400
401 pub fn bt_config(&self) -> &BacktestConfig {
403 &self.bt_config
404 }
405}
406
407fn build_grid(param_space: &[ParamRange]) -> Vec<HashMap<String, f64>> {
414 if param_space.is_empty() {
415 return vec![HashMap::new()];
416 }
417
418 let mut grid: Vec<HashMap<String, f64>> = vec![HashMap::new()];
419
420 for param in param_space {
421 let vals = param.values();
422 let mut new_grid = Vec::with_capacity(grid.len() * vals.len());
423 for existing in &grid {
424 for &v in &vals {
425 let mut m = existing.clone();
426 m.insert(param.name.clone(), v);
427 new_grid.push(m);
428 }
429 }
430 grid = new_grid;
431 }
432
433 grid
434}
435
436#[cfg(test)]
439mod tests {
440 use super::*;
441 use crate::backtest::{Signal, SignalDirection};
442 use crate::types::{NanoTimestamp, Price, Quantity, Symbol};
443 use rust_decimal::Decimal;
444 use rust_decimal_macros::dec;
445
446 fn make_bar(close: f64, ts: i64) -> OhlcvBar {
447 let sym = Symbol::new("TEST").unwrap();
448 let p = Price::new(Decimal::try_from(close).unwrap()).unwrap();
449 OhlcvBar {
450 symbol: sym,
451 open: p,
452 high: p,
453 low: p,
454 close: p,
455 volume: Quantity::new(dec!(1000)).unwrap(),
456 ts_open: NanoTimestamp::new(ts),
457 ts_close: NanoTimestamp::new(ts + 1),
458 tick_count: 1,
459 }
460 }
461
462 struct HoldAll;
464 impl Strategy for HoldAll {
465 fn on_bar(&mut self, _: &OhlcvBar) -> Option<Signal> {
466 None
467 }
468 }
469
470 struct BuyOnce {
472 bought: bool,
473 qty: f64,
474 }
475 impl BuyOnce {
476 fn from_params(params: &HashMap<String, f64>) -> Self {
477 Self {
478 bought: false,
479 qty: params.get("qty").copied().unwrap_or(1.0),
480 }
481 }
482 }
483 impl Strategy for BuyOnce {
484 fn on_bar(&mut self, _: &OhlcvBar) -> Option<Signal> {
485 if !self.bought {
486 self.bought = true;
487 let qty = Decimal::try_from(self.qty).unwrap_or(dec!(1));
488 return Some(Signal::new(SignalDirection::Buy, qty));
489 }
490 Some(Signal::hold())
491 }
492 }
493
494 #[test]
497 fn test_param_range_values_basic() {
498 let r = ParamRange { name: "x".to_owned(), min: 1.0, max: 3.0, step: 1.0 };
499 let vals = r.values();
500 assert_eq!(vals.len(), 3);
501 assert!((vals[0] - 1.0).abs() < 1e-10);
502 assert!((vals[2] - 3.0).abs() < 1e-10);
503 }
504
505 #[test]
506 fn test_param_range_single_value_when_min_equals_max() {
507 let r = ParamRange { name: "x".to_owned(), min: 5.0, max: 5.0, step: 1.0 };
508 let vals = r.values();
509 assert_eq!(vals.len(), 1);
510 assert!((vals[0] - 5.0).abs() < 1e-10);
511 }
512
513 #[test]
514 fn test_param_range_degenerate_step() {
515 let r = ParamRange { name: "x".to_owned(), min: 1.0, max: 5.0, step: 0.0 };
516 let vals = r.values();
517 assert_eq!(vals.len(), 1); }
519
520 #[test]
523 fn test_build_grid_empty_space() {
524 let grid = build_grid(&[]);
525 assert_eq!(grid.len(), 1);
526 assert!(grid[0].is_empty());
527 }
528
529 #[test]
530 fn test_build_grid_single_param() {
531 let params = vec![ParamRange { name: "p".to_owned(), min: 5.0, max: 15.0, step: 5.0 }];
532 let grid = build_grid(¶ms);
533 assert_eq!(grid.len(), 3); for m in &grid {
535 assert!(m.contains_key("p"));
536 }
537 }
538
539 #[test]
540 fn test_build_grid_two_params_cartesian() {
541 let params = vec![
542 ParamRange { name: "a".to_owned(), min: 1.0, max: 2.0, step: 1.0 },
543 ParamRange { name: "b".to_owned(), min: 10.0, max: 20.0, step: 10.0 },
544 ];
545 let grid = build_grid(¶ms);
546 assert_eq!(grid.len(), 4); }
548
549 #[test]
552 fn test_config_validation_zero_train() {
553 let cfg = WalkForwardConfig {
554 train_window: 0,
555 test_window: 20,
556 step: 20,
557 param_space: vec![],
558 };
559 assert!(cfg.validate().is_err());
560 }
561
562 #[test]
563 fn test_config_validation_zero_test() {
564 let cfg = WalkForwardConfig {
565 train_window: 60,
566 test_window: 0,
567 step: 20,
568 param_space: vec![],
569 };
570 assert!(cfg.validate().is_err());
571 }
572
573 #[test]
574 fn test_config_validation_zero_step() {
575 let cfg = WalkForwardConfig {
576 train_window: 60,
577 test_window: 20,
578 step: 0,
579 param_space: vec![],
580 };
581 assert!(cfg.validate().is_err());
582 }
583
584 #[test]
587 fn test_optimizer_too_few_bars() {
588 let bars: Vec<OhlcvBar> = (0..5).map(|i| make_bar(100.0, i)).collect();
589 let cfg = WalkForwardConfig {
590 train_window: 60,
591 test_window: 20,
592 step: 20,
593 param_space: vec![],
594 };
595 let bt_cfg = BacktestConfig::new(dec!(10_000), dec!(0)).unwrap();
596 let opt = WalkForwardOptimizer::new(cfg, bt_cfg).unwrap();
597 let result = opt.run(&bars, |_, _| Box::new(HoldAll));
598 assert!(result.is_err());
599 }
600
601 #[test]
602 fn test_optimizer_hold_strategy_returns_result() {
603 let bars: Vec<OhlcvBar> = (0..100).map(|i| make_bar(100.0 + i as f64 * 0.1, i)).collect();
604 let cfg = WalkForwardConfig {
605 train_window: 40,
606 test_window: 20,
607 step: 20,
608 param_space: vec![],
609 };
610 let bt_cfg = BacktestConfig::new(dec!(10_000), dec!(0)).unwrap();
611 let opt = WalkForwardOptimizer::new(cfg, bt_cfg).unwrap();
612 let result = opt.run(&bars, |_, _| Box::new(HoldAll)).unwrap();
613 assert!(!result.periods.is_empty());
614 assert_eq!(result.aggregate_sharpe, 0.0);
616 assert_eq!(result.stability_score, 0.0);
617 }
618
619 #[test]
620 fn test_optimizer_with_param_grid() {
621 let bars: Vec<OhlcvBar> = (0..120).map(|i| make_bar(100.0 + i as f64 * 0.5, i)).collect();
622 let cfg = WalkForwardConfig {
623 train_window: 50,
624 test_window: 20,
625 step: 20,
626 param_space: vec![
627 ParamRange { name: "qty".to_owned(), min: 1.0, max: 3.0, step: 1.0 },
628 ],
629 };
630 let bt_cfg = BacktestConfig::new(dec!(10_000), dec!(0)).unwrap();
631 let opt = WalkForwardOptimizer::new(cfg, bt_cfg).unwrap();
632 let result = opt.run(&bars, |_, params| Box::new(BuyOnce::from_params(params))).unwrap();
633 assert!(!result.periods.is_empty());
634 for p in &result.periods {
635 assert!(!p.best_params.is_empty());
636 assert!(p.best_params.contains_key("qty"));
637 }
638 }
639
640 #[test]
641 fn test_optimizer_stability_score_bounds() {
642 let bars: Vec<OhlcvBar> = (0..100).map(|i| make_bar(100.0, i)).collect();
643 let cfg = WalkForwardConfig {
644 train_window: 40,
645 test_window: 20,
646 step: 20,
647 param_space: vec![],
648 };
649 let bt_cfg = BacktestConfig::new(dec!(10_000), dec!(0)).unwrap();
650 let opt = WalkForwardOptimizer::new(cfg, bt_cfg).unwrap();
651 let result = opt.run(&bars, |_, _| Box::new(HoldAll)).unwrap();
652 assert!((0.0..=1.0).contains(&result.stability_score));
653 }
654
655 #[test]
656 fn test_wf_result_robustness_check() {
657 let result = WalkForwardResult {
658 periods: vec![],
659 aggregate_sharpe: 1.5,
660 stability_score: 0.8,
661 mean_oos_return: 0.05,
662 worst_oos_drawdown: 0.1,
663 };
664 assert!(result.is_robust(0.7));
665 assert!(!result.is_robust(0.9));
666 }
667
668 #[test]
669 fn test_wf_result_best_worst_period() {
670 let make_period = |oos_sharpe: f64| WfPeriod {
671 train_start: 0,
672 train_end: 50,
673 test_start: 50,
674 test_end: 70,
675 best_params: HashMap::new(),
676 in_sample_sharpe: 1.0,
677 out_of_sample_sharpe: oos_sharpe,
678 oos_result: crate::backtest::BacktestResult {
679 total_return: Decimal::ZERO,
680 sharpe_ratio: Decimal::ZERO,
681 max_drawdown: Decimal::ZERO,
682 win_rate: Decimal::ZERO,
683 trade_count: 0,
684 final_equity: dec!(10_000),
685 equity_curve: vec![],
686 },
687 };
688 let result = WalkForwardResult {
689 periods: vec![make_period(0.5), make_period(2.0), make_period(-0.3)],
690 aggregate_sharpe: 0.73,
691 stability_score: 0.67,
692 mean_oos_return: 0.0,
693 worst_oos_drawdown: 0.0,
694 };
695 assert!((result.best_period().unwrap().out_of_sample_sharpe - 2.0).abs() < 1e-10);
696 assert!((result.worst_period().unwrap().out_of_sample_sharpe + 0.3).abs() < 1e-10);
697 }
698
699 #[test]
700 fn test_optimizer_step_advances_window() {
701 let bars: Vec<OhlcvBar> = (0..150).map(|i| make_bar(100.0 + i as f64 * 0.1, i)).collect();
703 let cfg = WalkForwardConfig {
704 train_window: 50,
705 test_window: 30,
706 step: 10, param_space: vec![],
708 };
709 let bt_cfg = BacktestConfig::new(dec!(10_000), dec!(0)).unwrap();
710 let opt = WalkForwardOptimizer::new(cfg, bt_cfg).unwrap();
711 let result = opt.run(&bars, |_, _| Box::new(HoldAll)).unwrap();
712 assert!(result.periods.len() > 2);
714 }
715}