Skip to content

Commit 0b32b10

Browse files
committed
types(intro_sort): constrain items to Comparable
Bind every function in sorts/intro_sort.py to a Comparable Protocol so the signatures say "a list of items that can be compared with each other" instead of a bare list, and keep the element type in the return. Two of the hints were outright wrong rather than merely loose: - median_of_3 returned int but returns an element of the collection, so it is now T - partition took pivot: int but takes an element of the collection, so it is now T Both were only reachable with the wrong type through untyped callers. Doctests added for a comparable non-int type (strings) and for the failure mode on insertion_sort, heap_sort, median_of_3, partition and sort: mixing non-comparable items must raise TypeError rather than silently mis-sort. intro_sort had no test at all. tests/test_sorts.py now adds it to the shared battery and to the rejection check, plus a dedicated test that reaches the branches the battery cannot: every shared case is shorter than the 16-element threshold, so the battery only ever exercises insertion_sort. The new test sorts 17/32/100/500-element int and str inputs to take the quicksort branch, and drives intro_sort with a depth budget of 0 to take the heapsort branch. The RNG is seeded so the test is deterministic.
1 parent db6bf8a commit 0b32b10

2 files changed

Lines changed: 86 additions & 11 deletions

File tree

‎sorts/intro_sort.py‎

