Skip to content

Commit a282b36

Browse files
committed
gh-158600: Fix lost updates in concurrent set.intersection_update()
Calculate the intersection from a copy while holding the target set's critical section. This prevents concurrent intersection_update() calls from overwriting one another's results. Add free-threading tests for one-operand and multi-operand calls, with __iand__() as a control.
1 parent b42fdf6 commit a282b36

2 files changed

Lines changed: 108 additions & 6 deletions

File tree

‎Lib/test/test_free_threading/test_set.py‎

Lines changed: 92 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -148,6 +148,98 @@ def read_set():
148148
for t in threads:
149149
t.join()
150150

151+
def test_intersection_update_concurrent(self):
152+
"""Test one-operand intersection updates of one shared set."""
153+
NUM_ITERS = 10
154+
BLOCK_SIZE = self.SET_SIZE * 100
155+
156+
sources = [
157+
set(range(4 * BLOCK_SIZE)),
158+
set(range(3 * BLOCK_SIZE)),
159+
set(range(2 * BLOCK_SIZE)),
160+
set(range(BLOCK_SIZE)),
161+
]
162+
expected = set(range(BLOCK_SIZE))
163+
164+
for _ in range(NUM_ITERS):
165+
target = set(range(5 * BLOCK_SIZE))
166+
barrier = Barrier(len(sources), timeout=2)
167+
168+
def intersect(source):
169+
barrier.wait()
170+
target.intersection_update(source)
171+
172+
threads = [Thread(target=intersect, args=(source,))
173+
for source in sources]
174+
for thread in threads:
175+
thread.start()
176+
for thread in threads:
177+
thread.join()
178+
179+
self.assertEqual(target, expected)
180+
181+
def test_intersection_update_multiple_concurrent(self):
182+
"""Test multi-operand intersection updates of one shared set."""
183+
NUM_ITERS = 10
184+
BLOCK_SIZE = self.SET_SIZE * 100
185+
186+
evens = set(range(0, 4 * BLOCK_SIZE, 2))
187+
below_three_blocks = set(range(3 * BLOCK_SIZE))
188+
multiples_of_three = set(range(0, 2 * BLOCK_SIZE, 3))
189+
below_one_block = set(range(BLOCK_SIZE))
190+
expected = set(range(0, BLOCK_SIZE, 6))
191+
192+
for _ in range(NUM_ITERS):
193+
target = set(range(5 * BLOCK_SIZE))
194+
barrier = Barrier(2, timeout=2)
195+
196+
def intersect(first, second):
197+
barrier.wait()
198+
target.intersection_update(first, second)
199+
200+
threads = [
201+
Thread(target=intersect,
202+
args=(evens, below_three_blocks)),
203+
Thread(target=intersect,
204+
args=(multiples_of_three, below_one_block)),
205+
]
206+
for thread in threads:
207+
thread.start()
208+
for thread in threads:
209+
thread.join()
210+
211+
self.assertEqual(target, expected)
212+
213+
def test_iand_concurrent(self):
214+
"""Test concurrent &= operations on one shared set."""
215+
NUM_ITERS = 10
216+
BLOCK_SIZE = self.SET_SIZE * 100
217+
218+
sources = [
219+
set(range(4 * BLOCK_SIZE)),
220+
set(range(3 * BLOCK_SIZE)),
221+
set(range(2 * BLOCK_SIZE)),
222+
set(range(BLOCK_SIZE)),
223+
]
224+
expected = set(range(BLOCK_SIZE))
225+
226+
for _ in range(NUM_ITERS):
227+
target = set(range(5 * BLOCK_SIZE))
228+
barrier = Barrier(len(sources), timeout=2)
229+
230+
def intersect(source):
231+
barrier.wait()
232+
target.__iand__(source)
233+
234+
threads = [Thread(target=intersect, args=(source,))
235+
for source in sources]
236+
for thread in threads:
237+
thread.start()
238+
for thread in threads:
239+
thread.join()
240+
241+
self.assertEqual(target, expected)
242+
151243

152244
@threading_helper.requires_working_threading()
153245
class SmallSetTest(RaceTestBase, unittest.TestCase):

‎Objects/setobject.c‎

Lines changed: 16 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1861,15 +1861,25 @@ set_intersection_update_multi_impl(PySetObject *so, PyObject * const *others,
18611861
Py_ssize_t others_length)
18621862
/*[clinic end generated code: output=d768b5584675b48d input=782e422fc370e4fc]*/
18631863
{
1864-
PyObject *tmp;
1864+
PyObject *copy;
1865+
PyObject *result = NULL;
18651866

1866-
tmp = set_intersection_multi_impl(so, others, others_length);
1867-
if (tmp == NULL)
1868-
return NULL;
18691867
Py_BEGIN_CRITICAL_SECTION(so);
1870-
set_swap_bodies(so, (PySetObject *)tmp);
1868+
copy = set_copy_untracked_lock_held(so);
1869+
if (copy != NULL) {
1870+
result = set_intersection_multi_impl((PySetObject *)copy,
1871+
others, others_length);
1872+
Py_DECREF(copy);
1873+
if (result != NULL) {
1874+
set_swap_bodies(so, (PySetObject *)result);
1875+
}
1876+
}
18711877
Py_END_CRITICAL_SECTION();
1872-
Py_DECREF(tmp);
1878+
1879+
if (result == NULL) {
1880+
return NULL;
1881+
}
1882+
Py_DECREF(result);
18731883
Py_RETURN_NONE;
18741884
}
18751885

0 commit comments

Comments
 (0)