Skip to content

Commit 79b595e

Browse files
committed
Add base Websocket class
1 parent e3ea524 commit 79b595e

3 files changed

Lines changed: 86 additions & 59 deletions

File tree

Lines changed: 82 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,82 @@
1+
import logging
2+
import time
3+
from typing import Optional, cast
4+
5+
from pydantic import ValidationError
6+
7+
from homeassistant_api.errors import (
8+
ReceivingError,
9+
RequestError,
10+
)
11+
from homeassistant_api.models.websocket import (
12+
ErrorResponse,
13+
EventResponse,
14+
PingResponse,
15+
ResultResponse,
16+
)
17+
from homeassistant_api.utils import JSONType
18+
19+
logger = logging.getLogger(__name__)
20+
21+
22+
class RawBaseWebsocketClient:
23+
"""Shared methods for Websocket clients."""
24+
25+
api_url: str
26+
token: str
27+
_id_counter: int
28+
_result_responses: dict[int, Optional[ResultResponse]]
29+
_event_responses: dict[int, list[EventResponse]]
30+
_ping_responses: dict[int, PingResponse]
31+
32+
def __init__(self, api_url: str, token: str) -> None:
33+
self.api_url = api_url
34+
self.token = token.strip()
35+
36+
self._id_counter = 0
37+
self._result_responses = {} # id -> response
38+
self._event_responses = {} # id -> [response, ...]
39+
self._ping_responses = {} # id -> (sent, received)
40+
41+
def __repr__(self) -> str:
42+
return f"{self.__class__.__name__}({self.api_url!r})"
43+
44+
def _request_id(self) -> int:
45+
"""Get a unique id for a message."""
46+
self._id_counter += 1
47+
return self._id_counter
48+
49+
def check_success(self, data: dict[str, JSONType]) -> None:
50+
"""Check if a command message was successful."""
51+
try:
52+
error_resp = ErrorResponse.model_validate(data)
53+
raise RequestError(error_resp.error.code, error_resp.error.message)
54+
except ValidationError:
55+
pass
56+
57+
def handle_recv(self, data: dict[str, JSONType]) -> None:
58+
"""Handle a received message."""
59+
if "id" not in data:
60+
raise ReceivingError(
61+
"Received a message without an id outside the auth phase."
62+
)
63+
self.check_success(data)
64+
self.parse_response(data)
65+
66+
def parse_response(self, data: dict[str, JSONType]) -> None:
67+
data_id = cast(int, data["id"])
68+
if data.get("type") == "pong":
69+
logger.info("Received pong message")
70+
self._ping_responses[data_id].end = time.perf_counter_ns()
71+
elif data.get("type") == "result":
72+
logger.info("Received result message")
73+
if data.get("success"):
74+
self._result_responses[data_id] = ResultResponse.model_validate(data)
75+
else:
76+
error_resp = ErrorResponse.model_validate(data)
77+
raise RequestError(error_resp.error.code, error_resp.error.message)
78+
elif data.get("type") == "event":
79+
logger.info("Received event message %s", data["event"])
80+
self._event_responses[data_id].append(EventResponse.model_validate(data))
81+
else:
82+
raise ReceivingError(f"Received unexpected message type: {data}")

homeassistant_api/rawwebsocket.py

Lines changed: 3 additions & 57 deletions
Original file line numberDiff line numberDiff line change
@@ -8,25 +8,24 @@
88

99
from homeassistant_api.errors import (
1010
ReceivingError,
11-
RequestError,
1211
ResponseError,
1312
UnauthorizedError,
1413
)
1514
from homeassistant_api.models.websocket import (
1615
AuthInvalid,
1716
AuthOk,
1817
AuthRequired,
19-
ErrorResponse,
2018
EventResponse,
2119
PingResponse,
2220
ResultResponse,
2321
)
22+
from homeassistant_api.rawbasewebsocket import RawBaseWebsocketClient
2423
from homeassistant_api.utils import JSONType
2524

