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, sync::Arc};
17
18use arrow::{
19    array::{FixedSizeBinaryArray, FixedSizeBinaryBuilder, UInt64Array},
20    datatypes::{DataType, Field, Schema},
21    error::ArrowError,
22    record_batch::RecordBatch,
23};
24use nautilus_model::{data::prices::MarkPriceUpdate, types::fixed::PRECISION_BYTES};
25
26use super::{
27    DecodeDataFromRecordBatch, EncodingError, KEY_INSTRUMENT_ID, KEY_PRICE_PRECISION, decode_price,
28    extract_column, parse_price_metadata, validate_precision_bytes,
29};
30use crate::arrow::{ArrowSchemaProvider, Data, DecodeFromRecordBatch, EncodeToRecordBatch};
31
32impl ArrowSchemaProvider for MarkPriceUpdate {
33    fn get_schema(metadata: Option<HashMap<String, String>>) -> Schema {
34        let fields = vec![
35            Field::new("value", DataType::FixedSizeBinary(PRECISION_BYTES), false),
36            Field::new("ts_event", DataType::UInt64, false),
37            Field::new("ts_init", DataType::UInt64, false),
38        ];
39
40        match metadata {
41            Some(metadata) => Schema::new_with_metadata(fields, metadata),
42            None => Schema::new(fields),
43        }
44    }
45}
46
47impl EncodeToRecordBatch for MarkPriceUpdate {
48    fn encode_batch(
49        metadata: &HashMap<String, String>,
50        data: &[Self],
51    ) -> Result<RecordBatch, ArrowError> {
52        let mut value_builder = FixedSizeBinaryBuilder::with_capacity(data.len(), PRECISION_BYTES);
53        let mut ts_event_builder = UInt64Array::builder(data.len());
54        let mut ts_init_builder = UInt64Array::builder(data.len());
55
56        for update in data {
57            value_builder
58                .append_value(update.value.raw().to_le_bytes())
59                .unwrap();
60            ts_event_builder.append_value(update.ts_event.as_u64());
61            ts_init_builder.append_value(update.ts_init.as_u64());
62        }
63
64        RecordBatch::try_new(
65            Self::get_schema(Some(metadata.clone())).into(),
66            vec![
67                Arc::new(value_builder.finish()),
68                Arc::new(ts_event_builder.finish()),
69                Arc::new(ts_init_builder.finish()),
70            ],
71        )
72    }
73
74    fn metadata(&self) -> HashMap<String, String> {
75        let mut metadata = HashMap::new();
76        metadata.insert(
77            KEY_INSTRUMENT_ID.to_string(),
78            self.instrument_id.to_string(),
79        );
80        metadata.insert(
81            KEY_PRICE_PRECISION.to_string(),
82            self.value.precision.to_string(),
83        );
84        metadata
85    }
86}
87
88impl DecodeFromRecordBatch for MarkPriceUpdate {
89    fn decode_batch(
90        metadata: &HashMap<String, String>,
91        record_batch: RecordBatch,
92    ) -> Result<Vec<Self>, EncodingError> {
93        let (instrument_id, price_precision) = parse_price_metadata(metadata)?;
94        let cols = record_batch.columns();
95
96        let value_values = extract_column::<FixedSizeBinaryArray>(
97            cols,
98            "value",
99            0,
100            DataType::FixedSizeBinary(PRECISION_BYTES),
101        )?;
102        let ts_event_values = extract_column::<UInt64Array>(cols, "ts_event", 1, DataType::UInt64)?;
103        let ts_init_values = extract_column::<UInt64Array>(cols, "ts_init", 2, DataType::UInt64)?;
104
105        validate_precision_bytes(value_values, "value")?;
106
107        let result: Result<Vec<Self>, EncodingError> = (0..record_batch.num_rows())
108            .map(|row| {
109                let value = decode_price(value_values.value(row), price_precision, "value", row)?;
110                Ok(Self {
111                    instrument_id,
112                    value,
113                    ts_event: ts_event_values.value(row).into(),
114                    ts_init: ts_init_values.value(row).into(),
115                })
116            })
117            .collect();
118
119        result
120    }
121}
122
123impl DecodeDataFromRecordBatch for MarkPriceUpdate {
124    fn decode_data_batch(
125        metadata: &HashMap<String, String>,
126        record_batch: RecordBatch,
127    ) -> Result<Vec<Data>, EncodingError> {
128        let updates: Vec<Self> = Self::decode_batch(metadata, record_batch)?;
129        Ok(updates.into_iter().map(Data::from).collect())
130    }
131}
132
133#[cfg(test)]
134mod tests {
135    use std::sync::Arc;
136
137    use arrow::{array::Array, record_batch::RecordBatch};
138    use nautilus_model::{
139        identifiers::InstrumentId,
140        types::{Price, fixed::FIXED_SCALAR, price::PriceRaw},
141    };
142    use rstest::rstest;
143    use rust_decimal_macros::dec;
144
145    use super::*;
146    use crate::arrow::{fixed_size_binary, get_raw_price};
147
148    #[rstest]
149    fn test_get_schema() {
150        let instrument_id = InstrumentId::from("BTC-USDT.BINANCE");
151        let metadata = HashMap::from([
152            (KEY_INSTRUMENT_ID.to_string(), instrument_id.to_string()),
153            (KEY_PRICE_PRECISION.to_string(), "2".to_string()),
154        ]);
155        let schema = MarkPriceUpdate::get_schema(Some(metadata.clone()));
156
157        let expected_fields = vec![
158            Field::new("value", DataType::FixedSizeBinary(PRECISION_BYTES), false),
159            Field::new("ts_event", DataType::UInt64, false),
160            Field::new("ts_init", DataType::UInt64, false),
161        ];
162
163        let expected_schema = Schema::new_with_metadata(expected_fields, metadata);
164        assert_eq!(schema, expected_schema);
165    }
166
167    #[rstest]
168    fn test_get_schema_map() {
169        let schema_map = MarkPriceUpdate::get_schema_map();
170        let mut expected_map = HashMap::new();
171
172        let fixed_size_binary = format!("FixedSizeBinary({PRECISION_BYTES})");
173        expected_map.insert("value".to_string(), fixed_size_binary);
174        expected_map.insert("ts_event".to_string(), "UInt64".to_string());
175        expected_map.insert("ts_init".to_string(), "UInt64".to_string());
176        assert_eq!(schema_map, expected_map);
177    }
178
179    #[rstest]
180    fn test_encode_batch() {
181        let instrument_id = InstrumentId::from("BTC-USDT.BINANCE");
182        let metadata = HashMap::from([
183            (KEY_INSTRUMENT_ID.to_string(), instrument_id.to_string()),
184            (KEY_PRICE_PRECISION.to_string(), "2".to_string()),
185        ]);
186
187        let update1 = MarkPriceUpdate {
188            instrument_id,
189            value: Price::from("50200.00"),
190            ts_event: 1.into(),
191            ts_init: 3.into(),
192        };
193
194        let update2 = MarkPriceUpdate {
195            instrument_id,
196            value: Price::from("50300.00"),
197            ts_event: 2.into(),
198            ts_init: 4.into(),
199        };
200
201        let data = vec![update1, update2];
202        let record_batch = MarkPriceUpdate::encode_batch(&metadata, &data).unwrap();
203
204        let columns = record_batch.columns();
205        let value_values = columns[0]
206            .as_any()
207            .downcast_ref::<FixedSizeBinaryArray>()
208            .unwrap();
209        let ts_event_values = columns[1].as_any().downcast_ref::<UInt64Array>().unwrap();
210        let ts_init_values = columns[2].as_any().downcast_ref::<UInt64Array>().unwrap();
211
212        assert_eq!(columns.len(), 3);
213        assert_eq!(value_values.len(), 2);
214        assert_eq!(
215            get_raw_price(value_values.value(0)),
216            Price::from(dec!(50200.00).to_string()).raw()
217        );
218        assert_eq!(
219            get_raw_price(value_values.value(1)),
220            Price::from(dec!(50300.00).to_string()).raw()
221        );
222        assert_eq!(ts_event_values.len(), 2);
223        assert_eq!(ts_event_values.value(0), 1);
224        assert_eq!(ts_event_values.value(1), 2);
225        assert_eq!(ts_init_values.len(), 2);
226        assert_eq!(ts_init_values.value(0), 3);
227        assert_eq!(ts_init_values.value(1), 4);
228    }
229
230    #[rstest]
231    fn test_decode_batch() {
232        let instrument_id = InstrumentId::from("BTC-USDT.BINANCE");
233        let metadata = HashMap::from([
234            (KEY_INSTRUMENT_ID.to_string(), instrument_id.to_string()),
235            (KEY_PRICE_PRECISION.to_string(), "2".to_string()),
236        ]);
237
238        let raw_price1 = (50.20 * FIXED_SCALAR) as PriceRaw;
239        let raw_price2 = (50.30 * FIXED_SCALAR) as PriceRaw;
240        let value = fixed_size_binary(vec![&raw_price1.to_le_bytes(), &raw_price2.to_le_bytes()]);
241        let ts_event = UInt64Array::from(vec![1, 2]);
242        let ts_init = UInt64Array::from(vec![3, 4]);
243
244        let record_batch = RecordBatch::try_new(
245            MarkPriceUpdate::get_schema(Some(metadata.clone())).into(),
246            vec![Arc::new(value), Arc::new(ts_event), Arc::new(ts_init)],
247        )
248        .unwrap();
249
250        let decoded_data = MarkPriceUpdate::decode_batch(&metadata, record_batch).unwrap();
251
252        assert_eq!(decoded_data.len(), 2);
253        assert_eq!(decoded_data[0].instrument_id, instrument_id);
254        assert_eq!(decoded_data[0].value, Price::from_raw(raw_price1, 2));
255        assert_eq!(decoded_data[0].ts_event.as_u64(), 1);
256        assert_eq!(decoded_data[0].ts_init.as_u64(), 3);
257
258        assert_eq!(decoded_data[1].instrument_id, instrument_id);
259        assert_eq!(decoded_data[1].value, Price::from_raw(raw_price2, 2));
260        assert_eq!(decoded_data[1].ts_event.as_u64(), 2);
261        assert_eq!(decoded_data[1].ts_init.as_u64(), 4);
262    }
263
264    #[rstest]
265    fn test_decode_batch_invalid_value_returns_error() {
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 invalid_price: PriceRaw = PriceRaw::MAX - 1000;
273        let value = fixed_size_binary(vec![&invalid_price.to_le_bytes()]);
274        let ts_event = UInt64Array::from(vec![1]);
275        let ts_init = UInt64Array::from(vec![2]);
276
277        let record_batch = RecordBatch::try_new(
278            MarkPriceUpdate::get_schema(Some(metadata.clone())).into(),
279            vec![Arc::new(value), Arc::new(ts_event), Arc::new(ts_init)],
280        )
281        .unwrap();
282
283        let result = MarkPriceUpdate::decode_batch(&metadata, record_batch);
284        assert!(result.is_err());
285        let err = result.unwrap_err();
286        assert!(
287            err.to_string().contains("value") && err.to_string().contains("row 0"),
288            "Expected value error at row 0, was: {err}"
289        );
290    }
291
292    #[rstest]
293    fn test_decode_batch_missing_instrument_id_returns_error() {
294        let mut metadata = HashMap::from([
295            (
296                KEY_INSTRUMENT_ID.to_string(),
297                "BTC-USDT.BINANCE".to_string(),
298            ),
299            (KEY_PRICE_PRECISION.to_string(), "2".to_string()),
300        ]);
301
302        let raw_price = (50.20 * FIXED_SCALAR) as PriceRaw;
303        let value = fixed_size_binary(vec![&raw_price.to_le_bytes()]);
304        let ts_event = UInt64Array::from(vec![1]);
305        let ts_init = UInt64Array::from(vec![2]);
306
307        let record_batch = RecordBatch::try_new(
308            MarkPriceUpdate::get_schema(Some(metadata.clone())).into(),
309            vec![Arc::new(value), Arc::new(ts_event), Arc::new(ts_init)],
310        )
311        .unwrap();
312
313        metadata.remove(KEY_INSTRUMENT_ID);
314
315        let result = MarkPriceUpdate::decode_batch(&metadata, record_batch);
316        assert!(result.is_err());
317        let err = result.unwrap_err();
318        assert!(
319            err.to_string().contains("instrument_id"),
320            "Expected missing instrument_id error, was: {err}"
321        );
322    }
323
324    #[rstest]
325    fn test_encode_decode_round_trip() {
326        let instrument_id = InstrumentId::from("BTC-USDT.BINANCE");
327        let metadata = HashMap::from([
328            (KEY_INSTRUMENT_ID.to_string(), instrument_id.to_string()),
329            (KEY_PRICE_PRECISION.to_string(), "2".to_string()),
330        ]);
331
332        let update1 = MarkPriceUpdate {
333            instrument_id,
334            value: Price::from("50200.00"),
335            ts_event: 1_000_000_000.into(),
336            ts_init: 1_000_000_001.into(),
337        };
338
339        let update2 = MarkPriceUpdate {
340            instrument_id,
341            value: Price::from("50300.00"),
342            ts_event: 2_000_000_000.into(),
343            ts_init: 2_000_000_001.into(),
344        };
345
346        let original = vec![update1, update2];
347        let record_batch = MarkPriceUpdate::encode_batch(&metadata, &original).unwrap();
348        let decoded = MarkPriceUpdate::decode_batch(&metadata, record_batch).unwrap();
349
350        assert_eq!(decoded.len(), original.len());
351        for (orig, dec) in original.iter().zip(decoded.iter()) {
352            assert_eq!(dec.instrument_id, orig.instrument_id);
353            assert_eq!(dec.value, orig.value);
354            assert_eq!(dec.ts_event, orig.ts_event);
355            assert_eq!(dec.ts_init, orig.ts_init);
356        }
357    }
358}