1use serde::{Deserialize, Deserializer, Serialize};
4
5use super::journal::{StrategyJournalError, maximum_research_limits, validate_experiment_label};
6use super::{StrategyBacktestResult, StrategyDescriptor, StrategyResearchLimits};
7
8#[derive(Debug, Clone, PartialEq, Serialize)]
10pub struct StrategyComparisonMetrics {
11 pub total_positions: usize,
12 pub position_win_rate: f64,
13 pub total_pnl: f64,
14 pub max_drawdown: f64,
15 pub max_drawdown_pct: f64,
16 pub average_position_duration_secs: Option<i64>,
17}
18
19impl StrategyComparisonMetrics {
20 fn from_result(result: &StrategyBacktestResult) -> Self {
21 Self {
22 total_positions: result.replay.total_positions,
23 position_win_rate: result.replay.position_win_rate,
24 total_pnl: result.replay.total_pnl,
25 max_drawdown: result.replay.max_drawdown,
26 max_drawdown_pct: result.replay.max_drawdown_pct,
27 average_position_duration_secs: result
28 .replay
29 .duration_stats
30 .as_ref()
31 .map(|stats| stats.avg_duration_secs),
32 }
33 }
34
35 fn validate(&self) -> Result<(), StrategyExperimentError> {
36 for (field, value) in [
37 ("position_win_rate", self.position_win_rate),
38 ("total_pnl", self.total_pnl),
39 ("max_drawdown", self.max_drawdown),
40 ("max_drawdown_pct", self.max_drawdown_pct),
41 ] {
42 if !value.is_finite() {
43 return Err(StrategyExperimentError::NonFiniteMetric { field });
44 }
45 }
46 Ok(())
47 }
48}
49
50#[derive(Deserialize)]
51#[serde(deny_unknown_fields)]
52struct StrategyComparisonMetricsDef {
53 total_positions: usize,
54 position_win_rate: f64,
55 total_pnl: f64,
56 max_drawdown: f64,
57 max_drawdown_pct: f64,
58 average_position_duration_secs: Option<i64>,
59}
60
61impl<'de> Deserialize<'de> for StrategyComparisonMetrics {
62 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
63 where
64 D: Deserializer<'de>,
65 {
66 let value = StrategyComparisonMetricsDef::deserialize(deserializer)?;
67 let metrics = Self {
68 total_positions: value.total_positions,
69 position_win_rate: value.position_win_rate,
70 total_pnl: value.total_pnl,
71 max_drawdown: value.max_drawdown,
72 max_drawdown_pct: value.max_drawdown_pct,
73 average_position_duration_secs: value.average_position_duration_secs,
74 };
75 metrics.validate().map_err(serde::de::Error::custom)?;
76 Ok(metrics)
77 }
78}
79
80#[derive(Debug, Clone, PartialEq, Serialize)]
82pub struct StrategyComparisonSnapshot {
83 pub label: String,
84 pub descriptor: StrategyDescriptor,
85 pub metrics: StrategyComparisonMetrics,
86}
87
88impl StrategyComparisonSnapshot {
89 fn new(
90 label: impl Into<String>,
91 result: &StrategyBacktestResult,
92 limits: StrategyResearchLimits,
93 ) -> Result<Self, StrategyExperimentError> {
94 let label = label.into();
95 validate_experiment_label(&label, limits)?;
96 let metrics = StrategyComparisonMetrics::from_result(result);
97 metrics.validate()?;
98 Ok(Self {
99 label,
100 descriptor: result.descriptor.clone(),
101 metrics,
102 })
103 }
104}
105
106#[derive(Deserialize)]
107#[serde(deny_unknown_fields)]
108struct StrategyComparisonSnapshotDef {
109 label: String,
110 descriptor: StrategyDescriptor,
111 metrics: StrategyComparisonMetrics,
112}
113
114impl<'de> Deserialize<'de> for StrategyComparisonSnapshot {
115 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
116 where
117 D: Deserializer<'de>,
118 {
119 let value = StrategyComparisonSnapshotDef::deserialize(deserializer)?;
120 validate_experiment_label(&value.label, maximum_research_limits())
121 .map_err(serde::de::Error::custom)?;
122 Ok(Self {
123 label: value.label,
124 descriptor: value.descriptor,
125 metrics: value.metrics,
126 })
127 }
128}
129
130#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
132#[serde(deny_unknown_fields)]
133pub struct StrategyExperimentComparison {
134 pub baseline: StrategyComparisonSnapshot,
135 pub candidate: StrategyComparisonSnapshot,
136}
137
138impl StrategyExperimentComparison {
139 pub fn new(
140 baseline_label: impl Into<String>,
141 baseline: &StrategyBacktestResult,
142 candidate_label: impl Into<String>,
143 candidate: &StrategyBacktestResult,
144 limits: StrategyResearchLimits,
145 ) -> Result<Self, StrategyExperimentError> {
146 Ok(Self {
147 baseline: StrategyComparisonSnapshot::new(baseline_label, baseline, limits)?,
148 candidate: StrategyComparisonSnapshot::new(candidate_label, candidate, limits)?,
149 })
150 }
151}
152
153#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
155pub enum StrategyExperimentError {
156 #[error(transparent)]
157 Journal(#[from] StrategyJournalError),
158 #[error("comparison metric '{field}' must be finite")]
159 NonFiniteMetric { field: &'static str },
160}