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
35pub const SCHEMA_METADATA_FEATURE_SET_ID: &str = "dag_ml_data.feature_set_id";
36pub const SCHEMA_METADATA_REPRESENTATION_ID: &str = "dag_ml_data.representation_id";
37pub const OBSERVATION_ID_COLUMN: &str = "observation_id";
38
39/// Parse an Arrow IPC stream from any `Read` source into a
40/// `NumericFeatureBufferStore`. Each top-level record batch becomes one
41/// buffer; the batches' schemas must declare both
42/// `dag_ml_data.feature_set_id` and `dag_ml_data.representation_id` in
43/// their `metadata` map and expose an `observation_id` UTF-8 column.
44pub fn read_buffers_from_ipc_stream<R: Read>(reader: R) -> Result<NumericFeatureBufferStore> {
45    let stream = StreamReader::try_new(reader, None).map_err(arrow_error)?;
46    let mut matrices = Vec::new();
47    for batch in stream {
48        let batch = batch.map_err(arrow_error)?;
49        matrices.push(record_batch_to_matrix(&batch)?);
50    }
51    NumericFeatureBufferStore::from_f64_column_matrices(matrices)
52}
53
54/// Parse an Arrow IPC file (with footer + magic) from any `Read + Seek`
55/// source into a `NumericFeatureBufferStore`. Same per-batch contract as
56/// `read_buffers_from_ipc_stream`.
57pub fn read_buffers_from_ipc_file<R: Read + std::io::Seek>(
58    reader: R,
59) -> Result<NumericFeatureBufferStore> {
60    let file = FileReader::try_new(reader, None).map_err(arrow_error)?;
61    let mut matrices = Vec::new();
62    for batch in file {
63        let batch = batch.map_err(arrow_error)?;
64        matrices.push(record_batch_to_matrix(&batch)?);
65    }
66    NumericFeatureBufferStore::from_f64_column_matrices(matrices)
67}
68
69/// Convenience: read an Arrow IPC file from disk.
70pub fn read_buffers_from_ipc_path(path: &Path) -> Result<NumericFeatureBufferStore> {
71    let file = File::open(path).map_err(|error| {
72        DataError::Validation(format!(
73            "failed to open arrow IPC file at `{}`: {error}",
74            path.display()
75        ))
76    })?;
77    read_buffers_from_ipc_file(file)
78}
79
80fn record_batch_to_matrix(batch: &RecordBatch) -> Result<NumericFeatureMatrixF64Columnar> {
81    let schema = batch.schema();
82    let metadata = schema.metadata();
83    let feature_set_id = metadata
84        .get(SCHEMA_METADATA_FEATURE_SET_ID)
85        .ok_or_else(|| {
86            DataError::Validation(format!(
87                "arrow IPC batch is missing required schema metadata `{SCHEMA_METADATA_FEATURE_SET_ID}`"
88            ))
89        })?
90        .clone();
91    let representation_raw = metadata.get(SCHEMA_METADATA_REPRESENTATION_ID).ok_or_else(|| {
92        DataError::Validation(format!(
93            "arrow IPC batch `{feature_set_id}` is missing required schema metadata `{SCHEMA_METADATA_REPRESENTATION_ID}`"
94        ))
95    })?;
96    let representation_id = RepresentationId::new(representation_raw)?;
97
98    let observation_field_index = schema.index_of(OBSERVATION_ID_COLUMN).map_err(|_| {
99        DataError::Validation(format!(
100            "arrow IPC batch `{feature_set_id}` is missing the required `{OBSERVATION_ID_COLUMN}` column"
101        ))
102    })?;
103    let observation_array = batch
104        .column(observation_field_index)
105        .as_any()
106        .downcast_ref::<StringArray>()
107        .ok_or_else(|| {
108            DataError::Validation(format!(
109                "arrow IPC batch `{feature_set_id}` `{OBSERVATION_ID_COLUMN}` column must be Utf8"
110            ))
111        })?;
112    if observation_array.null_count() != 0 {
113        return Err(DataError::Validation(format!(
114            "arrow IPC batch `{feature_set_id}` `{OBSERVATION_ID_COLUMN}` column contains nulls"
115        )));
116    }
117    let mut observation_ids = Vec::with_capacity(observation_array.len());
118    for raw in observation_array.iter() {
119        let value = raw.ok_or_else(|| {
120            DataError::Validation(format!(
121                "arrow IPC batch `{feature_set_id}` `{OBSERVATION_ID_COLUMN}` column contains nulls"
122            ))
123        })?;
124        observation_ids.push(ObservationId::new(value)?);
125    }
126
127    let mut feature_names = Vec::new();
128    let mut columns = Vec::new();
129    let mut masks = Vec::new();
130    let mut any_null = false;
131    for (idx, field) in schema.fields().iter().enumerate() {
132        if idx == observation_field_index {
133            continue;
134        }
135        let column = batch.column(idx);
136        let float_column = column
137            .as_any()
138            .downcast_ref::<Float64Array>()
139            .ok_or_else(|| {
140                DataError::Validation(format!(
141                    "arrow IPC batch `{feature_set_id}` column `{}` must be Float64",
142                    field.name()
143                ))
144            })?;
145        let row_count = float_column.len();
146        let mut values = Vec::with_capacity(row_count);
147        let mut mask = Vec::with_capacity(row_count);
148        for row in 0..row_count {
149            if float_column.is_null(row) {
150                values.push(0.0);
151                mask.push(false);
152                any_null = true;
153            } else {
154                values.push(float_column.value(row));
155                mask.push(true);
156            }
157        }
158        feature_names.push(field.name().clone());
159        columns.push(values);
160        masks.push(mask);
161    }
162    if feature_names.is_empty() {
163        return Err(DataError::Validation(format!(
164            "arrow IPC batch `{feature_set_id}` has no feature columns besides `{OBSERVATION_ID_COLUMN}`"
165        )));
166    }
167    Ok(NumericFeatureMatrixF64Columnar {
168        feature_set_id,
169        representation_id,
170        feature_names,
171        observation_ids,
172        columns,
173        validity_masks: any_null.then_some(masks),
174    })
175}
176
177fn arrow_error(error: impl std::fmt::Display) -> DataError {
178    DataError::Validation(format!("arrow IPC reader error: {error}"))
179}
180
181#[cfg(test)]
182mod tests {
183    use super::*;
184    use arrow_array::{Float64Array, StringArray};
185    use arrow_ipc::writer::StreamWriter;
186    use arrow_schema::{DataType, Field, Schema};
187    use std::collections::HashMap;
188    use std::sync::Arc;
189
190    fn build_batch(
191        feature_set_id: &str,
192        observations: &[&str],
193        f0_values: &[Option<f64>],
194        f1_values: &[Option<f64>],
195    ) -> RecordBatch {
196        let mut metadata = HashMap::new();
197        metadata.insert(
198            SCHEMA_METADATA_FEATURE_SET_ID.to_string(),
199            feature_set_id.to_string(),
200        );
201        metadata.insert(
202            SCHEMA_METADATA_REPRESENTATION_ID.to_string(),
203            "tabular_numeric".to_string(),
204        );
205        let schema = Schema::new_with_metadata(
206            vec![
207                Field::new(OBSERVATION_ID_COLUMN, DataType::Utf8, false),
208                Field::new("f0", DataType::Float64, true),
209                Field::new("f1", DataType::Float64, true),
210            ],
211            metadata,
212        );
213        let observation_array = StringArray::from(observations.to_vec());
214        let f0_array = Float64Array::from(f0_values.to_vec());
215        let f1_array = Float64Array::from(f1_values.to_vec());
216        RecordBatch::try_new(
217            Arc::new(schema),
218            vec![
219                Arc::new(observation_array),
220                Arc::new(f0_array),
221                Arc::new(f1_array),
222            ],
223        )
224        .unwrap()
225    }
226
227    fn serialize_stream(batches: &[RecordBatch]) -> Vec<u8> {
228        let mut buffer = Vec::new();
229        {
230            let mut writer = StreamWriter::try_new(&mut buffer, &batches[0].schema()).unwrap();
231            for batch in batches {
232                writer.write(batch).unwrap();
233            }
234            writer.finish().unwrap();
235        }
236        buffer
237    }
238
239    #[test]
240    fn round_trip_columnar_buffer_through_ipc_stream() {
241        let batch = build_batch(
242            "x",
243            &["obs.A", "obs.B", "obs.C"],
244            &[Some(1.0), Some(2.0), Some(3.0)],
245            &[Some(10.0), Some(20.0), Some(30.0)],
246        );
247        let stream_bytes = serialize_stream(&[batch]);
248        let store = read_buffers_from_ipc_stream(stream_bytes.as_slice()).unwrap();
249        assert_eq!(store.len(), 1);
250        let manifest = &store.manifests().unwrap()[0];
251        assert_eq!(manifest.feature_set_id, "x");
252        assert_eq!(manifest.row_count, 3);
253        assert_eq!(manifest.feature_count, 2);
254    }
255
256    #[test]
257    fn null_columns_become_optional_values() {
258        let batch = build_batch(
259            "x",
260            &["obs.A", "obs.B"],
261            &[Some(1.0), None],
262            &[None, Some(20.0)],
263        );
264        let stream_bytes = serialize_stream(&[batch]);
265        let store = read_buffers_from_ipc_stream(stream_bytes.as_slice()).unwrap();
266        // Round-trip through the column-major typed input path preserves the
267        // fingerprint regardless of where nulls sit.
268        let manifest = &store.manifests().unwrap()[0];
269        assert_eq!(manifest.value_count, 4);
270    }
271
272    #[test]
273    fn rejects_batch_without_required_metadata() {
274        let schema = Schema::new(vec![
275            Field::new(OBSERVATION_ID_COLUMN, DataType::Utf8, false),
276            Field::new("f0", DataType::Float64, false),
277        ]);
278        let batch = RecordBatch::try_new(
279            Arc::new(schema),
280            vec![
281                Arc::new(StringArray::from(vec!["obs.A"])),
282                Arc::new(Float64Array::from(vec![1.0])),
283            ],
284        )
285        .unwrap();
286        let stream_bytes = serialize_stream(&[batch]);
287        let error = read_buffers_from_ipc_stream(stream_bytes.as_slice()).unwrap_err();
288        assert!(format!("{error}").contains(SCHEMA_METADATA_FEATURE_SET_ID));
289    }
290
291    #[test]
292    fn rejects_batch_without_observation_id_column() {
293        let mut metadata = HashMap::new();
294        metadata.insert(SCHEMA_METADATA_FEATURE_SET_ID.to_string(), "x".to_string());
295        metadata.insert(
296            SCHEMA_METADATA_REPRESENTATION_ID.to_string(),
297            "tabular_numeric".to_string(),
298        );
299        let schema =
300            Schema::new_with_metadata(vec![Field::new("f0", DataType::Float64, false)], metadata);
301        let batch = RecordBatch::try_new(
302            Arc::new(schema),
303            vec![Arc::new(Float64Array::from(vec![1.0, 2.0]))],
304        )
305        .unwrap();
306        let stream_bytes = serialize_stream(&[batch]);
307        let error = read_buffers_from_ipc_stream(stream_bytes.as_slice()).unwrap_err();
308        assert!(format!("{error}").contains(OBSERVATION_ID_COLUMN));
309    }
310
311    #[test]
312    fn rejects_batch_with_non_float64_feature_column() {
313        let mut metadata = HashMap::new();
314        metadata.insert(SCHEMA_METADATA_FEATURE_SET_ID.to_string(), "x".to_string());
315        metadata.insert(
316            SCHEMA_METADATA_REPRESENTATION_ID.to_string(),
317            "tabular_numeric".to_string(),
318        );
319        let schema = Schema::new_with_metadata(
320            vec![
321                Field::new(OBSERVATION_ID_COLUMN, DataType::Utf8, false),
322                Field::new("f0", DataType::Int32, false),
323            ],
324            metadata,
325        );
326        let batch = RecordBatch::try_new(
327            Arc::new(schema),
328            vec![
329                Arc::new(StringArray::from(vec!["obs.A"])),
330                Arc::new(arrow_array::Int32Array::from(vec![1])),
331            ],
332        )
333        .unwrap();
334        let stream_bytes = serialize_stream(&[batch]);
335        let error = read_buffers_from_ipc_stream(stream_bytes.as_slice()).unwrap_err();
336        assert!(format!("{error}").contains("must be Float64"));
337    }
338}