Skip to main content

moirai_iter/parallel/
sorting.rs

1//! Parallel slice sorting implementation.
2
3use moirai_core::error::ExecutorError;
4use moirai_executor::{HybridExecutor, SyncTask, global};
5use std::mem::MaybeUninit;
6
7/// Extension trait for parallel slice sorting.
8pub trait ParallelSliceMut<T: Send> {
9    /// Sorts the slice in parallel (stable).
10    fn par_sort(&mut self)
11    where
12        T: Ord;
13
14    /// Sorts the slice in parallel with a comparator (stable).
15    fn par_sort_by<F>(&mut self, compare: F)
16    where
17        F: Fn(&T, &T) -> std::cmp::Ordering + Sync + Send;
18
19    /// Sorts the slice in parallel with a key extraction function (stable).
20    fn par_sort_by_key<K, F>(&mut self, f: F)
21    where
22        F: Fn(&T) -> K + Sync + Send,
23        K: Ord + Send;
24
25    /// Sorts the slice in parallel (unstable).
26    fn par_sort_unstable(&mut self)
27    where
28        T: Ord;
29
30    /// Sorts the slice in parallel with a comparator (unstable).
31    fn par_sort_unstable_by<F>(&mut self, compare: F)
32    where
33        F: Fn(&T, &T) -> std::cmp::Ordering + Sync + Send;
34
35    /// Sorts the slice in parallel with a key extraction function (unstable).
36    fn par_sort_unstable_by_key<K, F>(&mut self, f: F)
37    where
38        F: Fn(&T) -> K + Sync + Send,
39        K: Ord + Send;
40}
41
42impl<T: Send> ParallelSliceMut<T> for [T] {
43    fn par_sort(&mut self)
44    where
45        T: Ord,
46    {
47        self.par_sort_by(T::cmp);
48    }
49
50    fn par_sort_by<F>(&mut self, compare: F)
51    where
52        F: Fn(&T, &T) -> std::cmp::Ordering + Sync + Send,
53    {
54        let executor = global();
55        let grain = fork_grain(executor, self.len(), STABLE_SEQUENTIAL_THRESHOLD);
56        par_merge_sort_impl(executor, self, &compare, grain);
57    }
58
59    fn par_sort_by_key<K, F>(&mut self, f: F)
60    where
61        F: Fn(&T) -> K + Sync + Send,
62        K: Ord + Send,
63    {
64        self.par_sort_by(move |a, b| f(a).cmp(&f(b)));
65    }
66
67    fn par_sort_unstable(&mut self)
68    where
69        T: Ord,
70    {
71        self.par_sort_unstable_by(T::cmp);
72    }
73
74    fn par_sort_unstable_by<F>(&mut self, compare: F)
75    where
76        F: Fn(&T, &T) -> std::cmp::Ordering + Sync + Send,
77    {
78        let executor = global();
79        let grain = fork_grain(executor, self.len(), UNSTABLE_SEQUENTIAL_THRESHOLD);
80        par_sort_unstable_by_impl(executor, self, &compare, grain);
81    }
82
83    fn par_sort_unstable_by_key<K, F>(&mut self, f: F)
84    where
85        F: Fn(&T) -> K + Sync + Send,
86        K: Ord + Send,
87    {
88        self.par_sort_unstable_by(move |a, b| f(a).cmp(&f(b)));
89    }
90}
91
92// Sequential thresholds keep worker dispatch on inputs large enough to amortize
93// task scheduling and merge/partition overhead.
94const STABLE_SEQUENTIAL_THRESHOLD: usize = 2048;
95const UNSTABLE_SEQUENTIAL_THRESHOLD: usize = 16_384;
96
97/// Segments per worker the recursion aims for before it stops forking.
98///
99/// The thresholds above are an absolute floor, not a granularity policy: on a
100/// large input they leave a leaf count proportional to `len`, and every leaf
101/// costs a scope — measurably more than the sort work it enables once the leaf
102/// is small relative to the machine. Oversubscribing the workers by this factor
103/// keeps enough independent segments for stealing to balance an uneven split
104/// (a straggler delays the phase by at most one segment) while making the fork
105/// count a function of machine width rather than input size.
106const SEGMENTS_PER_WORKER: usize = 8;
107
108/// Smallest sub-slice still worth handing to another lane.
109fn fork_grain(executor: &HybridExecutor, len: usize, sequential_threshold: usize) -> usize {
110    let workers = executor.config().worker_threads.max(1);
111    sequential_threshold.max(len.div_ceil(workers.saturating_mul(SEGMENTS_PER_WORKER)))
112}
113
114fn partition<T, F>(v: &mut [T], compare: &F) -> usize
115where
116    F: Fn(&T, &T) -> std::cmp::Ordering,
117{
118    let len = v.len();
119    if len <= 1 {
120        return 0;
121    }
122
123    let pivot_idx = len / 2;
124    v.swap(0, pivot_idx);
125
126    let mut i = 1;
127    let mut j = len - 1;
128
129    loop {
130        while i < len && compare(&v[i], &v[0]) == std::cmp::Ordering::Less {
131            i += 1;
132        }
133        while j > 0 && compare(&v[j], &v[0]) == std::cmp::Ordering::Greater {
134            j -= 1;
135        }
136        if i >= j {
137            break;
138        }
139        v.swap(i, j);
140        i += 1;
141        j -= 1;
142    }
143    v.swap(0, j);
144    j
145}
146
147/// Sort both halves concurrently: one on a scheduler lane, the other on the
148/// caller's.
149///
150/// The scheduler's scope is the fork-join primitive here rather than a plain
151/// thread pool because a worker that waits inside a scope *runs queued work*
152/// instead of parking (ADR-019). A pool without that property starves the
153/// moment recursion blocks every worker on a half that is still queued, which
154/// is what the deleted fork budget existed to prevent — at the cost of capping
155/// the whole work tree at the pool's width. Scoped jobs also borrow, so the
156/// halves cross the lane boundary as ordinary `&mut [T]` rather than as raw
157/// pointers laundered through a `'static` bound.
158///
159/// Each half is captured by unique borrow, never moved, so a job the scheduler
160/// refuses can still be run here: on refusal neither half has been touched.
161///
162/// The executor is a parameter rather than `global()` so the refusal path can
163/// be exercised against a shut-down executor in tests.
164///
165/// # Panics
166///
167/// Panics if the scheduled half panicked, propagating the failure on the
168/// caller's thread as rayon does.
169fn fork_join_halves<T, F, S>(
170    executor: &HybridExecutor,
171    left: &mut [T],
172    right: &mut [T],
173    compare: &F,
174    grain: usize,
175    sort: S,
176) where
177    T: Send,
178    F: Fn(&T, &T) -> std::cmp::Ordering + Sync + Send,
179    S: Fn(&HybridExecutor, &mut [T], &F, usize) + Copy + Send + Sync,
180{
181    // The job below *reborrows* its half rather than moving it: passing a
182    // `&mut` where a `&mut` is expected is an implicit reborrow, so the borrow
183    // ends when the scope joins and both bindings are usable again afterwards.
184    // That is what lets the refusal arm run a half the scheduler never ran.
185    let forked = executor.scope::<SyncTask, _>(|scope| {
186        scope.spawn(|_| sort(executor, left, compare, grain))?;
187        // Enter the scheduler before the caller takes its own half, so the two
188        // halves overlap instead of running back to back.
189        scope.flush()?;
190        sort(executor, right, compare, grain);
191        Ok(())
192    });
193
194    match forked {
195        Ok(()) => {}
196        // The scheduler refused the job and dropped it unexecuted, so `flush`
197        // returned before the caller's half started: neither half has run and
198        // both are still owned here. `ShuttingDown` and a full admission queue
199        // are the two ways that happens; running the work on the caller is the
200        // same answer `for_each_indexed` gives a rejected chunk.
201        Err(ExecutorError::ShuttingDown | ExecutorError::ResourceExhausted(_)) => {
202            sort(executor, left, compare, grain);
203            sort(executor, right, compare, grain);
204        }
205        Err(error) => panic!("invariant: scheduled sort half failed ({error})"),
206    }
207}
208
209fn par_sort_unstable_by_impl<T, F>(
210    executor: &HybridExecutor,
211    slice: &mut [T],
212    compare: &F,
213    grain: usize,
214) where
215    T: Send,
216    F: Fn(&T, &T) -> std::cmp::Ordering + Sync + Send,
217{
218    let len = slice.len();
219    if len <= grain {
220        slice.sort_unstable_by(compare);
221        return;
222    }
223
224    let pivot_idx = partition(slice, compare);
225    let (left, right) = slice.split_at_mut(pivot_idx);
226    let right = if right.is_empty() {
227        right
228    } else {
229        &mut right[1..] // Skip the pivot itself
230    };
231
232    fork_join_halves(
233        executor,
234        left,
235        right,
236        compare,
237        grain,
238        par_sort_unstable_by_impl,
239    );
240}
241
242fn par_merge_sort_impl<T, F>(executor: &HybridExecutor, slice: &mut [T], compare: &F, grain: usize)
243where
244    T: Send,
245    F: Fn(&T, &T) -> std::cmp::Ordering + Sync + Send,
246{
247    let len = slice.len();
248    if len <= grain {
249        slice.sort_by(compare);
250        return;
251    }
252
253    let mid = len / 2;
254    {
255        let (left, right) = slice.split_at_mut(mid);
256        fork_join_halves(executor, left, right, compare, grain, par_merge_sort_impl);
257    }
258
259    merge(slice, mid, compare);
260}
261
262/// Restores the slice on an early exit from [`merge`].
263///
264/// Every access to the slice goes through `base`, the one raw pointer derived
265/// from the caller's `&mut [T]`, and every access to the buffered left run goes
266/// through `left`, derived once from the buffer's allocation. Re-deriving a
267/// pointer from the `&mut` per element would retag it uniquely each time and
268/// invalidate the pointer the previous element used.
269struct MergeGuard<T> {
270    base: *mut T,
271    left: *const T,
272    /// Owns the allocation `left` points into; its elements are moved-from
273    /// copies, so it is never dropped element-wise.
274    _left_storage: Vec<MaybeUninit<T>>,
275    i: usize,
276    j: usize,
277    k: usize,
278    mid: usize,
279}
280
281impl<T> Drop for MergeGuard<T> {
282    fn drop(&mut self) {
283        let remaining = self.mid - self.i;
284        if remaining > 0 {
285            // SAFETY: drop runs only when merge bailed early with unconsumed
286            // left elements; the source range `i..mid` lies inside the
287            // buffered run and the destination `k..k + remaining` (which ends
288            // at `j`) was vacated by the consumed prefix, so the ranges are
289            // disjoint and both in bounds.
290            unsafe {
291                std::ptr::copy_nonoverlapping(
292                    self.left.add(self.i),
293                    self.base.add(self.k),
294                    remaining,
295                );
296            }
297        }
298    }
299}
300
301fn merge<T, F>(slice: &mut [T], mid: usize, compare: &F)
302where
303    T: Send,
304    F: Fn(&T, &T) -> std::cmp::Ordering,
305{
306    let len = slice.len();
307    if len <= 1 || mid == 0 || mid >= len {
308        return;
309    }
310
311    let base = slice.as_mut_ptr();
312    let mut left_vec: Vec<MaybeUninit<T>> = Vec::with_capacity(mid);
313    let left = left_vec.as_mut_ptr().cast::<T>();
314    // SAFETY: capacity mid was just reserved; copying mid initialized
315    // elements from the slice's left half makes them initialized owners, so
316    // set_len is honest and no value is duplicated (the slice side of these
317    // slots is logically moved out and never dropped twice — merge writes
318    // every slot before any later drop).
319    unsafe {
320        std::ptr::copy_nonoverlapping(base.cast_const(), left, mid);
321        left_vec.set_len(mid);
322    }
323
324    let mut guard = MergeGuard {
325        base,
326        left: left.cast_const(),
327        _left_storage: left_vec,
328        i: 0,
329        j: mid,
330        k: 0,
331        mid,
332    };
333
334    while guard.i < guard.mid && guard.j < len {
335        // SAFETY: i < mid bounds the left index and the buffer holds mid
336        // initialized values per the copy above; j < len bounds the right
337        // index inside the slice `base` addresses. Both references end with
338        // the comparison, before any write below.
339        let (left_val, right_val) =
340            unsafe { (&*guard.left.add(guard.i), &*guard.base.add(guard.j)) };
341
342        if compare(left_val, right_val) == std::cmp::Ordering::Greater {
343            // SAFETY: j < len and k <= j hold during the right-run advance,
344            // so forward `copy` handles the overlap correctly and both
345            // indices stay in bounds.
346            unsafe {
347                std::ptr::copy(guard.base.add(guard.j), guard.base.add(guard.k), 1);
348            }
349            guard.j += 1;
350        } else {
351            // SAFETY: i < mid bounds the source; k <= i + (j - mid) keeps the
352            // destination at or behind consumed positions, disjoint from the
353            // buffered left run, and within slice bounds.
354            unsafe {
355                std::ptr::copy_nonoverlapping(guard.left.add(guard.i), guard.base.add(guard.k), 1);
356            }
357            guard.i += 1;
358        }
359        guard.k += 1;
360    }
361}
362
363#[cfg(test)]
364mod tests {
365    use super::*;
366    use std::sync::atomic::{AtomicUsize, Ordering};
367
368    // Recursion depth well past any worker count. Under a runtime whose waiters
369    // park instead of helping, the forked halves have nobody left to run them
370    // and the sort never returns, so a regression trips nextest's terminate
371    // bound rather than failing an assertion. The deterministic single-worker
372    // proof lives at the scheduler layer, where the scope contract is owned
373    // (ADR-019); this input is what the sort itself can constrain, since it
374    // runs on the process-wide executor.
375    #[test]
376    fn deep_recursion_completes() {
377        let mut data: Vec<u64> = (0..1_048_576u64).rev().collect();
378        let grain = fork_grain(global(), data.len(), STABLE_SEQUENTIAL_THRESHOLD);
379        par_merge_sort_impl(global(), &mut data, &u64::cmp, grain);
380
381        assert!(
382            data.windows(2).all(|pair| pair[0] <= pair[1]),
383            "the sort must both finish and order the slice"
384        );
385    }
386
387    #[test]
388    fn deep_unstable_recursion_completes() {
389        let mut data: Vec<u64> = (0..1_048_576u64).rev().collect();
390        let grain = fork_grain(global(), data.len(), UNSTABLE_SEQUENTIAL_THRESHOLD);
391        par_sort_unstable_by_impl(global(), &mut data, &u64::cmp, grain);
392
393        assert!(
394            data.windows(2).all(|pair| pair[0] <= pair[1]),
395            "the sort must both finish and order the slice"
396        );
397    }
398
399    // A scheduler that refuses the forked half must not lose it. Sorting
400    // against a shut-down executor makes every fork take the refusal arm, so
401    // the whole recursion falls back to the caller's lane: slower, still
402    // correct. A refusal arm that dropped the half instead would leave the
403    // slice unsorted at every level.
404    #[test]
405    fn refused_forks_run_on_the_caller() {
406        let mut executor =
407            moirai_executor::HybridExecutor::new(moirai_core::executor::ExecutorConfig {
408                worker_threads: 2,
409                ..moirai_core::executor::ExecutorConfig::default()
410            })
411            .expect("build a local executor");
412        executor.shutdown().expect("shut the local executor down");
413        assert!(
414            executor.scope::<SyncTask, _>(|_| Ok(())).is_err(),
415            "precondition: the executor must refuse scopes, or the sorts below \
416             never reach the refusal arm"
417        );
418
419        let mut data: Vec<u64> = (0..16_384u64).rev().collect();
420        par_merge_sort_impl(&executor, &mut data, &u64::cmp, STABLE_SEQUENTIAL_THRESHOLD);
421        assert!(
422            data.windows(2).all(|pair| pair[0] <= pair[1]),
423            "a refused fork must still sort its half on the caller"
424        );
425
426        let mut data: Vec<u64> = (0..65_536u64).rev().collect();
427        par_sort_unstable_by_impl(
428            &executor,
429            &mut data,
430            &u64::cmp,
431            UNSTABLE_SEQUENTIAL_THRESHOLD,
432        );
433        assert!(
434            data.windows(2).all(|pair| pair[0] <= pair[1]),
435            "a refused fork must still sort its half on the caller"
436        );
437    }
438
439    // Sorts entered from *inside* a scheduler worker: the nested case, where a
440    // waiter that parks removes the last runner from the pool. Several of them
441    // run at once so the workers are saturated before any of them forks.
442    #[test]
443    fn nested_sorts_complete_from_scheduler_workers() {
444        const SORTS: usize = 8;
445
446        let mut inputs: Vec<Vec<u64>> =
447            (0..SORTS).map(|_| (0..65_536u64).rev().collect()).collect();
448
449        let slots: Vec<crate::base::SendPtr<Vec<u64>>> = inputs
450            .iter_mut()
451            .map(|input| crate::base::SendPtr(input as *mut Vec<u64>))
452            .collect();
453
454        moirai_executor::global()
455            .for_each_indexed::<SyncTask, _>(SORTS, |index| {
456                // Safety: each index owns exactly one element of `inputs`, and
457                // `for_each_indexed` joins every invocation before returning,
458                // so the borrows are disjoint and end before `inputs` is read.
459                let data = unsafe { &mut *slots[index].as_ptr() };
460                par_merge_sort_impl(
461                    global(),
462                    data.as_mut_slice(),
463                    &u64::cmp,
464                    STABLE_SEQUENTIAL_THRESHOLD,
465                );
466            })
467            .expect("nested sort fan-out must complete");
468
469        for input in &inputs {
470            assert!(
471                input.windows(2).all(|pair| pair[0] <= pair[1]),
472                "every nested sort must both finish and order its slice"
473            );
474        }
475    }
476
477    // The merge step in isolation: one thread, no scheduler, so a pointer
478    // derived from the slice and then invalidated by a later reborrow of it
479    // fails here deterministically under Miri. Equal keys pin stability: on a
480    // tie the left run's element must come first.
481    #[test]
482    fn merge_interleaves_two_sorted_runs_stably() {
483        let mut v = vec![
484            KeyVal { key: 1, val: 0 },
485            KeyVal { key: 3, val: 1 },
486            KeyVal { key: 5, val: 2 },
487            KeyVal { key: 1, val: 3 },
488            KeyVal { key: 3, val: 4 },
489            KeyVal { key: 4, val: 5 },
490        ];
491
492        merge(&mut v, 3, &|a, b| a.key.cmp(&b.key));
493
494        let order: Vec<(i32, usize)> = v.iter().map(|item| (item.key, item.val)).collect();
495        assert_eq!(order, [(1, 0), (1, 3), (3, 1), (3, 4), (4, 5), (5, 2)]);
496    }
497
498    #[test]
499    fn test_sorting_empty_and_single() {
500        let mut v: Vec<i32> = vec![];
501        v.par_sort();
502        assert!(v.is_empty());
503
504        let mut v = vec![42];
505        v.par_sort();
506        assert_eq!(v, vec![42]);
507
508        let mut v: Vec<i32> = vec![];
509        v.par_sort_unstable();
510        assert!(v.is_empty());
511
512        let mut v = vec![42];
513        v.par_sort_unstable();
514        assert_eq!(v, vec![42]);
515    }
516
517    #[test]
518    fn test_sorting_already_sorted_and_reverse() {
519        let mut v = vec![1, 2, 3, 4, 5, 6];
520        v.par_sort();
521        assert_eq!(v, vec![1, 2, 3, 4, 5, 6]);
522
523        let mut v = vec![6, 5, 4, 3, 2, 1];
524        v.par_sort();
525        assert_eq!(v, vec![1, 2, 3, 4, 5, 6]);
526
527        let mut v = vec![1, 2, 3, 4, 5, 6];
528        v.par_sort_unstable();
529        assert_eq!(v, vec![1, 2, 3, 4, 5, 6]);
530
531        let mut v = vec![6, 5, 4, 3, 2, 1];
532        v.par_sort_unstable();
533        assert_eq!(v, vec![1, 2, 3, 4, 5, 6]);
534    }
535
536    #[test]
537    fn test_sorting_duplicates() {
538        let mut v = vec![2, 2, 1, 1, 3, 3, 2, 2];
539        v.par_sort();
540        assert_eq!(v, vec![1, 1, 2, 2, 2, 2, 3, 3]);
541
542        let mut v = vec![2, 2, 1, 1, 3, 3, 2, 2];
543        v.par_sort_unstable();
544        assert_eq!(v, vec![1, 1, 2, 2, 2, 2, 3, 3]);
545    }
546
547    #[test]
548    fn test_sorting_large_random() {
549        // Simple deterministic LCG random generator for testing
550
551        let mut seed: u64 = 12345;
552        let mut random_u32 = move || {
553            seed = seed.wrapping_mul(1664525).wrapping_add(1013904223);
554            seed as u32
555        };
556
557        let mut original = Vec::new();
558        for _ in 0..5000 {
559            original.push(random_u32() % 10000);
560        }
561
562        let mut v1 = original.clone();
563        v1.par_sort();
564        let mut expected = original.clone();
565        expected.sort();
566        assert_eq!(v1, expected);
567
568        let mut v2 = original.clone();
569        v2.par_sort_unstable();
570        assert_eq!(v2, expected);
571    }
572
573    #[derive(Debug, Eq, PartialEq)]
574    struct KeyVal {
575        key: i32,
576        val: usize,
577    }
578
579    #[test]
580    fn test_sorting_stability() {
581        let mut v = vec![
582            KeyVal { key: 2, val: 0 },
583            KeyVal { key: 1, val: 1 },
584            KeyVal { key: 2, val: 2 },
585            KeyVal { key: 1, val: 3 },
586            KeyVal { key: 3, val: 4 },
587            KeyVal { key: 2, val: 5 },
588        ];
589
590        // Stable sort by key
591        v.par_sort_by(|a, b| a.key.cmp(&b.key));
592
593        assert_eq!(
594            v,
595            vec![
596                KeyVal { key: 1, val: 1 },
597                KeyVal { key: 1, val: 3 },
598                KeyVal { key: 2, val: 0 },
599                KeyVal { key: 2, val: 2 },
600                KeyVal { key: 2, val: 5 },
601                KeyVal { key: 3, val: 4 },
602            ]
603        );
604    }
605
606    #[test]
607    fn test_sorting_by_key() {
608        let mut v = [KeyVal { key: 2, val: 0 }, KeyVal { key: 1, val: 1 }];
609        v.par_sort_by_key(|item| item.key);
610        assert_eq!(v[0].key, 1);
611        assert_eq!(v[1].key, 2);
612
613        let mut v = [KeyVal { key: 2, val: 0 }, KeyVal { key: 1, val: 1 }];
614        v.par_sort_unstable_by_key(|item| item.key);
615        assert_eq!(v[0].key, 1);
616        assert_eq!(v[1].key, 2);
617    }
618
619    static DROP_COUNT: AtomicUsize = AtomicUsize::new(0);
620
621    #[derive(Debug, Clone, Eq, PartialEq)]
622    struct TrackedItem(i32);
623
624    impl Drop for TrackedItem {
625        fn drop(&mut self) {
626            DROP_COUNT.fetch_add(1, Ordering::SeqCst);
627        }
628    }
629
630    #[test]
631    fn test_panic_safety_no_double_drop() {
632        DROP_COUNT.store(0, Ordering::SeqCst);
633
634        let mut v = vec![
635            TrackedItem(3),
636            TrackedItem(1),
637            TrackedItem(2),
638            TrackedItem(4),
639        ];
640
641        let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
642            v.par_sort_by(|a, b| {
643                if a.0 == 2 || b.0 == 2 {
644                    panic!("simulated comparator panic");
645                }
646                a.0.cmp(&b.0)
647            });
648        }));
649
650        let payload = result.expect_err("a panicking comparator must unwind to the caller");
651        assert_eq!(
652            payload.downcast_ref::<&str>(),
653            Some(&"simulated comparator panic")
654        );
655        // Verify drop count matches number of elements exactly once when vector is dropped
656        drop(v);
657        assert_eq!(DROP_COUNT.load(Ordering::SeqCst), 4);
658    }
659}