Skip to main content

nautilus_serialization/arrow/
bar.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, str::FromStr, sync::Arc};
17
18use arrow::{
19    array::{Decimal128Array, UInt64Array},
20    datatypes::{DataType, Field, Schema},
21    error::ArrowError,
22    record_batch::RecordBatch,
23};
24use nautilus_model::data::{Bar, BarType};
25
26use super::{
27    DecodeDataFromRecordBatch, EncodingError, KEY_BAR_TYPE, KEY_IDENTIFIER, KEY_PRICE_PRECISION,
28    KEY_SIZE_PRECISION, decode_required_decimal_price, decode_required_decimal_quantity,
29    decode_required_timestamp, extract_column, fixed_decimal_data_type,
30    identifier_array_from_display, required_price_decimal_array, required_quantity_decimal_array,
31};
32use crate::arrow::{ArrowSchemaProvider, Data, DecodeFromRecordBatch, EncodeToRecordBatch};
33
34impl ArrowSchemaProvider for Bar {
35    fn get_schema(metadata: Option<HashMap<String, String>>) -> Schema {
36        let fields = vec![
37            Field::new("open", fixed_decimal_data_type(), true),
38            Field::new("high", fixed_decimal_data_type(), true),
39            Field::new("low", fixed_decimal_data_type(), true),
40            Field::new("close", fixed_decimal_data_type(), true),
41            Field::new("volume", fixed_decimal_data_type(), true),
42            Field::new("ts_event", crate::arrow::timestamp_data_type(), false),
43            Field::new("ts_init", crate::arrow::timestamp_data_type(), false),
44            Field::new(KEY_IDENTIFIER, DataType::Utf8, true),
45        ];
46
47        match metadata {
48            Some(metadata) => Schema::new_with_metadata(fields, metadata),
49            None => Schema::new(fields),
50        }
51    }
52}
53
54fn parse_metadata(metadata: &HashMap<String, String>) -> Result<(BarType, u8, u8), EncodingError> {
55    let bar_type_str = metadata
56        .get(KEY_BAR_TYPE)
57        .ok_or_else(|| EncodingError::MissingMetadata(KEY_BAR_TYPE))?;
58    let bar_type = BarType::from_str(bar_type_str)
59        .map_err(|e| EncodingError::ParseError(KEY_BAR_TYPE, e.to_string()))?;
60
61    let price_precision = metadata
62        .get(KEY_PRICE_PRECISION)
63        .ok_or_else(|| EncodingError::MissingMetadata(KEY_PRICE_PRECISION))?
64        .parse::<u8>()
65        .map_err(|e| EncodingError::ParseError(KEY_PRICE_PRECISION, e.to_string()))?;
66
67    let size_precision = metadata
68        .get(KEY_SIZE_PRECISION)
69        .ok_or_else(|| EncodingError::MissingMetadata(KEY_SIZE_PRECISION))?
70        .parse::<u8>()
71        .map_err(|e| EncodingError::ParseError(KEY_SIZE_PRECISION, e.to_string()))?;
72
73    Ok((bar_type, price_precision, size_precision))
74}
75
76impl EncodeToRecordBatch for Bar {
77    fn encode_batch<T>(
78        metadata: &HashMap<String, String>,
79        data: &[T],
80    ) -> Result<RecordBatch, ArrowError>
81    where
82        T: std::borrow::Borrow<Self>,
83    {
84        let mut ts_event_builder = UInt64Array::builder(data.len());
85        let mut ts_init_builder = UInt64Array::builder(data.len());
86
87        for bar in data.iter().map(std::borrow::Borrow::borrow) {
88            ts_event_builder.append_value(bar.ts_event.as_u64());
89            ts_init_builder.append_value(bar.ts_init.as_u64());
90        }
91
92        let ts_event_array = ts_event_builder.finish();
93        let ts_init_array = ts_init_builder.finish();
94
95        crate::arrow::record_batch_with_timestamps(
96            Self::get_schema(Some(metadata.clone())).into(),
97            vec![
98                Arc::new(required_price_decimal_array(
99                    data.iter()
100                        .map(std::borrow::Borrow::borrow)
101                        .map(|bar| bar.open.raw()),
102                    "open",
103                )?),
104                Arc::new(required_price_decimal_array(
105                    data.iter()
106                        .map(std::borrow::Borrow::borrow)
107                        .map(|bar| bar.high.raw()),
108                    "high",
109                )?),
110                Arc::new(required_price_decimal_array(
111                    data.iter()
112                        .map(std::borrow::Borrow::borrow)
113                        .map(|bar| bar.low.raw()),
114                    "low",
115                )?),
116                Arc::new(required_price_decimal_array(
117                    data.iter()
118                        .map(std::borrow::Borrow::borrow)
119                        .map(|bar| bar.close.raw()),
120                    "close",
121                )?),
122                Arc::new(required_quantity_decimal_array(
123                    data.iter()
124                        .map(std::borrow::Borrow::borrow)
125                        .map(|bar| bar.volume.raw()),
126                    "volume",
127                )?),
128                Arc::new(ts_event_array),
129                Arc::new(ts_init_array),
130                Arc::new(identifier_array_from_display(
131                    data.iter()
132                        .map(std::borrow::Borrow::borrow)
133                        .map(|bar| bar.bar_type),
134                )),
135            ],
136        )
137    }
138
139    fn metadata(&self) -> HashMap<String, String> {
140        Self::get_metadata(&self.bar_type, self.open.precision, self.volume.precision)
141    }
142}
143
144impl DecodeFromRecordBatch for Bar {
145    fn decode_batch(
146        metadata: &HashMap<String, String>,
147        record_batch: RecordBatch,
148    ) -> Result<Vec<Self>, EncodingError> {
149        let (bar_type, price_precision, size_precision) = parse_metadata(metadata)?;
150        let record_batch = crate::arrow::record_batch_with_u64_timestamps(&record_batch)?;
151        let record_batch = &record_batch;
152        let cols = record_batch.columns();
153
154        let open_values =
155            extract_column::<Decimal128Array>(cols, "open", 0, fixed_decimal_data_type())?;
156        let high_values =
157            extract_column::<Decimal128Array>(cols, "high", 1, fixed_decimal_data_type())?;
158        let low_values =
159            extract_column::<Decimal128Array>(cols, "low", 2, fixed_decimal_data_type())?;
160        let close_values =
161            extract_column::<Decimal128Array>(cols, "close", 3, fixed_decimal_data_type())?;
162        let volume_values =
163            extract_column::<Decimal128Array>(cols, "volume", 4, fixed_decimal_data_type())?;
164        let ts_event_values = extract_column::<UInt64Array>(cols, "ts_event", 5, DataType::UInt64)?;
165        let ts_init_values = extract_column::<UInt64Array>(cols, "ts_init", 6, DataType::UInt64)?;
166
167        let result: Result<Vec<Self>, EncodingError> = (0..record_batch.num_rows())
168            .map(|i| {
169                let open = decode_required_decimal_price(open_values, price_precision, "open", i)?;
170                let high = decode_required_decimal_price(high_values, price_precision, "high", i)?;
171                let low = decode_required_decimal_price(low_values, price_precision, "low", i)?;
172                let close =
173                    decode_required_decimal_price(close_values, price_precision, "close", i)?;
174                let volume =
175                    decode_required_decimal_quantity(volume_values, size_precision, "volume", i)?;
176                let ts_event = decode_required_timestamp(ts_event_values, "ts_event", i)?;
177                let ts_init = decode_required_timestamp(ts_init_values, "ts_init", i)?;
178
179                Ok(Self {
180                    bar_type,
181                    open,
182                    high,
183                    low,
184                    close,
185                    volume,
186                    ts_event,
187                    ts_init,
188                })
189            })
190            .collect();
191
192        result
193    }
194}
195
196impl DecodeDataFromRecordBatch for Bar {
197    fn decode_data_batch(
198        metadata: &HashMap<String, String>,
199        record_batch: RecordBatch,
200    ) -> Result<Vec<Data>, EncodingError> {
201        let bars: Vec<Self> = Self::decode_batch(metadata, record_batch)?;
202        Ok(bars.into_iter().map(Data::from).collect())
203    }
204}
205
206#[cfg(test)]
207mod tests {
208    use std::sync::Arc;
209
210    use arrow::array::{Array, TimestampNanosecondArray};
211    use nautilus_model::types::{
212        Price, Quantity, fixed::FIXED_SCALAR, price::PriceRaw, quantity::QuantityRaw,
213    };
214    use rstest::rstest;
215
216    use super::*;
217    use crate::arrow::{get_raw_price, get_raw_quantity};
218
219    #[rstest]
220    fn test_get_schema() {
221        let bar_type = BarType::from_str("AAPL.XNAS-1-MINUTE-LAST-INTERNAL").unwrap();
222        let metadata = Bar::get_metadata(&bar_type, 2, 0);
223        let schema = Bar::get_schema(Some(metadata.clone()));
224        let expected_fields = vec![
225            Field::new("open", fixed_decimal_data_type(), true),
226            Field::new("high", fixed_decimal_data_type(), true),
227            Field::new("low", fixed_decimal_data_type(), true),
228            Field::new("close", fixed_decimal_data_type(), true),
229            Field::new("volume", fixed_decimal_data_type(), true),
230            Field::new("ts_event", crate::arrow::timestamp_data_type(), false),
231            Field::new("ts_init", crate::arrow::timestamp_data_type(), false),
232            Field::new(KEY_IDENTIFIER, DataType::Utf8, true),
233        ];
234        let expected_schema = Schema::new_with_metadata(expected_fields, metadata);
235        assert_eq!(schema, expected_schema);
236    }
237
238    #[rstest]
239    fn test_get_schema_map() {
240        let schema_map = Bar::get_schema_map();
241        let mut expected_map = HashMap::new();
242        let fixed_size_binary = "Decimal128(38, 16)".to_string();
243        expected_map.insert("open".to_string(), fixed_size_binary.clone());
244        expected_map.insert("high".to_string(), fixed_size_binary.clone());
245        expected_map.insert("low".to_string(), fixed_size_binary.clone());
246        expected_map.insert("close".to_string(), fixed_size_binary.clone());
247        expected_map.insert("volume".to_string(), fixed_size_binary);
248        expected_map.insert(
249            "ts_event".to_string(),
250            "Timestamp(Nanosecond, Some(\"UTC\"))".to_string(),
251        );
252        expected_map.insert(
253            "ts_init".to_string(),
254            "Timestamp(Nanosecond, Some(\"UTC\"))".to_string(),
255        );
256        expected_map.insert(KEY_IDENTIFIER.to_string(), "Utf8".to_string());
257        assert_eq!(schema_map, expected_map);
258    }
259
260    #[rstest]
261    fn test_encode_batch() {
262        let bar_type = BarType::from_str("AAPL.XNAS-1-MINUTE-LAST-INTERNAL").unwrap();
263        let metadata = Bar::get_metadata(&bar_type, 2, 0);
264
265        let bar1 = Bar::new(
266            bar_type,
267            Price::from("100.10"),
268            Price::from("102.00"),
269            Price::from("100.00"),
270            Price::from("101.00"),
271            Quantity::from(1100),
272            1.into(),
273            3.into(),
274        );
275        let bar2 = Bar::new(
276            bar_type,
277            Price::from("100.00"),
278            Price::from("100.10"),
279            Price::from("100.00"),
280            Price::from("100.10"),
281            Quantity::from(1110),
282            2.into(),
283            4.into(),
284        );
285
286        let data = vec![bar1, bar2];
287        let record_batch = Bar::encode_batch(&metadata, &data).unwrap();
288
289        let columns = record_batch.columns();
290        let open_values = columns[0]
291            .as_any()
292            .downcast_ref::<Decimal128Array>()
293            .unwrap();
294        let high_values = columns[1]
295            .as_any()
296            .downcast_ref::<Decimal128Array>()
297            .unwrap();
298        let low_values = columns[2]
299            .as_any()
300            .downcast_ref::<Decimal128Array>()
301            .unwrap();
302        let close_values = columns[3]
303            .as_any()
304            .downcast_ref::<Decimal128Array>()
305            .unwrap();
306        let volume_values = columns[4]
307            .as_any()
308            .downcast_ref::<Decimal128Array>()
309            .unwrap();
310        let ts_event_values = columns[5]
311            .as_any()
312            .downcast_ref::<TimestampNanosecondArray>()
313            .unwrap();
314        let ts_init_values = columns[6]
315            .as_any()
316            .downcast_ref::<TimestampNanosecondArray>()
317            .unwrap();
318
319        assert_eq!(columns.len(), 8);
320        assert_eq!(open_values.len(), 2);
321        assert_eq!(
322            get_raw_price(open_values.value(0)),
323            (100.10 * FIXED_SCALAR) as PriceRaw
324        );
325        assert_eq!(
326            get_raw_price(open_values.value(1)),
327            (100.00 * FIXED_SCALAR) as PriceRaw
328        );
329        assert_eq!(high_values.len(), 2);
330        assert_eq!(
331            get_raw_price(high_values.value(0)),
332            (102.00 * FIXED_SCALAR) as PriceRaw
333        );
334        assert_eq!(
335            get_raw_price(high_values.value(1)),
336            (100.10 * FIXED_SCALAR) as PriceRaw
337        );
338        assert_eq!(low_values.len(), 2);
339        assert_eq!(
340            get_raw_price(low_values.value(0)),
341            (100.00 * FIXED_SCALAR) as PriceRaw
342        );
343        assert_eq!(
344            get_raw_price(low_values.value(1)),
345            (100.00 * FIXED_SCALAR) as PriceRaw
346        );
347        assert_eq!(close_values.len(), 2);
348        assert_eq!(
349            get_raw_price(close_values.value(0)),
350            (101.00 * FIXED_SCALAR) as PriceRaw
351        );
352        assert_eq!(
353            get_raw_price(close_values.value(1)),
354            (100.10 * FIXED_SCALAR) as PriceRaw
355        );
356        assert_eq!(volume_values.len(), 2);
357        assert_eq!(
358            get_raw_quantity(volume_values.value(0)),
359            (1100.0 * FIXED_SCALAR) as QuantityRaw
360        );
361        assert_eq!(
362            get_raw_quantity(volume_values.value(1)),
363            (1110.0 * FIXED_SCALAR) as QuantityRaw
364        );
365        assert_eq!(ts_event_values.len(), 2);
366        assert_eq!(ts_event_values.value(0), 1);
367        assert_eq!(ts_event_values.value(1), 2);
368        assert_eq!(ts_init_values.len(), 2);
369        assert_eq!(ts_init_values.value(0), 3);
370        assert_eq!(ts_init_values.value(1), 4);
371    }
372
373    #[rstest]
374    fn test_decode_batch() {
375        let bar_type = BarType::from_str("AAPL.XNAS-1-MINUTE-LAST-INTERNAL").unwrap();
376        let metadata = Bar::get_metadata(&bar_type, 2, 0);
377
378        let open = crate::arrow::test_support::decimal_array_from_bytes(vec![
379            &((100.10 * FIXED_SCALAR) as PriceRaw).to_le_bytes(),
380            &((10.00 * FIXED_SCALAR) as PriceRaw).to_le_bytes(),
381        ]);
382        let high = crate::arrow::test_support::decimal_array_from_bytes(vec![
383            &((102.00 * FIXED_SCALAR) as PriceRaw).to_le_bytes(),
384            &((10.00 * FIXED_SCALAR) as PriceRaw).to_le_bytes(),
385        ]);
386        let low = crate::arrow::test_support::decimal_array_from_bytes(vec![
387            &((100.00 * FIXED_SCALAR) as PriceRaw).to_le_bytes(),
388            &((10.00 * FIXED_SCALAR) as PriceRaw).to_le_bytes(),
389        ]);
390        let close = crate::arrow::test_support::decimal_array_from_bytes(vec![
391            &((101.00 * FIXED_SCALAR) as PriceRaw).to_le_bytes(),
392            &((10.01 * FIXED_SCALAR) as PriceRaw).to_le_bytes(),
393        ]);
394        let volume = crate::arrow::test_support::decimal_array_from_bytes(vec![
395            &((11.0 * FIXED_SCALAR) as QuantityRaw).to_le_bytes(),
396            &((10.0 * FIXED_SCALAR) as QuantityRaw).to_le_bytes(),
397        ]);
398        let ts_event = UInt64Array::from(vec![1, 2]);
399        let ts_init = UInt64Array::from(vec![3, 4]);
400
401        let record_batch = crate::arrow::record_batch_with_timestamps(
402            crate::arrow::schema_without_identifier_column(&Bar::get_schema(Some(
403                metadata.clone(),
404            )))
405            .into(),
406            vec![
407                Arc::new(open),
408                Arc::new(high),
409                Arc::new(low),
410                Arc::new(close),
411                Arc::new(volume),
412                Arc::new(ts_event),
413                Arc::new(ts_init),
414            ],
415        )
416        .unwrap();
417
418        let decoded_data = Bar::decode_batch(&metadata, record_batch).unwrap();
419        assert_eq!(decoded_data.len(), 2);
420    }
421
422    #[rstest]
423    fn test_decode_batch_invalid_price_returns_error() {
424        let bar_type = BarType::from_str("AAPL.XNAS-1-MINUTE-LAST-INTERNAL").unwrap();
425        let metadata = Bar::get_metadata(&bar_type, 2, 0);
426
427        let invalid_price: PriceRaw = PriceRaw::MAX - 1000;
428        let valid_price = (100.00 * FIXED_SCALAR) as PriceRaw;
429
430        let open = crate::arrow::test_support::decimal_array_from_bytes(vec![
431            &invalid_price.to_le_bytes(),
432        ]);
433        let high =
434            crate::arrow::test_support::decimal_array_from_bytes(vec![&valid_price.to_le_bytes()]);
435        let low =
436            crate::arrow::test_support::decimal_array_from_bytes(vec![&valid_price.to_le_bytes()]);
437        let close =
438            crate::arrow::test_support::decimal_array_from_bytes(vec![&valid_price.to_le_bytes()]);
439        let volume = crate::arrow::test_support::decimal_array_from_bytes(vec![
440            &((100.0 * FIXED_SCALAR) as QuantityRaw).to_le_bytes(),
441        ]);
442        let ts_event = UInt64Array::from(vec![1]);
443        let ts_init = UInt64Array::from(vec![2]);
444
445        let record_batch = crate::arrow::record_batch_with_timestamps(
446            crate::arrow::schema_without_identifier_column(&Bar::get_schema(Some(
447                metadata.clone(),
448            )))
449            .into(),
450            vec![
451                Arc::new(open),
452                Arc::new(high),
453                Arc::new(low),
454                Arc::new(close),
455                Arc::new(volume),
456                Arc::new(ts_event),
457                Arc::new(ts_init),
458            ],
459        )
460        .unwrap();
461
462        let result = Bar::decode_batch(&metadata, record_batch);
463        assert!(result.is_err());
464        let err = result.unwrap_err();
465        assert!(
466            err.to_string().contains("open") && err.to_string().contains("row 0"),
467            "Expected open error at row 0, was: {err}"
468        );
469    }
470
471    #[rstest]
472    fn test_decode_batch_missing_bar_type_returns_error() {
473        let bar_type = BarType::from_str("AAPL.XNAS-1-MINUTE-LAST-INTERNAL").unwrap();
474        let mut metadata = Bar::get_metadata(&bar_type, 2, 0);
475
476        let valid_price = (100.00 * FIXED_SCALAR) as PriceRaw;
477        let open =
478            crate::arrow::test_support::decimal_array_from_bytes(vec![&valid_price.to_le_bytes()]);
479        let high =
480            crate::arrow::test_support::decimal_array_from_bytes(vec![&valid_price.to_le_bytes()]);
481        let low =
482            crate::arrow::test_support::decimal_array_from_bytes(vec![&valid_price.to_le_bytes()]);
483        let close =
484            crate::arrow::test_support::decimal_array_from_bytes(vec![&valid_price.to_le_bytes()]);
485        let volume = crate::arrow::test_support::decimal_array_from_bytes(vec![
486            &((100.0 * FIXED_SCALAR) as QuantityRaw).to_le_bytes(),
487        ]);
488        let ts_event = UInt64Array::from(vec![1]);
489        let ts_init = UInt64Array::from(vec![2]);
490
491        let record_batch = crate::arrow::record_batch_with_timestamps(
492            crate::arrow::schema_without_identifier_column(&Bar::get_schema(Some(
493                metadata.clone(),
494            )))
495            .into(),
496            vec![
497                Arc::new(open),
498                Arc::new(high),
499                Arc::new(low),
500                Arc::new(close),
501                Arc::new(volume),
502                Arc::new(ts_event),
503                Arc::new(ts_init),
504            ],
505        )
506        .unwrap();
507
508        metadata.remove(KEY_BAR_TYPE);
509
510        let result = Bar::decode_batch(&metadata, record_batch);
511        assert!(result.is_err());
512        let err = result.unwrap_err();
513        assert!(
514            err.to_string().contains("bar_type"),
515            "Expected missing bar_type error, was: {err}"
516        );
517    }
518
519    #[rstest]
520    fn test_encode_decode_round_trip() {
521        let bar_type = BarType::from_str("AAPL.XNAS-1-MINUTE-LAST-INTERNAL").unwrap();
522        let metadata = Bar::get_metadata(&bar_type, 2, 0);
523
524        let bar1 = Bar::new(
525            bar_type,
526            Price::from("100.10"),
527            Price::from("102.00"),
528            Price::from("100.00"),
529            Price::from("101.00"),
530            Quantity::from(1100),
531            1_000_000_000.into(),
532            1_000_000_001.into(),
533        );
534
535        let bar2 = Bar::new(
536            bar_type,
537            Price::from("101.00"),
538            Price::from("103.00"),
539            Price::from("100.50"),
540            Price::from("102.50"),
541            Quantity::from(2200),
542            2_000_000_000.into(),
543            2_000_000_001.into(),
544        );
545
546        let original = vec![bar1, bar2];
547        let record_batch = Bar::encode_batch(&metadata, &original).unwrap();
548        let decoded = Bar::decode_batch(&metadata, record_batch).unwrap();
549
550        assert_eq!(decoded.len(), original.len());
551        for (orig, dec) in original.iter().zip(decoded.iter()) {
552            assert_eq!(dec.bar_type, orig.bar_type);
553            assert_eq!(dec.open, orig.open);
554            assert_eq!(dec.high, orig.high);
555            assert_eq!(dec.low, orig.low);
556            assert_eq!(dec.close, orig.close);
557            assert_eq!(dec.volume, orig.volume);
558            assert_eq!(dec.ts_event, orig.ts_event);
559            assert_eq!(dec.ts_init, orig.ts_init);
560        }
561    }
562}