Skip to main content

antecedent_data/
arrow_adapter.rs

1//! Arrow-backed adapters. Arrow types stay inside this module (ADR 0004).
2//!
3//! SPDX-License-Identifier: MIT OR Apache-2.0
4
5use std::num::NonZeroU32;
6use std::sync::Arc;
7
8use antecedent_core::{
9    CausalSchema, CausalSchemaBuilder, DiagnosticSet, MeasurementSpec, RoleHint, ScalarType,
10    SmallRoleSet, ValueType, VariableId,
11};
12use arrow_array::{Array, FixedSizeListArray, Float64Array, RecordBatch};
13
14use crate::arrow_ffi::{ArrowCColumn, float64_column_from_array};
15use crate::buffer::F64Buffer;
16use crate::column::{FixedVectorColumn, Float64Column, OwnedColumn, ValidityBitmap};
17use crate::dataset::TabularData;
18use crate::error::DataError;
19use crate::materialize::{MaterializationReason, materialization_diagnostic};
20use crate::storage::OwnedColumnarStorage;
21
22/// Result of loading Arrow input into library-owned storage.
23#[derive(Clone, Debug)]
24pub struct ArrowLoadResult {
25    /// Loaded tabular data.
26    pub data: TabularData,
27    /// Copy / materialization diagnostics.
28    pub diagnostics: DiagnosticSet,
29    /// Total bytes copied into owned buffers.
30    pub bytes_copied: u64,
31    /// Total bytes borrowed zero-copy from foreign buffers.
32    pub bytes_borrowed: u64,
33}
34
35/// Load float64 and fixed-size float64-list columns from an Arrow record batch.
36///
37/// Always copies into library-owned buffers (in-process `RecordBatch` path).
38/// Prefer [`tabular_from_arrow_c_columns`] for Arrow C Data Interface zero-copy.
39///
40/// # Errors
41///
42/// Unsupported column types, empty batches, schema construction failures, or more
43/// than `u32::MAX` columns.
44pub fn tabular_from_record_batch(batch: &RecordBatch) -> Result<ArrowLoadResult, DataError> {
45    if batch.num_columns() == 0 {
46        return Err(DataError::InvalidArgument { message: "record batch has no columns".into() });
47    }
48    let mut builder = CausalSchemaBuilder::new();
49    let mut columns = Vec::with_capacity(batch.num_columns());
50    let mut diagnostics = DiagnosticSet::new();
51    let mut bytes_copied = 0u64;
52    let n_rows = batch.num_rows();
53
54    for (i, field) in batch.schema().fields().iter().enumerate() {
55        let name = field.name().clone();
56        let id = VariableId::from_raw(u32::try_from(i).map_err(|_| {
57            DataError::InvalidArgument { message: "too many Arrow columns for VariableId".into() }
58        })?);
59        let array = batch.column(i);
60
61        if let Some(floats) = array.as_any().downcast_ref::<Float64Array>() {
62            builder
63                .add_variable(
64                    Arc::<str>::from(name),
65                    ValueType::Continuous,
66                    SmallRoleSet::from_hint(RoleHint::Context),
67                    None,
68                    None,
69                    MeasurementSpec::default(),
70                )
71                .map_err(|e| DataError::Schema(e.to_string()))?;
72            let (col, copied) = float64_owned_from_array(id, floats, n_rows)?;
73            bytes_copied += copied;
74            diagnostics.push(materialization_diagnostic(
75                MaterializationReason::ForeignBufferIncompatible,
76                copied,
77            ));
78            columns.push(OwnedColumn::Float64(col));
79            continue;
80        }
81
82        if let Some(list) = array.as_any().downcast_ref::<FixedSizeListArray>() {
83            let dim = usize::try_from(list.value_length()).map_err(|_| {
84                DataError::InvalidArgument { message: "FixedSizeList width must fit usize".into() }
85            })?;
86            if dim == 0 {
87                return Err(DataError::InvalidArgument {
88                    message: "FixedSizeList width must be > 0".into(),
89                });
90            }
91            let width = NonZeroU32::new(u32::try_from(dim).map_err(|_| {
92                DataError::InvalidArgument { message: "FixedSizeList width must fit u32".into() }
93            })?)
94            .ok_or(DataError::InvalidArgument {
95                message: "FixedSizeList width must be > 0".into(),
96            })?;
97            builder
98                .add_variable(
99                    Arc::<str>::from(name),
100                    ValueType::Vector { width, element: ScalarType::Float64 },
101                    SmallRoleSet::from_hint(RoleHint::Context),
102                    None,
103                    None,
104                    MeasurementSpec::default(),
105                )
106                .map_err(|e| DataError::Schema(e.to_string()))?;
107            let (col, copied) = fixed_vector_from_list(id, list, n_rows, dim)?;
108            bytes_copied += copied;
109            diagnostics.push(materialization_diagnostic(
110                MaterializationReason::ForeignBufferIncompatible,
111                copied,
112            ));
113            columns.push(OwnedColumn::FixedVector(col));
114            continue;
115        }
116
117        return Err(DataError::TypeMismatch { id, expected: "float64 or FixedSizeList<float64>" });
118    }
119
120    let schema: CausalSchema = builder.build().map_err(|e| DataError::Schema(e.to_string()))?;
121    let storage = OwnedColumnarStorage::try_new(schema, columns, None, None)?;
122    Ok(ArrowLoadResult {
123        data: TabularData::new(storage),
124        diagnostics,
125        bytes_copied,
126        bytes_borrowed: 0,
127    })
128}
129
130fn float64_owned_from_array(
131    id: VariableId,
132    floats: &Float64Array,
133    n_rows: usize,
134) -> Result<(Float64Column, u64), DataError> {
135    let mut values = Vec::with_capacity(n_rows);
136    let mut validity_bytes = vec![0u8; n_rows.div_ceil(8)];
137    for row in 0..n_rows {
138        if floats.is_null(row) {
139            values.push(0.0);
140        } else {
141            values.push(floats.value(row));
142            validity_bytes[row / 8] |= 1 << (row % 8);
143        }
144    }
145    let copied = (values.len() * core::mem::size_of::<f64>() + validity_bytes.len()) as u64;
146    let col = Float64Column::new(
147        id,
148        F64Buffer::owned(Arc::from(values)),
149        ValidityBitmap::from_bytes(validity_bytes, n_rows)?,
150    )?;
151    Ok((col, copied))
152}
153
154fn fixed_vector_from_list(
155    id: VariableId,
156    list: &FixedSizeListArray,
157    n_rows: usize,
158    dim: usize,
159) -> Result<(FixedVectorColumn, u64), DataError> {
160    let values = list.values();
161    let floats = values
162        .as_any()
163        .downcast_ref::<Float64Array>()
164        .ok_or(DataError::TypeMismatch { id, expected: "FixedSizeList<float64>" })?;
165    let mut flat = Vec::with_capacity(n_rows.saturating_mul(dim));
166    let mut validity_bytes = vec![0u8; n_rows.div_ceil(8)];
167    for row in 0..n_rows {
168        if list.is_null(row) {
169            flat.extend(std::iter::repeat_n(0.0, dim));
170            continue;
171        }
172        validity_bytes[row / 8] |= 1 << (row % 8);
173        let start = row.saturating_mul(dim);
174        for k in 0..dim {
175            let idx = start + k;
176            if floats.is_null(idx) {
177                flat.push(0.0);
178            } else {
179                flat.push(floats.value(idx));
180            }
181        }
182    }
183    let copied = (flat.len() * core::mem::size_of::<f64>() + validity_bytes.len()) as u64;
184    let col = FixedVectorColumn::new(
185        id,
186        dim,
187        Arc::from(flat),
188        ValidityBitmap::from_bytes(validity_bytes, n_rows)?,
189    )?;
190    Ok((col, copied))
191}
192
193/// Load float64 columns from Arrow C Data Interface exports, preferring zero-copy.
194///
195/// Consumes each [`ArrowCColumn`]'s FFI structs. Contiguous float64 value buffers
196/// are borrowed; validity bitmaps are copied into library storage.
197///
198/// # Errors
199///
200/// Empty input, non-float64 columns, CDI import failure, schema errors, or more
201/// than `u32::MAX` columns.
202pub fn tabular_from_arrow_c_columns(
203    columns: Vec<ArrowCColumn>,
204) -> Result<ArrowLoadResult, DataError> {
205    if columns.is_empty() {
206        return Err(DataError::InvalidArgument {
207            message: "Arrow CDI import needs ≥1 column".into(),
208        });
209    }
210    let mut builder = CausalSchemaBuilder::new();
211    let mut owned_cols = Vec::with_capacity(columns.len());
212    let mut diagnostics = DiagnosticSet::new();
213    let mut bytes_copied = 0u64;
214    let mut bytes_borrowed = 0u64;
215    let mut n_rows = None;
216
217    for (i, col) in columns.into_iter().enumerate() {
218        let name = col.name.clone();
219        builder
220            .add_variable(
221                Arc::<str>::from(name),
222                ValueType::Continuous,
223                SmallRoleSet::from_hint(RoleHint::Context),
224                None,
225                None,
226                MeasurementSpec::default(),
227            )
228            .map_err(|e| DataError::Schema(e.to_string()))?;
229
230        let array = col.into_array()?;
231        if let Some(n) = n_rows {
232            if array.len() != n {
233                return Err(DataError::LengthMismatch {
234                    expected: n,
235                    actual: array.len(),
236                    context: "Arrow CDI column lengths",
237                });
238            }
239        } else {
240            n_rows = Some(array.len());
241        }
242
243        let id = VariableId::from_raw(u32::try_from(i).map_err(|_| {
244            DataError::InvalidArgument { message: "too many Arrow columns for VariableId".into() }
245        })?);
246        let (owned, borrowed, copied, diag) = float64_column_from_array(id, array)?;
247        bytes_borrowed += borrowed;
248        bytes_copied += copied;
249        diagnostics.push(diag);
250        owned_cols.push(owned);
251    }
252
253    let schema: CausalSchema = builder.build().map_err(|e| DataError::Schema(e.to_string()))?;
254    let storage = OwnedColumnarStorage::try_new(schema, owned_cols, None, None)?;
255    Ok(ArrowLoadResult {
256        data: TabularData::new(storage),
257        diagnostics,
258        bytes_copied,
259        bytes_borrowed,
260    })
261}
262
263#[cfg(test)]
264mod tests {
265    use antecedent_core::VariableId;
266    use arrow_array::ffi::to_ffi;
267    use arrow_array::{Array, Float64Array};
268    use arrow_schema::{DataType, Field, Schema};
269
270    use super::*;
271    use crate::arrow_ffi::ArrowCColumn;
272    use crate::table::TableView;
273
274    #[test]
275    fn arrow_load_copies_and_exposes_table_view() {
276        let schema = Schema::new(vec![
277            Field::new("x", DataType::Float64, true),
278            Field::new("y", DataType::Float64, true),
279        ]);
280        let x = Float64Array::from(vec![Some(1.0), None, Some(3.0)]);
281        let y = Float64Array::from(vec![Some(4.0), Some(5.0), Some(6.0)]);
282        let batch = RecordBatch::try_new(Arc::new(schema), vec![Arc::new(x), Arc::new(y)]).unwrap();
283
284        let loaded = tabular_from_record_batch(&batch).unwrap();
285        assert!(loaded.bytes_copied > 0);
286        assert_eq!(loaded.bytes_borrowed, 0);
287        assert!(!loaded.diagnostics.is_empty());
288        assert_eq!(loaded.data.row_count(), 3);
289        let col = loaded.data.column(VariableId::from_raw(0)).unwrap();
290        match col {
291            crate::column::ColumnView::Float64(c) => {
292                assert!(c.validity.is_valid(0));
293                assert!(!c.validity.is_valid(1));
294                assert!((c.values[2] - 3.0).abs() < f64::EPSILON);
295                assert!(!c.values.is_foreign());
296            }
297            _ => panic!("expected float"),
298        }
299    }
300
301    #[test]
302    fn arrow_load_fixed_size_list_float64() {
303        use arrow_array::FixedSizeListArray;
304        use arrow_buffer::NullBuffer;
305
306        let values = Float64Array::from(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
307        let list = FixedSizeListArray::new(
308            Arc::new(Field::new("item", DataType::Float64, true)),
309            2,
310            Arc::new(values),
311            Some(NullBuffer::from(vec![true, false, true])),
312        );
313        let schema = Schema::new(vec![Field::new(
314            "v",
315            DataType::FixedSizeList(Arc::new(Field::new("item", DataType::Float64, true)), 2),
316            true,
317        )]);
318        let batch = RecordBatch::try_new(Arc::new(schema), vec![Arc::new(list)]).unwrap();
319        let loaded = tabular_from_record_batch(&batch).unwrap();
320        assert_eq!(loaded.data.row_count(), 3);
321        match loaded.data.column(VariableId::from_raw(0)).unwrap() {
322            crate::column::ColumnView::FixedVector(c) => {
323                assert_eq!(c.dim, 2);
324                assert!(c.validity.is_valid(0));
325                assert!(!c.validity.is_valid(1));
326                assert!(c.validity.is_valid(2));
327                assert!((c.values[0] - 1.0).abs() < f64::EPSILON);
328                assert!((c.values[1] - 2.0).abs() < f64::EPSILON);
329                assert!((c.values[4] - 5.0).abs() < f64::EPSILON);
330            }
331            _ => panic!("expected FixedVector"),
332        }
333    }
334
335    #[test]
336    fn arrow_cdi_zero_copy_borrows_values() {
337        let x = Float64Array::from(vec![1.0, 2.0, 3.0]);
338        let y = Float64Array::from(vec![4.0, 5.0, 6.0]);
339        let x_data = x.to_data();
340        let y_data = y.to_data();
341        let (x_arr, x_sch) = to_ffi(&x_data).unwrap();
342        let (y_arr, y_sch) = to_ffi(&y_data).unwrap();
343        let loaded = tabular_from_arrow_c_columns(vec![
344            ArrowCColumn { name: "x".into(), array: x_arr, schema: x_sch },
345            ArrowCColumn { name: "y".into(), array: y_arr, schema: y_sch },
346        ])
347        .unwrap();
348        assert!(loaded.bytes_borrowed > 0);
349        assert_eq!(loaded.data.row_count(), 3);
350        let col = loaded.data.column(VariableId::from_raw(0)).unwrap();
351        match col {
352            crate::column::ColumnView::Float64(c) => {
353                assert!(c.values.is_foreign());
354                assert!((c.values[0] - 1.0).abs() < f64::EPSILON);
355                assert!((c.values[2] - 3.0).abs() < f64::EPSILON);
356            }
357            _ => panic!("expected float"),
358        }
359    }
360}