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