@@ -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