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, 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}