Skip to main content

dag_ml_data_arrow/
lib.rs

1//! Apache Arrow IPC feature buffer reader for dag-ml-data.
2//!
3//! This crate is a workspace member so the workspace `cargo build` /
4//! `cargo test` commands always compile it, but the C ABI wiring in
5//! `dag-ml-data-capi` is gated behind the `arrow-ipc` feature, so the
6//! Arrow dependency is **opt-in at the cdylib boundary** for hosts that
7//! ship `libdag_ml_data_capi.so`. Workspace consumers running the gate
8//! still pay the compile-time cost.
9//!
10//! The reader maps each Arrow `RecordBatch` whose schema metadata carries
11//! `dag_ml_data.feature_set_id` to a `NumericFeatureMatrixF64Columnar`,
12//! then resolves the whole stream to a `NumericFeatureBufferStore`. The
13//! mapping is total:
14//!
15//! - `Float64Array` column → `NumericFeatureMatrixF64Columnar.columns[i]`
16//! - Arrow validity bitmap → per-column `validity_masks[i]`
17//! - `observation_id` `Utf8Array` column → `observation_ids`
18//! - feature column names → `feature_names`
19//!
20//! Mandatory schema metadata keys:
21//! - `dag_ml_data.feature_set_id` — the feature-set id of the buffer
22//! - `dag_ml_data.representation_id` — the representation id of the buffer
23
24use std::fs::File;
25use std::io::Read;
26use std::path::Path;
27
28use arrow_array::{Array, Float64Array, RecordBatch, StringArray};
29use arrow_ipc::reader::{FileReader, StreamReader};
30
31use dag_ml_data_core::buffer::{NumericFeatureBufferStore, NumericFeatureMatrixF64Columnar};
32use dag_ml_data_core::error::{DataError, Result};
33use dag_ml_data_core::ids::{ObservationId, RepresentationId};
34
35/// Arrow schema-metadata key carrying the source feature-set id.
36pub const SCHEMA_METADATA_FEATURE_SET_ID: &str = "dag_ml_data.feature_set_id";
37/// Arrow schema-metadata key carrying the representation id.
38pub const SCHEMA_METADATA_REPRESENTATION_ID: &str = "dag_ml_data.representation_id";
39/// Column name holding per-row observation ids in each record batch.
40pub const OBSERVATION_ID_COLUMN: &str = "observation_id";
41
42/// Parse an Arrow IPC stream from any `Read` source into a
43/// `NumericFeatureBufferStore`. Each top-level record batch becomes one
44/// buffer; the batches' schemas must declare both
45/// `dag_ml_data.feature_set_id` and `dag_ml_data.representation_id` in
46/// their `metadata` map and expose an `observation_id` UTF-8 column.
47pub fn read_buffers_from_ipc_stream<R: Read>(reader: R) -> Result<NumericFeatureBufferStore> {
48    let stream = StreamReader::try_new(reader, None).map_err(arrow_error)?;
49    read_batches(stream)
50}
51
52/// Parse an Arrow IPC file (with footer + magic) from any `Read + Seek`
53/// source into a `NumericFeatureBufferStore`. Same per-batch contract as
54/// `read_buffers_from_ipc_stream`.
55pub fn read_buffers_from_ipc_file<R: Read + std::io::Seek>(
56    reader: R,
57) -> Result<NumericFeatureBufferStore> {
58    let file = FileReader::try_new(reader, None).map_err(arrow_error)?;
59    read_batches(file)
60}
61
62fn read_batches(
63    batches: impl Iterator<Item = std::result::Result<RecordBatch, arrow_schema::ArrowError>>,
64) -> Result<NumericFeatureBufferStore> {
65    let mut matrices = std::collections::BTreeMap::<String, NumericFeatureMatrixF64Columnar>::new();
66    for batch in batches {
67        let next = record_batch_to_matrix(&batch.map_err(arrow_error)?)?;
68        if next.observation_ids.is_empty() {
69            continue;
70        }
71        let Some(matrix) = matrices.get_mut(&next.feature_set_id) else {
72            matrices.insert(next.feature_set_id.clone(), next);
73            continue;
74        };
75        if matrix.representation_id != next.representation_id
76            || matrix.feature_names != next.feature_names
77        {
78            return Err(DataError::Validation(
79                "arrow IPC batches use incompatible feature schemas".into(),
80            ));
81        }
82        let prior_rows = matrix.observation_ids.len();
83        let next_rows = next.observation_ids.len();
84        if matrix.validity_masks.is_none() && next.validity_masks.is_some() {
85            matrix.validity_masks = Some(vec![vec![true; prior_rows]; matrix.columns.len()]);
86        }
87        if let Some(masks) = &mut matrix.validity_masks {
88            for (idx, mask) in masks.iter_mut().enumerate() {
89                if let Some(next_masks) = &next.validity_masks {
90                    mask.extend_from_slice(&next_masks[idx]);
91                } else {
92                    mask.extend(std::iter::repeat_n(true, next_rows));
93                }
94            }
95        }
96        matrix.observation_ids.extend(next.observation_ids);
97        for (column, next_column) in matrix.columns.iter_mut().zip(next.columns) {
98            column.extend(next_column);
99        }
100    }
101    NumericFeatureBufferStore::from_f64_column_matrices(matrices.into_values().collect())
102}
103
104/// Convenience: read an Arrow IPC file from disk.
105pub fn read_buffers_from_ipc_path(path: &Path) -> Result<NumericFeatureBufferStore> {
106    let file = File::open(path).map_err(|error| {
107        DataError::Validation(format!(
108            "failed to open arrow IPC file at `{}`: {error}",
109            path.display()
110        ))
111    })?;
112    read_buffers_from_ipc_file(file)
113}
114
115fn record_batch_to_matrix(batch: &RecordBatch) -> Result<NumericFeatureMatrixF64Columnar> {
116    let schema = batch.schema();
117    let metadata = schema.metadata();
118    let feature_set_id = metadata
119        .get(SCHEMA_METADATA_FEATURE_SET_ID)
120        .ok_or_else(|| {
121            DataError::Validation(format!(
122                "arrow IPC batch is missing required schema metadata `{SCHEMA_METADATA_FEATURE_SET_ID}`"
123            ))
124        })?
125        .clone();
126    let representation_raw = metadata.get(SCHEMA_METADATA_REPRESENTATION_ID).ok_or_else(|| {
127        DataError::Validation(format!(
128            "arrow IPC batch `{feature_set_id}` is missing required schema metadata `{SCHEMA_METADATA_REPRESENTATION_ID}`"
129        ))
130    })?;
131    let representation_id = RepresentationId::new(representation_raw)?;
132
133    let observation_field_index = schema.index_of(OBSERVATION_ID_COLUMN).map_err(|_| {
134        DataError::Validation(format!(
135            "arrow IPC batch `{feature_set_id}` is missing the required `{OBSERVATION_ID_COLUMN}` column"
136        ))
137    })?;
138    let observation_array = batch
139        .column(observation_field_index)
140        .as_any()
141        .downcast_ref::<StringArray>()
142        .ok_or_else(|| {
143            DataError::Validation(format!(
144                "arrow IPC batch `{feature_set_id}` `{OBSERVATION_ID_COLUMN}` column must be Utf8"
145            ))
146        })?;
147    if observation_array.null_count() != 0 {
148        return Err(DataError::Validation(format!(
149            "arrow IPC batch `{feature_set_id}` `{OBSERVATION_ID_COLUMN}` column contains nulls"
150        )));
151    }
152    let mut observation_ids = Vec::with_capacity(observation_array.len());
153    for raw in observation_array.iter() {
154        let value = raw.ok_or_else(|| {
155            DataError::Validation(format!(
156                "arrow IPC batch `{feature_set_id}` `{OBSERVATION_ID_COLUMN}` column contains nulls"
157            ))
158        })?;
159        observation_ids.push(ObservationId::new(value)?);
160    }
161
162    let mut feature_names = Vec::new();
163    let mut columns = Vec::new();
164    let mut masks = Vec::new();
165    let mut any_null = false;
166    for (idx, field) in schema.fields().iter().enumerate() {
167        if idx == observation_field_index {
168            continue;
169        }
170        let column = batch.column(idx);
171        let float_column = column
172            .as_any()
173            .downcast_ref::<Float64Array>()
174            .ok_or_else(|| {
175                DataError::Validation(format!(
176                    "arrow IPC batch `{feature_set_id}` column `{}` must be Float64",
177                    field.name()
178                ))
179            })?;
180        let row_count = float_column.len();
181        let mut values = Vec::with_capacity(row_count);
182        let mut mask = Vec::with_capacity(row_count);
183        for row in 0..row_count {
184            if float_column.is_null(row) {
185                values.push(0.0);
186                mask.push(false);
187                any_null = true;
188            } else {
189                values.push(float_column.value(row));
190                mask.push(true);
191            }
192        }
193        feature_names.push(field.name().clone());
194        columns.push(values);
195        masks.push(mask);
196    }
197    if feature_names.is_empty() {
198        return Err(DataError::Validation(format!(
199            "arrow IPC batch `{feature_set_id}` has no feature columns besides `{OBSERVATION_ID_COLUMN}`"
200        )));
201    }
202    Ok(NumericFeatureMatrixF64Columnar {
203        feature_set_id,
204        representation_id,
205        feature_names,
206        observation_ids,
207        columns,
208        validity_masks: any_null.then_some(masks),
209    })
210}
211
212fn arrow_error(error: impl std::fmt::Display) -> DataError {
213    DataError::Validation(format!("arrow IPC reader error: {error}"))
214}
215
216#[cfg(test)]
217mod tests {
218    use super::*;
219    use arrow_array::{Float64Array, StringArray};
220    use arrow_ipc::writer::StreamWriter;
221    use arrow_schema::{DataType, Field, Schema};
222    use std::collections::HashMap;
223    use std::sync::Arc;
224
225    fn build_batch(
226        feature_set_id: &str,
227        observations: &[&str],
228        f0_values: &[Option<f64>],
229        f1_values: &[Option<f64>],
230    ) -> RecordBatch {
231        let mut metadata = HashMap::new();
232        metadata.insert(
233            SCHEMA_METADATA_FEATURE_SET_ID.to_string(),
234            feature_set_id.to_string(),
235        );
236        metadata.insert(
237            SCHEMA_METADATA_REPRESENTATION_ID.to_string(),
238            "tabular_numeric".to_string(),
239        );
240        let schema = Schema::new_with_metadata(
241            vec![
242                Field::new(OBSERVATION_ID_COLUMN, DataType::Utf8, false),
243                Field::new("f0", DataType::Float64, true),
244                Field::new("f1", DataType::Float64, true),
245            ],
246            metadata,
247        );
248        let observation_array = StringArray::from(observations.to_vec());
249        let f0_array = Float64Array::from(f0_values.to_vec());
250        let f1_array = Float64Array::from(f1_values.to_vec());
251        RecordBatch::try_new(
252            Arc::new(schema),
253            vec![
254                Arc::new(observation_array),
255                Arc::new(f0_array),
256                Arc::new(f1_array),
257            ],
258        )
259        .unwrap()
260    }
261
262    fn serialize_stream(batches: &[RecordBatch]) -> Vec<u8> {
263        let mut buffer = Vec::new();
264        {
265            let mut writer = StreamWriter::try_new(&mut buffer, &batches[0].schema()).unwrap();
266            for batch in batches {
267                writer.write(batch).unwrap();
268            }
269            writer.finish().unwrap();
270        }
271        buffer
272    }
273
274    #[test]
275    fn round_trip_columnar_buffer_through_ipc_stream() {
276        let batch = build_batch(
277            "x",
278            &["obs.A", "obs.B", "obs.C"],
279            &[Some(1.0), Some(2.0), Some(3.0)],
280            &[Some(10.0), Some(20.0), Some(30.0)],
281        );
282        let stream_bytes = serialize_stream(&[batch]);
283        let store = read_buffers_from_ipc_stream(stream_bytes.as_slice()).unwrap();
284        assert_eq!(store.len(), 1);
285        let manifest = &store.manifests().unwrap()[0];
286        assert_eq!(manifest.feature_set_id, "x");
287        assert_eq!(manifest.row_count, 3);
288        assert_eq!(manifest.feature_count, 2);
289    }
290
291    #[test]
292    fn null_columns_become_optional_values() {
293        let batch = build_batch(
294            "x",
295            &["obs.A", "obs.B"],
296            &[Some(1.0), None],
297            &[None, Some(20.0)],
298        );
299        let stream_bytes = serialize_stream(&[batch]);
300        let store = read_buffers_from_ipc_stream(stream_bytes.as_slice()).unwrap();
301        // Round-trip through the column-major typed input path preserves the
302        // fingerprint regardless of where nulls sit.
303        let manifest = &store.manifests().unwrap()[0];
304        assert_eq!(manifest.value_count, 4);
305    }
306
307    #[test]
308    fn rejects_batch_without_required_metadata() {
309        let schema = Schema::new(vec![
310            Field::new(OBSERVATION_ID_COLUMN, DataType::Utf8, false),
311            Field::new("f0", DataType::Float64, false),
312        ]);
313        let batch = RecordBatch::try_new(
314            Arc::new(schema),
315            vec![
316                Arc::new(StringArray::from(vec!["obs.A"])),
317                Arc::new(Float64Array::from(vec![1.0])),
318            ],
319        )
320        .unwrap();
321        let stream_bytes = serialize_stream(&[batch]);
322        let error = read_buffers_from_ipc_stream(stream_bytes.as_slice()).unwrap_err();
323        assert!(format!("{error}").contains(SCHEMA_METADATA_FEATURE_SET_ID));
324    }
325
326    #[test]
327    fn rejects_batch_without_observation_id_column() {
328        let mut metadata = HashMap::new();
329        metadata.insert(SCHEMA_METADATA_FEATURE_SET_ID.to_string(), "x".to_string());
330        metadata.insert(
331            SCHEMA_METADATA_REPRESENTATION_ID.to_string(),
332            "tabular_numeric".to_string(),
333        );
334        let schema =
335            Schema::new_with_metadata(vec![Field::new("f0", DataType::Float64, false)], metadata);
336        let batch = RecordBatch::try_new(
337            Arc::new(schema),
338            vec![Arc::new(Float64Array::from(vec![1.0, 2.0]))],
339        )
340        .unwrap();
341        let stream_bytes = serialize_stream(&[batch]);
342        let error = read_buffers_from_ipc_stream(stream_bytes.as_slice()).unwrap_err();
343        assert!(format!("{error}").contains(OBSERVATION_ID_COLUMN));
344    }
345
346    #[test]
347    fn rejects_batch_with_non_float64_feature_column() {
348        let mut metadata = HashMap::new();
349        metadata.insert(SCHEMA_METADATA_FEATURE_SET_ID.to_string(), "x".to_string());
350        metadata.insert(
351            SCHEMA_METADATA_REPRESENTATION_ID.to_string(),
352            "tabular_numeric".to_string(),
353        );
354        let schema = Schema::new_with_metadata(
355            vec![
356                Field::new(OBSERVATION_ID_COLUMN, DataType::Utf8, false),
357                Field::new("f0", DataType::Int32, false),
358            ],
359            metadata,
360        );
361        let batch = RecordBatch::try_new(
362            Arc::new(schema),
363            vec![
364                Arc::new(StringArray::from(vec!["obs.A"])),
365                Arc::new(arrow_array::Int32Array::from(vec![1])),
366            ],
367        )
368        .unwrap();
369        let stream_bytes = serialize_stream(&[batch]);
370        let error = read_buffers_from_ipc_stream(stream_bytes.as_slice()).unwrap_err();
371        assert!(format!("{error}").contains("must be Float64"));
372    }
373}