1use 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#[expect(clippy::missing_panics_doc)] pub 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
194pub 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(), ×tamp_data_type());
347 assert_eq!(schema.field(10).data_type(), ×tamp_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}