Skip to main content

engine/
observation.rs

1use crate::types::{MarketState, Observation};
2use crate::portfolio::Portfolio;
3
4/// Controls what ground truth data from MarketState is exposed to an agent
5#[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    /// Pass all ground truth data (transparent lens)
15    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    /// Block all ground truth data (opaque lens)
25    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    /// Build an Observation from ground truth MarketState and agent Portfolio
35    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}