1use std::{
2 fs::File,
3 path::{Path, PathBuf},
4 sync::Arc,
5};
6
7use arrow::datatypes::SchemaRef;
8use parquet::{
9 arrow::{ArrowWriter, arrow_reader::ParquetRecordBatchReaderBuilder},
10 file::properties::WriterProperties,
11};
12
13use crate::{
14 LadduDataError, LadduDataResult,
15 data::EventBatch,
16 io::{
17 DataFragment, EventSink, EventSource, FragmentedSource, OutputMode, OutputPath, ReadPlan,
18 SinkState, SourceBuild, SourceBuildOptions, SourceCapabilities, WritePlan, build_source,
19 fragmented_batches, sink_error, source_error,
20 },
21 schema::{
22 Precision, Schema, SchemaColumnNames, SchemaInferenceOptions, SchemaWriteOptions,
23 WriteWeightColumn,
24 },
25};
26
27mod decode;
28mod encode;
29
30#[cfg(test)]
31use decode::record_batch_to_event_batch;
32use decode::{arrow_columns, open_parquet_fragment_reader, parquet_arrow_schema};
33use encode::{arrow_schema_from_event_schema, event_batch_to_record_batch};
34
35#[derive(Clone, Debug)]
37pub struct ParquetSource {
38 files: Arc<[Arc<PathBuf>]>,
39 schema: Arc<Schema>,
40 options: ParquetReadOptions,
41}
42
43#[derive(Clone, Debug)]
45pub struct ParquetReadOptions {
46 pub infer_schema: bool,
48 pub validate_all_files: bool,
50 pub null_handling: NullHandling,
52 pub sort_glob: bool,
54 pub schema_inference: SchemaInferenceOptions,
56}
57
58impl Default for ParquetReadOptions {
59 fn default() -> Self {
60 Self {
61 infer_schema: true,
62 validate_all_files: true,
63 null_handling: NullHandling::Error,
64 sort_glob: true,
65 schema_inference: SchemaInferenceOptions::default(),
66 }
67 }
68}
69
70#[derive(Clone, Copy, Debug)]
72pub enum NullHandling {
73 Error,
75 NaN,
77}
78
79#[derive(Clone, Debug)]
81pub struct ParquetFragmentKey {
82 pub file: Arc<PathBuf>,
84 pub row_group: usize,
86}
87
88impl ParquetSource {
89 pub fn open(pattern: impl AsRef<str>) -> LadduDataResult<Self> {
96 Self::builder(pattern).build()
97 }
98
99 pub fn builder(pattern: impl AsRef<str>) -> ParquetSourceBuilder {
101 ParquetSourceBuilder {
102 pattern: pattern.as_ref().to_owned(),
103 schema: None,
104 options: ParquetReadOptions::default(),
105 }
106 }
107
108 pub fn files(&self) -> &[Arc<PathBuf>] {
110 &self.files
111 }
112}
113
114pub struct ParquetSourceBuilder {
116 pattern: String,
117 schema: Option<Arc<Schema>>,
118 options: ParquetReadOptions,
119}
120
121impl ParquetSourceBuilder {
122 pub fn schema(mut self, schema: Arc<Schema>) -> Self {
124 self.schema = Some(schema);
125 self.options.infer_schema = false;
126 self
127 }
128
129 pub fn infer_schema(mut self, value: bool) -> Self {
131 self.options.infer_schema = value;
132 self
133 }
134
135 pub fn require_weight(mut self, value: bool) -> Self {
137 self.options.schema_inference.require_weight = value;
138 self
139 }
140
141 pub fn validate_all_files(mut self, value: bool) -> Self {
143 self.options.validate_all_files = value;
144 self
145 }
146
147 pub fn nulls_as_nan(mut self) -> Self {
149 self.options.null_handling = NullHandling::NaN;
150 self
151 }
152
153 pub fn error_on_nulls(mut self) -> Self {
155 self.options.null_handling = NullHandling::Error;
156 self
157 }
158
159 pub fn sort_glob(mut self, value: bool) -> Self {
161 self.options.sort_glob = value;
162 self
163 }
164
165 pub fn schema_inference(mut self, options: SchemaInferenceOptions) -> Self {
167 self.options.schema_inference = options;
168 self
169 }
170
171 pub fn build(self) -> LadduDataResult<ParquetSource> {
178 let ParquetSourceBuilder {
179 pattern,
180 schema: explicit_schema,
181 options,
182 } = self;
183 let infer_options = options.schema_inference.clone();
184 let validate_options = options.schema_inference.clone();
185 let SourceBuild { files, schema, .. } = build_source(
186 SourceBuildOptions {
187 pattern: &pattern,
188 sort: options.sort_glob,
189 format: "parquet",
190 explicit_schema,
191 infer_schema: options.infer_schema,
192 validate_all_files: options.validate_all_files,
193 },
194 |_| Ok(()),
195 move |path, _| {
196 let arrow_schema = parquet_arrow_schema(path)?;
197 Schema::infer_from_columns(arrow_columns(&arrow_schema), &infer_options)
198 },
199 move |path, schema, _| {
200 let arrow_schema = parquet_arrow_schema(path)?;
201 schema.validate_required_columns(arrow_columns(&arrow_schema), &validate_options)
202 },
203 )?;
204
205 Ok(ParquetSource {
206 files,
207 schema,
208 options,
209 })
210 }
211}
212
213impl EventSource for ParquetSource {
214 fn schema(&self) -> LadduDataResult<Arc<Schema>> {
215 Ok(Arc::clone(&self.schema))
216 }
217
218 fn capabilities(&self) -> SourceCapabilities {
219 SourceCapabilities {
220 exact_len: true,
221 exact_weighted_total: false,
222 random_access: false,
223 deterministic_partitioning: true,
224 predicate_pushdown: false,
225 projection_pushdown: true,
226 streaming: true,
227 }
228 }
229
230 fn num_events(&self) -> LadduDataResult<Option<u64>> {
231 Ok(Some(self.fragments()?.iter().map(|f| f.rows).sum()))
232 }
233
234 fn batches(
235 &self,
236 plan: ReadPlan,
237 ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
238 fragmented_batches(Arc::new(self.clone()), plan)
239 }
240}
241
242impl FragmentedSource for ParquetSource {
243 type Key = ParquetFragmentKey;
244
245 fn fragments(&self) -> LadduDataResult<Vec<DataFragment<Self::Key>>> {
246 let mut fragments = Vec::new();
247 let mut global_start = 0_u64;
248
249 for path in self.files.iter() {
250 let resource = path.as_ref().display().to_string();
251 let file = File::open(path.as_ref())
252 .map_err(|error| source_error("open Parquet file", &resource, error))?;
253
254 let builder = ParquetRecordBatchReaderBuilder::try_new(file)
255 .map_err(|error| source_error("read Parquet metadata", &resource, error))?;
256
257 let metadata = builder.metadata();
258
259 for row_group in 0..metadata.num_row_groups() {
260 let rows = metadata.row_group(row_group).num_rows() as u64;
261
262 fragments.push(DataFragment {
263 key: ParquetFragmentKey {
264 file: Arc::clone(path),
265 row_group,
266 },
267 global_start,
268 rows,
269 });
270
271 global_start += rows;
272 }
273 }
274
275 Ok(fragments)
276 }
277
278 fn read_fragment_range(
279 &self,
280 key: &Self::Key,
281 local_start: usize,
282 local_len: usize,
283 chunk_size: Option<usize>,
284 ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
285 open_parquet_fragment_reader(
286 Arc::clone(&self.schema),
287 self.options.clone(),
288 key.clone(),
289 local_start,
290 local_len,
291 chunk_size,
292 )
293 }
294}
295
296pub struct ParquetSink {
298 output: OutputPath,
299 writer: Option<ArrowWriter<File>>,
300 arrow_schema: Option<SchemaRef>,
301 event_schema: Option<Arc<Schema>>,
302 options: ParquetWriteOptions,
303 resolved_path: Option<PathBuf>,
304 state: SinkState,
305}
306
307#[derive(Clone, Debug, Default)]
309pub struct ParquetWriteOptions {
310 pub writer_properties: Option<WriterProperties>,
312 pub schema_write: SchemaWriteOptions,
314}
315
316impl ParquetSink {
317 pub fn create(path: impl Into<PathBuf>) -> Self {
319 Self::builder(path).build()
320 }
321
322 pub fn builder(path: impl Into<PathBuf>) -> ParquetSinkBuilder {
324 ParquetSinkBuilder {
325 output: OutputPath::new(path),
326 options: ParquetWriteOptions::default(),
327 }
328 }
329
330 pub fn resolved_path(&self) -> Option<&Path> {
332 self.resolved_path.as_deref()
333 }
334}
335
336pub struct ParquetSinkBuilder {
338 output: OutputPath,
339 options: ParquetWriteOptions,
340}
341
342impl ParquetSinkBuilder {
343 pub fn output_mode(mut self, mode: OutputMode) -> Self {
345 self.output = self.output.with_mode(mode);
346 self
347 }
348
349 pub fn single_file(self) -> Self {
351 self.output_mode(OutputMode::SingleFile)
352 }
353
354 pub fn per_rank_files(self) -> Self {
356 self.output_mode(OutputMode::PerRankFiles)
357 }
358
359 pub fn auto_output(self) -> Self {
361 self.output_mode(OutputMode::Auto)
362 }
363
364 pub fn schema_write(mut self, options: SchemaWriteOptions) -> Self {
366 self.options.schema_write = options;
367 self
368 }
369
370 pub fn column_names(mut self, column_names: SchemaColumnNames) -> Self {
372 self.options.schema_write.column_names = column_names;
373 self
374 }
375
376 pub fn precision(mut self, precision: Precision) -> Self {
378 self.options.schema_write.precision = precision;
379 self
380 }
381
382 pub fn writer_properties(mut self, props: WriterProperties) -> Self {
384 self.options.writer_properties = Some(props);
385 self
386 }
387
388 pub fn write_weight_column(mut self, value: WriteWeightColumn) -> Self {
390 self.options.schema_write.write_weight_column = value;
391 self
392 }
393
394 pub fn build(self) -> ParquetSink {
396 ParquetSink {
397 output: self.output,
398 writer: None,
399 arrow_schema: None,
400 event_schema: None,
401 options: self.options,
402 resolved_path: None,
403 state: SinkState::Idle,
404 }
405 }
406}
407
408impl EventSink for ParquetSink {
409 fn begin(&mut self, schema: Arc<Schema>, plan: WritePlan) -> LadduDataResult<()> {
410 match self.state {
411 SinkState::Idle => {}
412 SinkState::Writing => {
413 return Err(LadduDataError::Sink(
414 "parquet sink already initialized".into(),
415 ));
416 }
417 SinkState::Failed => {
418 return Err(LadduDataError::Sink(
419 "parquet sink requires abort after failure".into(),
420 ));
421 }
422 }
423
424 let path = self.output.resolve(plan, "parquet")?;
425 OutputPath::create_parent_dirs(&path)?;
426
427 let arrow_schema = Arc::new(arrow_schema_from_event_schema(
428 &schema,
429 self.options.schema_write.write_weight_column,
430 &self.options.schema_write,
431 ));
432
433 let file = File::create(&path)
434 .map_err(|error| sink_error("create Parquet file", path.display(), error))?;
435
436 let writer = ArrowWriter::try_new(
437 file,
438 Arc::clone(&arrow_schema),
439 self.options.writer_properties.clone(),
440 )
441 .map_err(|error| sink_error("initialize Parquet writer", path.display(), error))?;
442
443 self.arrow_schema = Some(arrow_schema);
444 self.event_schema = Some(schema);
445 self.writer = Some(writer);
446 self.resolved_path = Some(path);
447 self.state = SinkState::Writing;
448
449 Ok(())
450 }
451
452 fn write_batch(&mut self, batch: &EventBatch) -> LadduDataResult<()> {
453 if !matches!(self.state, SinkState::Writing) {
454 return Err(LadduDataError::Sink(
455 match self.state {
456 SinkState::Idle => "parquet sink not initialized",
457 SinkState::Failed => "parquet sink requires abort after failure",
458 SinkState::Writing => unreachable!(),
459 }
460 .into(),
461 ));
462 }
463
464 let arrow_schema = self
465 .arrow_schema
466 .as_ref()
467 .ok_or_else(|| LadduDataError::Sink("parquet sink not initialized".into()))?;
468
469 let event_schema = self
470 .event_schema
471 .as_ref()
472 .ok_or_else(|| LadduDataError::Sink("parquet sink not initialized".into()))?;
473 if event_schema.as_ref() != batch.schema().as_ref() {
474 return Err(LadduDataError::Sink(
475 "batch schema does not match parquet sink schema".into(),
476 ));
477 }
478
479 let rb = match event_batch_to_record_batch(
480 batch,
481 Arc::clone(arrow_schema),
482 self.options.schema_write.write_weight_column,
483 &self.options.schema_write,
484 self.options.schema_write.precision,
485 ) {
486 Ok(rb) => rb,
487 Err(error) => {
488 self.state = SinkState::Failed;
489 let resource = self.resolved_path.as_deref().map_or_else(
490 || "Parquet sink".to_owned(),
491 |path| path.display().to_string(),
492 );
493 return Err(sink_error("encode Parquet batch", resource, error));
494 }
495 };
496
497 let resource = self.resolved_path.as_deref().map_or_else(
498 || "Parquet sink".to_owned(),
499 |path| path.display().to_string(),
500 );
501 let result = self
502 .writer
503 .as_mut()
504 .ok_or_else(|| LadduDataError::Sink("parquet sink not initialized".into()))?
505 .write(&rb)
506 .map_err(|error| sink_error("write Parquet batch", resource, error));
507 if result.is_err() {
508 self.state = SinkState::Failed;
509 }
510 result
511 }
512
513 fn finish(&mut self) -> LadduDataResult<()> {
514 if matches!(self.state, SinkState::Idle) {
515 return Ok(());
516 }
517 if matches!(self.state, SinkState::Failed) {
518 return Err(LadduDataError::Sink(
519 "parquet sink requires abort after failure".into(),
520 ));
521 }
522
523 let resource = self.resolved_path.as_deref().map_or_else(
524 || "Parquet sink".to_owned(),
525 |path| path.display().to_string(),
526 );
527 if let Some(writer) = self.writer.take()
528 && let Err(error) = writer.close()
529 {
530 self.state = SinkState::Failed;
531 return Err(sink_error("finalize Parquet file", resource, error));
532 }
533
534 self.arrow_schema = None;
535 self.event_schema = None;
536 self.state = SinkState::Idle;
537 Ok(())
538 }
539
540 fn abort(&mut self) -> LadduDataResult<()> {
541 self.writer.take();
545 self.arrow_schema = None;
546 self.event_schema = None;
547 self.state = SinkState::Idle;
548 Ok(())
549 }
550}
551
552#[cfg(test)]
553mod tests {
554 use arrow::{
555 array::{ArrayRef, Float32Array, Float64Array},
556 datatypes::{DataType, Field, Schema as ArrowSchema},
557 record_batch::RecordBatch,
558 };
559 use laddu_physics::vectors::RealVec4;
560
561 use super::*;
562 use crate::data::{Dataset, EventBatchBuilder};
563
564 fn temp_path(ext: &str) -> PathBuf {
565 static NEXT_PATH: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
566
567 let nanos = std::time::SystemTime::now()
568 .duration_since(std::time::UNIX_EPOCH)
569 .unwrap()
570 .as_nanos();
571 let sequence = NEXT_PATH.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
572
573 std::env::temp_dir().join(format!(
574 "laddu-parquet-test-{}-{nanos}-{sequence}.{ext}",
575 std::process::id()
576 ))
577 }
578
579 fn v(x: f64) -> RealVec4 {
580 RealVec4 {
581 e: x + 0.3,
582 px: x,
583 py: x + 0.1,
584 pz: x + 0.2,
585 }
586 }
587
588 fn schema() -> Arc<Schema> {
589 Arc::new(Schema::new(["p"], ["mass"], true).unwrap())
590 }
591
592 fn batch() -> EventBatch {
593 let schema = schema();
594 let mut builder = EventBatchBuilder::new(schema);
595
596 for i in 0..4 {
597 builder
598 .push_weighted([v(i as f64)], [100.0 + i as f64], 10.0 + i as f64)
599 .unwrap();
600 }
601
602 builder.finish().unwrap()
603 }
604
605 #[test]
606 fn parquet_sink_and_source_roundtrip_with_f32_write_and_schema_inference() {
607 let path = temp_path("parquet");
608 let batch = batch();
609
610 let mut sink = ParquetSink::builder(path.clone())
611 .precision(Precision::F32)
612 .build();
613
614 sink.begin(Arc::clone(batch.schema()), WritePlan::default())
615 .unwrap();
616 sink.write_batch(&batch).unwrap();
617 sink.finish().unwrap();
618
619 let source = ParquetSource::builder(path.to_str().unwrap())
620 .infer_schema(true)
621 .validate_all_files(true)
622 .build()
623 .unwrap();
624
625 let inferred = source.schema().unwrap();
626 assert_eq!(
627 inferred
628 .p4s()
629 .iter()
630 .map(|n| n.to_string())
631 .collect::<Vec<_>>(),
632 vec!["p"]
633 );
634 assert_eq!(
635 inferred
636 .scalars()
637 .iter()
638 .map(|n| n.to_string())
639 .collect::<Vec<_>>(),
640 vec!["mass"]
641 );
642 assert!(inferred.has_weight());
643
644 let read_batches: Vec<EventBatch> = source
645 .batches(ReadPlan {
646 chunk_size: Some(2),
647 #[cfg(feature = "mpi")]
648 distribution: Default::default(),
649 })
650 .unwrap()
651 .map(Result::unwrap)
652 .collect();
653
654 assert_eq!(
655 read_batches.iter().map(EventBatch::len).collect::<Vec<_>>(),
656 vec![2, 2]
657 );
658
659 let read = EventBatch::concat(&read_batches).unwrap();
660
661 assert_eq!(read.scalar_column(0), &[100.0, 101.0, 102.0, 103.0]);
662 assert_eq!(read.weights_column().unwrap(), &[10.0, 11.0, 12.0, 13.0]);
663 assert!((read.p4_at(0, 2).e - 2.3).abs() < 1.0e-6);
664
665 let _ = std::fs::remove_file(path);
666 }
667
668 #[test]
669 fn parquet_source_require_weight_fails_when_written_without_weight_column() {
670 let path = temp_path("parquet");
671 let schema = Arc::new(Schema::new(["p"], ["mass"], false).unwrap());
672
673 let mut builder = EventBatchBuilder::new(Arc::clone(&schema));
674 builder.push([v(1.0)], [5.0]).unwrap();
675 let batch = builder.finish().unwrap();
676
677 let mut sink = ParquetSink::builder(path.clone())
678 .write_weight_column(WriteWeightColumn::OnlyIfPresent)
679 .build();
680
681 sink.begin(schema, WritePlan::default()).unwrap();
682 sink.write_batch(&batch).unwrap();
683 sink.finish().unwrap();
684
685 let err = ParquetSource::builder(path.to_str().unwrap())
686 .require_weight(true)
687 .build()
688 .unwrap_err();
689
690 assert!(matches!(err, LadduDataError::MissingColumn(name) if name.as_ref() == "weight"));
691
692 let _ = std::fs::remove_file(path);
693 }
694
695 #[test]
696 fn record_batch_to_event_batch_handles_nulls_as_error_or_nan_and_reads_f32_as_f64() {
697 let arrow_schema = Arc::new(ArrowSchema::new(vec![
698 Field::new("p_e", DataType::Float64, true),
699 Field::new("p_px", DataType::Float32, true),
700 Field::new("p_py", DataType::Float64, true),
701 Field::new("p_pz", DataType::Float32, true),
702 Field::new("mass", DataType::Float32, true),
703 Field::new("weight", DataType::Float64, true),
704 ]));
705
706 let rb = RecordBatch::try_new(
707 Arc::clone(&arrow_schema),
708 vec![
709 Arc::new(Float64Array::from(vec![Some(1.0), None])) as ArrayRef,
710 Arc::new(Float32Array::from(vec![Some(0.1), Some(0.2)])) as ArrayRef,
711 Arc::new(Float64Array::from(vec![Some(0.3), Some(0.4)])) as ArrayRef,
712 Arc::new(Float32Array::from(vec![Some(0.5), Some(0.6)])) as ArrayRef,
713 Arc::new(Float32Array::from(vec![Some(2.0), Some(3.0)])) as ArrayRef,
714 Arc::new(Float64Array::from(vec![Some(4.0), Some(5.0)])) as ArrayRef,
715 ],
716 )
717 .unwrap();
718
719 let schema = schema();
720
721 let error_options = ParquetReadOptions {
722 null_handling: NullHandling::Error,
723 ..ParquetReadOptions::default()
724 };
725
726 let err = record_batch_to_event_batch(rb.clone(), Arc::clone(&schema), &error_options)
727 .unwrap_err();
728
729 assert!(matches!(err, LadduDataError::Source(msg) if msg.contains("null in column p_e")));
730
731 let nan_options = ParquetReadOptions {
732 null_handling: NullHandling::NaN,
733 ..ParquetReadOptions::default()
734 };
735
736 let batch = record_batch_to_event_batch(rb, schema, &nan_options).unwrap();
737
738 assert!(batch.p4_at(0, 1).e.is_nan());
739 assert!((batch.p4_at(0, 1).px - 0.2).abs() < 1.0e-6);
740 assert_eq!(batch.scalar_column(0), &[2.0, 3.0]);
741 assert_eq!(batch.weights_column().unwrap(), &[4.0, 5.0]);
742 }
743
744 #[test]
745 fn parquet_sink_rejects_batches_with_different_schema() {
746 let path = temp_path("parquet");
747 let batch = batch();
748
749 let mut sink = ParquetSink::builder(path.clone()).build();
750 sink.begin(Arc::clone(batch.schema()), WritePlan::default())
751 .unwrap();
752
753 let other_schema = Arc::new(Schema::new(["q"], ["mass"], true).unwrap());
754 let mut builder = EventBatchBuilder::new(other_schema);
755 builder.push_weighted([v(1.0)], [1.0], 1.0).unwrap();
756 let other = builder.finish().unwrap();
757
758 let err = sink.write_batch(&other).unwrap_err();
759 assert!(matches!(err, LadduDataError::Sink(msg) if msg.contains("schema")));
760
761 sink.finish().unwrap();
762 let _ = std::fs::remove_file(path);
763 }
764
765 #[test]
766 fn dataset_write_to_parquet_applies_dataset_transformations_before_writing() {
767 let path = temp_path("parquet");
768
769 let dataset = Dataset::from_batch(batch()).filter(|ev| ev.scalar(0) >= 102.0);
770
771 let mut sink = ParquetSink::builder(path.clone()).build();
772 dataset.write_to(&mut sink).unwrap();
773
774 let source = ParquetSource::open(path.to_str().unwrap()).unwrap();
775 let read = EventBatch::concat(
776 &source
777 .batches(ReadPlan::default())
778 .unwrap()
779 .map(Result::unwrap)
780 .collect::<Vec<_>>(),
781 )
782 .unwrap();
783
784 assert_eq!(read.scalar_column(0), &[102.0, 103.0]);
785 assert_eq!(read.weights_column().unwrap(), &[12.0, 13.0]);
786
787 let _ = std::fs::remove_file(path);
788 }
789}