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        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
179/// Decodes a `RecordBatch` into a vector of [`DatabentoStatistics`].
180///
181/// # Errors
182///
183/// Returns an `EncodingError` if decoding fails.
184pub 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
257/// Encodes a vector of [`DatabentoStatistics`] into an Arrow `RecordBatch`.
258///
259/// # Errors
260///
261/// Returns an error if `data` is empty or encoding fails.
262// Guarded by empty check
263pub 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(), &timestamp_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        // Encode with FIXED_PRECISION metadata (as if from an all-None chunk)
577        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}