nautilus_serialization/arrow/display/
account_state.rs1use std::sync::Arc;
19
20use arrow::{
21 array::{BooleanBuilder, StringBuilder, TimestampNanosecondBuilder},
22 datatypes::Schema,
23 error::ArrowError,
24 record_batch::RecordBatch,
25};
26use nautilus_model::events::AccountState;
27
28use super::{bool_field, timestamp_field, unix_nanos_to_i64, utf8_field};
29use crate::arrow::timestamp_data_type;
30
31#[must_use]
33pub fn account_state_schema() -> Schema {
34 Schema::new(vec![
35 utf8_field("account_id", false),
36 utf8_field("account_type", false),
37 utf8_field("base_currency", true),
38 utf8_field("balances", false),
39 utf8_field("margins", false),
40 bool_field("is_reported", false),
41 utf8_field("event_id", false),
42 timestamp_field("ts_event", false),
43 timestamp_field("ts_init", false),
44 ])
45}
46
47fn balances_to_json(state: &AccountState) -> String {
48 let entries: Vec<serde_json::Value> = state
49 .balances
50 .iter()
51 .map(|b| {
52 serde_json::json!({
53 "currency": b.currency.to_string(),
54 "total": b.total.as_f64(),
55 "locked": b.locked.as_f64(),
56 "free": b.free.as_f64(),
57 })
58 })
59 .collect();
60 serde_json::to_string(&entries).unwrap_or_default()
61}
62
63fn margins_to_json(state: &AccountState) -> String {
64 let entries: Vec<serde_json::Value> = state
65 .margins
66 .iter()
67 .map(|m| {
68 serde_json::json!({
69 "instrument_id": m.instrument_id.map(|id| id.to_string()),
70 "currency": m.currency.to_string(),
71 "initial": m.initial.as_f64(),
72 "maintenance": m.maintenance.as_f64(),
73 })
74 })
75 .collect();
76 serde_json::to_string(&entries).unwrap_or_default()
77}
78
79pub fn encode_account_states(data: &[AccountState]) -> Result<RecordBatch, ArrowError> {
92 let mut account_id = StringBuilder::new();
93 let mut account_type = StringBuilder::new();
94 let mut base_currency = StringBuilder::new();
95 let mut balances = StringBuilder::new();
96 let mut margins = StringBuilder::new();
97 let mut is_reported = BooleanBuilder::with_capacity(data.len());
98 let mut event_id = StringBuilder::new();
99 let mut ts_event =
100 TimestampNanosecondBuilder::with_capacity(data.len()).with_data_type(timestamp_data_type());
101 let mut ts_init =
102 TimestampNanosecondBuilder::with_capacity(data.len()).with_data_type(timestamp_data_type());
103
104 for state in data {
105 account_id.append_value(state.account_id);
106 account_type.append_value(format!("{}", state.account_type));
107 base_currency.append_option(state.base_currency.map(|v| v.to_string()));
108 balances.append_value(balances_to_json(state));
109 margins.append_value(margins_to_json(state));
110 is_reported.append_value(state.is_reported);
111 event_id.append_value(state.event_id.to_string());
112 ts_event.append_value(unix_nanos_to_i64(state.ts_event.as_u64()));
113 ts_init.append_value(unix_nanos_to_i64(state.ts_init.as_u64()));
114 }
115
116 RecordBatch::try_new(
117 Arc::new(account_state_schema()),
118 vec![
119 Arc::new(account_id.finish()),
120 Arc::new(account_type.finish()),
121 Arc::new(base_currency.finish()),
122 Arc::new(balances.finish()),
123 Arc::new(margins.finish()),
124 Arc::new(is_reported.finish()),
125 Arc::new(event_id.finish()),
126 Arc::new(ts_event.finish()),
127 Arc::new(ts_init.finish()),
128 ],
129 )
130}
131
132#[cfg(test)]
133mod tests {
134 use arrow::{
135 array::{Array, BooleanArray, StringArray, TimestampNanosecondArray},
136 datatypes::{DataType, TimeUnit},
137 };
138 use nautilus_core::UUID4;
139 use nautilus_model::{
140 enums::AccountType,
141 identifiers::AccountId,
142 types::{AccountBalance, Currency, Money},
143 };
144 use rstest::rstest;
145
146 use super::*;
147
148 fn make_account_state(ts: u64) -> AccountState {
149 let currency = Currency::USD();
150 let balance = AccountBalance::new(
151 Money::new(10_000.0, currency),
152 Money::new(1_000.0, currency),
153 Money::new(9_000.0, currency),
154 );
155 AccountState {
156 account_id: AccountId::from("SIM-001"),
157 account_type: AccountType::Cash,
158 base_currency: Some(currency),
159 balances: vec![balance],
160 margins: vec![],
161 is_reported: false,
162 event_id: UUID4::default(),
163 ts_event: ts.into(),
164 ts_init: (ts + 1).into(),
165 info: None,
166 }
167 }
168
169 #[rstest]
170 fn test_encode_account_states_schema() {
171 let batch = encode_account_states(&[]).unwrap();
172 let schema = batch.schema();
173 let fields = schema.fields();
174 assert_eq!(fields.len(), 9);
175 assert_eq!(fields[0].name(), "account_id");
176 assert_eq!(fields[0].data_type(), &DataType::Utf8);
177 assert_eq!(fields[5].name(), "is_reported");
178 assert_eq!(fields[5].data_type(), &DataType::Boolean);
179 assert_eq!(fields[7].name(), "ts_event");
180 assert_eq!(
181 fields[7].data_type(),
182 &DataType::Timestamp(TimeUnit::Nanosecond, Some("UTC".into()))
183 );
184 }
185
186 #[rstest]
187 fn test_encode_account_states_values() {
188 let states = vec![make_account_state(1_000_000)];
189 let batch = encode_account_states(&states).unwrap();
190
191 assert_eq!(batch.num_rows(), 1);
192
193 let account_id_col = batch
194 .column(0)
195 .as_any()
196 .downcast_ref::<StringArray>()
197 .unwrap();
198 let is_reported_col = batch
199 .column(5)
200 .as_any()
201 .downcast_ref::<BooleanArray>()
202 .unwrap();
203 let ts_event_col = batch
204 .column(7)
205 .as_any()
206 .downcast_ref::<TimestampNanosecondArray>()
207 .unwrap();
208 let balances_col = batch
209 .column(3)
210 .as_any()
211 .downcast_ref::<StringArray>()
212 .unwrap();
213
214 assert_eq!(account_id_col.value(0), "SIM-001");
215 assert!(!is_reported_col.value(0));
216 assert_eq!(ts_event_col.value(0), 1_000_000);
217
218 let balances: Vec<serde_json::Value> = serde_json::from_str(balances_col.value(0)).unwrap();
219 assert_eq!(balances.len(), 1);
220 assert_eq!(balances[0]["currency"], "USD");
221 assert!((balances[0]["total"].as_f64().unwrap() - 10_000.0).abs() < 1e-9);
222 assert!((balances[0]["locked"].as_f64().unwrap() - 1_000.0).abs() < 1e-9);
223 assert!((balances[0]["free"].as_f64().unwrap() - 9_000.0).abs() < 1e-9);
224 }
225
226 #[rstest]
227 fn test_encode_account_states_empty() {
228 let batch = encode_account_states(&[]).unwrap();
229 assert_eq!(batch.num_rows(), 0);
230 assert_eq!(batch.schema().fields().len(), 9);
231 }
232
233 #[rstest]
234 fn test_encode_account_states_null_base_currency() {
235 let mut state = make_account_state(1_000);
236 state.base_currency = None;
237 let batch = encode_account_states(&[state]).unwrap();
238
239 let base_currency_col = batch
240 .column(2)
241 .as_any()
242 .downcast_ref::<StringArray>()
243 .unwrap();
244 assert!(base_currency_col.is_null(0));
245 }
246}