Skip to main content

laddu_data/io/
mod.rs

1use serde::{Deserialize, Serialize};
2use std::fmt::Display;
3use std::sync::Arc;
4
5use crate::{LadduDataError, LadduDataResult, data::EventBatch, schema::Schema};
6
7/// In-memory event sources and sinks.
8pub mod memory;
9mod output;
10/// Parquet event sources and sinks.
11pub mod parquet;
12/// ROOT event sources.
13pub mod root;
14
15mod source;
16
17pub use output::{OutputMode, OutputPath};
18
19pub(crate) use source::{SourceBuild, SourceBuildOptions, build_source};
20
21/// Adds stable operation/resource context to a source failure.
22pub(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
30/// Adds stable operation/resource context to a sink failure.
31pub(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/// Distribution of event I/O across MPI ranks.
63#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize)]
64pub enum Distribution {
65    /// Single-process I/O.
66    #[default]
67    Serial,
68    /// MPI-distributed I/O with explicit rank metadata.
69    Mpi {
70        /// Zero-based rank.
71        rank: usize,
72        /// Number of ranks.
73        nranks: usize,
74        /// Work-partitioning strategy.
75        partitioning: Partitioning,
76    },
77}
78
79#[cfg(feature = "mpi")]
80impl Distribution {
81    /// Creates serial distribution.
82    pub fn serial() -> Self {
83        Self::Serial
84    }
85
86    /// Creates MPI distribution from a communicator.
87    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    /// Returns the current rank.
99    pub fn rank(self) -> usize {
100        match self {
101            Self::Serial => 0,
102            Self::Mpi { rank, .. } => rank,
103        }
104    }
105
106    /// Returns the number of ranks.
107    pub fn nranks(self) -> usize {
108        match self {
109            Self::Serial => 1,
110            Self::Mpi { nranks, .. } => nranks,
111        }
112    }
113
114    /// Returns the partitioning strategy.
115    pub fn partitioning(self) -> Partitioning {
116        match self {
117            Self::Serial => Partitioning::Contiguous,
118            Self::Mpi { partitioning, .. } => partitioning,
119        }
120    }
121
122    /// Returns this distribution with a new partitioning strategy.
123    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/// Strategy for partitioning input rows across ranks.
136#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
137pub enum Partitioning {
138    /// Each rank reads a contiguous global row range.
139    #[default]
140    Contiguous,
141
142    /// Each rank reads whole source fragments, such as files or row groups, round-robin.
143    FileGroups,
144
145    /// Rank r keeps rows where global_row % nranks == r.
146    /// Deterministic, but usually slower.
147    Rows,
148}
149
150/// Options controlling event-source reads.
151#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize)]
152pub struct ReadPlan {
153    /// Optional maximum output batch size.
154    pub chunk_size: Option<usize>,
155
156    #[cfg(feature = "mpi")]
157    /// MPI distribution.
158    pub distribution: Distribution,
159}
160
161impl ReadPlan {
162    /// Creates a serial read plan.
163    pub fn serial() -> Self {
164        Self::default()
165    }
166
167    /// Returns the current rank.
168    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    /// Returns the number of ranks.
181    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    /// Returns whether reads are distributed.
194    pub fn is_distributed(&self) -> bool {
195        self.nranks() > 1
196    }
197
198    /// Returns the low-level fragment-partitioning strategy.
199    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/// Options controlling event-sink writes.
217#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize)]
218pub struct WritePlan {
219    #[cfg(feature = "mpi")]
220    /// MPI distribution.
221    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    /// Returns the current rank.
236    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    /// Returns the number of ranks.
249    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    /// Returns whether writes are distributed.
262    pub fn is_distributed(&self) -> bool {
263        self.nranks() > 1
264    }
265}
266
267/// Low-level assignment of source fragments or rows.
268#[derive(Clone, Copy, Debug, Serialize, Deserialize)]
269pub enum FragmentPartitioning {
270    /// Contiguous global row ranges.
271    Contiguous,
272    /// Whole fragments assigned round-robin.
273    RoundRobinFragments,
274    /// Individual rows assigned by global index modulo rank count.
275    StridedRows,
276}
277
278/// Optional performance and planning capabilities of an [`EventSource`].
279#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize)]
280pub struct SourceCapabilities {
281    /// Exact event count is available cheaply.
282    pub exact_len: bool,
283    /// Exact weighted total is available cheaply.
284    pub exact_weighted_total: bool,
285    /// Arbitrary row ranges can be read.
286    pub random_access: bool,
287    /// Distributed row assignment is deterministic.
288    pub deterministic_partitioning: bool,
289    /// Filters can be pushed into the source.
290    pub predicate_pushdown: bool,
291    /// Column projection can be pushed into the source.
292    pub projection_pushdown: bool,
293    /// Batches can be streamed without full materialization.
294    pub streaming: bool,
295}
296
297/// Sendable iterator of fallible event batches.
298pub type EventBatchIter = Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>;
299
300/// Thread-safe producer of schema-compatible event batches.
301///
302/// Repeated reads with the same plan must replay identical ordered rows and
303/// contents. Batch boundaries may vary. This guarantees that separately
304/// collected row-data columns, expression results, and weights stay aligned.
305pub trait EventSource: Send + Sync {
306    /// Returns the source schema.
307    ///
308    /// # Errors
309    ///
310    /// Returns [`LadduDataError`] when source metadata cannot be read or
311    /// interpreted as a logical schema.
312    fn schema(&self) -> LadduDataResult<Arc<Schema>>;
313
314    /// Returns optional source capabilities.
315    fn capabilities(&self) -> SourceCapabilities {
316        SourceCapabilities::default()
317    }
318
319    /// Returns the exact event count when cheaply available.
320    ///
321    /// # Errors
322    ///
323    /// Returns [`LadduDataError`] when the source cannot read the metadata
324    /// needed to determine its event count.
325    fn num_events(&self) -> LadduDataResult<Option<u64>> {
326        Ok(None)
327    }
328
329    /// Returns the exact sum of event weights when cheaply available.
330    ///
331    /// # Errors
332    ///
333    /// Returns [`LadduDataError`] when source weights or their metadata cannot
334    /// be read.
335    fn weighted_total(&self) -> LadduDataResult<Option<f64>> {
336        Ok(None)
337    }
338
339    /// Opens a batch iterator using `plan`.
340    ///
341    /// # Errors
342    ///
343    /// Returns [`LadduDataError`] when `plan` is invalid or the source cannot
344    /// initialize the requested read.
345    fn batches(&self, plan: ReadPlan) -> LadduDataResult<EventBatchIter>;
346}
347
348/// Consumer of schema-compatible event batches.
349pub trait EventSink: Send {
350    /// Returns whether written batches remain resident in memory.
351    fn retains_batches(&self) -> bool {
352        false
353    }
354
355    /// Begins a write operation.
356    ///
357    /// # Errors
358    ///
359    /// Returns [`LadduDataError`] when the plan or schema is unsupported or
360    /// output initialization fails.
361    fn begin(&mut self, schema: Arc<Schema>, plan: WritePlan) -> LadduDataResult<()>;
362
363    /// Writes one batch.
364    ///
365    /// # Errors
366    ///
367    /// Returns [`LadduDataError`] when the batch schema is incompatible or the
368    /// output cannot be written.
369    fn write_batch(&mut self, batch: &EventBatch) -> LadduDataResult<()>;
370
371    /// Finishes and flushes the write operation.
372    ///
373    /// # Errors
374    ///
375    /// Returns [`LadduDataError`] when buffered output cannot be finalized.
376    fn finish(&mut self) -> LadduDataResult<()>;
377
378    /// Aborts the current write operation and releases any backend resources.
379    ///
380    /// Aborting is idempotent: calling it while the sink is idle is a no-op.
381    /// Backends leave any partially written output in place; callers should
382    /// treat such files as incomplete.
383    ///
384    /// # Errors
385    ///
386    /// Returns [`LadduDataError`] when backend resources cannot be released.
387    fn abort(&mut self) -> LadduDataResult<()> {
388        Ok(())
389    }
390}
391
392/// Internal lifecycle shared by concrete event sinks.
393#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
394pub(crate) enum SinkState {
395    #[default]
396    Idle,
397    Writing,
398    Failed,
399}
400
401/// Metadata describing one addressable source fragment.
402#[derive(Clone, Debug)]
403pub struct DataFragment<K> {
404    /// Source-specific fragment key.
405    pub key: K,
406    /// Global row offset.
407    pub global_start: u64,
408    /// Number of rows.
409    pub rows: u64,
410}
411
412/// Planned read of one source fragment.
413#[derive(Clone, Debug)]
414pub struct FragmentRead<K> {
415    /// Source-specific fragment key.
416    pub key: K,
417    /// Rows selected from the fragment.
418    pub selection: FragmentSelection,
419}
420
421/// Row selection within one source fragment.
422#[derive(Clone, Copy, Debug)]
423pub enum FragmentSelection {
424    /// Contiguous local row range.
425    Range {
426        /// First local row.
427        local_start: usize,
428        /// Number of local rows.
429        local_len: usize,
430    },
431    /// Rows assigned by global index modulo rank count.
432    StridedRows {
433        /// Global offset of the fragment.
434        global_start: u64,
435        /// Number of rows in the fragment.
436        rows: usize,
437        /// Current rank.
438        rank: usize,
439        /// Number of ranks.
440        nranks: usize,
441    },
442}
443
444/// Event source composed of independently addressable fragments.
445pub trait FragmentedSource: Send + Sync {
446    /// Source-specific fragment key.
447    type Key: Clone + Send + Sync + 'static;
448
449    /// Lists all fragments in global row order.
450    ///
451    /// # Errors
452    ///
453    /// Returns [`LadduDataError`] when fragment metadata cannot be read.
454    fn fragments(&self) -> LadduDataResult<Vec<DataFragment<Self::Key>>>;
455
456    /// Reads a contiguous range within one fragment.
457    ///
458    /// # Errors
459    ///
460    /// Returns [`LadduDataError`] when the key, range, or chunk size is invalid
461    /// or fragment data cannot be read.
462    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
471/// Creates a planned batch iterator for a fragmented source.
472///
473/// # Errors
474///
475/// Returns [`LadduDataError`] when the read plan is invalid, fragment metadata
476/// cannot be loaded, or the iterator cannot be initialized.
477pub 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
490/// Assigns source fragments or rows according to a read plan.
491///
492/// # Errors
493///
494/// Returns [`LadduDataError`] when rank settings are invalid or fragment sizes
495/// cannot be represented on this platform.
496pub 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                            // Opening one fragment is an item-level failure. Keep
709                            // the next fragment available for a later call.
710                            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        // Global rows are 1..=9 because global_start = 1.
1055        // Rank 1 of 3 keeps global rows 1, 4, 7.
1056        // Those correspond to local scalar ids 0, 3, 6.
1057        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}