Skip to content

Commit 8747434

Browse files
nightcitybladenightcitybladecclausspre-commit-ci[bot]
authored
types(power_sort): preserve input item type (#15488)
* types(power_sort): preserve input item type * Enhance power_sort docstring with more examples Added examples for using power_sort with tuples and lists. * Update power_sort.py * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Co-authored-by: nightcityblade <nightcityblade@gmail.com> Co-authored-by: Christian Clauss <cclauss@me.com> Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
1 parent 32e2d50 commit 8747434

2 files changed

Lines changed: 33 additions & 20 deletions

File tree

‎sorts/power_sort.py‎

Lines changed: 30 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -25,14 +25,12 @@
2525
python power_sort.py
2626
"""
2727

28-
from __future__ import annotations
29-
30-
from collections.abc import Callable
28+
from collections.abc import Callable, Iterable
3129
from typing import Any
3230

3331

34-
def _find_run(
35-
arr: list, start: int, end: int, key: Callable[[Any], Any] | None = None
32+
def _find_run[T](
33+
arr: list[T], start: int, end: int, key: Callable[[Any], Any] | None = None
3634
) -> int:
3735
"""
3836
Detect a run (ascending or descending sequence) starting at 'start'.
@@ -67,7 +65,7 @@ def _find_run(
6765
if start >= end - 1:
6866
return start + 1
6967

70-
key_func = key if key else lambda element: element
68+
key_func = key or (lambda element: element)
7169
run_end = start + 1
7270

7371
# Check if run is ascending or descending
@@ -134,8 +132,8 @@ def _node_power(total_length: int, b1: int, n1: int, b2: int, n2: int) -> int:
134132
return power
135133

136134

137-
def _merge(
138-
arr: list,
135+
def _merge[T](
136+
arr: list[T],
139137
start1: int,
140138
end1: int,
141139
end2: int,
@@ -165,7 +163,7 @@ def _merge(
165163
>>> arr
166164
[1, 2, 3, 5, 6, 7]
167165
"""
168-
key_func = key if key else lambda element: element
166+
key_func = key or (lambda element: element)
169167

170168
# Copy the runs to temporary storage
171169
left = arr[start1:end1]
@@ -196,12 +194,12 @@ def _merge(
196194
k += 1
197195

198196

199-
def power_sort(
200-
collection: list,
197+
def power_sort[T](
198+
collection: Iterable[T],
201199
*,
202200
key: Callable[[Any], Any] | None = None,
203201
reverse: bool = False,
204-
) -> list:
202+
) -> list[T]:
205203
"""
206204
Sort a list using the PowerSort algorithm.
207205
@@ -243,26 +241,38 @@ def power_sort(
243241
['apple', 'banana', 'cherry']
244242
>>> power_sort([3.14, 2.71, 1.41, 1.73])
245243
[1.41, 1.73, 2.71, 3.14]
244+
>>> power_sort(value for value in [3, 1, 2]) # list
245+
[1, 2, 3]
246+
>>> power_sort(value for value in (3, 1, 2)) # tuple
247+
[1, 2, 3]
246248
>>> power_sort([5, 2, 8, 1, 9], reverse=True)
247249
[9, 8, 5, 2, 1]
250+
>>> power_sort(['apple', 'pie', 'a', 'longer'])
251+
['a', 'apple', 'longer', 'pie']
248252
>>> power_sort(['apple', 'pie', 'a', 'longer'], key=len)
249253
['a', 'pie', 'apple', 'longer']
254+
>>> power_sort(['apple', 'pie', 'a', 'longer'], reverse=True)
255+
['pie', 'longer', 'apple', 'a']
256+
>>> power_sort(['apple', 'pie', 'a', 'longer'], key=len, reverse=True) # Fix me!
257+
['a', 'pie', 'apple', 'longer']
250258
>>> power_sort([(1, 'b'), (2, 'a'), (1, 'a')], key=lambda x: x[0])
251259
[(1, 'b'), (1, 'a'), (2, 'a')]
252260
>>> power_sort([1, 2, 3, 2, 1, 2, 3, 4])
253261
[1, 1, 2, 2, 2, 3, 3, 4]
254-
>>> result = power_sort(list(range(100)))
255-
>>> result == list(range(100))
262+
>>> power_sort(list(range(100))) == list(range(100))
256263
True
257-
>>> result = power_sort(list(reversed(range(50))))
258-
>>> result == list(range(50))
264+
>>> power_sort(list(reversed(range(50)))) == list(range(50))
259265
True
266+
>>> power_sort([1, "a"])
267+
Traceback (most recent call last):
268+
...
269+
TypeError: '<' not supported between instances of 'str' and 'int'
260270
"""
261-
if len(collection) <= 1:
262-
return collection
263-
264-
# Make a copy to avoid modifying the original if it's immutable
271+
# Make a copy so any iterable is accepted and the original is not modified.
265272
arr = list(collection)
273+
if len(arr) <= 1:
274+
return arr
275+
266276
total_length = len(arr)
267277

268278
# Adjust key function for reverse sorting

‎tests/test_sorts.py‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,7 @@
4444
from sorts.odd_even_transposition_single_threaded import odd_even_transposition
4545
from sorts.pancake_sort import pancake_sort
4646
from sorts.patience_sort import patience_sort
47+
from sorts.power_sort import power_sort
4748
from sorts.quick_sort import quick_sort
4849
from sorts.quick_sort_3_partition import three_way_radix_quicksort
4950
from sorts.recursive_insertion_sort import rec_insertion_sort
@@ -117,6 +118,7 @@ def test_intro_sort_heap_fallback_preserves_surrounding_items(max_depth: int) ->
117118
odd_even_transposition,
118119
pancake_sort,
119120
patience_sort,
121+
power_sort,
120122
quick_sort,
121123
recursive_quick_sort,
122124
reverse_selection_sort,
@@ -198,6 +200,7 @@ def test_rec_insertion_sort(case) -> None:
198200
odd_even_transposition,
199201
pancake_sort,
200202
patience_sort,
203+
power_sort,
201204
recursive_quick_sort,
202205
reverse_selection_sort,
203206
reversort,

0 commit comments

Comments
 (0)