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