Skip to main content

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