use moirai_core::error::ExecutorError;
use moirai_executor::{global, HybridExecutor, SyncTask};
use std::mem::MaybeUninit;
pub trait ParallelSliceMut<T: Send> {
fn par_sort(&mut self)
where
T: Ord;
fn par_sort_by<F>(&mut self, compare: F)
where
F: Fn(&T, &T) -> std::cmp::Ordering + Sync + Send;
fn par_sort_by_key<K, F>(&mut self, f: F)
where
F: Fn(&T) -> K + Sync + Send,
K: Ord + Send;
fn par_sort_unstable(&mut self)
where
T: Ord;
fn par_sort_unstable_by<F>(&mut self, compare: F)
where
F: Fn(&T, &T) -> std::cmp::Ordering + Sync + Send;
fn par_sort_unstable_by_key<K, F>(&mut self, f: F)
where
F: Fn(&T) -> K + Sync + Send,
K: Ord + Send;
}
impl<T: Send> ParallelSliceMut<T> for [T] {
fn par_sort(&mut self)
where
T: Ord,
{
self.par_sort_by(T::cmp);
}
fn par_sort_by<F>(&mut self, compare: F)
where
F: Fn(&T, &T) -> std::cmp::Ordering + Sync + Send,
{
let executor = global();
let grain = fork_grain(executor, self.len(), STABLE_SEQUENTIAL_THRESHOLD);
par_merge_sort_impl(executor, self, &compare, grain);
}
fn par_sort_by_key<K, F>(&mut self, f: F)
where
F: Fn(&T) -> K + Sync + Send,
K: Ord + Send,
{
self.par_sort_by(move |a, b| f(a).cmp(&f(b)));
}
fn par_sort_unstable(&mut self)
where
T: Ord,
{
self.par_sort_unstable_by(T::cmp);
}
fn par_sort_unstable_by<F>(&mut self, compare: F)
where
F: Fn(&T, &T) -> std::cmp::Ordering + Sync + Send,
{
let executor = global();
let grain = fork_grain(executor, self.len(), UNSTABLE_SEQUENTIAL_THRESHOLD);
par_sort_unstable_by_impl(executor, self, &compare, grain);
}
fn par_sort_unstable_by_key<K, F>(&mut self, f: F)
where
F: Fn(&T) -> K + Sync + Send,
K: Ord + Send,
{
self.par_sort_unstable_by(move |a, b| f(a).cmp(&f(b)));
}
}
const STABLE_SEQUENTIAL_THRESHOLD: usize = 2048;
const UNSTABLE_SEQUENTIAL_THRESHOLD: usize = 16_384;
const SEGMENTS_PER_WORKER: usize = 8;
fn fork_grain(executor: &HybridExecutor, len: usize, sequential_threshold: usize) -> usize {
let workers = executor.config().worker_threads.max(1);
sequential_threshold.max(len.div_ceil(workers.saturating_mul(SEGMENTS_PER_WORKER)))
}
fn partition<T, F>(v: &mut [T], compare: &F) -> usize
where
F: Fn(&T, &T) -> std::cmp::Ordering,
{
let len = v.len();
if len <= 1 {
return 0;
}
let pivot_idx = len / 2;
v.swap(0, pivot_idx);
let mut i = 1;
let mut j = len - 1;
loop {
while i < len && compare(&v[i], &v[0]) == std::cmp::Ordering::Less {
i += 1;
}
while j > 0 && compare(&v[j], &v[0]) == std::cmp::Ordering::Greater {
j -= 1;
}
if i >= j {
break;
}
v.swap(i, j);
i += 1;
j -= 1;
}
v.swap(0, j);
j
}
fn fork_join_halves<T, F, S>(
executor: &HybridExecutor,
left: &mut [T],
right: &mut [T],
compare: &F,
grain: usize,
sort: S,
) where
T: Send,
F: Fn(&T, &T) -> std::cmp::Ordering + Sync + Send,
S: Fn(&HybridExecutor, &mut [T], &F, usize) + Copy + Send + Sync,
{
let forked = executor.scope::<SyncTask, _>(|scope| {
scope.spawn(|_| sort(executor, left, compare, grain))?;
scope.flush()?;
sort(executor, right, compare, grain);
Ok(())
});
match forked {
Ok(()) => {}
Err(ExecutorError::ShuttingDown | ExecutorError::ResourceExhausted(_)) => {
sort(executor, left, compare, grain);
sort(executor, right, compare, grain);
}
Err(error) => panic!("invariant: scheduled sort half failed ({error})"),
}
}
fn par_sort_unstable_by_impl<T, F>(
executor: &HybridExecutor,
slice: &mut [T],
compare: &F,
grain: usize,
) where
T: Send,
F: Fn(&T, &T) -> std::cmp::Ordering + Sync + Send,
{
let len = slice.len();
if len <= grain {
slice.sort_unstable_by(compare);
return;
}
let pivot_idx = partition(slice, compare);
let (left, right) = slice.split_at_mut(pivot_idx);
let right = if right.is_empty() {
right
} else {
&mut right[1..] };
fork_join_halves(
executor,
left,
right,
compare,
grain,
par_sort_unstable_by_impl,
);
}
fn par_merge_sort_impl<T, F>(executor: &HybridExecutor, slice: &mut [T], compare: &F, grain: usize)
where
T: Send,
F: Fn(&T, &T) -> std::cmp::Ordering + Sync + Send,
{
let len = slice.len();
if len <= grain {
slice.sort_by(compare);
return;
}
let mid = len / 2;
{
let (left, right) = slice.split_at_mut(mid);
fork_join_halves(executor, left, right, compare, grain, par_merge_sort_impl);
}
merge(slice, mid, compare);
}
struct MergeGuard<'a, T> {
slice: &'a mut [T],
left_vec: Vec<MaybeUninit<T>>,
i: usize,
j: usize,
k: usize,
mid: usize,
}
impl<'a, T> Drop for MergeGuard<'a, T> {
fn drop(&mut self) {
let remaining = self.mid - self.i;
if remaining > 0 {
unsafe {
std::ptr::copy_nonoverlapping(
self.left_vec.as_ptr().add(self.i),
self.slice.as_mut_ptr().add(self.k).cast::<MaybeUninit<T>>(),
remaining,
);
}
}
}
}
fn merge<T, F>(slice: &mut [T], mid: usize, compare: &F)
where
T: Send,
F: Fn(&T, &T) -> std::cmp::Ordering,
{
let len = slice.len();
if len <= 1 || mid == 0 || mid >= len {
return;
}
let mut left_vec: Vec<MaybeUninit<T>> = Vec::with_capacity(mid);
unsafe {
std::ptr::copy_nonoverlapping(
slice.as_ptr().cast::<MaybeUninit<T>>(),
left_vec.as_mut_ptr(),
mid,
);
left_vec.set_len(mid);
}
let mut guard = MergeGuard {
slice,
left_vec,
i: 0,
j: mid,
k: 0,
mid,
};
while guard.i < guard.mid && guard.j < len {
let left_val = unsafe { &*guard.left_vec.as_ptr().add(guard.i).cast::<T>() };
let right_val = &guard.slice[guard.j];
if compare(left_val, right_val) == std::cmp::Ordering::Greater {
unsafe {
std::ptr::copy(
guard.slice.as_ptr().add(guard.j),
guard.slice.as_mut_ptr().add(guard.k),
1,
);
}
guard.j += 1;
} else {
unsafe {
std::ptr::copy_nonoverlapping(
guard.left_vec.as_ptr().add(guard.i),
guard
.slice
.as_mut_ptr()
.add(guard.k)
.cast::<MaybeUninit<T>>(),
1,
);
}
guard.i += 1;
}
guard.k += 1;
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
#[test]
fn deep_recursion_completes() {
let mut data: Vec<u64> = (0..1_048_576u64).rev().collect();
let grain = fork_grain(global(), data.len(), STABLE_SEQUENTIAL_THRESHOLD);
par_merge_sort_impl(global(), &mut data, &u64::cmp, grain);
assert!(
data.windows(2).all(|pair| pair[0] <= pair[1]),
"the sort must both finish and order the slice"
);
}
#[test]
fn deep_unstable_recursion_completes() {
let mut data: Vec<u64> = (0..1_048_576u64).rev().collect();
let grain = fork_grain(global(), data.len(), UNSTABLE_SEQUENTIAL_THRESHOLD);
par_sort_unstable_by_impl(global(), &mut data, &u64::cmp, grain);
assert!(
data.windows(2).all(|pair| pair[0] <= pair[1]),
"the sort must both finish and order the slice"
);
}
#[test]
fn refused_forks_run_on_the_caller() {
let mut executor =
moirai_executor::HybridExecutor::new(moirai_core::executor::ExecutorConfig {
worker_threads: 2,
..moirai_core::executor::ExecutorConfig::default()
})
.expect("build a local executor");
executor.shutdown().expect("shut the local executor down");
assert!(
executor.scope::<SyncTask, _>(|_| Ok(())).is_err(),
"precondition: the executor must refuse scopes, or the sorts below \
never reach the refusal arm"
);
let mut data: Vec<u64> = (0..16_384u64).rev().collect();
par_merge_sort_impl(&executor, &mut data, &u64::cmp, STABLE_SEQUENTIAL_THRESHOLD);
assert!(
data.windows(2).all(|pair| pair[0] <= pair[1]),
"a refused fork must still sort its half on the caller"
);
let mut data: Vec<u64> = (0..65_536u64).rev().collect();
par_sort_unstable_by_impl(
&executor,
&mut data,
&u64::cmp,
UNSTABLE_SEQUENTIAL_THRESHOLD,
);
assert!(
data.windows(2).all(|pair| pair[0] <= pair[1]),
"a refused fork must still sort its half on the caller"
);
}
#[test]
fn nested_sorts_complete_from_scheduler_workers() {
const SORTS: usize = 8;
let mut inputs: Vec<Vec<u64>> =
(0..SORTS).map(|_| (0..65_536u64).rev().collect()).collect();
let slots: Vec<crate::base::SendPtr<Vec<u64>>> = inputs
.iter_mut()
.map(|input| crate::base::SendPtr(input as *mut Vec<u64>))
.collect();
moirai_executor::global()
.for_each_indexed::<SyncTask, _>(SORTS, |index| {
let data = unsafe { &mut *slots[index].as_ptr() };
par_merge_sort_impl(
global(),
data.as_mut_slice(),
&u64::cmp,
STABLE_SEQUENTIAL_THRESHOLD,
);
})
.expect("nested sort fan-out must complete");
for input in &inputs {
assert!(
input.windows(2).all(|pair| pair[0] <= pair[1]),
"every nested sort must both finish and order its slice"
);
}
}
#[test]
fn test_sorting_empty_and_single() {
let mut v: Vec<i32> = vec![];
v.par_sort();
assert!(v.is_empty());
let mut v = vec![42];
v.par_sort();
assert_eq!(v, vec![42]);
let mut v: Vec<i32> = vec![];
v.par_sort_unstable();
assert!(v.is_empty());
let mut v = vec![42];
v.par_sort_unstable();
assert_eq!(v, vec![42]);
}
#[test]
fn test_sorting_already_sorted_and_reverse() {
let mut v = vec![1, 2, 3, 4, 5, 6];
v.par_sort();
assert_eq!(v, vec![1, 2, 3, 4, 5, 6]);
let mut v = vec![6, 5, 4, 3, 2, 1];
v.par_sort();
assert_eq!(v, vec![1, 2, 3, 4, 5, 6]);
let mut v = vec![1, 2, 3, 4, 5, 6];
v.par_sort_unstable();
assert_eq!(v, vec![1, 2, 3, 4, 5, 6]);
let mut v = vec![6, 5, 4, 3, 2, 1];
v.par_sort_unstable();
assert_eq!(v, vec![1, 2, 3, 4, 5, 6]);
}
#[test]
fn test_sorting_duplicates() {
let mut v = vec![2, 2, 1, 1, 3, 3, 2, 2];
v.par_sort();
assert_eq!(v, vec![1, 1, 2, 2, 2, 2, 3, 3]);
let mut v = vec![2, 2, 1, 1, 3, 3, 2, 2];
v.par_sort_unstable();
assert_eq!(v, vec![1, 1, 2, 2, 2, 2, 3, 3]);
}
#[test]
fn test_sorting_large_random() {
let mut seed: u64 = 12345;
let mut random_u32 = move || {
seed = seed.wrapping_mul(1664525).wrapping_add(1013904223);
seed as u32
};
let mut original = Vec::new();
for _ in 0..5000 {
original.push(random_u32() % 10000);
}
let mut v1 = original.clone();
v1.par_sort();
let mut expected = original.clone();
expected.sort();
assert_eq!(v1, expected);
let mut v2 = original.clone();
v2.par_sort_unstable();
assert_eq!(v2, expected);
}
#[derive(Debug, Eq, PartialEq)]
struct KeyVal {
key: i32,
val: usize,
}
#[test]
fn test_sorting_stability() {
let mut v = vec![
KeyVal { key: 2, val: 0 },
KeyVal { key: 1, val: 1 },
KeyVal { key: 2, val: 2 },
KeyVal { key: 1, val: 3 },
KeyVal { key: 3, val: 4 },
KeyVal { key: 2, val: 5 },
];
v.par_sort_by(|a, b| a.key.cmp(&b.key));
assert_eq!(
v,
vec![
KeyVal { key: 1, val: 1 },
KeyVal { key: 1, val: 3 },
KeyVal { key: 2, val: 0 },
KeyVal { key: 2, val: 2 },
KeyVal { key: 2, val: 5 },
KeyVal { key: 3, val: 4 },
]
);
}
#[test]
fn test_sorting_by_key() {
let mut v = [KeyVal { key: 2, val: 0 }, KeyVal { key: 1, val: 1 }];
v.par_sort_by_key(|item| item.key);
assert_eq!(v[0].key, 1);
assert_eq!(v[1].key, 2);
let mut v = [KeyVal { key: 2, val: 0 }, KeyVal { key: 1, val: 1 }];
v.par_sort_unstable_by_key(|item| item.key);
assert_eq!(v[0].key, 1);
assert_eq!(v[1].key, 2);
}
static DROP_COUNT: AtomicUsize = AtomicUsize::new(0);
#[derive(Debug, Clone, Eq, PartialEq)]
struct TrackedItem(i32);
impl Drop for TrackedItem {
fn drop(&mut self) {
DROP_COUNT.fetch_add(1, Ordering::SeqCst);
}
}
#[test]
fn test_panic_safety_no_double_drop() {
DROP_COUNT.store(0, Ordering::SeqCst);
let mut v = vec![
TrackedItem(3),
TrackedItem(1),
TrackedItem(2),
TrackedItem(4),
];
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
v.par_sort_by(|a, b| {
if a.0 == 2 || b.0 == 2 {
panic!("simulated comparator panic");
}
a.0.cmp(&b.0)
});
}));
assert!(result.is_err());
drop(v);
assert_eq!(DROP_COUNT.load(Ordering::SeqCst), 4);
}
}