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";
36pub const SCHEMA_METADATA_REPRESENTATION_ID: &str = "dag_ml_data.representation_id";
37pub const OBSERVATION_ID_COLUMN: &str = "observation_id";
38
39pub 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
54pub 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
69pub 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 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}