diff --git a/Lib/test/test_free_threading/test_set.py b/Lib/test/test_free_threading/test_set.py index 9dd3d68d5dad135..c3e11dbbb8e1d7f 100644 --- a/Lib/test/test_free_threading/test_set.py +++ b/Lib/test/test_free_threading/test_set.py @@ -148,6 +148,187 @@ def read_set(): for t in threads: t.join() + def test_intersection_update_concurrent(self): + """Test one-operand intersection updates of one shared set.""" + NUM_ITERS = 10 + BLOCK_SIZE = self.SET_SIZE * 100 + + sources = [ + set(range(4 * BLOCK_SIZE)), + set(range(3 * BLOCK_SIZE)), + set(range(2 * BLOCK_SIZE)), + set(range(BLOCK_SIZE)), + ] + expected = set(range(BLOCK_SIZE)) + + for _ in range(NUM_ITERS): + target = set(range(5 * BLOCK_SIZE)) + barrier = Barrier(len(sources), timeout=2) + + def intersect(source): + barrier.wait() + target.intersection_update(source) + + threads = [Thread(target=intersect, args=(source,)) + for source in sources] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + + self.assertEqual(target, expected) + + def test_intersection_update_multiple_concurrent(self): + """Test multi-operand intersection updates of one shared set.""" + NUM_ITERS = 10 + BLOCK_SIZE = self.SET_SIZE * 100 + + evens = set(range(0, 4 * BLOCK_SIZE, 2)) + below_three_blocks = set(range(3 * BLOCK_SIZE)) + multiples_of_three = set(range(0, 2 * BLOCK_SIZE, 3)) + below_one_block = set(range(BLOCK_SIZE)) + expected = set(range(0, BLOCK_SIZE, 6)) + + for _ in range(NUM_ITERS): + target = set(range(5 * BLOCK_SIZE)) + barrier = Barrier(2, timeout=2) + + def intersect(first, second): + barrier.wait() + target.intersection_update(first, second) + + threads = [ + Thread(target=intersect, + args=(evens, below_three_blocks)), + Thread(target=intersect, + args=(multiples_of_three, below_one_block)), + ] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + + self.assertEqual(target, expected) + + def test_intersection_update_three_operands_concurrent(self): + """Test three-operand intersection updates of one shared set.""" + NUM_ITERS = 10 + BLOCK_SIZE = self.SET_SIZE * 100 + + updates = [ + ( + set(range(0, 10 * BLOCK_SIZE, 2)), + set(range(0, 9 * BLOCK_SIZE, 3)), + set(range(0, 8 * BLOCK_SIZE, 5)), + ), + ( + set(range(0, 6 * BLOCK_SIZE, 5)), + set(range(0, 5 * BLOCK_SIZE, 2)), + set(range(0, 4 * BLOCK_SIZE)), + ), + ( + set(range(0, 3 * BLOCK_SIZE, 3)), + set(range(0, 2 * BLOCK_SIZE, 5)), + set(range(0, BLOCK_SIZE)), + ), + ( + set(range(0, 9 * BLOCK_SIZE, 2)), + set(range(0, 8 * BLOCK_SIZE, 3)), + set(range(0, 7 * BLOCK_SIZE)), + ), + ] + expected = set(range(0, BLOCK_SIZE, 30)) + + for _ in range(NUM_ITERS): + target = set(range(10 * BLOCK_SIZE)) + barrier = Barrier(len(updates), timeout=2) + + def intersect(first, second, third): + barrier.wait() + target.intersection_update(first, second, third) + + threads = [Thread(target=intersect, args=operands) + for operands in updates] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + + self.assertEqual(target, expected) + + def test_intersection_update_suspended_lock(self): + """Test an update while a later operand's lock is held.""" + NUM_ITERS = 200 + BLOCK_SIZE = self.SET_SIZE * 1_000 + HOLD_WORK = 10_000 + + initial = set(range(BLOCK_SIZE)) + blocked_operand = set(initial) + evens = set(range(0, BLOCK_SIZE, 2)) + threes = set(range(0, BLOCK_SIZE, 3)) + expected = set(range(0, BLOCK_SIZE, 6)) + + class SlowEmpty: + def __iter__(self): + for _ in range(HOLD_WORK): + pass + return iter(()) + + # Repeat because the lock suspension depends on thread scheduling. + for _ in range(NUM_ITERS): + target = set(initial) + + def hold_operand(): + blocked_operand.intersection(SlowEmpty()) + + def first_update(): + target.intersection_update(evens, blocked_operand) + + def second_update(): + target.intersection_update(threes) + + threads = [ + Thread(target=hold_operand), + Thread(target=first_update), + Thread(target=second_update), + ] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + + self.assertEqual(target, expected) + + def test_iand_concurrent(self): + """Test concurrent &= operations on one shared set.""" + NUM_ITERS = 10 + BLOCK_SIZE = self.SET_SIZE * 100 + + sources = [ + set(range(4 * BLOCK_SIZE)), + set(range(3 * BLOCK_SIZE)), + set(range(2 * BLOCK_SIZE)), + set(range(BLOCK_SIZE)), + ] + expected = set(range(BLOCK_SIZE)) + + for _ in range(NUM_ITERS): + target = set(range(5 * BLOCK_SIZE)) + barrier = Barrier(len(sources), timeout=2) + + def intersect(source): + barrier.wait() + target.__iand__(source) + + threads = [Thread(target=intersect, args=(source,)) + for source in sources] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + + self.assertEqual(target, expected) + @threading_helper.requires_working_threading() class SmallSetTest(RaceTestBase, unittest.TestCase): diff --git a/Misc/NEWS.d/next/Core_and_Builtins/2026-10-03-18-48-41.gh-issue-158600.RZt2lC.rst b/Misc/NEWS.d/next/Core_and_Builtins/2026-10-03-18-48-41.gh-issue-158600.RZt2lC.rst new file mode 100644 index 000000000000000..e8d58f88deb2d79 --- /dev/null +++ b/Misc/NEWS.d/next/Core_and_Builtins/2026-10-03-18-48-41.gh-issue-158600.RZt2lC.rst @@ -0,0 +1 @@ +Fix a race condition in :meth:`set.intersection_update` on free-threaded builds that could cause concurrent updates to the same set to be lost. diff --git a/Objects/setobject.c b/Objects/setobject.c index 8fdd1eb26118c0a..3e8b35e042e15f1 100644 --- a/Objects/setobject.c +++ b/Objects/setobject.c @@ -1861,15 +1861,53 @@ set_intersection_update_multi_impl(PySetObject *so, PyObject * const *others, Py_ssize_t others_length) /*[clinic end generated code: output=d768b5584675b48d input=782e422fc370e4fc]*/ { - PyObject *tmp; + Py_ssize_t i; + PyObject *other; + PyObject *result; - tmp = set_intersection_multi_impl(so, others, others_length); - if (tmp == NULL) + if (others_length == 0) { + Py_RETURN_NONE; + } + + other = others[0]; + Py_BEGIN_CRITICAL_SECTION2(so, other); + result = set_intersection(so, other); + if (result != NULL) { + for (i = 1; i < others_length; i++) { + PyObject *newresult; + + other = others[i]; + Py_BEGIN_CRITICAL_SECTION(other); + newresult = set_intersection((PySetObject *)result, other); + Py_END_CRITICAL_SECTION(); + if (newresult == NULL) { + Py_CLEAR(result); + break; + } + Py_SETREF(result, newresult); + } + if (result != NULL) { + if (others_length == 1) { + set_swap_bodies(so, (PySetObject *)result); + } + else { + /* A later operand's critical section may have suspended + so's lock. Preserve concurrent removals by intersecting + the result with so's current contents. */ + PyObject *updated = set_intersection_update(so, result); + if (updated == NULL) { + Py_CLEAR(result); + } + Py_XDECREF(updated); + } + } + } + Py_END_CRITICAL_SECTION2(); + + if (result == NULL) { return NULL; - Py_BEGIN_CRITICAL_SECTION(so); - set_swap_bodies(so, (PySetObject *)tmp); - Py_END_CRITICAL_SECTION(); - Py_DECREF(tmp); + } + Py_DECREF(result); Py_RETURN_NONE; }