Skip to main content

nautilus_databento/arrow/
statistics.rs

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