Skip to content

Commit a5ca6a9

Browse files
committed
sorts: type bitonic_sort's three functions for any comparable item
Part of #15234.
1 parent ba2d8ef commit a5ca6a9

1 file changed

Lines changed: 30 additions & 3 deletions

File tree

‎sorts/bitonic_sort.py‎

Lines changed: 30 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -6,8 +6,16 @@
66

77
from __future__ import annotations
88

9+
from typing import Protocol
910

10-
def comp_and_swap(array: list[int], index1: int, index2: int, direction: int) -> None:
11+
12+
class Comparable(Protocol):
13+
def __lt__(self, other: object, /) -> bool: ...
14+
15+
16+
def comp_and_swap[T: Comparable](
17+
array: list[T], index1: int, index2: int, direction: int
18+
) -> None:
1119
"""Compare the value at given index1 and index2 of the array and swap them as per
1220
the given direction.
1321
@@ -31,14 +39,26 @@ def comp_and_swap(array: list[int], index1: int, index2: int, direction: int) ->
3139
>>> comp_and_swap(arr, 0, 3, 0)
3240
>>> arr
3341
[12, 42, -21, 1]
42+
43+
>>> arr = ["d", "a", "c", "b"]
44+
>>> comp_and_swap(arr, 0, 1, 1)
45+
>>> arr
46+
['a', 'd', 'c', 'b']
47+
48+
>>> comp_and_swap([1, "a"], 0, 1, 1) # doctest: +IGNORE_EXCEPTION_DETAIL
49+
Traceback (most recent call last):
50+
...
51+
TypeError: '<' not supported between instances of 'str' and 'int'
3452
"""
3553
if (direction == 1 and array[index1] > array[index2]) or (
3654
direction == 0 and array[index1] < array[index2]
3755
):
3856
array[index1], array[index2] = array[index2], array[index1]
3957

4058

41-
def bitonic_merge(array: list[int], low: int, length: int, direction: int) -> None:
59+
def bitonic_merge[T: Comparable](
60+
array: list[T], low: int, length: int, direction: int
61+
) -> None:
4262
"""
4363
It recursively sorts a bitonic sequence in ascending order, if direction = 1, and in
4464
descending if direction = 0.
@@ -62,7 +82,9 @@ def bitonic_merge(array: list[int], low: int, length: int, direction: int) -> No
6282
bitonic_merge(array, low + middle, middle, direction)
6383

6484

65-
def bitonic_sort(array: list[int], low: int, length: int, direction: int) -> None:
85+
def bitonic_sort[T: Comparable](
86+
array: list[T], low: int, length: int, direction: int
87+
) -> None:
6688
"""
6789
This function first produces a bitonic sequence by recursively sorting its two
6890
halves in opposite sorting orders, and then calls bitonic_merge to make them in the
@@ -76,6 +98,11 @@ def bitonic_sort(array: list[int], low: int, length: int, direction: int) -> Non
7698
>>> bitonic_sort(arr, 0, 8, 0)
7799
>>> arr
78100
[145, 92, 34, 12, 0, -23, -121, -167]
101+
102+
>>> arr = [3.3, 1.1, 4.4, 2.2]
103+
>>> bitonic_sort(arr, 0, 4, 1)
104+
>>> arr
105+
[1.1, 2.2, 3.3, 4.4]
79106
"""
80107
if length > 1:
81108
middle = int(length / 2)

0 commit comments

Comments
 (0)