1use 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}