1use 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#[derive(Clone, Debug)]
24pub struct ArrowLoadResult {
25 pub data: TabularData,
27 pub diagnostics: DiagnosticSet,
29 pub bytes_copied: u64,
31 pub bytes_borrowed: u64,
33}
34
35pub 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
193pub 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}