finance_query/backtesting/strategy/
builder.rs1use std::collections::HashSet;
26
27use crate::backtesting::condition::{Condition, HtfIndicatorSpec};
28use crate::backtesting::signal::Signal;
29use crate::indicators::Indicator;
30
31use super::{Strategy, StrategyContext};
32
33struct BoxedCondition {
35 evaluate_fn: Box<dyn Fn(&StrategyContext) -> bool + Send + Sync>,
36 required_indicators: Vec<(String, Indicator)>,
37 htf_requirements: Vec<HtfIndicatorSpec>,
38 tracks_position_extremes: bool,
39 description: String,
40}
41
42impl BoxedCondition {
43 fn new<C: Condition>(cond: C) -> Self {
44 let required_indicators = cond.required_indicators();
45 let htf_requirements = cond.htf_requirements();
46 let tracks_position_extremes = cond.tracks_position_extremes();
47 let description = cond.description();
48 Self {
49 evaluate_fn: Box::new(move |ctx| cond.evaluate(ctx)),
50 required_indicators,
51 htf_requirements,
52 tracks_position_extremes,
53 description,
54 }
55 }
56
57 fn evaluate(&self, ctx: &StrategyContext) -> bool {
58 (self.evaluate_fn)(ctx)
59 }
60
61 fn required_indicators(&self) -> &[(String, Indicator)] {
62 &self.required_indicators
63 }
64
65 fn htf_requirements(&self) -> &[HtfIndicatorSpec] {
66 &self.htf_requirements
67 }
68
69 fn tracks_position_extremes(&self) -> bool {
70 self.tracks_position_extremes
71 }
72
73 fn description(&self) -> &str {
74 &self.description
75 }
76}
77
78pub struct StrategyBuilder<E = (), X = ()> {
88 name: String,
89 entry_condition: E,
90 exit_condition: X,
91 short_entry_condition: Option<BoxedCondition>,
92 short_exit_condition: Option<BoxedCondition>,
93 regime_filter: Option<BoxedCondition>,
94 warmup_override: Option<usize>,
95}
96
97impl StrategyBuilder<(), ()> {
98 pub fn new(name: impl Into<String>) -> Self {
106 Self {
107 name: name.into(),
108 entry_condition: (),
109 exit_condition: (),
110 short_entry_condition: None,
111 short_exit_condition: None,
112 regime_filter: None,
113 warmup_override: None,
114 }
115 }
116}
117
118impl<X> StrategyBuilder<(), X> {
119 pub fn entry<C: Condition>(self, condition: C) -> StrategyBuilder<C, X> {
128 StrategyBuilder {
129 name: self.name,
130 entry_condition: condition,
131 exit_condition: self.exit_condition,
132 short_entry_condition: self.short_entry_condition,
133 short_exit_condition: self.short_exit_condition,
134 regime_filter: self.regime_filter,
135 warmup_override: self.warmup_override,
136 }
137 }
138}
139
140impl<E> StrategyBuilder<E, ()> {
141 pub fn exit<C: Condition>(self, condition: C) -> StrategyBuilder<E, C> {
151 StrategyBuilder {
152 name: self.name,
153 entry_condition: self.entry_condition,
154 exit_condition: condition,
155 short_entry_condition: self.short_entry_condition,
156 short_exit_condition: self.short_exit_condition,
157 regime_filter: self.regime_filter,
158 warmup_override: self.warmup_override,
159 }
160 }
161}
162
163impl<E, X> StrategyBuilder<E, X> {
164 pub fn regime_filter<C: Condition>(mut self, condition: C) -> Self {
188 self.regime_filter = Some(BoxedCondition::new(condition));
189 self
190 }
191}
192
193impl<E: Condition, X: Condition> StrategyBuilder<E, X> {
194 pub fn with_short<SE: Condition, SX: Condition>(mut self, entry: SE, exit: SX) -> Self {
209 self.short_entry_condition = Some(BoxedCondition::new(entry));
210 self.short_exit_condition = Some(BoxedCondition::new(exit));
211 self
212 }
213
214 pub fn warmup(mut self, bars: usize) -> Self {
230 self.warmup_override = Some(bars);
231 self
232 }
233
234 pub fn build(self) -> CustomStrategy<E, X> {
245 CustomStrategy {
246 name: self.name,
247 entry_condition: self.entry_condition,
248 exit_condition: self.exit_condition,
249 short_entry_condition: self.short_entry_condition,
250 short_exit_condition: self.short_exit_condition,
251 regime_filter: self.regime_filter,
252 warmup_override: self.warmup_override,
253 }
254 }
255}
256
257pub struct CustomStrategy<E: Condition, X: Condition> {
262 name: String,
263 entry_condition: E,
264 exit_condition: X,
265 short_entry_condition: Option<BoxedCondition>,
266 short_exit_condition: Option<BoxedCondition>,
267 regime_filter: Option<BoxedCondition>,
272 warmup_override: Option<usize>,
276}
277
278impl<E: Condition, X: Condition> Strategy for CustomStrategy<E, X> {
279 fn name(&self) -> &str {
280 &self.name
281 }
282
283 fn required_indicators(&self) -> Vec<(String, Indicator)> {
284 let mut indicators = self.entry_condition.required_indicators();
285 indicators.extend(self.exit_condition.required_indicators());
286
287 if let Some(ref se) = self.short_entry_condition {
288 indicators.extend(se.required_indicators().iter().cloned());
289 }
290 if let Some(ref sx) = self.short_exit_condition {
291 indicators.extend(sx.required_indicators().iter().cloned());
292 }
293 if let Some(ref rf) = self.regime_filter {
294 indicators.extend(rf.required_indicators().iter().cloned());
295 }
296
297 let mut seen: Vec<(String, Indicator)> = Vec::new();
300 indicators.retain(|item| {
301 let is_new = !seen.contains(item);
302 if is_new {
303 seen.push(item.clone());
304 }
305 is_new
306 });
307
308 indicators
309 }
310
311 fn htf_requirements(&self) -> Vec<HtfIndicatorSpec> {
312 let mut reqs = self.entry_condition.htf_requirements();
313 reqs.extend(self.exit_condition.htf_requirements());
314
315 if let Some(ref se) = self.short_entry_condition {
316 reqs.extend(se.htf_requirements().iter().cloned());
317 }
318 if let Some(ref sx) = self.short_exit_condition {
319 reqs.extend(sx.htf_requirements().iter().cloned());
320 }
321 if let Some(ref rf) = self.regime_filter {
322 reqs.extend(rf.htf_requirements().iter().cloned());
323 }
324
325 let mut seen = HashSet::new();
327 reqs.retain(|spec| seen.insert(spec.htf_key.clone()));
328 reqs
329 }
330
331 fn tracks_position_extremes(&self) -> bool {
332 self.entry_condition.tracks_position_extremes()
333 || self.exit_condition.tracks_position_extremes()
334 || self
335 .short_entry_condition
336 .as_ref()
337 .is_some_and(|c| c.tracks_position_extremes())
338 || self
339 .short_exit_condition
340 .as_ref()
341 .is_some_and(|c| c.tracks_position_extremes())
342 || self
343 .regime_filter
344 .as_ref()
345 .is_some_and(|c| c.tracks_position_extremes())
346 }
347
348 fn warmup_period(&self) -> usize {
349 if let Some(n) = self.warmup_override {
351 return n;
352 }
353
354 let max_warmup = self
358 .required_indicators()
359 .iter()
360 .map(|(_, indicator)| indicator.warmup_bars())
361 .max()
362 .unwrap_or(1);
363
364 max_warmup + 1
365 }
366
367 fn on_candle(&self, ctx: &StrategyContext) -> Signal {
368 let candle = ctx.current_candle();
369
370 if ctx.is_long() && self.exit_condition.evaluate(ctx) {
372 return Signal::exit(candle.timestamp, candle.close)
373 .with_reason(self.exit_condition.description());
374 }
375
376 if ctx.is_short()
377 && let Some(ref exit) = self.short_exit_condition
378 && exit.evaluate(ctx)
379 {
380 return Signal::exit(candle.timestamp, candle.close)
381 .with_reason(exit.description().to_string());
382 }
383
384 if !ctx.has_position() {
386 let regime_ok = self
388 .regime_filter
389 .as_ref()
390 .is_none_or(|rf| rf.evaluate(ctx));
391
392 if regime_ok {
393 if self.entry_condition.evaluate(ctx) {
395 return Signal::long(candle.timestamp, candle.close)
396 .with_reason(self.entry_condition.description());
397 }
398
399 if let Some(ref entry) = self.short_entry_condition
401 && entry.evaluate(ctx)
402 {
403 return Signal::short(candle.timestamp, candle.close)
404 .with_reason(entry.description().to_string());
405 }
406 }
407 }
408
409 Signal::hold()
410 }
411}
412
413#[cfg(test)]
414mod tests {
415 use std::collections::HashMap;
416
417 use super::*;
418 use crate::backtesting::condition::{always_false, always_true};
419 use crate::backtesting::signal::SignalDirection;
420 use crate::models::chart::Candle;
421
422 fn make_candle(ts: i64, close: f64) -> Candle {
423 Candle {
424 timestamp: ts,
425 open: close,
426 high: close,
427 low: close,
428 close,
429 volume: 1000,
430 adj_close: None,
431 provider_id: None,
432 }
433 }
434
435 fn make_ctx<'a>(
436 candles: &'a [Candle],
437 indicators: &'a HashMap<String, Vec<Option<f64>>>,
438 ) -> StrategyContext<'a> {
439 StrategyContext {
440 candles,
441 index: 0,
442 position: None,
443 equity: 10_000.0,
444 indicators,
445 extremes: None,
446 indicator_index: None,
447 }
448 }
449
450 #[test]
451 fn test_strategy_builder() {
452 let strategy = StrategyBuilder::new("Test Strategy")
453 .entry(always_true())
454 .exit(always_false())
455 .build();
456
457 assert_eq!(strategy.name(), "Test Strategy");
458 }
459
460 #[test]
461 fn test_strategy_builder_with_short() {
462 let strategy = StrategyBuilder::new("Test Strategy")
463 .entry(always_true())
464 .exit(always_false())
465 .with_short(always_false(), always_true())
466 .build();
467
468 assert_eq!(strategy.name(), "Test Strategy");
469 assert!(strategy.short_entry_condition.is_some());
470 assert!(strategy.short_exit_condition.is_some());
471 }
472
473 #[test]
474 fn test_required_indicators_deduplication() {
475 use crate::backtesting::condition::Above;
476 use crate::backtesting::refs::rsi;
477
478 let entry = Above::new(rsi(14), 70.0);
480 let exit = Above::new(rsi(14), 30.0);
481
482 let strategy = StrategyBuilder::new("Test").entry(entry).exit(exit).build();
483
484 let indicators = strategy.required_indicators();
485 assert_eq!(indicators.len(), 1);
487 assert_eq!(indicators[0].0, "rsi_14");
488 }
489
490 #[test]
493 fn test_regime_filter_suppresses_entry_when_false() {
494 let strategy = StrategyBuilder::new("Regime Test")
495 .regime_filter(always_false()) .entry(always_true())
497 .exit(always_false())
498 .build();
499
500 let candles = vec![make_candle(1, 100.0)];
501 let indicators = HashMap::new();
502 let ctx = make_ctx(&candles, &indicators);
503
504 assert_eq!(strategy.on_candle(&ctx).direction, SignalDirection::Hold);
506 }
507
508 #[test]
509 fn test_regime_filter_allows_entry_when_true() {
510 let strategy = StrategyBuilder::new("Regime Test")
511 .regime_filter(always_true()) .entry(always_true())
513 .exit(always_false())
514 .build();
515
516 let candles = vec![make_candle(1, 100.0)];
517 let indicators = HashMap::new();
518 let ctx = make_ctx(&candles, &indicators);
519
520 assert_eq!(strategy.on_candle(&ctx).direction, SignalDirection::Long);
521 }
522
523 #[test]
524 fn test_no_regime_filter_behaves_normally() {
525 let strategy = StrategyBuilder::new("No Regime")
526 .entry(always_true())
527 .exit(always_false())
528 .build();
529
530 let candles = vec![make_candle(1, 100.0)];
531 let indicators = HashMap::new();
532 let ctx = make_ctx(&candles, &indicators);
533
534 assert_eq!(strategy.on_candle(&ctx).direction, SignalDirection::Long);
535 }
536
537 #[test]
538 fn test_regime_filter_does_not_block_exit() {
539 use crate::backtesting::position::{Position, PositionSide};
540
541 let strategy = StrategyBuilder::new("Regime Exit Test")
542 .regime_filter(always_false()) .entry(always_false())
544 .exit(always_true()) .build();
546
547 let candles = vec![make_candle(1, 100.0)];
548 let indicators = HashMap::new();
549
550 let position = Position::new(
552 PositionSide::Long,
553 1,
554 90.0,
555 10.0,
556 0.0,
557 Signal::long(1, 90.0),
558 );
559
560 let ctx = StrategyContext {
561 candles: &candles,
562 index: 0,
563 position: Some(&position),
564 equity: 10_000.0,
565 indicators: &indicators,
566 extremes: None,
567 indicator_index: None,
568 };
569
570 assert_eq!(strategy.on_candle(&ctx).direction, SignalDirection::Exit);
572 }
573
574 #[test]
575 fn test_regime_filter_indicators_included_in_required() {
576 use crate::backtesting::refs::{IndicatorRefExt, sma};
577 use crate::indicators::Indicator;
578
579 let strategy = StrategyBuilder::new("Regime Indicators")
580 .regime_filter(sma(200).above_ref(sma(400)))
581 .entry(always_true())
582 .exit(always_false())
583 .build();
584
585 let indicators = strategy.required_indicators();
586 let keys: Vec<&str> = indicators.iter().map(|(k, _)| k.as_str()).collect();
587
588 assert!(
589 keys.contains(&"sma_200"),
590 "sma_200 must be in required_indicators"
591 );
592 assert!(
593 keys.contains(&"sma_400"),
594 "sma_400 must be in required_indicators"
595 );
596
597 let sma_200 = indicators.iter().find(|(k, _)| k == "sma_200").unwrap();
599 assert!(matches!(sma_200.1, Indicator::Sma(200)));
600 }
601
602 #[test]
603 fn test_regime_filter_callable_before_entry() {
604 let strategy = StrategyBuilder::new("Order Test")
606 .regime_filter(always_true())
607 .entry(always_true())
608 .exit(always_false())
609 .build();
610
611 assert!(strategy.regime_filter.is_some());
612 }
613
614 #[test]
615 fn test_regime_filter_warmup_accounts_for_filter_indicators() {
616 use crate::backtesting::refs::{IndicatorRefExt, sma};
617
618 let strategy = StrategyBuilder::new("Warmup Test")
619 .regime_filter(sma(400).above_ref(sma(200)))
620 .entry(always_true())
621 .exit(always_false())
622 .build();
623
624 assert!(
626 strategy.warmup_period() >= 401,
627 "warmup_period must account for sma(400): got {}",
628 strategy.warmup_period()
629 );
630 }
631}