diff --git a/homeassistant_api/asyncwebsocket.py b/homeassistant_api/asyncwebsocket.py index d82f0fe..e0f6caa 100644 --- a/homeassistant_api/asyncwebsocket.py +++ b/homeassistant_api/asyncwebsocket.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio import contextlib import json import logging @@ -54,6 +55,7 @@ class AsyncWebsocketClient(BaseWebsocketClient): _session: AsyncSession _ws: AsyncExtensionFromHTTP + _reader_task: asyncio.Task[None] def __init__( self, @@ -65,6 +67,9 @@ def __init__( ) -> None: super().__init__(api_url, token, max_size=max_size) self._session = session if session is not None else AsyncSession() + self._new_message = asyncio.Condition() + self._fatal_error: BaseException | None = None + self._id_errors: dict[int, BaseException] = {} async def __aenter__(self) -> Self: await self._session.__aenter__() @@ -74,8 +79,11 @@ async def __aenter__(self) -> Self: msg = "Server did not upgrade to WebSocket" raise ReceivingError(msg) self._ws = resp.extension + # Authenticate before starting the reader task: these messages have no "id" + # and aren't dispatched through it, so nothing else may read the socket yet. okay = await self.authentication_phase() logger.info("Authenticated with Home Assistant (%s)", okay.ha_version) + self._reader_task = asyncio.create_task(self._reader_loop()) await self.supported_features_phase() return self @@ -85,10 +93,50 @@ async def __aexit__( exc_value: BaseException | None, traceback: TracebackType | None, ) -> None: + self._reader_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await self._reader_task await self._ws.close() del self._ws await self._session.__aexit__(exc_type, exc_value, traceback) + async def _reader_loop(self) -> None: + """ + Continuously read messages off the socket and dispatch them by id. + + This is the only task allowed to touch the socket directly, so reads can + never interleave or starve a waiting `recv()` call. + """ + while True: + try: + data = await self._async_recv() + except Exception as e: # noqa: BLE001 + # No message id is available yet, so there's no single owner to + # blame this on: the connection is dead for every pending recv(). + async with self._new_message: + self._fatal_error = e + self._new_message.notify_all() + return + await self._dispatch(data) + + async def _dispatch(self, data: dict[str, Any]) -> None: + """Hand a received message to `handle_recv`, routing any failure to its owner.""" + try: + self.handle_recv(data) + except Exception as e: # noqa: BLE001 + data_id = data.get("id") + async with self._new_message: + if isinstance(data_id, int): + # The message carried its own id, so only the caller waiting + # on that id needs to see the failure. + self._id_errors[data_id] = e + else: + self._fatal_error = e + self._new_message.notify_all() + return + async with self._new_message: + self._new_message.notify_all() + async def _async_send(self, data: dict[str, Any]) -> None: """Send a message to the websocket server.""" if data.get("type") != "auth": @@ -147,20 +195,26 @@ async def recv( msg_id: int, ) -> EventResponse | ResultResponse | PingResponse | None: """Receive a response to a message from the websocket server.""" - while True: - ## have we received a message with the id we're looking for? - if self._result_responses.get(msg_id) is not None: - return self._result_responses.pop(msg_id) - if self._event_responses.get(msg_id, []): - return self._event_responses[msg_id].pop(0) - if ( - self._ping_responses.get(msg_id) is not None - and self._ping_responses[msg_id].end is not None - ): - return self._ping_responses.pop(msg_id) - - ## if not, keep receiving messages until we do - self.handle_recv(await self._async_recv()) + async with self._new_message: + while True: + ## have we received a message with the id we're looking for? + if self._result_responses.get(msg_id) is not None: + return self._result_responses.pop(msg_id) + if self._event_responses.get(msg_id, []): + return self._event_responses[msg_id].pop(0) + if ( + self._ping_responses.get(msg_id) is not None + and self._ping_responses[msg_id].end is not None + ): + return self._ping_responses.pop(msg_id) + ## did our own request fail, or did the connection die outright? + if msg_id in self._id_errors: + raise self._id_errors.pop(msg_id) + if self._fatal_error is not None: + raise self._fatal_error + + ## if not, wait for the reader task to dispatch another message + await self._new_message.wait() async def recv_result(self, msg_id: int) -> ResultResponse: """Receive a ResultResponse, raising TypeError if the response is not a ResultResponse.""" diff --git a/pyproject.toml b/pyproject.toml index 6ba9028..706ffb5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -43,12 +43,13 @@ dev = [ "aiosqlite>=0.22", "prek>=0.3.8", "pre-commit>=4.5.1", - "nimax", + "nimax>=1.1.1", ] [tool.nimax] cassette_library_dir = "tests/cassettes" match_on = ["method", "uri"] +ws_id_extractor = "id" [tool.pytest.ini_options] asyncio_mode = "auto"