1use std::{collections::HashMap, sync::Arc};
17
18use arrow::{
19 array::{Decimal128Array, Int8Array, TimestampNanosecondArray},
20 datatypes::{DataType, Field, Schema},
21 error::ArrowError,
22 record_batch::RecordBatch,
23};
24use nautilus_model::{
25 data::{Data, custom::CustomData},
26 enums::OrderSide,
27};
28use nautilus_serialization::arrow::{
29 ArrowSchemaProvider, DecodeDataFromRecordBatch, EncodeToRecordBatch, EncodingError,
30 decode_decimal_price, decode_decimal_quantity, decode_timestamp, enum_dictionary_array,
31 enum_dictionary_data_type, extract_column, fixed_decimal_data_type, price_decimal_array,
32 quantity_decimal_array, timestamp_array, timestamp_data_type,
33};
34
35use super::{EnumColumn, parse_metadata};
36use crate::types::DatabentoImbalance;
37
38impl ArrowSchemaProvider for DatabentoImbalance {
39 fn get_schema(metadata: Option<HashMap<String, String>>) -> Schema {
40 let fields = vec![
41 Field::new("ref_price", fixed_decimal_data_type(), true),
42 Field::new("cont_book_clr_price", fixed_decimal_data_type(), true),
43 Field::new("auct_interest_clr_price", fixed_decimal_data_type(), true),
44 Field::new("paired_qty", fixed_decimal_data_type(), false),
45 Field::new("total_imbalance_qty", fixed_decimal_data_type(), false),
46 Field::new("side", enum_dictionary_data_type(), false),
47 Field::new("significant_imbalance", DataType::Int8, false),
48 Field::new("ts_event", timestamp_data_type(), false),
49 Field::new("ts_recv", timestamp_data_type(), false),
50 Field::new("ts_init", timestamp_data_type(), false),
51 ];
52
53 match metadata {
54 Some(metadata) => Schema::new_with_metadata(fields, metadata),
55 None => Schema::new(fields),
56 }
57 }
58}
59
60impl EncodeToRecordBatch for DatabentoImbalance {
61 #[expect(clippy::unnecessary_cast)] fn encode_batch<T>(
63 metadata: &HashMap<String, String>,
64 data: &[T],
65 ) -> Result<RecordBatch, ArrowError>
66 where
67 T: std::borrow::Borrow<Self>,
68 {
69 let mut significant_imbalance_builder = Int8Array::builder(data.len());
70
71 for item in data.iter().map(std::borrow::Borrow::borrow) {
72 significant_imbalance_builder.append_value(item.significant_imbalance as i8);
73 }
74
75 RecordBatch::try_new(
76 Self::get_schema(Some(metadata.clone())).into(),
77 vec![
78 Arc::new(price_decimal_array(
79 data.iter().map(|item| item.borrow().ref_price.raw()),
80 "ref_price",
81 )?),
82 Arc::new(price_decimal_array(
83 data.iter()
84 .map(|item| item.borrow().cont_book_clr_price.raw()),
85 "cont_book_clr_price",
86 )?),
87 Arc::new(price_decimal_array(
88 data.iter()
89 .map(|item| item.borrow().auct_interest_clr_price.raw()),
90 "auct_interest_clr_price",
91 )?),
92 Arc::new(quantity_decimal_array(
93 data.iter().map(|item| item.borrow().paired_qty.raw()),
94 "paired_qty",
95 )?),
96 Arc::new(quantity_decimal_array(
97 data.iter()
98 .map(|item| item.borrow().total_imbalance_qty.raw()),
99 "total_imbalance_qty",
100 )?),
101 Arc::new(enum_dictionary_array(data.iter().map(|item| {
102 item.borrow()
103 .side
104 .map_or_else(|| "NO_ORDER_SIDE".to_string(), |side| side.to_string())
105 }))?),
106 Arc::new(significant_imbalance_builder.finish()),
107 Arc::new(timestamp_array(
108 data.iter().map(|item| item.borrow().ts_event.as_u64()),
109 )?),
110 Arc::new(timestamp_array(
111 data.iter().map(|item| item.borrow().ts_recv.as_u64()),
112 )?),
113 Arc::new(timestamp_array(
114 data.iter().map(|item| item.borrow().ts_init.as_u64()),
115 )?),
116 ],
117 )
118 }
119
120 fn metadata(&self) -> HashMap<String, String> {
121 let mut metadata = Self::get_metadata(
122 &self.instrument_id,
123 self.ref_price.precision,
124 self.paired_qty.precision,
125 );
126 metadata.insert("type_name".to_string(), "DatabentoImbalance".to_string());
127 metadata
128 }
129}
130
131impl DecodeDataFromRecordBatch for DatabentoImbalance {
132 fn decode_data_batch(
133 metadata: &HashMap<String, String>,
134 record_batch: RecordBatch,
135 ) -> Result<Vec<Data>, EncodingError> {
136 let items = decode_imbalance_batch(metadata, &record_batch)?;
137 Ok(items
138 .into_iter()
139 .map(|item| Data::Custom(CustomData::from_arc(Arc::new(item))))
140 .collect())
141 }
142}
143
144pub fn decode_imbalance_batch(
150 metadata: &HashMap<String, String>,
151 record_batch: &RecordBatch,
152) -> Result<Vec<DatabentoImbalance>, EncodingError> {
153 let (instrument_id, price_precision, size_precision) = parse_metadata(metadata)?;
154 let cols = record_batch.columns();
155
156 let decimal_type = fixed_decimal_data_type();
157 let ref_price_values =
158 extract_column::<Decimal128Array>(cols, "ref_price", 0, decimal_type.clone())?;
159 let cont_book_clr_price_values =
160 extract_column::<Decimal128Array>(cols, "cont_book_clr_price", 1, decimal_type.clone())?;
161 let auct_interest_clr_price_values = extract_column::<Decimal128Array>(
162 cols,
163 "auct_interest_clr_price",
164 2,
165 decimal_type.clone(),
166 )?;
167 let paired_qty_values =
168 extract_column::<Decimal128Array>(cols, "paired_qty", 3, decimal_type.clone())?;
169 let total_imbalance_qty_values =
170 extract_column::<Decimal128Array>(cols, "total_imbalance_qty", 4, decimal_type)?;
171 let significant_imbalance_values =
172 extract_column::<Int8Array>(cols, "significant_imbalance", 6, DataType::Int8)?;
173 let side_column = EnumColumn::try_from_column(&cols[5], "side", 5)?;
174 let ts_event_values =
175 extract_column::<TimestampNanosecondArray>(cols, "ts_event", 7, timestamp_data_type())?;
176 let ts_recv_values =
177 extract_column::<TimestampNanosecondArray>(cols, "ts_recv", 8, timestamp_data_type())?;
178 let ts_init_values =
179 extract_column::<TimestampNanosecondArray>(cols, "ts_init", 9, timestamp_data_type())?;
180
181 (0..record_batch.num_rows())
182 .map(|row| {
183 let ref_price =
184 decode_decimal_price(ref_price_values, price_precision, "ref_price", row)?;
185 let cont_book_clr_price = decode_decimal_price(
186 cont_book_clr_price_values,
187 price_precision,
188 "cont_book_clr_price",
189 row,
190 )?;
191 let auct_interest_clr_price = decode_decimal_price(
192 auct_interest_clr_price_values,
193 price_precision,
194 "auct_interest_clr_price",
195 row,
196 )?;
197 let paired_qty =
198 decode_decimal_quantity(paired_qty_values, size_precision, "paired_qty", row)?;
199 let total_imbalance_qty = decode_decimal_quantity(
200 total_imbalance_qty_values,
201 size_precision,
202 "total_imbalance_qty",
203 row,
204 )?;
205 let side = side_column.decode_optional(row, "NO_ORDER_SIDE", |value| match value {
206 1 => Some(OrderSide::Buy),
207 2 => Some(OrderSide::Sell),
208 _ => None,
209 })?;
210 let significant_imbalance = significant_imbalance_values.value(row) as std::ffi::c_char;
211
212 Ok(DatabentoImbalance {
213 instrument_id,
214 ref_price,
215 cont_book_clr_price,
216 auct_interest_clr_price,
217 paired_qty,
218 total_imbalance_qty,
219 side,
220 significant_imbalance,
221 ts_event: decode_timestamp(ts_event_values, "ts_event", row)?.into(),
222 ts_recv: decode_timestamp(ts_recv_values, "ts_recv", row)?.into(),
223 ts_init: decode_timestamp(ts_init_values, "ts_init", row)?.into(),
224 })
225 })
226 .collect()
227}
228
229pub fn imbalance_to_arrow_record_batch(
236 data: &[DatabentoImbalance],
237) -> Result<RecordBatch, EncodingError> {
238 if data.is_empty() {
239 return Err(EncodingError::EmptyData);
240 }
241
242 let metadata = DatabentoImbalance::chunk_metadata(data);
243 DatabentoImbalance::encode_batch(&metadata, data).map_err(EncodingError::ArrowError)
244}
245
246#[cfg(test)]
247mod tests {
248 use arrow::array::UInt8Array;
249 use nautilus_model::{
250 enums::OrderSide,
251 identifiers::InstrumentId,
252 types::{PRICE_UNDEF, Price, Quantity},
253 };
254 use nautilus_serialization::arrow::{
255 ArrowSchemaProvider, EncodeToRecordBatch, KEY_INSTRUMENT_ID, KEY_PRICE_PRECISION,
256 KEY_SIZE_PRECISION,
257 };
258 use rstest::rstest;
259
260 use super::*;
261
262 fn test_metadata() -> HashMap<String, String> {
263 HashMap::from([
264 (KEY_INSTRUMENT_ID.to_string(), "AAPL.XNAS".to_string()),
265 (KEY_PRICE_PRECISION.to_string(), "2".to_string()),
266 (KEY_SIZE_PRECISION.to_string(), "0".to_string()),
267 ])
268 }
269
270 fn test_imbalance(instrument_id: InstrumentId) -> DatabentoImbalance {
271 DatabentoImbalance::new(
272 instrument_id,
273 Price::from("100.50"),
274 Price::from("100.45"),
275 Price::from("100.55"),
276 Quantity::from("1000"),
277 Quantity::from("500"),
278 Some(OrderSide::Buy),
279 b'Y' as std::ffi::c_char,
280 1.into(),
281 2.into(),
282 3.into(),
283 )
284 }
285
286 #[rstest]
287 fn test_undefined_prices_round_trip() {
288 let mut value = test_imbalance(InstrumentId::from("AAPL.XNAS"));
289 value.ref_price = Price::from_raw(PRICE_UNDEF, 0);
290 value.cont_book_clr_price = value.ref_price;
291 value.auct_interest_clr_price = value.ref_price;
292 let metadata = test_metadata();
293 let batch =
294 DatabentoImbalance::encode_batch(&metadata, std::slice::from_ref(&value)).unwrap();
295 let decoded = decode_imbalance_batch(&metadata, &batch).unwrap();
296
297 assert_eq!(decoded, vec![value]);
298 }
299
300 #[rstest]
301 fn test_get_schema() {
302 let schema = DatabentoImbalance::get_schema(None);
303 assert_eq!(schema.fields().len(), 10);
304 assert_eq!(schema.field(0).name(), "ref_price");
305 assert_eq!(schema.field(5).name(), "side");
306 assert_eq!(schema.field(9).name(), "ts_init");
307 assert_eq!(schema.field(0).data_type(), &fixed_decimal_data_type());
308 assert_eq!(schema.field(5).data_type(), &enum_dictionary_data_type());
309 assert_eq!(schema.field(8).data_type(), ×tamp_data_type());
310 }
311
312 #[rstest]
313 fn test_encode_batch() {
314 let instrument_id = InstrumentId::from("AAPL.XNAS");
315 let metadata = test_metadata();
316 let data = vec![test_imbalance(instrument_id)];
317 let batch = DatabentoImbalance::encode_batch(&metadata, &data).unwrap();
318
319 assert_eq!(batch.num_rows(), 1);
320 assert_eq!(batch.num_columns(), 10);
321 }
322
323 #[rstest]
324 fn test_encode_decode_round_trip() {
325 let instrument_id = InstrumentId::from("AAPL.XNAS");
326 let metadata = test_metadata();
327 let original = vec![test_imbalance(instrument_id)];
328 let batch = DatabentoImbalance::encode_batch(&metadata, &original).unwrap();
329 let decoded = decode_imbalance_batch(&metadata, &batch).unwrap();
330
331 assert_eq!(decoded.len(), 1);
332 assert_eq!(decoded[0].instrument_id, instrument_id);
333 assert_eq!(decoded[0].ref_price, original[0].ref_price);
334 assert_eq!(
335 decoded[0].cont_book_clr_price,
336 original[0].cont_book_clr_price
337 );
338 assert_eq!(
339 decoded[0].auct_interest_clr_price,
340 original[0].auct_interest_clr_price
341 );
342 assert_eq!(decoded[0].paired_qty, original[0].paired_qty);
343 assert_eq!(
344 decoded[0].total_imbalance_qty,
345 original[0].total_imbalance_qty
346 );
347 assert_eq!(decoded[0].side, original[0].side);
348 assert_eq!(
349 decoded[0].significant_imbalance,
350 original[0].significant_imbalance
351 );
352 assert_eq!(decoded[0].ts_event, original[0].ts_event);
353 assert_eq!(decoded[0].ts_recv, original[0].ts_recv);
354 assert_eq!(decoded[0].ts_init, original[0].ts_init);
355 }
356
357 #[rstest]
358 fn test_decode_legacy_side_column() {
359 let instrument_id = InstrumentId::from("AAPL.XNAS");
360 let metadata = test_metadata();
361 let original = test_imbalance(instrument_id);
362 let batch =
363 DatabentoImbalance::encode_batch(&metadata, std::slice::from_ref(&original)).unwrap();
364 let mut fields = batch.schema().fields().to_vec();
365 fields[5] = Arc::new(Field::new("side", DataType::UInt8, false));
366 let mut columns = batch.columns().to_vec();
367 columns[5] = Arc::new(UInt8Array::from(vec![
368 original.side.map_or(0, |side| side as u8),
369 ]));
370 let legacy_batch = RecordBatch::try_new(
371 Arc::new(Schema::new_with_metadata(fields, metadata.clone())),
372 columns,
373 )
374 .unwrap();
375
376 let decoded = decode_imbalance_batch(&metadata, &legacy_batch).unwrap();
377
378 assert_eq!(decoded, vec![original]);
379 }
380
381 #[rstest]
382 fn test_encode_decode_multiple_rows() {
383 let instrument_id = InstrumentId::from("AAPL.XNAS");
384 let metadata = test_metadata();
385 let imb1 = test_imbalance(instrument_id);
386 let mut imb2 = test_imbalance(instrument_id);
387 imb2.side = Some(OrderSide::Sell);
388 imb2.ref_price = Price::from("101.00");
389 imb2.ts_event = 100.into();
390 let mut imb3 = test_imbalance(instrument_id);
391 imb3.side = None;
392 imb3.significant_imbalance = b'N' as std::ffi::c_char;
393 let original = vec![imb1, imb2, imb3];
394
395 let batch = DatabentoImbalance::encode_batch(&metadata, &original).unwrap();
396 assert_eq!(batch.num_rows(), 3);
397
398 let decoded = decode_imbalance_batch(&metadata, &batch).unwrap();
399 assert_eq!(decoded.len(), 3);
400 for (orig, dec) in original.iter().zip(decoded.iter()) {
401 assert_eq!(dec.instrument_id, orig.instrument_id);
402 assert_eq!(dec.ref_price, orig.ref_price);
403 assert_eq!(dec.side, orig.side);
404 assert_eq!(dec.significant_imbalance, orig.significant_imbalance);
405 assert_eq!(dec.ts_event, orig.ts_event);
406 }
407 }
408
409 #[rstest]
410 fn test_imbalance_to_arrow_record_batch_round_trip() {
411 let instrument_id = InstrumentId::from("AAPL.XNAS");
412 let original = vec![test_imbalance(instrument_id)];
413 let batch = imbalance_to_arrow_record_batch(&original).unwrap();
414 let metadata = batch.schema().metadata().clone();
415 let decoded = decode_imbalance_batch(&metadata, &batch).unwrap();
416
417 assert_eq!(decoded.len(), 1);
418 assert_eq!(decoded[0].ref_price, original[0].ref_price);
419 assert_eq!(decoded[0].paired_qty, original[0].paired_qty);
420 }
421
422 #[rstest]
423 fn test_get_schema_with_metadata() {
424 let metadata = test_metadata();
425 let schema = DatabentoImbalance::get_schema(Some(metadata.clone()));
426 assert_eq!(schema.metadata(), &metadata);
427 assert_eq!(schema.fields().len(), 10);
428 }
429
430 #[rstest]
431 fn test_imbalance_to_arrow_record_batch_empty() {
432 let result = imbalance_to_arrow_record_batch(&[]);
433 assert!(result.is_err());
434 }
435
436 #[rstest]
437 fn test_decode_missing_metadata_returns_error() {
438 let instrument_id = InstrumentId::from("AAPL.XNAS");
439 let metadata = test_metadata();
440 let data = vec![test_imbalance(instrument_id)];
441 let batch = DatabentoImbalance::encode_batch(&metadata, &data).unwrap();
442
443 let empty_metadata = HashMap::new();
444 let result = decode_imbalance_batch(&empty_metadata, &batch);
445 assert!(result.is_err());
446 }
447
448 #[rstest]
449 fn test_decode_data_batch_produces_custom_data() {
450 let instrument_id = InstrumentId::from("AAPL.XNAS");
451 let metadata = test_metadata();
452 let original = vec![test_imbalance(instrument_id)];
453 let batch = DatabentoImbalance::encode_batch(&metadata, &original).unwrap();
454 let data_vec = DatabentoImbalance::decode_data_batch(&metadata, batch).unwrap();
455
456 assert_eq!(data_vec.len(), 1);
457 match &data_vec[0] {
458 Data::Custom(custom) => {
459 assert_eq!(custom.data.type_name(), "DatabentoImbalance");
460 let imbalance = custom
461 .data
462 .as_any()
463 .downcast_ref::<DatabentoImbalance>()
464 .unwrap();
465 assert_eq!(imbalance.instrument_id, instrument_id);
466 assert_eq!(imbalance.ref_price, original[0].ref_price);
467 assert_eq!(imbalance.paired_qty, original[0].paired_qty);
468 assert_eq!(imbalance.side, original[0].side);
469 assert_eq!(imbalance.ts_event, original[0].ts_event);
470 assert_eq!(imbalance.ts_init, original[0].ts_init);
471 }
472 other => panic!("Expected Data::Custom, was {other:?}"),
473 }
474 }
475
476 #[rstest]
477 fn test_decode_data_batch_multiple_rows() {
478 let instrument_id = InstrumentId::from("AAPL.XNAS");
479 let metadata = test_metadata();
480 let mut imb2 = test_imbalance(instrument_id);
481 imb2.side = Some(OrderSide::Sell);
482 imb2.ts_event = 100.into();
483 let original = vec![test_imbalance(instrument_id), imb2];
484 let batch = DatabentoImbalance::encode_batch(&metadata, &original).unwrap();
485 let data_vec = DatabentoImbalance::decode_data_batch(&metadata, batch).unwrap();
486
487 assert_eq!(data_vec.len(), 2);
488 for (i, data) in data_vec.iter().enumerate() {
489 match data {
490 Data::Custom(custom) => {
491 let imbalance = custom
492 .data
493 .as_any()
494 .downcast_ref::<DatabentoImbalance>()
495 .unwrap();
496 assert_eq!(imbalance.instrument_id, original[i].instrument_id);
497 assert_eq!(imbalance.side, original[i].side);
498 assert_eq!(imbalance.ts_event, original[i].ts_event);
499 }
500 other => panic!("Expected Data::Custom, was {other:?}"),
501 }
502 }
503 }
504
505 #[rstest]
506 fn test_ipc_stream_round_trip() {
507 use std::io::Cursor;
508
509 use arrow::ipc::{reader::StreamReader, writer::StreamWriter};
510
511 let instrument_id = InstrumentId::from("AAPL.XNAS");
512 let original = vec![test_imbalance(instrument_id), {
513 let mut imb = test_imbalance(instrument_id);
514 imb.side = Some(OrderSide::Sell);
515 imb.ref_price = Price::from("101.25");
516 imb.ts_event = 100.into();
517 imb
518 }];
519 let batch = imbalance_to_arrow_record_batch(&original).unwrap();
520
521 let mut cursor = Cursor::new(Vec::new());
522 {
523 let mut writer = StreamWriter::try_new(&mut cursor, &batch.schema()).unwrap();
524 writer.write(&batch).unwrap();
525 writer.finish().unwrap();
526 }
527
528 let buffer = cursor.into_inner();
529 let reader = StreamReader::try_new(Cursor::new(buffer), None).unwrap();
530 let mut decoded = Vec::new();
531
532 for batch_result in reader {
533 let batch = batch_result.unwrap();
534 let metadata = batch.schema().metadata().clone();
535 decoded.extend(decode_imbalance_batch(&metadata, &batch).unwrap());
536 }
537
538 assert_eq!(decoded.len(), 2);
539 for (orig, dec) in original.iter().zip(decoded.iter()) {
540 assert_eq!(dec, orig);
541 }
542 }
543}