Skip to main content

moirai_iter/parallel/
sources.rs

1use super::{
2    CollectConsumer, Consumer, IndexedParallelIterator, IntoParallelIterator,
3    IntoParallelRefIterator, ParallelExtend, ParallelIterator,
4};
5use moirai_executor::{SyncTask, global};
6use std::ops::ControlFlow;
7use std::sync::Mutex;
8
9/// Minimum source size for scheduler-backed non-indexed driving.
10///
11/// Smaller sources stay on the existing recursive consumer path so dispatch
12/// overhead does not dominate the work. Larger vector-backed sources split at
13/// each drive level and run one branch through the nesting-safe scheduler
14/// scope; child drives stop at the same threshold.
15pub(super) const PARALLEL_DRIVE_THRESHOLD: usize = 1024;
16
17fn drive_split<I, C, R>(left: I, right: I, left_consumer: C, right_consumer: C) -> R
18where
19    I: ParallelIterator,
20    C: Consumer<I::Item, Result = R> + Send + Sync,
21    R: Send,
22{
23    let left_result = Mutex::new(None);
24    let left_branch = Mutex::new(Some((left, left_consumer)));
25    let right_branch = Mutex::new(Some((right, right_consumer)));
26    let mut right_result = None;
27
28    let scope_result = global().scope::<SyncTask, _>(|scope| {
29        scope.spawn(|_| {
30            let (left, left_consumer) = left_branch
31                .lock()
32                .unwrap_or_else(std::sync::PoisonError::into_inner)
33                .take()
34                .expect("parallel iterator left branch must be claimed once");
35            let result = left_consumer.consume(left);
36            *left_result
37                .lock()
38                .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(result);
39        })?;
40        // Flush before consuming the caller branch so the two branches overlap
41        // whenever scheduler admission succeeds. A refused job is run inline
42        // by the scope, preserving the every-job-runs contract under pressure.
43        scope.flush()?;
44        let (right, right_consumer) = right_branch
45            .lock()
46            .unwrap_or_else(std::sync::PoisonError::into_inner)
47            .take()
48            .expect("parallel iterator right branch must be claimed once");
49        right_result = Some(right_consumer.consume(right));
50        Ok(())
51    });
52
53    // `drive` is an infallible terminal API. If shutdown rejects the scoped
54    // branch, recover the still-unclaimed branch and finish both halves on the
55    // caller rather than dropping work or panicking after a partial drive.
56    if let Err(error) = scope_result {
57        match error {
58            moirai_core::ExecutorError::ShuttingDown
59            | moirai_core::ExecutorError::ResourceExhausted(_) => {
60                let fallback = left_branch
61                    .lock()
62                    .unwrap_or_else(std::sync::PoisonError::into_inner)
63                    .take();
64                if let Some((left, left_consumer)) = fallback {
65                    let result = left_consumer.consume(left);
66                    *left_result
67                        .lock()
68                        .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(result);
69                }
70                if right_result.is_none() {
71                    let fallback = right_branch
72                        .lock()
73                        .unwrap_or_else(std::sync::PoisonError::into_inner)
74                        .take();
75                    if let Some((right, right_consumer)) = fallback {
76                        right_result = Some(right_consumer.consume(right));
77                    }
78                }
79            }
80            error => panic!("moirai global executor: parallel iterator drive: {error}"),
81        }
82    }
83
84    let left_result = left_result
85        .into_inner()
86        .unwrap_or_else(std::sync::PoisonError::into_inner)
87        .expect("parallel iterator left branch must complete");
88    let right_result = right_result.expect("parallel iterator right branch must complete");
89    C::combine(left_result, right_result)
90}
91
92fn move_vec_items_into<T>(source: Vec<T>, target: &mut Vec<T>) {
93    target.clear();
94    let len = source.len();
95    if target.capacity() < len {
96        *target = source;
97        return;
98    }
99
100    // Consuming `source` moves every element without a `Clone` bound and
101    // releases its backing allocation while retaining `target`'s capacity.
102    // The prior `ManuallyDrop` copy leaked the source buffer.
103    target.extend(source);
104}
105
106/// Parallel iterator over a vector.
107pub struct VecParIter<T> {
108    data: Vec<T>,
109}
110
111impl<T> VecParIter<T> {
112    /// Create a parallel iterator over the given vector.
113    pub fn new(data: Vec<T>) -> Self {
114        Self { data }
115    }
116
117    pub(in crate::parallel) fn into_vec(self) -> Vec<T> {
118        self.data
119    }
120}
121
122impl<T: Send + Sync + 'static> ParallelIterator for VecParIter<T> {
123    type Item = T;
124
125    fn seq_items(self) -> Vec<Self::Item> {
126        self.into_vec()
127    }
128
129    fn seq_iter(self) -> impl Iterator<Item = Self::Item> {
130        self.into_vec().into_iter()
131    }
132
133    fn seq_try_fold<Acc, B, FoldFn>(self, init: Acc, fold_fn: FoldFn) -> ControlFlow<B, Acc>
134    where
135        FoldFn: FnMut(Acc, Self::Item) -> ControlFlow<B, Acc>,
136    {
137        self.into_vec().into_iter().try_fold(init, fold_fn)
138    }
139
140    fn drive<C, R>(self, consumer: C) -> R
141    where
142        C: Consumer<Self::Item, Result = R> + Send + Sync,
143        R: Send,
144    {
145        // At or below the dispatch threshold a shard is consumed in one
146        // sequential pass. The superseded shape kept splitting to
147        // single-element shards, which bought no parallelism below the
148        // threshold — the scheduler is only engaged above it — and cost one
149        // consumer split and one combine per element.
150        if self.data.len() <= PARALLEL_DRIVE_THRESHOLD {
151            return consumer.consume(self);
152        }
153
154        // Owned elements have no safe zero-copy split: handing a shard its own
155        // range of a `Vec<T>` without moving the elements needs either raw
156        // pointer reads or an `Option` slot per element, and `Option<T>` is only
157        // niche-packed when `T` has a spare value — for a plain scalar it
158        // doubles the buffer and adds a write per element read back out, which
159        // measured worse than the copy it replaced. Splitting therefore still
160        // copies, but only down to the threshold, so copy traffic is
161        // proportional to `log(len / threshold)` levels rather than `log(len)`.
162        let mut data = self.data;
163        let mid = data.len() / 2;
164        let right_data = data.split_off(mid);
165        let left_data = std::mem::take(&mut data);
166
167        let (left_consumer, right_consumer) = consumer.split_at(left_data.len());
168
169        drive_split(
170            VecParIter::new(left_data),
171            VecParIter::new(right_data),
172            left_consumer,
173            right_consumer,
174        )
175    }
176}
177
178impl<T: Send + Sync + 'static> IndexedParallelIterator for VecParIter<T> {
179    fn len(&self) -> usize {
180        self.data.len()
181    }
182
183    fn collect_into_vec(self, target: &mut Vec<Self::Item>) {
184        move_vec_items_into(self.data, target);
185    }
186}
187
188/// Range parallel iterator.
189pub struct RangeParIter<T> {
190    start: T,
191    end: T,
192}
193
194impl<T> RangeParIter<T>
195where
196    T: Send + Sync + Clone + 'static + PartialOrd + std::ops::Add<Output = T> + From<u8>,
197{
198    /// Create a parallel iterator over the half-open range `start..end`.
199    pub fn new(start: T, end: T) -> Self {
200        Self { start, end }
201    }
202}
203
204impl<T> ParallelIterator for RangeParIter<T>
205where
206    T: Send + Sync + Clone + 'static + PartialOrd + std::ops::Add<Output = T> + From<u8>,
207{
208    type Item = T;
209
210    fn seq_items(self) -> Vec<Self::Item> {
211        let mut items = Vec::new();
212        let mut current = self.start;
213        while current < self.end {
214            items.push(current.clone());
215            current = current + T::from(1u8);
216        }
217        items
218    }
219
220    fn seq_try_fold<Acc, B, FoldFn>(self, init: Acc, mut fold_fn: FoldFn) -> ControlFlow<B, Acc>
221    where
222        FoldFn: FnMut(Acc, Self::Item) -> ControlFlow<B, Acc>,
223    {
224        let mut accumulator = init;
225        let mut current = self.start;
226        while current < self.end {
227            accumulator = fold_fn(accumulator, current.clone())?;
228            current = current + T::from(1u8);
229        }
230        ControlFlow::Continue(accumulator)
231    }
232
233    fn drive<C, R>(self, consumer: C) -> R
234    where
235        C: Consumer<Self::Item, Result = R> + Send + Sync,
236        R: Send,
237    {
238        let mut items = Vec::new();
239        let mut current = self.start;
240        while current < self.end {
241            items.push(current.clone());
242            current = current + T::from(1u8);
243        }
244
245        VecParIter::new(items).drive(consumer)
246    }
247}
248
249impl IndexedParallelIterator for RangeParIter<usize> {
250    fn len(&self) -> usize {
251        self.end.saturating_sub(self.start)
252    }
253
254    fn collect_into_vec(self, target: &mut Vec<Self::Item>) {
255        target.clear();
256        target.extend(self.start..self.end);
257    }
258}
259
260/// Sequential iterator adapter for compatibility.
261pub struct SequentialAdapter<I> {
262    iter: I,
263}
264
265impl<I> SequentialAdapter<I> {
266    pub(super) fn new(iter: I) -> Self {
267        Self { iter }
268    }
269}
270
271impl<I> IntoIterator for SequentialAdapter<I>
272where
273    I: ParallelIterator,
274{
275    type Item = I::Item;
276    type IntoIter = std::vec::IntoIter<I::Item>;
277
278    fn into_iter(self) -> Self::IntoIter {
279        self.iter.seq_items().into_iter()
280    }
281}
282
283/// Adapter that drives a sequential iterator through the parallel-consumer
284/// machinery as a single shard.
285pub struct SequentialIterAdapter<I> {
286    iter: I,
287}
288
289impl<I> SequentialIterAdapter<I> {
290    /// Wrap a sequential iterator so the parallel consumers can drive it.
291    pub fn new(iter: I) -> Self {
292        Self { iter }
293    }
294}
295
296impl<I> ParallelIterator for SequentialIterAdapter<I>
297where
298    I: Iterator + Send,
299    I::Item: Send + Sync + 'static,
300{
301    type Item = I::Item;
302
303    fn seq_items(self) -> Vec<Self::Item> {
304        self.iter.collect()
305    }
306
307    fn seq_try_fold<Acc, B, FoldFn>(self, mut init: Acc, mut fold_fn: FoldFn) -> ControlFlow<B, Acc>
308    where
309        FoldFn: FnMut(Acc, Self::Item) -> ControlFlow<B, Acc>,
310    {
311        let mut iter = self.iter;
312        for item in iter.by_ref() {
313            init = fold_fn(init, item)?;
314        }
315        ControlFlow::Continue(init)
316    }
317
318    fn drive<C, R>(self, consumer: C) -> R
319    where
320        C: Consumer<Self::Item, Result = R> + Send + Sync,
321        R: Send,
322    {
323        let items: Vec<Self::Item> = self.iter.collect();
324        consumer.consume(VecParIter::new(items))
325    }
326}
327
328impl<I> IndexedParallelIterator for SequentialIterAdapter<I>
329where
330    I: ExactSizeIterator + Send,
331    I::Item: Send + Sync + 'static,
332{
333    fn len(&self) -> usize {
334        self.iter.len()
335    }
336
337    fn collect_into_vec(self, target: &mut Vec<Self::Item>) {
338        target.clear();
339        target.extend(self.iter);
340    }
341}
342
343impl<T: Send + Sync + 'static> IntoParallelIterator for Vec<T> {
344    type Item = T;
345    type Iter = VecParIter<T>;
346
347    fn into_par_iter(self) -> Self::Iter {
348        VecParIter::new(self)
349    }
350}
351
352impl<'data, T: Send + Sync + 'data> IntoParallelRefIterator<'data> for Vec<T> {
353    type Item = &'data T;
354    type Iter = VecRefParIter<'data, T>;
355
356    fn par_iter(&'data self) -> Self::Iter {
357        VecRefParIter::new(self)
358    }
359}
360
361/// Parallel iterator over vector references.
362pub struct VecRefParIter<'data, T> {
363    data: &'data Vec<T>,
364}
365
366impl<'data, T> VecRefParIter<'data, T> {
367    fn new(data: &'data Vec<T>) -> Self {
368        Self { data }
369    }
370
371    pub(in crate::parallel) fn into_slice(self) -> &'data [T] {
372        self.data.as_slice()
373    }
374
375    /// Return matching logical positions without materializing borrowed items.
376    pub fn positions<F>(self, predicate: F) -> VecRefPositions<'data, T, F>
377    where
378        F: Fn(&'data T) -> bool + Send + Sync + Clone,
379    {
380        VecRefPositions {
381            data: self.data,
382            predicate,
383        }
384    }
385}
386
387impl<'data, T: Send + Sync + 'data> ParallelIterator for VecRefParIter<'data, T> {
388    type Item = &'data T;
389
390    fn seq_items(self) -> Vec<Self::Item> {
391        self.into_slice().iter().collect()
392    }
393
394    fn seq_iter(self) -> impl Iterator<Item = Self::Item> {
395        self.into_slice().iter()
396    }
397
398    fn seq_try_fold<Acc, B, FoldFn>(self, init: Acc, fold_fn: FoldFn) -> ControlFlow<B, Acc>
399    where
400        FoldFn: FnMut(Acc, Self::Item) -> ControlFlow<B, Acc>,
401    {
402        self.into_slice().iter().try_fold(init, fold_fn)
403    }
404
405    fn drive<C, R>(self, consumer: C) -> R
406    where
407        C: Consumer<Self::Item, Result = R> + Send + Sync,
408        R: Send,
409    {
410        // Drive the backing storage directly. Collecting `Vec<&T>` first cost
411        // one pointer per element before any work started, and every split
412        // below then copied halves of that pointer vector.
413        SliceParIter::new(self.data.as_slice()).drive(consumer)
414    }
415}
416
417/// Borrowed shard addressed as a subslice of one shared slice.
418///
419/// Splitting is `slice::split_at`, so neither an element nor a reference to one
420/// is copied at any depth of the drive recursion.
421struct SliceParIter<'data, T> {
422    data: &'data [T],
423}
424
425impl<'data, T> SliceParIter<'data, T> {
426    fn new(data: &'data [T]) -> Self {
427        Self { data }
428    }
429}
430
431impl<'data, T: Send + Sync + 'data> ParallelIterator for SliceParIter<'data, T> {
432    type Item = &'data T;
433
434    fn seq_items(self) -> Vec<Self::Item> {
435        self.data.iter().collect()
436    }
437
438    fn seq_iter(self) -> impl Iterator<Item = Self::Item> {
439        self.data.iter()
440    }
441
442    fn seq_try_fold<Acc, B, FoldFn>(self, init: Acc, fold_fn: FoldFn) -> ControlFlow<B, Acc>
443    where
444        FoldFn: FnMut(Acc, Self::Item) -> ControlFlow<B, Acc>,
445    {
446        self.data.iter().try_fold(init, fold_fn)
447    }
448
449    fn drive<C, R>(self, consumer: C) -> R
450    where
451        C: Consumer<Self::Item, Result = R> + Send + Sync,
452        R: Send,
453    {
454        if self.data.len() <= PARALLEL_DRIVE_THRESHOLD {
455            return consumer.consume(self);
456        }
457
458        let mid = self.data.len() / 2;
459        let (left_data, right_data) = self.data.split_at(mid);
460        let (left_consumer, right_consumer) = consumer.split_at(left_data.len());
461
462        drive_split(
463            SliceParIter::new(left_data),
464            SliceParIter::new(right_data),
465            left_consumer,
466            right_consumer,
467        )
468    }
469}
470
471impl<'data, T: Send + Sync + 'data> IndexedParallelIterator for VecRefParIter<'data, T> {
472    fn len(&self) -> usize {
473        self.data.len()
474    }
475
476    fn collect_into_vec(self, target: &mut Vec<Self::Item>) {
477        target.clear();
478        target.extend(self.data.iter());
479    }
480}
481
482/// Position stream over borrowed vector storage.
483pub struct VecRefPositions<'data, T, F> {
484    data: &'data Vec<T>,
485    predicate: F,
486}
487
488impl<'data, T, F> ParallelIterator for VecRefPositions<'data, T, F>
489where
490    T: Send + Sync + 'data,
491    F: Fn(&'data T) -> bool + Send + Sync + Clone,
492{
493    type Item = usize;
494
495    fn seq_items(self) -> Vec<Self::Item> {
496        self.data
497            .iter()
498            .enumerate()
499            .filter_map(|(index, item)| (self.predicate)(item).then_some(index))
500            .collect()
501    }
502
503    fn seq_try_fold<Acc, B, FoldFn>(self, init: Acc, fold_fn: FoldFn) -> ControlFlow<B, Acc>
504    where
505        FoldFn: FnMut(Acc, Self::Item) -> ControlFlow<B, Acc>,
506    {
507        self.data
508            .iter()
509            .enumerate()
510            .filter_map(|(index, item)| (self.predicate)(item).then_some(index))
511            .try_fold(init, fold_fn)
512    }
513
514    /// # Why this stays sequential
515    ///
516    /// The yielded item is a logical index, so this stream needs the offset
517    /// documented as absent on the `Enumerate` adapter.
518    fn drive<C, R>(self, consumer: C) -> R
519    where
520        C: Consumer<Self::Item, Result = R> + Send + Sync,
521        R: Send,
522    {
523        consumer.consume(VecParIter::new(self.seq_items()))
524    }
525}
526
527/// A parallel iterator specifically for reference vectors.
528pub struct RefVecParIter<'a, T> {
529    data: Vec<&'a T>,
530}
531
532impl<'a, T> RefVecParIter<'a, T> {
533    fn new(data: Vec<&'a T>) -> Self {
534        Self { data }
535    }
536}
537
538impl<'a, T: Send + Sync> ParallelIterator for RefVecParIter<'a, T> {
539    type Item = &'a T;
540
541    fn seq_items(self) -> Vec<Self::Item> {
542        self.data
543    }
544
545    fn seq_iter(self) -> impl Iterator<Item = Self::Item> {
546        self.data.into_iter()
547    }
548
549    fn seq_try_fold<Acc, B, FoldFn>(self, init: Acc, fold_fn: FoldFn) -> ControlFlow<B, Acc>
550    where
551        FoldFn: FnMut(Acc, Self::Item) -> ControlFlow<B, Acc>,
552    {
553        self.data.into_iter().try_fold(init, fold_fn)
554    }
555
556    fn drive<C, R>(self, consumer: C) -> R
557    where
558        C: Consumer<Self::Item, Result = R> + Send + Sync,
559        R: Send,
560    {
561        // Sequential base case: consumer.consume(RefVecParIter) is safe because
562        // the base consumers terminate on the item stream rather than driving
563        // it again.
564        if self.data.len() <= PARALLEL_DRIVE_THRESHOLD {
565            return consumer.consume(self);
566        }
567
568        let mut data = self.data;
569        let mid = data.len() / 2;
570        let right_data = data.split_off(mid);
571        let left_data = std::mem::take(&mut data);
572
573        let (left_consumer, right_consumer) = consumer.split_at(left_data.len());
574
575        drive_split(
576            RefVecParIter::new(left_data),
577            RefVecParIter::new(right_data),
578            left_consumer,
579            right_consumer,
580        )
581    }
582}
583
584impl<'a, T: Send + Sync> IndexedParallelIterator for RefVecParIter<'a, T> {
585    fn len(&self) -> usize {
586        self.data.len()
587    }
588
589    fn collect_into_vec(self, target: &mut Vec<Self::Item>) {
590        move_vec_items_into(self.data, target);
591    }
592}
593
594impl IntoParallelIterator for std::ops::Range<usize> {
595    type Item = usize;
596    type Iter = RangeParIter<usize>;
597
598    fn into_par_iter(self) -> Self::Iter {
599        RangeParIter::new(self.start, self.end)
600    }
601}
602
603impl<T: Send + Sync> ParallelExtend<T> for Vec<T> {
604    fn par_extend<I>(&mut self, par_iter: I)
605    where
606        I: ParallelIterator<Item = T>,
607    {
608        // Drive through CollectConsumer rather than calling collect::<Vec<_>>(), which
609        // would call par_extend again and create infinite mutual recursion.
610        self.extend(par_iter.drive(CollectConsumer::new()));
611    }
612}