1use moirai_core::error::ExecutorError;
4use moirai_executor::{HybridExecutor, SyncTask, global};
5use std::mem::MaybeUninit;
6
7pub trait ParallelSliceMut<T: Send> {
9 fn par_sort(&mut self)
11 where
12 T: Ord;
13
14 fn par_sort_by<F>(&mut self, compare: F)
16 where
17 F: Fn(&T, &T) -> std::cmp::Ordering + Sync + Send;
18
19 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 fn par_sort_unstable(&mut self)
27 where
28 T: Ord;
29
30 fn par_sort_unstable_by<F>(&mut self, compare: F)
32 where
33 F: Fn(&T, &T) -> std::cmp::Ordering + Sync + Send;
34
35 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
92const STABLE_SEQUENTIAL_THRESHOLD: usize = 2048;
95const UNSTABLE_SEQUENTIAL_THRESHOLD: usize = 16_384;
96
97const SEGMENTS_PER_WORKER: usize = 8;
107
108fn 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
147fn 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 let forked = executor.scope::<SyncTask, _>(|scope| {
186 scope.spawn(|_| sort(executor, left, compare, grain))?;
187 scope.flush()?;
190 sort(executor, right, compare, grain);
191 Ok(())
192 });
193
194 match forked {
195 Ok(()) => {}
196 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..] };
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
262struct MergeGuard<T> {
270 base: *mut T,
271 left: *const T,
272 _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 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 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 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 unsafe {
347 std::ptr::copy(guard.base.add(guard.j), guard.base.add(guard.k), 1);
348 }
349 guard.j += 1;
350 } else {
351 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 #[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 #[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 #[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 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 #[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 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 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 drop(v);
657 assert_eq!(DROP_COUNT.load(Ordering::SeqCst), 4);
658 }
659}