1use crate::{IntoParIter, NumThreads, Par, Params, ThreadPool, runner::ParRunner};
2use alloc::vec::Vec;
3use core::{mem::MaybeUninit, num::NonZeroUsize};
4
5struct SortChunkTask<T> {
7 ptr: *mut T,
8 len: usize,
9}
10unsafe impl<T: Send> Send for SortChunkTask<T> {}
11unsafe impl<T: Send> Sync for SortChunkTask<T> {}
12
13impl<T: Ord> SortChunkTask<T> {
14 #[inline]
15 fn execute(self) {
16 if self.len > 1 {
17 let s = unsafe { core::slice::from_raw_parts_mut(self.ptr, self.len) };
18 s.sort_unstable();
19 }
20 }
21}
22
23fn find_split<T: Ord>(a: &[T], b: &[T], r: usize) -> (usize, usize) {
26 let mut low = r.saturating_sub(b.len());
27 let mut high = a.len().min(r);
28 while low < high {
29 let mid = low + (high - low) / 2;
30 let b_idx = r - mid - 1;
31 match a[mid] < b[b_idx] {
32 true => low = mid + 1,
33 false => high = mid,
34 }
35 }
36 (low, r - low)
37}
38
39struct MergeSubTask<T> {
41 src_a: *const T,
42 len_a: usize,
43 src_b: *const T,
44 len_b: usize,
45 dst: *mut T,
46}
47unsafe impl<T: Send> Send for MergeSubTask<T> {}
48unsafe impl<T: Send> Sync for MergeSubTask<T> {}
49
50impl<T: Ord> MergeSubTask<T> {
51 #[inline]
52 fn execute(self) {
53 if self.len_a == 0 {
54 if self.len_b > 0 {
55 unsafe {
56 core::ptr::copy_nonoverlapping(self.src_b, self.dst, self.len_b);
57 }
58 }
59 return;
60 }
61 if self.len_b == 0 {
62 unsafe {
63 core::ptr::copy_nonoverlapping(self.src_a, self.dst, self.len_a);
64 }
65 return;
66 }
67
68 let last_a = unsafe { &*self.src_a.add(self.len_a - 1) };
70 let first_b = unsafe { &*self.src_b };
71 if last_a <= first_b {
72 unsafe {
73 core::ptr::copy_nonoverlapping(self.src_a, self.dst, self.len_a);
74 core::ptr::copy_nonoverlapping(self.src_b, self.dst.add(self.len_a), self.len_b);
75 }
76 return;
77 }
78
79 let last_b = unsafe { &*self.src_b.add(self.len_b - 1) };
81 let first_a = unsafe { &*self.src_a };
82 if last_b < first_a {
83 unsafe {
84 core::ptr::copy_nonoverlapping(self.src_b, self.dst, self.len_b);
85 core::ptr::copy_nonoverlapping(self.src_a, self.dst.add(self.len_b), self.len_a);
86 }
87 return;
88 }
89
90 unsafe {
92 let mut ptr_a = self.src_a;
93 let end_a = self.src_a.add(self.len_a);
94 let mut ptr_b = self.src_b;
95 let end_b = self.src_b.add(self.len_b);
96 let mut out = self.dst;
97
98 while ptr_a < end_a && ptr_b < end_b {
99 match *ptr_b < *ptr_a {
100 true => {
101 core::ptr::copy_nonoverlapping(ptr_b, out, 1);
102 ptr_b = ptr_b.add(1);
103 }
104 false => {
105 core::ptr::copy_nonoverlapping(ptr_a, out, 1);
106 ptr_a = ptr_a.add(1);
107 }
108 }
109 out = out.add(1);
110 }
111
112 if ptr_a < end_a {
113 let rem = end_a.offset_from(ptr_a) as usize;
114 core::ptr::copy_nonoverlapping(ptr_a, out, rem);
115 } else if ptr_b < end_b {
116 let rem = end_b.offset_from(ptr_b) as usize;
117 core::ptr::copy_nonoverlapping(ptr_b, out, rem);
118 }
119 }
120 }
121}
122
123struct CopyChunkTask<T> {
125 src: *const T,
126 dst: *mut T,
127 len: usize,
128}
129unsafe impl<T: Send> Send for CopyChunkTask<T> {}
130unsafe impl<T: Send> Sync for CopyChunkTask<T> {}
131
132impl<T> CopyChunkTask<T> {
133 #[inline]
134 fn execute(self) {
135 if self.len > 0 {
136 unsafe {
137 core::ptr::copy_nonoverlapping(self.src, self.dst, self.len);
138 }
139 }
140 }
141}
142
143pub fn par_experimental_sort<T, R>(slice: &mut [T], runner: &mut R, params: Params)
156where
157 T: Ord + Send,
158 R: ParRunner,
159{
160 let n = slice.len();
161 if n <= 1 {
162 return;
163 }
164
165 if params.is_sequential() {
166 slice.sort_unstable();
167 return;
168 }
169
170 let par_num_threads = match params.num_threads {
171 NumThreads::Auto => NonZeroUsize::MAX,
172 NumThreads::Max(x) => x,
173 };
174 let max_threads: usize = runner.pool().max_num_threads().min(par_num_threads).into();
175
176 if max_threads <= 1 || n < 1024 {
177 slice.sort_unstable();
178 return;
179 }
180
181 let mut num_chunks = (max_threads * 2).next_power_of_two();
183 while num_chunks > 2 && n / num_chunks < 512 {
184 num_chunks /= 2;
185 }
186
187 if n / num_chunks < 64 {
188 slice.sort_unstable();
189 return;
190 }
191
192 let orig_ptr = slice.as_mut_ptr();
193
194 let mut sort_tasks = Vec::with_capacity(num_chunks);
196 for i in 0..num_chunks {
197 let start = i * n / num_chunks;
198 let end = (i + 1) * n / num_chunks;
199 let len = end - start;
200 sort_tasks.push(SortChunkTask {
201 ptr: unsafe { orig_ptr.add(start) },
202 len,
203 });
204 }
205
206 let par = sort_tasks.into_par().runner(&mut *runner);
207 let par = params.apply(par);
208 par.for_each(|task| task.execute());
209
210 let mut aux = Vec::<MaybeUninit<T>>::with_capacity(n);
212 let aux_ptr = aux.as_mut_ptr() as *mut T;
213
214 let mut current_chunks = num_chunks;
216 let mut pass = 0;
217
218 while current_chunks > 1 {
219 let (src_base, dst_base) = if pass % 2 == 0 {
220 (orig_ptr, aux_ptr)
221 } else {
222 (aux_ptr, orig_ptr)
223 };
224
225 let next_chunks = current_chunks / 2;
226 let stride = num_chunks / current_chunks;
227 let tasks_per_pair = (max_threads / next_chunks).max(1);
228 let mut merge_tasks = Vec::with_capacity(next_chunks * tasks_per_pair);
229
230 for pair_idx in 0..next_chunks {
231 let chunk_a_idx = pair_idx * 2;
232 let chunk_b_idx = chunk_a_idx + 1;
233
234 let start_a = (chunk_a_idx * stride) * n / num_chunks;
235 let end_a = ((chunk_a_idx + 1) * stride) * n / num_chunks;
236 let len_a = end_a - start_a;
237
238 let start_b = (chunk_b_idx * stride) * n / num_chunks;
239 let end_b = ((chunk_b_idx + 1) * stride) * n / num_chunks;
240 let len_b = end_b - start_b;
241
242 let slice_a = unsafe { core::slice::from_raw_parts(src_base.add(start_a), len_a) };
243 let slice_b = unsafe { core::slice::from_raw_parts(src_base.add(start_b), len_b) };
244
245 let pair_total = len_a + len_b;
246
247 let mut prev_a = 0;
248 let mut prev_b = 0;
249 let mut prev_r = 0;
250
251 for t in 1..=tasks_per_pair {
252 let r = t * pair_total / tasks_per_pair;
253 let (curr_a, curr_b) = if t == tasks_per_pair {
254 (len_a, len_b)
255 } else {
256 find_split(slice_a, slice_b, r)
257 };
258
259 let sub_len_a = curr_a - prev_a;
260 let sub_len_b = curr_b - prev_b;
261 let sub_dst = unsafe { dst_base.add(start_a + prev_r) };
262
263 merge_tasks.push(MergeSubTask {
264 src_a: unsafe { slice_a.as_ptr().add(prev_a) },
265 len_a: sub_len_a,
266 src_b: unsafe { slice_b.as_ptr().add(prev_b) },
267 len_b: sub_len_b,
268 dst: sub_dst,
269 });
270
271 prev_a = curr_a;
272 prev_b = curr_b;
273 prev_r = r;
274 }
275 }
276
277 let par = merge_tasks.into_par().runner(&mut *runner);
278 let par = params.apply(par);
279 par.for_each(|task| task.execute());
280
281 current_chunks = next_chunks;
282 pass += 1;
283 }
284
285 if pass % 2 != 0 {
287 let copy_chunks = max_threads.min(n);
288 let mut copy_tasks = Vec::with_capacity(copy_chunks);
289 for i in 0..copy_chunks {
290 let start = i * n / copy_chunks;
291 let end = (i + 1) * n / copy_chunks;
292 let len = end - start;
293 copy_tasks.push(CopyChunkTask {
294 src: unsafe { aux_ptr.add(start) },
295 dst: unsafe { orig_ptr.add(start) },
296 len,
297 });
298 }
299
300 let par = copy_tasks.into_par().runner(&mut *runner);
301 let par = params.apply(par);
302 par.for_each(|task| task.execute());
303 }
304}
305
306#[cfg(test)]
307mod tests {
308 use super::*;
309 use crate::runner::default_runner;
310 use alloc::string::String;
311
312 #[test]
313 fn test_sort_empty_and_single() {
314 let mut runner = default_runner();
315 let mut empty: [i32; 0] = [];
316 par_experimental_sort(&mut empty, &mut runner, Params::default());
317
318 let mut single = [42];
319 par_experimental_sort(&mut single, &mut runner, Params::default());
320 assert_eq!(single, [42]);
321 }
322
323 #[test]
324 fn test_sort_small_slices() {
325 let mut runner = default_runner();
326 let mut data = [9, 3, 7, 1, 5, 2, 8, 4, 6];
327 par_experimental_sort(&mut data, &mut runner, Params::default());
328 assert_eq!(data, [1, 2, 3, 4, 5, 6, 7, 8, 9]);
329 }
330
331 #[test]
332 fn test_sort_medium_and_large_random() {
333 let mut runner = default_runner();
334
335 for size in [500, 1024, 2048, 5000, 20000, 50000] {
336 let mut data: Vec<i32> = (0..size as u64)
337 .map(|i| (i.wrapping_mul(1103515245).wrapping_add(12345) & 0x7FFFFFFF) as i32)
338 .collect();
339 let mut expected = data.clone();
340 expected.sort_unstable();
341
342 par_experimental_sort(&mut data, &mut runner, Params::default());
343 assert_eq!(data, expected, "Failed for size {}", size);
344 }
345 }
346
347 #[test]
348 fn test_sort_sorted_and_reversed() {
349 let mut runner = default_runner();
350 let size = 10000;
351
352 let mut sorted: Vec<i32> = (0..size).collect();
353 par_experimental_sort(&mut sorted, &mut runner, Params::default());
354 assert!(sorted.windows(2).all(|w| w[0] <= w[1]));
355
356 let mut reversed: Vec<i32> = (0..size).rev().collect();
357 par_experimental_sort(&mut reversed, &mut runner, Params::default());
358 assert!(reversed.windows(2).all(|w| w[0] <= w[1]));
359 }
360
361 #[test]
362 fn test_sort_high_duplicates() {
363 let mut runner = default_runner();
364 let size = 20000;
365 let mut data: Vec<i32> = (0..size).map(|i| i % 7).collect();
366 let mut expected = data.clone();
367 expected.sort_unstable();
368
369 par_experimental_sort(&mut data, &mut runner, Params::default());
370 assert_eq!(data, expected);
371 }
372
373 #[test]
374 fn test_sort_non_copy_types() {
375 let mut runner = default_runner();
376 let size = 5000;
377 let mut data: Vec<String> = (0..size)
378 .map(|i| alloc::format!("item_{:06}", (i * 7919) % size))
379 .collect();
380 let mut expected = data.clone();
381 expected.sort();
382
383 par_experimental_sort(&mut data, &mut runner, Params::default());
384 assert_eq!(data, expected);
385 }
386
387 #[test]
388 fn test_sort_num_threads_configs() {
389 for nt in [1, 2, 4, 8] {
390 let mut runner = default_runner();
391 let size = 10000;
392 let mut data: Vec<i32> = (0..size as u64)
393 .map(|i| (i.wrapping_mul(2654435761) & 0x7FFFFFFF) as i32)
394 .collect();
395 let mut expected = data.clone();
396 expected.sort_unstable();
397
398 let params = Params::default().with_num_threads(nt);
399 par_experimental_sort(&mut data, &mut runner, params);
400 assert_eq!(data, expected, "Failed for num_threads = {}", nt);
401 }
402 }
403}