1use 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";
37pub const SCHEMA_METADATA_REPRESENTATION_ID: &str = "dag_ml_data.representation_id";
39pub const OBSERVATION_ID_COLUMN: &str = "observation_id";
41
42pub 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
52pub 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
104pub 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 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}