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 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
57pub 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
72pub 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 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}