| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent fb87aaa commit 773dbb9
3 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -346,9 +346,8 @@ class Semaphore(_ContextManagerMixin, mixins._LoopBoundMixin): | |||
| 346 | 346 | def __init__(self, value=1): | |
| 347 | 347 | if value < 0: | |
| 348 | 348 | raise ValueError("Semaphore initial value must be >= 0") | |
| 349 | + self._waiters = None | ||
| 349 | 350 | self._value = value | |
| 350 | - self._waiters = collections.deque() | ||
| 351 | - self._wakeup_scheduled = False | ||
| 352 | 351 | ||
| 353 | 352 | def __repr__(self): | |
| 354 | 353 | res = super().__repr__() | |
@@ -357,16 +356,8 @@ def __repr__(self): | |||
| 357 | 356 | extra = f'{extra}, waiters:{len(self._waiters)}' | |
| 358 | 357 | return f'<{res[1:-1]} [{extra}]>' | |
| 359 | 358 | ||
| 360 | - def _wake_up_next(self): | ||
| 361 | - while self._waiters: | ||
| 362 | - waiter = self._waiters.popleft() | ||
| 363 | - if not waiter.done(): | ||
| 364 | - waiter.set_result(None) | ||
| 365 | - self._wakeup_scheduled = True | ||
| 366 | - return | ||
| 367 | - | ||
| 368 | 359 | def locked(self): | |
| 369 | - """Returns True if semaphore can not be acquired immediately.""" | ||
| 360 | + """Returns True if semaphore counter is zero.""" | ||
| 370 | 361 | return self._value == 0 | |
| 371 | 362 | ||
| 372 | 363 | async def acquire(self): | |
@@ -378,28 +369,57 @@ async def acquire(self): | |||
| 378 | 369 | called release() to make it larger than 0, and then return | |
| 379 | 370 | True. | |
| 380 | 371 | """ | |
| 381 | - # _wakeup_scheduled is set if *another* task is scheduled to wakeup | ||
| 382 | - # but its acquire() is not resumed yet | ||
| 383 | - while self._wakeup_scheduled or self._value <= 0: | ||
| 384 | - fut = self._get_loop().create_future() | ||
| 385 | - self._waiters.append(fut) | ||
| 372 | + if (not self.locked() and (self._waiters is None or | ||
| 373 | + all(w.cancelled() for w in self._waiters))): | ||
| 374 | + self._value -= 1 | ||
| 375 | + return True | ||
| 376 | + | ||
| 377 | + if self._waiters is None: | ||
| 378 | + self._waiters = collections.deque() | ||
| 379 | + fut = self._get_loop().create_future() | ||
| 380 | + self._waiters.append(fut) | ||
| 381 | + | ||
| 382 | + # Finally block should be called before the CancelledError | ||
| 383 | + # handling as we don't want CancelledError to call | ||
| 384 | + # _wake_up_first() and attempt to wake up itself. | ||
| 385 | + try: | ||
| 386 | 386 | try: | |
| 387 | 387 | await fut | |
| 388 | - # reset _wakeup_scheduled *after* waiting for a future | ||
| 389 | - self._wakeup_scheduled = False | ||
| 390 | - except exceptions.CancelledError: | ||
| 391 | - self._wake_up_next() | ||
| 392 | - raise | ||
| 388 | + finally: | ||
| 389 | + self._waiters.remove(fut) | ||
| 390 | + except exceptions.CancelledError: | ||
| 391 | + if not self.locked(): | ||
| 392 | + self._wake_up_first() | ||
| 393 | + raise | ||
| 394 | + | ||
| 393 | 395 | self._value -= 1 | |
| 396 | + if not self.locked(): | ||
| 397 | + self._wake_up_first() | ||
| 394 | 398 | return True | |
| 395 | 399 | ||
| 396 | 400 | def release(self): | |
| 397 | 401 | """Release a semaphore, incrementing the internal counter by one. | |
| 402 | + | ||
| 398 | 403 | When it was zero on entry and another coroutine is waiting for it to | |
| 399 | 404 | become larger than zero again, wake up that coroutine. | |
| 400 | 405 | """ | |
| 401 | 406 | self._value += 1 | |
| 402 | - self._wake_up_next() | ||
| 407 | + self._wake_up_first() | ||
| 408 | + | ||
| 409 | + def _wake_up_first(self): | ||
| 410 | + """Wake up the first waiter if it isn't done.""" | ||
| 411 | + if not self._waiters: | ||
| 412 | + return | ||
| 413 | + try: | ||
| 414 | + fut = next(iter(self._waiters)) | ||
| 415 | + except StopIteration: | ||
| 416 | + return | ||
| 417 | + | ||
| 418 | + # .done() necessarily means that a waiter will wake up later on and | ||
| 419 | + # either take the lock, or, if it was cancelled and lock wasn't | ||
| 420 | + # taken already, will hit this again and wake up a new waiter. | ||
| 421 | + if not fut.done(): | ||
| 422 | + fut.set_result(True) | ||
| 403 | 423 | ||
| 404 | 424 | ||
| 405 | 425 | class BoundedSemaphore(Semaphore): | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -5,6 +5,7 @@ | |||
| 5 | 5 | import re | |
| 6 | 6 | ||
| 7 | 7 | import asyncio | |
| 8 | + import collections | ||
| 8 | 9 | ||
| 9 | 10 | STR_RGX_REPR = ( | |
| 10 | 11 | r'^<(?P<class>.*?) object at (?P<address>.*?)' | |
@@ -774,6 +775,9 @@ async def test_repr(self): | |||
| 774 | 775 | self.assertTrue('waiters' not in repr(sem)) | |
| 775 | 776 | self.assertTrue(RGX_REPR.match(repr(sem))) | |
| 776 | 777 | ||
| 778 | + if sem._waiters is None: | ||
| 779 | + sem._waiters = collections.deque() | ||
| 780 | + | ||
| 777 | 781 | sem._waiters.append(mock.Mock()) | |
| 778 | 782 | self.assertTrue('waiters:1' in repr(sem)) | |
| 779 | 783 | self.assertTrue(RGX_REPR.match(repr(sem))) | |
@@ -842,6 +846,7 @@ async def c4(result): | |||
| 842 | 846 | sem.release() | |
| 843 | 847 | self.assertEqual(2, sem._value) | |
| 844 | 848 | ||
| 849 | + await asyncio.sleep(0) | ||
| 845 | 850 | await asyncio.sleep(0) | |
| 846 | 851 | self.assertEqual(0, sem._value) | |
| 847 | 852 | self.assertEqual(3, len(result)) | |
@@ -884,6 +889,7 @@ async def test_acquire_cancel_before_awoken(self): | |||
| 884 | 889 | t2.cancel() | |
| 885 | 890 | sem.release() | |
| 886 | 891 | ||
| 892 | + await asyncio.sleep(0) | ||
| 887 | 893 | await asyncio.sleep(0) | |
| 888 | 894 | num_done = sum(t.done() for t in [t3, t4]) | |
| 889 | 895 | self.assertEqual(num_done, 1) | |
@@ -904,9 +910,32 @@ async def test_acquire_hang(self): | |||
| 904 | 910 | t1.cancel() | |
| 905 | 911 | sem.release() | |
| 906 | 912 | await asyncio.sleep(0) | |
| 913 | + await asyncio.sleep(0) | ||
| 907 | 914 | self.assertTrue(sem.locked()) | |
| 908 | 915 | self.assertTrue(t2.done()) | |
| 909 | 916 | ||
| 917 | + async def test_acquire_no_hang(self): | ||
| 918 | + | ||
| 919 | + sem = asyncio.Semaphore(1) | ||
| 920 | + | ||
| 921 | + async def c1(): | ||
| 922 | + async with sem: | ||
| 923 | + await asyncio.sleep(0) | ||
| 924 | + t2.cancel() | ||
| 925 | + | ||
| 926 | + async def c2(): | ||
| 927 | + async with sem: | ||
| 928 | + self.assertFalse(True) | ||
| 929 | + | ||
| 930 | + t1 = asyncio.create_task(c1()) | ||
| 931 | + t2 = asyncio.create_task(c2()) | ||
| 932 | + | ||
| 933 | + r1, r2 = await asyncio.gather(t1, t2, return_exceptions=True) | ||
| 934 | + self.assertTrue(r1 is None) | ||
| 935 | + self.assertTrue(isinstance(r2, asyncio.CancelledError)) | ||
| 936 | + | ||
| 937 | + await asyncio.wait_for(sem.acquire(), timeout=1.0) | ||
| 938 | + | ||
| 910 | 939 | def test_release_not_acquired(self): | |
| 911 | 940 | sem = asyncio.BoundedSemaphore() | |
| 912 | 941 | ||
@@ -945,6 +974,77 @@ async def coro(tag): | |||
| 945 | 974 | result | |
| 946 | 975 | ) | |
| 947 | 976 | ||
| 977 | + async def test_acquire_fifo_order_2(self): | ||
| 978 | + sem = asyncio.Semaphore(1) | ||
| 979 | + result = [] | ||
| 980 | + | ||
| 981 | + async def c1(result): | ||
| 982 | + await sem.acquire() | ||
| 983 | + result.append(1) | ||
| 984 | + return True | ||
| 985 | + | ||
| 986 | + async def c2(result): | ||
| 987 | + await sem.acquire() | ||
| 988 | + result.append(2) | ||
| 989 | + sem.release() | ||
| 990 | + await sem.acquire() | ||
| 991 | + result.append(4) | ||
| 992 | + return True | ||
| 993 | + | ||
| 994 | + async def c3(result): | ||
| 995 | + await sem.acquire() | ||
| 996 | + result.append(3) | ||
| 997 | + return True | ||
| 998 | + | ||
| 999 | + t1 = asyncio.create_task(c1(result)) | ||
| 1000 | + t2 = asyncio.create_task(c2(result)) | ||
| 1001 | + t3 = asyncio.create_task(c3(result)) | ||
| 1002 | + | ||
| 1003 | + await asyncio.sleep(0) | ||
| 1004 | + | ||
| 1005 | + sem.release() | ||
| 1006 | + sem.release() | ||
| 1007 | + | ||
| 1008 | + tasks = [t1, t2, t3] | ||
| 1009 | + await asyncio.gather(*tasks) | ||
| 1010 | + self.assertEqual([1, 2, 3, 4], result) | ||
| 1011 | + | ||
| 1012 | + async def test_acquire_fifo_order_3(self): | ||
| 1013 | + sem = asyncio.Semaphore(0) | ||
| 1014 | + result = [] | ||
| 1015 | + | ||
| 1016 | + async def c1(result): | ||
| 1017 | + await sem.acquire() | ||
| 1018 | + result.append(1) | ||
| 1019 | + return True | ||
| 1020 | + | ||
| 1021 | + async def c2(result): | ||
| 1022 | + await sem.acquire() | ||
| 1023 | + result.append(2) | ||
| 1024 | + return True | ||
| 1025 | + | ||
| 1026 | + async def c3(result): | ||
| 1027 | + await sem.acquire() | ||
| 1028 | + result.append(3) | ||
| 1029 | + return True | ||
| 1030 | + | ||
| 1031 | + t1 = asyncio.create_task(c1(result)) | ||
| 1032 | + t2 = asyncio.create_task(c2(result)) | ||
| 1033 | + t3 = asyncio.create_task(c3(result)) | ||
| 1034 | + | ||
| 1035 | + await asyncio.sleep(0) | ||
| 1036 | + | ||
| 1037 | + t1.cancel() | ||
| 1038 | + | ||
| 1039 | + await asyncio.sleep(0) | ||
| 1040 | + | ||
| 1041 | + sem.release() | ||
| 1042 | + sem.release() | ||
| 1043 | + | ||
| 1044 | + tasks = [t1, t2, t3] | ||
| 1045 | + await asyncio.gather(*tasks, return_exceptions=True) | ||
| 1046 | + self.assertEqual([2, 3], result) | ||
| 1047 | + | ||
| 948 | 1048 | ||
| 949 | 1049 | class BarrierTests(unittest.IsolatedAsyncioTestCase): | |
| 950 | 1050 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1 @@ | |||
| 1 | + Fix broken :class:`asyncio.Semaphore` when acquire is cancelled. | ||
| Back | FazBrowse Home | New Git URL |
0 commit comments