1use 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}