Skip to content

Commit b3f113d

Browse files
tayfuryldzcclauss
andauthored
types: constrain tim sort items to Comparable (#15432)
Co-authored-by: TayfurYldz <238304586+TayfurYldz@users.noreply.github.com> Co-authored-by: Christian Clauss <cclauss@me.com>
1 parent 18c82c9 commit b3f113d

2 files changed

Lines changed: 20 additions & 8 deletions

File tree

‎sorts/tim_sort.py‎

Lines changed: 17 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,12 @@
1-
from typing import Any
1+
from collections.abc import Sequence
2+
from typing import Any, Protocol
23

34

4-
def binary_search(lst: list[Any], item: Any, start: int, end: int) -> int:
5+
class Comparable(Protocol):
6+
def __lt__(self, other: Any, /) -> bool: ...
7+
8+
9+
def binary_search[T: Comparable](lst: list[T], item: T, start: int, end: int) -> int:
510
""">>> binary_search([1, 3, 5], 4, 0, 2)
611
2
712
>>> binary_search([1, 3, 5], 0, 0, 2)
@@ -30,20 +35,20 @@ def binary_search(lst: list[Any], item: Any, start: int, end: int) -> int:
3035
Space: ``O(log n)`` due to recursion depth.
3136
"""
3237
if start == end:
33-
return start if lst[start] > item else start + 1
38+
return start if item < lst[start] else start + 1
3439
if start > end:
3540
return start
3641

3742
mid = (start + end) // 2
3843
if lst[mid] < item:
3944
return binary_search(lst, item, mid + 1, end)
40-
elif lst[mid] > item:
45+
elif item < lst[mid]:
4146
return binary_search(lst, item, start, mid - 1)
4247
else:
4348
return mid
4449

4550

46-
def insertion_sort(lst: list[Any]) -> list[Any]:
51+
def insertion_sort[T: Comparable](lst: list[T]) -> list[T]:
4752
""">>> insertion_sort([3, 2, 1])
4853
[1, 2, 3]
4954
@@ -74,7 +79,7 @@ def insertion_sort(lst: list[Any]) -> list[Any]:
7479
return lst
7580

7681

77-
def merge(left: list[Any], right: list[Any]) -> list[Any]:
82+
def merge[T: Comparable](left: list[T], right: list[T]) -> list[T]:
7883
""">>> merge([1, 4], [2, 3])
7984
[1, 2, 3, 4]
8085
@@ -104,7 +109,7 @@ def merge(left: list[Any], right: list[Any]) -> list[Any]:
104109
return [right[0], *merge(left, right[1:])]
105110

106111

107-
def tim_sort(lst: list[Any] | tuple[Any, ...] | str) -> list[Any]:
112+
def tim_sort[T: Comparable](lst: Sequence[T]) -> list[T]:
108113
"""
109114
Sort and return the input using a TimSort-like approach: detect
110115
runs, sort each run with insertion sort, then merge the runs.
@@ -125,14 +130,18 @@ def tim_sort(lst: list[Any] | tuple[Any, ...] | str) -> list[Any]:
125130
True
126131
>>> tim_sort([3, 2, 1]) == sorted([3, 2, 1])
127132
True
133+
>>> tim_sort([1, "a"])
134+
Traceback (most recent call last):
135+
...
136+
TypeError: '<' not supported between instances of 'str' and 'int'
128137
129138
"""
130139
if not lst:
131140
return []
132141
length = len(lst)
133142
runs, sorted_runs = [], []
134143
new_run = [lst[0]]
135-
sorted_array: list[Any] = []
144+
sorted_array: list[T] = []
136145
i = 1
137146
while i < length:
138147
if lst[i] < lst[i - 1]:

‎tests/test_sorts.py‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -51,6 +51,7 @@
5151
from sorts.shrink_shell_sort import shell_sort as shrink_shell_sort
5252
from sorts.stooge_sort import stooge_sort
5353
from sorts.strand_sort import strand_sort
54+
from sorts.tim_sort import tim_sort
5455
from sorts.unknown_sort import merge_sort as unknown_sort
5556

5657

@@ -90,6 +91,7 @@ def test_heap_sort() -> None:
9091
shrink_shell_sort,
9192
stooge_sort,
9293
strand_sort,
94+
tim_sort,
9395
unknown_sort,
9496
)
9597

@@ -163,6 +165,7 @@ def test_rec_insertion_sort(case) -> None:
163165
selection_sort,
164166
shrink_shell_sort,
165167
strand_sort,
168+
tim_sort,
166169
unknown_sort,
167170
],
168171
ids=lambda f: f.__name__,

0 commit comments

Comments
 (0)