Skip to content

Commit ae02fb8

Browse files
committed
fix: close the platform events websocket iterator on shutdown
1 parent 90d1725 commit ae02fb8

2 files changed

Lines changed: 64 additions & 28 deletions

File tree

src/apify/events/_apify_event_manager.py

Lines changed: 34 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
import contextlib
55
import time
66
from logging import getLogger
7-
from typing import TYPE_CHECKING, Annotated, Self
7+
from typing import TYPE_CHECKING, Annotated, Self, cast
88

99
import websockets.asyncio.client
1010
import websockets.client
@@ -20,7 +20,7 @@
2020
from apify.events._types import DeprecatedEvent, EventMessage, SystemInfoEventData, UnknownEvent
2121

2222
if TYPE_CHECKING:
23-
from collections.abc import Generator
23+
from collections.abc import AsyncGenerator, Generator
2424
from types import TracebackType
2525

2626
from crawlee.events._event_manager import EventManagerOptions
@@ -149,31 +149,38 @@ async def _process_platform_messages(self, ws_url: str) -> None:
149149

150150
try:
151151
# Used as an async iterator, `connect` reconnects with exponential backoff on failed connection attempts.
152-
async for websocket in websockets.asyncio.client.connect(
153-
ws_url, process_exception=self._process_connection_exception
154-
):
155-
self._platform_events_websocket = websocket
156-
if self._connected_to_platform_websocket and not self._connected_to_platform_websocket.done():
157-
self._connected_to_platform_websocket.set_result(True)
158-
else:
159-
logger.info('Reconnected to the platform events websocket.')
160-
161-
connection_opened_at = time.monotonic()
162-
connection_lost = await self._consume_messages(websocket)
163-
164-
if not self._should_reconnect_after_close(websocket, connection_lost=connection_lost):
165-
break
166-
167-
# Reconnect a healthy connection immediately; back off only on repeated rapid drops. The first
168-
# rapid drop reconnects once without delay (it only primes the backoff generator), and each
169-
# subsequent consecutive rapid drop then sleeps for the next backoff interval. A healthy
170-
# connection resets the generator, so the next rapid drop again gets that one free retry.
171-
if time.monotonic() - connection_opened_at >= self._HEALTHY_CONNECTION_MIN_DURATION:
172-
backoff_delays = None
173-
elif backoff_delays is None:
174-
backoff_delays = websockets.client.backoff()
175-
else:
176-
await asyncio.sleep(next(backoff_delays))
152+
connector = websockets.asyncio.client.connect(ws_url, process_exception=self._process_connection_exception)
153+
154+
# That iterator is an async generator whose cleanup closes the current connection, and neither `break` nor
155+
# cancelling this task closes it. Left suspended, it is closed by the garbage collector instead - on Actor
156+
# exit that lands in `asyncio.run` teardown, too late to be awaited, which then breaks the event loop's
157+
# async generator shutdown. `websockets` types the iterator as a plain `AsyncIterator`, hence the cast.
158+
connections = cast('AsyncGenerator[websockets.asyncio.client.ClientConnection]', aiter(connector))
159+
160+
async with contextlib.aclosing(connections):
161+
async for websocket in connections:
162+
self._platform_events_websocket = websocket
163+
if self._connected_to_platform_websocket and not self._connected_to_platform_websocket.done():
164+
self._connected_to_platform_websocket.set_result(True)
165+
else:
166+
logger.info('Reconnected to the platform events websocket.')
167+
168+
connection_opened_at = time.monotonic()
169+
connection_lost = await self._consume_messages(websocket)
170+
171+
if not self._should_reconnect_after_close(websocket, connection_lost=connection_lost):
172+
break
173+
174+
# Reconnect a healthy connection immediately; back off only on repeated rapid drops. The first
175+
# rapid drop reconnects once without delay (it only primes the backoff generator), and each
176+
# subsequent consecutive rapid drop then sleeps for the next backoff interval. A healthy
177+
# connection resets the generator, so the next rapid drop again gets that one free retry.
178+
if time.monotonic() - connection_opened_at >= self._HEALTHY_CONNECTION_MIN_DURATION:
179+
backoff_delays = None
180+
elif backoff_delays is None:
181+
backoff_delays = websockets.client.backoff()
182+
else:
183+
await asyncio.sleep(next(backoff_delays))
177184
except Exception:
178185
logger.exception('Error in websocket connection')
179186
if self._connected_to_platform_websocket is not None and not self._connected_to_platform_websocket.done():

tests/unit/events/test_apify_event_manager.py

Lines changed: 30 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313

1414
import pytest
1515
import websockets
16+
import websockets.asyncio.client
1617
import websockets.asyncio.server
1718

1819
from crawlee.events._types import Event
@@ -24,7 +25,7 @@
2425
from apify.events._types import SystemInfoEventData
2526

2627
if TYPE_CHECKING:
27-
from collections.abc import AsyncGenerator, Awaitable, Callable
28+
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable
2829

2930

3031
DUMMY_SYSTEM_INFO = {
@@ -595,6 +596,34 @@ async def test_shutdown_during_reconnect_backoff_is_clean(monkeypatch: pytest.Mo
595596
assert persist_state_task is None or persist_state_task.done()
596597

597598

599+
async def test_exit_closes_the_reconnecting_iterator(monkeypatch: pytest.MonkeyPatch) -> None:
600+
"""Test that exiting closes the `connect` async iterator itself, rather than leaving it to the garbage collector."""
601+
# Keeping a reference to every iterator prevents the garbage collector from finalizing it, so the assertion below
602+
# holds only if the event manager closes the iterator on its own.
603+
iterators: list[Any] = []
604+
original_aiter = websockets.asyncio.client.connect.__aiter__
605+
606+
def recording_aiter(
607+
connector: websockets.asyncio.client.connect,
608+
) -> AsyncIterator[websockets.asyncio.client.ClientConnection]:
609+
iterator = original_aiter(connector)
610+
iterators.append(iterator)
611+
return iterator
612+
613+
monkeypatch.setattr(websockets.asyncio.client.connect, '__aiter__', recording_aiter)
614+
615+
async with (
616+
_platform_ws_server(monkeypatch) as (_, client_connected),
617+
ApifyEventManager(Configuration.get_global_configuration()),
618+
):
619+
await asyncio.wait_for(client_connected.wait(), timeout=10)
620+
621+
# An iterator left suspended keeps its websocket open, and closing it is then deferred to the garbage collector.
622+
# On Actor exit that lands in `asyncio.run` teardown, too late to be awaited, breaking async generator shutdown.
623+
assert iterators
624+
assert all(iterator.ag_frame is None for iterator in iterators)
625+
626+
598627
async def test_malformed_message_logs_exception(
599628
monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
600629
) -> None:

0 commit comments

Comments
 (0)