1use std::sync::Arc;
26
27use arrow::{
28 array::{
29 Array, FixedSizeListArray, GenericListArray, GenericListViewArray, LargeListArray,
30 LargeListViewArray, ListArray, ListViewArray, OffsetSizeTrait,
31 },
32 datatypes::{DataType as ArrowDataType, Schema},
33 record_batch::RecordBatch,
34};
35use nautilus_model::data::{
36 ArrowDecoder, ArrowEncoder, CustomData, CustomDataTrait, Data, DataType,
37 decode_custom_from_arrow, ensure_arrow_registered, ensure_custom_data_json_registered,
38 get_arrow_schema, validate_custom_arrow_schema,
39};
40
41use super::{ArrowSchemaProvider, DecodeDataFromRecordBatch, EncodeToRecordBatch};
42
43pub trait CustomDataSerialize: CustomDataTrait {
50 fn schema(&self) -> anyhow::Result<arrow::datatypes::Schema>;
55
56 fn encode_record_batch(
61 &self,
62 items: &[Arc<dyn CustomDataTrait>],
63 ) -> anyhow::Result<RecordBatch>;
64}
65
66pub fn ensure_custom_data_registered<T>()
75where
76 T: CustomDataTrait
77 + ArrowSchemaProvider
78 + EncodeToRecordBatch
79 + DecodeDataFromRecordBatch
80 + Clone
81 + Send
82 + Sync
83 + 'static,
84{
85 let type_name = T::type_name_static();
86
87 if let Some(schema) = get_arrow_schema(type_name) {
89 assert_custom_schema(type_name, &schema);
90 return;
91 }
92
93 let _ = ensure_custom_data_json_registered::<T>();
94
95 let schema = Arc::new(T::get_schema(None));
96 assert_custom_schema(type_name, &schema);
97
98 let encoder: ArrowEncoder = Box::new(|items: &[Arc<dyn CustomDataTrait>]| {
99 let typed: Result<Vec<T>, _> = items
100 .iter()
101 .map(|b| {
102 b.as_any()
103 .downcast_ref::<T>()
104 .cloned()
105 .ok_or_else(|| anyhow::anyhow!("Expected {}", T::type_name_static()))
106 })
107 .collect();
108 let typed = typed?;
109 let metadata = typed
110 .first()
111 .map(EncodeToRecordBatch::metadata)
112 .unwrap_or_default();
113 EncodeToRecordBatch::encode_batch(&metadata, &typed).map_err(|e| anyhow::anyhow!("{e}"))
114 });
115
116 let decoder: ArrowDecoder = Box::new(|metadata, batch| {
117 T::decode_data_batch(metadata, batch).map_err(|e| anyhow::anyhow!("{e}"))
118 });
119
120 let _ = ensure_arrow_registered(type_name, schema, encoder, decoder);
121}
122
123fn assert_custom_schema(type_name: &str, schema: &Schema) {
124 validate_custom_arrow_schema(type_name, schema, true).unwrap();
125}
126
127#[cfg(test)]
128mod tests {
129 use std::sync::Arc;
130
131 use arrow::{
132 array::{Array, ArrayData, FixedSizeListArray, Int64Array, ListArray, ListViewArray},
133 buffer::{Buffer, NullBuffer},
134 datatypes::{DataType, Field, Schema},
135 record_batch::RecordBatch,
136 };
137 use rstest::rstest;
138
139 use super::{
140 assert_custom_schema, validate_required_list_child, validate_required_list_values,
141 };
142
143 #[rstest]
144 fn custom_schema_allows_vec_u8_binary() {
145 let schema = Schema::new(vec![Field::new("payload", DataType::Binary, false)]);
146
147 assert_custom_schema("BinaryPayload", &schema);
148 }
149
150 #[rstest]
151 #[should_panic(
152 expected = "custom write schema `OpaquePayload` contains opaque byte field `item`: FixedSizeBinary(8)"
153 )]
154 fn custom_schema_rejects_nested_opaque_bytes() {
155 let child = Field::new("item", DataType::FixedSizeBinary(8), false);
156 let schema = Schema::new(vec![Field::new(
157 "values",
158 DataType::List(std::sync::Arc::new(child)),
159 false,
160 )]);
161
162 assert_custom_schema("OpaquePayload", &schema);
163 }
164
165 #[rstest]
166 fn required_list_rejects_null_child_value() {
167 let item = Field::new("item", DataType::Int64, false);
168 let values = Int64Array::from(vec![Some(1), None]);
169
170 let error = validate_required_list_child("values", &item, &values, std::iter::once((0, 2)))
171 .unwrap_err();
172
173 assert_eq!(
174 error.to_string(),
175 "Error parsing `custom_data`: field 'values': required list element 1 is null"
176 );
177 }
178
179 #[rstest]
180 fn required_fixed_size_list_allows_null_child_under_null_outer_row() {
181 let child = Arc::new(Field::new("item", DataType::Int64, false));
182 let values = Int64Array::from(vec![Some(1), Some(2), None, None, Some(5), Some(6)]);
183 let nulls = NullBuffer::from(vec![true, false, true]);
184 let list =
185 FixedSizeListArray::try_new(Arc::clone(&child), 2, Arc::new(values), Some(nulls))
186 .unwrap();
187 let schema = Schema::new(vec![Field::new(
188 "values",
189 DataType::FixedSizeList(child, 2),
190 true,
191 )]);
192 let batch = RecordBatch::try_new(Arc::new(schema), vec![Arc::new(list)]).unwrap();
193
194 validate_required_list_values(&batch).unwrap();
195 }
196
197 #[rstest]
198 #[allow(
199 unsafe_code,
200 reason = "storage can declare a non-nullable child while slots hold nulls; arrow's safe constructors reject building such data"
201 )]
202 fn required_list_allows_unreferenced_null_child_after_slice() {
203 let child = Arc::new(Field::new("item", DataType::Int64, false));
204 let values = Int64Array::from(vec![Some(1), Some(2), None]);
205 let data = unsafe {
206 ArrayData::builder(DataType::List(Arc::clone(&child)))
207 .len(3)
208 .add_buffer(Buffer::from_slice_ref([0i32, 1, 2, 3]))
209 .add_child_data(values.into_data())
210 .build_unchecked()
211 };
212 let list = ListArray::from(data).slice(0, 2);
213 let schema = Schema::new(vec![Field::new("values", DataType::List(child), false)]);
214 let batch = RecordBatch::try_new(Arc::new(schema), vec![Arc::new(list)]).unwrap();
215
216 validate_required_list_values(&batch).unwrap();
217 }
218
219 #[rstest]
220 #[allow(
221 unsafe_code,
222 reason = "storage can declare a non-nullable child while slots hold nulls; arrow's safe constructors reject building such data"
223 )]
224 fn required_list_rejects_null_child_referenced_by_valid_row() {
225 let child = Arc::new(Field::new("item", DataType::Int64, false));
226 let values = Int64Array::from(vec![Some(1), None]);
227 let data = unsafe {
228 ArrayData::builder(DataType::List(Arc::clone(&child)))
229 .len(2)
230 .add_buffer(Buffer::from_slice_ref([0i32, 1, 2]))
231 .add_child_data(values.into_data())
232 .build_unchecked()
233 };
234 let list = ListArray::from(data);
235 let schema = Schema::new(vec![Field::new("values", DataType::List(child), false)]);
236 let batch = RecordBatch::try_new(Arc::new(schema), vec![Arc::new(list)]).unwrap();
237
238 let error = validate_required_list_values(&batch).unwrap_err();
239
240 assert_eq!(
241 error.to_string(),
242 "Error parsing `custom_data`: field 'values': required list element 1 is null"
243 );
244 }
245
246 #[rstest]
247 fn required_list_view_rejects_referenced_null_child_value() {
248 let child = Arc::new(Field::new("item", DataType::Int64, false));
249 let values = Int64Array::from(vec![Some(1), None, Some(3)]);
250 let data = ArrayData::builder(DataType::ListView(Arc::clone(&child)))
251 .len(2)
252 .add_buffer(Buffer::from_slice_ref([0i32, 1]))
253 .add_buffer(Buffer::from_slice_ref([1i32, 2]))
254 .add_child_data(values.into_data())
255 .build()
256 .unwrap();
257 let list = ListViewArray::from(data);
258 let schema = Schema::new(vec![Field::new("values", DataType::ListView(child), false)]);
259 let batch = RecordBatch::try_new(Arc::new(schema), vec![Arc::new(list)]).unwrap();
260
261 let error = validate_required_list_values(&batch).unwrap_err();
262
263 assert_eq!(
264 error.to_string(),
265 "Error parsing `custom_data`: field 'values': required list element 1 is null"
266 );
267 }
268}
269
270#[derive(Debug)]
278pub struct CustomDataDecoder;
279
280impl ArrowSchemaProvider for CustomDataDecoder {
281 fn get_schema(
282 metadata: Option<std::collections::HashMap<String, String>>,
283 ) -> arrow::datatypes::Schema {
284 if let Some(metadata) = metadata
285 && let Some(type_name) = metadata.get("type_name")
286 && let Some(schema) = get_arrow_schema(type_name)
287 {
288 let schema = (*schema).clone();
289 let mut fields = schema.fields().iter().cloned().collect::<Vec<_>>();
290 if schema.field_with_name("data_type").is_err() {
291 fields.push(Arc::new(arrow::datatypes::Field::new(
292 "data_type",
293 arrow::datatypes::DataType::Utf8,
294 false,
295 )));
296 }
297 let mut merged_metadata = schema.metadata().clone();
298 merged_metadata.extend(metadata);
299 return arrow::datatypes::Schema::new_with_metadata(fields, merged_metadata);
300 }
301
302 arrow::datatypes::Schema::new(vec![arrow::datatypes::Field::new(
304 "dummy",
305 arrow::datatypes::DataType::Int64,
306 true,
307 )])
308 }
309}
310
311fn strip_data_type_column(
314 batch: &RecordBatch,
315) -> Result<(RecordBatch, Option<DataType>), super::EncodingError> {
316 use super::extract_column_string;
317
318 let Some(data_type_col_idx) = batch
319 .schema()
320 .fields()
321 .iter()
322 .position(|f| f.name() == "data_type")
323 else {
324 return Ok((batch.clone(), None));
325 };
326
327 if batch.num_rows() == 0 {
328 return Ok((batch.clone(), None));
329 }
330
331 let cols = batch.columns();
332 let data_type = if cols[data_type_col_idx].is_null(0) {
333 None
334 } else {
335 let string_col =
336 extract_column_string(cols, "data_type", data_type_col_idx).map_err(|e| {
337 super::EncodingError::ParseError("custom_data", format!("data_type column: {e}"))
338 })?;
339 let first_value = string_col.value(0);
340 Some(
341 DataType::from_persistence_json(first_value)
342 .map_err(|e| super::EncodingError::ParseError("custom_data", e.to_string()))?,
343 )
344 };
345
346 let new_fields: Vec<_> = batch
347 .schema()
348 .fields()
349 .iter()
350 .enumerate()
351 .filter(|(i, _)| *i != data_type_col_idx)
352 .map(|(_, f)| f.clone())
353 .collect();
354 let new_columns: Vec<Arc<dyn arrow::array::Array>> = batch
355 .columns()
356 .iter()
357 .enumerate()
358 .filter(|(i, _)| *i != data_type_col_idx)
359 .map(|(_, c)| Arc::clone(c))
360 .collect();
361 let new_schema =
362 arrow::datatypes::Schema::new_with_metadata(new_fields, batch.schema().metadata().clone());
363 let stripped_batch = RecordBatch::try_new(Arc::new(new_schema), new_columns)
364 .map_err(|e| super::EncodingError::ParseError("custom_data", e.to_string()))?;
365
366 Ok((stripped_batch, data_type))
367}
368
369impl CustomDataDecoder {
370 pub fn decode_custom_batch(
377 metadata: &std::collections::HashMap<String, String>,
378 record_batch: &RecordBatch,
379 ) -> Result<Vec<CustomData>, super::EncodingError> {
380 let type_name = metadata
381 .get("type_name")
382 .cloned()
383 .unwrap_or_else(|| "Unknown".to_string());
384
385 let (batch_to_decode, restored_data_type) = strip_data_type_column(record_batch)?;
386 validate_required_list_values(&batch_to_decode)?;
387
388 if batch_to_decode.num_rows() == 0 {
389 return Ok(Vec::new());
390 }
391
392 let data = match decode_custom_from_arrow(&type_name, metadata, batch_to_decode) {
393 Ok(Some(d)) => d,
394 Ok(None) => {
395 return Err(super::EncodingError::ParseError(
396 "custom_data",
397 format!(
398 "unknown custom data type '{type_name}'; only Rust-registered types are supported"
399 ),
400 ));
401 }
402 Err(e) => {
403 return Err(super::EncodingError::ParseError(
404 "custom_data",
405 format!("decode_custom_from_arrow: {e}"),
406 ));
407 }
408 };
409
410 data.into_iter()
411 .map(|d| match d {
412 Data::Custom(c) => Ok(match &restored_data_type {
413 Some(dt) => CustomData::new(Arc::clone(&c.data), dt.clone()),
414 None => c,
415 }),
416 _ => Err(super::EncodingError::ParseError(
417 "custom_data",
418 format!("registered decoder for '{type_name}' yielded a non-custom row"),
419 )),
420 })
421 .collect()
422 }
423}
424
425fn validate_required_list_values(batch: &RecordBatch) -> Result<(), super::EncodingError> {
426 let schema = batch.schema();
427 for (field, column) in schema.fields().iter().zip(batch.columns()) {
428 match field.data_type() {
429 ArrowDataType::List(child) => {
430 let list = downcast_list::<ListArray>(field.name(), column.as_ref(), "list")?;
431 validate_required_list_child(
432 field.name(),
433 child,
434 list.values().as_ref(),
435 list_ranges(list),
436 )?;
437 }
438 ArrowDataType::LargeList(child) => {
439 let list =
440 downcast_list::<LargeListArray>(field.name(), column.as_ref(), "large-list")?;
441 validate_required_list_child(
442 field.name(),
443 child,
444 list.values().as_ref(),
445 list_ranges(list),
446 )?;
447 }
448 ArrowDataType::ListView(child) => {
449 let list =
450 downcast_list::<ListViewArray>(field.name(), column.as_ref(), "list-view")?;
451 validate_required_list_child(
452 field.name(),
453 child,
454 list.values().as_ref(),
455 list_view_ranges(list),
456 )?;
457 }
458 ArrowDataType::LargeListView(child) => {
459 let list = downcast_list::<LargeListViewArray>(
460 field.name(),
461 column.as_ref(),
462 "large-list-view",
463 )?;
464 validate_required_list_child(
465 field.name(),
466 child,
467 list.values().as_ref(),
468 list_view_ranges(list),
469 )?;
470 }
471 ArrowDataType::FixedSizeList(child, _) => {
472 let list = downcast_list::<FixedSizeListArray>(
473 field.name(),
474 column.as_ref(),
475 "fixed-size-list",
476 )?;
477 validate_required_list_child(
478 field.name(),
479 child,
480 list.values().as_ref(),
481 fixed_size_list_ranges(list),
482 )?;
483 }
484 _ => {}
485 }
486 }
487 Ok(())
488}
489
490fn downcast_list<'a, T: Array + 'static>(
491 field_name: &str,
492 column: &'a dyn Array,
493 kind: &str,
494) -> Result<&'a T, super::EncodingError> {
495 column.as_any().downcast_ref::<T>().ok_or_else(|| {
496 super::EncodingError::ParseError(
497 "custom_data",
498 format!("field '{field_name}' is not a {kind} array"),
499 )
500 })
501}
502
503fn validate_required_list_child(
504 field_name: &str,
505 child: &arrow::datatypes::Field,
506 values: &dyn Array,
507 ranges: impl Iterator<Item = (usize, usize)>,
508) -> Result<(), super::EncodingError> {
509 if child.is_nullable() || values.null_count() == 0 {
510 return Ok(());
511 }
512
513 for (start, end) in ranges {
514 if let Some(index) = (start..end).find(|index| values.is_null(*index)) {
515 return Err(super::EncodingError::ParseError(
516 "custom_data",
517 format!("field '{field_name}': required list element {index} is null"),
518 ));
519 }
520 }
521 Ok(())
522}
523
524fn list_ranges<O: OffsetSizeTrait>(
525 list: &GenericListArray<O>,
526) -> impl Iterator<Item = (usize, usize)> + '_ {
527 let offsets = list.value_offsets();
528 (0..list.len())
529 .filter(|row| list.is_valid(*row))
530 .map(move |row| (offsets[row].as_usize(), offsets[row + 1].as_usize()))
531}
532
533fn list_view_ranges<O: OffsetSizeTrait>(
534 list: &GenericListViewArray<O>,
535) -> impl Iterator<Item = (usize, usize)> + '_ {
536 let offsets = list.value_offsets();
537 let sizes = list.value_sizes();
538 (0..list.len())
539 .filter(|row| list.is_valid(*row))
540 .map(move |row| {
541 let start = offsets[row].as_usize();
542 (start, start + sizes[row].as_usize())
543 })
544}
545
546fn fixed_size_list_ranges(list: &FixedSizeListArray) -> impl Iterator<Item = (usize, usize)> + '_ {
547 let length =
548 usize::try_from(list.value_length()).expect("fixed-size list length is non-negative");
549 (0..list.len())
550 .filter(|row| list.is_valid(*row))
551 .map(move |row| (row * length, (row + 1) * length))
552}
553
554impl DecodeDataFromRecordBatch for CustomDataDecoder {
555 fn decode_data_batch(
556 metadata: &std::collections::HashMap<String, String>,
557 record_batch: RecordBatch,
558 ) -> Result<Vec<Data>, super::EncodingError> {
559 Ok(Self::decode_custom_batch(metadata, &record_batch)?
560 .into_iter()
561 .map(Data::Custom)
562 .collect())
563 }
564}