1use 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::IndexPriceUpdate, 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 IndexPriceUpdate {
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 IndexPriceUpdate {
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 IndexPriceUpdate {
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 IndexPriceUpdate {
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 = IndexPriceUpdate::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 = IndexPriceUpdate::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 = IndexPriceUpdate {
188 instrument_id,
189 value: Price::from("50000.00"),
190 ts_event: 1.into(),
191 ts_init: 3.into(),
192 };
193
194 let update2 = IndexPriceUpdate {
195 instrument_id,
196 value: Price::from("51000.00"),
197 ts_event: 2.into(),
198 ts_init: 4.into(),
199 };
200
201 let data = vec![update1, update2];
202 let record_batch = IndexPriceUpdate::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!(50000.00).to_string()).raw()
217 );
218 assert_eq!(
219 get_raw_price(value_values.value(1)),
220 Price::from(dec!(51000.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.00 * FIXED_SCALAR) as PriceRaw;
239 let raw_price2 = (51.00 * 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 IndexPriceUpdate::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 = IndexPriceUpdate::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 IndexPriceUpdate::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 = IndexPriceUpdate::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.00 * 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 IndexPriceUpdate::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 = IndexPriceUpdate::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 = IndexPriceUpdate {
333 instrument_id,
334 value: Price::from("50000.00"),
335 ts_event: 1_000_000_000.into(),
336 ts_init: 1_000_000_001.into(),
337 };
338
339 let update2 = IndexPriceUpdate {
340 instrument_id,
341 value: Price::from("51000.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 = IndexPriceUpdate::encode_batch(&metadata, &original).unwrap();
348 let decoded = IndexPriceUpdate::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}