@@ -210,6 +210,52 @@ def intersect(first, second):
210210
211211 self .assertEqual (target , expected )
212212
213+ def test_intersection_update_three_operands_concurrent (self ):
214+ """Test three-operand intersection updates of one shared set."""
215+ NUM_ITERS = 10
216+ BLOCK_SIZE = self .SET_SIZE * 100
217+
218+ updates = [
219+ (
220+ set (range (0 , 10 * BLOCK_SIZE , 2 )),
221+ set (range (0 , 9 * BLOCK_SIZE , 3 )),
222+ set (range (0 , 8 * BLOCK_SIZE , 5 )),
223+ ),
224+ (
225+ set (range (0 , 6 * BLOCK_SIZE , 5 )),
226+ set (range (0 , 5 * BLOCK_SIZE , 2 )),
227+ set (range (0 , 4 * BLOCK_SIZE )),
228+ ),
229+ (
230+ set (range (0 , 3 * BLOCK_SIZE , 3 )),
231+ set (range (0 , 2 * BLOCK_SIZE , 5 )),
232+ set (range (0 , BLOCK_SIZE )),
233+ ),
234+ (
235+ set (range (0 , 9 * BLOCK_SIZE , 2 )),
236+ set (range (0 , 8 * BLOCK_SIZE , 3 )),
237+ set (range (0 , 7 * BLOCK_SIZE )),
238+ ),
239+ ]
240+ expected = set (range (0 , BLOCK_SIZE , 30 ))
241+
242+ for _ in range (NUM_ITERS ):
243+ target = set (range (10 * BLOCK_SIZE ))
244+ barrier = Barrier (len (updates ), timeout = 2 )
245+
246+ def intersect (first , second , third ):
247+ barrier .wait ()
248+ target .intersection_update (first , second , third )
249+
250+ threads = [Thread (target = intersect , args = operands )
251+ for operands in updates ]
252+ for thread in threads :
253+ thread .start ()
254+ for thread in threads :
255+ thread .join ()
256+
257+ self .assertEqual (target , expected )
258+
213259 def test_intersection_update_suspended_lock (self ):
214260 """Test an update while a later operand's lock is held."""
215261 NUM_ITERS = 200
0 commit comments