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
« 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."""
6from __future__ import annotations
8import functools
9import itertools
10import json
11import os
12from pathlib import Path
13from typing import Any, Self
14from urllib.parse import urlsplit, urlunsplit
16import niquests
19class HaReplError(Exception):
20 """Connection, auth or command failure (not an error in the user's code)."""
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
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 )
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, "", ""))
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, "", ""))
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
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 )
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
118 async def __aenter__(self) -> Self:
119 return await self.open()
121 async def __aexit__(self, *exc: object) -> None:
122 await self.close()
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
148 async def close(self) -> None:
149 await self._ws.close()
150 await self._session.close()
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)
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"]