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, UnixNanos};
29use nautilus_model::{
30 enums::AssetClass,
31 identifiers::{InstrumentId, Symbol},
32 instruments::commodity::Commodity,
33 types::{money::Money, price::Price, quantity::Quantity},
34};
35use rust_decimal::Decimal;
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 Commodity {
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("asset_class", DataType::Utf8, false),
50 Field::new("quote_currency", DataType::Utf8, false),
51 Field::new("price_precision", DataType::UInt8, false),
52 Field::new("size_precision", DataType::UInt8, false),
53 Field::new("price_increment", DataType::Utf8, false),
54 Field::new("size_increment", DataType::Utf8, false),
55 Field::new("lot_size", DataType::Utf8, true), Field::new("max_quantity", DataType::Utf8, true), Field::new("min_quantity", DataType::Utf8, true), Field::new("max_notional", DataType::Utf8, true), Field::new("min_notional", DataType::Utf8, true), Field::new("max_price", DataType::Utf8, true), Field::new("min_price", DataType::Utf8, true), Field::new("margin_init", DataType::Utf8, false),
63 Field::new("margin_maint", DataType::Utf8, false),
64 Field::new("maker_fee", DataType::Utf8, false),
65 Field::new("taker_fee", DataType::Utf8, false),
66 Field::new("tick_scheme", DataType::Utf8, true),
67 Field::new("info", DataType::Binary, true), Field::new("ts_event", DataType::UInt64, false),
69 Field::new("ts_init", DataType::UInt64, false),
70 ];
71
72 let mut final_metadata = HashMap::new();
73 final_metadata.insert(KEY_CLASS.to_string(), "Commodity".to_string());
74
75 if let Some(meta) = metadata {
76 final_metadata.extend(meta);
77 }
78
79 Schema::new_with_metadata(fields, final_metadata)
80 }
81}
82
83impl EncodeToRecordBatch for Commodity {
84 fn encode_batch(
85 #[allow(unused)] metadata: &HashMap<String, String>,
86 data: &[Self],
87 ) -> Result<RecordBatch, ArrowError> {
88 let mut id_builder = StringBuilder::new();
89 let mut raw_symbol_builder = StringBuilder::new();
90 let mut asset_class_builder = StringBuilder::new();
91 let mut quote_currency_builder = StringBuilder::new();
92 let mut price_precision_builder = UInt8Array::builder(data.len());
93 let mut size_precision_builder = UInt8Array::builder(data.len());
94 let mut price_increment_builder = StringBuilder::new();
95 let mut size_increment_builder = StringBuilder::new();
96 let mut lot_size_builder = StringBuilder::new();
97 let mut max_quantity_builder = StringBuilder::new();
98 let mut min_quantity_builder = StringBuilder::new();
99 let mut max_notional_builder = StringBuilder::new();
100 let mut min_notional_builder = StringBuilder::new();
101 let mut max_price_builder = StringBuilder::new();
102 let mut min_price_builder = StringBuilder::new();
103 let mut margin_init_builder = StringBuilder::new();
104 let mut margin_maint_builder = StringBuilder::new();
105 let mut maker_fee_builder = StringBuilder::new();
106 let mut taker_fee_builder = StringBuilder::new();
107 let mut tick_scheme_builder = StringBuilder::new();
108 let mut info_builder = BinaryBuilder::new();
109 let mut ts_event_builder = UInt64Array::builder(data.len());
110 let mut ts_init_builder = UInt64Array::builder(data.len());
111
112 for commodity in data {
113 id_builder.append_value(commodity.id.to_string());
114 raw_symbol_builder.append_value(commodity.raw_symbol);
115 asset_class_builder.append_value(commodity.asset_class);
116 quote_currency_builder.append_value(commodity.quote_currency.to_string());
117 price_precision_builder.append_value(commodity.price_precision);
118 size_precision_builder.append_value(commodity.size_precision);
119 price_increment_builder.append_value(commodity.price_increment.to_string());
120 size_increment_builder.append_value(commodity.size_increment.to_string());
121
122 if let Some(lot_size) = commodity.lot_size {
123 lot_size_builder.append_value(lot_size.to_string());
124 } else {
125 lot_size_builder.append_null();
126 }
127
128 if let Some(max_qty) = commodity.max_quantity {
129 max_quantity_builder.append_value(max_qty.to_string());
130 } else {
131 max_quantity_builder.append_null();
132 }
133
134 if let Some(min_qty) = commodity.min_quantity {
135 min_quantity_builder.append_value(min_qty.to_string());
136 } else {
137 min_quantity_builder.append_null();
138 }
139
140 if let Some(max_not) = commodity.max_notional {
141 max_notional_builder.append_value(max_not.to_string());
142 } else {
143 max_notional_builder.append_null();
144 }
145
146 if let Some(min_not) = commodity.min_notional {
147 min_notional_builder.append_value(min_not.to_string());
148 } else {
149 min_notional_builder.append_null();
150 }
151
152 if let Some(max_p) = commodity.max_price {
153 max_price_builder.append_value(max_p.to_string());
154 } else {
155 max_price_builder.append_null();
156 }
157
158 if let Some(min_p) = commodity.min_price {
159 min_price_builder.append_value(min_p.to_string());
160 } else {
161 min_price_builder.append_null();
162 }
163
164 margin_init_builder.append_value(commodity.margin_init.to_string());
165 margin_maint_builder.append_value(commodity.margin_maint.to_string());
166 maker_fee_builder.append_value(commodity.maker_fee.to_string());
167 taker_fee_builder.append_value(commodity.taker_fee.to_string());
168
169 if let Some(tick_scheme) = commodity.tick_scheme {
170 tick_scheme_builder.append_value(tick_scheme);
171 } else {
172 tick_scheme_builder.append_null();
173 }
174
175 if let Some(ref info) = commodity.info {
177 match serde_json::to_vec(info) {
178 Ok(json_bytes) => {
179 info_builder.append_value(json_bytes);
180 }
181 Err(e) => {
182 return Err(ArrowError::InvalidArgumentError(format!(
183 "Failed to serialize info dict to JSON: {e}"
184 )));
185 }
186 }
187 } else {
188 info_builder.append_null();
189 }
190
191 ts_event_builder.append_value(commodity.ts_event.as_u64());
192 ts_init_builder.append_value(commodity.ts_init.as_u64());
193 }
194
195 let mut final_metadata = metadata.clone();
196 final_metadata.insert(KEY_CLASS.to_string(), "Commodity".to_string());
197
198 RecordBatch::try_new(
199 Self::get_schema(Some(final_metadata)).into(),
200 vec![
201 Arc::new(id_builder.finish()),
202 Arc::new(raw_symbol_builder.finish()),
203 Arc::new(asset_class_builder.finish()),
204 Arc::new(quote_currency_builder.finish()),
205 Arc::new(price_precision_builder.finish()),
206 Arc::new(size_precision_builder.finish()),
207 Arc::new(price_increment_builder.finish()),
208 Arc::new(size_increment_builder.finish()),
209 Arc::new(lot_size_builder.finish()),
210 Arc::new(max_quantity_builder.finish()),
211 Arc::new(min_quantity_builder.finish()),
212 Arc::new(max_notional_builder.finish()),
213 Arc::new(min_notional_builder.finish()),
214 Arc::new(max_price_builder.finish()),
215 Arc::new(min_price_builder.finish()),
216 Arc::new(margin_init_builder.finish()),
217 Arc::new(margin_maint_builder.finish()),
218 Arc::new(maker_fee_builder.finish()),
219 Arc::new(taker_fee_builder.finish()),
220 Arc::new(tick_scheme_builder.finish()),
221 Arc::new(info_builder.finish()),
222 Arc::new(ts_event_builder.finish()),
223 Arc::new(ts_init_builder.finish()),
224 ],
225 )
226 }
227
228 fn metadata(&self) -> HashMap<String, String> {
229 let mut metadata = HashMap::new();
230 metadata.insert(KEY_INSTRUMENT_ID.to_string(), self.id.to_string());
231 metadata.insert(
232 KEY_PRICE_PRECISION.to_string(),
233 self.price_precision.to_string(),
234 );
235 metadata
236 }
237}
238
239pub fn decode_commodity_batch(
249 #[allow(unused)] metadata: &HashMap<String, String>,
250 record_batch: &RecordBatch,
251) -> Result<Vec<Commodity>, EncodingError> {
252 let cols = record_batch.columns();
253 let num_rows = record_batch.num_rows();
254
255 let id_values = extract_column::<StringArray>(cols, "id", 0, DataType::Utf8)?;
256 let raw_symbol_values = extract_column::<StringArray>(cols, "raw_symbol", 1, DataType::Utf8)?;
257 let asset_class_values = extract_column::<StringArray>(cols, "asset_class", 2, DataType::Utf8)?;
258 let quote_currency_values =
259 extract_column::<StringArray>(cols, "quote_currency", 3, DataType::Utf8)?;
260 let price_precision_values =
261 extract_column::<UInt8Array>(cols, "price_precision", 4, DataType::UInt8)?;
262 let size_precision_values =
263 extract_column::<UInt8Array>(cols, "size_precision", 5, DataType::UInt8)?;
264 let price_increment_values =
265 extract_column::<StringArray>(cols, "price_increment", 6, DataType::Utf8)?;
266 let size_increment_values =
267 extract_column::<StringArray>(cols, "size_increment", 7, DataType::Utf8)?;
268 let lot_size_values = cols
269 .get(8)
270 .ok_or_else(|| EncodingError::MissingColumn("lot_size", 8))?;
271 let max_quantity_values = cols
272 .get(9)
273 .ok_or_else(|| EncodingError::MissingColumn("max_quantity", 9))?;
274 let min_quantity_values = cols
275 .get(10)
276 .ok_or_else(|| EncodingError::MissingColumn("min_quantity", 10))?;
277 let max_notional_values = cols
278 .get(11)
279 .ok_or_else(|| EncodingError::MissingColumn("max_notional", 11))?;
280 let min_notional_values = cols
281 .get(12)
282 .ok_or_else(|| EncodingError::MissingColumn("min_notional", 12))?;
283 let max_price_values = cols
284 .get(13)
285 .ok_or_else(|| EncodingError::MissingColumn("max_price", 13))?;
286 let min_price_values = cols
287 .get(14)
288 .ok_or_else(|| EncodingError::MissingColumn("min_price", 14))?;
289 let margin_init_values =
290 extract_column::<StringArray>(cols, "margin_init", 15, DataType::Utf8)?;
291 let margin_maint_values =
292 extract_column::<StringArray>(cols, "margin_maint", 16, DataType::Utf8)?;
293 let maker_fee_values = extract_column::<StringArray>(cols, "maker_fee", 17, DataType::Utf8)?;
294 let taker_fee_values = extract_column::<StringArray>(cols, "taker_fee", 18, DataType::Utf8)?;
295 let tick_scheme_values = extract_optional_string_column_by_name(record_batch, "tick_scheme")?;
296 let info_values =
297 extract_column_by_name_or_index::<BinaryArray>(record_batch, "info", 19, DataType::Binary)?;
298 let ts_event_values = extract_column_by_name_or_index::<UInt64Array>(
299 record_batch,
300 "ts_event",
301 20,
302 DataType::UInt64,
303 )?;
304 let ts_init_values = extract_column_by_name_or_index::<UInt64Array>(
305 record_batch,
306 "ts_init",
307 21,
308 DataType::UInt64,
309 )?;
310
311 let mut result = Vec::with_capacity(num_rows);
312
313 for i in 0..num_rows {
314 let id = InstrumentId::from_str(id_values.value(i))
315 .map_err(|e| EncodingError::ParseError("id", format!("row {i}: {e}")))?;
316 let raw_symbol = Symbol::from(raw_symbol_values.value(i));
317 let asset_class = AssetClass::from_str(asset_class_values.value(i))
318 .map_err(|e| EncodingError::ParseError("asset_class", format!("row {i}: {e}")))?;
319 let quote_currency = super::decode_currency(
320 quote_currency_values.value(i),
321 "quote_currency",
322 "commodity.quote_currency",
323 i,
324 )?;
325 let price_prec = price_precision_values.value(i);
326 let size_prec = size_precision_values.value(i);
327
328 let price_increment = Price::from_str(price_increment_values.value(i))
329 .map_err(|e| EncodingError::ParseError("price_increment", format!("row {i}: {e}")))?;
330 let size_increment = Quantity::from_str(size_increment_values.value(i))
331 .map_err(|e| EncodingError::ParseError("size_increment", format!("row {i}: {e}")))?;
332
333 let lot_size = if lot_size_values.is_null(i) {
334 None
335 } else {
336 let lot_size_str = lot_size_values
337 .as_any()
338 .downcast_ref::<StringArray>()
339 .ok_or_else(|| {
340 EncodingError::ParseError("lot_size", format!("row {i}: invalid type"))
341 })?
342 .value(i);
343 Some(
344 Quantity::from_str(lot_size_str)
345 .map_err(|e| EncodingError::ParseError("lot_size", format!("row {i}: {e}")))?,
346 )
347 };
348
349 let max_quantity =
350 if max_quantity_values.is_null(i) {
351 None
352 } else {
353 let max_qty_str = max_quantity_values
354 .as_any()
355 .downcast_ref::<StringArray>()
356 .ok_or_else(|| {
357 EncodingError::ParseError("max_quantity", format!("row {i}: invalid type"))
358 })?
359 .value(i);
360 Some(Quantity::from_str(max_qty_str).map_err(|e| {
361 EncodingError::ParseError("max_quantity", format!("row {i}: {e}"))
362 })?)
363 };
364
365 let min_quantity =
366 if min_quantity_values.is_null(i) {
367 None
368 } else {
369 let min_qty_str = min_quantity_values
370 .as_any()
371 .downcast_ref::<StringArray>()
372 .ok_or_else(|| {
373 EncodingError::ParseError("min_quantity", format!("row {i}: invalid type"))
374 })?
375 .value(i);
376 Some(Quantity::from_str(min_qty_str).map_err(|e| {
377 EncodingError::ParseError("min_quantity", format!("row {i}: {e}"))
378 })?)
379 };
380
381 let max_notional =
382 if max_notional_values.is_null(i) {
383 None
384 } else {
385 let max_not_str = max_notional_values
386 .as_any()
387 .downcast_ref::<StringArray>()
388 .ok_or_else(|| {
389 EncodingError::ParseError("max_notional", format!("row {i}: invalid type"))
390 })?
391 .value(i);
392 Some(Money::from_str(max_not_str).map_err(|e| {
393 EncodingError::ParseError("max_notional", format!("row {i}: {e}"))
394 })?)
395 };
396
397 let min_notional =
398 if min_notional_values.is_null(i) {
399 None
400 } else {
401 let min_not_str = min_notional_values
402 .as_any()
403 .downcast_ref::<StringArray>()
404 .ok_or_else(|| {
405 EncodingError::ParseError("min_notional", format!("row {i}: invalid type"))
406 })?
407 .value(i);
408 Some(Money::from_str(min_not_str).map_err(|e| {
409 EncodingError::ParseError("min_notional", format!("row {i}: {e}"))
410 })?)
411 };
412
413 let max_price = if max_price_values.is_null(i) {
414 None
415 } else {
416 let max_p_str = max_price_values
417 .as_any()
418 .downcast_ref::<StringArray>()
419 .ok_or_else(|| {
420 EncodingError::ParseError("max_price", format!("row {i}: invalid type"))
421 })?
422 .value(i);
423 Some(
424 Price::from_str(max_p_str)
425 .map_err(|e| EncodingError::ParseError("max_price", format!("row {i}: {e}")))?,
426 )
427 };
428
429 let min_price = if min_price_values.is_null(i) {
430 None
431 } else {
432 let min_p_str = min_price_values
433 .as_any()
434 .downcast_ref::<StringArray>()
435 .ok_or_else(|| {
436 EncodingError::ParseError("min_price", format!("row {i}: invalid type"))
437 })?
438 .value(i);
439 Some(
440 Price::from_str(min_p_str)
441 .map_err(|e| EncodingError::ParseError("min_price", format!("row {i}: {e}")))?,
442 )
443 };
444
445 let margin_init = Decimal::from_str(margin_init_values.value(i))
446 .map_err(|e| EncodingError::ParseError("margin_init", format!("row {i}: {e}")))?;
447 let margin_maint = Decimal::from_str(margin_maint_values.value(i))
448 .map_err(|e| EncodingError::ParseError("margin_maint", format!("row {i}: {e}")))?;
449 let maker_fee = Decimal::from_str(maker_fee_values.value(i))
450 .map_err(|e| EncodingError::ParseError("maker_fee", format!("row {i}: {e}")))?;
451 let taker_fee = Decimal::from_str(taker_fee_values.value(i))
452 .map_err(|e| EncodingError::ParseError("taker_fee", format!("row {i}: {e}")))?;
453
454 let info = if info_values.is_null(i) {
456 None
457 } else {
458 let info_bytes = info_values
459 .as_any()
460 .downcast_ref::<BinaryArray>()
461 .ok_or_else(|| EncodingError::ParseError("info", format!("row {i}: invalid type")))?
462 .value(i);
463
464 match serde_json::from_slice::<Params>(info_bytes) {
465 Ok(info_dict) => Some(info_dict),
466 Err(e) => {
467 return Err(EncodingError::ParseError(
468 "info",
469 format!("row {i}: failed to deserialize JSON: {e}"),
470 ));
471 }
472 }
473 };
474
475 let ts_event = UnixNanos::from(ts_event_values.value(i));
476 let ts_init = UnixNanos::from(ts_init_values.value(i));
477
478 let tick_scheme = optional_ustr_value(tick_scheme_values, i);
479
480 let commodity = Commodity::builder()
481 .instrument_id(id)
482 .raw_symbol(raw_symbol)
483 .asset_class(asset_class)
484 .quote_currency(quote_currency)
485 .price_precision(price_prec)
486 .size_precision(size_prec)
487 .price_increment(price_increment)
488 .size_increment(size_increment)
489 .maybe_lot_size(lot_size)
490 .maybe_max_quantity(max_quantity)
491 .maybe_min_quantity(min_quantity)
492 .maybe_max_notional(max_notional)
493 .maybe_min_notional(min_notional)
494 .maybe_max_price(max_price)
495 .maybe_min_price(min_price)
496 .margin_init(margin_init)
497 .margin_maint(margin_maint)
498 .maker_fee(maker_fee)
499 .taker_fee(taker_fee)
500 .maybe_tick_scheme(tick_scheme)
501 .maybe_info(info)
502 .ts_event(ts_event)
503 .ts_init(ts_init)
504 .build()
505 .map_err(|e| super::instrument_validation_error::<Commodity>(i, e))?;
506
507 result.push(commodity);
508 }
509
510 Ok(result)
511}