1use std::{collections::HashMap, sync::Arc};
17
18use arrow::{
19 array::{
20 Array, Decimal128Array, Int32Array, TimestampNanosecondArray, UInt8Array, UInt16Array,
21 UInt32Array,
22 },
23 datatypes::{DataType, Field, Schema},
24 error::ArrowError,
25 record_batch::RecordBatch,
26};
27use databento::dbn;
28use nautilus_model::{
29 data::{Data, custom::CustomData},
30 identifiers::InstrumentId,
31 types::{PRICE_UNDEF, QUANTITY_UNDEF, fixed::FIXED_PRECISION},
32};
33use nautilus_serialization::arrow::{
34 ArrowSchemaProvider, DecodeDataFromRecordBatch, EncodeToRecordBatch, EncodingError,
35 KEY_TYPE_NAME, decode_decimal_price, decode_decimal_quantity, decode_timestamp,
36 enum_dictionary_array, enum_dictionary_data_type, extract_column, fixed_decimal_data_type,
37 optional_timestamp_array, price_decimal_array, quantity_decimal_array, timestamp_array,
38 timestamp_data_type,
39};
40
41use super::{EnumColumn, parse_metadata};
42use crate::{
43 enums::{DatabentoStatisticType, DatabentoStatisticUpdateAction},
44 types::DatabentoStatistics,
45};
46
47impl ArrowSchemaProvider for DatabentoStatistics {
48 fn get_schema(metadata: Option<HashMap<String, String>>) -> Schema {
49 let fields = vec![
50 Field::new("stat_type", enum_dictionary_data_type(), false),
51 Field::new("update_action", enum_dictionary_data_type(), false),
52 Field::new("price", fixed_decimal_data_type(), true),
53 Field::new("quantity", fixed_decimal_data_type(), true),
54 Field::new("channel_id", DataType::UInt16, false),
55 Field::new("stat_flags", DataType::UInt8, false),
56 Field::new("sequence", DataType::UInt32, false),
57 Field::new("ts_ref", timestamp_data_type(), true),
58 Field::new("ts_in_delta", DataType::Int32, false),
59 Field::new("ts_event", timestamp_data_type(), false),
60 Field::new("ts_recv", timestamp_data_type(), false),
61 Field::new("ts_init", timestamp_data_type(), false),
62 ];
63
64 match metadata {
65 Some(metadata) => Schema::new_with_metadata(fields, metadata),
66 None => Schema::new(fields),
67 }
68 }
69}
70
71impl EncodeToRecordBatch for DatabentoStatistics {
72 fn encode_batch<T>(
73 metadata: &HashMap<String, String>,
74 data: &[T],
75 ) -> Result<RecordBatch, ArrowError>
76 where
77 T: std::borrow::Borrow<Self>,
78 {
79 let mut channel_id_builder = UInt16Array::builder(data.len());
80 let mut stat_flags_builder = UInt8Array::builder(data.len());
81 let mut sequence_builder = UInt32Array::builder(data.len());
82 let mut ts_in_delta_builder = Int32Array::builder(data.len());
83
84 for item in data.iter().map(std::borrow::Borrow::borrow) {
85 channel_id_builder.append_value(item.channel_id);
86 stat_flags_builder.append_value(item.stat_flags);
87 sequence_builder.append_value(item.sequence);
88 ts_in_delta_builder.append_value(item.ts_in_delta);
89 }
90
91 RecordBatch::try_new(
92 Self::get_schema(Some(metadata.clone())).into(),
93 vec![
94 Arc::new(enum_dictionary_array(
95 data.iter().map(|item| item.borrow().stat_type),
96 )?),
97 Arc::new(enum_dictionary_array(
98 data.iter().map(|item| item.borrow().update_action),
99 )?),
100 Arc::new(price_decimal_array(
101 data.iter()
102 .map(|item| item.borrow().price.map_or(PRICE_UNDEF, |value| value.raw())),
103 "price",
104 )?),
105 Arc::new(quantity_decimal_array(
106 data.iter().map(|item| {
107 item.borrow()
108 .quantity
109 .map_or(QUANTITY_UNDEF, |value| value.raw())
110 }),
111 "quantity",
112 )?),
113 Arc::new(channel_id_builder.finish()),
114 Arc::new(stat_flags_builder.finish()),
115 Arc::new(sequence_builder.finish()),
116 Arc::new(optional_timestamp_array(data.iter().map(|item| {
117 let value = item.borrow().ts_ref.as_u64();
118 (value != dbn::UNDEF_TIMESTAMP).then_some(value)
119 }))?),
120 Arc::new(ts_in_delta_builder.finish()),
121 Arc::new(timestamp_array(
122 data.iter().map(|item| item.borrow().ts_event.as_u64()),
123 )?),
124 Arc::new(timestamp_array(
125 data.iter().map(|item| item.borrow().ts_recv.as_u64()),
126 )?),
127 Arc::new(timestamp_array(
128 data.iter().map(|item| item.borrow().ts_init.as_u64()),
129 )?),
130 ],
131 )
132 }
133
134 fn metadata(&self) -> HashMap<String, String> {
135 statistics_metadata(
136 &self.instrument_id,
137 self.price.map_or(FIXED_PRECISION, |p| p.precision),
138 self.quantity.map_or(FIXED_PRECISION, |q| q.precision),
139 )
140 }
141
142 fn chunk_metadata<T>(chunk: &[T]) -> HashMap<String, String>
143 where
144 T: std::borrow::Borrow<Self>,
145 {
146 let first = chunk
147 .first()
148 .map(std::borrow::Borrow::borrow)
149 .expect("Chunk should have at least one element to encode");
150
151 let price_precision = chunk
152 .iter()
153 .map(std::borrow::Borrow::borrow)
154 .find_map(|s| s.price.map(|p| p.precision))
155 .unwrap_or(FIXED_PRECISION);
156 let size_precision = chunk
157 .iter()
158 .map(std::borrow::Borrow::borrow)
159 .find_map(|s| s.quantity.map(|q| q.precision))
160 .unwrap_or(FIXED_PRECISION);
161
162 statistics_metadata(&first.instrument_id, price_precision, size_precision)
163 }
164}
165
166impl DecodeDataFromRecordBatch for DatabentoStatistics {
167 fn decode_data_batch(
168 metadata: &HashMap<String, String>,
169 record_batch: RecordBatch,
170 ) -> Result<Vec<Data>, EncodingError> {
171 let items = decode_statistics_batch(metadata, &record_batch)?;
172 Ok(items
173 .into_iter()
174 .map(|item| Data::Custom(CustomData::from_arc(Arc::new(item))))
175 .collect())
176 }
177}
178
179pub fn decode_statistics_batch(
185 metadata: &HashMap<String, String>,
186 record_batch: &RecordBatch,
187) -> Result<Vec<DatabentoStatistics>, EncodingError> {
188 let (instrument_id, price_precision, size_precision) = parse_metadata(metadata)?;
189 let cols = record_batch.columns();
190
191 let price_values =
192 extract_column::<Decimal128Array>(cols, "price", 2, fixed_decimal_data_type())?;
193 let quantity_values =
194 extract_column::<Decimal128Array>(cols, "quantity", 3, fixed_decimal_data_type())?;
195 let channel_id_values = extract_column::<UInt16Array>(cols, "channel_id", 4, DataType::UInt16)?;
196 let stat_flags_values = extract_column::<UInt8Array>(cols, "stat_flags", 5, DataType::UInt8)?;
197 let sequence_values = extract_column::<UInt32Array>(cols, "sequence", 6, DataType::UInt32)?;
198 let ts_ref_values =
199 extract_column::<TimestampNanosecondArray>(cols, "ts_ref", 7, timestamp_data_type())?;
200 let ts_in_delta_values = extract_column::<Int32Array>(cols, "ts_in_delta", 8, DataType::Int32)?;
201 let ts_event_values =
202 extract_column::<TimestampNanosecondArray>(cols, "ts_event", 9, timestamp_data_type())?;
203 let ts_recv_values =
204 extract_column::<TimestampNanosecondArray>(cols, "ts_recv", 10, timestamp_data_type())?;
205 let ts_init_values =
206 extract_column::<TimestampNanosecondArray>(cols, "ts_init", 11, timestamp_data_type())?;
207 let stat_type_column = EnumColumn::try_from_column(&cols[0], "stat_type", 0)?;
208 let update_action_column = EnumColumn::try_from_column(&cols[1], "update_action", 1)?;
209
210 (0..record_batch.num_rows())
211 .map(|row| {
212 let stat_type = stat_type_column.decode::<DatabentoStatisticType>(row)?;
213 let update_action =
214 update_action_column.decode::<DatabentoStatisticUpdateAction>(row)?;
215
216 let price = (!price_values.is_null(row))
217 .then(|| decode_decimal_price(price_values, price_precision, "price", row))
218 .transpose()?;
219 let quantity = (!quantity_values.is_null(row))
220 .then(|| decode_decimal_quantity(quantity_values, size_precision, "quantity", row))
221 .transpose()?;
222
223 Ok(DatabentoStatistics {
224 instrument_id,
225 stat_type,
226 update_action,
227 price,
228 quantity,
229 channel_id: channel_id_values.value(row),
230 stat_flags: stat_flags_values.value(row),
231 sequence: sequence_values.value(row),
232 ts_ref: if ts_ref_values.is_null(row) {
233 dbn::UNDEF_TIMESTAMP.into()
234 } else {
235 decode_timestamp(ts_ref_values, "ts_ref", row)?.into()
236 },
237 ts_in_delta: ts_in_delta_values.value(row),
238 ts_event: decode_timestamp(ts_event_values, "ts_event", row)?.into(),
239 ts_recv: decode_timestamp(ts_recv_values, "ts_recv", row)?.into(),
240 ts_init: decode_timestamp(ts_init_values, "ts_init", row)?.into(),
241 })
242 })
243 .collect()
244}
245
246fn statistics_metadata(
247 instrument_id: &InstrumentId,
248 price_precision: u8,
249 size_precision: u8,
250) -> HashMap<String, String> {
251 let mut metadata =
252 DatabentoStatistics::get_metadata(instrument_id, price_precision, size_precision);
253 metadata.insert(KEY_TYPE_NAME.to_string(), "DatabentoStatistics".to_string());
254 metadata
255}
256
257pub fn statistics_to_arrow_record_batch(
264 data: &[DatabentoStatistics],
265) -> Result<RecordBatch, EncodingError> {
266 if data.is_empty() {
267 return Err(EncodingError::EmptyData);
268 }
269
270 let metadata = DatabentoStatistics::chunk_metadata(data);
271 DatabentoStatistics::encode_batch(&metadata, data).map_err(EncodingError::ArrowError)
272}
273
274#[cfg(test)]
275mod tests {
276 use std::collections::HashMap;
277
278 use nautilus_model::{
279 identifiers::InstrumentId,
280 types::{Price, Quantity},
281 };
282 use nautilus_serialization::arrow::{
283 ArrowSchemaProvider, EncodeToRecordBatch, KEY_INSTRUMENT_ID, KEY_PRICE_PRECISION,
284 KEY_SIZE_PRECISION,
285 };
286 use rstest::rstest;
287
288 use super::*;
289
290 fn test_metadata() -> HashMap<String, String> {
291 HashMap::from([
292 (KEY_INSTRUMENT_ID.to_string(), "ESM4.GLBX".to_string()),
293 (KEY_PRICE_PRECISION.to_string(), "2".to_string()),
294 (KEY_SIZE_PRECISION.to_string(), "0".to_string()),
295 ])
296 }
297
298 fn test_statistics(instrument_id: InstrumentId) -> DatabentoStatistics {
299 DatabentoStatistics::new(
300 instrument_id,
301 DatabentoStatisticType::OpeningPrice,
302 DatabentoStatisticUpdateAction::Added,
303 Some(Price::from("5000.50")),
304 Some(Quantity::from("100")),
305 1,
306 0,
307 42,
308 1_000_000_000.into(),
309 500,
310 2_000_000_000.into(),
311 3_000_000_000.into(),
312 4_000_000_000.into(),
313 )
314 }
315
316 #[rstest]
317 fn test_get_schema() {
318 let schema = DatabentoStatistics::get_schema(None);
319 assert_eq!(schema.fields().len(), 12);
320 assert_eq!(schema.field(0).name(), "stat_type");
321 assert_eq!(schema.field(11).name(), "ts_init");
322 assert_eq!(schema.field(0).data_type(), &enum_dictionary_data_type());
323 assert_eq!(schema.field(2).data_type(), &fixed_decimal_data_type());
324 assert!(schema.field(2).is_nullable());
325 assert!(schema.field(7).is_nullable());
326 assert_eq!(schema.field(10).data_type(), ×tamp_data_type());
327 }
328
329 #[rstest]
330 fn test_encode_batch() {
331 let instrument_id = InstrumentId::from("ESM4.GLBX");
332 let metadata = test_metadata();
333 let data = vec![test_statistics(instrument_id)];
334 let batch = DatabentoStatistics::encode_batch(&metadata, &data).unwrap();
335
336 assert_eq!(batch.num_rows(), 1);
337 assert_eq!(batch.num_columns(), 12);
338 }
339
340 #[rstest]
341 fn test_encode_decode_round_trip() {
342 let instrument_id = InstrumentId::from("ESM4.GLBX");
343 let metadata = test_metadata();
344 let original = vec![test_statistics(instrument_id)];
345 let batch = DatabentoStatistics::encode_batch(&metadata, &original).unwrap();
346 let decoded = decode_statistics_batch(&metadata, &batch).unwrap();
347
348 assert_eq!(decoded.len(), 1);
349 assert_eq!(decoded[0].instrument_id, instrument_id);
350 assert_eq!(decoded[0].stat_type, original[0].stat_type);
351 assert_eq!(decoded[0].update_action, original[0].update_action);
352 assert_eq!(decoded[0].price, original[0].price);
353 assert_eq!(decoded[0].quantity, original[0].quantity);
354 assert_eq!(decoded[0].channel_id, original[0].channel_id);
355 assert_eq!(decoded[0].stat_flags, original[0].stat_flags);
356 assert_eq!(decoded[0].sequence, original[0].sequence);
357 assert_eq!(decoded[0].ts_ref, original[0].ts_ref);
358 assert_eq!(decoded[0].ts_in_delta, original[0].ts_in_delta);
359 assert_eq!(decoded[0].ts_event, original[0].ts_event);
360 assert_eq!(decoded[0].ts_recv, original[0].ts_recv);
361 assert_eq!(decoded[0].ts_init, original[0].ts_init);
362 }
363
364 #[rstest]
365 fn test_decode_legacy_enum_columns() {
366 let instrument_id = InstrumentId::from("ESM4.GLBX");
367 let metadata = test_metadata();
368 let original = test_statistics(instrument_id);
369 let batch =
370 DatabentoStatistics::encode_batch(&metadata, std::slice::from_ref(&original)).unwrap();
371 let mut fields = batch.schema().fields().to_vec();
372 fields[0] = Arc::new(Field::new("stat_type", DataType::UInt8, false));
373 fields[1] = Arc::new(Field::new("update_action", DataType::UInt8, false));
374 let mut columns = batch.columns().to_vec();
375 columns[0] = Arc::new(UInt8Array::from(vec![original.stat_type as u8]));
376 columns[1] = Arc::new(UInt8Array::from(vec![original.update_action as u8]));
377 let legacy_batch = RecordBatch::try_new(
378 Arc::new(Schema::new_with_metadata(fields, metadata.clone())),
379 columns,
380 )
381 .unwrap();
382
383 let decoded = decode_statistics_batch(&metadata, &legacy_batch).unwrap();
384
385 assert_eq!(decoded, vec![original]);
386 }
387
388 #[rstest]
389 fn test_encode_decode_round_trip_with_none_values() {
390 let instrument_id = InstrumentId::from("ESM4.GLBX");
391 let metadata = test_metadata();
392 let stats = DatabentoStatistics::new(
393 instrument_id,
394 DatabentoStatisticType::ClearedVolume,
395 DatabentoStatisticUpdateAction::Added,
396 None,
397 None,
398 1,
399 0,
400 42,
401 dbn::UNDEF_TIMESTAMP.into(),
402 500,
403 2_000_000_000.into(),
404 3_000_000_000.into(),
405 4_000_000_000.into(),
406 );
407 let original = vec![stats];
408 let batch = DatabentoStatistics::encode_batch(&metadata, &original).unwrap();
409 let decoded = decode_statistics_batch(&metadata, &batch).unwrap();
410
411 assert_eq!(decoded.len(), 1);
412 let ts_ref = batch
413 .column(7)
414 .as_any()
415 .downcast_ref::<TimestampNanosecondArray>()
416 .unwrap();
417 assert!(ts_ref.is_null(0));
418 assert_eq!(decoded[0].price, None);
419 assert_eq!(decoded[0].quantity, None);
420 assert_eq!(decoded[0].ts_ref.as_u64(), dbn::UNDEF_TIMESTAMP);
421 }
422
423 #[rstest]
424 fn test_chunk_metadata_uses_first_non_none_precision() {
425 let instrument_id = InstrumentId::from("ESM4.GLBX");
426 let none_stats = DatabentoStatistics::new(
427 instrument_id,
428 DatabentoStatisticType::ClearedVolume,
429 DatabentoStatisticUpdateAction::Added,
430 None,
431 None,
432 1,
433 0,
434 42,
435 1_000_000_000.into(),
436 500,
437 2_000_000_000.into(),
438 3_000_000_000.into(),
439 4_000_000_000.into(),
440 );
441 let some_stats = test_statistics(instrument_id);
442 let data = vec![none_stats, some_stats];
443
444 let batch = statistics_to_arrow_record_batch(&data).unwrap();
445 let metadata = batch.schema().metadata().clone();
446 let decoded = decode_statistics_batch(&metadata, &batch).unwrap();
447
448 assert_eq!(decoded.len(), 2);
449 assert_eq!(decoded[0].price, None);
450 assert_eq!(decoded[0].quantity, None);
451 assert_eq!(decoded[1].price, data[1].price);
452 assert_eq!(decoded[1].quantity, data[1].quantity);
453 }
454
455 #[rstest]
456 fn test_encode_decode_multiple_rows() {
457 let instrument_id = InstrumentId::from("ESM4.GLBX");
458 let metadata = test_metadata();
459 let stats1 = test_statistics(instrument_id);
460 let stats2 = DatabentoStatistics::new(
461 instrument_id,
462 DatabentoStatisticType::ClearedVolume,
463 DatabentoStatisticUpdateAction::Added,
464 Some(Price::from("5100.25")),
465 None,
466 2,
467 1,
468 43,
469 2_000_000_000.into(),
470 600,
471 3_000_000_000.into(),
472 4_000_000_000.into(),
473 5_000_000_000.into(),
474 );
475 let stats3 = DatabentoStatistics::new(
476 instrument_id,
477 DatabentoStatisticType::OpeningPrice,
478 DatabentoStatisticUpdateAction::Added,
479 None,
480 Some(Quantity::from("200")),
481 3,
482 0,
483 44,
484 3_000_000_000.into(),
485 700,
486 4_000_000_000.into(),
487 5_000_000_000.into(),
488 6_000_000_000.into(),
489 );
490 let original = vec![stats1, stats2, stats3];
491
492 let batch = DatabentoStatistics::encode_batch(&metadata, &original).unwrap();
493 assert_eq!(batch.num_rows(), 3);
494
495 let decoded = decode_statistics_batch(&metadata, &batch).unwrap();
496 assert_eq!(decoded.len(), 3);
497 for (orig, dec) in original.iter().zip(decoded.iter()) {
498 assert_eq!(dec.instrument_id, orig.instrument_id);
499 assert_eq!(dec.stat_type, orig.stat_type);
500 assert_eq!(dec.price, orig.price);
501 assert_eq!(dec.quantity, orig.quantity);
502 assert_eq!(dec.channel_id, orig.channel_id);
503 assert_eq!(dec.sequence, orig.sequence);
504 }
505 }
506
507 #[rstest]
508 fn test_statistics_to_arrow_record_batch_round_trip() {
509 let instrument_id = InstrumentId::from("ESM4.GLBX");
510 let original = vec![test_statistics(instrument_id)];
511 let batch = statistics_to_arrow_record_batch(&original).unwrap();
512 let metadata = batch.schema().metadata().clone();
513 let decoded = decode_statistics_batch(&metadata, &batch).unwrap();
514
515 assert_eq!(decoded.len(), 1);
516 assert_eq!(decoded[0].price, original[0].price);
517 assert_eq!(decoded[0].quantity, original[0].quantity);
518 }
519
520 #[rstest]
521 fn test_chunk_metadata_all_none_uses_fixed_precision() {
522 use nautilus_model::types::fixed::FIXED_PRECISION;
523
524 let instrument_id = InstrumentId::from("ESM4.GLBX");
525 let stats = DatabentoStatistics::new(
526 instrument_id,
527 DatabentoStatisticType::ClearedVolume,
528 DatabentoStatisticUpdateAction::Added,
529 None,
530 None,
531 1,
532 0,
533 42,
534 1_000_000_000.into(),
535 500,
536 2_000_000_000.into(),
537 3_000_000_000.into(),
538 4_000_000_000.into(),
539 );
540 let data = vec![stats];
541 let metadata = DatabentoStatistics::chunk_metadata(&data);
542
543 assert_eq!(
544 metadata.get(KEY_PRICE_PRECISION).unwrap(),
545 &FIXED_PRECISION.to_string(),
546 );
547 assert_eq!(
548 metadata.get(KEY_SIZE_PRECISION).unwrap(),
549 &FIXED_PRECISION.to_string(),
550 );
551 }
552
553 #[rstest]
554 fn test_all_none_metadata_decodes_real_prices_correctly() {
555 use nautilus_model::types::fixed::FIXED_PRECISION;
556
557 let instrument_id = InstrumentId::from("ESM4.GLBX");
558 let price = Price::from("5000.50");
559 let quantity = Quantity::from("100");
560 let stats = DatabentoStatistics::new(
561 instrument_id,
562 DatabentoStatisticType::OpeningPrice,
563 DatabentoStatisticUpdateAction::Added,
564 Some(price),
565 Some(quantity),
566 1,
567 0,
568 42,
569 1_000_000_000.into(),
570 500,
571 2_000_000_000.into(),
572 3_000_000_000.into(),
573 4_000_000_000.into(),
574 );
575
576 let metadata = HashMap::from([
578 (KEY_INSTRUMENT_ID.to_string(), "ESM4.GLBX".to_string()),
579 (KEY_PRICE_PRECISION.to_string(), FIXED_PRECISION.to_string()),
580 (KEY_SIZE_PRECISION.to_string(), FIXED_PRECISION.to_string()),
581 ]);
582
583 let batch = DatabentoStatistics::encode_batch(&metadata, &[stats]).unwrap();
584 let decoded = decode_statistics_batch(&metadata, &batch).unwrap();
585
586 assert_eq!(decoded.len(), 1);
587 assert_eq!(decoded[0].price.unwrap().as_f64(), price.as_f64());
588 assert_eq!(decoded[0].quantity.unwrap().as_f64(), quantity.as_f64());
589 }
590
591 #[rstest]
592 fn test_get_schema_with_metadata() {
593 let metadata = test_metadata();
594 let schema = DatabentoStatistics::get_schema(Some(metadata.clone()));
595 assert_eq!(schema.metadata(), &metadata);
596 assert_eq!(schema.fields().len(), 12);
597 }
598
599 #[rstest]
600 fn test_decode_missing_metadata_returns_error() {
601 let instrument_id = InstrumentId::from("ESM4.GLBX");
602 let metadata = test_metadata();
603 let data = vec![test_statistics(instrument_id)];
604 let batch = DatabentoStatistics::encode_batch(&metadata, &data).unwrap();
605
606 let empty_metadata = HashMap::new();
607 let result = decode_statistics_batch(&empty_metadata, &batch);
608 assert!(result.is_err());
609 }
610
611 #[rstest]
612 fn test_statistics_to_arrow_record_batch_empty() {
613 let result = statistics_to_arrow_record_batch(&[]);
614 assert!(result.is_err());
615 }
616
617 #[rstest]
618 fn test_decode_data_batch_produces_custom_data() {
619 let instrument_id = InstrumentId::from("ESM4.GLBX");
620 let metadata = test_metadata();
621 let original = vec![test_statistics(instrument_id)];
622 let batch = DatabentoStatistics::encode_batch(&metadata, &original).unwrap();
623 let data_vec = DatabentoStatistics::decode_data_batch(&metadata, batch).unwrap();
624
625 assert_eq!(data_vec.len(), 1);
626 match &data_vec[0] {
627 Data::Custom(custom) => {
628 assert_eq!(custom.data.type_name(), "DatabentoStatistics");
629 let stats = custom
630 .data
631 .as_any()
632 .downcast_ref::<DatabentoStatistics>()
633 .unwrap();
634 assert_eq!(stats.instrument_id, instrument_id);
635 assert_eq!(stats.stat_type, original[0].stat_type);
636 assert_eq!(stats.price, original[0].price);
637 assert_eq!(stats.quantity, original[0].quantity);
638 assert_eq!(stats.ts_event, original[0].ts_event);
639 assert_eq!(stats.ts_init, original[0].ts_init);
640 }
641 other => panic!("Expected Data::Custom, was {other:?}"),
642 }
643 }
644
645 #[rstest]
646 fn test_decode_data_batch_multiple_rows() {
647 let instrument_id = InstrumentId::from("ESM4.GLBX");
648 let metadata = test_metadata();
649 let stats2 = DatabentoStatistics::new(
650 instrument_id,
651 DatabentoStatisticType::ClearedVolume,
652 DatabentoStatisticUpdateAction::Added,
653 None,
654 Some(Quantity::from("200")),
655 2,
656 1,
657 43,
658 2_000_000_000.into(),
659 600,
660 3_000_000_000.into(),
661 4_000_000_000.into(),
662 5_000_000_000.into(),
663 );
664 let original = vec![test_statistics(instrument_id), stats2];
665 let batch = DatabentoStatistics::encode_batch(&metadata, &original).unwrap();
666 let data_vec = DatabentoStatistics::decode_data_batch(&metadata, batch).unwrap();
667
668 assert_eq!(data_vec.len(), 2);
669 for (i, data) in data_vec.iter().enumerate() {
670 match data {
671 Data::Custom(custom) => {
672 let stats = custom
673 .data
674 .as_any()
675 .downcast_ref::<DatabentoStatistics>()
676 .unwrap();
677 assert_eq!(stats.instrument_id, original[i].instrument_id);
678 assert_eq!(stats.stat_type, original[i].stat_type);
679 assert_eq!(stats.price, original[i].price);
680 assert_eq!(stats.quantity, original[i].quantity);
681 }
682 other => panic!("Expected Data::Custom, was {other:?}"),
683 }
684 }
685 }
686
687 #[rstest]
688 fn test_ipc_stream_round_trip() {
689 use std::io::Cursor;
690
691 use arrow::ipc::{reader::StreamReader, writer::StreamWriter};
692
693 let instrument_id = InstrumentId::from("ESM4.GLBX");
694 let original = vec![
695 test_statistics(instrument_id),
696 DatabentoStatistics::new(
697 instrument_id,
698 DatabentoStatisticType::ClearedVolume,
699 DatabentoStatisticUpdateAction::Added,
700 None,
701 Some(Quantity::from("200")),
702 2,
703 1,
704 43,
705 2_000_000_000.into(),
706 600,
707 3_000_000_000.into(),
708 4_000_000_000.into(),
709 5_000_000_000.into(),
710 ),
711 ];
712 let batch = statistics_to_arrow_record_batch(&original).unwrap();
713
714 let mut cursor = Cursor::new(Vec::new());
715 {
716 let mut writer = StreamWriter::try_new(&mut cursor, &batch.schema()).unwrap();
717 writer.write(&batch).unwrap();
718 writer.finish().unwrap();
719 }
720
721 let buffer = cursor.into_inner();
722 let reader = StreamReader::try_new(Cursor::new(buffer), None).unwrap();
723 let mut decoded = Vec::new();
724
725 for batch_result in reader {
726 let batch = batch_result.unwrap();
727 let metadata = batch.schema().metadata().clone();
728 decoded.extend(decode_statistics_batch(&metadata, &batch).unwrap());
729 }
730
731 assert_eq!(decoded.len(), 2);
732 for (orig, dec) in original.iter().zip(decoded.iter()) {
733 assert_eq!(dec, orig);
734 }
735 }
736}