diff --git a/sorts/strand_sort.py b/sorts/strand_sort.py index 4cadd396178e..ee41aa5df51d 100644 --- a/sorts/strand_sort.py +++ b/sorts/strand_sort.py @@ -1,7 +1,19 @@ import operator +from typing import Protocol, TypeVar -def strand_sort(arr: list, reverse: bool = False, solution: list | None = None) -> list: +class Comparable(Protocol): + def __lt__(self, other: object, /) -> bool: ... + + def __gt__(self, other: object, /) -> bool: ... + + +T = TypeVar("T", bound=Comparable) + + +def strand_sort[T]( + arr: list[T], reverse: bool = False, solution: list[T] | None = None +) -> list[T]: """ Strand sort implementation source: https://en.wikipedia.org/wiki/Strand_sort @@ -16,6 +28,10 @@ def strand_sort(arr: list, reverse: bool = False, solution: list | None = None) >>> strand_sort([4, 2, 5, 3, 0, 1], reverse=True) [5, 4, 3, 2, 1, 0] + + >>> strand_sort(["banana", "apple", "cherry"]) + ['apple', 'banana', 'cherry'] + """ _operator = operator.lt if reverse else operator.gt solution = solution or [] diff --git a/tests/test_sorts.py b/tests/test_sorts.py index f2b14be44efe..e3e91838a792 100644 --- a/tests/test_sorts.py +++ b/tests/test_sorts.py @@ -142,6 +142,7 @@ def test_rec_insertion_sort(case) -> None: pancake_sort, selection_sort, shrink_shell_sort, + strand_sort, ], ids=lambda f: f.__name__, )