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 schema.validate_column_names(&self.options.schema_write.column_names)?;
411 match self.state {
412 SinkState::Idle => {}
413 SinkState::Writing => {
414 return Err(LadduDataError::Sink(
415 "parquet sink already initialized".into(),
416 ));
417 }
418 SinkState::Failed => {
419 return Err(LadduDataError::Sink(
420 "parquet sink requires abort after failure".into(),
421 ));
422 }
423 }
424
425 let path = self.output.resolve(plan, "parquet")?;
426 OutputPath::create_parent_dirs(&path)?;
427
428 let arrow_schema = Arc::new(arrow_schema_from_event_schema(
429 &schema,
430 self.options.schema_write.write_weight_column,
431 &self.options.schema_write,
432 ));
433
434 let file = File::create(&path)
435 .map_err(|error| sink_error("create Parquet file", path.display(), error))?;
436
437 let writer = ArrowWriter::try_new(
438 file,
439 Arc::clone(&arrow_schema),
440 self.options.writer_properties.clone(),
441 )
442 .map_err(|error| sink_error("initialize Parquet writer", path.display(), error))?;
443
444 self.arrow_schema = Some(arrow_schema);
445 self.event_schema = Some(schema);
446 self.writer = Some(writer);
447 self.resolved_path = Some(path);
448 self.state = SinkState::Writing;
449
450 Ok(())
451 }
452
453 fn write_batch(&mut self, batch: &EventBatch) -> LadduDataResult<()> {
454 if !matches!(self.state, SinkState::Writing) {
455 return Err(LadduDataError::Sink(
456 match self.state {
457 SinkState::Idle => "parquet sink not initialized",
458 SinkState::Failed => "parquet sink requires abort after failure",
459 SinkState::Writing => unreachable!(),
460 }
461 .into(),
462 ));
463 }
464
465 let arrow_schema = self
466 .arrow_schema
467 .as_ref()
468 .ok_or_else(|| LadduDataError::Sink("parquet sink not initialized".into()))?;
469
470 let event_schema = self
471 .event_schema
472 .as_ref()
473 .ok_or_else(|| LadduDataError::Sink("parquet sink not initialized".into()))?;
474 if event_schema.as_ref() != batch.schema().as_ref() {
475 return Err(LadduDataError::Sink(
476 "batch schema does not match parquet sink schema".into(),
477 ));
478 }
479
480 let rb = match event_batch_to_record_batch(
481 batch,
482 Arc::clone(arrow_schema),
483 self.options.schema_write.write_weight_column,
484 &self.options.schema_write,
485 self.options.schema_write.precision,
486 ) {
487 Ok(rb) => rb,
488 Err(error) => {
489 self.state = SinkState::Failed;
490 let resource = self.resolved_path.as_deref().map_or_else(
491 || "Parquet sink".to_owned(),
492 |path| path.display().to_string(),
493 );
494 return Err(sink_error("encode Parquet batch", resource, error));
495 }
496 };
497
498 let resource = self.resolved_path.as_deref().map_or_else(
499 || "Parquet sink".to_owned(),
500 |path| path.display().to_string(),
501 );
502 let result = self
503 .writer
504 .as_mut()
505 .ok_or_else(|| LadduDataError::Sink("parquet sink not initialized".into()))?
506 .write(&rb)
507 .map_err(|error| sink_error("write Parquet batch", resource, error));
508 if result.is_err() {
509 self.state = SinkState::Failed;
510 }
511 result
512 }
513
514 fn finish(&mut self) -> LadduDataResult<()> {
515 if matches!(self.state, SinkState::Idle) {
516 return Ok(());
517 }
518 if matches!(self.state, SinkState::Failed) {
519 return Err(LadduDataError::Sink(
520 "parquet sink requires abort after failure".into(),
521 ));
522 }
523
524 let resource = self.resolved_path.as_deref().map_or_else(
525 || "Parquet sink".to_owned(),
526 |path| path.display().to_string(),
527 );
528 if let Some(writer) = self.writer.take()
529 && let Err(error) = writer.close()
530 {
531 self.state = SinkState::Failed;
532 return Err(sink_error("finalize Parquet file", resource, error));
533 }
534
535 self.arrow_schema = None;
536 self.event_schema = None;
537 self.state = SinkState::Idle;
538 Ok(())
539 }
540
541 fn abort(&mut self) -> LadduDataResult<()> {
542 self.writer.take();
546 self.arrow_schema = None;
547 self.event_schema = None;
548 self.state = SinkState::Idle;
549 Ok(())
550 }
551}
552
553#[cfg(test)]
554mod tests {
555 use arrow::{
556 array::{ArrayRef, Float32Array, Float64Array},
557 datatypes::{DataType, Field, Schema as ArrowSchema},
558 record_batch::RecordBatch,
559 };
560 use laddu_physics::vectors::RealVec4;
561
562 use super::*;
563 use crate::data::{Dataset, EventBatchBuilder};
564
565 fn temp_path(ext: &str) -> PathBuf {
566 static NEXT_PATH: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
567
568 let nanos = std::time::SystemTime::now()
569 .duration_since(std::time::UNIX_EPOCH)
570 .unwrap()
571 .as_nanos();
572 let sequence = NEXT_PATH.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
573
574 std::env::temp_dir().join(format!(
575 "laddu-parquet-test-{}-{nanos}-{sequence}.{ext}",
576 std::process::id()
577 ))
578 }
579
580 fn v(x: f64) -> RealVec4 {
581 RealVec4 {
582 e: x + 0.3,
583 px: x,
584 py: x + 0.1,
585 pz: x + 0.2,
586 }
587 }
588
589 fn schema() -> Arc<Schema> {
590 Arc::new(Schema::new(["p"], ["mass"], true).unwrap())
591 }
592
593 fn batch() -> EventBatch {
594 let schema = schema();
595 let mut builder = EventBatchBuilder::new(schema);
596
597 for i in 0..4 {
598 builder
599 .push_weighted([v(i as f64)], [100.0 + i as f64], 10.0 + i as f64)
600 .unwrap();
601 }
602
603 builder.finish().unwrap()
604 }
605
606 #[test]
607 fn parquet_sink_and_source_roundtrip_with_f32_write_and_schema_inference() {
608 let path = temp_path("parquet");
609 let batch = batch();
610
611 let mut sink = ParquetSink::builder(path.clone())
612 .precision(Precision::F32)
613 .build();
614
615 sink.begin(Arc::clone(batch.schema()), WritePlan::default())
616 .unwrap();
617 sink.write_batch(&batch).unwrap();
618 sink.finish().unwrap();
619
620 let source = ParquetSource::builder(path.to_str().unwrap())
621 .infer_schema(true)
622 .validate_all_files(true)
623 .build()
624 .unwrap();
625
626 let inferred = source.schema().unwrap();
627 assert_eq!(
628 inferred
629 .p4s()
630 .iter()
631 .map(|n| n.to_string())
632 .collect::<Vec<_>>(),
633 vec!["p"]
634 );
635 assert_eq!(
636 inferred
637 .scalars()
638 .iter()
639 .map(|n| n.to_string())
640 .collect::<Vec<_>>(),
641 vec!["mass"]
642 );
643 assert!(inferred.has_weight());
644
645 let read_batches: Vec<EventBatch> = source
646 .batches(ReadPlan {
647 chunk_size: Some(2),
648 #[cfg(feature = "mpi")]
649 distribution: Default::default(),
650 })
651 .unwrap()
652 .map(Result::unwrap)
653 .collect();
654
655 assert_eq!(
656 read_batches.iter().map(EventBatch::len).collect::<Vec<_>>(),
657 vec![2, 2]
658 );
659
660 let read = EventBatch::concat(&read_batches).unwrap();
661
662 assert_eq!(read.scalar_column(0), &[100.0, 101.0, 102.0, 103.0]);
663 assert_eq!(read.weights_column().unwrap(), &[10.0, 11.0, 12.0, 13.0]);
664 assert!((read.p4_at(0, 2).e - 2.3).abs() < 1.0e-6);
665
666 let _ = std::fs::remove_file(path);
667 }
668
669 #[test]
670 fn parquet_source_require_weight_fails_when_written_without_weight_column() {
671 let path = temp_path("parquet");
672 let schema = Arc::new(Schema::new(["p"], ["mass"], false).unwrap());
673
674 let mut builder = EventBatchBuilder::new(Arc::clone(&schema));
675 builder.push([v(1.0)], [5.0]).unwrap();
676 let batch = builder.finish().unwrap();
677
678 let mut sink = ParquetSink::builder(path.clone())
679 .write_weight_column(WriteWeightColumn::OnlyIfPresent)
680 .build();
681
682 sink.begin(schema, WritePlan::default()).unwrap();
683 sink.write_batch(&batch).unwrap();
684 sink.finish().unwrap();
685
686 let err = ParquetSource::builder(path.to_str().unwrap())
687 .require_weight(true)
688 .build()
689 .unwrap_err();
690
691 assert!(matches!(err, LadduDataError::MissingColumn(name) if name.as_ref() == "weight"));
692
693 let _ = std::fs::remove_file(path);
694 }
695
696 #[test]
697 fn record_batch_to_event_batch_handles_nulls_as_error_or_nan_and_reads_f32_as_f64() {
698 let arrow_schema = Arc::new(ArrowSchema::new(vec![
699 Field::new("p_e", DataType::Float64, true),
700 Field::new("p_px", DataType::Float32, true),
701 Field::new("p_py", DataType::Float64, true),
702 Field::new("p_pz", DataType::Float32, true),
703 Field::new("mass", DataType::Float32, true),
704 Field::new("weight", DataType::Float64, true),
705 ]));
706
707 let rb = RecordBatch::try_new(
708 Arc::clone(&arrow_schema),
709 vec![
710 Arc::new(Float64Array::from(vec![Some(1.0), None])) as ArrayRef,
711 Arc::new(Float32Array::from(vec![Some(0.1), Some(0.2)])) as ArrayRef,
712 Arc::new(Float64Array::from(vec![Some(0.3), Some(0.4)])) as ArrayRef,
713 Arc::new(Float32Array::from(vec![Some(0.5), Some(0.6)])) as ArrayRef,
714 Arc::new(Float32Array::from(vec![Some(2.0), Some(3.0)])) as ArrayRef,
715 Arc::new(Float64Array::from(vec![Some(4.0), Some(5.0)])) as ArrayRef,
716 ],
717 )
718 .unwrap();
719
720 let schema = schema();
721
722 let error_options = ParquetReadOptions {
723 null_handling: NullHandling::Error,
724 ..ParquetReadOptions::default()
725 };
726
727 let err = record_batch_to_event_batch(rb.clone(), Arc::clone(&schema), &error_options)
728 .unwrap_err();
729
730 assert!(matches!(err, LadduDataError::Source(msg) if msg.contains("null in column p_e")));
731
732 let nan_options = ParquetReadOptions {
733 null_handling: NullHandling::NaN,
734 ..ParquetReadOptions::default()
735 };
736
737 let batch = record_batch_to_event_batch(rb, schema, &nan_options).unwrap();
738
739 assert!(batch.p4_at(0, 1).e.is_nan());
740 assert!((batch.p4_at(0, 1).px - 0.2).abs() < 1.0e-6);
741 assert_eq!(batch.scalar_column(0), &[2.0, 3.0]);
742 assert_eq!(batch.weights_column().unwrap(), &[4.0, 5.0]);
743 }
744
745 #[test]
746 fn parquet_sink_rejects_batches_with_different_schema() {
747 let path = temp_path("parquet");
748 let batch = batch();
749
750 let mut sink = ParquetSink::builder(path.clone()).build();
751 sink.begin(Arc::clone(batch.schema()), WritePlan::default())
752 .unwrap();
753
754 let other_schema = Arc::new(Schema::new(["q"], ["mass"], true).unwrap());
755 let mut builder = EventBatchBuilder::new(other_schema);
756 builder.push_weighted([v(1.0)], [1.0], 1.0).unwrap();
757 let other = builder.finish().unwrap();
758
759 let err = sink.write_batch(&other).unwrap_err();
760 assert!(matches!(err, LadduDataError::Sink(msg) if msg.contains("schema")));
761
762 sink.finish().unwrap();
763 let _ = std::fs::remove_file(path);
764 }
765
766 #[test]
767 fn dataset_write_to_parquet_applies_dataset_transformations_before_writing() {
768 let path = temp_path("parquet");
769
770 let dataset = Dataset::from_batch(batch()).filter(|ev| ev.scalar(0) >= 102.0);
771
772 let mut sink = ParquetSink::builder(path.clone()).build();
773 dataset.write_to(&mut sink).unwrap();
774
775 let source = ParquetSource::open(path.to_str().unwrap()).unwrap();
776 let read = EventBatch::concat(
777 &source
778 .batches(ReadPlan::default())
779 .unwrap()
780 .map(Result::unwrap)
781 .collect::<Vec<_>>(),
782 )
783 .unwrap();
784
785 assert_eq!(read.scalar_column(0), &[102.0, 103.0]);
786 assert_eq!(read.weights_column().unwrap(), &[12.0, 13.0]);
787
788 let _ = std::fs::remove_file(path);
789 }
790}