From 466816007988be58e17979126bc327f1a566f152 Mon Sep 17 00:00:00 2001 From: Simon Arlott Date: Fri, 4 Sep 2026 23:01:41 +0100 Subject: [PATCH] Fix handling of multiple async readers If there are multiple async tasks trying to read from the WebSocket, messages spread across multiple chunks won't be handled correctly because they each independently try to combine chunks. Only one of the multiple tasks can be successful, leaving the others starved of responses because they'll all try to read one message and won't react if a new message has already been received for them by another task. Use a lock to ensure that only one of the tasks can be reading at any one time, and wake up the other tasks whenever there's a new message. Closed WebSockets are handled by having each task becoming the active reader in turn to discover that the connection has been closed. --- homeassistant_api/asyncwebsocket.py | 44 ++++++++++++++++++++--------- 1 file changed, 31 insertions(+), 13 deletions(-) diff --git a/homeassistant_api/asyncwebsocket.py b/homeassistant_api/asyncwebsocket.py index d82f0fe..c28e1da 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 @@ -74,6 +75,8 @@ async def __aenter__(self) -> Self: msg = "Server did not upgrade to WebSocket" raise ReceivingError(msg) self._ws = resp.extension + self._read_lock = asyncio.Condition() + self._read_active = False okay = await self.authentication_phase() logger.info("Authenticated with Home Assistant (%s)", okay.ha_version) await self.supported_features_phase() @@ -148,19 +151,34 @@ async def recv( ) -> 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()) + await self._read_lock.acquire() + try: + ## 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, ensuring that + ## only one task is actively reading at any time + if not self._read_active: + self._read_active = True + self._read_lock.release() + try: + self.handle_recv(await self._async_recv()) + finally: + await self._read_lock.acquire() + self._read_active = False + self._read_lock.notify_all() + else: + await self._read_lock.wait() + finally: + self._read_lock.release() async def recv_result(self, msg_id: int) -> ResultResponse: """Receive a ResultResponse, raising TypeError if the response is not a ResultResponse."""