Skip to main content

qs_backtest/
mtm.rs

1//! Deterministic mark-to-market output collection.
2
3use std::collections::{BTreeMap, BTreeSet};
4
5use serde::{Deserialize, Deserializer, Serialize, de::Error as _};
6use thiserror::Error;
7
8use crate::portfolio::EquityPoint;
9
10pub const DEFAULT_MTM_MAX_POINTS: usize = 4_096;
11pub const MIN_MTM_MAX_POINTS: usize = 8;
12pub const MAX_MTM_MAX_POINTS: usize = 16_384;
13
14/// Controls how many exact mark-to-market observations are included in output artifacts.
15#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
16#[serde(rename_all = "snake_case")]
17pub enum MtmOutputPolicy {
18    None,
19    Bounded { max_points: usize },
20    Full,
21}
22
23impl Default for MtmOutputPolicy {
24    fn default() -> Self {
25        Self::Bounded {
26            max_points: DEFAULT_MTM_MAX_POINTS,
27        }
28    }
29}
30
31impl MtmOutputPolicy {
32    pub fn validate(&self) -> Result<(), MtmOutputPolicyError> {
33        if let Self::Bounded { max_points } = *self
34            && !(MIN_MTM_MAX_POINTS..=MAX_MTM_MAX_POINTS).contains(&max_points)
35        {
36            return Err(MtmOutputPolicyError::InvalidMaxPoints { max_points });
37        }
38        Ok(())
39    }
40
41    pub fn max_points(self) -> Option<usize> {
42        match self {
43            Self::None => Some(0),
44            Self::Bounded { max_points } => Some(max_points),
45            Self::Full => None,
46        }
47    }
48}
49
50#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)]
51#[serde(rename_all = "snake_case")]
52enum MtmOutputPolicyRepr {
53    None,
54    Bounded { max_points: usize },
55    Full,
56}
57
58impl From<MtmOutputPolicyRepr> for MtmOutputPolicy {
59    fn from(value: MtmOutputPolicyRepr) -> Self {
60        match value {
61            MtmOutputPolicyRepr::None => Self::None,
62            MtmOutputPolicyRepr::Bounded { max_points } => Self::Bounded { max_points },
63            MtmOutputPolicyRepr::Full => Self::Full,
64        }
65    }
66}
67
68impl<'de> Deserialize<'de> for MtmOutputPolicy {
69    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
70    where
71        D: Deserializer<'de>,
72    {
73        let policy = Self::from(MtmOutputPolicyRepr::deserialize(deserializer)?);
74        policy.validate().map_err(D::Error::custom)?;
75        Ok(policy)
76    }
77}
78
79/// Validation failure for bounded mark-to-market output configuration.
80#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
81pub enum MtmOutputPolicyError {
82    #[error(
83        "MTM max_points must be between {MIN_MTM_MAX_POINTS} and {MAX_MTM_MAX_POINTS}, got {max_points}"
84    )]
85    InvalidMaxPoints { max_points: usize },
86}
87
88/// Counts describing the relationship between exact observations and emitted points.
89#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
90#[serde(default)]
91pub struct MtmOutputSummary {
92    pub policy: MtmOutputPolicy,
93    pub observed_points: u64,
94    pub retained_points: u64,
95    pub omitted_points: u64,
96}
97
98/// Streaming collector that bounds output while retaining significant observations.
99#[derive(Debug, Clone)]
100pub struct MtmCurveCollector {
101    policy: MtmOutputPolicy,
102    observed_points: u64,
103    points: BTreeMap<u64, EquityPoint>,
104    eviction_index: BTreeSet<(u64, u64)>,
105    first_sequence: Option<u64>,
106    last_sequence: Option<u64>,
107    min_equity: Option<(f64, u64)>,
108    max_equity: Option<(f64, u64)>,
109    max_drawdown: Option<(f64, u64)>,
110}
111
112impl Default for MtmCurveCollector {
113    fn default() -> Self {
114        Self::new(MtmOutputPolicy::default()).expect("default MTM output policy is valid")
115    }
116}
117
118impl MtmCurveCollector {
119    pub fn new(policy: MtmOutputPolicy) -> Result<Self, MtmOutputPolicyError> {
120        policy.validate()?;
121        Ok(Self {
122            policy,
123            observed_points: 0,
124            points: BTreeMap::new(),
125            eviction_index: BTreeSet::new(),
126            first_sequence: None,
127            last_sequence: None,
128            min_equity: None,
129            max_equity: None,
130            max_drawdown: None,
131        })
132    }
133
134    pub fn policy(&self) -> MtmOutputPolicy {
135        self.policy
136    }
137
138    pub fn observe(&mut self, point: EquityPoint) -> u64 {
139        self.push(point)
140    }
141
142    pub fn push(&mut self, mut point: EquityPoint) -> u64 {
143        let sequence = self.observed_points;
144        self.observed_points = self.observed_points.saturating_add(1);
145        if point.observation_sequence.is_none() {
146            point.observation_sequence = Some(sequence);
147        }
148
149        self.update_pins(sequence, &point);
150        if !matches!(self.policy, MtmOutputPolicy::None) {
151            self.points.insert(sequence, point);
152            if matches!(self.policy, MtmOutputPolicy::Bounded { .. }) {
153                self.eviction_index
154                    .insert((sample_priority(sequence), sequence));
155            }
156            self.enforce_bound();
157        }
158        sequence
159    }
160
161    pub fn extend(&mut self, points: impl IntoIterator<Item = EquityPoint>) {
162        for point in points {
163            self.push(point);
164        }
165    }
166
167    pub fn summary(&self) -> MtmOutputSummary {
168        let retained_points = self.points.len() as u64;
169        MtmOutputSummary {
170            policy: self.policy,
171            observed_points: self.observed_points,
172            retained_points,
173            omitted_points: self.observed_points.saturating_sub(retained_points),
174        }
175    }
176
177    pub fn retained_points(&self) -> Vec<EquityPoint> {
178        self.points.values().cloned().collect()
179    }
180
181    pub fn into_curve(self) -> Vec<EquityPoint> {
182        self.points.into_values().collect()
183    }
184
185    pub fn into_parts(self) -> (Vec<EquityPoint>, MtmOutputSummary) {
186        let summary = self.summary();
187        (self.into_curve(), summary)
188    }
189
190    fn update_pins(&mut self, sequence: u64, point: &EquityPoint) {
191        self.first_sequence.get_or_insert(sequence);
192        self.last_sequence = Some(sequence);
193
194        if let Some(equity) = point.equity.filter(|value| value.is_finite()) {
195            if self.min_equity.is_none_or(|(minimum, _)| equity < minimum) {
196                self.min_equity = Some((equity, sequence));
197            }
198            if self.max_equity.is_none_or(|(maximum, _)| equity > maximum) {
199                self.max_equity = Some((equity, sequence));
200            }
201        }
202
203        let drawdown = point
204            .drawdown
205            .filter(|value| value.is_finite())
206            .or_else(|| point.max_drawdown.filter(|value| value.is_finite()));
207        if let Some(drawdown) = drawdown
208            && self
209                .max_drawdown
210                .is_none_or(|(maximum, _)| drawdown > maximum)
211        {
212            self.max_drawdown = Some((drawdown, sequence));
213        }
214    }
215
216    fn enforce_bound(&mut self) {
217        let MtmOutputPolicy::Bounded { max_points } = self.policy else {
218            return;
219        };
220        while self.points.len() > max_points {
221            let (_, evicted) = self
222                .eviction_index
223                .iter()
224                .rev()
225                .copied()
226                .find(|(_, sequence)| !self.is_pinned(*sequence))
227                .expect("a valid bounded policy always leaves an unpinned point");
228            self.points.remove(&evicted);
229            self.eviction_index
230                .remove(&(sample_priority(evicted), evicted));
231        }
232    }
233
234    fn is_pinned(&self, sequence: u64) -> bool {
235        [
236            self.first_sequence,
237            self.last_sequence,
238            self.min_equity.map(|(_, sequence)| sequence),
239            self.max_equity.map(|(_, sequence)| sequence),
240            self.max_drawdown.map(|(_, sequence)| sequence),
241        ]
242        .into_iter()
243        .flatten()
244        .any(|pinned| pinned == sequence)
245    }
246}
247
248fn sample_priority(sequence: u64) -> u64 {
249    let mut value = sequence.wrapping_add(0x9e3779b97f4a7c15);
250    value = (value ^ (value >> 30)).wrapping_mul(0xbf58476d1ce4e5b9);
251    value = (value ^ (value >> 27)).wrapping_mul(0x94d049bb133111eb);
252    value ^ (value >> 31)
253}
254
255#[cfg(test)]
256mod tests {
257    use super::*;
258    use chrono::{Duration, NaiveDate};
259
260    fn point(sequence: u64, equity: f64, drawdown: f64) -> EquityPoint {
261        EquityPoint {
262            ts: NaiveDate::from_ymd_opt(2026, 1, 1)
263                .unwrap()
264                .and_hms_opt(0, 0, 0)
265                .unwrap()
266                + Duration::seconds(sequence as i64),
267            equity: Some(equity),
268            drawdown: Some(drawdown),
269            ..EquityPoint::default()
270        }
271    }
272
273    #[test]
274    fn policy_defaults_and_validates_bounded_limits() {
275        assert_eq!(
276            MtmOutputPolicy::default(),
277            MtmOutputPolicy::Bounded {
278                max_points: DEFAULT_MTM_MAX_POINTS
279            }
280        );
281        assert!(
282            MtmOutputPolicy::Bounded {
283                max_points: MIN_MTM_MAX_POINTS
284            }
285            .validate()
286            .is_ok()
287        );
288        assert!(
289            MtmOutputPolicy::Bounded {
290                max_points: MAX_MTM_MAX_POINTS
291            }
292            .validate()
293            .is_ok()
294        );
295        assert!(matches!(
296            MtmOutputPolicy::Bounded { max_points: 7 }.validate(),
297            Err(MtmOutputPolicyError::InvalidMaxPoints { max_points: 7 })
298        ));
299        assert!(
300            serde_json::from_str::<MtmOutputPolicy>(r#"{"bounded":{"max_points":16385}}"#).is_err()
301        );
302    }
303
304    #[test]
305    fn none_and_full_policies_report_exact_counts() {
306        let mut none = MtmCurveCollector::new(MtmOutputPolicy::None).unwrap();
307        let mut full = MtmCurveCollector::new(MtmOutputPolicy::Full).unwrap();
308        for sequence in 0..3 {
309            let point = point(sequence, 100.0 + sequence as f64, 0.0);
310            none.push(point.clone());
311            full.push(point);
312        }
313
314        let (none_curve, none_summary) = none.into_parts();
315        assert!(none_curve.is_empty());
316        assert_eq!(none_summary.observed_points, 3);
317        assert_eq!(none_summary.retained_points, 0);
318        assert_eq!(none_summary.omitted_points, 3);
319
320        let (full_curve, full_summary) = full.into_parts();
321        assert_eq!(full_curve.len(), 3);
322        assert_eq!(full_summary.retained_points, 3);
323        assert_eq!(full_summary.omitted_points, 0);
324        assert_eq!(full_curve[2].observation_sequence, Some(2));
325    }
326
327    #[test]
328    fn bounded_output_is_deterministic_and_pins_significant_points() {
329        let policy = MtmOutputPolicy::Bounded { max_points: 8 };
330        let mut first = MtmCurveCollector::new(policy).unwrap();
331        let mut second = MtmCurveCollector::new(policy).unwrap();
332
333        for sequence in 0..100 {
334            let mut equity = 100.0 + (sequence % 7) as f64;
335            let mut drawdown = 0.0;
336            if sequence == 10 {
337                equity = 1_000.0;
338            } else if sequence == 20 {
339                equity = -1_000.0;
340            } else if sequence == 30 {
341                drawdown = 500.0;
342            }
343            let point = point(sequence, equity, drawdown);
344            first.push(point.clone());
345            second.push(point);
346        }
347
348        assert_eq!(first.eviction_index.len(), 8);
349        assert_eq!(second.eviction_index.len(), 8);
350        let (first_curve, summary) = first.into_parts();
351        let second_curve = second.into_curve();
352        assert_eq!(first_curve, second_curve);
353        assert_eq!(first_curve.len(), 8);
354        assert_eq!(summary.observed_points, 100);
355        assert_eq!(summary.retained_points, 8);
356        assert_eq!(summary.omitted_points, 92);
357
358        let sequences: Vec<_> = first_curve
359            .iter()
360            .filter_map(|point| point.observation_sequence)
361            .collect();
362        assert_eq!(sequences, vec![0, 10, 20, 21, 30, 48, 68, 99]);
363        for pinned in [0, 10, 20, 30, 99] {
364            assert!(sequences.contains(&pinned), "missing pinned point {pinned}");
365        }
366        assert!(sequences.windows(2).all(|pair| pair[0] < pair[1]));
367    }
368}