Skip to content

Commit aeb4f45

Browse files
nickita-khylkouskimarandaneto
authored andcommitted
fix: clean up async consumer waiters on cancellation
1 parent a1002c5 commit aeb4f45

3 files changed

Lines changed: 71 additions & 9 deletions

File tree

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,5 @@
1+
---
2+
pypi/posthog: patch
3+
---
4+
5+
Clean up queue and flush waiters when the async capture consumer is cancelled.

‎posthog/_async_consumer.py‎

Lines changed: 11 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -117,15 +117,17 @@ def request_flush(self) -> None:
117117
async def _get_or_flush(self, timeout: float) -> tuple[Any, bool]:
118118
get_task = asyncio.create_task(self.queue.get())
119119
flush_task = asyncio.create_task(self._flush_event.wait())
120-
done, pending = await asyncio.wait(
121-
{get_task, flush_task},
122-
timeout=timeout,
123-
return_when=asyncio.FIRST_COMPLETED,
124-
)
125-
for task in pending:
126-
task.cancel()
127-
if pending:
128-
await asyncio.gather(*pending, return_exceptions=True)
120+
try:
121+
done, _ = await asyncio.wait(
122+
{get_task, flush_task},
123+
timeout=timeout,
124+
return_when=asyncio.FIRST_COMPLETED,
125+
)
126+
finally:
127+
for task in (get_task, flush_task):
128+
if not task.done():
129+
task.cancel()
130+
await asyncio.gather(get_task, flush_task, return_exceptions=True)
129131

130132
if get_task in done:
131133
return get_task.result(), False

‎posthog/test/test_async_consumer.py‎

Lines changed: 55 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -152,3 +152,58 @@ async def test_request_stops_after_configured_retry_limit():
152152

153153
assert batch_post.await_count == 3
154154
assert [call.args[0] for call in sleep.await_args_list] == [1, 2]
155+
156+
157+
@pytest.mark.asyncio
158+
@pytest.mark.parametrize("run_worker", [False, True], ids=["wait", "worker"])
159+
async def test_get_or_flush_cancels_waiters_on_cancellation(run_worker):
160+
consumer = make_consumer(retries=0)
161+
consumer.flush_interval = 60
162+
wait_started = asyncio.Event()
163+
waiters = []
164+
real_wait = asyncio.wait
165+
166+
async def observe_wait(tasks, **kwargs):
167+
waiters.extend(tasks)
168+
wait_started.set()
169+
return await real_wait(tasks, **kwargs)
170+
171+
with mock.patch("posthog._async_consumer.asyncio.wait", side_effect=observe_wait):
172+
task = asyncio.create_task(
173+
consumer.run() if run_worker else consumer._get_or_flush(60)
174+
)
175+
try:
176+
await wait_started.wait()
177+
task.cancel()
178+
with pytest.raises(asyncio.CancelledError):
179+
await task
180+
181+
assert len(waiters) == 2
182+
assert all(waiter.cancelled() for waiter in waiters)
183+
finally:
184+
task.cancel()
185+
for waiter in waiters:
186+
waiter.cancel()
187+
await asyncio.gather(task, *waiters, return_exceptions=True)
188+
189+
190+
@pytest.mark.asyncio
191+
@pytest.mark.parametrize(
192+
("queued", "flush"), [(True, False), (False, True), (False, False), (True, True)]
193+
)
194+
async def test_get_or_flush_preserves_results_and_cleans_up_waiters(queued, flush):
195+
consumer = make_consumer(retries=0)
196+
event = {"event": "test"}
197+
if queued:
198+
consumer.queue.put_nowait(event)
199+
if flush:
200+
consumer.request_flush()
201+
tasks_before = asyncio.all_tasks()
202+
203+
result = await consumer._get_or_flush(60 if queued or flush else 0)
204+
205+
assert result == (event if queued else None, flush and not queued)
206+
assert consumer._flush_event.is_set() == (queued and flush)
207+
assert not (asyncio.all_tasks() - tasks_before)
208+
if queued:
209+
consumer.queue.task_done()

0 commit comments

Comments
 (0)