Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
312 changes: 312 additions & 0 deletions Lib/test/test_free_threading/test_set.py
Original file line number Diff line number Diff line change
Expand Up @@ -148,6 +148,318 @@ def read_set():
for t in threads:
t.join()

def test_pop_concurrent(self):
"""Test set.pop() from several threads."""
NUM_THREADS = 4
NUM_ITERS = 20

for _ in range(NUM_ITERS):
items = set(range(self.SET_SIZE))
s = set(items)
barrier = Barrier(NUM_THREADS, timeout=2)
popped = [[] for _ in range(NUM_THREADS)]

def pop_set(out):
barrier.wait()
while True:
try:
out.append(s.pop())
except KeyError:
break

threads = [Thread(target=pop_set, args=(out,)) for out in popped]
for t in threads:
t.start()
for t in threads:
t.join()

all_popped = [item for out in popped for item in out]
self.assertEqual(len(all_popped), len(items))
self.assertEqual(set(all_popped), items)
self.assertEqual(len(s), 0)

def test_add_concurrent(self):
"""Test set.add() with disjoint inputs from several threads."""
NUM_THREADS = 4
NUM_ITERS = 20

for _ in range(NUM_ITERS):
inputs = [
range(i * self.SET_SIZE, (i + 1) * self.SET_SIZE)
for i in range(NUM_THREADS)
]
expected = set().union(*inputs)
s = set()
barrier = Barrier(NUM_THREADS, timeout=2)

def add_items(items):
barrier.wait()
for item in items:
s.add(item)

threads = [Thread(target=add_items, args=(items,))
for items in inputs]
for t in threads:
t.start()
for t in threads:
t.join()

self.assertEqual(s, expected)

def test_remove_discard_concurrent(self):
"""Test set.remove() and set.discard() from several threads."""
NUM_THREADS = 4
NUM_ITERS = 20

for method_name in ("remove", "discard"):
with self.subTest(method=method_name):
for _ in range(NUM_ITERS):
inputs = [
range(i * self.SET_SIZE, (i + 1) * self.SET_SIZE)
for i in range(NUM_THREADS)
]
untouched = set(range(
NUM_THREADS * self.SET_SIZE,
(NUM_THREADS + 1) * self.SET_SIZE,
))
s = set().union(untouched, *inputs)
barrier = Barrier(NUM_THREADS, timeout=2)

def remove_items(items):
barrier.wait()
method = getattr(s, method_name)
for item in items:
method(item)

threads = [Thread(target=remove_items, args=(items,))
for items in inputs]
for t in threads:
t.start()
for t in threads:
t.join()

self.assertEqual(s, untouched)

def test_copy_clear_concurrent(self):
"""Test set.copy() while another thread clears the set."""
NUM_ITERS = 20

for _ in range(NUM_ITERS):
items = set(range(self.SET_SIZE))
s = set(items)
copies = []
barrier = Barrier(2, timeout=2)

def copy_set():
barrier.wait()
copies.append(s.copy())

def clear_set():
barrier.wait()
s.clear()

threads = [Thread(target=copy_set), Thread(target=clear_set)]
for t in threads:
t.start()
for t in threads:
t.join()

self.assertIn(copies[0], (items, set()))
self.assertEqual(s, set())

def test_update_concurrent(self):
"""Test updates of one shared set from disjoint source sets."""
NUM_THREADS = 4
NUM_ITERS = 20

sources = [
set(range(i * self.SET_SIZE, (i + 1) * self.SET_SIZE))
for i in range(NUM_THREADS)
]
original_sources = [set(source) for source in sources]
initial = set(range(
NUM_THREADS * self.SET_SIZE,
(NUM_THREADS + 1) * self.SET_SIZE,
))
expected = set().union(initial, *sources)

for _ in range(NUM_ITERS):
s = set(initial)
barrier = Barrier(NUM_THREADS, timeout=2)

def update_set(source):
barrier.wait()
s.update(source)

threads = [Thread(target=update_set, args=(source,))
for source in sources]
for t in threads:
t.start()
for t in threads:
t.join()

self.assertEqual(s, expected)
self.assertEqual(sources, original_sources)

def test_update_opposing(self):
"""Test opposing updates of two shared sets."""
NUM_ITERS = 20

