Skip to content

Commit bcf24f3

Browse files
fix(concurrency): recheck throttling after queued permit admission (#150)
Signed-off-by: Rudy Celekli <47457359+rudycelekli@users.noreply.github.com> Co-authored-by: n-papaioannou <n.papaioannou@ime.life>
1 parent dc4d745 commit bcf24f3

2 files changed

Lines changed: 75 additions & 0 deletions

File tree

‎ifixai/core/concurrency.py‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -101,6 +101,11 @@ async def acquire(self):
101101
async with self._throttle_cond:
102102
await self._throttle_cond.wait_for(lambda: not self._throttled)
103103
async with self._semaphore:
104+
# This call may have queued before a 429 started the cooldown.
105+
# Recheck after admission so releasing an in-flight permit cannot
106+
# silently bypass the shared throttle.
107+
async with self._throttle_cond:
108+
await self._throttle_cond.wait_for(lambda: not self._throttled)
104109
yield
105110

106111
async def reserve_judge_call(self) -> None:
Lines changed: 70 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,70 @@
1+
import asyncio
2+
3+
import pytest
4+
5+
from ifixai.core.concurrency import ConcurrencyGovernor
6+
7+
8+
@pytest.mark.asyncio
9+
async def test_waiter_queued_before_rate_limit_does_not_issue_during_cooldown():
10+
governor = ConcurrencyGovernor(1)
11+
entered = asyncio.Event()
12+
13+
async def wait():
14+
async with governor.acquire():
15+
entered.set()
16+
17+
async with governor.acquire():
18+
waiter = asyncio.create_task(wait())
19+
await asyncio.sleep(0)
20+
await governor.on_rate_limit()
21+
try:
22+
await asyncio.sleep(0.01)
23+
assert not entered.is_set(), (
24+
"a queued request bypassed the active rate-limit cooldown"
25+
)
26+
async with governor._throttle_cond:
27+
governor._throttled = False
28+
governor._throttle_cond.notify_all()
29+
await asyncio.wait_for(waiter, 1)
30+
assert entered.is_set()
31+
finally:
32+
if governor._recovery_task:
33+
governor._recovery_task.cancel()
34+
await asyncio.gather(governor._recovery_task, return_exceptions=True)
35+
waiter.cancel()
36+
await asyncio.gather(waiter, return_exceptions=True)
37+
38+
39+
@pytest.mark.asyncio
40+
async def test_cancelled_cooldown_waiter_does_not_leak_permit():
41+
governor = ConcurrencyGovernor(1)
42+
entered = asyncio.Event()
43+
44+
async def acquire_again():
45+
async with governor.acquire():
46+
entered.set()
47+
48+
async with governor.acquire():
49+
waiter = asyncio.create_task(acquire_again())
50+
await asyncio.sleep(0)
51+
await governor.on_rate_limit()
52+
try:
53+
# The queued waiter acquires the released semaphore, then must wait at
54+
# the second throttle gate while it holds the permit.
55+
await asyncio.sleep(0)
56+
await asyncio.sleep(0)
57+
assert governor._semaphore._value == 0
58+
assert not entered.is_set()
59+
waiter.cancel()
60+
await asyncio.gather(waiter, return_exceptions=True)
61+
async with governor._throttle_cond:
62+
governor._throttled = False
63+
governor._throttle_cond.notify_all()
64+
await asyncio.wait_for(acquire_again(), 1)
65+
assert entered.is_set()
66+
finally:
67+
waiter.cancel()
68+
await asyncio.gather(waiter, return_exceptions=True)
69+
governor._recovery_task.cancel()
70+
await asyncio.gather(governor._recovery_task, return_exceptions=True)

0 commit comments

Comments
 (0)