Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
82 changes: 68 additions & 14 deletions homeassistant_api/asyncwebsocket.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from __future__ import annotations

import asyncio
import contextlib
import json
import logging
Expand Down Expand Up @@ -54,6 +55,7 @@
class AsyncWebsocketClient(BaseWebsocketClient):
_session: AsyncSession
_ws: AsyncExtensionFromHTTP
_reader_task: asyncio.Task[None]

def __init__(
self,
Expand All @@ -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__()
Expand All @@ -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

Expand All @@ -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":
Expand Down Expand Up @@ -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."""
Expand Down
3 changes: 2 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
Loading