Skip to content

Commit f59bddb

Browse files
wanjinhao1cclauss
andauthored
fix(sorts): make smoothsort generic over Comparable items (#15441)
* fix(sorts): make smoothsort generic over Comparable items * The TypeVar is no longer needed. --------- Co-authored-by: Christian Clauss <cclauss@me.com>
1 parent 2525255 commit f59bddb

2 files changed

Lines changed: 27 additions & 10 deletions

File tree

‎sorts/smoothsort.py‎

Lines changed: 25 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -10,14 +10,21 @@
1010
https://www.cs.utexas.edu/~EWD/ewd07xx/EWD796a.PDF
1111
"""
1212

13+
from typing import Any, Protocol
14+
15+
16+
class Comparable(Protocol):
17+
def __lt__(self, other: Any, /) -> bool: ...
18+
19+
1320
# Precomputed Leonardo numbers: L(0)=1, L(1)=1, L(k)=L(k-1)+L(k-2)+1.
1421
# 46 values comfortably cover all practical list sizes.
1522
_LEONARDO: list[int] = [1, 1]
1623
while _LEONARDO[-1] < 2**31:
1724
_LEONARDO.append(_LEONARDO[-1] + _LEONARDO[-2] + 1)
1825

1926

20-
def _sift(seq: list[int], root: int, order: int) -> None:
27+
def _sift[T: Comparable](seq: list[T], root: int, order: int) -> None:
2128
"""
2229
Restore the max-heap property within a Leonardo tree of the given ``order``.
2330
@@ -59,20 +66,20 @@ def _sift(seq: list[int], root: int, order: int) -> None:
5966
right = root - 1 # right child root
6067
left = root - 1 - _LEONARDO[order - 2] # left child root
6168

62-
if seq[left] >= seq[right] and seq[left] > seq[root]:
69+
if not (seq[left] < seq[right]) and seq[root] < seq[left]:
6370
seq[root], seq[left] = seq[left], seq[root]
6471
root = left
6572
order -= 1
66-
elif seq[right] > seq[left] and seq[right] > seq[root]:
73+
elif seq[left] < seq[right] and seq[root] < seq[right]:
6774
seq[root], seq[right] = seq[right], seq[root]
6875
root = right
6976
order -= 2
7077
else:
7178
break
7279

7380

74-
def _trinkle(
75-
seq: list[int],
81+
def _trinkle[T: Comparable](
82+
seq: list[T],
7683
pos: int,
7784
heap_sizes: list[int],
7885
idx: int,
@@ -105,14 +112,14 @@ def _trinkle(
105112
"""
106113
while idx > 0:
107114
prev_root = pos - _LEONARDO[heap_sizes[idx]]
108-
if seq[pos] >= seq[prev_root]:
115+
if not (seq[pos] < seq[prev_root]):
109116
break
110-
# Only swap if prev_root is also >= its own children; otherwise
117+
# Only swap if prev_root is also > its own children; otherwise
111118
# moving it would break the heap on the left side.
112119
if heap_sizes[idx] > 1:
113120
right = pos - 1
114121
left = pos - 1 - _LEONARDO[heap_sizes[idx] - 2]
115-
if seq[prev_root] <= seq[right] or seq[prev_root] <= seq[left]:
122+
if not (seq[right] < seq[prev_root]) or not (seq[left] < seq[prev_root]):
116123
break
117124
seq[pos], seq[prev_root] = seq[prev_root], seq[pos]
118125
pos = prev_root
@@ -121,7 +128,7 @@ def _trinkle(
121128
_sift(seq, pos, heap_sizes[idx])
122129

123130

124-
def smoothsort(seq: list[int]) -> list[int]:
131+
def smoothsort[T: Comparable](seq: list[T]) -> list[T]:
125132
"""
126133
Sort a list in-place using the Smoothsort algorithm and return it.
127134
@@ -131,7 +138,7 @@ def smoothsort(seq: list[int]) -> list[int]:
131138
whose structure mirrors the sorted prefix of the sequence.
132139
133140
Args:
134-
seq: A list of integers to sort.
141+
seq: A list of mutually comparable items to sort.
135142
136143
Returns:
137144
The same list object, sorted in ascending order.
@@ -147,10 +154,18 @@ def smoothsort(seq: list[int]) -> list[int]:
147154
[1, 2, 3, 4, 5]
148155
>>> smoothsort([3, 3, 2, 1, 2])
149156
[1, 2, 2, 3, 3]
157+
>>> smoothsort(["d", "a", "c", "b"])
158+
['a', 'b', 'c', 'd']
159+
>>> smoothsort([2.5, -1, 0.0])
160+
[-1, 0.0, 2.5]
150161
>>> smoothsort([1, 2, 3, 4, 5])
151162
[1, 2, 3, 4, 5]
152163
>>> smoothsort([-3, 0, -1, 5, 2])
153164
[-3, -1, 0, 2, 5]
165+
>>> smoothsort([1, "a"])
166+
Traceback (most recent call last):
167+
...
168+
TypeError: '<' not supported between instances of 'str' and 'int'
154169
"""
155170
n = len(seq)
156171
if n < 2:

‎tests/test_sorts.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -51,6 +51,7 @@
5151
from sorts.selection_sort import selection_sort
5252
from sorts.shell_sort import shell_sort
5353
from sorts.shrink_shell_sort import shell_sort as shrink_shell_sort
54+
from sorts.smoothsort import smoothsort
5455
from sorts.stooge_sort import stooge_sort
5556
from sorts.strand_sort import strand_sort
5657
from sorts.tim_sort import tim_sort
@@ -92,6 +93,7 @@ def test_heap_sort() -> None:
9293
selection_sort,
9394
shell_sort,
9495
shrink_shell_sort,
96+
smoothsort,
9597
stooge_sort,
9698
strand_sort,
9799
three_way_radix_quicksort,

0 commit comments

Comments
 (0)