fix(tests): await gateway drain completion

This commit is contained in:
Teknium 2026-08-08 18:05:14 -07:00
parent 6601330e0a
commit e2157b8697
1 changed files with 17 additions and 13 deletions

View File

@ -59,7 +59,9 @@ def _make_event(text="hi", chat_id="42"):
return MessageEvent(
text=text,
message_type=MessageType.TEXT,
source=SessionSource(platform=Platform.TELEGRAM, chat_id=chat_id, chat_type="dm"),
source=SessionSource(
platform=Platform.TELEGRAM, chat_id=chat_id, chat_type="dm"
),
)
@ -81,10 +83,13 @@ async def test_pending_drain_keeps_active_session_guard_live():
# pending message arrives.
first_started = asyncio.Event()
release_first = asyncio.Event()
second_processed = asyncio.Event()
async def handler(event):
first_started.set()
await release_first.wait()
if event.text == "M2":
second_processed.set()
return "done"
adapter._message_handler = handler
@ -125,8 +130,8 @@ async def test_pending_drain_keeps_active_session_guard_live():
"the old Event may have waiters that now won't be signaled"
)
# Finish drain.
await asyncio.sleep(0.1)
# Finish drain without relying on scheduler speed.
await asyncio.wait_for(second_processed.wait(), timeout=2.0)
await adapter.cancel_background_tasks()
@ -139,9 +144,12 @@ async def test_finally_cleanup_drains_late_arrival_pending():
sk = _sk()
processed = []
late_processed = asyncio.Event()
async def handler(event):
processed.append(event.text)
if event.text == "LATE":
late_processed.set()
return "ok"
adapter._message_handler = handler
@ -169,11 +177,8 @@ async def test_finally_cleanup_drains_late_arrival_pending():
# Send M1.
await adapter.handle_message(_make_event(text="M1"))
# Drain: wait for M1 to finish and the late-drain task to process LATE.
for _ in range(50): # up to ~0.5s
if "LATE" in processed:
break
await asyncio.sleep(0.01)
# Drain: wait for the late-drain task itself to process LATE.
await asyncio.wait_for(late_processed.wait(), timeout=2.0)
await adapter.cancel_background_tasks()
@ -198,11 +203,10 @@ async def test_no_pending_cleans_up_normally():
await adapter.handle_message(_make_event(text="solo"))
# Wait for background task to finish.
for _ in range(50):
if sk not in adapter._active_sessions:
break
await asyncio.sleep(0.01)
# Await the task that owns this session rather than sampling cleanup after
# an arbitrary wall-clock delay.
owner_task = adapter._session_tasks[sk]
await asyncio.wait_for(asyncio.shield(owner_task), timeout=2.0)
assert sk not in adapter._active_sessions, (
"_active_sessions was not cleaned up after a normal turn with no pending"