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