Skip to main content

nautilus_databento/arrow/
imbalance.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::{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)] // c_char is u8 on some targets
81    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
163/// Decodes a `RecordBatch` into a vector of [`DatabentoImbalance`].
164///
165/// # Errors
166///
167/// Returns an `EncodingError` if decoding fails.
168pub 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
281/// Encodes a vector of [`DatabentoImbalance`] into an Arrow `RecordBatch`.
282///
283/// # Errors
284///
285/// Returns an error if `data` is empty or encoding fails.
286// Guarded by empty check
287pub 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}