1use 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#[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#[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#[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#[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}