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