Skip to content

Commit 1cac88d

Browse files
committed
Test concurrent set mutations, copying, and clearing. Cover both shared-target
and opposing-target update operations.
1 parent 96a0a5c commit 1cac88d

1 file changed

Lines changed: 282 additions & 0 deletions

File tree

‎Lib/test/test_free_threading/test_set.py‎

Lines changed: 282 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -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()
183465
class SmallSetTest(RaceTestBase, unittest.TestCase):

0 commit comments

Comments
 (0)