Skip to main content

moirai_iter/parallel/
sources.rs

1use super::{
2    CollectConsumer, Consumer, IndexedParallelIterator, IntoParallelIterator,
3    IntoParallelRefIterator, ParallelExtend, ParallelIterator,
4};
5use moirai_executor::{global, SyncTask};
6use std::sync::Mutex;
7
8/// Minimum source size for scheduler-backed non-indexed driving.
9///
10/// Smaller sources stay on the existing recursive consumer path so dispatch
11/// overhead does not dominate the work. Larger vector-backed sources split at
12/// each drive level and run one branch through the nesting-safe scheduler
13/// scope; child drives stop at the same threshold.
14const PARALLEL_DRIVE_THRESHOLD: usize = 1024;
15
16fn drive_split<I, C, R>(left: I, right: I, left_consumer: C, right_consumer: C) -> R
17where
18    I: ParallelIterator,
19    C: Consumer<I::Item, Result = R> + Send + Sync,
20    R: Send,
21{
22    let left_result = Mutex::new(None);
23    let left_branch = Mutex::new(Some((left, left_consumer)));
24    let right_branch = Mutex::new(Some((right, right_consumer)));
25    let mut right_result = None;
26
27    let scope_result = global().scope::<SyncTask, _>(|scope| {
28        scope.spawn(|_| {
29            let (left, left_consumer) = left_branch
30                .lock()
31                .unwrap_or_else(std::sync::PoisonError::into_inner)
32                .take()
33                .expect("parallel iterator left branch must be claimed once");
34            let result = left_consumer.consume(left);
35            *left_result
36                .lock()
37                .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(result);
38        })?;
39        // Flush before consuming the caller branch so the two branches overlap
40        // whenever scheduler admission succeeds. A refused job is run inline
41        // by the scope, preserving the every-job-runs contract under pressure.
42        scope.flush()?;
43        let (right, right_consumer) = right_branch
44            .lock()
45            .unwrap_or_else(std::sync::PoisonError::into_inner)
46            .take()
47            .expect("parallel iterator right branch must be claimed once");
48        right_result = Some(right_consumer.consume(right));
49        Ok(())
50    });
51
52    // `drive` is an infallible terminal API. If shutdown rejects the scoped
53    // branch, recover the still-unclaimed branch and finish both halves on the
54    // caller rather than dropping work or panicking after a partial drive.
55    if let Err(error) = scope_result {
56        match error {
57            moirai_core::ExecutorError::ShuttingDown
58            | moirai_core::ExecutorError::ResourceExhausted(_) => {
59                let fallback = left_branch
60                    .lock()
61                    .unwrap_or_else(std::sync::PoisonError::into_inner)
62                    .take();
63                if let Some((left, left_consumer)) = fallback {
64                    let result = left_consumer.consume(left);
65                    *left_result
66                        .lock()
67                        .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(result);
68                }
69                if right_result.is_none() {
70                    let fallback = right_branch
71                        .lock()
72                        .unwrap_or_else(std::sync::PoisonError::into_inner)
73                        .take();
74                    if let Some((right, right_consumer)) = fallback {
75                        right_result = Some(right_consumer.consume(right));
76                    }
77                }
78            }
79            error => panic!("moirai global executor: parallel iterator drive: {error}"),
80        }
81    }
82
83    let left_result = left_result
84        .into_inner()
85        .unwrap_or_else(std::sync::PoisonError::into_inner)
86        .expect("parallel iterator left branch must complete");
87    let right_result = right_result.expect("parallel iterator right branch must complete");
88    C::combine(left_result, right_result)
89}
90
91fn move_vec_items_into<T>(source: Vec<T>, target: &mut Vec<T>) {
92    target.clear();
93    let len = source.len();
94    if target.capacity() < len {
95        *target = source;
96        return;
97    }
98
99    // Consuming `source` moves every element without a `Clone` bound and
100    // releases its backing allocation while retaining `target`'s capacity.
101    // The prior `ManuallyDrop` copy leaked the source buffer.
102    target.extend(source);
103}
104
105/// Parallel iterator over a vector.
106pub struct VecParIter<T> {
107    data: Vec<T>,
108}
109
110impl<T> VecParIter<T> {
111    /// Create a parallel iterator over the given vector.
112    pub fn new(data: Vec<T>) -> Self {
113        Self { data }
114    }
115
116    pub(in crate::parallel) fn into_vec(self) -> Vec<T> {
117        self.data
118    }
119}
120
121impl<T: Send + Sync + 'static> ParallelIterator for VecParIter<T> {
122    type Item = T;
123
124    fn seq_items(self) -> Vec<Self::Item> {
125        self.data
126    }
127
128    fn drive<C, R>(mut self, consumer: C) -> R
129    where
130        C: Consumer<Self::Item, Result = R> + Send + Sync,
131        R: Send,
132    {
133        if self.data.len() <= 1 {
134            return consumer.consume(SequentialIterAdapter::new(self.data.into_iter()));
135        }
136
137        let total_len = self.data.len();
138        let mid = total_len / 2;
139        let right_data = self.data.split_off(mid);
140        let left_data = std::mem::take(&mut self.data);
141
142        let (left_consumer, right_consumer) = consumer.split_at(left_data.len());
143
144        if total_len > PARALLEL_DRIVE_THRESHOLD {
145            return drive_split(
146                VecParIter::new(left_data),
147                VecParIter::new(right_data),
148                left_consumer,
149                right_consumer,
150            );
151        }
152
153        let left_result = left_consumer.consume(VecParIter::new(left_data));
154        let right_result = right_consumer.consume(VecParIter::new(right_data));
155
156        C::combine(left_result, right_result)
157    }
158}
159
160impl<T: Send + Sync + 'static> IndexedParallelIterator for VecParIter<T> {
161    fn len(&self) -> usize {
162        self.data.len()
163    }
164
165    fn collect_into_vec(self, target: &mut Vec<Self::Item>) {
166        move_vec_items_into(self.data, target);
167    }
168}
169
170/// Range parallel iterator.
171pub struct RangeParIter<T> {
172    start: T,
173    end: T,
174}
175
176impl<T> RangeParIter<T>
177where
178    T: Send + Sync + Clone + 'static + PartialOrd + std::ops::Add<Output = T> + From<u8>,
179{
180    /// Create a parallel iterator over the half-open range `start..end`.
181    pub fn new(start: T, end: T) -> Self {
182        Self { start, end }
183    }
184}
185
186impl<T> ParallelIterator for RangeParIter<T>
187where
188    T: Send + Sync + Clone + 'static + PartialOrd + std::ops::Add<Output = T> + From<u8>,
189{
190    type Item = T;
191
192    fn seq_items(self) -> Vec<Self::Item> {
193        let mut items = Vec::new();
194        let mut current = self.start;
195        while current < self.end {
196            items.push(current.clone());
197            current = current + T::from(1u8);
198        }
199        items
200    }
201
202    fn drive<C, R>(self, consumer: C) -> R
203    where
204        C: Consumer<Self::Item, Result = R> + Send + Sync,
205        R: Send,
206    {
207        let mut items = Vec::new();
208        let mut current = self.start;
209        while current < self.end {
210            items.push(current.clone());
211            current = current + T::from(1u8);
212        }
213
214        VecParIter::new(items).drive(consumer)
215    }
216}
217
218impl IndexedParallelIterator for RangeParIter<usize> {
219    fn len(&self) -> usize {
220        self.end.saturating_sub(self.start)
221    }
222
223    fn collect_into_vec(self, target: &mut Vec<Self::Item>) {
224        target.clear();
225        target.extend(self.start..self.end);
226    }
227}
228
229/// Sequential iterator adapter for compatibility.
230pub struct SequentialAdapter<I> {
231    iter: I,
232}
233
234impl<I> SequentialAdapter<I> {
235    pub(super) fn new(iter: I) -> Self {
236        Self { iter }
237    }
238}
239
240/// Adapter that drives a sequential iterator through the parallel-consumer
241/// machinery as a single shard.
242pub struct SequentialIterAdapter<I> {
243    iter: I,
244}
245
246impl<I> SequentialIterAdapter<I> {
247    pub(super) fn new(iter: I) -> Self {
248        Self { iter }
249    }
250}
251
252impl<I> ParallelIterator for SequentialIterAdapter<I>
253where
254    I: Iterator + Send,
255    I::Item: Send + Sync + 'static,
256{
257    type Item = I::Item;
258
259    fn seq_items(self) -> Vec<Self::Item> {
260        self.iter.collect()
261    }
262
263    fn drive<C, R>(self, consumer: C) -> R
264    where
265        C: Consumer<Self::Item, Result = R> + Send + Sync,
266        R: Send,
267    {
268        let items: Vec<Self::Item> = self.iter.collect();
269        consumer.consume(VecParIter::new(items))
270    }
271}
272
273impl<I> IndexedParallelIterator for SequentialIterAdapter<I>
274where
275    I: ExactSizeIterator + Send,
276    I::Item: Send + Sync + 'static,
277{
278    fn len(&self) -> usize {
279        self.iter.len()
280    }
281
282    fn collect_into_vec(self, target: &mut Vec<Self::Item>) {
283        target.clear();
284        target.extend(self.iter);
285    }
286}
287
288impl<T: Send + Sync + 'static> IntoParallelIterator for Vec<T> {
289    type Item = T;
290    type Iter = VecParIter<T>;
291
292    fn into_par_iter(self) -> Self::Iter {
293        VecParIter::new(self)
294    }
295}
296
297impl<'data, T: Send + Sync + 'data> IntoParallelRefIterator<'data> for Vec<T> {
298    type Item = &'data T;
299    type Iter = VecRefParIter<'data, T>;
300
301    fn par_iter(&'data self) -> Self::Iter {
302        VecRefParIter::new(self)
303    }
304}
305
306/// Parallel iterator over vector references.
307pub struct VecRefParIter<'data, T> {
308    data: &'data Vec<T>,
309}
310
311impl<'data, T> VecRefParIter<'data, T> {
312    fn new(data: &'data Vec<T>) -> Self {
313        Self { data }
314    }
315
316    pub(in crate::parallel) fn into_slice(self) -> &'data [T] {
317        self.data.as_slice()
318    }
319
320    /// Return matching logical positions without materializing borrowed items.
321    pub fn positions<F>(self, predicate: F) -> VecRefPositions<'data, T, F>
322    where
323        F: Fn(&'data T) -> bool + Send + Sync + Clone,
324    {
325        VecRefPositions {
326            data: self.data,
327            predicate,
328        }
329    }
330}
331
332impl<'data, T: Send + Sync + 'data> ParallelIterator for VecRefParIter<'data, T> {
333    type Item = &'data T;
334
335    fn seq_items(self) -> Vec<Self::Item> {
336        self.data.iter().collect()
337    }
338
339    fn drive<C, R>(self, consumer: C) -> R
340    where
341        C: Consumer<Self::Item, Result = R> + Send + Sync,
342        R: Send,
343    {
344        let refs: Vec<&'data T> = self.data.iter().collect();
345        consumer.consume(RefVecParIter::new(refs))
346    }
347}
348
349impl<'data, T: Send + Sync + 'data> IndexedParallelIterator for VecRefParIter<'data, T> {
350    fn len(&self) -> usize {
351        self.data.len()
352    }
353
354    fn collect_into_vec(self, target: &mut Vec<Self::Item>) {
355        target.clear();
356        target.extend(self.data.iter());
357    }
358}
359
360/// Position stream over borrowed vector storage.
361pub struct VecRefPositions<'data, T, F> {
362    data: &'data Vec<T>,
363    predicate: F,
364}
365
366impl<'data, T, F> ParallelIterator for VecRefPositions<'data, T, F>
367where
368    T: Send + Sync + 'data,
369    F: Fn(&'data T) -> bool + Send + Sync + Clone,
370{
371    type Item = usize;
372
373    fn seq_items(self) -> Vec<Self::Item> {
374        self.data
375            .iter()
376            .enumerate()
377            .filter_map(|(index, item)| (self.predicate)(item).then_some(index))
378            .collect()
379    }
380
381    fn drive<C, R>(self, consumer: C) -> R
382    where
383        C: Consumer<Self::Item, Result = R> + Send + Sync,
384        R: Send,
385    {
386        consumer.consume(VecParIter::new(self.seq_items()))
387    }
388}
389
390/// A parallel iterator specifically for reference vectors.
391pub struct RefVecParIter<'a, T> {
392    data: Vec<&'a T>,
393}
394
395impl<'a, T> RefVecParIter<'a, T> {
396    fn new(data: Vec<&'a T>) -> Self {
397        Self { data }
398    }
399}
400
401impl<'a, T: Send + Sync> ParallelIterator for RefVecParIter<'a, T> {
402    type Item = &'a T;
403
404    fn seq_items(self) -> Vec<Self::Item> {
405        self.data
406    }
407
408    fn drive<C, R>(mut self, consumer: C) -> R
409    where
410        C: Consumer<Self::Item, Result = R> + Send + Sync,
411        R: Send,
412    {
413        // Sequential base case: consumer.consume(RefVecParIter) is safe because
414        // CollectConsumer::consume now calls seq_items(), terminating the chain.
415        if self.data.len() <= 1 {
416            return consumer.consume(RefVecParIter::new(self.data));
417        }
418
419        let total_len = self.data.len();
420        let mid = total_len / 2;
421        let right_data = self.data.split_off(mid);
422        let left_data = std::mem::take(&mut self.data);
423
424        let (left_consumer, right_consumer) = consumer.split_at(left_data.len());
425
426        if total_len > PARALLEL_DRIVE_THRESHOLD {
427            return drive_split(
428                RefVecParIter::new(left_data),
429                RefVecParIter::new(right_data),
430                left_consumer,
431                right_consumer,
432            );
433        }
434
435        let left_result = left_consumer.consume(RefVecParIter::new(left_data));
436        let right_result = right_consumer.consume(RefVecParIter::new(right_data));
437
438        C::combine(left_result, right_result)
439    }
440}
441
442impl<'a, T: Send + Sync> IndexedParallelIterator for RefVecParIter<'a, T> {
443    fn len(&self) -> usize {
444        self.data.len()
445    }
446
447    fn collect_into_vec(self, target: &mut Vec<Self::Item>) {
448        move_vec_items_into(self.data, target);
449    }
450}
451
452impl IntoParallelIterator for std::ops::Range<usize> {
453    type Item = usize;
454    type Iter = RangeParIter<usize>;
455
456    fn into_par_iter(self) -> Self::Iter {
457        RangeParIter::new(self.start, self.end)
458    }
459}
460
461impl<T: Send + Sync> ParallelExtend<T> for Vec<T> {
462    fn par_extend<I>(&mut self, par_iter: I)
463    where
464        I: ParallelIterator<Item = T>,
465    {
466        // Drive through CollectConsumer rather than calling collect::<Vec<_>>(), which
467        // would call par_extend again and create infinite mutual recursion.
468        self.extend(par_iter.drive(CollectConsumer::new()));
469    }
470}