1use std::{cmp::Ordering, ptr::copy_nonoverlapping};
2
3pub trait RadixSortKey: Copy {
4 const BYTES: usize;
5 fn radix_byte(self, byte: usize) -> usize;
6}
7
8macro_rules! unsigned_radix_sort_key {
9 ($($t:ty),* $(,)?) => {
10 $(
11 impl RadixSortKey for $t {
12 const BYTES: usize = (<$t>::BITS / 8) as usize;
13
14 fn radix_byte(self, byte: usize) -> usize {
15 ((self >> (byte * 8)) & 0xff) as usize
16 }
17 }
18 )*
19 };
20}
21
22macro_rules! signed_radix_sort_key {
23 ($($t:ty => $u:ty),* $(,)?) => {
24 $(
25 impl RadixSortKey for $t {
26 const BYTES: usize = (<$t>::BITS / 8) as usize;
27
28 fn radix_byte(self, byte: usize) -> usize {
29 let key = self as $u ^ (1 << (<$t>::BITS - 1));
30 ((key >> (byte * 8)) & 0xff) as usize
31 }
32 }
33 )*
34 };
35}
36
37unsigned_radix_sort_key!(u8, u16, u32, u64, u128, usize);
38signed_radix_sort_key!(
39 i8 => u8,
40 i16 => u16,
41 i32 => u32,
42 i64 => u64,
43 i128 => u128,
44 isize => usize,
45);
46
47pub trait SliceSortExt<T> {
48 fn bubble_sort(&mut self)
49 where
50 T: Ord;
51 fn bubble_sort_by<F>(&mut self, compare: F)
52 where
53 F: FnMut(&T, &T) -> Ordering;
54 fn merge_sort(&mut self)
55 where
56 T: Ord;
57 fn merge_sort_by<F>(&mut self, compare: F)
58 where
59 F: FnMut(&T, &T) -> Ordering;
60 fn insertion_sort(&mut self)
61 where
62 T: Ord;
63 fn insertion_sort_by<F>(&mut self, compare: F)
64 where
65 F: FnMut(&T, &T) -> Ordering;
66 fn radix_sort_by_key<K>(&mut self, key: impl FnMut(&T) -> K)
67 where
68 T: Clone,
69 K: RadixSortKey;
70}
71impl<T> SliceSortExt<T> for [T] {
72 fn bubble_sort(&mut self)
73 where
74 T: Ord,
75 {
76 bubble_sort(self, |a, b| a.lt(b));
77 }
78 fn bubble_sort_by<F>(&mut self, mut compare: F)
79 where
80 F: FnMut(&T, &T) -> Ordering,
81 {
82 bubble_sort(self, |a, b| compare(a, b) == Ordering::Less);
83 }
84 fn merge_sort(&mut self)
85 where
86 T: Ord,
87 {
88 merge_sort(self, |a, b| a.lt(b));
89 }
90 fn merge_sort_by<F>(&mut self, mut compare: F)
91 where
92 F: FnMut(&T, &T) -> Ordering,
93 {
94 merge_sort(self, |a, b| compare(a, b) == Ordering::Less);
95 }
96 fn insertion_sort(&mut self)
97 where
98 T: Ord,
99 {
100 insertion_sort(self, |a, b| a.lt(b));
101 }
102 fn insertion_sort_by<F>(&mut self, mut compare: F)
103 where
104 F: FnMut(&T, &T) -> Ordering,
105 {
106 insertion_sort(self, |a, b| compare(a, b) == Ordering::Less);
107 }
108 fn radix_sort_by_key<K>(&mut self, key: impl FnMut(&T) -> K)
109 where
110 T: Clone,
111 K: RadixSortKey,
112 {
113 radix_sort_by_key(self, key);
114 }
115}
116
117fn radix_sort_by_key<T, K>(values: &mut [T], mut key: impl FnMut(&T) -> K)
118where
119 T: Clone,
120 K: RadixSortKey,
121{
122 if values.len() <= 1 {
123 return;
124 }
125 let mut histograms = vec![[0usize; 256]; K::BYTES];
126 for value in values.iter() {
127 let key = key(value);
128 for (byte, histogram) in histograms.iter_mut().enumerate() {
129 histogram[key.radix_byte(byte)] += 1;
130 }
131 }
132 for histogram in histograms.iter_mut() {
133 let mut position = 0;
134 for count in histogram.iter_mut() {
135 let next = position + *count;
136 *count = position;
137 position = next;
138 }
139 }
140 let mut ends = [values.len(); 256];
141 ends[..255].copy_from_slice(&histograms[0][1..]);
142 let mut buffer = Vec::with_capacity(values.len());
143 {
144 let spare = buffer.spare_capacity_mut();
145 let positions = &mut histograms[0];
146 for value in values.iter() {
147 let bucket = key(value).radix_byte(0);
148 let position = &mut positions[bucket];
149 assert!(*position < ends[bucket]);
150 spare[*position].write(value.clone());
151 *position += 1;
152 }
153 }
154 unsafe { buffer.set_len(values.len()) };
156 for (byte, positions) in histograms.iter_mut().enumerate().skip(1) {
157 macro_rules! distribute {
158 ($source:expr, $destination:expr) => {{
159 let source = $source;
160 let destination = $destination;
161 for value in source {
162 let bucket = key(value).radix_byte(byte);
163 let position = &mut positions[bucket];
164 destination[*position].clone_from(value);
165 *position += 1;
166 }
167 }};
168 }
169 if byte % 2 == 0 {
170 distribute!(&*values, &mut buffer);
171 } else {
172 distribute!(&buffer, &mut *values);
173 }
174 }
175 if K::BYTES % 2 == 1 {
176 values.clone_from_slice(&buffer);
177 }
178}
179
180fn bubble_sort<T, F>(v: &mut [T], mut is_less: F)
181where
182 F: FnMut(&T, &T) -> bool,
183{
184 let len = v.len();
185 if len <= 1 {
186 return;
187 }
188 for i in 0..len - 1 {
189 for j in 0..len - i - 1 {
190 unsafe {
191 if is_less(v.get_unchecked(j + 1), v.get_unchecked(j)) {
192 v.swap(j, j + 1);
193 }
194 }
195 }
196 }
197}
198
199unsafe fn merge<T, F>(v: &mut [T], mut mid: usize, buf: *mut T, is_less: &mut F)
200where
201 F: FnMut(&T, &T) -> bool,
202{
203 unsafe {
204 let len = v.len();
205 let v = v.as_mut_ptr();
206 let (v_mid, v_end) = (v.add(mid), v.add(len));
207
208 copy_nonoverlapping(v, buf, mid);
209 let mut start = buf;
210 let end = buf.add(mid);
211 let mut dest = v;
212
213 let left = &mut start;
214 let mut right = v_mid;
215 while *left < end && right < v_end {
216 let to_copy = if is_less(&*right, &**left) {
217 get_and_increment(&mut right)
218 } else {
219 mid -= 1;
220 get_and_increment(left)
221 };
222 copy_nonoverlapping(to_copy, get_and_increment(&mut dest), 1);
223 }
224
225 copy_nonoverlapping(start, dest, mid);
227 }
228
229 unsafe fn get_and_increment<T>(ptr: &mut *mut T) -> *mut T {
230 let old = *ptr;
231 *ptr = unsafe { ptr.offset(1) };
232 old
233 }
234}
235
236fn merge_sort<T, F>(v: &mut [T], mut is_less: F)
237where
238 F: FnMut(&T, &T) -> bool,
239{
240 let len = v.len();
241 if len <= 1 {
242 return;
243 }
244 let mut buf = Vec::with_capacity(len / 2);
245 let mut runs: Vec<Run> = vec![];
246 let mut end = len;
247 while end > 0 {
248 let start = end - 1;
249 let mut left = Run {
250 start,
251 len: end - start,
252 };
253 end = start;
254
255 while let Some(right) = runs.pop_if(|right| left.start == 0 || right.len <= left.len) {
256 unsafe {
257 merge(
258 &mut v[left.start..right.start + right.len],
259 left.len,
260 buf.as_mut_ptr(),
261 &mut is_less,
262 );
263 }
264 left = Run {
265 start: left.start,
266 len: left.len + right.len,
267 };
268 }
269 runs.push(left);
270 }
271
272 debug_assert!(runs.len() == 1 && runs[0].start == 0 && runs[0].len == len);
273
274 #[derive(Clone, Copy)]
275 struct Run {
276 start: usize,
277 len: usize,
278 }
279}
280
281fn insertion_sort<T, F>(v: &mut [T], mut is_less: F)
282where
283 F: FnMut(&T, &T) -> bool,
284{
285 for i in 1..v.len() {
286 let x = &v[i];
287 let p = v[..i].partition_point(|y| is_less(y, x));
288 v[p..=i].rotate_right(1);
289 }
290}
291
292#[cfg(test)]
293mod tests {
294 use super::*;
295 use crate::tools::{
296 Xorshift,
297 testutil::{exhaustive_sequences, sample_usize},
298 };
299
300 #[test]
301 fn test_comparison_sorts() {
302 let mut rng = Xorshift::default();
303 let mut cases: Vec<_> = exhaustive_sequences(-1..=1, 0..=8).collect();
304 for n in sample_usize(&mut rng, 16, 0..=3000, 100) {
305 let bound = rng.random(0..=1000i32);
306 let values: Vec<_> = rng.random_iter(-bound..=bound).take(n).collect();
307 cases.push(vec![0; n]);
308 cases.push((0..n as i32).collect());
309 cases.push((0..n as i32).rev().collect());
310 cases.push(values);
311 }
312 for values in cases {
313 let mut expected = values.clone();
314 expected.sort();
315 let mut actual = values.clone();
316 actual.bubble_sort();
317 assert_eq!(actual, expected);
318 let mut actual = values.clone();
319 actual.merge_sort();
320 assert_eq!(actual, expected);
321 let mut actual = values;
322 actual.insertion_sort();
323 assert_eq!(actual, expected);
324 }
325 }
326
327 #[test]
328 fn test_large_comparison_sorts() {
329 let mut rng = Xorshift::default();
330 for n in sample_usize(&mut rng, 16, 0..=100_000, 10) {
331 let values: Vec<i32> = rng.random_iter(..).take(n).collect();
332 let mut expected = values.clone();
333 expected.sort();
334 let mut actual = values.clone();
335 actual.merge_sort();
336 assert_eq!(actual, expected);
337 let mut actual = values;
338 actual.insertion_sort();
339 assert_eq!(actual, expected);
340 }
341 }
342
343 #[test]
344 fn test_radix_sort() {
345 let mut rng = Xorshift::default();
346 macro_rules! test_types {
347 ($($t:ty),* $(,)?) => {
348 $(
349 for _ in 0..20 {
350 let n = rng.random(0..1000);
351 let mut actual: Vec<($t, usize)> =
352 rng.random_iter(..).take(n).zip(0..).collect();
353 let mut expected = actual.clone();
354 actual.radix_sort_by_key(|&(key, _)| key);
355 expected.sort_by_key(|&(key, _)| key);
356 assert_eq!(actual, expected);
357 }
358 )*
359 };
360 }
361 test_types!(u8, u16, u32, u64, u128, usize);
362 test_types!(i8, i16, i32, i64, i128, isize);
363 for _ in 0..20 {
364 let n = rng.random(0..1000);
365 let mut actual: Vec<_> = rng
366 .random_iter(0u32..64)
367 .take(n)
368 .zip(0..)
369 .map(|(key, index)| (key, index.to_string()))
370 .collect();
371 let mut expected = actual.clone();
372 actual.radix_sort_by_key(|value| value.0);
373 expected.sort_by_key(|value| value.0);
374 assert_eq!(actual, expected);
375 }
376 }
377}