Skip to main content

nautilus_binance/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::{Array, Decimal128Array, UInt64Array},
20    datatypes::{DataType, Field, Schema},
21    error::ArrowError,
22    record_batch::RecordBatch,
23};
24use nautilus_model::data::{Data, bar::BarType, custom::CustomData};
25use nautilus_serialization::arrow::{
26    ArrowSchemaProvider, DecodeDataFromRecordBatch, EncodeToRecordBatch, EncodingError,
27    FIXED_DECIMAL_PRECISION, FIXED_DECIMAL_SCALE, KEY_PRICE_PRECISION, KEY_SIZE_PRECISION,
28    StringColumnRef, decimal_to_arrow, decode_decimal, decode_decimal_price,
29    decode_decimal_quantity, extract_column, extract_decimal_column, fixed_decimal_data_type,
30    price_decimal_array, quantity_decimal_array, record_batch_with_timestamps,
31    record_batch_with_u64_timestamps, timestamp_data_type,
32};
33use rust_decimal::Decimal;
34
35use crate::common::bar::BinanceBar;
36
37const KEY_BAR_TYPE: &str = "bar_type";
38
39fn parse_metadata(metadata: &HashMap<String, String>) -> Result<(BarType, u8, u8), EncodingError> {
40    let bar_type_str = metadata
41        .get(KEY_BAR_TYPE)
42        .ok_or_else(|| EncodingError::MissingMetadata(KEY_BAR_TYPE))?;
43    let bar_type = BarType::from_str(bar_type_str)
44        .map_err(|e| EncodingError::ParseError(KEY_BAR_TYPE, e.to_string()))?;
45
46    let price_precision = metadata
47        .get(KEY_PRICE_PRECISION)
48        .ok_or_else(|| EncodingError::MissingMetadata(KEY_PRICE_PRECISION))?
49        .parse::<u8>()
50        .map_err(|e| EncodingError::ParseError(KEY_PRICE_PRECISION, e.to_string()))?;
51
52    let size_precision = metadata
53        .get(KEY_SIZE_PRECISION)
54        .ok_or_else(|| EncodingError::MissingMetadata(KEY_SIZE_PRECISION))?
55        .parse::<u8>()
56        .map_err(|e| EncodingError::ParseError(KEY_SIZE_PRECISION, e.to_string()))?;
57
58    Ok((bar_type, price_precision, size_precision))
59}
60
61impl ArrowSchemaProvider for BinanceBar {
62    fn get_schema(metadata: Option<HashMap<String, String>>) -> Schema {
63        let fields = vec![
64            Field::new("open", fixed_decimal_data_type(), true),
65            Field::new("high", fixed_decimal_data_type(), true),
66            Field::new("low", fixed_decimal_data_type(), true),
67            Field::new("close", fixed_decimal_data_type(), true),
68            Field::new("volume", fixed_decimal_data_type(), true),
69            Field::new("quote_volume", fixed_decimal_data_type(), false),
70            Field::new("count", DataType::UInt64, false),
71            Field::new("taker_buy_base_volume", fixed_decimal_data_type(), false),
72            Field::new("taker_buy_quote_volume", fixed_decimal_data_type(), false),
73            Field::new("ts_event", timestamp_data_type(), false),
74            Field::new("ts_init", timestamp_data_type(), false),
75        ];
76
77        match metadata {
78            Some(metadata) => Schema::new_with_metadata(fields, metadata),
79            None => Schema::new(fields),
80        }
81    }
82}
83
84impl EncodeToRecordBatch for BinanceBar {
85    fn encode_batch<T>(
86        metadata: &HashMap<String, String>,
87        data: &[T],
88    ) -> Result<RecordBatch, ArrowError>
89    where
90        T: std::borrow::Borrow<Self>,
91    {
92        let mut count_builder = UInt64Array::builder(data.len());
93        let mut ts_event_builder = UInt64Array::builder(data.len());
94        let mut ts_init_builder = UInt64Array::builder(data.len());
95
96        for bar in data.iter().map(std::borrow::Borrow::borrow) {
97            count_builder.append_value(bar.count);
98            ts_event_builder.append_value(bar.ts_event.as_u64());
99            ts_init_builder.append_value(bar.ts_init.as_u64());
100        }
101
102        record_batch_with_timestamps(
103            Self::get_schema(Some(metadata.clone())).into(),
104            vec![
105                Arc::new(price_decimal_array(
106                    data.iter()
107                        .map(std::borrow::Borrow::borrow)
108                        .map(|bar| bar.open.raw()),
109                    "open",
110                )?),
111                Arc::new(price_decimal_array(
112                    data.iter()
113                        .map(std::borrow::Borrow::borrow)
114                        .map(|bar| bar.high.raw()),
115                    "high",
116                )?),
117                Arc::new(price_decimal_array(
118                    data.iter()
119                        .map(std::borrow::Borrow::borrow)
120                        .map(|bar| bar.low.raw()),
121                    "low",
122                )?),
123                Arc::new(price_decimal_array(
124                    data.iter()
125                        .map(std::borrow::Borrow::borrow)
126                        .map(|bar| bar.close.raw()),
127                    "close",
128                )?),
129                Arc::new(quantity_decimal_array(
130                    data.iter()
131                        .map(std::borrow::Borrow::borrow)
132                        .map(|bar| bar.volume.raw()),
133                    "volume",
134                )?),
135                Arc::new(decimal_array(
136                    data.iter()
137                        .map(std::borrow::Borrow::borrow)
138                        .map(|bar| &bar.quote_volume),
139                    "quote_volume",
140                )?),
141                Arc::new(count_builder.finish()),
142                Arc::new(decimal_array(
143                    data.iter()
144                        .map(std::borrow::Borrow::borrow)
145                        .map(|bar| &bar.taker_buy_base_volume),
146                    "taker_buy_base_volume",
147                )?),
148                Arc::new(decimal_array(
149                    data.iter()
150                        .map(std::borrow::Borrow::borrow)
151                        .map(|bar| &bar.taker_buy_quote_volume),
152                    "taker_buy_quote_volume",
153                )?),
154                Arc::new(ts_event_builder.finish()),
155                Arc::new(ts_init_builder.finish()),
156            ],
157        )
158    }
159
160    fn metadata(&self) -> HashMap<String, String> {
161        let mut metadata = Self::get_metadata(&self.bar_type);
162        metadata.insert(
163            KEY_PRICE_PRECISION.to_string(),
164            self.open.precision.to_string(),
165        );
166        metadata.insert(
167            KEY_SIZE_PRECISION.to_string(),
168            self.volume.precision.to_string(),
169        );
170        metadata
171    }
172}
173
174/// Encodes a vector of [`BinanceBar`] into an Arrow `RecordBatch`.
175///
176/// # Errors
177///
178/// Returns an error if `data` is empty or encoding fails.
179#[expect(clippy::missing_panics_doc)] // Guarded by empty check
180pub fn binance_bar_to_arrow_record_batch(
181    data: &[BinanceBar],
182) -> Result<RecordBatch, EncodingError> {
183    if data.is_empty() {
184        return Err(EncodingError::EmptyData);
185    }
186
187    let first = data
188        .first()
189        .expect("Chunk should have at least one element to encode");
190    let metadata = first.metadata();
191    BinanceBar::encode_batch(&metadata, data).map_err(EncodingError::ArrowError)
192}
193
194/// Decodes a `RecordBatch` into a vector of [`BinanceBar`].
195///
196/// # Errors
197///
198/// Returns an `EncodingError` if decoding fails.
199pub fn decode_binance_bar_batch(
200    metadata: &HashMap<String, String>,
201    record_batch: &RecordBatch,
202) -> Result<Vec<BinanceBar>, EncodingError> {
203    let (bar_type, price_precision, size_precision) = parse_metadata(metadata)?;
204    let record_batch = record_batch_with_u64_timestamps(record_batch)?;
205    let cols = record_batch.columns();
206
207    let open_values =
208        extract_column::<Decimal128Array>(cols, "open", 0, fixed_decimal_data_type())?;
209    let high_values =
210        extract_column::<Decimal128Array>(cols, "high", 1, fixed_decimal_data_type())?;
211    let low_values = extract_column::<Decimal128Array>(cols, "low", 2, fixed_decimal_data_type())?;
212    let close_values =
213        extract_column::<Decimal128Array>(cols, "close", 3, fixed_decimal_data_type())?;
214    let volume_values =
215        extract_column::<Decimal128Array>(cols, "volume", 4, fixed_decimal_data_type())?;
216    let count_values = extract_column::<UInt64Array>(cols, "count", 6, DataType::UInt64)?;
217    let ts_event_values = extract_column::<UInt64Array>(cols, "ts_event", 9, DataType::UInt64)?;
218    let ts_init_values = extract_column::<UInt64Array>(cols, "ts_init", 10, DataType::UInt64)?;
219
220    (0..record_batch.num_rows())
221        .map(|row| {
222            let open = decode_decimal_price(open_values, price_precision, "open", row)?;
223            let high = decode_decimal_price(high_values, price_precision, "high", row)?;
224            let low = decode_decimal_price(low_values, price_precision, "low", row)?;
225            let close = decode_decimal_price(close_values, price_precision, "close", row)?;
226            let volume = decode_decimal_quantity(volume_values, size_precision, "volume", row)?;
227            let quote_volume = decode_decimal_column(&record_batch, "quote_volume", row)?;
228            let taker_buy_base_volume =
229                decode_decimal_column(&record_batch, "taker_buy_base_volume", row)?;
230            let taker_buy_quote_volume =
231                decode_decimal_column(&record_batch, "taker_buy_quote_volume", row)?;
232
233            Ok(BinanceBar::new(
234                bar_type,
235                open,
236                high,
237                low,
238                close,
239                volume,
240                quote_volume,
241                count_values.value(row),
242                taker_buy_base_volume,
243                taker_buy_quote_volume,
244                ts_event_values.value(row).into(),
245                ts_init_values.value(row).into(),
246            ))
247        })
248        .collect()
249}
250
251fn decimal_array<'a>(
252    values: impl IntoIterator<Item = &'a Decimal>,
253    field: &'static str,
254) -> Result<Decimal128Array, ArrowError> {
255    let values = values
256        .into_iter()
257        .map(|value| decimal_to_arrow(value, field))
258        .collect::<Result<Vec<_>, _>>()?;
259    Decimal128Array::from(values)
260        .with_precision_and_scale(FIXED_DECIMAL_PRECISION, FIXED_DECIMAL_SCALE)
261}
262
263fn decode_decimal_column(
264    record_batch: &RecordBatch,
265    field: &'static str,
266    row: usize,
267) -> Result<Decimal, EncodingError> {
268    let index = record_batch.schema().index_of(field)?;
269    let column = record_batch
270        .columns()
271        .get(index)
272        .ok_or(EncodingError::MissingColumn(field, index))?;
273    if column.data_type() == &fixed_decimal_data_type() {
274        let values = extract_decimal_column(record_batch, field)?;
275        return decode_decimal(values, field, row);
276    }
277    let values = StringColumnRef::try_from_array(column.as_ref()).ok_or_else(|| {
278        EncodingError::ParseError(
279            field,
280            format!(
281                "expected Decimal128(38, 16) or legacy string, was {}",
282                column.data_type()
283            ),
284        )
285    })?;
286
287    if values.is_null(row) {
288        return Err(EncodingError::ParseError(
289            field,
290            format!("row {row}: required decimal is null"),
291        ));
292    }
293    Decimal::from_str(values.value(row))
294        .map_err(|e| EncodingError::ParseError(field, format!("row {row}: {e}")))
295}
296
297impl DecodeDataFromRecordBatch for BinanceBar {
298    fn decode_data_batch(
299        metadata: &HashMap<String, String>,
300        record_batch: RecordBatch,
301    ) -> Result<Vec<Data>, EncodingError> {
302        let items = decode_binance_bar_batch(metadata, &record_batch)?;
303        Ok(items
304            .into_iter()
305            .map(|item| Data::Custom(CustomData::from_arc(Arc::new(item))))
306            .collect())
307    }
308}
309
310#[cfg(test)]
311mod tests {
312    use arrow::array::StringArray;
313    use nautilus_model::types::{Price, Quantity};
314    use rstest::rstest;
315    use rust_decimal_macros::dec;
316
317    use super::*;
318
319    fn stub_binance_bar() -> BinanceBar {
320        BinanceBar::new(
321            BarType::from("BTCUSDT.BINANCE-1-MINUTE-LAST-EXTERNAL"),
322            Price::from("0.01634790"),
323            Price::from("0.01640000"),
324            Price::from("0.01575800"),
325            Price::from("0.01577100"),
326            Quantity::from("148976.11427815"),
327            dec!(2434.19055334),
328            100,
329            dec!(1756.87402397),
330            dec!(28.46694368),
331            1_650_000_000_000_000_000u64.into(),
332            1_650_000_000_000_000_000u64.into(),
333        )
334    }
335
336    #[rstest]
337    fn test_get_schema() {
338        let schema = BinanceBar::get_schema(None);
339        assert_eq!(schema.fields().len(), 11);
340        assert_eq!(schema.field(0).name(), "open");
341        assert_eq!(schema.field(0).data_type(), &fixed_decimal_data_type());
342        assert_eq!(schema.field(5).name(), "quote_volume");
343        assert_eq!(schema.field(5).data_type(), &fixed_decimal_data_type());
344        assert_eq!(schema.field(6).name(), "count");
345        assert_eq!(schema.field(6).data_type(), &DataType::UInt64);
346        assert_eq!(schema.field(9).data_type(), &timestamp_data_type());
347        assert_eq!(schema.field(10).data_type(), &timestamp_data_type());
348    }
349
350    #[rstest]
351    fn test_encode_decode_round_trip() {
352        let bar = stub_binance_bar();
353        let metadata = bar.metadata();
354        let data = vec![bar.clone()];
355
356        let record_batch = BinanceBar::encode_batch(&metadata, &data).unwrap();
357        let decoded = decode_binance_bar_batch(&metadata, &record_batch).unwrap();
358
359        assert_eq!(decoded.len(), 1);
360        assert_eq!(decoded[0], bar);
361    }
362
363    #[rstest]
364    fn test_encode_decode_multiple_bars() {
365        let bar1 = stub_binance_bar();
366        let bar2 = BinanceBar::new(
367            BarType::from("BTCUSDT.BINANCE-1-MINUTE-LAST-EXTERNAL"),
368            Price::from("0.01700000"),
369            Price::from("0.01710000"),
370            Price::from("0.01690000"),
371            Price::from("0.01695000"),
372            Quantity::from("50000.00000000"),
373            dec!(1000.00000000),
374            50,
375            dec!(500.00000000),
376            dec!(10.00000000),
377            1_650_000_060_000_000_000u64.into(),
378            1_650_000_060_000_000_000u64.into(),
379        );
380
381        let metadata = bar1.metadata();
382        let data = vec![bar1.clone(), bar2.clone()];
383
384        let record_batch = BinanceBar::encode_batch(&metadata, &data).unwrap();
385        let decoded = decode_binance_bar_batch(&metadata, &record_batch).unwrap();
386
387        assert_eq!(decoded.len(), 2);
388        assert_eq!(decoded[0], bar1);
389        assert_eq!(decoded[1], bar2);
390    }
391
392    #[rstest]
393    fn test_decode_data_batch_returns_custom_data() {
394        let bar = stub_binance_bar();
395        let metadata = bar.metadata();
396        let data = vec![bar];
397
398        let record_batch = BinanceBar::encode_batch(&metadata, &data).unwrap();
399        let decoded = BinanceBar::decode_data_batch(&metadata, record_batch).unwrap();
400
401        assert_eq!(decoded.len(), 1);
402        assert!(matches!(decoded[0], Data::Custom(_)));
403    }
404
405    #[rstest]
406    fn test_decode_legacy_string_decimal_columns() {
407        let bar = stub_binance_bar();
408        let metadata = bar.metadata();
409        let batch = legacy_string_batch(&bar, Some("2434.19055334"));
410
411        let decoded = decode_binance_bar_batch(&metadata, &batch).unwrap();
412
413        assert_eq!(decoded, vec![bar]);
414    }
415
416    #[rstest]
417    fn test_decode_legacy_string_decimal_rejects_null() {
418        let bar = stub_binance_bar();
419        let metadata = bar.metadata();
420        let batch = legacy_string_batch(&bar, None);
421
422        let error = decode_binance_bar_batch(&metadata, &batch).unwrap_err();
423
424        assert_eq!(
425            error.to_string(),
426            "Error parsing `quote_volume`: row 0: required decimal is null",
427        );
428    }
429
430    fn legacy_string_batch(bar: &BinanceBar, quote_volume: Option<&str>) -> RecordBatch {
431        let metadata = bar.metadata();
432        let batch = BinanceBar::encode_batch(&metadata, &[bar]).unwrap();
433        let mut fields = batch
434            .schema()
435            .fields()
436            .iter()
437            .map(|field| field.as_ref().clone())
438            .collect::<Vec<_>>();
439        fields[5] = Field::new("quote_volume", DataType::Utf8, true);
440        fields[7] = Field::new("taker_buy_base_volume", DataType::Utf8, false);
441        fields[8] = Field::new("taker_buy_quote_volume", DataType::Utf8, false);
442        let schema = Schema::new_with_metadata(fields, metadata);
443        let mut columns = batch.columns().to_vec();
444        columns[5] = Arc::new(StringArray::from(vec![quote_volume]));
445        columns[7] = Arc::new(StringArray::from(vec![Some("1756.87402397")]));
446        columns[8] = Arc::new(StringArray::from(vec![Some("28.46694368")]));
447
448        RecordBatch::try_new(Arc::new(schema), columns).unwrap()
449    }
450}