use super::{
CollectConsumer, Consumer, IndexedParallelIterator, IntoParallelIterator,
IntoParallelRefIterator, ParallelExtend, ParallelIterator,
};
use moirai_executor::{global, SyncTask};
use std::sync::Mutex;
const PARALLEL_DRIVE_THRESHOLD: usize = 1024;
fn drive_split<I, C, R>(left: I, right: I, left_consumer: C, right_consumer: C) -> R
where
I: ParallelIterator,
C: Consumer<I::Item, Result = R> + Send + Sync,
R: Send,
{
let left_result = Mutex::new(None);
let left_branch = Mutex::new(Some((left, left_consumer)));
let right_branch = Mutex::new(Some((right, right_consumer)));
let mut right_result = None;
let scope_result = global().scope::<SyncTask, _>(|scope| {
scope.spawn(|_| {
let (left, left_consumer) = left_branch
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.take()
.expect("parallel iterator left branch must be claimed once");
let result = left_consumer.consume(left);
*left_result
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner) = Some(result);
})?;
scope.flush()?;
let (right, right_consumer) = right_branch
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.take()
.expect("parallel iterator right branch must be claimed once");
right_result = Some(right_consumer.consume(right));
Ok(())
});
if let Err(error) = scope_result {
match error {
moirai_core::ExecutorError::ShuttingDown
| moirai_core::ExecutorError::ResourceExhausted(_) => {
let fallback = left_branch
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.take();
if let Some((left, left_consumer)) = fallback {
let result = left_consumer.consume(left);
*left_result
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner) = Some(result);
}
if right_result.is_none() {
let fallback = right_branch
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.take();
if let Some((right, right_consumer)) = fallback {
right_result = Some(right_consumer.consume(right));
}
}
}
error => panic!("moirai global executor: parallel iterator drive: {error}"),
}
}
let left_result = left_result
.into_inner()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.expect("parallel iterator left branch must complete");
let right_result = right_result.expect("parallel iterator right branch must complete");
C::combine(left_result, right_result)
}
fn move_vec_items_into<T>(source: Vec<T>, target: &mut Vec<T>) {
target.clear();
let len = source.len();
if target.capacity() < len {
*target = source;
return;
}
target.extend(source);
}
pub struct VecParIter<T> {
data: Vec<T>,
}
impl<T> VecParIter<T> {
pub fn new(data: Vec<T>) -> Self {
Self { data }
}
pub(in crate::parallel) fn into_vec(self) -> Vec<T> {
self.data
}
}
impl<T: Send + Sync + 'static> ParallelIterator for VecParIter<T> {
type Item = T;
fn seq_items(self) -> Vec<Self::Item> {
self.data
}
fn drive<C, R>(mut self, consumer: C) -> R
where
C: Consumer<Self::Item, Result = R> + Send + Sync,
R: Send,
{
if self.data.len() <= 1 {
return consumer.consume(SequentialIterAdapter::new(self.data.into_iter()));
}
let total_len = self.data.len();
let mid = total_len / 2;
let right_data = self.data.split_off(mid);
let left_data = std::mem::take(&mut self.data);
let (left_consumer, right_consumer) = consumer.split_at(left_data.len());
if total_len > PARALLEL_DRIVE_THRESHOLD {
return drive_split(
VecParIter::new(left_data),
VecParIter::new(right_data),
left_consumer,
right_consumer,
);
}
let left_result = left_consumer.consume(VecParIter::new(left_data));
let right_result = right_consumer.consume(VecParIter::new(right_data));
C::combine(left_result, right_result)
}
}
impl<T: Send + Sync + 'static> IndexedParallelIterator for VecParIter<T> {
fn len(&self) -> usize {
self.data.len()
}
fn collect_into_vec(self, target: &mut Vec<Self::Item>) {
move_vec_items_into(self.data, target);
}
}
pub struct RangeParIter<T> {
start: T,
end: T,
}
impl<T> RangeParIter<T>
where
T: Send + Sync + Clone + 'static + PartialOrd + std::ops::Add<Output = T> + From<u8>,
{
pub fn new(start: T, end: T) -> Self {
Self { start, end }
}
}
impl<T> ParallelIterator for RangeParIter<T>
where
T: Send + Sync + Clone + 'static + PartialOrd + std::ops::Add<Output = T> + From<u8>,
{
type Item = T;
fn seq_items(self) -> Vec<Self::Item> {
let mut items = Vec::new();
let mut current = self.start;
while current < self.end {
items.push(current.clone());
current = current + T::from(1u8);
}
items
}
fn drive<C, R>(self, consumer: C) -> R
where
C: Consumer<Self::Item, Result = R> + Send + Sync,
R: Send,
{
let mut items = Vec::new();
let mut current = self.start;
while current < self.end {
items.push(current.clone());
current = current + T::from(1u8);
}
VecParIter::new(items).drive(consumer)
}
}
impl IndexedParallelIterator for RangeParIter<usize> {
fn len(&self) -> usize {
self.end.saturating_sub(self.start)
}
fn collect_into_vec(self, target: &mut Vec<Self::Item>) {
target.clear();
target.extend(self.start..self.end);
}
}
pub struct SequentialAdapter<I> {
iter: I,
}
impl<I> SequentialAdapter<I> {
pub(super) fn new(iter: I) -> Self {
Self { iter }
}
}
pub struct SequentialIterAdapter<I> {
iter: I,
}
impl<I> SequentialIterAdapter<I> {
pub(super) fn new(iter: I) -> Self {
Self { iter }
}
}
impl<I> ParallelIterator for SequentialIterAdapter<I>
where
I: Iterator + Send,
I::Item: Send + Sync + 'static,
{
type Item = I::Item;
fn seq_items(self) -> Vec<Self::Item> {
self.iter.collect()
}
fn drive<C, R>(self, consumer: C) -> R
where
C: Consumer<Self::Item, Result = R> + Send + Sync,
R: Send,
{
let items: Vec<Self::Item> = self.iter.collect();
consumer.consume(VecParIter::new(items))
}
}
impl<I> IndexedParallelIterator for SequentialIterAdapter<I>
where
I: ExactSizeIterator + Send,
I::Item: Send + Sync + 'static,
{
fn len(&self) -> usize {
self.iter.len()
}
fn collect_into_vec(self, target: &mut Vec<Self::Item>) {
target.clear();
target.extend(self.iter);
}
}
impl<T: Send + Sync + 'static> IntoParallelIterator for Vec<T> {
type Item = T;
type Iter = VecParIter<T>;
fn into_par_iter(self) -> Self::Iter {
VecParIter::new(self)
}
}
impl<'data, T: Send + Sync + 'data> IntoParallelRefIterator<'data> for Vec<T> {
type Item = &'data T;
type Iter = VecRefParIter<'data, T>;
fn par_iter(&'data self) -> Self::Iter {
VecRefParIter::new(self)
}
}
pub struct VecRefParIter<'data, T> {
data: &'data Vec<T>,
}
impl<'data, T> VecRefParIter<'data, T> {
fn new(data: &'data Vec<T>) -> Self {
Self { data }
}
pub(in crate::parallel) fn into_slice(self) -> &'data [T] {
self.data.as_slice()
}
pub fn positions<F>(self, predicate: F) -> VecRefPositions<'data, T, F>
where
F: Fn(&'data T) -> bool + Send + Sync + Clone,
{
VecRefPositions {
data: self.data,
predicate,
}
}
}
impl<'data, T: Send + Sync + 'data> ParallelIterator for VecRefParIter<'data, T> {
type Item = &'data T;
fn seq_items(self) -> Vec<Self::Item> {
self.data.iter().collect()
}
fn drive<C, R>(self, consumer: C) -> R
where
C: Consumer<Self::Item, Result = R> + Send + Sync,
R: Send,
{
let refs: Vec<&'data T> = self.data.iter().collect();
consumer.consume(RefVecParIter::new(refs))
}
}
impl<'data, T: Send + Sync + 'data> IndexedParallelIterator for VecRefParIter<'data, T> {
fn len(&self) -> usize {
self.data.len()
}
fn collect_into_vec(self, target: &mut Vec<Self::Item>) {
target.clear();
target.extend(self.data.iter());
}
}
pub struct VecRefPositions<'data, T, F> {
data: &'data Vec<T>,
predicate: F,
}
impl<'data, T, F> ParallelIterator for VecRefPositions<'data, T, F>
where
T: Send + Sync + 'data,
F: Fn(&'data T) -> bool + Send + Sync + Clone,
{
type Item = usize;
fn seq_items(self) -> Vec<Self::Item> {
self.data
.iter()
.enumerate()
.filter_map(|(index, item)| (self.predicate)(item).then_some(index))
.collect()
}
fn drive<C, R>(self, consumer: C) -> R
where
C: Consumer<Self::Item, Result = R> + Send + Sync,
R: Send,
{
consumer.consume(VecParIter::new(self.seq_items()))
}
}
pub struct RefVecParIter<'a, T> {
data: Vec<&'a T>,
}
impl<'a, T> RefVecParIter<'a, T> {
fn new(data: Vec<&'a T>) -> Self {
Self { data }
}
}
impl<'a, T: Send + Sync> ParallelIterator for RefVecParIter<'a, T> {
type Item = &'a T;
fn seq_items(self) -> Vec<Self::Item> {
self.data
}
fn drive<C, R>(mut self, consumer: C) -> R
where
C: Consumer<Self::Item, Result = R> + Send + Sync,
R: Send,
{
if self.data.len() <= 1 {
return consumer.consume(RefVecParIter::new(self.data));
}
let total_len = self.data.len();
let mid = total_len / 2;
let right_data = self.data.split_off(mid);
let left_data = std::mem::take(&mut self.data);
let (left_consumer, right_consumer) = consumer.split_at(left_data.len());
if total_len > PARALLEL_DRIVE_THRESHOLD {
return drive_split(
RefVecParIter::new(left_data),
RefVecParIter::new(right_data),
left_consumer,
right_consumer,
);
}
let left_result = left_consumer.consume(RefVecParIter::new(left_data));
let right_result = right_consumer.consume(RefVecParIter::new(right_data));
C::combine(left_result, right_result)
}
}
impl<'a, T: Send + Sync> IndexedParallelIterator for RefVecParIter<'a, T> {
fn len(&self) -> usize {
self.data.len()
}
fn collect_into_vec(self, target: &mut Vec<Self::Item>) {
move_vec_items_into(self.data, target);
}
}
impl IntoParallelIterator for std::ops::Range<usize> {
type Item = usize;
type Iter = RangeParIter<usize>;
fn into_par_iter(self) -> Self::Iter {
RangeParIter::new(self.start, self.end)
}
}
impl<T: Send + Sync> ParallelExtend<T> for Vec<T> {
fn par_extend<I>(&mut self, par_iter: I)
where
I: ParallelIterator<Item = T>,
{
self.extend(par_iter.drive(CollectConsumer::new()));
}
}