Coverage for src/homeassistant_repl/client.py: 50%

110 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-10-05 13:50 +0000

1"""Minimal Home Assistant websocket API client - built on niquests' websocket 

2extension (a GET against a ws(s):// URL upgrades in place and hands back the 

3extension on the response), the same HTTP stack `rest.py`'s REST client uses, 

4rather than pulling in the separate `websockets` library for this one job.""" 

5 

6from __future__ import annotations 

7 

8import functools 

9import itertools 

10import json 

11import os 

12from pathlib import Path 

13from typing import Any, Self 

14from urllib.parse import urlsplit, urlunsplit 

15 

16import niquests 

17 

18 

19class HaReplError(Exception): 

20 """Connection, auth or command failure (not an error in the user's code).""" 

21 

22 

23@functools.cache 

24def _dotenv() -> dict[str, str]: 

25 """Best-effort KEY=VALUE pairs from a `.env` file in the current 

26 directory, if one exists - checked only once real environment variables 

27 come up empty, see resolve_url()/resolve_token(). Deliberately minimal 

28 rather than pulling in python-dotenv: comments, blank lines, an optional 

29 "export " prefix and quoted values are all a token/url file needs. 

30 Cached since it's read-only for the life of the process and both 

31 resolve_url() and resolve_token() may call it. 

32 """ 

33 path = Path(".env") 

34 if not path.is_file(): 

35 return {} 

36 values: dict[str, str] = {} 

37 for raw_line in path.read_text(encoding="utf-8").splitlines(): 

38 line = raw_line.strip().removeprefix("export ") 

39 if not line or line.startswith("#"): 

40 continue 

41 key, sep, value = line.partition("=") 

42 if not sep: 42 ↛ 43line 42 didn't jump to line 43 because the condition on line 42 was never true

43 continue 

44 values[key.strip()] = value.strip().strip("'\"") 

45 return values 

46 

47 

48def _default_url() -> str: 

49 return ( 

50 os.environ.get("HASS_SERVER") 

51 or _dotenv().get("HASS_SERVER") 

52 or "http://homeassistant.local:8123" 

53 ) 

54 

55 

56def resolve_url(url: str | None) -> str: 

57 """Accept http(s)://host:8123 or ws(s)://... and return the websocket endpoint.""" 

58 if not url: 58 ↛ 63line 58 didn't jump to line 63 because the condition on line 58 was always true

59 # Inside a Home Assistant add-on (e.g. Studio Code Server) talk via the supervisor. 

60 if os.environ.get("SUPERVISOR_TOKEN"): 60 ↛ 61line 60 didn't jump to line 61 because the condition on line 60 was never true

61 return "ws://supervisor/core/websocket" 

62 url = _default_url() 

63 parts = urlsplit(url) 

64 scheme = {"http": "ws", "https": "wss"}.get(parts.scheme, parts.scheme) 

65 path = parts.path.rstrip("/") 

66 if not path.endswith("/websocket"): 66 ↛ 68line 66 didn't jump to line 68 because the condition on line 66 was always true

67 path += "/api/websocket" 

68 return urlunsplit((scheme, parts.netloc, path, "", "")) 

69 

70 

71def resolve_rest_url(url: str | None) -> str: 

72 """Accept http(s)://host:8123, ws(s)://... or an existing Client.url, and 

73 return the REST API base (scheme normalised to http(s), path ending in 

74 `/api`) that homeassistant_api.AsyncClient expects.""" 

75 if not url: 75 ↛ 76line 75 didn't jump to line 76 because the condition on line 75 was never true

76 if os.environ.get("SUPERVISOR_TOKEN"): 

77 return "http://supervisor/core/api" 

78 url = _default_url() 

79 parts = urlsplit(url) 

80 scheme = {"ws": "http", "wss": "https"}.get(parts.scheme, parts.scheme) 

81 path = parts.path.rstrip("/").removesuffix("/websocket") 

