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