Skip to main content

nautilus_serialization/arrow/display/
account_state.rs

1// -------------------------------------------------------------------------------------------------
2//  Copyright (C) 2015-2026 Nautech Systems Pty Ltd. All rights reserved.
3//  https://nautechsystems.io
4//
5//  Licensed under the GNU Lesser General Public License Version 3.0 (the "License");
6//  You may not use this file except in compliance with the License.
7//  You may obtain a copy of the License at https://www.gnu.org/licenses/lgpl-3.0.en.html
8//
9//  Unless required by applicable law or agreed to in writing, software
10//  distributed under the License is distributed on an "AS IS" BASIS,
11//  WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12//  See the License for the specific language governing permissions and
13//  limitations under the License.
14// -------------------------------------------------------------------------------------------------
15
16//! Display-mode Arrow encoder for [`AccountState`].
17
18use 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/// Returns the display-mode Arrow schema for [`AccountState`].
32#[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
79/// Encodes account state snapshots as a display-friendly Arrow [`RecordBatch`].
80///
81/// Emits `Utf8` columns for identifiers and JSON-serialized balances/margins,
82/// `Timestamp(Nanosecond)` columns for event and init times, and a `Boolean`
83/// column for `is_reported`. Balances and margins are serialized as JSON arrays
84/// with `f64` amounts for display readability.
85///
86/// Returns an empty [`RecordBatch`] with the correct schema when `data` is empty.
87///
88/// # Errors
89///
90/// Returns an [`ArrowError`] if the Arrow `RecordBatch` cannot be constructed.
91pub 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}