66
77from __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