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