Skip to main content

nautilus_serialization/arrow/
close.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, 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}