2625
logger = logging.getLogger(__name__)
2726

2827

29-
class RawWebsocketClient:
28+
class RawWebsocketClient(RawBaseWebsocketClient):
3029
api_url: str
3130
token: str
3231
_conn: Optional[ws.ClientConnection]
@@ -36,22 +35,9 @@ def __init__(
3635
api_url: str,
3736
token: str,
3837
) -> None:
39-
self.api_url = api_url
40-
self.token = token.strip()
38+
super().__init__(api_url, token)
4139
self._conn = None
4240

43-
self._id_counter = 0
44-
self._result_responses: dict[int, Optional[ResultResponse]] = (
45-
{}
46-
) # id -> response
47-
self._event_responses: dict[int, list[EventResponse]] = (
48-
{}
49-
) # id -> [response, ...]
50-
self._ping_responses: dict[int, PingResponse] = {} # id -> (sent, received)
51-
52-
def __repr__(self) -> str:
53-
return f"{self.__class__.__name__}({self.api_url!r})"
54-
5541
def __enter__(self):
5642
self._conn = ws.connect(self.api_url)
5743
self._conn.__enter__()
@@ -66,11 +52,6 @@ def __exit__(self, exc_type, exc_value, traceback):
6652
self._conn.__exit__(exc_type, exc_value, traceback)
6753
self._conn = None
6854

69-
def _request_id(self) -> int:
70-
"""Get a unique id for a message."""
71-
self._id_counter += 1
72-
return self._id_counter
73-
7455
def _send(self, data: dict[str, JSONType]) -> None:
7556
"""Send a message to the websocket server."""
7657
logger.debug(f"Sending message: {data}")
@@ -112,41 +93,6 @@ def send(self, type: str, include_id: bool = True, **data: Any) -> int:
11293
return data["id"]
11394
return -1 # non-command messages don't have an id
11495

115-
def check_success(self, data: dict[str, JSONType]) -> None:
116-
"""Check if a command message was successful."""
117-
try:
118-
error_resp = ErrorResponse.model_validate(data)
119-
raise RequestError(error_resp.error.code, error_resp.error.message)
120-
except ValidationError:
121-
pass
122-
123-
def handle_recv(self, data: dict[str, JSONType]) -> None:
124-
"""Handle a received message."""
125-
if "id" not in data:
126-
raise ReceivingError(
127-
"Received a message without an id outside the auth phase."
128-
)
129-
self.check_success(data)
130-
self.parse_response(data)
131-
132-
def parse_response(self, data: dict[str, JSONType]) -> None:
133-
data_id = cast(int, data["id"])
134-
if data.get("type") == "pong":
135-
logger.info("Received pong message")
136-
self._ping_responses[data_id].end = time.perf_counter_ns()
137-
elif data.get("type") == "result":
138-
logger.info("Received result message")
139-
if data.get("success"):
140-
self._result_responses[data_id] = ResultResponse.model_validate(data)
141-
else:
142-
error_resp = ErrorResponse.model_validate(data)
143-
raise RequestError(error_resp.error.code, error_resp.error.message)
144-
elif data.get("type") == "event":
145-
logger.info("Received event message %s", data["event"])
146-
self._event_responses[data_id].append(EventResponse.model_validate(data))
147-
else:
148-
raise ReceivingError(f"Received unexpected message type: {data}")
149-
15096
def recv(self, id: int) -> Union[EventResponse, ResultResponse, PingResponse]:
15197
"""Receive a response to a message from the websocket server."""
15298
while True:

homeassistant_api/websocket.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -12,10 +12,9 @@
1212
ResultResponse,
1313
TemplateEvent,
1414
)
15+
from homeassistant_api.rawwebsocket import RawWebsocketClient
1516
from homeassistant_api.utils import JSONType, prepare_entity_id
1617

17-
from .rawwebsocket import RawWebsocketClient
18-
1918
logger = logging.getLogger(__name__)
2019

2120

0 commit comments

Comments
 (0)