Skip to main content

nautilus_serialization/arrow/
custom.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 code 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
16//! Custom data: registration and dynamic decoding.
17//!
18//! - **Registration:** Call [`ensure_custom_data_registered::<T>()`] once (e.g. before using the
19//!   catalog) for each custom data type `T` using the `#[arrow_custom_data]` macro. When Python
20//!   support is enabled, also call `nautilus_model::data::register_rust_extractor::<T>()`.
21//! - **Decoder:** [`CustomDataDecoder`] provides [`ArrowSchemaProvider`] and
22//!   [`DecodeDataFromRecordBatch`] for Parquet-backed custom data decoded at runtime by type name.
23//!   Types must be registered via [`ensure_custom_data_registered::<T>()`] before use.
24
25use 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
43/// Trait for custom data types that support Arrow schema and record batch encoding.
44/// Used as a type bound by the `#[arrow_custom_data]` macro; catalog encoding goes through
45/// the registry, not this trait directly.
46///
47/// Implemented by the `#[arrow_custom_data]` macro for Rust custom data types. Python custom
48/// types use the registry encoder registered by `register_custom_data_class` instead.
49pub trait CustomDataSerialize: CustomDataTrait {
50    /// Returns the Arrow schema for this custom data type.
51    ///
52    /// # Errors
53    /// Returns an error if schema construction fails.
54    fn schema(&self) -> anyhow::Result<arrow::datatypes::Schema>;
55
56    /// Encodes a batch of custom data items to an Arrow RecordBatch.
57    ///
58    /// # Errors
59    /// Returns an error if encoding fails (e.g. type mismatch or Arrow error).
60    fn encode_record_batch(
61        &self,
62        items: &[Arc<dyn CustomDataTrait>],
63    ) -> anyhow::Result<RecordBatch>;
64}
65
66/// Registers a custom data type in the JSON and Arrow registries. Call once per type
67/// (e.g. at catalog decode or before querying custom data).
68///
69/// Each distinct type `T` is registered at most once (per process). Safe to call
70/// multiple times for the same `T`.
71///
72/// When Python support is enabled, also call
73/// `nautilus_model::data::register_rust_extractor::<T>()` for types exposed to Python.
74pub 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    // Skip if already registered
88    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/// Decoder for custom data types that are identified at runtime by metadata (e.g. `type_name`).
271///
272/// Only Rust-registered custom types (e.g. `RustTestCustomData`, `MacroYieldCurveData`) can be
273/// decoded. Unknown types return an error.
274///
275/// **Important:** The caller must ensure that any Rust custom data types are registered
276/// via [`ensure_custom_data_registered::<T>()`] before use.
277#[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        // Unknown type - return minimal schema (caller should not use this for decode)
303        arrow::datatypes::Schema::new(vec![arrow::datatypes::Field::new(
304            "dummy",
305            arrow::datatypes::DataType::Int64,
306            true,
307        )])
308    }
309}
310
311/// Strips the data_type column from a record batch and returns the parsed DataType.
312/// Returns (batch, None) if there is no data_type column.
313fn 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    /// Decodes a `RecordBatch` into typed [`CustomData`] values.
371    ///
372    /// # Errors
373    ///
374    /// Returns an `EncodingError` if the type is unregistered, decoding fails, or a registered
375    /// decoder yields a non-custom row.
376    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}