Skip to content

Commit b6519ed

Browse files
committed
fix(sorts): propagate parallel odd-even comparison errors
1 parent 32e2d50 commit b6519ed

2 files changed

Lines changed: 147 additions & 45 deletions

File tree

‎sorts/odd_even_transposition_parallel.py‎

Lines changed: 92 additions & 45 deletions
Original file line numberDiff line numberDiff line change
@@ -7,15 +7,13 @@
77
This implementation represents each variable in the list with a process and
88
each process communicates with its neighboring processes in the list to perform
99
comparisons.
10-
They are synchronized with locks and message passing but other forms of
10+
They are synchronized with message passing but other forms of
1111
synchronization could be used.
1212
"""
1313

1414
import multiprocessing as mp
15-
16-
# lock used to ensure that two processes do not access a pipe at the same time
17-
# NOTE This breaks testing on build runner. May work better locally
18-
# process_lock = mp.Lock()
15+
from multiprocessing.connection import Connection, wait
16+
from typing import Any, Protocol
1917

2018
"""
2119
The function run by the processes that sorts the list
@@ -29,46 +27,51 @@
2927
"""
3028

3129

32-
def oe_process(
33-
position,
34-
value,
35-
l_send,
36-
r_send,
37-
lr_cv,
38-
rr_cv,
39-
result_pipe,
40-
multiprocessing_context,
41-
) -> None:
42-
process_lock = multiprocessing_context.Lock()
30+
class Comparable(Protocol):
31+
def __lt__(self, other: Any, /) -> bool: ...
32+
4333

34+
def oe_process[T: Comparable](
35+
position: int,
36+
value: T,
37+
l_send: tuple[Connection, Connection] | None,
38+
r_send: tuple[Connection, Connection] | None,
39+
lr_cv: tuple[Connection, Connection] | None,
40+
rr_cv: tuple[Connection, Connection] | None,
41+
result_pipe: tuple[Connection, Connection],
42+
) -> None:
4443
# we perform n swaps since after n swaps we know we are sorted
4544
# we *could* stop early if we are sorted already, but it takes as long to
4645
# find out we are sorted as it does to sort the list with this algorithm
47-
for i in range(10):
48-
if (i + position) % 2 == 0 and r_send is not None:
49-
# send your value to your right neighbor
50-
with process_lock:
46+
try:
47+
for i in range(10):
48+
if (i + position) % 2 == 0 and r_send is not None and rr_cv is not None:
49+
# send your value to your right neighbor
5150
r_send[1].send(value)
5251

53-
# receive your right neighbor's value
54-
with process_lock:
52+
# receive your right neighbor's value
5553
temp = rr_cv[0].recv()
5654

57-
# take the lower value since you are on the left
58-
value = min(value, temp)
59-
elif (i + position) % 2 != 0 and l_send is not None:
60-
# send your value to your left neighbor
61-
with process_lock:
55+
# take the lower value since you are on the left
56+
value = temp if temp < value else value
57+
elif (i + position) % 2 != 0 and l_send is not None and lr_cv is not None:
58+
# send your value to your left neighbor
6259
l_send[1].send(value)
6360

64-
# receive your left neighbor's value
65-
with process_lock:
61+
# receive your left neighbor's value
6662
temp = lr_cv[0].recv()
6763

68-
# take the higher value since you are on the right
69-
value = max(value, temp)
70-
# after all swaps are performed, send the values back to main
71-
result_pipe[1].send(value)
64+
# take the higher value since you are on the right
65+
value = temp if value < temp else value
66+
# after all swaps are performed, send the values back to main
67+
result_pipe[1].send((value, None))
68+
except Exception as error: # noqa: BLE001 -- propagate worker errors to the caller
69+
result_pipe[1].send((None, error))
70+
finally:
71+
for pipe in (l_send, r_send, lr_cv, rr_cv, result_pipe):
72+
if pipe is not None:
73+
for connection in pipe:
74+
connection.close()
7275

7376

7477
"""
@@ -78,7 +81,7 @@ def oe_process(
7881
"""
7982

8083

81-
def odd_even_transposition(arr):
84+
def odd_even_transposition[T: Comparable](arr: list[T]) -> list[T]:
8285
"""
8386
>>> odd_even_transposition(list(range(10)[::-1])) == sorted(list(range(10)[::-1]))
8487
True
@@ -98,6 +101,17 @@ def odd_even_transposition(arr):
98101
>>> unsorted_list = [-442, -98, -554, 266, -491, 985, -53, -529, 82, -429]
99102
>>> odd_even_transposition(unsorted_list) == sorted(unsorted_list + [1])
100103
False
104+
>>> values = ["c", "a", "b"]
105+
>>> odd_even_transposition(values) is values
106+
True
107+
>>> values
108+
['a', 'b', 'c']
109+
>>> odd_even_transposition([2.5, -1, 0.0])
110+
[-1, 0.0, 2.5]
111+
>>> odd_even_transposition([1, "a"]) # doctest: +IGNORE_EXCEPTION_DETAIL
112+
Traceback (most recent call last):
113+
...
114+
TypeError: '<' not supported between instances of 'str' and 'int'
101115
"""
102116
# spawn method is considered safer than fork
103117
multiprocessing_context = mp.get_context("spawn")
@@ -112,6 +126,7 @@ def odd_even_transposition(arr):
112126
# of the loop
113127
temp_rs = multiprocessing_context.Pipe()
114128
temp_rr = multiprocessing_context.Pipe()
129+
neighbor_pipes = [temp_rs, temp_rr]
115130
process_array_.append(
116131
multiprocessing_context.Process(
117132
target=oe_process,
@@ -123,7 +138,6 @@ def odd_even_transposition(arr):
123138
None,
124139
temp_rr,
125140
result_pipe[0],
126-
multiprocessing_context,
127141
),
128142
)
129143
)
@@ -133,6 +147,7 @@ def odd_even_transposition(arr):
133147
for i in range(1, len(arr) - 1):
134148
temp_rs = multiprocessing_context.Pipe()
135149
temp_rr = multiprocessing_context.Pipe()
150+
neighbor_pipes.extend((temp_rs, temp_rr))
136151
process_array_.append(
137152
multiprocessing_context.Process(
138153
target=oe_process,
@@ -144,7 +159,6 @@ def odd_even_transposition(arr):
144159
temp_lr,
145160
temp_rr,
146161
result_pipe[i],
147-
multiprocessing_context,
148162
),
149163
)
150164
)
@@ -162,19 +176,52 @@ def odd_even_transposition(arr):
162176
temp_lr,
163177
None,
164178
result_pipe[len(arr) - 1],
165-
multiprocessing_context,
166179
),
167180
)
168181
)
169182

170-
# start the processes
171-
for p in process_array_:
172-
p.start()
173-
174-
# wait for the processes to end and write their values to the list
175-
for p in range(len(result_pipe)):
176-
arr[p] = result_pipe[p][0].recv()
177-
process_array_[p].join()
183+
started_processes = []
184+
try:
185+
for process in process_array_:
186+
process.start()
187+
started_processes.append(process)
188+
189+
pending = {pipe[0]: position for position, pipe in enumerate(result_pipe)}
190+
sentinels = {
191+
process.sentinel: position
192+
for position, process in enumerate(process_array_)
193+
}
194+
values = list(arr)
195+
while pending:
196+
ready = set(wait([*pending, *sentinels]))
197+
for connection in pending.keys() & ready:
198+
position = pending.pop(connection)
199+
value, error = connection.recv()
200+
if error is not None:
201+
raise error
202+
values[position] = value
203+
for process_sentinel in sentinels.keys() & ready:
204+
position = sentinels.pop(process_sentinel)
205+
connection = result_pipe[position][0]
206+
if connection in pending and not connection.poll():
207+
raise RuntimeError(
208+
"Sorting worker exited without returning a result"
209+
)
210+
211+
# Do not partially overwrite the input if another worker fails.
212+
arr[:] = values
213+
except BaseException:
214+
for process in started_processes:
215+
if process.is_alive():
216+
process.terminate()
217+
raise
218+
finally:
219+
for process in started_processes:
220+
process.join()
221+
process.close()
222+
for pipe in result_pipe + neighbor_pipes:
223+
for connection in pipe:
224+
connection.close()
178225
return arr
179226

180227

‎tests/test_sorts.py‎

Lines changed: 55 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,9 @@
1717
separately below.
1818
"""
1919

20+
import multiprocessing as mp
21+
import os
22+
import signal
2023
from dataclasses import dataclass
2124
from typing import NamedTuple
2225

@@ -41,6 +44,9 @@
4144
from sorts.merge_insertion_sort import merge_insertion_sort
4245
from sorts.merge_sort import merge_sort
4346
from sorts.odd_even_sort import odd_even_sort
47+
from sorts.odd_even_transposition_parallel import (
48+
odd_even_transposition as parallel_odd_even_transposition,
49+
)
4450
from sorts.odd_even_transposition_single_threaded import odd_even_transposition
4551
from sorts.pancake_sort import pancake_sort
4652
from sorts.patience_sort import patience_sort
@@ -241,3 +247,52 @@ def test_bitonic_sort_comparable_items() -> None:
241247

242248
with pytest.raises(TypeError):
243249
bitonic_sort([1, "two", 3, "four"], 0, 4, 1)
250+
251+
252+
def _check_parallel_odd_even_transposition(case: list[object], rejects: bool) -> None:
253+
# Give this probe and its workers a process group that the test alone owns.
254+
os.setsid()
255+
collection = list(case)
256+
if rejects:
257+
with pytest.raises(TypeError):
258+
parallel_odd_even_transposition(collection)
259+
assert collection == case
260+
else:
261+
assert parallel_odd_even_transposition(collection) is collection
262+
assert collection == sorted(case)
263+
assert not mp.active_children()
264+
265+
266+
@pytest.mark.skipif(
267+
os.name != "posix", reason="timeout cleanup requires process groups"
268+
)
269+
@pytest.mark.parametrize(
270+
("case", "rejects"),
271+
[
272+
(["c", "a", "b"], False),
273+
([2.5, -1, 0.0], False),
274+
([Person(cost=100.0), Person(cost=-100.0), Person(name="Al")], False),
275+
([Dog(weight=15.5), Dog(weight=15.1), Dog(name="Buddy")], False),
276+
([1, "a"], True),
277+
([3, 2, "a", 1], True),
278+
],
279+
)
280+
def test_parallel_odd_even_transposition(case: list[object], rejects: bool) -> None:
281+
process = mp.get_context("spawn").Process(
282+
target=_check_parallel_odd_even_transposition, args=(case, rejects)
283+
)
284+
process.start()
285+
try:
286+
process.join(timeout=10)
287+
assert not process.is_alive(), (
288+
"parallel sorting did not finish within 10 seconds"
289+
)
290+
assert process.exitcode == 0
291+
finally:
292+
if process.is_alive():
293+
try:
294+
os.killpg(process.pid, signal.SIGKILL)
295+
except ProcessLookupError:
296+
process.kill()
297+
process.join()
298+
process.close()

0 commit comments

Comments
 (0)