1use serde::{Deserialize, Serialize};
2use std::fmt::Display;
3use std::sync::Arc;
4
5use crate::{LadduDataError, LadduDataResult, data::EventBatch, schema::Schema};
6
7pub mod memory;
9mod output;
10pub mod parquet;
12pub mod root;
14
15mod source;
16
17pub use output::{OutputMode, OutputPath};
18
19pub(crate) use source::{SourceBuild, SourceBuildOptions, build_source};
20
21pub(crate) fn source_error(
23 operation: impl AsRef<str>,
24 resource: impl Display,
25 cause: impl Display,
26) -> LadduDataError {
27 LadduDataError::Source(format!("{} `{resource}`: {cause}", operation.as_ref()))
28}
29
30pub(crate) fn sink_error(
32 operation: impl AsRef<str>,
33 resource: impl Display,
34 cause: impl Display,
35) -> LadduDataError {
36 LadduDataError::Sink(format!("{} `{resource}`: {cause}", operation.as_ref()))
37}
38
39#[cfg(test)]
40mod context_tests {
41 use super::*;
42
43 #[test]
44 fn source_and_sink_context_have_stable_operation_resource_format() {
45 assert!(matches!(
46 source_error("read ROOT tree", "events.root::events", "branch failed"),
47 LadduDataError::Source(message)
48 if message == "read ROOT tree `events.root::events`: branch failed"
49 ));
50 assert!(matches!(
51 sink_error("write Parquet file", "events.parquet", "disk full"),
52 LadduDataError::Sink(message)
53 if message == "write Parquet file `events.parquet`: disk full"
54 ));
55 }
56}
57
58#[cfg(test)]
59mod contract_tests;
60
61#[cfg(feature = "mpi")]
62#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize)]
64pub enum Distribution {
65 #[default]
67 Serial,
68 Mpi {
70 rank: usize,
72 nranks: usize,
74 partitioning: Partitioning,
76 },
77}
78
79#[cfg(feature = "mpi")]
80impl Distribution {
81 pub fn serial() -> Self {
83 Self::Serial
84 }
85
86 pub fn from_world<C>(world: &C) -> Self
88 where
89 C: mpi::topology::Communicator,
90 {
91 Self::Mpi {
92 rank: world.rank() as usize,
93 nranks: world.size() as usize,
94 partitioning: Partitioning::default(),
95 }
96 }
97
98 pub fn rank(self) -> usize {
100 match self {
101 Self::Serial => 0,
102 Self::Mpi { rank, .. } => rank,
103 }
104 }
105
106 pub fn nranks(self) -> usize {
108 match self {
109 Self::Serial => 1,
110 Self::Mpi { nranks, .. } => nranks,
111 }
112 }
113
114 pub fn partitioning(self) -> Partitioning {
116 match self {
117 Self::Serial => Partitioning::Contiguous,
118 Self::Mpi { partitioning, .. } => partitioning,
119 }
120 }
121
122 pub fn with_partitioning(self, partitioning: Partitioning) -> Self {
124 match self {
125 Self::Serial => Self::Serial,
126 Self::Mpi { rank, nranks, .. } => Self::Mpi {
127 rank,
128 nranks,
129 partitioning,
130 },
131 }
132 }
133}
134
135#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
137pub enum Partitioning {
138 #[default]
140 Contiguous,
141
142 FileGroups,
144
145 Rows,
148}
149
150#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize)]
152pub struct ReadPlan {
153 pub chunk_size: Option<usize>,
155
156 #[cfg(feature = "mpi")]
157 pub distribution: Distribution,
159}
160
161impl ReadPlan {
162 pub fn serial() -> Self {
164 Self::default()
165 }
166
167 pub fn rank(&self) -> usize {
169 #[cfg(feature = "mpi")]
170 {
171 self.distribution.rank()
172 }
173
174 #[cfg(not(feature = "mpi"))]
175 {
176 0
177 }
178 }
179
180 pub fn nranks(&self) -> usize {
182 #[cfg(feature = "mpi")]
183 {
184 self.distribution.nranks()
185 }
186
187 #[cfg(not(feature = "mpi"))]
188 {
189 1
190 }
191 }
192
193 pub fn is_distributed(&self) -> bool {
195 self.nranks() > 1
196 }
197
198 pub fn fragment_partitioning(&self) -> FragmentPartitioning {
200 #[cfg(feature = "mpi")]
201 {
202 match self.distribution.partitioning() {
203 Partitioning::Contiguous => FragmentPartitioning::Contiguous,
204 Partitioning::FileGroups => FragmentPartitioning::RoundRobinFragments,
205 Partitioning::Rows => FragmentPartitioning::StridedRows,
206 }
207 }
208
209 #[cfg(not(feature = "mpi"))]
210 {
211 FragmentPartitioning::Contiguous
212 }
213 }
214}
215
216#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize)]
218pub struct WritePlan {
219 #[cfg(feature = "mpi")]
220 pub distribution: Distribution,
222}
223
224impl From<ReadPlan> for WritePlan {
225 #[cfg_attr(not(feature = "mpi"), allow(unused_variables))]
226 fn from(plan: ReadPlan) -> Self {
227 Self {
228 #[cfg(feature = "mpi")]
229 distribution: plan.distribution,
230 }
231 }
232}
233
234impl WritePlan {
235 pub fn rank(&self) -> usize {
237 #[cfg(feature = "mpi")]
238 {
239 self.distribution.rank()
240 }
241
242 #[cfg(not(feature = "mpi"))]
243 {
244 0
245 }
246 }
247
248 pub fn nranks(&self) -> usize {
250 #[cfg(feature = "mpi")]
251 {
252 self.distribution.nranks()
253 }
254
255 #[cfg(not(feature = "mpi"))]
256 {
257 1
258 }
259 }
260
261 pub fn is_distributed(&self) -> bool {
263 self.nranks() > 1
264 }
265}
266
267#[derive(Clone, Copy, Debug, Serialize, Deserialize)]
269pub enum FragmentPartitioning {
270 Contiguous,
272 RoundRobinFragments,
274 StridedRows,
276}
277
278#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize)]
280pub struct SourceCapabilities {
281 pub exact_len: bool,
283 pub exact_weighted_total: bool,
285 pub random_access: bool,
287 pub deterministic_partitioning: bool,
289 pub predicate_pushdown: bool,
291 pub projection_pushdown: bool,
293 pub streaming: bool,
295}
296
297pub type EventBatchIter = Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>;
299
300pub trait EventSource: Send + Sync {
306 fn schema(&self) -> LadduDataResult<Arc<Schema>>;
313
314 fn capabilities(&self) -> SourceCapabilities {
316 SourceCapabilities::default()
317 }
318
319 fn num_events(&self) -> LadduDataResult<Option<u64>> {
326 Ok(None)
327 }
328
329 fn weighted_total(&self) -> LadduDataResult<Option<f64>> {
336 Ok(None)
337 }
338
339 fn batches(&self, plan: ReadPlan) -> LadduDataResult<EventBatchIter>;
346}
347
348pub trait EventSink: Send {
350 fn retains_batches(&self) -> bool {
352 false
353 }
354
355 fn begin(&mut self, schema: Arc<Schema>, plan: WritePlan) -> LadduDataResult<()>;
362
363 fn write_batch(&mut self, batch: &EventBatch) -> LadduDataResult<()>;
370
371 fn finish(&mut self) -> LadduDataResult<()>;
377
378 fn abort(&mut self) -> LadduDataResult<()> {
388 Ok(())
389 }
390}
391
392#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
394pub(crate) enum SinkState {
395 #[default]
396 Idle,
397 Writing,
398 Failed,
399}
400
401#[derive(Clone, Debug)]
403pub struct DataFragment<K> {
404 pub key: K,
406 pub global_start: u64,
408 pub rows: u64,
410}
411
412#[derive(Clone, Debug)]
414pub struct FragmentRead<K> {
415 pub key: K,
417 pub selection: FragmentSelection,
419}
420
421#[derive(Clone, Copy, Debug)]
423pub enum FragmentSelection {
424 Range {
426 local_start: usize,
428 local_len: usize,
430 },
431 StridedRows {
433 global_start: u64,
435 rows: usize,
437 rank: usize,
439 nranks: usize,
441 },
442}
443
444pub trait FragmentedSource: Send + Sync {
446 type Key: Clone + Send + Sync + 'static;
448
449 fn fragments(&self) -> LadduDataResult<Vec<DataFragment<Self::Key>>>;
455
456 fn read_fragment_range(
463 &self,
464 key: &Self::Key,
465 local_start: usize,
466 local_len: usize,
467 chunk_size: Option<usize>,
468 ) -> LadduDataResult<EventBatchIter>;
469}
470
471pub fn fragmented_batches<S>(source: Arc<S>, plan: ReadPlan) -> LadduDataResult<EventBatchIter>
478where
479 S: FragmentedSource + 'static,
480{
481 let iter = FragmentBatchIter::new(source, plan)?;
482
483 if plan.chunk_size.is_none() {
484 Ok(Box::new(CoalescedBatchIter::new(iter)))
485 } else {
486 Ok(Box::new(iter))
487 }
488}
489
490pub fn plan_fragments<K: Clone>(
497 fragments: &[DataFragment<K>],
498 plan: ReadPlan,
499) -> LadduDataResult<Vec<FragmentRead<K>>> {
500 let total_rows: u64 = fragments.iter().map(|f| f.rows).sum();
501 let rank = plan.rank();
502 let nranks = plan.nranks();
503
504 if nranks == 1 {
505 return fragments
506 .iter()
507 .map(|f| {
508 Ok(FragmentRead {
509 key: f.key.clone(),
510 selection: FragmentSelection::Range {
511 local_start: 0,
512 local_len: usize_from_u64(f.rows)?,
513 },
514 })
515 })
516 .collect();
517 }
518
519 match plan.fragment_partitioning() {
520 FragmentPartitioning::Contiguous => contiguous_plan(fragments, total_rows, rank, nranks),
521 FragmentPartitioning::RoundRobinFragments => {
522 round_robin_fragment_plan(fragments, rank, nranks)
523 }
524 FragmentPartitioning::StridedRows => strided_row_plan(fragments, rank, nranks),
525 }
526}
527
528fn contiguous_plan<K: Clone>(
529 fragments: &[DataFragment<K>],
530 total_rows: u64,
531 rank: usize,
532 nranks: usize,
533) -> LadduDataResult<Vec<FragmentRead<K>>> {
534 let rank_start = total_rows * rank as u64 / nranks as u64;
535 let rank_end = total_rows * (rank as u64 + 1) / nranks as u64;
536
537 let mut out = Vec::new();
538
539 for f in fragments {
540 let frag_start = f.global_start;
541 let frag_end = f.global_start + f.rows;
542
543 let start = rank_start.max(frag_start);
544 let end = rank_end.min(frag_end);
545
546 if start < end {
547 out.push(FragmentRead {
548 key: f.key.clone(),
549 selection: FragmentSelection::Range {
550 local_start: usize_from_u64(start - frag_start)?,
551 local_len: usize_from_u64(end - start)?,
552 },
553 });
554 }
555 }
556
557 Ok(out)
558}
559
560fn round_robin_fragment_plan<K: Clone>(
561 fragments: &[DataFragment<K>],
562 rank: usize,
563 nranks: usize,
564) -> LadduDataResult<Vec<FragmentRead<K>>> {
565 let mut out = Vec::new();
566
567 for (i, f) in fragments.iter().enumerate() {
568 if i % nranks == rank {
569 out.push(FragmentRead {
570 key: f.key.clone(),
571 selection: FragmentSelection::Range {
572 local_start: 0,
573 local_len: usize_from_u64(f.rows)?,
574 },
575 });
576 }
577 }
578
579 Ok(out)
580}
581
582fn strided_row_plan<K: Clone>(
583 fragments: &[DataFragment<K>],
584 rank: usize,
585 nranks: usize,
586) -> LadduDataResult<Vec<FragmentRead<K>>> {
587 fragments
588 .iter()
589 .map(|f| {
590 Ok(FragmentRead {
591 key: f.key.clone(),
592 selection: FragmentSelection::StridedRows {
593 global_start: f.global_start,
594 rows: usize_from_u64(f.rows)?,
595 rank,
596 nranks,
597 },
598 })
599 })
600 .collect()
601}
602
603fn usize_from_u64(value: u64) -> LadduDataResult<usize> {
604 usize::try_from(value).map_err(|_| LadduDataError::InvalidArgument("row count exceeds usize"))
605}
606
607pub(crate) struct FragmentBatchIter<S>
608where
609 S: FragmentedSource,
610{
611 source: Arc<S>,
612 reads: Vec<FragmentRead<S::Key>>,
613 state: FragmentBatchState,
614 chunk_size: Option<usize>,
615}
616
617enum FragmentBatchState {
618 NeedFragment {
619 index: usize,
620 },
621 Reading {
622 iter: EventBatchIter,
623 next_index: usize,
624 },
625 Done,
626}
627
628impl<S> FragmentBatchIter<S>
629where
630 S: FragmentedSource,
631{
632 pub(crate) fn new(source: Arc<S>, plan: ReadPlan) -> LadduDataResult<Self> {
633 let fragments = source.fragments()?;
634 let reads = plan_fragments(&fragments, plan)?;
635
636 Ok(Self {
637 source,
638 reads,
639 state: FragmentBatchState::NeedFragment { index: 0 },
640 chunk_size: plan.chunk_size,
641 })
642 }
643
644 fn open_fragment(&self, read: FragmentRead<S::Key>) -> LadduDataResult<EventBatchIter> {
645 match read.selection {
646 FragmentSelection::Range {
647 local_start,
648 local_len,
649 } => {
650 self.source
651 .read_fragment_range(&read.key, local_start, local_len, self.chunk_size)
652 }
653
654 FragmentSelection::StridedRows {
655 global_start,
656 rows,
657 rank,
658 nranks,
659 } => {
660 let inner = self
661 .source
662 .read_fragment_range(&read.key, 0, rows, self.chunk_size);
663
664 inner.and_then(|iter| {
665 let iter = StridedRowsBatchIter::new(iter, global_start, rank, nranks)?;
666 Ok(Box::new(iter) as EventBatchIter)
667 })
668 }
669 }
670 }
671}
672
673impl<S> Iterator for FragmentBatchIter<S>
674where
675 S: FragmentedSource,
676{
677 type Item = LadduDataResult<EventBatch>;
678
679 fn next(&mut self) -> Option<Self::Item> {
680 loop {
681 let state = std::mem::replace(&mut self.state, FragmentBatchState::Done);
682 match state {
683 FragmentBatchState::Done => return None,
684 FragmentBatchState::Reading {
685 mut iter,
686 next_index,
687 } => match iter.next() {
688 Some(batch) => {
689 self.state = FragmentBatchState::Reading { iter, next_index };
690 return Some(batch);
691 }
692 None => {
693 self.state = FragmentBatchState::NeedFragment { index: next_index };
694 }
695 },
696 FragmentBatchState::NeedFragment { index } => {
697 let Some(read) = self.reads.get(index).cloned() else {
698 self.state = FragmentBatchState::Done;
699 return None;
700 };
701
702 let next_index = index.saturating_add(1);
703 match self.open_fragment(read) {
704 Ok(iter) => {
705 self.state = FragmentBatchState::Reading { iter, next_index };
706 }
707 Err(err) => {
708 self.state = FragmentBatchState::NeedFragment { index: next_index };
711 return Some(Err(err));
712 }
713 }
714 }
715 }
716 }
717 }
718}
719
720pub(crate) struct SliceBatchIter<I> {
721 inner: I,
722 start: usize,
723 end: usize,
724 state: SliceBatchState,
725}
726
727#[derive(Clone, Copy)]
728enum SliceBatchState {
729 Reading { consumed: usize },
730 Done,
731}
732
733impl<I> SliceBatchIter<I> {
734 pub(crate) fn new(inner: I, start: usize, len: usize) -> LadduDataResult<Self> {
735 let end = start
736 .checked_add(len)
737 .ok_or(LadduDataError::InvalidArgument(
738 "slice range overflows usize",
739 ))?;
740 Ok(Self {
741 inner,
742 start,
743 end,
744 state: SliceBatchState::Reading { consumed: 0 },
745 })
746 }
747}
748
749impl<I> Iterator for SliceBatchIter<I>
750where
751 I: Iterator<Item = LadduDataResult<EventBatch>>,
752{
753 type Item = LadduDataResult<EventBatch>;
754
755 fn next(&mut self) -> Option<Self::Item> {
756 loop {
757 let SliceBatchState::Reading { consumed } = self.state else {
758 return None;
759 };
760
761 if consumed >= self.end {
762 self.state = SliceBatchState::Done;
763 return None;
764 }
765
766 let Some(item) = self.inner.next() else {
767 self.state = SliceBatchState::Done;
768 return None;
769 };
770 let batch = match item {
771 Ok(batch) => batch,
772 Err(err) => return Some(Err(err)),
773 };
774
775 let batch_start = consumed;
776 let batch_end = batch_start + batch.len();
777 self.state = SliceBatchState::Reading {
778 consumed: batch_end,
779 };
780
781 let lo = self.start.max(batch_start);
782 let hi = self.end.min(batch_end);
783
784 if lo >= hi {
785 continue;
786 }
787
788 let local_lo = lo - batch_start;
789 let local_hi = hi - batch_start;
790
791 return Some(Ok(batch.slice(local_lo, local_hi)));
792 }
793 }
794}
795
796pub(crate) struct StridedRowsBatchIter<I> {
797 inner: I,
798 global_start: u64,
799 rank: usize,
800 nranks: usize,
801 state: StridedRowsBatchState,
802}
803
804#[derive(Clone, Copy)]
805enum StridedRowsBatchState {
806 Reading { consumed: u64 },
807 Done,
808}
809
810impl<I> StridedRowsBatchIter<I> {
811 pub(crate) fn new(
812 inner: I,
813 global_start: u64,
814 rank: usize,
815 nranks: usize,
816 ) -> LadduDataResult<Self> {
817 if nranks == 0 {
818 return Err(LadduDataError::InvalidArgument("nranks must be nonzero"));
819 }
820 if rank >= nranks {
821 return Err(LadduDataError::InvalidArgument(
822 "rank must be less than nranks",
823 ));
824 }
825 Ok(Self {
826 inner,
827 global_start,
828 rank,
829 nranks,
830 state: StridedRowsBatchState::Reading { consumed: 0 },
831 })
832 }
833}
834
835impl<I> Iterator for StridedRowsBatchIter<I>
836where
837 I: Iterator<Item = LadduDataResult<EventBatch>>,
838{
839 type Item = LadduDataResult<EventBatch>;
840
841 fn next(&mut self) -> Option<Self::Item> {
842 loop {
843 let StridedRowsBatchState::Reading { consumed } = self.state else {
844 return None;
845 };
846
847 let Some(item) = self.inner.next() else {
848 self.state = StridedRowsBatchState::Done;
849 return None;
850 };
851 let batch = match item {
852 Ok(batch) => batch,
853 Err(err) => return Some(Err(err)),
854 };
855
856 let batch_global_start = self.global_start + consumed;
857 self.state = StridedRowsBatchState::Reading {
858 consumed: consumed.saturating_add(batch.len() as u64),
859 };
860
861 let rows: Vec<usize> = (0..batch.len())
862 .filter(|&i| {
863 ((batch_global_start + i as u64) % self.nranks as u64) == self.rank as u64
864 })
865 .collect();
866
867 if rows.is_empty() {
868 continue;
869 }
870
871 return Some(Ok(batch.select(&rows)));
872 }
873 }
874}
875
876pub(crate) struct CoalescedBatchIter<I> {
877 state: CoalescedBatchState<I>,
878}
879
880enum CoalescedBatchState<I> {
881 Reading(I),
882 Done,
883}
884
885impl<I> CoalescedBatchIter<I> {
886 pub(crate) fn new(inner: I) -> Self {
887 Self {
888 state: CoalescedBatchState::Reading(inner),
889 }
890 }
891}
892
893impl<I> Iterator for CoalescedBatchIter<I>
894where
895 I: Iterator<Item = LadduDataResult<EventBatch>>,
896{
897 type Item = LadduDataResult<EventBatch>;
898
899 fn next(&mut self) -> Option<Self::Item> {
900 let CoalescedBatchState::Reading(mut inner) =
901 std::mem::replace(&mut self.state, CoalescedBatchState::Done)
902 else {
903 return None;
904 };
905 let mut batches = Vec::new();
906
907 for batch in &mut inner {
908 match batch {
909 Ok(batch) => batches.push(batch),
910 Err(err) => return Some(Err(err)),
911 }
912 }
913
914 if batches.is_empty() {
915 None
916 } else {
917 Some(EventBatch::concat(&batches))
918 }
919 }
920}
921
922#[cfg(test)]
923mod tests {
924 use super::*;
925 use crate::{
926 data::{EventBatch, EventBatchBuilder},
927 schema::Schema,
928 };
929 use std::path::PathBuf;
930
931 fn v(x: f64) -> RealVec4 {
932 RealVec4 {
933 e: x,
934 px: x,
935 py: x,
936 pz: x,
937 }
938 }
939
940 fn schema() -> Arc<Schema> {
941 Arc::new(Schema::new(["p"], ["id"], true).unwrap())
942 }
943
944 fn batch(start: usize, len: usize) -> EventBatch {
945 let schema = schema();
946 let mut builder = EventBatchBuilder::with_capacity(schema, len);
947
948 for i in start..start + len {
949 builder
950 .push_weighted([v(i as f64)], [i as f64], 100.0 + i as f64)
951 .unwrap();
952 }
953
954 builder.finish().unwrap()
955 }
956
957 fn concat_values(batches: Vec<EventBatch>) -> Vec<f64> {
958 EventBatch::concat(&batches)
959 .unwrap()
960 .scalar_column(0)
961 .to_vec()
962 }
963
964 struct FragmentOpenFailureSource;
965
966 impl FragmentedSource for FragmentOpenFailureSource {
967 type Key = usize;
968
969 fn fragments(&self) -> LadduDataResult<Vec<DataFragment<Self::Key>>> {
970 Ok(vec![
971 DataFragment {
972 key: 0,
973 global_start: 0,
974 rows: 1,
975 },
976 DataFragment {
977 key: 1,
978 global_start: 1,
979 rows: 1,
980 },
981 ])
982 }
983
984 fn read_fragment_range(
985 &self,
986 key: &Self::Key,
987 _local_start: usize,
988 _local_len: usize,
989 _chunk_size: Option<usize>,
990 ) -> LadduDataResult<EventBatchIter> {
991 if *key == 0 {
992 return Err(LadduDataError::Source(
993 "first fragment failed to open".into(),
994 ));
995 }
996
997 Ok(Box::new(vec![Ok(batch(10, 1))].into_iter()))
998 }
999 }
1000
1001 #[test]
1002 fn fragment_open_error_is_an_item_and_later_fragments_remain_readable() {
1003 let mut iter = FragmentBatchIter::new(
1004 Arc::new(FragmentOpenFailureSource),
1005 ReadPlan {
1006 chunk_size: Some(1),
1007 #[cfg(feature = "mpi")]
1008 distribution: Default::default(),
1009 },
1010 )
1011 .unwrap();
1012
1013 assert!(matches!(
1014 iter.next(),
1015 Some(Err(LadduDataError::Source(message))) if message == "first fragment failed to open"
1016 ));
1017 assert_eq!(iter.next().unwrap().unwrap().scalar_column(0), &[10.0]);
1018 assert!(iter.next().is_none());
1019 assert!(iter.next().is_none());
1020 }
1021
1022 #[test]
1023 fn slice_batch_iter_slices_across_batch_boundaries_without_losing_alignment() {
1024 let inner = vec![Ok(batch(0, 3)), Ok(batch(3, 2)), Ok(batch(5, 4))].into_iter();
1025
1026 let out: Vec<EventBatch> = SliceBatchIter::new(inner, 2, 5)
1027 .unwrap()
1028 .map(Result::unwrap)
1029 .collect();
1030
1031 let values = concat_values(out);
1032 assert_eq!(values, vec![2.0, 3.0, 4.0, 5.0, 6.0]);
1033 }
1034
1035 #[test]
1036 fn slice_batch_iter_skips_empty_batches_and_has_a_repeated_terminal_none() {
1037 let inner = vec![Ok(batch(0, 0)), Ok(batch(0, 2))].into_iter();
1038 let mut iter = SliceBatchIter::new(inner, 0, 2).unwrap();
1039
1040 assert_eq!(iter.next().unwrap().unwrap().len(), 2);
1041 assert!(iter.next().is_none());
1042 assert!(iter.next().is_none());
1043 }
1044
1045 #[test]
1046 fn strided_rows_batch_iter_uses_global_row_numbers_across_batches() {
1047 let inner = vec![Ok(batch(0, 4)), Ok(batch(4, 5))].into_iter();
1048
1049 let out: Vec<EventBatch> = StridedRowsBatchIter::new(inner, 1, 1, 3)
1050 .unwrap()
1051 .map(Result::unwrap)
1052 .collect();
1053
1054 assert_eq!(concat_values(out), vec![0.0, 3.0, 6.0]);
1058 }
1059
1060 #[test]
1061 fn coalesced_batch_iter_concatenates_successes_and_propagates_first_error() {
1062 let success_inner = vec![Ok(batch(0, 2)), Ok(batch(2, 3))].into_iter();
1063 let mut success = CoalescedBatchIter::new(success_inner);
1064
1065 let merged = success.next().unwrap().unwrap();
1066 assert_eq!(merged.scalar_column(0), &[0.0, 1.0, 2.0, 3.0, 4.0]);
1067 assert!(success.next().is_none());
1068
1069 let error_inner = vec![
1070 Ok(batch(0, 1)),
1071 Err(LadduDataError::Source("boom".into())),
1072 Ok(batch(1, 1)),
1073 ]
1074 .into_iter();
1075
1076 let mut error_iter = CoalescedBatchIter::new(error_inner);
1077 let err = error_iter.next().unwrap().unwrap_err();
1078
1079 assert!(matches!(err, LadduDataError::Source(msg) if msg == "boom"));
1080 assert!(error_iter.next().is_none());
1081 assert!(error_iter.next().is_none());
1082
1083 let mut empty =
1084 CoalescedBatchIter::new(Vec::<LadduDataResult<EventBatch>>::new().into_iter());
1085 assert!(empty.next().is_none());
1086 assert!(empty.next().is_none());
1087 }
1088
1089 #[test]
1090 fn output_path_resolves_single_file_and_per_rank_names() {
1091 let plan = WritePlan::default();
1092
1093 let single = OutputPath::new(PathBuf::from("events.parquet"))
1094 .resolve(plan, "parquet")
1095 .unwrap();
1096
1097 assert_eq!(single, PathBuf::from("events.parquet"));
1098
1099 let per_rank_with_extension = OutputPath::new(PathBuf::from("events.parquet"))
1100 .with_mode(OutputMode::PerRankFiles)
1101 .resolve(plan, "parquet")
1102 .unwrap();
1103
1104 assert_eq!(
1105 per_rank_with_extension,
1106 PathBuf::from("events.rank00000-of00001.parquet")
1107 );
1108
1109 let per_rank_without_extension = OutputPath::new(PathBuf::from("events"))
1110 .with_mode(OutputMode::PerRankFiles)
1111 .resolve(plan, "root")
1112 .unwrap();
1113
1114 assert_eq!(
1115 per_rank_without_extension,
1116 PathBuf::from("events").join("part-rank00000-of00001.root")
1117 );
1118 }
1119
1120 #[test]
1121 fn plan_fragments_serial_mode_keeps_all_fragments_in_order() {
1122 let fragments = vec![
1123 DataFragment {
1124 key: "a",
1125 global_start: 0,
1126 rows: 2,
1127 },
1128 DataFragment {
1129 key: "b",
1130 global_start: 2,
1131 rows: 3,
1132 },
1133 ];
1134
1135 let reads = plan_fragments(&fragments, ReadPlan::default()).unwrap();
1136
1137 assert_eq!(reads.len(), 2);
1138
1139 match &reads[0].selection {
1140 FragmentSelection::Range {
1141 local_start,
1142 local_len,
1143 } => {
1144 assert_eq!((*local_start, *local_len), (0, 2));
1145 }
1146 _ => panic!("expected range read"),
1147 }
1148
1149 match &reads[1].selection {
1150 FragmentSelection::Range {
1151 local_start,
1152 local_len,
1153 } => {
1154 assert_eq!((*local_start, *local_len), (0, 3));
1155 }
1156 _ => panic!("expected range read"),
1157 }
1158 }
1159
1160 use laddu_physics::vectors::RealVec4;
1161 #[cfg(feature = "mpi")]
1162 use mpi::traits::*;
1163 #[cfg(feature = "mpi")]
1164 use mpi_test::mpi_test;
1165
1166 #[cfg(feature = "mpi")]
1167 fn distributed_plan(
1168 partitioning: Partitioning,
1169 world: &impl mpi::topology::Communicator,
1170 ) -> ReadPlan {
1171 ReadPlan {
1172 chunk_size: None,
1173 distribution: Distribution::from_world(world).with_partitioning(partitioning),
1174 }
1175 }
1176
1177 #[cfg(feature = "mpi")]
1178 fn expected_contiguous_global_range(total_rows: u64, rank: usize, nranks: usize) -> (u64, u64) {
1179 let start = total_rows * rank as u64 / nranks as u64;
1180 let end = total_rows * (rank as u64 + 1) / nranks as u64;
1181 (start, end)
1182 }
1183
1184 #[cfg(feature = "mpi")]
1185 #[mpi_test(np = [2, 3, 4])]
1186 fn mpi_contiguous_plan_assigns_disjoint_ranges_covering_all_rows() {
1187 let universe = mpi::initialize().unwrap();
1188 let world = universe.world();
1189
1190 let rank = world.rank() as usize;
1191 let nranks = world.size() as usize;
1192
1193 let fragments = vec![
1194 DataFragment {
1195 key: "a",
1196 global_start: 0,
1197 rows: 4,
1198 },
1199 DataFragment {
1200 key: "b",
1201 global_start: 4,
1202 rows: 5,
1203 },
1204 DataFragment {
1205 key: "c",
1206 global_start: 9,
1207 rows: 3,
1208 },
1209 ];
1210
1211 let total_rows = fragments.iter().map(|f| f.rows).sum::<u64>();
1212 let plan = distributed_plan(Partitioning::Contiguous, &world);
1213 let reads = plan_fragments(&fragments, plan).unwrap();
1214
1215 let assigned_rows: u64 = reads
1216 .iter()
1217 .map(|read| match read.selection {
1218 FragmentSelection::Range { local_len, .. } => local_len as u64,
1219 FragmentSelection::StridedRows { .. } => panic!("expected range selection"),
1220 })
1221 .sum();
1222
1223 let (expected_start, expected_end) =
1224 expected_contiguous_global_range(total_rows, rank, nranks);
1225
1226 assert_eq!(assigned_rows, expected_end - expected_start);
1227
1228 for read in reads {
1229 let fragment = fragments
1230 .iter()
1231 .find(|fragment| fragment.key == read.key)
1232 .unwrap();
1233
1234 match read.selection {
1235 FragmentSelection::Range {
1236 local_start,
1237 local_len,
1238 } => {
1239 let global_start = fragment.global_start + local_start as u64;
1240 let global_end = global_start + local_len as u64;
1241
1242 assert!(expected_start <= global_start);
1243 assert!(global_end <= expected_end);
1244 assert!(fragment.global_start <= global_start);
1245 assert!(global_end <= fragment.global_start + fragment.rows);
1246 }
1247 FragmentSelection::StridedRows { .. } => panic!("expected range selection"),
1248 }
1249 }
1250 }
1251
1252 #[cfg(feature = "mpi")]
1253 #[mpi_test(np = [2, 3])]
1254 fn mpi_file_group_plan_assigns_fragment_by_rank_round_robin() {
1255 let universe = mpi::initialize().unwrap();
1256 let world = universe.world();
1257
1258 let rank = world.rank() as usize;
1259 let nranks = world.size() as usize;
1260
1261 let fragments = (0..8)
1262 .map(|i| DataFragment {
1263 key: i,
1264 global_start: 10 * i as u64,
1265 rows: 10,
1266 })
1267 .collect::<Vec<_>>();
1268
1269 let plan = distributed_plan(Partitioning::FileGroups, &world);
1270 let reads = plan_fragments(&fragments, plan).unwrap();
1271
1272 let keys = reads.iter().map(|read| read.key).collect::<Vec<_>>();
1273 let expected = (0..8).filter(|i| i % nranks == rank).collect::<Vec<_>>();
1274
1275 assert_eq!(keys, expected);
1276
1277 for read in reads {
1278 match read.selection {
1279 FragmentSelection::Range {
1280 local_start,
1281 local_len,
1282 } => {
1283 assert_eq!(local_start, 0);
1284 assert_eq!(local_len, 10);
1285 }
1286 FragmentSelection::StridedRows { .. } => panic!("expected range selection"),
1287 }
1288 }
1289 }
1290
1291 #[cfg(feature = "mpi")]
1292 #[mpi_test(np = [2, 3, 4])]
1293 fn mpi_rows_plan_assigns_strided_row_selection_with_world_rank() {
1294 let universe = mpi::initialize().unwrap();
1295 let world = universe.world();
1296
1297 let rank = world.rank() as usize;
1298 let nranks = world.size() as usize;
1299
1300 let fragments = vec![
1301 DataFragment {
1302 key: "a",
1303 global_start: 0,
1304 rows: 4,
1305 },
1306 DataFragment {
1307 key: "b",
1308 global_start: 4,
1309 rows: 5,
1310 },
1311 ];
1312
1313 let plan = distributed_plan(Partitioning::Rows, &world);
1314 let reads = plan_fragments(&fragments, plan).unwrap();
1315
1316 assert_eq!(reads.len(), fragments.len());
1317
1318 for (read, fragment) in reads.iter().zip(fragments.iter()) {
1319 assert_eq!(read.key, fragment.key);
1320
1321 match read.selection {
1322 FragmentSelection::StridedRows {
1323 global_start,
1324 rows,
1325 rank: selected_rank,
1326 nranks: selected_nranks,
1327 } => {
1328 assert_eq!(global_start, fragment.global_start);
1329 assert_eq!(rows, fragment.rows as usize);
1330 assert_eq!(selected_rank, rank);
1331 assert_eq!(selected_nranks, nranks);
1332 }
1333 FragmentSelection::Range { .. } => panic!("expected strided selection"),
1334 }
1335 }
1336 }
1337
1338 #[cfg(feature = "mpi")]
1339 #[mpi_test(np = [2, 3])]
1340 fn mpi_read_plan_and_write_plan_reflect_world_distribution() {
1341 let universe = mpi::initialize().unwrap();
1342 let world = universe.world();
1343
1344 let read_plan = ReadPlan {
1345 chunk_size: Some(7),
1346 distribution: Distribution::from_world(&world).with_partitioning(Partitioning::Rows),
1347 };
1348
1349 assert!(read_plan.is_distributed());
1350 assert_eq!(read_plan.rank(), world.rank() as usize);
1351 assert_eq!(read_plan.nranks(), world.size() as usize);
1352
1353 match read_plan.fragment_partitioning() {
1354 FragmentPartitioning::StridedRows => {}
1355 _ => panic!("expected strided row partitioning"),
1356 }
1357
1358 let write_plan = WritePlan::from(read_plan);
1359
1360 assert!(write_plan.is_distributed());
1361 assert_eq!(write_plan.rank(), world.rank() as usize);
1362 assert_eq!(write_plan.nranks(), world.size() as usize);
1363 }
1364}