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