Skip to content

Commit 474ad67

Browse files
committed
fix: reconnect to platform events websocket after connection drop
1 parent 2cc5602 commit 474ad67

2 files changed

Lines changed: 70 additions & 55 deletions

File tree

src/apify/events/_apify_event_manager.py

Lines changed: 54 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
from typing import TYPE_CHECKING, Annotated, Self
66

77
import websockets.asyncio.client
8+
import websockets.exceptions
89
from pydantic import Discriminator, TypeAdapter
910
from typing_extensions import Unpack, override
1011

@@ -91,49 +92,70 @@ async def __aexit__(
9192
exc_value: BaseException | None,
9293
exc_traceback: TracebackType | None,
9394
) -> None:
94-
if self._platform_events_websocket:
95-
await self._platform_events_websocket.close()
96-
95+
# Cancel the message-processing task first so that closing the websocket below is not treated
96+
# as a dropped connection and followed by a reconnect attempt.
9797
if self._process_platform_messages_task and not self._process_platform_messages_task.done():
9898
self._process_platform_messages_task.cancel()
9999
with contextlib.suppress(asyncio.CancelledError):
100100
await self._process_platform_messages_task
101101

102+
if self._platform_events_websocket:
103+
await self._platform_events_websocket.close()
104+
102105
await super().__aexit__(exc_type, exc_value, exc_traceback)
103106