Lines changed: 52 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -2,12 +2,25 @@
22
Introspective Sort is a hybrid sort (Quick Sort + Heap Sort + Insertion Sort)
33
if the size of the list is under 16, use insertion sort
44
https://en.wikipedia.org/wiki/Introsort
5+
6+
For doctests run following command:
7+
python3 -m doctest -v intro_sort.py
8+
9+
For manual testing run:
10+
python3 intro_sort.py
511
"""
612

713
import math
14+
from typing import Any, Protocol
15+
16+
17+
class Comparable(Protocol):
18+
def __lt__(self, other: Any, /) -> bool: ...
819

920

10-
def insertion_sort(array: list, start: int = 0, end: int = 0) -> list:
21+
def insertion_sort[T: Comparable](
22+
array: list[T], start: int = 0, end: int = 0
23+
) -> list[T]:
1124
"""
1225
>>> array = [4, 2, 6, 8, 1, 7, 8, 22, 14, 56, 27, 79, 23, 45, 14, 12]
1326
>>> insertion_sort(array, 0, len(array))
@@ -24,6 +37,10 @@ def insertion_sort(array: list, start: int = 0, end: int = 0) -> list:
2437
>>> array = [73.568, 73.56, -45.03, 1.7, 0, 89.45]
2538
>>> insertion_sort(array, 0, len(array))
2639
[-45.03, 0, 1.7, 73.56, 73.568, 89.45]
40+
>>> insertion_sort([1, "a"]) # doctest: +ELLIPSIS
41+
Traceback (most recent call last):
42+
...
43+
TypeError: ...
2744
"""
2845
end = end or len(array)
2946
for i in range(start, end):
@@ -36,7 +53,9 @@ def insertion_sort(array: list, start: int = 0, end: int = 0) -> list:
3653
return array
3754

3855

39-
def heapify(array: list, index: int, heap_size: int) -> None: # Max Heap
56+
def heapify[T: Comparable](
57+
array: list[T], index: int, heap_size: int
58+
) -> None: # Max Heap
4059
"""
4160
>>> array = [4, 2, 6, 8, 1, 7, 8, 22, 14, 56, 27, 79, 23, 45, 14, 12]
4261
>>> heapify(array, len(array) // 2, len(array))
@@ -56,7 +75,7 @@ def heapify(array: list, index: int, heap_size: int) -> None: # Max Heap
5675
heapify(array, largest, heap_size)
5776

5877

59-
def heap_sort(array: list) -> list:
78+
def heap_sort[T: Comparable](array: list[T]) -> list[T]:
6079
"""
6180
>>> heap_sort([4, 2, 6, 8, 1, 7, 8, 22, 14, 56, 27, 79, 23, 45, 14, 12])
6281
[1, 2, 4, 6, 7, 8, 8, 12, 14, 14, 22, 23, 27, 45, 56, 79]
@@ -66,6 +85,12 @@ def heap_sort(array: list) -> list:
6685
['b', 'b', 'd', 'e', 'e', 'f', 'g', 'p', 's', 'u', 'v', 'x', 'z']
6786
>>> heap_sort([6.2, -45.54, 8465.20, 758.56, -457.0, 0, 1, 2.879, 1.7, 11.7])
6887
[-457.0, -45.54, 0, 1, 1.7, 2.879, 6.2, 11.7, 758.56, 8465.2]
88+
>>> heap_sort(["banana", "apple", "cherry"])
89+
['apple', 'banana', 'cherry']
90+
>>> heap_sort([1, "a"]) # doctest: +ELLIPSIS
91+
Traceback (most recent call last):
92+
...
93+
TypeError: ...
6994
"""
7095
n = len(array)
7196

@@ -79,9 +104,9 @@ def heap_sort(array: list) -> list:
79104
return array
80105

81106

82-
def median_of_3(
83-
array: list, first_index: int, middle_index: int, last_index: int
84-
) -> int:
107+
def median_of_3[T: Comparable](
108+
array: list[T], first_index: int, middle_index: int, last_index: int
109+
) -> T:
85110
"""
86111
>>> array = [4, 2, 6, 8, 1, 7, 8, 22, 14, 56, 27, 79, 23, 45, 14, 12]
87112
>>> median_of_3(array, 0, ((len(array) - 0) // 2) + 1, len(array) - 1)
@@ -92,6 +117,12 @@ def median_of_3(
92117
>>> array = [4, 2, 6, 8, 1, 7, 8, 22, 15, 14, 27, 79, 23, 45, 14, 16]
93118
>>> median_of_3(array, 0, ((len(array) - 0) // 2) + 1, len(array) - 1)
94119
14
120+
>>> median_of_3(["b", "m", "z"], 0, 1, 2)
121+
'm'
122+
>>> median_of_3([1, "a", 2], 0, 1, 2) # doctest: +ELLIPSIS
123+
Traceback (most recent call last):
124+
...
125+
TypeError: ...
95126
"""
96127
if (array[first_index] > array[middle_index]) != (
97128
array[first_index] > array[last_index]
@@ -105,7 +136,7 @@ def median_of_3(
105136
return array[last_index]
106137

107138

108-
def partition(array: list, low: int, high: int, pivot: int) -> int:
139+
def partition[T: Comparable](array: list[T], low: int, high: int, pivot: T) -> int:
109140
"""
110141
>>> array = [4, 2, 6, 8, 1, 7, 8, 22, 14, 56, 27, 79, 23, 45, 14, 12]
111142
>>> partition(array, 0, len(array), 12)
@@ -119,6 +150,10 @@ def partition(array: list, low: int, high: int, pivot: int) -> int:
119150
>>> array = [6.2, -45.54, 8465.20, 758.56, -457.0, 0, 1, 2.879, 1.7, 11.7]
120151
>>> partition(array, 0, len(array), 2.879)
121152
6
153+
>>> partition([1, "a", 2], 0, 3, 2) # doctest: +ELLIPSIS
154+
Traceback (most recent call last):
155+
...
156+
TypeError: ...
122157
"""
123158
i = low
124159
j = high
@@ -134,7 +169,7 @@ def partition(array: list, low: int, high: int, pivot: int) -> int:
134169
i += 1
135170

136171

137-
def sort(array: list) -> list:
172+
def sort[T: Comparable](array: list[T]) -> list[T]:
138173
"""
139174
:param collection: some mutable ordered collection with heterogeneous
140175
comparable items inside
@@ -155,6 +190,12 @@ def sort(array: list) -> list:
155190
[0.3, 1.0, 1.7, 2.1, 3.3]
156191
>>> sort(['d', 'a', 'b', 'e', 'c'])
157192
['a', 'b', 'c', 'd', 'e']
193+
>>> sort(["banana", "apple", "cherry"])
194+
['apple', 'banana', 'cherry']
195+
>>> sort([1, "a"]) # doctest: +ELLIPSIS
196+
Traceback (most recent call last):
197+
...
198+
TypeError: ...
158199
"""
159200
if len(array) == 0:
160201
return array
@@ -163,9 +204,9 @@ def sort(array: list) -> list:
163204
return intro_sort(array, 0, len(array), size_threshold, max_depth)
164205

165206

166-
def intro_sort(
167-
array: list, start: int, end: int, size_threshold: int, max_depth: int
168-
) -> list:
207+
def intro_sort[T: Comparable](
208+
array: list[T], start: int, end: int, size_threshold: int, max_depth: int
209+
) -> list[T]:
169210
"""
170211
>>> array = [4, 2, 6, 8, 1, 7, 8, 22, 14, 56, 27, 79, 23, 45, 14, 12]
171212
>>> max_depth = 2 * math.ceil(math.log2(len(array)))

‎tests/test_sorts.py‎

Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
separately below.
1818
"""
1919

20+
import random
2021
from dataclasses import dataclass
2122
from typing import NamedTuple
2223

@@ -34,6 +35,7 @@
3435
from sorts.gnome_sort import gnome_sort
3536
from sorts.heap_sort import heap_sort
3637
from sorts.insertion_sort import insertion_sort
38+
from sorts.intro_sort import intro_sort, sort
3739
from sorts.iterative_merge_sort import iter_merge_sort
3840
from sorts.merge_sort import merge_sort
3941
from sorts.odd_even_sort import odd_even_sort
@@ -59,6 +61,36 @@ def test_heap_sort() -> None:
5961
assert heap_sort([5, 4, 3, 2, 1]) == [1, 2, 3, 4, 5]
6062

6163

64+
def test_intro_sort_comparable_items() -> None:
65+
"""``intro_sort`` agrees with the built-in once it leaves the small-input
66+
insertion-sort base case, and both of its remaining branches are reached.
67+
68+
Every case in the shared ``CASES`` battery is shorter than the 16-element
69+
threshold, so the battery alone only ever exercises ``insertion_sort``.
70+
The inputs below are large enough to take the quicksort branch, and the
71+
last two force the heapsort branch by exhausting the recursion budget.
72+
"""
73+
rng = random.Random(20260923)
74+
alphabet = "abcdefghijklmnopqrstuvwxyz"
75+
76+
for size in (17, 32, 100, 500):
77+
numbers = [rng.randint(-1000, 1000) for _ in range(size)]
78+
assert sort(numbers) == sorted(numbers)
79+
80+
letters = rng.choices(alphabet, k=size)
81+
assert sort(letters) == sorted(letters)
82+
83+
# a depth budget of 0 sends the algorithm straight to the heapsort branch
84+
values = [rng.randint(-1000, 1000) for _ in range(50)]
85+
assert intro_sort(list(values), 0, len(values), 1, 0) == sorted(values)
86+
87+
# a generous budget keeps it on the quicksort branch instead
88+
assert intro_sort(list(values), 0, len(values), 1, 64) == sorted(values)
89+
90+
with pytest.raises(TypeError):
91+
sort([1, "a"])
92+
93+
6294
SORTS = (
6395
binary_insertion_sort,
6496
bubble_sort_iterative,
@@ -83,6 +115,7 @@ def test_heap_sort() -> None:
83115
selection_sort,
84116
shell_sort,
85117
shrink_shell_sort,
118+
sort,
86119
stooge_sort,
87120
strand_sort,
88121
)
@@ -152,6 +185,7 @@ def test_rec_insertion_sort(case) -> None:
152185
reversort,
153186
selection_sort,
154187
shrink_shell_sort,
188+
sort,
155189
strand_sort,
156190
],
157191
ids=lambda f: f.__name__,

0 commit comments

Comments
 (0)