for _ in range(NUM_ITERS):
left = set(range(self.SET_SIZE))
right = set(range(self.SET_SIZE, self.SET_SIZE * 2))
expected = set(range(self.SET_SIZE * 2))
barrier = Barrier(2, timeout=2)

def update_set(target, source):
barrier.wait()
target.update(source)

threads = [
Thread(target=update_set, args=(left, right)),
Thread(target=update_set, args=(right, left)),
]
for t in threads:
t.start()
for t in threads:
t.join()

self.assertEqual(left, expected)
self.assertEqual(right, expected)

def test_difference_update_concurrent(self):
"""Test set.difference_update() with disjoint source sets."""
NUM_THREADS = 4
NUM_ITERS = 20

sources = [
set(range(i * self.SET_SIZE, (i + 1) * self.SET_SIZE))
for i in range(NUM_THREADS)
]
original_sources = [set(source) for source in sources]
untouched = set(range(
NUM_THREADS * self.SET_SIZE,
(NUM_THREADS + 1) * self.SET_SIZE,
))
items = set().union(untouched, *sources)

for _ in range(NUM_ITERS):
s = set(items)
barrier = Barrier(NUM_THREADS, timeout=2)

def difference_update(source):
barrier.wait()
s.difference_update(source)

threads = [Thread(target=difference_update, args=(source,))
for source in sources]
for t in threads:
t.start()
for t in threads:
t.join()

self.assertEqual(s, untouched)
self.assertEqual(sources, original_sources)

def test_difference_update_opposing(self):
"""Test opposing difference updates of two shared sets."""
NUM_ITERS = 20

for _ in range(NUM_ITERS):
left = set(range(self.SET_SIZE))
right = set(range(1, self.SET_SIZE + 1))
original_left = set(left)
original_right = set(right)
left_only = {0}
right_only = {self.SET_SIZE}
barrier = Barrier(2, timeout=2)

def difference_update(target, source):
barrier.wait()
target.difference_update(source)

threads = [
Thread(target=difference_update, args=(left, right)),
Thread(target=difference_update, args=(right, left)),
]
for t in threads:
t.start()
for t in threads:
t.join()

actual = (left, right)
expected = [
(left_only, original_right),
(original_left, right_only),
]
self.assertIn(actual, expected)

def test_symmetric_difference_update_concurrent(self):
"""Test symmetric difference updates with disjoint source sets."""
NUM_THREADS = 4
NUM_ITERS = 20

sources = [
set(range(i * self.SET_SIZE, (i + 1) * self.SET_SIZE))
for i in range(NUM_THREADS)
]
items = set().union(*sources)
initial = {item for item in items if item % 2 == 0}
expected = {item for item in items if item % 2 == 1}

for _ in range(NUM_ITERS):
s = set(initial)
barrier = Barrier(NUM_THREADS, timeout=2)

def symmetric_difference_update(source):
barrier.wait()
s.symmetric_difference_update(source)

threads = [
Thread(target=symmetric_difference_update, args=(source,))
for source in sources
]
for t in threads:
t.start()
for t in threads:
t.join()

self.assertEqual(s, expected)

def test_symmetric_difference_update_opposing(self):
"""Test opposing symmetric difference updates of shared sets."""
NUM_ITERS = 20

for _ in range(NUM_ITERS):
left = set(range(self.SET_SIZE))
right = set(range(self.SET_SIZE, self.SET_SIZE * 2))
original_left = set(left)
original_right = set(right)
union = set(range(self.SET_SIZE * 2))
barrier = Barrier(2, timeout=2)

def symmetric_update(target, source):
barrier.wait()
target.symmetric_difference_update(source)

threads = [
Thread(target=symmetric_update, args=(left, right)),
Thread(target=symmetric_update, args=(right, left)),
]
for t in threads:
t.start()
for t in threads:
t.join()

actual = (left, right)
expected = [
(union, original_left),
(original_right, union),
]
self.assertIn(actual, expected)

# TODO: test_intersection_update_concurrent
# TODO: test_intersection_update_opposing

@threading_helper.requires_working_threading()
class SmallSetTest(RaceTestBase, unittest.TestCase):
Expand Down
Loading