test(call): wait on the provider instead of counting event-loop yields

D2 raced the in-flight winner by yielding the event loop up to 1000 times
and asserting the call had reached the fake provider by then. An idle loop
returns from sleep(0) immediately, so that budget is spent at a rate set by
the machine rather than by the call's progress: locally ~260 of the ~290
yields were sub-0.1ms no-ops, and on CI the budget ran out before the first
call arrived.

FakeProvider now signals arrivals on an asyncio.Event, and the test waits on
it with a wall-clock ceiling.
This commit is contained in:
SToneX
2026-08-25 11:32:26 +08:00
parent 65d32eb810
commit 06873fbca2
2 changed files with 24 additions and 5 deletions
+23
View File
@@ -29,12 +29,34 @@ class FakeProvider:
def __init__(self) -> None:
self.hits: list[Hit] = []
self._arrived = asyncio.Event()
self.app = FastAPI()
self.app.add_api_route(
"/{path:path}", self._respond,
methods=["GET", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"],
)
async def wait_for_hits(self, count: int, timeout: float = 30.0) -> None:
"""Block until ``count`` requests have arrived, or fail the wait.
Callers racing an in-flight call must wait on this rather than yielding
the event loop a fixed number of times: an idle loop returns from
``sleep(0)`` immediately, so a yield budget is spent at a rate set by
the machine, not by the call's progress.
"""
try:
async with asyncio.timeout(timeout):
while True:
self._arrived.clear()
if len(self.hits) >= count:
return
await self._arrived.wait()
except TimeoutError:
raise AssertionError(
f"only {len(self.hits)} of {count} calls reached the fake provider "
f"in {timeout}s",
) from None
async def _respond(self, request: Request) -> Response:
body = await request.body()
headers = {key.lower(): value for key, value in request.headers.items()}
@@ -47,6 +69,7 @@ class FakeProvider:
headers=headers,
body=body,
))
self._arrived.set()
if delay := headers.get("x-fake-sleep"):
await asyncio.sleep(float(delay))
+1 -5
View File
@@ -364,11 +364,7 @@ async def test_d2_concurrent_same_key_loser_gets_409(
first_task = asyncio.create_task(
matrix_clients.get(f"/call/{EP}?aweme_id=7", headers=headers),
)
for _ in range(1_000):
if len(fake_provider.hits) > before.hit_count:
break
await asyncio.sleep(0)
assert len(fake_provider.hits) == before.hit_count + 1, "first call never reached fake provider"
await fake_provider.wait_for_hits(before.hit_count + 1)
loser = await matrix_clients.get(f"/call/{EP}?aweme_id=7", headers=headers)
winner = await first_task