77This implementation represents each variable in the list with a process and
88each process communicates with its neighboring processes in the list to perform
99comparisons.
10- They are synchronized with locks and message passing but other forms of
10+ They are synchronized with message passing but other forms of
1111synchronization could be used.
1212"""
1313
1414import 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"""
2119The function run by the processes that sorts the list
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
0 commit comments