1use std::{collections::HashMap, sync::Arc};
17
18use arrow::{
19 array::{FixedSizeBinaryArray, FixedSizeBinaryBuilder, UInt8Array, UInt64Array},
20 datatypes::{DataType, Field, Schema},
21 error::ArrowError,
22 record_batch::RecordBatch,
23};
24use nautilus_model::{
25 data::close::InstrumentClose,
26 enums::{FromU8, InstrumentCloseType},
27 types::fixed::PRECISION_BYTES,
28};
29
30use super::{
31 DecodeDataFromRecordBatch, EncodingError, decode_price, extract_column, parse_price_metadata,
32 validate_precision_bytes,
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(
40 "close_price",
41 DataType::FixedSizeBinary(PRECISION_BYTES),
42 false,
43 ),
44 Field::new("close_type", DataType::UInt8, false),
45 Field::new("ts_event", DataType::UInt64, false),
46 Field::new("ts_init", DataType::UInt64, false),
47 ];
48
49 match metadata {
50 Some(metadata) => Schema::new_with_metadata(fields, metadata),
51 None => Schema::new(fields),
52 }
53 }
54}
55
56impl EncodeToRecordBatch for InstrumentClose {
57 fn encode_batch(
58 metadata: &HashMap<String, String>,
59 data: &[Self],
60 ) -> Result<RecordBatch, ArrowError> {
61 let mut close_price_builder =
62 FixedSizeBinaryBuilder::with_capacity(data.len(), PRECISION_BYTES);
63 let mut close_type_builder = UInt8Array::builder(data.len());
64 let mut ts_event_builder = UInt64Array::builder(data.len());
65 let mut ts_init_builder = UInt64Array::builder(data.len());
66
67 for item in data {
68 close_price_builder
69 .append_value(item.close_price.raw().to_le_bytes())
70 .unwrap();
71 close_type_builder.append_value(item.close_type as u8);
72 ts_event_builder.append_value(item.ts_event.as_u64());
73 ts_init_builder.append_value(item.ts_init.as_u64());
74 }
75
76 RecordBatch::try_new(
77 Self::get_schema(Some(metadata.clone())).into(),
78 vec![
79 Arc::new(close_price_builder.finish()),
80 Arc::new(close_type_builder.finish()),
81 Arc::new(ts_event_builder.finish()),
82 Arc::new(ts_init_builder.finish()),
83 ],
84 )
85 }
86
87 fn metadata(&self) -> HashMap<String, String> {
88 Self::get_metadata(&self.instrument_id, self.close_price.precision)
89 }
90}
91
92impl DecodeFromRecordBatch for InstrumentClose {
93 fn decode_batch(
94 metadata: &HashMap<String, String>,
95 record_batch: RecordBatch,
96 ) -> Result<Vec<Self>, EncodingError> {
97 let (instrument_id, price_precision) = parse_price_metadata(metadata)?;
98 let cols = record_batch.columns();
99
100 let close_price_values = extract_column::<FixedSizeBinaryArray>(
101 cols,
102 "close_price",
103 0,
104 DataType::FixedSizeBinary(PRECISION_BYTES),
105 )?;
106 let close_type_values =
107 extract_column::<UInt8Array>(cols, "close_type", 1, DataType::UInt8)?;
108 let ts_event_values = extract_column::<UInt64Array>(cols, "ts_event", 2, DataType::UInt64)?;
109 let ts_init_values = extract_column::<UInt64Array>(cols, "ts_init", 3, DataType::UInt64)?;
110
111 validate_precision_bytes(close_price_values, "close_price")?;
112
113 let result: Result<Vec<Self>, EncodingError> = (0..record_batch.num_rows())
114 .map(|row| {
115 let close_price = decode_price(
116 close_price_values.value(row),
117 price_precision,
118 "close_price",
119 row,
120 )?;
121 let close_type_value = close_type_values.value(row);
122 let close_type =
123 InstrumentCloseType::from_u8(close_type_value).ok_or_else(|| {
124 EncodingError::ParseError(
125 stringify!(InstrumentCloseType),
126 format!("Invalid enum value, was {close_type_value}"),
127 )
128 })?;
129 Ok(Self {
130 instrument_id,
131 close_price,
132 close_type,
133 ts_event: ts_event_values.value(row).into(),
134 ts_init: ts_init_values.value(row).into(),
135 })
136 })
137 .collect();
138
139 result
140 }
141}
142
143impl DecodeDataFromRecordBatch for InstrumentClose {
144 fn decode_data_batch(
145 metadata: &HashMap<String, String>,
146 record_batch: RecordBatch,
147 ) -> Result<Vec<Data>, EncodingError> {
148 let items: Vec<Self> = Self::decode_batch(metadata, record_batch)?;
149 Ok(items.into_iter().map(Data::from).collect())
150 }
151}
152
153#[cfg(test)]
154mod tests {
155 use std::sync::Arc;
156
157 use arrow::{array::Array, record_batch::RecordBatch};
158 use nautilus_model::{
159 identifiers::InstrumentId,
160 types::{Price, fixed::FIXED_SCALAR, price::PriceRaw},
161 };
162 use rstest::rstest;
163
164 use super::*;
165 use crate::arrow::{KEY_INSTRUMENT_ID, KEY_PRICE_PRECISION, fixed_size_binary, get_raw_price};
166
167 #[rstest]
168 fn test_get_schema() {
169 let instrument_id = InstrumentId::from("AAPL.XNAS");
170 let metadata = HashMap::from([
171 (KEY_INSTRUMENT_ID.to_string(), instrument_id.to_string()),
172 (KEY_PRICE_PRECISION.to_string(), "2".to_string()),
173 ]);
174 let schema = InstrumentClose::get_schema(Some(metadata.clone()));
175
176 let expected_fields = vec![
177 Field::new(
178 "close_price",
179 DataType::FixedSizeBinary(PRECISION_BYTES),
180 false,
181 ),
182 Field::new("close_type", DataType::UInt8, false),
183 Field::new("ts_event", DataType::UInt64, false),
184 Field::new("ts_init", DataType::UInt64, false),
185 ];
186
187 let expected_schema = Schema::new_with_metadata(expected_fields, metadata);
188 assert_eq!(schema, expected_schema);
189 }
190
191 #[rstest]
192 fn test_get_schema_map() {
193 let schema_map = InstrumentClose::get_schema_map();
194 let mut expected_map = HashMap::new();
195
196 let fixed_size_binary = format!("FixedSizeBinary({PRECISION_BYTES})");
197 expected_map.insert("close_price".to_string(), fixed_size_binary);
198 expected_map.insert("close_type".to_string(), "UInt8".to_string());
199 expected_map.insert("ts_event".to_string(), "UInt64".to_string());
200 expected_map.insert("ts_init".to_string(), "UInt64".to_string());
201 assert_eq!(schema_map, expected_map);
202 }
203
204 #[rstest]
205 fn test_encode_batch() {
206 let instrument_id = InstrumentId::from("AAPL.XNAS");
207 let metadata = HashMap::from([
208 (KEY_INSTRUMENT_ID.to_string(), instrument_id.to_string()),
209 (KEY_PRICE_PRECISION.to_string(), "2".to_string()),
210 ]);
211
212 let close1 = InstrumentClose {
213 instrument_id,
214 close_price: Price::from("150.50"),
215 close_type: InstrumentCloseType::EndOfSession,
216 ts_event: 1.into(),
217 ts_init: 3.into(),
218 };
219
220 let close2 = InstrumentClose {
221 instrument_id,
222 close_price: Price::from("151.25"),
223 close_type: InstrumentCloseType::ContractExpired,
224 ts_event: 2.into(),
225 ts_init: 4.into(),
226 };
227
228 let data = vec![close1, close2];
229 let record_batch = InstrumentClose::encode_batch(&metadata, &data).unwrap();
230
231 let columns = record_batch.columns();
232 let close_price_values = columns[0]
233 .as_any()
234 .downcast_ref::<FixedSizeBinaryArray>()
235 .unwrap();
236 let close_type_values = columns[1].as_any().downcast_ref::<UInt8Array>().unwrap();
237 let ts_event_values = columns[2].as_any().downcast_ref::<UInt64Array>().unwrap();
238 let ts_init_values = columns[3].as_any().downcast_ref::<UInt64Array>().unwrap();
239
240 assert_eq!(columns.len(), 4);
241 assert_eq!(close_price_values.len(), 2);
242 assert_eq!(
243 get_raw_price(close_price_values.value(0)),
244 (150.50 * FIXED_SCALAR) as PriceRaw
245 );
246 assert_eq!(
247 get_raw_price(close_price_values.value(1)),
248 (151.25 * FIXED_SCALAR) as PriceRaw
249 );
250 assert_eq!(close_type_values.len(), 2);
251 assert_eq!(
252 close_type_values.value(0),
253 InstrumentCloseType::EndOfSession as u8
254 );
255 assert_eq!(
256 close_type_values.value(1),
257 InstrumentCloseType::ContractExpired as u8
258 );
259 assert_eq!(ts_event_values.len(), 2);
260 assert_eq!(ts_event_values.value(0), 1);
261 assert_eq!(ts_event_values.value(1), 2);
262 assert_eq!(ts_init_values.len(), 2);
263 assert_eq!(ts_init_values.value(0), 3);
264 assert_eq!(ts_init_values.value(1), 4);
265 }
266
267 #[rstest]
268 fn test_decode_batch() {
269 let instrument_id = InstrumentId::from("AAPL.XNAS");
270 let metadata = HashMap::from([
271 (KEY_INSTRUMENT_ID.to_string(), instrument_id.to_string()),
272 (KEY_PRICE_PRECISION.to_string(), "2".to_string()),
273 ]);
274
275 let raw_price1 = (150.50 * FIXED_SCALAR) as PriceRaw;
276 let raw_price2 = (151.25 * FIXED_SCALAR) as PriceRaw;
277 let close_price =
278 fixed_size_binary(vec![&raw_price1.to_le_bytes(), &raw_price2.to_le_bytes()]);
279 let close_type = UInt8Array::from(vec![
280 InstrumentCloseType::EndOfSession as u8,
281 InstrumentCloseType::ContractExpired as u8,
282 ]);
283 let ts_event = UInt64Array::from(vec![1, 2]);
284 let ts_init = UInt64Array::from(vec![3, 4]);
285
286 let record_batch = RecordBatch::try_new(
287 InstrumentClose::get_schema(Some(metadata.clone())).into(),
288 vec![
289 Arc::new(close_price),
290 Arc::new(close_type),
291 Arc::new(ts_event),
292 Arc::new(ts_init),
293 ],
294 )
295 .unwrap();
296
297 let decoded_data = InstrumentClose::decode_batch(&metadata, record_batch).unwrap();
298
299 assert_eq!(decoded_data.len(), 2);
300 assert_eq!(decoded_data[0].instrument_id, instrument_id);
301 assert_eq!(decoded_data[0].close_price, Price::from_raw(raw_price1, 2));
302 assert_eq!(
303 decoded_data[0].close_type,
304 InstrumentCloseType::EndOfSession
305 );
306 assert_eq!(decoded_data[0].ts_event.as_u64(), 1);
307 assert_eq!(decoded_data[0].ts_init.as_u64(), 3);
308
309 assert_eq!(decoded_data[1].instrument_id, instrument_id);
310 assert_eq!(decoded_data[1].close_price, Price::from_raw(raw_price2, 2));
311 assert_eq!(
312 decoded_data[1].close_type,
313 InstrumentCloseType::ContractExpired
314 );
315 assert_eq!(decoded_data[1].ts_event.as_u64(), 2);
316 assert_eq!(decoded_data[1].ts_init.as_u64(), 4);
317 }
318
319 #[rstest]
320 fn test_decode_batch_invalid_close_price_returns_error() {
321 let instrument_id = InstrumentId::from("AAPL.XNAS");
322 let metadata = HashMap::from([
323 (KEY_INSTRUMENT_ID.to_string(), instrument_id.to_string()),
324 (KEY_PRICE_PRECISION.to_string(), "2".to_string()),
325 ]);
326
327 let invalid_price: PriceRaw = PriceRaw::MAX - 1000;
328 let close_price = fixed_size_binary(vec![&invalid_price.to_le_bytes()]);
329 let close_type = UInt8Array::from(vec![InstrumentCloseType::EndOfSession as u8]);
330 let ts_event = UInt64Array::from(vec![1]);
331 let ts_init = UInt64Array::from(vec![2]);
332
333 let record_batch = RecordBatch::try_new(
334 InstrumentClose::get_schema(Some(metadata.clone())).into(),
335 vec![
336 Arc::new(close_price),
337 Arc::new(close_type),
338 Arc::new(ts_event),
339 Arc::new(ts_init),
340 ],
341 )
342 .unwrap();
343
344 let result = InstrumentClose::decode_batch(&metadata, record_batch);
345 assert!(result.is_err());
346 let err = result.unwrap_err();
347 assert!(
348 err.to_string().contains("close_price") && err.to_string().contains("row 0"),
349 "Expected close_price error at row 0, was: {err}"
350 );
351 }
352
353 #[rstest]
354 fn test_decode_batch_invalid_close_type_returns_error() {
355 let instrument_id = InstrumentId::from("AAPL.XNAS");
356 let metadata = HashMap::from([
357 (KEY_INSTRUMENT_ID.to_string(), instrument_id.to_string()),
358 (KEY_PRICE_PRECISION.to_string(), "2".to_string()),
359 ]);
360
361 let raw_price = (150.50 * FIXED_SCALAR) as PriceRaw;
362 let close_price = fixed_size_binary(vec![&raw_price.to_le_bytes()]);
363 let close_type = UInt8Array::from(vec![99]);
364 let ts_event = UInt64Array::from(vec![1]);
365 let ts_init = UInt64Array::from(vec![2]);
366
367 let record_batch = RecordBatch::try_new(
368 InstrumentClose::get_schema(Some(metadata.clone())).into(),
369 vec![
370 Arc::new(close_price),
371 Arc::new(close_type),
372 Arc::new(ts_event),
373 Arc::new(ts_init),
374 ],
375 )
376 .unwrap();
377
378 let result = InstrumentClose::decode_batch(&metadata, record_batch);
379 assert!(result.is_err());
380 let err = result.unwrap_err();
381 assert!(
382 err.to_string().contains("InstrumentCloseType"),
383 "Expected InstrumentCloseType error, was: {err}"
384 );
385 }
386
387 #[rstest]
388 fn test_decode_batch_missing_instrument_id_returns_error() {
389 let instrument_id = InstrumentId::from("AAPL.XNAS");
390 let mut metadata = HashMap::from([
391 (KEY_INSTRUMENT_ID.to_string(), instrument_id.to_string()),
392 (KEY_PRICE_PRECISION.to_string(), "2".to_string()),
393 ]);
394
395 let raw_price = (150.50 * FIXED_SCALAR) as PriceRaw;
396 let close_price = fixed_size_binary(vec![&raw_price.to_le_bytes()]);
397 let close_type = UInt8Array::from(vec![InstrumentCloseType::EndOfSession as u8]);
398 let ts_event = UInt64Array::from(vec![1]);
399 let ts_init = UInt64Array::from(vec![2]);
400
401 let record_batch = RecordBatch::try_new(
402 InstrumentClose::get_schema(Some(metadata.clone())).into(),
403 vec![
404 Arc::new(close_price),
405 Arc::new(close_type),
406 Arc::new(ts_event),
407 Arc::new(ts_init),
408 ],
409 )
410 .unwrap();
411
412 metadata.remove(KEY_INSTRUMENT_ID);
413
414 let result = InstrumentClose::decode_batch(&metadata, record_batch);
415 assert!(result.is_err());
416 let err = result.unwrap_err();
417 assert!(
418 err.to_string().contains("instrument_id"),
419 "Expected missing instrument_id error, was: {err}"
420 );
421 }
422
423 #[rstest]
424 fn test_encode_decode_round_trip() {
425 let instrument_id = InstrumentId::from("AAPL.XNAS");
426 let metadata = HashMap::from([
427 (KEY_INSTRUMENT_ID.to_string(), instrument_id.to_string()),
428 (KEY_PRICE_PRECISION.to_string(), "2".to_string()),
429 ]);
430
431 let close1 = InstrumentClose {
432 instrument_id,
433 close_price: Price::from("150.50"),
434 close_type: InstrumentCloseType::EndOfSession,
435 ts_event: 1_000_000_000.into(),
436 ts_init: 1_000_000_001.into(),
437 };
438
439 let close2 = InstrumentClose {
440 instrument_id,
441 close_price: Price::from("151.25"),
442 close_type: InstrumentCloseType::ContractExpired,
443 ts_event: 2_000_000_000.into(),
444 ts_init: 2_000_000_001.into(),
445 };
446
447 let original = vec![close1, close2];
448 let record_batch = InstrumentClose::encode_batch(&metadata, &original).unwrap();
449 let decoded = InstrumentClose::decode_batch(&metadata, record_batch).unwrap();
450
451 assert_eq!(decoded.len(), original.len());
452 for (orig, dec) in original.iter().zip(decoded.iter()) {
453 assert_eq!(dec.instrument_id, orig.instrument_id);
454 assert_eq!(dec.close_price, orig.close_price);
455 assert_eq!(dec.close_type, orig.close_type);
456 assert_eq!(dec.ts_event, orig.ts_event);
457 assert_eq!(dec.ts_init, orig.ts_init);
458 }
459 }
460}