diff --git a/sorts/odd_even_transposition_parallel.py b/sorts/odd_even_transposition_parallel.py index 747899725094..1661a63bd9de 100644 --- a/sorts/odd_even_transposition_parallel.py +++ b/sorts/odd_even_transposition_parallel.py @@ -7,15 +7,14 @@ This implementation represents each variable in the list with a process and each process communicates with its neighboring processes in the list to perform comparisons. -They are synchronized with locks and message passing but other forms of +They are synchronized with message passing but other forms of synchronization could be used. """ import multiprocessing as mp - -# lock used to ensure that two processes do not access a pipe at the same time -# NOTE This breaks testing on build runner. May work better locally -# process_lock = mp.Lock() +from multiprocessing.connection import Connection, wait +from multiprocessing.reduction import ForkingPickler +from typing import Any, Protocol """ The function run by the processes that sorts the list @@ -29,46 +28,56 @@ """ -def oe_process( - position, - value, - l_send, - r_send, - lr_cv, - rr_cv, - result_pipe, - multiprocessing_context, -) -> None: - process_lock = multiprocessing_context.Lock() +class Comparable(Protocol): + def __lt__(self, other: Any, /) -> bool: ... + +def oe_process[T: Comparable]( + position: int, + value: T, + l_send: tuple[Connection, Connection] | None, + r_send: tuple[Connection, Connection] | None, + lr_cv: tuple[Connection, Connection] | None, + rr_cv: tuple[Connection, Connection] | None, + result_pipe: tuple[Connection, Connection], +) -> None: # we perform n swaps since after n swaps we know we are sorted # we *could* stop early if we are sorted already, but it takes as long to # find out we are sorted as it does to sort the list with this algorithm - for i in range(10): - if (i + position) % 2 == 0 and r_send is not None: - # send your value to your right neighbor - with process_lock: + try: + for i in range(10): + if (i + position) % 2 == 0 and r_send is not None and rr_cv is not None: + # send your value to your right neighbor r_send[1].send(value) - # receive your right neighbor's value - with process_lock: + # receive your right neighbor's value temp = rr_cv[0].recv() - # take the lower value since you are on the left - value = min(value, temp) - elif (i + position) % 2 != 0 and l_send is not None: - # send your value to your left neighbor - with process_lock: + # take the lower value since you are on the left + value = temp if temp < value else value + elif (i + position) % 2 != 0 and l_send is not None and lr_cv is not None: + # send your value to your left neighbor l_send[1].send(value) - # receive your left neighbor's value - with process_lock: + # receive your left neighbor's value temp = lr_cv[0].recv() - # take the higher value since you are on the right - value = max(value, temp) - # after all swaps are performed, send the values back to main - result_pipe[1].send(value) + # take the higher value since you are on the right + value = temp if value < temp else value + # after all swaps are performed, send the values back to main + result_pipe[1].send((value, None)) + except Exception as error: # noqa: BLE001 -- propagate worker errors to the caller + try: + payload = ForkingPickler.dumps((None, error)) + except Exception: # noqa: BLE001 -- user exceptions can fail to pickle + fallback = RuntimeError(f"{type(error).__name__}: {error}") + payload = ForkingPickler.dumps((None, fallback)) + result_pipe[1].send_bytes(payload) + finally: + for pipe in (l_send, r_send, lr_cv, rr_cv, result_pipe): + if pipe is not None: + for connection in pipe: + connection.close() """ @@ -78,8 +87,16 @@ def oe_process( """ -def odd_even_transposition(arr): +def odd_even_transposition[T: Comparable](arr: list[T]) -> list[T]: """ + Sort in place, propagating worker errors to the caller. Unpickleable + exceptions become RuntimeError with the original type name and message. + + >>> odd_even_transposition([]) + [] + >>> values = [42] + >>> odd_even_transposition(values) is values + True >>> odd_even_transposition(list(range(10)[::-1])) == sorted(list(range(10)[::-1])) True >>> odd_even_transposition(["a", "x", "c"]) == sorted(["x", "a", "c"]) @@ -98,7 +115,21 @@ def odd_even_transposition(arr): >>> unsorted_list = [-442, -98, -554, 266, -491, 985, -53, -529, 82, -429] >>> odd_even_transposition(unsorted_list) == sorted(unsorted_list + [1]) False + >>> values = ["c", "a", "b"] + >>> odd_even_transposition(values) is values + True + >>> values + ['a', 'b', 'c'] + >>> odd_even_transposition([2.5, -1, 0.0]) + [-1, 0.0, 2.5] + >>> odd_even_transposition([1, "a"]) # doctest: +IGNORE_EXCEPTION_DETAIL + Traceback (most recent call last): + ... + TypeError: '<' not supported between instances of 'str' and 'int' """ + if len(arr) < 2: + return arr + # spawn method is considered safer than fork multiprocessing_context = mp.get_context("spawn") @@ -112,6 +143,7 @@ def odd_even_transposition(arr): # of the loop temp_rs = multiprocessing_context.Pipe() temp_rr = multiprocessing_context.Pipe() + neighbor_pipes = [temp_rs, temp_rr] process_array_.append( multiprocessing_context.Process( target=oe_process, @@ -123,7 +155,6 @@ def odd_even_transposition(arr): None, temp_rr, result_pipe[0], - multiprocessing_context, ), ) ) @@ -133,6 +164,7 @@ def odd_even_transposition(arr): for i in range(1, len(arr) - 1): temp_rs = multiprocessing_context.Pipe() temp_rr = multiprocessing_context.Pipe() + neighbor_pipes.extend((temp_rs, temp_rr)) process_array_.append( multiprocessing_context.Process( target=oe_process, @@ -144,7 +176,6 @@ def odd_even_transposition(arr): temp_lr, temp_rr, result_pipe[i], - multiprocessing_context, ), ) ) @@ -162,19 +193,52 @@ def odd_even_transposition(arr): temp_lr, None, result_pipe[len(arr) - 1], - multiprocessing_context, ), ) ) - # start the processes - for p in process_array_: - p.start() + started_processes = [] + try: + for process in process_array_: + process.start() + started_processes.append(process) + + pending = {pipe[0]: position for position, pipe in enumerate(result_pipe)} + sentinels = { + process.sentinel: position + for position, process in enumerate(process_array_) + } + values = list(arr) + while pending: + ready = set(wait([*pending, *sentinels])) + for connection in pending.keys() & ready: + position = pending.pop(connection) + value, error = connection.recv() + if error is not None: + raise error + values[position] = value + for process_sentinel in sentinels.keys() & ready: + position = sentinels.pop(process_sentinel) + connection = result_pipe[position][0] + if connection in pending and not connection.poll(): + raise RuntimeError( + "Sorting worker exited without returning a result" + ) - # wait for the processes to end and write their values to the list - for p in range(len(result_pipe)): - arr[p] = result_pipe[p][0].recv() - process_array_[p].join() + # Do not partially overwrite the input if another worker fails. + arr[:] = values + except BaseException: + for process in started_processes: + if process.is_alive(): + process.terminate() + raise + finally: + for process in started_processes: + process.join() + process.close() + for pipe in result_pipe + neighbor_pipes: + for connection in pipe: + connection.close() return arr diff --git a/tests/test_sorts.py b/tests/test_sorts.py index 5cfbacf30e6c..b59e7dd03431 100644 --- a/tests/test_sorts.py +++ b/tests/test_sorts.py @@ -17,8 +17,12 @@ separately below. """ +import multiprocessing as mp +import os +import signal from dataclasses import dataclass from typing import NamedTuple +from unittest.mock import patch import pytest @@ -41,6 +45,9 @@ from sorts.merge_insertion_sort import merge_insertion_sort from sorts.merge_sort import merge_sort from sorts.odd_even_sort import odd_even_sort +from sorts.odd_even_transposition_parallel import ( + odd_even_transposition as parallel_odd_even_transposition, +) from sorts.odd_even_transposition_single_threaded import odd_even_transposition from sorts.pancake_sort import pancake_sort from sorts.patience_sort import patience_sort @@ -244,3 +251,80 @@ def test_bitonic_sort_comparable_items() -> None: with pytest.raises(TypeError): bitonic_sort([1, "two", 3, "four"], 0, 4, 1) + + +@dataclass +class UnpicklableComparison: + value: int + + def __lt__(self, _other: object, /) -> bool: + class LocalComparisonError(TypeError): + pass + + raise LocalComparisonError("comparison failed") + + +def _check_parallel_odd_even_transposition( + case: list[object], error: type[Exception] | None +) -> None: + # Give this probe and its workers a process group that the test alone owns. + os.setsid() + collection = list(case) + if error is not None: + with pytest.raises(error) as caught: + parallel_odd_even_transposition(collection) + if isinstance(case[0], UnpicklableComparison): + assert "LocalComparisonError: comparison failed" in str(caught.value) + assert collection == case + elif len(collection) < 2: + with patch( + "multiprocessing.process.BaseProcess.start", + side_effect=AssertionError("trivial input must not start a worker"), + ): + assert parallel_odd_even_transposition(collection) is collection + assert all(item is original for item, original in zip(collection, case)) + else: + assert parallel_odd_even_transposition(collection) is collection + assert collection == sorted(case) + assert not mp.active_children() + + +@pytest.mark.skipif( + os.name != "posix", reason="timeout cleanup requires process groups" +) +@pytest.mark.parametrize( + ("case", "error"), + [ + ([], None), + ([1], None), + ([Person()], None), + (["c", "a", "b"], None), + ([2.5, -1, 0.0], None), + ([Person(cost=100.0), Person(cost=-100.0), Person(name="Al")], None), + ([Dog(weight=15.5), Dog(weight=15.1), Dog(name="Buddy")], None), + ([1, "a"], TypeError), + ([3, 2, "a", 1], TypeError), + ([UnpicklableComparison(3), UnpicklableComparison(2)], RuntimeError), + ], +) +def test_parallel_odd_even_transposition( + case: list[object], error: type[Exception] | None +) -> None: + process = mp.get_context("spawn").Process( + target=_check_parallel_odd_even_transposition, args=(case, error) + ) + process.start() + try: + process.join(timeout=10) + assert not process.is_alive(), ( + "parallel sorting did not finish within 10 seconds" + ) + assert process.exitcode == 0 + finally: + if process.is_alive(): + try: + os.killpg(process.pid, signal.SIGKILL) + except ProcessLookupError: + process.kill() + process.join() + process.close()