1use crate::types::{MarketState, Observation};
2use crate::portfolio::Portfolio;
3
4#[derive(Debug, Clone)]
6pub struct ObservationFilter {
7 pub pass_volatility: bool,
8 pub pass_drift: bool,
9 pub pass_all_parameters: bool,
10 pub allowed_parameters: Vec<String>,
11}
12
13impl ObservationFilter {
14 pub fn transparent() -> Self {
16 Self {
17 pass_volatility: true,
18 pass_drift: true,
19 pass_all_parameters: true,
20 allowed_parameters: Vec::new(),
21 }
22 }
23
24 pub fn opaque() -> Self {
26 Self {
27 pass_volatility: false,
28 pass_drift: false,
29 pass_all_parameters: false,
30 allowed_parameters: Vec::new(),
31 }
32 }
33
34 pub fn build(&self, state: &MarketState, portfolio: &Portfolio) -> Observation {
36 let volatility = if self.pass_volatility { state.true_volatility } else { None };
37 let drift = if self.pass_drift { state.true_drift } else { None };
38
39 let parameters = if self.pass_all_parameters {
40 state.parameters.clone()
41 } else if !self.allowed_parameters.is_empty() {
42 state.parameters.as_ref().and_then(|p| {
43 let filtered: std::collections::HashMap<String, f64> = p.iter()
44 .filter(|(k, _)| self.allowed_parameters.contains(k))
45 .map(|(k, v)| (k.clone(), *v))
46 .collect();
47 if filtered.is_empty() { None } else { Some(filtered) }
48 })
49 } else {
50 None
51 };
52
53 Observation {
54 timestamp: state.timestamp,
55 best_bid: state.best_bid,
56 best_ask: state.best_ask,
57 last_price: state.last_price,
58 portfolio: portfolio.snapshot(),
59 volatility,
60 drift,
61 parameters,
62 }
63 }
64}
65
66impl Default for ObservationFilter {
67 fn default() -> Self {
68 Self::transparent()
69 }
70}
71
72#[cfg(test)]
73mod tests {
74 use super::*;
75 use std::collections::HashMap;
76
77 fn sample_state() -> MarketState {
78 let mut params = HashMap::new();
79 params.insert("kappa".to_string(), 2.0);
80 params.insert("theta".to_string(), 0.04);
81
82 MarketState {
83 timestamp: 0.0,
84 best_bid: 99.95,
85 best_ask: 100.05,
86 last_price: Some(100.0),
87 true_volatility: Some(0.2),
88 true_drift: Some(0.05),
89 parameters: Some(params),
90 }
91 }
92
93 fn sample_portfolio() -> Portfolio {
94 let mut p = Portfolio::new(10000.0, 0.001);
95 p.position = 5.0;
96 p
97 }
98
99 #[test]
100 fn test_transparent_passes_everything() {
101 let filter = ObservationFilter::transparent();
102 let obs = filter.build(&sample_state(), &sample_portfolio());
103
104 assert_eq!(obs.best_bid, 99.95);
105 assert_eq!(obs.best_ask, 100.05);
106 assert_eq!(obs.portfolio.cash, 10000.0);
107 assert_eq!(obs.portfolio.position, 5.0);
108 assert!(obs.volatility.is_some());
109 assert_eq!(obs.volatility.unwrap(), 0.2);
110 assert!(obs.drift.is_some());
111 assert!(obs.parameters.is_some());
112 assert_eq!(obs.parameters.unwrap().len(), 2);
113 }
114
115 #[test]
116 fn test_opaque_blocks_ground_truth() {
117 let filter = ObservationFilter::opaque();
118 let obs = filter.build(&sample_state(), &sample_portfolio());
119
120 assert_eq!(obs.best_bid, 99.95);
121 assert_eq!(obs.best_ask, 100.05);
122 assert_eq!(obs.portfolio.cash, 10000.0);
123 assert_eq!(obs.portfolio.position, 5.0);
124 assert!(obs.volatility.is_none());
125 assert!(obs.drift.is_none());
126 assert!(obs.parameters.is_none());
127 }
128
129 #[test]
130 fn test_partial_filter_with_whitelist() {
131 let filter = ObservationFilter {
132 pass_volatility: true,
133 pass_drift: false,
134 pass_all_parameters: false,
135 allowed_parameters: vec!["kappa".to_string()],
136 };
137 let obs = filter.build(&sample_state(), &sample_portfolio());
138
139 assert!(obs.volatility.is_some());
140 assert!(obs.drift.is_none());
141 let params = obs.parameters.unwrap();
142 assert_eq!(params.len(), 1);
143 assert!(params.contains_key("kappa"));
144 assert!(!params.contains_key("theta"));
145 }
146
147 #[test]
148 fn test_mid_price() {
149 let filter = ObservationFilter::transparent();
150 let obs = filter.build(&sample_state(), &sample_portfolio());
151 assert!((obs.mid_price() - 100.0).abs() < 1e-10);
152 }
153
154 #[test]
155 fn test_default_is_transparent() {
156 let filter = ObservationFilter::default();
157 assert!(filter.pass_volatility);
158 assert!(filter.pass_drift);
159 assert!(filter.pass_all_parameters);
160 }
161}