Skip to content

Commit aeba150

Browse files
committed
fix(sorts): make smoothsort generic over Comparable items
1 parent 2525255 commit aeba150

2 files changed

Lines changed: 30 additions & 10 deletions

File tree

‎sorts/smoothsort.py‎

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

13+
from typing import Any, Protocol, TypeVar
14+
15+
16+
class Comparable(Protocol):
17+
def __lt__(self, other: Any, /) -> bool: ...
18+
19+
20+
T = TypeVar("T", bound=Comparable)
21+
22+
1323
# Precomputed Leonardo numbers: L(0)=1, L(1)=1, L(k)=L(k-1)+L(k-2)+1.
1424
# 46 values comfortably cover all practical list sizes.
1525
_LEONARDO: list[int] = [1, 1]
1626
while _LEONARDO[-1] < 2**31:
1727
_LEONARDO.append(_LEONARDO[-1] + _LEONARDO[-2] + 1)
1828

1929

20-
def _sift(seq: list[int], root: int, order: int) -> None:
30+
def _sift[T: Comparable](seq: list[T], root: int, order: int) -> None:
2131
"""
2232
Restore the max-heap property within a Leonardo tree of the given ``order``.
2333
@@ -59,20 +69,20 @@ def _sift(seq: list[int], root: int, order: int) -> None:
5969
right = root - 1 # right child root
6070
left = root - 1 - _LEONARDO[order - 2] # left child root
6171

62-
if seq[left] >= seq[right] and seq[left] > seq[root]:
72+
if not (seq[left] < seq[right]) and seq[root] < seq[left]:
6373
seq[root], seq[left] = seq[left], seq[root]
6474
root = left
6575
order -= 1
66-
elif seq[right] > seq[left] and seq[right] > seq[root]:
76+
elif seq[left] < seq[right] and seq[root] < seq[right]:
6777
seq[root], seq[right] = seq[right], seq[root]
6878
root = right
6979
order -= 2
7080
else:
7181
break
7282

7383

74-
def _trinkle(
75-
seq: list[int],
84+
def _trinkle[T: Comparable](
85+
seq: list[T],
7686
pos: int,
7787
heap_sizes: list[int],
7888
idx: int,
@@ -105,14 +115,14 @@ def _trinkle(
105115
"""
106116
while idx > 0:
107117
prev_root = pos - _LEONARDO[heap_sizes[idx]]
108-
if seq[pos] >= seq[prev_root]:
118+
if not (seq[pos] < seq[prev_root]):
109119
break
110-
# Only swap if prev_root is also >= its own children; otherwise
120+
# Only swap if prev_root is also > its own children; otherwise
111121
# moving it would break the heap on the left side.
112122
if heap_sizes[idx] > 1:
113123
right = pos - 1
114124
left = pos - 1 - _LEONARDO[heap_sizes[idx] - 2]
115-
if seq[prev_root] <= seq[right] or seq[prev_root] <= seq[left]:
125+
if not (seq[right] < seq[prev_root]) or not (seq[left] < seq[prev_root]):
116126
break
117127
seq[pos], seq[prev_root] = seq[prev_root], seq[pos]
118128
pos = prev_root
@@ -121,7 +131,7 @@ def _trinkle(
121131
_sift(seq, pos, heap_sizes[idx])
122132

123133

124-
def smoothsort(seq: list[int]) -> list[int]:
134+
def smoothsort[T: Comparable](seq: list[T]) -> list[T]:
125135
"""
126136
Sort a list in-place using the Smoothsort algorithm and return it.
127137
@@ -131,7 +141,7 @@ def smoothsort(seq: list[int]) -> list[int]:
131141
whose structure mirrors the sorted prefix of the sequence.
132142
133143
Args:
134-
seq: A list of integers to sort.
144+
seq: A list of mutually comparable items to sort.
135145
136146
Returns:
137147
The same list object, sorted in ascending order.
@@ -147,10 +157,18 @@ def smoothsort(seq: list[int]) -> list[int]:
147157
[1, 2, 3, 4, 5]
148158
>>> smoothsort([3, 3, 2, 1, 2])
149159
[1, 2, 2, 3, 3]
160+
>>> smoothsort(["d", "a", "c", "b"])
161+
['a', 'b', 'c', 'd']
162+
>>> smoothsort([2.5, -1, 0.0])
163+
[-1, 0.0, 2.5]
150164
>>> smoothsort([1, 2, 3, 4, 5])
151165
[1, 2, 3, 4, 5]
152166
>>> smoothsort([-3, 0, -1, 5, 2])
153167
[-3, -1, 0, 2, 5]
168+
>>> smoothsort([1, "a"])
169+
Traceback (most recent call last):
170+
...
171+
TypeError: '<' not supported between instances of 'str' and 'int'
154172
"""
155173
n = len(seq)
156174
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)