diff --git a/packages/reflex-base/news/6791.bugfix.md b/packages/reflex-base/news/6791.bugfix.md new file mode 100644 index 00000000000..ac1acba2fed --- /dev/null +++ b/packages/reflex-base/news/6791.bugfix.md @@ -0,0 +1 @@ +Drain queued same-token events during graceful event processor shutdown. diff --git a/packages/reflex-base/src/reflex_base/event/processor/event_processor.py b/packages/reflex-base/src/reflex_base/event/processor/event_processor.py index 7d9296fe4dd..8e5460ca6f1 100644 --- a/packages/reflex-base/src/reflex_base/event/processor/event_processor.py +++ b/packages/reflex-base/src/reflex_base/event/processor/event_processor.py @@ -21,7 +21,6 @@ from reflex.istate.manager import StateManager from reflex.utils import console from reflex_base.event.context import EventContext -from reflex_base.event.processor.compat import as_completed from reflex_base.event.processor.future import EventFuture from reflex_base.event.processor.timeout import DrainTimeoutManager from reflex_base.registry import RegisteredEventHandler, RegistrationContext @@ -238,19 +237,26 @@ async def _stop_tasks(self, timeout: float | None = None) -> None: queue to drain before cancelling tasks. If None, the processor will not wait and will cancel tasks immediately. """ - finished_tasks = set() # Graceful drain time, wait for tasks to finish and handle any exceptions. - if timeout is not None and self._tasks: - with contextlib.suppress(asyncio.TimeoutError): - async for task in as_completed(self._tasks.values(), timeout=timeout): + if timeout is not None: + deadline = time.monotonic() + timeout + while self._tasks: + remaining_time = deadline - time.monotonic() + if remaining_time <= 0: + break + finished_tasks, _ = await asyncio.wait( + tuple(self._tasks.values()), + timeout=remaining_time, + return_when=asyncio.FIRST_COMPLETED, + ) + if not finished_tasks: + break + for task in finished_tasks: # Exceptions are handled in _finish_task and ignored here. - with contextlib.suppress(Exception): + with contextlib.suppress(Exception, asyncio.CancelledError): await task - finished_tasks.add(task) # Cancel all outstanding event handler tasks. - outstanding_tasks = [ - task for task in self._tasks.values() if task not in finished_tasks - ] + outstanding_tasks = list(self._tasks.values()) for task in outstanding_tasks: task.cancel() # Wait for all tasks to finish and log any exceptions that were raised. diff --git a/tests/units/reflex_base/event/processor/test_event_processor.py b/tests/units/reflex_base/event/processor/test_event_processor.py index bcc1108be98..c85ab9782e5 100644 --- a/tests/units/reflex_base/event/processor/test_event_processor.py +++ b/tests/units/reflex_base/event/processor/test_event_processor.py @@ -356,6 +356,32 @@ async def test_multiple_futures_cancelled_on_stop(processor: EventProcessor): assert ep._futures == {} +async def test_stop_drains_same_token_sequential_backlog( + processor: EventProcessor, + token: str, +): + """Graceful stop drains queued same-token events within the shutdown budget. + + Args: + processor: The event processor fixture. + token: The client token. + """ + processor.configure() + async with processor as ep: + futures = [ + await ep.enqueue( + token, Event.from_event_type(slow_logging_event(str(i)))[0] + ) + for i in range(10) + ] + + assert [entry["value"] for entry in _CALL_LOG] == [str(i) for i in range(10)] + assert all(future.done() and not future.cancelled() for future in futures) + assert ep._tasks == {} + assert ep._token_queues == {} + assert ep._futures == {} + + async def test_cancel_future_before_task_starts( mock_event_processor: EventProcessor, token: str,