diff --git a/sorts/bitonic_sort.py b/sorts/bitonic_sort.py index 600f8139603a..a2a974ff1430 100644 --- a/sorts/bitonic_sort.py +++ b/sorts/bitonic_sort.py @@ -6,8 +6,16 @@ from __future__ import annotations +from typing import Protocol -def comp_and_swap(array: list[int], index1: int, index2: int, direction: int) -> None: + +class Comparable(Protocol): + def __lt__(self, other: object, /) -> bool: ... + + +def comp_and_swap[T: Comparable]( + array: list[T], index1: int, index2: int, direction: int +) -> None: """Compare the value at given index1 and index2 of the array and swap them as per the given direction. @@ -31,6 +39,16 @@ def comp_and_swap(array: list[int], index1: int, index2: int, direction: int) -> >>> comp_and_swap(arr, 0, 3, 0) >>> arr [12, 42, -21, 1] + + >>> arr = ["d", "a", "c", "b"] + >>> comp_and_swap(arr, 0, 1, 1) + >>> arr + ['a', 'd', 'c', 'b'] + + >>> comp_and_swap([1, "a"], 0, 1, 1) # doctest: +IGNORE_EXCEPTION_DETAIL + Traceback (most recent call last): + ... + TypeError: '<' not supported between instances of 'str' and 'int' """ if (direction == 1 and array[index1] > array[index2]) or ( direction == 0 and array[index1] < array[index2] @@ -38,7 +56,9 @@ def comp_and_swap(array: list[int], index1: int, index2: int, direction: int) -> array[index1], array[index2] = array[index2], array[index1] -def bitonic_merge(array: list[int], low: int, length: int, direction: int) -> None: +def bitonic_merge[T: Comparable]( + array: list[T], low: int, length: int, direction: int +) -> None: """ It recursively sorts a bitonic sequence in ascending order, if direction = 1, and in descending if direction = 0. @@ -62,7 +82,9 @@ def bitonic_merge(array: list[int], low: int, length: int, direction: int) -> No bitonic_merge(array, low + middle, middle, direction) -def bitonic_sort(array: list[int], low: int, length: int, direction: int) -> None: +def bitonic_sort[T: Comparable]( + array: list[T], low: int, length: int, direction: int +) -> None: """ This function first produces a bitonic sequence by recursively sorting its two 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 >>> bitonic_sort(arr, 0, 8, 0) >>> arr [145, 92, 34, 12, 0, -23, -121, -167] + + >>> arr = [3.3, 1.1, 4.4, 2.2] + >>> bitonic_sort(arr, 0, 4, 1) + >>> arr + [1.1, 2.2, 3.3, 4.4] """ if length > 1: middle = int(length / 2)