104107
async def _process_platform_messages(self, ws_url: str) -> None:
108+
def process_exception(exc: Exception) -> Exception | None:
109+
# Until the first connection succeeds, treat every error as fatal so that `__aenter__` fails fast.
110+
# Afterwards, treat every error as transient — the reconnect iterator keeps retrying with backoff
111+
# so that platform events (e.g. `MIGRATING`) are not missed for the rest of the run.
112+
if self._connected_to_platform_websocket is None or not self._connected_to_platform_websocket.done():
113+
return exc
114+
return None
115+
105116
try:
106-
async with websockets.asyncio.client.connect(ws_url) as websocket:
117+
async for websocket in websockets.asyncio.client.connect(ws_url, process_exception=process_exception):
107118
self._platform_events_websocket = websocket
108-
if self._connected_to_platform_websocket is not None:
109-
self._connected_to_platform_websocket.set_result(True)
110-
111-
async for message in websocket:
112-
try:
113-
parsed_message = event_data_adapter.validate_json(message)
114-
115-
if isinstance(parsed_message, DeprecatedEvent):
116-
continue
117-
118-
if isinstance(parsed_message, UnknownEvent):
119-
logger.info(
120-
f'Unknown message received: event_name={parsed_message.name}, '
121-
f'event_data={parsed_message.data}'
119+
connected_future = self._connected_to_platform_websocket
120+
if connected_future is not None and not connected_future.done():
121+
connected_future.set_result(True)
122+
123+
try:
124+
async for message in websocket:
125+
try:
126+
parsed_message = event_data_adapter.validate_json(message)
127+
128+
if isinstance(parsed_message, DeprecatedEvent):
129+
continue
130+
131+
if isinstance(parsed_message, UnknownEvent):
132+
logger.info(
133+
f'Unknown message received: event_name={parsed_message.name}, '
134+
f'event_data={parsed_message.data}'
135+
)
136+
continue
137+
138+
self.emit(
139+
event=parsed_message.name,
140+
event_data=parsed_message.data
141+
if not isinstance(parsed_message.data, SystemInfoEventData)
142+
else parsed_message.data.to_crawlee_format(self._configuration.dedicated_cpus or 1),
122143
)
123-
continue
124-
125-
self.emit(
126-
event=parsed_message.name,
127-
event_data=parsed_message.data
128-
if not isinstance(parsed_message.data, SystemInfoEventData)
129-
else parsed_message.data.to_crawlee_format(self._configuration.dedicated_cpus or 1),
130-
)
131-
132-
if parsed_message.name == Event.MIGRATING:
133-
await self._emit_persist_state_event_rec_task.stop()
134-
self.emit(event=Event.PERSIST_STATE, event_data=EventPersistStateData(is_migrating=True))
135-
except Exception:
136-
logger.exception('Cannot parse Actor event', extra={'raw_message': message})
144+
145+
if parsed_message.name == Event.MIGRATING:
146+
await self._emit_persist_state_event_rec_task.stop()
147+
self.emit(
148+
event=Event.PERSIST_STATE, event_data=EventPersistStateData(is_migrating=True)
149+
)
150+
except Exception:
151+
logger.exception('Cannot parse Actor event', extra={'raw_message': message})
152+
except websockets.exceptions.ConnectionClosed:
153+
pass
154+
155+
logger.warning(
156+
f'Connection to platform events websocket was closed '
157+
f'(code={websocket.close_code}, reason={websocket.close_reason!r}), reconnecting...'
158+
)
137159
except Exception:
138160
logger.exception('Error in websocket connection')
139161
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: 16 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -12,7 +12,6 @@
1212
import pytest
1313
import websockets
1414
import websockets.asyncio.server
15-
import websockets.exceptions
1615

1716
from crawlee.events._types import Event
1817

@@ -320,38 +319,32 @@ def migrating_listener(data: Any) -> None:
320319
assert len(migration_persist_events) >= 1
321320

322321

323-
async def test_websocket_mid_stream_disconnect_does_not_raise_invalid_state_error(
324-
monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
325-
) -> None:
326-
"""Regression: a mid-stream websocket disconnect after a successful connect must not raise InvalidStateError.
327-
328-
The `_connected_to_platform_websocket` future is resolved to `True` on successful connect. If the websocket
329-
later drops, the outer `except` in `_process_platform_messages` must not call `set_result(False)` on the
330-
already-resolved future.
331-
"""
322+
async def test_websocket_reconnects_after_connection_drop(monkeypatch: pytest.MonkeyPatch) -> None:
323+
"""Test that after a mid-stream websocket drop, the manager reconnects and keeps receiving platform events."""
332324
async with (
333325
_platform_ws_server(monkeypatch) as (connected_ws_clients, client_connected),
334326
ApifyEventManager(Configuration.get_global_configuration()) as event_manager,
335327
):
336328
await client_connected.wait()
329+
aborting_calls: list[Any] = []
337330

338-
# Force an abnormal close from the server so the client's `async for` raises ConnectionClosedError.
331+
def listener(data: Any) -> None:
332+
aborting_calls.append(data)
333+
334+
event_manager.on(event=Event.ABORTING, listener=listener)
335+
336+
# Drop the connection abnormally from the server side.
337+
client_connected.clear()
339338
for ws in list(connected_ws_clients):
340339
await ws.close(code=1011, reason='Simulated server error')
341340

342-
task = event_manager._process_platform_messages_task
343-
assert task is not None
344-
await asyncio.wait_for(asyncio.shield(task), timeout=2.0)
345-
346-
exc = task.exception()
347-
assert not isinstance(exc, asyncio.InvalidStateError), f'Task raised InvalidStateError: {exc}'
341+
# The event manager should reconnect on its own.
342+
await asyncio.wait_for(client_connected.wait(), timeout=5.0)
348343

349-
# Confirm the test actually exercised the disconnect path — the outer `except` in
350-
# `_process_platform_messages` should have logged a `ConnectionClosedError`.
351-
logged_exc_types = [
352-
record.exc_info[0] for record in caplog.records if record.exc_info and record.exc_info[0] is not None
353-
]
354-
assert any(issubclass(exc_type, websockets.exceptions.ConnectionClosedError) for exc_type in logged_exc_types)
344+
# Events sent over the new connection must still be received.
345+
websockets.broadcast(connected_ws_clients, json.dumps({'name': 'aborting'}))
346+
await poll_until_condition(lambda: bool(aborting_calls), poll_interval=0.05)
347+
assert len(aborting_calls) == 1
355348

356349

357350
async def test_malformed_message_logs_exception(

0 commit comments

Comments
 (0)