@@ -178,6 +178,288 @@ def pop_set(out):
178178 self .assertEqual (set (all_popped ), items )
179179 self .assertEqual (len (s ), 0 )
180180
181+ def test_add_concurrent (self ):
182+ """Test set.add() with disjoint inputs from several threads."""
183+ NUM_THREADS = 4
184+ NUM_ITERS = 20
185+
186+ for _ in range (NUM_ITERS ):
187+ inputs = [
188+ range (i * self .SET_SIZE , (i + 1 ) * self .SET_SIZE )
189+ for i in range (NUM_THREADS )
190+ ]
191+ expected = set ().union (* inputs )
192+ s = set ()
193+ barrier = Barrier (NUM_THREADS , timeout = 2 )
194+
195+ def add_items (items ):
196+ barrier .wait ()
197+ for item in items :
198+ s .add (item )
199+
200+ threads = [Thread (target = add_items , args = (items ,))
201+ for items in inputs ]
202+ for t in threads :
203+ t .start ()
204+ for t in threads :
205+ t .join ()
206+
207+ self .assertEqual (s , expected )
208+
209+ def test_remove_discard_concurrent (self ):
210+ """Test set.remove() and set.discard() from several threads."""
211+ NUM_THREADS = 4
212+ NUM_ITERS = 20
213+
214+ for method_name in ("remove" , "discard" ):
215+ with self .subTest (method = method_name ):
216+ for _ in range (NUM_ITERS ):
217+ inputs = [
218+ range (i * self .SET_SIZE , (i + 1 ) * self .SET_SIZE )
219+ for i in range (NUM_THREADS )
220+ ]
221+ untouched = set (range (
222+ NUM_THREADS * self .SET_SIZE ,
223+ (NUM_THREADS + 1 ) * self .SET_SIZE ,
224+ ))
225+ s = set ().union (untouched , * inputs )
226+ barrier = Barrier (NUM_THREADS , timeout = 2 )
227+
228+ def remove_items (items ):
229+ barrier .wait ()
230+ method = getattr (s , method_name )
231+ for item in items :
232+ method (item )
233+
234+ threads = [Thread (target = remove_items , args = (items ,))
235+ for items in inputs ]
236+ for t in threads :
237+ t .start ()
238+ for t in threads :
239+ t .join ()
240+
241+ self .assertEqual (s , untouched )
242+
243+ def test_copy_clear_concurrent (self ):
244+ """Test set.copy() while another thread clears the set."""
245+ NUM_ITERS = 20
246+
247+ for _ in range (NUM_ITERS ):
248+ items = set (range (self .SET_SIZE ))
249+ s = set (items )
250+ copies = []
251+ barrier = Barrier (2 , timeout = 2 )
252+
253+ def copy_set ():
254+ barrier .wait ()
255+ copies .append (s .copy ())
256+
257+ def clear_set ():
258+ barrier .wait ()
259+ s .clear ()
260+
261+ threads = [Thread (target = copy_set ), Thread (target = clear_set )]
262+ for t in threads :
263+ t .start ()
264+ for t in threads :
265+ t .join ()
266+
267+ self .assertIn (copies [0 ], (items , set ()))
268+ self .assertEqual (s , set ())
269+
270+ def test_update_concurrent (self ):
271+ """Test updates of one shared set from disjoint source sets."""
272+ NUM_THREADS = 4
273+ NUM_ITERS = 20
274+
275+ sources = [
276+ set (range (i * self .SET_SIZE , (i + 1 ) * self .SET_SIZE ))
277+ for i in range (NUM_THREADS )
278+ ]
279+ original_sources = [set (source ) for source in sources ]
280+ initial = set (range (
281+ NUM_THREADS * self .SET_SIZE ,
282+ (NUM_THREADS + 1 ) * self .SET_SIZE ,
283+ ))
284+ expected = set ().union (initial , * sources )
285+
286+ for _ in range (NUM_ITERS ):
287+ s = set (initial )
288+ barrier = Barrier (NUM_THREADS , timeout = 2 )
289+
290+ def update_set (source ):
291+ barrier .wait ()
292+ s .update (source )
293+
294+ threads = [Thread (target = update_set , args = (source ,))
295+ for source in sources ]
296+ for t in threads :
297+ t .start ()
298+ for t in threads :
299+ t .join ()
300+
301+ self .assertEqual (s , expected )
302+ self .assertEqual (sources , original_sources )
303+
304+ def test_update_opposing (self ):
305+ """Test opposing updates of two shared sets."""
306+ NUM_ITERS = 20
307+
308+ for _ in range (NUM_ITERS ):
309+ left = set (range (self .SET_SIZE ))
310+ right = set (range (self .SET_SIZE , self .SET_SIZE * 2 ))
311+ expected = set (range (self .SET_SIZE * 2 ))
312+ barrier = Barrier (2 , timeout = 2 )
313+
314+ def update_set (target , source ):
315+ barrier .wait ()
316+ target .update (source )
317+
318+ threads = [
319+ Thread (target = update_set , args = (left , right )),
320+ Thread (target = update_set , args = (right , left )),
321+ ]
322+ for t in threads :
323+ t .start ()
324+ for t in threads :
325+ t .join ()
326+
327+ self .assertEqual (left , expected )
328+ self .assertEqual (right , expected )
329+
330+ def test_difference_update_concurrent (self ):
331+ """Test set.difference_update() with disjoint source sets."""
332+ NUM_THREADS = 4
333+ NUM_ITERS = 20
334+
335+ sources = [
336+ set (range (i * self .SET_SIZE , (i + 1 ) * self .SET_SIZE ))
337+ for i in range (NUM_THREADS )
338+ ]
339+ original_sources = [set (source ) for source in sources ]
340+ untouched = set (range (
341+ NUM_THREADS * self .SET_SIZE ,
342+ (NUM_THREADS + 1 ) * self .SET_SIZE ,
343+ ))
344+ items = set ().union (untouched , * sources )
345+
346+ for _ in range (NUM_ITERS ):
347+ s = set (items )
348+ barrier = Barrier (NUM_THREADS , timeout = 2 )
349+
350+ def difference_update (source ):
351+ barrier .wait ()
352+ s .difference_update (source )
353+
354+ threads = [Thread (target = difference_update , args = (source ,))
355+ for source in sources ]
356+ for t in threads :
357+ t .start ()
358+ for t in threads :
359+ t .join ()
360+
361+ self .assertEqual (s , untouched )
362+ self .assertEqual (sources , original_sources )
363+
364+ def test_difference_update_opposing (self ):
365+ """Test opposing difference updates of two shared sets."""
366+ NUM_ITERS = 20
367+
368+ for _ in range (NUM_ITERS ):
369+ left = set (range (self .SET_SIZE ))
370+ right = set (range (1 , self .SET_SIZE + 1 ))
371+ original_left = set (left )
372+ original_right = set (right )
373+ left_only = {0 }
374+ right_only = {self .SET_SIZE }
375+ barrier = Barrier (2 , timeout = 2 )
376+
377+ def difference_update (target , source ):
378+ barrier .wait ()
379+ target .difference_update (source )
380+
381+ threads = [
382+ Thread (target = difference_update , args = (left , right )),
383+ Thread (target = difference_update , args = (right , left )),
384+ ]
385+ for t in threads :
386+ t .start ()
387+ for t in threads :
388+ t .join ()
389+
390+ actual = (left , right )
391+ expected = [
392+ (left_only , original_right ),
393+ (original_left , right_only ),
394+ ]
395+ self .assertIn (actual , expected )
396+
397+ def test_symmetric_difference_update_concurrent (self ):
398+ """Test symmetric difference updates with disjoint source sets."""
399+ NUM_THREADS = 4
400+ NUM_ITERS = 20
401+
402+ sources = [
403+ set (range (i * self .SET_SIZE , (i + 1 ) * self .SET_SIZE ))
404+ for i in range (NUM_THREADS )
405+ ]
406+ items = set ().union (* sources )
407+ initial = {item for item in items if item % 2 == 0 }
408+ expected = {item for item in items if item % 2 == 1 }
409+
410+ for _ in range (NUM_ITERS ):
411+ s = set (initial )
412+ barrier = Barrier (NUM_THREADS , timeout = 2 )
413+
414+ def symmetric_difference_update (source ):
415+ barrier .wait ()
416+ s .symmetric_difference_update (source )
417+
418+ threads = [
419+ Thread (target = symmetric_difference_update , args = (source ,))
420+ for source in sources
421+ ]
422+ for t in threads :
423+ t .start ()
424+ for t in threads :
425+ t .join ()
426+
427+ self .assertEqual (s , expected )
428+
429+ def test_symmetric_difference_update_opposing (self ):
430+ """Test opposing symmetric difference updates of shared sets."""
431+ NUM_ITERS = 20
432+
433+ for _ in range (NUM_ITERS ):
434+ left = set (range (self .SET_SIZE ))
435+ right = set (range (self .SET_SIZE , self .SET_SIZE * 2 ))
436+ original_left = set (left )
437+ original_right = set (right )
438+ union = set (range (self .SET_SIZE * 2 ))
439+ barrier = Barrier (2 , timeout = 2 )
440+
441+ def symmetric_update (target , source ):
442+ barrier .wait ()
443+ target .symmetric_difference_update (source )
444+
445+ threads = [
446+ Thread (target = symmetric_update , args = (left , right )),
447+ Thread (target = symmetric_update , args = (right , left )),
448+ ]
449+ for t in threads :
450+ t .start ()
451+ for t in threads :
452+ t .join ()
453+
454+ actual = (left , right )
455+ expected = [
456+ (union , original_left ),
457+ (original_right , union ),
458+ ]
459+ self .assertIn (actual , expected )
460+
461+ # TODO: test_intersection_update_concurrent
462+ # TODO: test_intersection_update_opposing
181463
182464@threading_helper .requires_working_threading ()
183465class SmallSetTest (RaceTestBase , unittest .TestCase ):
0 commit comments