From cd402bd5ca6b37c089dc7298a1a17e3b45c92ca4 Mon Sep 17 00:00:00 2001 From: AuroraAeon Date: Wed, 23 Sep 2026 17:24:16 +0800 Subject: [PATCH] types(unknown_sort): constrain items to Comparable Bind merge_sort's element type to a Comparable Protocol so the signature says "a list of items that can be compared with each other" instead of a bare list, and keeps the element type in the return. start/end are now annotated list[T] so the generic survives the local bindings. Adds doctests for a comparable non-int type (strings, floats) and for the failure mode: mixing non-comparable items must raise TypeError rather than silently mis-sort. The test battery picks the sort up for the shared cases and for the rejection check. It is imported under an alias because sorts/merge_sort.py already exports a merge_sort; the parametrize id therefore shows up as merge_sort0/merge_sort1, the same way shrink_shell_sort already shares the shell_sort id. --- sorts/unknown_sort.py | 22 ++++++++++++++++++++-- tests/test_sorts.py | 3 +++ 2 files changed, 23 insertions(+), 2 deletions(-) diff --git a/sorts/unknown_sort.py b/sorts/unknown_sort.py index 3545da68ea80..de9c1e73ad7f 100644 --- a/sorts/unknown_sort.py +++ b/sorts/unknown_sort.py @@ -5,8 +5,14 @@ already O(n) """ +from typing import Any, Protocol -def merge_sort(collection: list) -> list: + +class Comparable(Protocol): + def __lt__(self, other: Any, /) -> bool: ... + + +def merge_sort[T: Comparable](collection: list[T]) -> list[T]: """Pure implementation of the fastest merge sort algorithm in Python :param collection: some mutable ordered collection with heterogeneous @@ -22,8 +28,20 @@ def merge_sort(collection: list) -> list: >>> merge_sort([-2, -5, -45]) [-45, -5, -2] + + >>> merge_sort(["banana", "apple", "cherry"]) + ['apple', 'banana', 'cherry'] + + >>> merge_sort([3.14, 1.5, 2.7]) + [1.5, 2.7, 3.14] + + >>> merge_sort([1, "a"]) # doctest: +ELLIPSIS + Traceback (most recent call last): + ... + TypeError: ... """ - start, end = [], [] + start: list[T] = [] + end: list[T] = [] while len(collection) > 1: min_one, max_one = min(collection), max(collection) start.append(min_one) diff --git a/tests/test_sorts.py b/tests/test_sorts.py index d18be1c22b83..6dadd6127ec0 100644 --- a/tests/test_sorts.py +++ b/tests/test_sorts.py @@ -49,6 +49,7 @@ from sorts.shrink_shell_sort import shell_sort as shrink_shell_sort from sorts.stooge_sort import stooge_sort from sorts.strand_sort import strand_sort +from sorts.unknown_sort import merge_sort as unknown_merge_sort def test_heap_sort() -> None: @@ -85,6 +86,7 @@ def test_heap_sort() -> None: shrink_shell_sort, stooge_sort, strand_sort, + unknown_merge_sort, ) @@ -153,6 +155,7 @@ def test_rec_insertion_sort(case) -> None: selection_sort, shrink_shell_sort, strand_sort, + unknown_merge_sort, ], ids=lambda f: f.__name__, )