Skip to main content

rustpython_vm/
sorting.rs

1// TODO: MERGESTATE_TEMP_SIZE unused — buf is a dynamic Vec, not a fixed stack array.
2const MIN_GALLOP: usize = 7;
3const MAX_MINRUN: usize = 64;
4
5enum LoBreakout {
6    Succeed,
7    CopyB,
8}
9
10enum HiBreakout {
11    Succeed,
12    CopyA,
13}
14
15#[derive(Clone, Copy)]
16struct Run {
17    base: usize,
18    len: usize,
19    power: u32,
20}
21
22struct MergeState<T> {
23    buf: Vec<T>,
24    min_gallop: usize,
25    pending: Vec<Run>,
26}
27
28impl<T: Clone> MergeState<T> {
29    fn merge_lo<E, F>(
30        &mut self,
31        values: &mut [T],
32        is_lt: &mut F,
33        start_a: usize,
34        mut len_a: usize,
35        start_b: usize,
36        mut len_b: usize,
37    ) -> Result<(), E>
38    where
39        F: FnMut(&T, &T) -> Result<bool, E>,
40    {
41        debug_assert!(len_a > 0);
42        debug_assert!(len_b > 0);
43        debug_assert!(start_a + len_a == start_b);
44
45        self.buf.clear();
46        self.buf
47            .extend_from_slice(&values[start_a..start_a + len_a]);
48
49        let mut cursor_a = 0;
50        let mut cursor_b = start_b;
51        let mut dest = start_a;
52
53        values[dest] = values[cursor_b].clone();
54        dest += 1;
55        cursor_b += 1;
56        len_b -= 1;
57
58        if len_b == 0 {
59            values[dest..dest + len_a].clone_from_slice(&self.buf[cursor_a..cursor_a + len_a]);
60            return Ok(());
61        }
62        if len_a == 1 {
63            copy_within_clone(values, cursor_b, dest, len_b);
64            values[dest + len_b] = self.buf[cursor_a].clone();
65            return Ok(());
66        }
67
68        let mut min_gallop = self.min_gallop;
69
70        let breakout: Result<LoBreakout, E> = 'merging: loop {
71            let mut a_count = 0;
72            let mut b_count = 0;
73
74            loop {
75                let b_wins = match is_lt(&values[cursor_b], &self.buf[cursor_a]) {
76                    Ok(v) => v,
77                    Err(e) => break 'merging Err(e),
78                };
79                if b_wins {
80                    values[dest] = values[cursor_b].clone();
81                    dest += 1;
82                    cursor_b += 1;
83                    len_b -= 1;
84                    b_count += 1;
85                    a_count = 0;
86                    if len_b == 0 {
87                        break 'merging Ok(LoBreakout::Succeed);
88                    }
89                    if b_count >= min_gallop {
90                        break;
91                    }
92                } else {
93                    values[dest] = self.buf[cursor_a].clone();
94                    dest += 1;
95                    cursor_a += 1;
96                    len_a -= 1;
97                    a_count += 1;
98                    b_count = 0;
99                    if len_a == 1 {
100                        break 'merging Ok(LoBreakout::CopyB);
101                    }
102                    if a_count >= min_gallop {
103                        break;
104                    }
105                }
106            }
107
108            min_gallop += 1;
109            loop {
110                if min_gallop > 1 {
111                    min_gallop -= 1;
112                }
113                self.min_gallop = min_gallop;
114                let mut k =
115                    match gallop_right(&self.buf, is_lt, &values[cursor_b], cursor_a, len_a, 0) {
116                        Ok(k) => k,
117                        Err(e) => break 'merging Err(e),
118                    };
119                a_count = k;
120                if k > 0 {
121                    values[dest..dest + k].clone_from_slice(&self.buf[cursor_a..cursor_a + k]);
122                    dest += k;
123                    cursor_a += k;
124                    len_a -= k;
125                    if len_a == 1 {
126                        break 'merging Ok(LoBreakout::CopyB);
127                    }
128                    if len_a == 0 {
129                        break 'merging Ok(LoBreakout::Succeed);
130                    }
131                }
132                values[dest] = values[cursor_b].clone();
133                dest += 1;
134                cursor_b += 1;
135                len_b -= 1;
136                if len_b == 0 {
137                    break 'merging Ok(LoBreakout::Succeed);
138                }
139                k = match gallop_left(values, is_lt, &self.buf[cursor_a], cursor_b, len_b, 0) {
140                    Ok(k) => k,
141                    Err(e) => break 'merging Err(e),
142                };
143                b_count = k;
144                if k > 0 {
145                    copy_within_clone(values, cursor_b, dest, k);
146                    dest += k;
147                    cursor_b += k;
148                    len_b -= k;
149                    if len_b == 0 {
150                        break 'merging Ok(LoBreakout::Succeed);
151                    }
152                }
153                values[dest] = self.buf[cursor_a].clone();
154                dest += 1;
155                cursor_a += 1;
156                len_a -= 1;
157                if len_a == 1 {
158                    break 'merging Ok(LoBreakout::CopyB);
159                }
160                if a_count < MIN_GALLOP && b_count < MIN_GALLOP {
161                    break;
162                }
163            }
164
165            min_gallop += 1;
166            self.min_gallop = min_gallop;
167        };
168
169        match breakout {
170            Ok(LoBreakout::CopyB) => {
171                copy_within_clone(values, cursor_b, dest, len_b);
172                values[dest + len_b] = self.buf[cursor_a].clone();
173                Ok(())
174            }
175            other => {
176                if len_a > 0 {
177                    values[dest..dest + len_a]
178                        .clone_from_slice(&self.buf[cursor_a..cursor_a + len_a]);
179                }
180                other.map(|_| ())
181            }
182        }
183    }
184
185    fn merge_hi<E, F>(
186        &mut self,
187        values: &mut [T],
188        is_lt: &mut F,
189        start_a: usize,
190        mut len_a: usize,
191        start_b: usize,
192        mut len_b: usize,
193    ) -> Result<(), E>
194    where
195        F: FnMut(&T, &T) -> Result<bool, E>,
196    {
197        debug_assert!(len_a > 0);
198        debug_assert!(len_b > 0);
199        debug_assert!(start_a + len_a == start_b);
200
201        self.buf.clear();
202        self.buf
203            .extend_from_slice(&values[start_b..start_b + len_b]);
204
205        let mut dest = start_b + len_b - 1;
206        let mut cursor_a = start_a + len_a - 1;
207        let mut cursor_b = len_b - 1;
208
209        values[dest] = values[cursor_a].clone();
210        dest -= 1;
211        cursor_a -= 1;
212        len_a -= 1;
213
214        if len_a == 0 {
215            values[(dest - len_b + 1)..=dest].clone_from_slice(&self.buf[0..len_b]);
216            return Ok(());
217        }
218        if len_b == 1 {
219            let src = cursor_a + 1 - len_a;
220            let dst = dest + 1 - len_a;
221            copy_within_clone(values, src, dst, len_a);
222            values[dst - 1] = self.buf[cursor_b].clone();
223            return Ok(());
224        }
225
226        let mut min_gallop = self.min_gallop;
227        let breakout: Result<HiBreakout, E> = 'merging: loop {
228            let mut a_count = 0;
229            let mut b_count = 0;
230
231            loop {
232                let b_wins = match is_lt(&self.buf[cursor_b], &values[cursor_a]) {
233                    Ok(v) => v,
234                    Err(e) => break 'merging Err(e),
235                };
236                if b_wins {
237                    values[dest] = values[cursor_a].clone();
238                    dest -= 1;
239                    len_a -= 1;
240
241                    if len_a == 0 {
242                        break 'merging Ok(HiBreakout::Succeed);
243                    }
244
245                    cursor_a -= 1;
246                    a_count += 1;
247                    b_count = 0;
248
249                    if a_count >= min_gallop {
250                        break;
251                    }
252                } else {
253                    values[dest] = self.buf[cursor_b].clone();
254                    dest -= 1;
255                    cursor_b -= 1;
256                    len_b -= 1;
257                    b_count += 1;
258                    a_count = 0;
259                    if len_b == 1 {
260                        break 'merging Ok(HiBreakout::CopyA);
261                    }
262                    if b_count >= min_gallop {
263                        break;
264                    }
265                }
266            }
267
268            min_gallop += 1;
269            loop {
270                if min_gallop > 1 {
271                    min_gallop -= 1;
272                }
273                self.min_gallop = min_gallop;
274                let mut k = match gallop_right(
275                    values,
276                    is_lt,
277                    &self.buf[cursor_b],
278                    start_a,
279                    len_a,
280                    len_a - 1,
281                ) {
282                    Ok(k) => k,
283                    Err(e) => break 'merging Err(e),
284                };
285                k = len_a - k;
286                a_count = k;
287                if k > 0 {
288                    copy_within_clone(values, cursor_a + 1 - k, dest + 1 - k, k);
289                    dest -= k;
290                    len_a -= k;
291                    if len_a == 0 {
292                        break 'merging Ok(HiBreakout::Succeed);
293                    }
294                    cursor_a -= k;
295                }
296                values[dest] = self.buf[cursor_b].clone();
297                dest -= 1;
298                cursor_b -= 1;
299                len_b -= 1;
300                if len_b == 1 {
301                    break 'merging Ok(HiBreakout::CopyA);
302                }
303                k = match gallop_left(&self.buf, is_lt, &values[cursor_a], 0, len_b, len_b - 1) {
304                    Ok(k) => k,
305                    Err(e) => break 'merging Err(e),
306                };
307                k = len_b - k;
308                b_count = k;
309                if k > 0 {
310                    values[dest + 1 - k..=dest]
311                        .clone_from_slice(&self.buf[cursor_b + 1 - k..=cursor_b]);
312                    dest -= k;
313                    len_b -= k;
314
315                    if len_b == 0 {
316                        break 'merging Ok(HiBreakout::Succeed);
317                    }
318                    cursor_b -= k;
319
320                    if len_b == 1 {
321                        break 'merging Ok(HiBreakout::CopyA);
322                    }
323                }
324                values[dest] = values[cursor_a].clone();
325                dest -= 1;
326                len_a -= 1;
327
328                if len_a == 0 {
329                    break 'merging Ok(HiBreakout::Succeed);
330                }
331
332                cursor_a -= 1;
333
334                if a_count < MIN_GALLOP && b_count < MIN_GALLOP {
335                    break;
336                }
337            }
338            min_gallop += 1;
339            self.min_gallop = min_gallop;
340        };
341
342        match breakout {
343            Ok(HiBreakout::CopyA) => {
344                let src = cursor_a + 1 - len_a;
345                let dst = dest + 1 - len_a;
346                copy_within_clone(values, src, dst, len_a);
347                values[dst - 1] = self.buf[cursor_b].clone();
348                Ok(())
349            }
350            other => {
351                if len_b > 0 {
352                    values[(dest + 1) - len_b..=dest].clone_from_slice(&self.buf[0..len_b]);
353                }
354                other.map(|_| ())
355            }
356        }
357    }
358
359    fn merge_at<E, F>(&mut self, values: &mut [T], is_lt: &mut F, i: usize) -> Result<(), E>
360    where
361        F: FnMut(&T, &T) -> Result<bool, E>,
362    {
363        debug_assert!(self.pending.len() >= 2);
364        debug_assert!(i == self.pending.len() - 2 || i == self.pending.len() - 3);
365
366        let mut start_a = self.pending[i].base;
367        let mut len_a = self.pending[i].len;
368        let start_b = self.pending[i + 1].base;
369        let mut len_b = self.pending[i + 1].len;
370
371        debug_assert!(len_a > 0);
372        debug_assert!(len_b > 0);
373        debug_assert!(start_a + len_a == start_b);
374
375        self.pending[i].len = len_a + len_b;
376        self.pending.remove(i + 1);
377
378        let k = gallop_right(values, is_lt, &values[start_b], start_a, len_a, 0)?;
379        start_a += k;
380        len_a -= k;
381
382        if len_a == 0 {
383            return Ok(());
384        }
385
386        len_b = gallop_left(
387            values,
388            is_lt,
389            &values[start_a + len_a - 1],
390            start_b,
391            len_b,
392            len_b - 1,
393        )?;
394
395        if len_b == 0 {
396            return Ok(());
397        }
398
399        if len_a <= len_b {
400            self.merge_lo(values, is_lt, start_a, len_a, start_b, len_b)?;
401        } else {
402            self.merge_hi(values, is_lt, start_a, len_a, start_b, len_b)?;
403        }
404        Ok(())
405    }
406
407    fn found_new_run<E, F>(
408        &mut self,
409        new_run_len: usize,
410        values: &mut [T],
411        is_lt: &mut F,
412    ) -> Result<(), E>
413    where
414        F: FnMut(&T, &T) -> Result<bool, E>,
415    {
416        if !self.pending.is_empty() {
417            let last = self.pending.len() - 1;
418            let s1 = self.pending[last].base;
419            let n1 = self.pending[last].len;
420            let power = powerloop(s1, n1, new_run_len, values.len());
421
422            while self.pending.len() > 1 && self.pending[self.pending.len() - 2].power > power {
423                self.merge_at(values, is_lt, self.pending.len() - 2)?;
424            }
425
426            debug_assert!(
427                self.pending.len() < 2 || self.pending[self.pending.len() - 2].power < power
428            );
429            let last = self.pending.len() - 1;
430            self.pending[last].power = power;
431        }
432        Ok(())
433    }
434
435    fn push_run(&mut self, base: usize, len: usize) {
436        self.pending.push(Run {
437            base,
438            len,
439            power: 0,
440        })
441    }
442
443    fn merge_force_collapse<E, F>(&mut self, values: &mut [T], is_lt: &mut F) -> Result<(), E>
444    where
445        F: FnMut(&T, &T) -> Result<bool, E>,
446    {
447        while self.pending.len() > 1 {
448            let mut n = self.pending.len() - 2;
449            if n > 0 && self.pending[n - 1].len < self.pending[n + 1].len {
450                n -= 1;
451            }
452            self.merge_at(values, is_lt, n)?;
453        }
454        Ok(())
455    }
456}
457
458fn binary_insertion_sort<T, E, F>(values: &mut [T], is_lt: &mut F, start: usize) -> Result<(), E>
459where
460    F: FnMut(&T, &T) -> Result<bool, E>,
461{
462    for i in start..values.len() {
463        let mut l = 0;
464        let mut r = i;
465
466        while l < r {
467            let m = (l + r) / 2;
468            if is_lt(&values[i], &values[m])? {
469                r = m;
470            } else {
471                l = m + 1;
472            }
473        }
474        values[l..=i].rotate_right(1);
475    }
476    Ok(())
477}
478
479fn copy_within_clone<T: Clone>(values: &mut [T], src: usize, dest: usize, n: usize) {
480    if dest <= src {
481        for k in 0..n {
482            values[dest + k] = values[src + k].clone();
483        }
484    } else {
485        for k in (0..n).rev() {
486            values[dest + k] = values[src + k].clone();
487        }
488    }
489}
490
491fn count_run<T, E, F>(values: &[T], is_lt: &mut F) -> Result<(usize, bool), E>
492where
493    F: FnMut(&T, &T) -> Result<bool, E>,
494{
495    let n = values.len();
496    if n == 1 {
497        return Ok((1, false));
498    }
499    let mut i = 2;
500    let descending = is_lt(&values[1], &values[0])?;
501    if descending {
502        while i < n && is_lt(&values[i], &values[i - 1])? {
503            i += 1;
504        }
505    } else {
506        while i < n && !is_lt(&values[i], &values[i - 1])? {
507            i += 1;
508        }
509    }
510    Ok((i, descending))
511}
512
513fn gallop_left<T, E, F>(
514    values: &[T],
515    is_lt: &mut F,
516    key: &T,
517    base: usize,
518    len: usize,
519    hint: usize,
520) -> Result<usize, E>
521where
522    F: FnMut(&T, &T) -> Result<bool, E>,
523{
524    debug_assert!(hint < len);
525    let mut lastofs: isize = 0;
526    let mut ofs: isize = 1;
527    let hint_i = hint as isize;
528    let len_i = len as isize;
529
530    if is_lt(&values[base + hint], key)? {
531        let maxofs = len_i - hint_i;
532        while ofs < maxofs && is_lt(&values[base + hint + ofs as usize], key)? {
533            lastofs = ofs;
534            ofs = (ofs * 2) + 1;
535        }
536        if ofs > maxofs {
537            ofs = maxofs;
538        }
539        lastofs += hint_i;
540        ofs += hint_i;
541    } else {
542        let maxofs = hint_i + 1;
543        while ofs < maxofs && !is_lt(&values[base + (hint_i - ofs) as usize], key)? {
544            lastofs = ofs;
545            ofs = (ofs * 2) + 1;
546        }
547        if ofs > maxofs {
548            ofs = maxofs;
549        }
550        (lastofs, ofs) = (hint_i - ofs, hint_i - lastofs);
551    }
552    lastofs += 1;
553    while lastofs < ofs {
554        let m = lastofs + ((ofs - lastofs) / 2);
555        if is_lt(&values[base + m as usize], key)? {
556            lastofs = m + 1;
557        } else {
558            ofs = m;
559        }
560    }
561    Ok(ofs as usize)
562}
563
564fn gallop_right<T, E, F>(
565    values: &[T],
566    is_lt: &mut F,
567    key: &T,
568    base: usize,
569    len: usize,
570    hint: usize,
571) -> Result<usize, E>
572where
573    F: FnMut(&T, &T) -> Result<bool, E>,
574{
575    debug_assert!(hint < len);
576    let mut lastofs: isize = 0;
577    let mut ofs: isize = 1;
578    let hint_i = hint as isize;
579    let len_i = len as isize;
580
581    if is_lt(key, &values[base + hint])? {
582        let maxofs = hint_i + 1;
583        while ofs < maxofs && is_lt(key, &values[base + (hint_i - ofs) as usize])? {
584            lastofs = ofs;
585            ofs = (ofs * 2) + 1;
586        }
587        if ofs > maxofs {
588            ofs = maxofs;
589        }
590        (lastofs, ofs) = (hint_i - ofs, hint_i - lastofs);
591    } else {
592        let maxofs = len_i - hint_i;
593        while ofs < maxofs && !is_lt(key, &values[base + hint + ofs as usize])? {
594            lastofs = ofs;
595            ofs = (ofs * 2) + 1;
596        }
597        if ofs > maxofs {
598            ofs = maxofs;
599        }
600        lastofs += hint_i;
601        ofs += hint_i;
602    }
603    lastofs += 1;
604    while lastofs < ofs {
605        let m = lastofs + ((ofs - lastofs) / 2);
606        if is_lt(key, &values[base + m as usize])? {
607            ofs = m;
608        } else {
609            lastofs = m + 1;
610        }
611    }
612    Ok(ofs as usize)
613}
614
615// TODO: consider CPython 3.12+'s incremental minrun (mr_current/mr_e/mr_mask)
616//       for a more precise minrun; current bit-shift version is the classic one.
617fn merge_compute_minrun(mut n: usize) -> usize {
618    let mut r = 0;
619    while n >= MAX_MINRUN {
620        r |= n & 1;
621        n >>= 1;
622    }
623    n + r
624}
625
626fn powerloop(s1: usize, n1: usize, n2: usize, n: usize) -> u32 {
627    let mut result: u32 = 0;
628    let mut a = 2 * s1 + n1;
629    let mut b = a + n1 + n2;
630
631    loop {
632        result += 1;
633        if a >= n {
634            debug_assert!(b >= a);
635            a -= n;
636            b -= n;
637        } else if b >= n {
638            break;
639        }
640        debug_assert!(a < b && b < n);
641        a <<= 1;
642        b <<= 1;
643    }
644    result
645}
646
647/// Stable adaptive mergesort (Tim Peters' timsort with powersort's
648/// merge-ordering policy, matching CPython 3.11+). `is_lt` provides comparison.
649pub fn timsort<T, E, F>(values: &mut [T], is_lt: &mut F) -> Result<(), E>
650where
651    T: Clone,
652    F: FnMut(&T, &T) -> Result<bool, E>,
653{
654    let n = values.len();
655    let mut ms = MergeState {
656        buf: Vec::new(),
657        min_gallop: MIN_GALLOP,
658        pending: Vec::new(),
659    };
660
661    if n < 2 {
662        return Ok(());
663    }
664
665    if n < MAX_MINRUN {
666        let (l, desc) = count_run(values, is_lt)?;
667        if desc {
668            values[0..l].reverse();
669        }
670        binary_insertion_sort(values, is_lt, l)?;
671        return Ok(());
672    }
673
674    let minrun = merge_compute_minrun(n);
675    let mut lo = 0;
676
677    while lo < n {
678        let (mut l, desc) = count_run(&values[lo..n], is_lt)?;
679        if desc {
680            values[lo..lo + l].reverse();
681        }
682        if l < minrun {
683            let force = minrun.min(n - lo);
684            binary_insertion_sort(&mut values[lo..lo + force], is_lt, l)?;
685            l = force;
686        }
687        ms.found_new_run(l, values, is_lt)?;
688        ms.push_run(lo, l);
689        lo += l;
690    }
691    ms.merge_force_collapse(values, is_lt)?;
692    debug_assert!(ms.pending.len() == 1 && ms.pending[0].len == n);
693    Ok(())
694}
695
696#[cfg(test)]
697mod tests {
698    use super::*;
699
700    fn sort(mut v: Vec<i32>) -> Vec<i32> {
701        timsort(&mut v, &mut |a: &i32, b: &i32| Ok::<bool, ()>(a < b)).unwrap();
702        v
703    }
704
705    #[test]
706    fn basic_examples() {
707        assert_eq!(sort(vec![3, 1, 2]), vec![1, 2, 3]);
708        assert_eq!(sort(Vec::<i32>::new()), Vec::<i32>::new());
709        assert_eq!(sort(vec![1]), vec![1]);
710        assert_eq!(sort(vec![2, 1]), vec![1, 2]);
711    }
712
713    #[test]
714    fn five_elements_forwards_and_backwards() {
715        assert_eq!(sort(vec![1, 2, 3, 4, 5]), vec![1, 2, 3, 4, 5]);
716        assert_eq!(sort(vec![5, 4, 3, 2, 1]), vec![1, 2, 3, 4, 5]);
717    }
718
719    #[test]
720    fn six_elements_with_duplicates() {
721        assert_eq!(sort(vec![3, 1, 3, 1, 2, 2]), vec![1, 1, 2, 2, 3, 3]);
722    }
723
724    #[test]
725    fn one_thousand_elements() {
726        let v: Vec<i32> = (0..1000).rev().collect(); // 999..0
727        let sorted: Vec<i32> = (0..1000).collect();
728        assert_eq!(sort(v), sorted);
729    }
730
731    #[test]
732    fn pseudorandom_collection() {
733        let v: Vec<i32> = (0..500).map(|i| (i * 7919) % 500).collect();
734        let mut expected = v.clone();
735        expected.sort();
736        assert_eq!(sort(v), expected);
737    }
738}