1use 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), 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),
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
219pub 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 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}