Skip to main content

nautilus_serialization/arrow/instrument/
equity.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//! Arrow serialization for Equity instruments.
17
18use std::{borrow::Borrow, collections::HashMap, str::FromStr, sync::Arc};
19
20use arrow::{
21    array::{Array, StringArray, StringBuilder, UInt8Array, UInt64Array},
22    datatypes::{DataType, Field, Schema},
23    error::ArrowError,
24    record_batch::RecordBatch,
25};
26use nautilus_core::Params;
27use nautilus_model::{
28    identifiers::{InstrumentId, Symbol},
29    instruments::equity::Equity,
30    types::{price::Price, quantity::Quantity},
31};
32use rust_decimal::Decimal;
33use ustr::Ustr;
34
35use super::KEY_CLASS;
36use crate::arrow::{
37    ArrowSchemaProvider, EncodeToRecordBatch, EncodingError, KEY_INSTRUMENT_ID,
38    KEY_PRICE_PRECISION, extract_column, extract_column_by_name_or_index,
39    extract_optional_string_column_by_name, json_string_field, optional_ustr_value,
40    record_batch_with_timestamps, record_batch_with_u64_timestamps, timestamp_data_type,
41};
42
43impl ArrowSchemaProvider for Equity {
44    fn get_schema(metadata: Option<HashMap<String, String>>) -> Schema {
45        let fields = vec![
46            Field::new("id", DataType::Utf8, false),
47            Field::new("raw_symbol", DataType::Utf8, false),
48            Field::new("currency", DataType::Utf8, false),
49            Field::new("price_precision", DataType::UInt8, false),
50            Field::new("price_increment", DataType::Utf8, false),
51            Field::new("lot_size", DataType::Utf8, true), // nullable
52            Field::new("isin", DataType::Utf8, true),     // nullable
53            Field::new("max_quantity", DataType::Utf8, true), // nullable
54            Field::new("min_quantity", DataType::Utf8, true), // nullable
55            Field::new("max_price", DataType::Utf8, true), // nullable
56            Field::new("min_price", DataType::Utf8, true), // nullable
57            Field::new("margin_init", DataType::Utf8, false),
58            Field::new("margin_maint", DataType::Utf8, false),
59            Field::new("maker_fee", DataType::Utf8, false),
60            Field::new("taker_fee", DataType::Utf8, false),
61            Field::new("tick_scheme", DataType::Utf8, true),
62            json_string_field("info", true),
63            Field::new("ts_event", timestamp_data_type(), false),
64            Field::new("ts_init", timestamp_data_type(), false),
65        ];
66
67        let mut final_metadata = HashMap::new();
68        final_metadata.insert(KEY_CLASS.to_string(), "Equity".to_string());
69
70        if let Some(meta) = metadata {
71            final_metadata.extend(meta);
72        }
73
74        Schema::new_with_metadata(fields, final_metadata)
75    }
76}
77
78impl EncodeToRecordBatch for Equity {
79    fn encode_batch<T>(
80        #[allow(unused)] metadata: &HashMap<String, String>,
81        data: &[T],
82    ) -> Result<RecordBatch, ArrowError>
83    where
84        T: std::borrow::Borrow<Self>,
85    {
86        let mut id_builder = StringBuilder::new();
87        let mut raw_symbol_builder = StringBuilder::new();
88        let mut currency_builder = StringBuilder::new();
89        let mut price_precision_builder = UInt8Array::builder(data.len());
90        let mut price_increment_builder = StringBuilder::new();
91        let mut lot_size_builder = StringBuilder::new();
92        let mut isin_builder = StringBuilder::new();
93        let mut max_quantity_builder = StringBuilder::new();
94        let mut min_quantity_builder = StringBuilder::new();
95        let mut max_price_builder = StringBuilder::new();
96        let mut min_price_builder = StringBuilder::new();
97        let mut margin_init_builder = StringBuilder::new();
98        let mut margin_maint_builder = StringBuilder::new();
99        let mut maker_fee_builder = StringBuilder::new();
100        let mut taker_fee_builder = StringBuilder::new();
101        let mut tick_scheme_builder = StringBuilder::new();
102        let mut info_builder = StringBuilder::new();
103        let mut ts_event_builder = UInt64Array::builder(data.len());
104        let mut ts_init_builder = UInt64Array::builder(data.len());
105
106        for equity in data.iter().map(Borrow::borrow) {
107            id_builder.append_value(equity.id.to_string());
108            raw_symbol_builder.append_value(equity.raw_symbol);
109            currency_builder.append_value(equity.currency.to_string());
110            price_precision_builder.append_value(equity.price_precision);
111            price_increment_builder.append_value(equity.price_increment.to_string());
112
113            if let Some(lot_size) = equity.lot_size {
114                lot_size_builder.append_value(lot_size.to_string());
115            } else {
116                lot_size_builder.append_null();
117            }
118
119            if let Some(isin) = equity.isin {
120                isin_builder.append_value(isin);
121            } else {
122                isin_builder.append_null();
123            }
124
125            if let Some(max_qty) = equity.max_quantity {
126                max_quantity_builder.append_value(max_qty.to_string());
127            } else {
128                max_quantity_builder.append_null();
129            }
130
131            if let Some(min_qty) = equity.min_quantity {
132                min_quantity_builder.append_value(min_qty.to_string());
133            } else {
134                min_quantity_builder.append_null();
135            }
136
137            if let Some(max_p) = equity.max_price {
138                max_price_builder.append_value(max_p.to_string());
139            } else {
140                max_price_builder.append_null();
141            }
142
143            if let Some(min_p) = equity.min_price {
144                min_price_builder.append_value(min_p.to_string());
145            } else {
146                min_price_builder.append_null();
147            }
148
149            margin_init_builder.append_value(equity.margin_init.to_string());
150            margin_maint_builder.append_value(equity.margin_maint.to_string());
151            maker_fee_builder.append_value(equity.maker_fee.to_string());
152            taker_fee_builder.append_value(equity.taker_fee.to_string());
153
154            if let Some(tick_scheme) = equity.tick_scheme {
155                tick_scheme_builder.append_value(tick_scheme);
156            } else {
157                tick_scheme_builder.append_null();
158            }
159
160            if let Some(ref info) = equity.info {
161                match serde_json::to_string(info) {
162                    Ok(json) => {
163                        info_builder.append_value(json);
164                    }
165                    Err(e) => {
166                        return Err(ArrowError::InvalidArgumentError(format!(
167                            "Failed to serialize info dict to JSON: {e}"
168                        )));
169                    }
170                }
171            } else {
172                info_builder.append_null();
173            }
174
175            ts_event_builder.append_value(equity.ts_event.as_u64());
176            ts_init_builder.append_value(equity.ts_init.as_u64());
177        }
178
179        let mut final_metadata = metadata.clone();
180        final_metadata.insert(KEY_CLASS.to_string(), "Equity".to_string());
181
182        record_batch_with_timestamps(
183            Self::get_schema(Some(final_metadata)).into(),
184            vec![
185                Arc::new(id_builder.finish()),
186                Arc::new(raw_symbol_builder.finish()),
187                Arc::new(currency_builder.finish()),
188                Arc::new(price_precision_builder.finish()),
189                Arc::new(price_increment_builder.finish()),
190                Arc::new(lot_size_builder.finish()),
191                Arc::new(isin_builder.finish()),
192                Arc::new(max_quantity_builder.finish()),
193                Arc::new(min_quantity_builder.finish()),
194                Arc::new(max_price_builder.finish()),
195                Arc::new(min_price_builder.finish()),
196                Arc::new(margin_init_builder.finish()),
197                Arc::new(margin_maint_builder.finish()),
198                Arc::new(maker_fee_builder.finish()),
199                Arc::new(taker_fee_builder.finish()),
200                Arc::new(tick_scheme_builder.finish()),
201                Arc::new(info_builder.finish()),
202                Arc::new(ts_event_builder.finish()),
203                Arc::new(ts_init_builder.finish()),
204            ],
205        )
206    }
207
208    fn metadata(&self) -> HashMap<String, String> {
209        let mut metadata = HashMap::new();
210        metadata.insert(KEY_INSTRUMENT_ID.to_string(), self.id.to_string());
211        metadata.insert(
212            KEY_PRICE_PRECISION.to_string(),
213            self.price_precision.to_string(),
214        );
215        metadata
216    }
217}
218
219/// Decodes [`Equity`] instruments from a record batch.
220///
221/// Not a [`DecodeFromRecordBatch`] implementation because that trait requires `Into<Data>`.
222///
223/// # Errors
224///
225/// Returns an `EncodingError` if the record batch cannot be decoded.
226///
227/// [`DecodeFromRecordBatch`]: crate::arrow::DecodeFromRecordBatch
228pub fn decode_equity_batch(
229    #[allow(unused)] metadata: &HashMap<String, String>,
230    record_batch: &RecordBatch,
231) -> Result<Vec<Equity>, EncodingError> {
232    let record_batch = record_batch_with_u64_timestamps(record_batch)?;
233    let record_batch = &record_batch;
234    let cols = record_batch.columns();
235    let num_rows = record_batch.num_rows();
236
237    // Read precision from data columns (it's in the schema)
238    let id_values = extract_column::<StringArray>(cols, "id", 0, DataType::Utf8)?;
239    let raw_symbol_values = extract_column::<StringArray>(cols, "raw_symbol", 1, DataType::Utf8)?;
240    let currency_values = extract_column::<StringArray>(cols, "currency", 2, DataType::Utf8)?;
241    let price_precision_values =
242        extract_column::<UInt8Array>(cols, "price_precision", 3, DataType::UInt8)?;
243    let price_increment_values =
244        extract_column::<StringArray>(cols, "price_increment", 4, DataType::Utf8)?;
245    let lot_size_values = cols
246        .get(5)
247        .ok_or_else(|| EncodingError::MissingColumn("lot_size", 5))?;
248    let isin_values = cols
249        .get(6)
250        .ok_or_else(|| EncodingError::MissingColumn("isin", 6))?;
251    let max_quantity_values = cols
252        .get(7)
253        .ok_or_else(|| EncodingError::MissingColumn("max_quantity", 7))?;
254    let min_quantity_values = cols
255        .get(8)
256        .ok_or_else(|| EncodingError::MissingColumn("min_quantity", 8))?;
257    let max_price_values = cols
258        .get(9)
259        .ok_or_else(|| EncodingError::MissingColumn("max_price", 9))?;
260    let min_price_values = cols
261        .get(10)
262        .ok_or_else(|| EncodingError::MissingColumn("min_price", 10))?;
263    let margin_init_values =
264        extract_column::<StringArray>(cols, "margin_init", 11, DataType::Utf8)?;
265    let margin_maint_values =
266        extract_column::<StringArray>(cols, "margin_maint", 12, DataType::Utf8)?;
267    let maker_fee_values = extract_column::<StringArray>(cols, "maker_fee", 13, DataType::Utf8)?;
268    let taker_fee_values = extract_column::<StringArray>(cols, "taker_fee", 14, DataType::Utf8)?;
269    let tick_scheme_values = extract_optional_string_column_by_name(record_batch, "tick_scheme")?;
270    let info_values =
271        extract_column_by_name_or_index::<StringArray>(record_batch, "info", 15, DataType::Utf8)?;
272    let ts_event_values = extract_column_by_name_or_index::<UInt64Array>(
273        record_batch,
274        "ts_event",
275        16,
276        DataType::UInt64,
277    )?;
278    let ts_init_values = extract_column_by_name_or_index::<UInt64Array>(
279        record_batch,
280        "ts_init",
281        17,
282        DataType::UInt64,
283    )?;
284
285    let mut result = Vec::with_capacity(num_rows);
286
287    for i in 0..num_rows {
288        let id = InstrumentId::from_str(id_values.value(i))
289            .map_err(|e| EncodingError::ParseError("id", format!("row {i}: {e}")))?;
290        let raw_symbol = Symbol::from(raw_symbol_values.value(i));
291        let currency =
292            super::decode_currency(currency_values.value(i), "currency", "equity.currency", i)?;
293        let price_prec = price_precision_values.value(i);
294
295        let price_increment = Price::from_str(price_increment_values.value(i))
296            .map_err(|e| EncodingError::ParseError("price_increment", format!("row {i}: {e}")))?;
297
298        let lot_size = if lot_size_values.is_null(i) {
299            None
300        } else {
301            let lot_size_str = lot_size_values
302                .as_any()
303                .downcast_ref::<StringArray>()
304                .ok_or_else(|| {
305                    EncodingError::ParseError("lot_size", format!("row {i}: invalid type"))
306                })?
307                .value(i);
308            Some(
309                Quantity::from_str(lot_size_str)
310                    .map_err(|e| EncodingError::ParseError("lot_size", format!("row {i}: {e}")))?,
311            )
312        };
313
314        let isin = if isin_values.is_null(i) {
315            None
316        } else {
317            let isin_str = isin_values
318                .as_any()
319                .downcast_ref::<StringArray>()
320                .ok_or_else(|| EncodingError::ParseError("isin", format!("row {i}: invalid type")))?
321                .value(i);
322            Some(Ustr::from(isin_str))
323        };
324
325        let max_quantity =
326            if max_quantity_values.is_null(i) {
327                None
328            } else {
329                let max_qty_str = max_quantity_values
330                    .as_any()
331                    .downcast_ref::<StringArray>()
332                    .ok_or_else(|| {
333                        EncodingError::ParseError("max_quantity", format!("row {i}: invalid type"))
334                    })?
335                    .value(i);
336                Some(Quantity::from_str(max_qty_str).map_err(|e| {
337                    EncodingError::ParseError("max_quantity", format!("row {i}: {e}"))
338                })?)
339            };
340
341        let min_quantity =
342            if min_quantity_values.is_null(i) {
343                None
344            } else {
345                let min_qty_str = min_quantity_values
346                    .as_any()
347                    .downcast_ref::<StringArray>()
348                    .ok_or_else(|| {
349                        EncodingError::ParseError("min_quantity", format!("row {i}: invalid type"))
350                    })?
351                    .value(i);
352                Some(Quantity::from_str(min_qty_str).map_err(|e| {
353                    EncodingError::ParseError("min_quantity", format!("row {i}: {e}"))
354                })?)
355            };
356
357        let max_price = if max_price_values.is_null(i) {
358            None
359        } else {
360            let max_p_str = max_price_values
361                .as_any()
362                .downcast_ref::<StringArray>()
363                .ok_or_else(|| {
364                    EncodingError::ParseError("max_price", format!("row {i}: invalid type"))
365                })?
366                .value(i);
367            Some(
368                Price::from_str(max_p_str)
369                    .map_err(|e| EncodingError::ParseError("max_price", format!("row {i}: {e}")))?,
370            )
371        };
372
373        let min_price = if min_price_values.is_null(i) {
374            None
375        } else {
376            let min_p_str = min_price_values
377                .as_any()
378                .downcast_ref::<StringArray>()
379                .ok_or_else(|| {
380                    EncodingError::ParseError("min_price", format!("row {i}: invalid type"))
381                })?
382                .value(i);
383            Some(
384                Price::from_str(min_p_str)
385                    .map_err(|e| EncodingError::ParseError("min_price", format!("row {i}: {e}")))?,
386            )
387        };
388
389        let margin_init = Decimal::from_str(margin_init_values.value(i))
390            .map_err(|e| EncodingError::ParseError("margin_init", format!("row {i}: {e}")))?;
391        let margin_maint = Decimal::from_str(margin_maint_values.value(i))
392            .map_err(|e| EncodingError::ParseError("margin_maint", format!("row {i}: {e}")))?;
393        let maker_fee = Decimal::from_str(maker_fee_values.value(i))
394            .map_err(|e| EncodingError::ParseError("maker_fee", format!("row {i}: {e}")))?;
395        let taker_fee = Decimal::from_str(taker_fee_values.value(i))
396            .map_err(|e| EncodingError::ParseError("taker_fee", format!("row {i}: {e}")))?;
397
398        let info = if info_values.is_null(i) {
399            None
400        } else {
401            let info_json = info_values
402                .as_any()
403                .downcast_ref::<StringArray>()
404                .ok_or_else(|| EncodingError::ParseError("info", format!("row {i}: invalid type")))?
405                .value(i);
406
407            match serde_json::from_str::<Params>(info_json) {
408                Ok(info_dict) => Some(info_dict),
409                Err(e) => {
410                    return Err(EncodingError::ParseError(
411                        "info",
412                        format!("row {i}: failed to deserialize JSON: {e}"),
413                    ));
414                }
415            }
416        };
417
418        let ts_event = nautilus_core::UnixNanos::from(ts_event_values.value(i));
419        let ts_init = nautilus_core::UnixNanos::from(ts_init_values.value(i));
420
421        let tick_scheme = optional_ustr_value(tick_scheme_values, i);
422
423        let equity = Equity::builder()
424            .instrument_id(id)
425            .raw_symbol(raw_symbol)
426            .maybe_isin(isin)
427            .currency(currency)
428            .price_precision(price_prec)
429            .price_increment(price_increment)
430            .maybe_lot_size(lot_size)
431            .maybe_max_quantity(max_quantity)
432            .maybe_min_quantity(min_quantity)
433            .maybe_max_price(max_price)
434            .maybe_min_price(min_price)
435            .margin_init(margin_init)
436            .margin_maint(margin_maint)
437            .maker_fee(maker_fee)
438            .taker_fee(taker_fee)
439            .maybe_tick_scheme(tick_scheme)
440            .maybe_info(info)
441            .ts_event(ts_event)
442            .ts_init(ts_init)
443            .build()
444            .map_err(|e| super::instrument_validation_error::<Equity>(i, e))?;
445
446        result.push(equity);
447    }
448
449    Ok(result)
450}