diff --git a/sorts/bitonic_sort.py b/sorts/bitonic_sort.py index 600f8139603a..f2cf04d8c681 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. @@ -32,13 +40,15 @@ def comp_and_swap(array: list[int], index1: int, index2: int, direction: int) -> >>> arr [12, 42, -21, 1] """ - if (direction == 1 and array[index1] > array[index2]) or ( + if (direction == 1 and array[index2] < array[index1]) or ( direction == 0 and array[index1] < array[index2] ): 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 +72,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 +88,22 @@ 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 = ["banana", "apple", "cherry", "date"] + >>> bitonic_sort(arr, 0, 4, 1) + >>> arr + ['apple', 'banana', 'cherry', 'date'] + + >>> arr = [3, 1.5, 2, 4.5] + >>> bitonic_sort(arr, 0, 4, 1) + >>> arr + [1.5, 2, 3, 4.5] + + >>> arr = [1, "two", 3, "four"] + >>> bitonic_sort(arr, 0, 4, 1) + Traceback (most recent call last): + ... + TypeError: '<' not supported between instances of 'str' and 'int' """ if length > 1: middle = int(length / 2) diff --git a/tests/test_sorts.py b/tests/test_sorts.py index a4a8ed7411bf..aa869eb582a9 100644 --- a/tests/test_sorts.py +++ b/tests/test_sorts.py @@ -183,3 +183,18 @@ def test_bogo_sort_comparable_items() -> None: with pytest.raises(TypeError): bogo_sort([1, "a"]) + + +def test_bitonic_sort_comparable_items() -> None: + from sorts.bitonic_sort import bitonic_sort + + strings = ["banana", "apple", "cherry", "date"] + bitonic_sort(strings, 0, len(strings), 1) + assert strings == ["apple", "banana", "cherry", "date"] + + numbers = [3, 1.5, 2, 4.5] + bitonic_sort(numbers, 0, len(numbers), 1) + assert numbers == [1.5, 2, 3, 4.5] + + with pytest.raises(TypeError): + bitonic_sort([1, "two", 3, "four"], 0, 4, 1)