82 if not path.endswith("/api"): 82 ↛ 84line 82 didn't jump to line 84 because the condition on line 82 was always true

83 path += "/api" 

84 return urlunsplit((scheme, parts.netloc, path, "", "")) 

85 

86 

87def resolve_token(token: str | None) -> str: 

88 token = ( 

89 token 

90 or os.environ.get("HASS_TOKEN") 

91 or os.environ.get("SUPERVISOR_TOKEN") 

92 or _dotenv().get("HASS_TOKEN") 

93 ) 

94 if not token: 

95 raise HaReplError( 

96 "No access token: set HASS_TOKEN or pass --token " 

97 "(create one under Profile > Security > Long-lived access tokens)" 

98 ) 

99 return token 

100 

101 

102class Client: 

103 def __init__(self, url: str, token: str) -> None: 

104 self.url = url 

105 self._token = token 

106 self._ids = itertools.count(1) 

107 self._session: Any = None # set by open() - AsyncSession 

108 self._ws: Any = ( 

109 None # set by open() - the extension: send_payload()/next_payload()/close() 

110 ) 

111 

112 @property 

113 def token(self) -> str: 

114 """The resolved access token - also what a REST call (see rest.py's 

115 `hass_api()`) against the same instance should authenticate with.""" 

116 return self._token 

117 

118 async def __aenter__(self) -> Self: 

119 return await self.open() 

120 

121 async def __aexit__(self, *exc: object) -> None: 

122 await self.close() 

123 

124 async def open(self) -> Self: 

125 """Connect and authenticate - what `async with Client(...)` does, 

126 exposed directly for callers that want to keep the connection open 

127 past the enclosing scope (e.g. `connect()` in __init__.py).""" 

128 self._session = niquests.AsyncSession() 

129 try: 

130 resp = await self._session.get(self.url) 

131 resp.raise_for_status() 

132 except niquests.exceptions.RequestException as err: 

133 raise HaReplError(f"Cannot connect to {self.url}: {err}") from err 

134 if resp.extension is None: 

135 raise HaReplError(f"Server did not upgrade to WebSocket: {self.url}") 

136 self._ws = resp.extension 

137 msg = await self._recv() 

138 if msg.get("type") != "auth_required": 

139 raise HaReplError(f"Unexpected greeting: {msg}") 

140 await self._ws.send_payload( 

141 json.dumps({"type": "auth", "access_token": self._token}) 

142 ) 

143 msg = await self._recv() 

144 if msg.get("type") != "auth_ok": 

145 raise HaReplError(f"Authentication failed: {msg.get('message', msg)}") 

146 return self 

147 

148 async def close(self) -> None: 

149 await self._ws.close() 

150 await self._session.close() 

151 

152 async def _recv(self) -> dict[str, Any]: 

153 try: 

154 payload = await self._ws.next_payload() 

155 except niquests.exceptions.RequestException as err: 

156 raise HaReplError(f"Connection closed: {err}") from err 

157 if payload is None: 

158 raise HaReplError("Connection closed") 

159 return json.loads(payload) 

160 

161 async def call(self, type_: str, **payload: Any) -> Any: 

162 msg_id = next(self._ids) 

163 await self._ws.send_payload( 

164 json.dumps({"id": msg_id, "type": type_, **payload}) 

165 ) 

166 while True: 

167 msg = await self._recv() 

168 if msg.get("id") != msg_id or msg.get("type") != "result": 

169 continue 

170 if not msg["success"]: 

171 error = msg.get("error", {}) 

172 if error.get("code") == "unknown_command": 

173 raise HaReplError( 

174 f"{type_} not available: is the ha_repl_server integration " 

175 "installed and `ha_repl_server:` in configuration.yaml?" 

176 ) 

177 raise HaReplError(f"{type_} failed: {error.get('message', error)}") 

178 return msg["result"]