|
5 | 5 | from typing import TYPE_CHECKING, Annotated, Self |
6 | 6 |
|
7 | 7 | import websockets.asyncio.client |
| 8 | +import websockets.exceptions |
8 | 9 | from pydantic import Discriminator, TypeAdapter |
9 | 10 | from typing_extensions import Unpack, override |
10 | 11 |
|
@@ -91,49 +92,70 @@ async def __aexit__( |
91 | 92 | exc_value: BaseException | None, |
92 | 93 | exc_traceback: TracebackType | None, |
93 | 94 | ) -> 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. |
97 | 97 | if self._process_platform_messages_task and not self._process_platform_messages_task.done(): |
98 | 98 | self._process_platform_messages_task.cancel() |
99 | 99 | with contextlib.suppress(asyncio.CancelledError): |
100 | 100 | await self._process_platform_messages_task |
101 | 101 |
|
| 102 | + if self._platform_events_websocket: |
| 103 | + await self._platform_events_websocket.close() |
| 104 | + |
102 | 105 | await super().__aexit__(exc_type, exc_value, exc_traceback) |
103 | 106 |
|
104 | 107 | 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 | + |
105 | 116 | 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): |
107 | 118 | 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), |
122 | 143 | ) |
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 | + ) |
137 | 159 | except Exception: |
138 | 160 | logger.exception('Error in websocket connection') |
139 | 161 | if self._connected_to_platform_websocket is not None and not self._connected_to_platform_websocket.done(): |
|
0 commit comments