mirror of
https://github.com/alexhopeoconnor/arduino-home-assistant.git
synced 2026-10-04 02:48:13 +10:00
test: extract reusable HA MQTT contract testkit
This commit is contained in:
@@ -0,0 +1,18 @@
|
||||
# HA/MQTT contract testkit
|
||||
|
||||
This small Python package supplies generic test transport primitives for a
|
||||
disposable Home Assistant + MQTT contract environment:
|
||||
|
||||
- Home Assistant onboarding, MQTT config-entry setup, WebSocket registry access,
|
||||
REST state/service helpers, and readiness waiting;
|
||||
- retained MQTT publishing and discovery observation; and
|
||||
- JSON/JSONL artifact output plus bounded retry helpers.
|
||||
|
||||
It deliberately contains no ArduinoHA discovery fixture, DeviceFramework
|
||||
import, entity expectation, or release policy. ArduinoHA owns its migration
|
||||
fixture in `tests/ha-contract`; DeviceFramework owns its hardware fixture and
|
||||
Docker adapter in its own repository. Other projects can reuse this package by
|
||||
providing their own retained MQTT messages and assertions.
|
||||
|
||||
The package is copied into a test container as a local build context and is not
|
||||
part of the Arduino firmware library export.
|
||||
@@ -0,0 +1,17 @@
|
||||
[build-system]
|
||||
requires = ["setuptools>=68"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "ha-mqtt-contract-testkit"
|
||||
version = "0.1.0"
|
||||
description = "Generic Home Assistant MQTT contract-test support"
|
||||
requires-python = ">=3.11"
|
||||
dependencies = [
|
||||
"paho-mqtt==2.1.0",
|
||||
"requests==2.32.3",
|
||||
"websocket-client==1.8.0",
|
||||
]
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
where = ["src"]
|
||||
@@ -0,0 +1,21 @@
|
||||
"""Generic Home Assistant MQTT contract-test support.
|
||||
|
||||
This package deliberately knows nothing about ArduinoHA, DeviceFramework, or
|
||||
any particular discovery schema. Consumers own their fixtures and assertions;
|
||||
the package owns only HA/MQTT transport, readiness, and artifact primitives.
|
||||
"""
|
||||
|
||||
from .artifacts import write_json_artifact, write_json_lines
|
||||
from .home_assistant import HomeAssistantClient
|
||||
from .mqtt import MqttObserver, RetainedPublisher
|
||||
from .wait import ContractError, wait_until
|
||||
|
||||
__all__ = [
|
||||
"ContractError",
|
||||
"HomeAssistantClient",
|
||||
"MqttObserver",
|
||||
"RetainedPublisher",
|
||||
"wait_until",
|
||||
"write_json_artifact",
|
||||
"write_json_lines",
|
||||
]
|
||||
@@ -0,0 +1,30 @@
|
||||
"""Small artifact writers which keep container tests independent of host tools."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import pathlib
|
||||
from collections.abc import Iterable
|
||||
from typing import Any
|
||||
|
||||
|
||||
def _path(root: str | pathlib.Path, name: str) -> pathlib.Path:
|
||||
path = pathlib.Path(root) / name
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
return path
|
||||
|
||||
|
||||
def write_json_artifact(root: str | pathlib.Path, name: str, value: Any) -> pathlib.Path:
|
||||
"""Write a deterministic UTF-8 JSON artifact and return its path."""
|
||||
path = _path(root, name)
|
||||
path.write_text(json.dumps(value, indent=2, sort_keys=True) + "\n", encoding="utf-8")
|
||||
return path
|
||||
|
||||
|
||||
def write_json_lines(root: str | pathlib.Path, name: str, values: Iterable[Any]) -> pathlib.Path:
|
||||
"""Write a deterministic JSONL artifact and return its path."""
|
||||
path = _path(root, name)
|
||||
with path.open("w", encoding="utf-8") as output:
|
||||
for value in values:
|
||||
output.write(json.dumps(value, sort_keys=True) + "\n")
|
||||
return path
|
||||
@@ -0,0 +1,245 @@
|
||||
"""Home Assistant onboarding, MQTT setup, registry, and service helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import pathlib
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import requests
|
||||
import websocket
|
||||
|
||||
from .wait import ContractError, wait_until
|
||||
|
||||
|
||||
class HomeAssistantClient:
|
||||
"""Authenticated HA API client with a deliberately small contract surface."""
|
||||
|
||||
def __init__(self, base_url: str, token: str):
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.token = token
|
||||
self._socket: websocket.WebSocket | None = None
|
||||
self._next_id = 1
|
||||
|
||||
@property
|
||||
def _headers(self) -> dict[str, str]:
|
||||
return {"Authorization": f"Bearer {self.token}"}
|
||||
|
||||
@classmethod
|
||||
def bootstrap(
|
||||
cls,
|
||||
base_url: str,
|
||||
state_dir: str | pathlib.Path,
|
||||
*,
|
||||
owner_name: str = "HA MQTT Contract Owner",
|
||||
username: str = "ha-mqtt-contract",
|
||||
password: str = "ha-mqtt-contract-password",
|
||||
) -> "HomeAssistantClient":
|
||||
"""Create or load an isolated owner token for a disposable HA volume."""
|
||||
base_url = base_url.rstrip("/")
|
||||
state = pathlib.Path(state_dir)
|
||||
token_file = state / "ha-token"
|
||||
cls.wait_until_ready(base_url)
|
||||
if token_file.exists():
|
||||
return cls(base_url, token_file.read_text(encoding="utf-8").strip())
|
||||
|
||||
client_id = "http://ha-mqtt-contract.local/"
|
||||
user = {
|
||||
"client_id": client_id,
|
||||
"name": owner_name,
|
||||
"username": username,
|
||||
"password": password,
|
||||
"language": "en",
|
||||
}
|
||||
created = cls._response_json(
|
||||
requests.post(f"{base_url}/api/onboarding/users", json=user, timeout=10),
|
||||
"Home Assistant onboarding user creation",
|
||||
)
|
||||
auth_code = created.get("auth_code")
|
||||
if not auth_code:
|
||||
raise ContractError("Home Assistant onboarding did not return an auth_code")
|
||||
token_response = cls._response_json(
|
||||
requests.post(
|
||||
f"{base_url}/auth/token",
|
||||
data={"client_id": client_id, "grant_type": "authorization_code", "code": auth_code},
|
||||
timeout=10,
|
||||
),
|
||||
"Home Assistant token exchange",
|
||||
)
|
||||
token = token_response.get("access_token")
|
||||
if not token:
|
||||
raise ContractError("Home Assistant token exchange did not return an access_token")
|
||||
|
||||
client = cls(base_url, token)
|
||||
for path, payload in (
|
||||
("/api/onboarding/core_config", {}),
|
||||
("/api/onboarding/analytics", {"preferences": {}}),
|
||||
):
|
||||
response = requests.post(f"{base_url}{path}", headers=client._headers, json=payload, timeout=10)
|
||||
if response.status_code not in (200, 201, 400, 404):
|
||||
raise ContractError(
|
||||
f"Home Assistant onboarding step {path} failed: {response.status_code} {response.text}"
|
||||
)
|
||||
state.mkdir(parents=True, exist_ok=True)
|
||||
token_file.write_text(token, encoding="utf-8")
|
||||
return client
|
||||
|
||||
@classmethod
|
||||
def from_state(cls, base_url: str, state_dir: str | pathlib.Path) -> "HomeAssistantClient":
|
||||
token_file = pathlib.Path(state_dir) / "ha-token"
|
||||
if not token_file.exists():
|
||||
raise ContractError("Home Assistant token has not been bootstrapped")
|
||||
return cls(base_url, token_file.read_text(encoding="utf-8").strip())
|
||||
|
||||
@staticmethod
|
||||
def _response_json(response: requests.Response, context: str) -> dict[str, Any]:
|
||||
if not response.ok:
|
||||
raise ContractError(f"{context} failed ({response.status_code}): {response.text}")
|
||||
try:
|
||||
value = response.json()
|
||||
except ValueError as error:
|
||||
raise ContractError(f"{context} returned invalid JSON: {error}") from error
|
||||
if not isinstance(value, dict):
|
||||
raise ContractError(f"{context} returned unexpected JSON: {value}")
|
||||
return value
|
||||
|
||||
@classmethod
|
||||
def wait_until_ready(cls, base_url: str, *, timeout: float = 180) -> None:
|
||||
def ready() -> bool:
|
||||
response = requests.get(f"{base_url.rstrip('/')}/api/", timeout=5)
|
||||
return response.status_code in (200, 401)
|
||||
|
||||
wait_until("Home Assistant HTTP API", ready, timeout=timeout)
|
||||
|
||||
def configure_mqtt(self, host: str, port: int) -> None:
|
||||
"""Create HA's MQTT config entry, negotiating current/older form schemas."""
|
||||
flow = self._response_json(
|
||||
requests.post(
|
||||
f"{self.base_url}/api/config/config_entries/flow",
|
||||
headers=self._headers,
|
||||
json={"handler": "mqtt"},
|
||||
timeout=15,
|
||||
),
|
||||
"MQTT config-entry flow creation",
|
||||
)
|
||||
if flow.get("type") == "create_entry":
|
||||
return
|
||||
if flow.get("type") != "form" or not flow.get("flow_id"):
|
||||
raise ContractError(f"unexpected MQTT config-entry flow result: {flow}")
|
||||
user_input: dict[str, Any] = {"broker": host, "port": port}
|
||||
schema_names = {
|
||||
field.get("name") for field in flow.get("data_schema", []) if isinstance(field, dict)
|
||||
}
|
||||
if "other_settings" in schema_names:
|
||||
user_input["other_settings"] = {
|
||||
"set_client_cert": False,
|
||||
"set_ca_cert": "off",
|
||||
"transport": "tcp",
|
||||
}
|
||||
configured = self._response_json(
|
||||
requests.post(
|
||||
f"{self.base_url}/api/config/config_entries/flow/{flow['flow_id']}",
|
||||
headers=self._headers,
|
||||
json=user_input,
|
||||
timeout=15,
|
||||
),
|
||||
"MQTT config-entry flow configuration",
|
||||
)
|
||||
if configured.get("type") != "create_entry":
|
||||
raise ContractError(f"MQTT config-entry flow did not create an entry: {configured}")
|
||||
time.sleep(3)
|
||||
|
||||
def _ensure_socket(self) -> websocket.WebSocket:
|
||||
if self._socket is not None:
|
||||
return self._socket
|
||||
scheme = "wss" if self.base_url.startswith("https://") else "ws"
|
||||
address = self.base_url.split("://", 1)[1]
|
||||
socket = websocket.create_connection(f"{scheme}://{address}/api/websocket", timeout=15)
|
||||
required = json.loads(socket.recv())
|
||||
if required.get("type") != "auth_required":
|
||||
socket.close()
|
||||
raise ContractError(f"unexpected Home Assistant WebSocket greeting: {required}")
|
||||
socket.send(json.dumps({"type": "auth", "access_token": self.token}))
|
||||
authenticated = json.loads(socket.recv())
|
||||
if authenticated.get("type") != "auth_ok":
|
||||
socket.close()
|
||||
raise ContractError(f"Home Assistant WebSocket authentication failed: {authenticated}")
|
||||
self._socket = socket
|
||||
return socket
|
||||
|
||||
def close(self) -> None:
|
||||
if self._socket is not None:
|
||||
self._socket.close()
|
||||
self._socket = None
|
||||
|
||||
def call(self, message_type: str, **kwargs: Any) -> Any:
|
||||
socket = self._ensure_socket()
|
||||
message_id = self._next_id
|
||||
self._next_id += 1
|
||||
socket.send(json.dumps({"id": message_id, "type": message_type, **kwargs}))
|
||||
while True:
|
||||
result = json.loads(socket.recv())
|
||||
if result.get("id") != message_id:
|
||||
continue
|
||||
if not result.get("success"):
|
||||
raise ContractError(f"WebSocket {message_type} failed: {result}")
|
||||
return result.get("result")
|
||||
|
||||
def entity_registry(self) -> list[dict[str, Any]]:
|
||||
result = self.call("config/entity_registry/list")
|
||||
if not isinstance(result, list):
|
||||
raise ContractError(f"entity registry returned unexpected value: {result}")
|
||||
return result
|
||||
|
||||
# A concise compatibility spelling for contract suites; this remains schema-neutral.
|
||||
def registry_entries(self) -> list[dict[str, Any]]:
|
||||
return self.entity_registry()
|
||||
|
||||
def device_registry(self) -> list[dict[str, Any]]:
|
||||
result = self.call("config/device_registry/list")
|
||||
if not isinstance(result, list):
|
||||
raise ContractError(f"device registry returned unexpected value: {result}")
|
||||
return result
|
||||
|
||||
def entities_for_device(self, device_id: str) -> list[dict[str, Any]]:
|
||||
return [entry for entry in self.entity_registry() if entry.get("device_id") == device_id]
|
||||
|
||||
def wait_for_entity(self, unique_id: str, *, timeout: float = 60) -> dict[str, Any]:
|
||||
def find() -> dict[str, Any] | None:
|
||||
matches = [entry for entry in self.entity_registry() if entry.get("unique_id") == unique_id]
|
||||
if len(matches) > 1:
|
||||
raise ContractError(f"duplicate entity-registry entries for {unique_id}: {matches}")
|
||||
return matches[0] if matches else None
|
||||
|
||||
return wait_until(f"entity registry entry {unique_id}", find, timeout=timeout)
|
||||
|
||||
def state(self, entity_id: str) -> dict[str, Any] | None:
|
||||
response = requests.get(f"{self.base_url}/api/states/{entity_id}", headers=self._headers, timeout=10)
|
||||
if response.status_code == 404:
|
||||
return None
|
||||
return self._response_json(response, f"state lookup for {entity_id}")
|
||||
|
||||
def wait_for_state(self, entity_id: str, state: str, *, timeout: float = 60) -> dict[str, Any]:
|
||||
return wait_until(
|
||||
f"Home Assistant state {entity_id}={state}",
|
||||
lambda: value if (value := self.state(entity_id)) and value.get("state") == state else None,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
def call_service(self, domain: str, service: str, data: dict[str, Any]) -> list[dict[str, Any]]:
|
||||
response = requests.post(
|
||||
f"{self.base_url}/api/services/{domain}/{service}",
|
||||
headers=self._headers,
|
||||
json=data,
|
||||
timeout=15,
|
||||
)
|
||||
if not response.ok:
|
||||
raise ContractError(f"service {domain}.{service} failed ({response.status_code}): {response.text}")
|
||||
try:
|
||||
result = response.json()
|
||||
except ValueError as error:
|
||||
raise ContractError(f"service {domain}.{service} returned invalid JSON: {error}") from error
|
||||
if not isinstance(result, list):
|
||||
raise ContractError(f"service {domain}.{service} returned unexpected JSON: {result}")
|
||||
return result
|
||||
@@ -0,0 +1,122 @@
|
||||
"""MQTT publishing and observation helpers for retained-discovery contracts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Iterable
|
||||
from typing import Any
|
||||
|
||||
import paho.mqtt.client as mqtt
|
||||
|
||||
from .wait import ContractError, wait_until
|
||||
|
||||
|
||||
class _MqttClient:
|
||||
def __init__(self, host: str, port: int, *, client_prefix: str):
|
||||
self._connected = threading.Event()
|
||||
self._client = mqtt.Client(
|
||||
mqtt.CallbackAPIVersion.VERSION2,
|
||||
client_id=f"{client_prefix}-{uuid.uuid4()}",
|
||||
)
|
||||
self._client.on_connect = self._on_connect
|
||||
self._client.connect(host, port, keepalive=15)
|
||||
self._client.loop_start()
|
||||
if not self._connected.wait(timeout=15):
|
||||
self.close()
|
||||
raise ContractError(f"MQTT client {client_prefix} did not connect within 15 seconds")
|
||||
|
||||
def _on_connect(self, _client: mqtt.Client, _userdata: Any, _flags: Any, reason_code: Any, _properties: Any) -> None:
|
||||
if reason_code == 0:
|
||||
self._connected.set()
|
||||
|
||||
def close(self) -> None:
|
||||
self._client.loop_stop()
|
||||
self._client.disconnect()
|
||||
|
||||
|
||||
class RetainedPublisher(_MqttClient):
|
||||
"""A QoS-1 retained publisher used by schema/migration fixtures."""
|
||||
|
||||
def __init__(self, host: str, port: int):
|
||||
super().__init__(host, port, client_prefix="contract-publisher")
|
||||
|
||||
def publish(self, topic: str, payload: str, *, retain: bool = True) -> None:
|
||||
info = self._client.publish(topic, payload, qos=1, retain=retain)
|
||||
info.wait_for_publish(timeout=10)
|
||||
if not info.is_published():
|
||||
raise ContractError(f"MQTT publish timed out for {topic}")
|
||||
|
||||
def retained_payload(self, topic: str, *, timeout: float = 30) -> str:
|
||||
received: list[str] = []
|
||||
|
||||
def on_message(_client: mqtt.Client, _userdata: Any, message: mqtt.MQTTMessage) -> None:
|
||||
if message.topic == topic:
|
||||
received.append(message.payload.decode("utf-8"))
|
||||
|
||||
self._client.on_message = on_message
|
||||
self._client.subscribe(topic, qos=1)
|
||||
try:
|
||||
return wait_until(
|
||||
f"retained MQTT message {topic}",
|
||||
lambda: received[0] if received else None,
|
||||
timeout=timeout,
|
||||
)
|
||||
finally:
|
||||
self._client.unsubscribe(topic)
|
||||
self._client.on_message = None
|
||||
|
||||
|
||||
class MqttObserver(_MqttClient):
|
||||
"""Observe MQTT traffic without imposing a discovery-schema interpretation."""
|
||||
|
||||
def __init__(self, host: str, port: int, topics: Iterable[str] = ("homeassistant/#",)):
|
||||
self._messages: list[dict[str, Any]] = []
|
||||
self._lock = threading.Lock()
|
||||
self._subscription_ack = threading.Event()
|
||||
super().__init__(host, port, client_prefix="contract-observer")
|
||||
self._client.on_message = self._on_message
|
||||
self._client.on_subscribe = self._on_subscribe
|
||||
for topic in topics:
|
||||
result, _mid = self._client.subscribe(topic, qos=1)
|
||||
if result != mqtt.MQTT_ERR_SUCCESS:
|
||||
self.close()
|
||||
raise ContractError(f"unable to subscribe to MQTT topic {topic}")
|
||||
if not self._subscription_ack.wait(timeout=15):
|
||||
self.close()
|
||||
raise ContractError(f"MQTT subscription to {topic} was not acknowledged within 15 seconds")
|
||||
self._subscription_ack.clear()
|
||||
|
||||
def _on_subscribe(
|
||||
self,
|
||||
_client: mqtt.Client,
|
||||
_userdata: Any,
|
||||
_mid: int,
|
||||
_reason_codes: Any,
|
||||
_properties: Any,
|
||||
) -> None:
|
||||
self._subscription_ack.set()
|
||||
|
||||
def _on_message(self, _client: mqtt.Client, _userdata: Any, message: mqtt.MQTTMessage) -> None:
|
||||
with self._lock:
|
||||
self._messages.append(
|
||||
{
|
||||
"topic": message.topic,
|
||||
"payload": message.payload.decode("utf-8", errors="replace"),
|
||||
"qos": message.qos,
|
||||
"retain": message.retain,
|
||||
"received_at": time.time(),
|
||||
}
|
||||
)
|
||||
|
||||
def messages(self) -> list[dict[str, Any]]:
|
||||
with self._lock:
|
||||
return list(self._messages)
|
||||
|
||||
def wait_for_topic(self, topic: str, *, timeout: float = 60) -> dict[str, Any]:
|
||||
return wait_until(
|
||||
f"MQTT topic {topic}",
|
||||
lambda: next((entry for entry in reversed(self.messages()) if entry["topic"] == topic), None),
|
||||
timeout=timeout,
|
||||
)
|
||||
@@ -0,0 +1,36 @@
|
||||
"""Bounded wait helpers shared by integration contracts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from typing import TypeVar
|
||||
|
||||
|
||||
class ContractError(RuntimeError):
|
||||
"""An external Home Assistant or MQTT contract could not be satisfied."""
|
||||
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
def wait_until(
|
||||
description: str,
|
||||
predicate: Callable[[], T | None | bool],
|
||||
*,
|
||||
timeout: float = 90,
|
||||
interval: float = 1,
|
||||
) -> T:
|
||||
"""Return the first truthy predicate result or raise a useful timeout."""
|
||||
deadline = time.monotonic() + timeout
|
||||
last_error: Exception | None = None
|
||||
while time.monotonic() < deadline:
|
||||
try:
|
||||
result = predicate()
|
||||
if result:
|
||||
return result # type: ignore[return-value]
|
||||
except Exception as error: # Services are expected to be starting.
|
||||
last_error = error
|
||||
time.sleep(interval)
|
||||
suffix = f" (last error: {last_error})" if last_error else ""
|
||||
raise ContractError(f"timed out waiting for {description}{suffix}")
|
||||
@@ -1,8 +1,10 @@
|
||||
FROM python:3.12-slim
|
||||
|
||||
WORKDIR /tests
|
||||
COPY requirements.txt .
|
||||
COPY test-support/ha-mqtt-contract /opt/ha-mqtt-contract
|
||||
COPY tests/ha-contract/requirements.txt .
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
COPY test_contract.py .
|
||||
RUN pip install --no-cache-dir /opt/ha-mqtt-contract
|
||||
COPY tests/ha-contract/test_contract.py .
|
||||
|
||||
CMD ["python", "test_contract.py"]
|
||||
|
||||
@@ -44,6 +44,10 @@ disabled device components are cleaned. The test uses only ephemeral named volum
|
||||
`down -v` removes its broker data, Home Assistant config, owner token, and
|
||||
registry state.
|
||||
|
||||
## Shared support
|
||||
|
||||
The schema-neutral [HA/MQTT contract testkit](../../test-support/ha-mqtt-contract/README.md) owns only Docker-side Home Assistant onboarding, MQTT transport, registry/service access, retries, and artifact helpers. It has no ArduinoHA discovery assertions. This suite owns ArduinoHA migration fixtures; DeviceFramework carries its own fixtures and hardware adapter while reusing the testkit.
|
||||
|
||||
## Scope
|
||||
|
||||
The firmware's native Unity suite covers JSON escaping, invalid topic tokens,
|
||||
|
||||
@@ -23,7 +23,8 @@ services:
|
||||
|
||||
tests:
|
||||
build:
|
||||
context: .
|
||||
context: ../..
|
||||
dockerfile: tests/ha-contract/Dockerfile
|
||||
environment:
|
||||
HA_URL: http://homeassistant:8123
|
||||
MQTT_HOST: mqtt
|
||||
|
||||
@@ -10,11 +10,13 @@ import json
|
||||
import os
|
||||
import pathlib
|
||||
import time
|
||||
import uuid
|
||||
|
||||
import paho.mqtt.client as mqtt
|
||||
import requests
|
||||
import websocket
|
||||
from ha_mqtt_contract import (
|
||||
ContractError,
|
||||
HomeAssistantClient,
|
||||
RetainedPublisher as SharedRetainedPublisher,
|
||||
wait_until,
|
||||
)
|
||||
|
||||
|
||||
HA_URL = os.environ.get("HA_URL", "http://homeassistant:8123").rstrip("/")
|
||||
@@ -28,201 +30,8 @@ EDGE_DEVICE_ID = "contract_edge_device"
|
||||
EXPECT_DISABLED_CLEANUP = os.environ.get("CONTRACT_EXPECT_DISABLED_CLEANUP") == "1"
|
||||
|
||||
|
||||
class ContractFailure(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
def fail(message):
|
||||
raise ContractFailure(message)
|
||||
|
||||
|
||||
def wait_until(description, predicate, timeout=90, interval=1):
|
||||
deadline = time.monotonic() + timeout
|
||||
last_error = None
|
||||
while time.monotonic() < deadline:
|
||||
try:
|
||||
result = predicate()
|
||||
if result:
|
||||
return result
|
||||
except Exception as error: # HA is expected to be starting initially.
|
||||
last_error = error
|
||||
time.sleep(interval)
|
||||
suffix = f" (last error: {last_error})" if last_error else ""
|
||||
fail(f"timed out waiting for {description}{suffix}")
|
||||
|
||||
|
||||
def wait_for_home_assistant():
|
||||
def ready():
|
||||
response = requests.get(f"{HA_URL}/api/", timeout=5)
|
||||
return response.status_code in (200, 401)
|
||||
|
||||
wait_until("Home Assistant HTTP API", ready, timeout=180)
|
||||
|
||||
|
||||
def response_json(response, context):
|
||||
if not response.ok:
|
||||
fail(f"{context} failed ({response.status_code}): {response.text}")
|
||||
try:
|
||||
return response.json()
|
||||
except ValueError as error:
|
||||
fail(f"{context} returned invalid JSON: {error}")
|
||||
|
||||
|
||||
def onboarding_token():
|
||||
"""Create an ephemeral owner token once and share it across restart checks."""
|
||||
if TOKEN_FILE.exists():
|
||||
return TOKEN_FILE.read_text(encoding="utf-8").strip()
|
||||
|
||||
client_id = "http://contract-tests.local/"
|
||||
user = {
|
||||
"client_id": client_id,
|
||||
"name": "Contract Test Owner",
|
||||
"username": "contract-owner",
|
||||
"password": "contract-test-password",
|
||||
"language": "en",
|
||||
}
|
||||
response = requests.post(f"{HA_URL}/api/onboarding/users", json=user, timeout=10)
|
||||
created = response_json(response, "Home Assistant onboarding user creation")
|
||||
auth_code = created.get("auth_code")
|
||||
if not auth_code:
|
||||
fail("Home Assistant onboarding did not return an auth_code")
|
||||
|
||||
token_response = requests.post(
|
||||
f"{HA_URL}/auth/token",
|
||||
data={
|
||||
"client_id": client_id,
|
||||
"grant_type": "authorization_code",
|
||||
"code": auth_code,
|
||||
},
|
||||
timeout=10,
|
||||
)
|
||||
token = response_json(token_response, "Home Assistant token exchange").get("access_token")
|
||||
if not token:
|
||||
fail("Home Assistant token exchange did not return an access_token")
|
||||
|
||||
headers = {"Authorization": f"Bearer {token}"}
|
||||
# These operations are idempotent across HA versions that still expose them.
|
||||
for path, payload in (
|
||||
("/api/onboarding/core_config", {}),
|
||||
("/api/onboarding/analytics", {"preferences": {}}),
|
||||
):
|
||||
response = requests.post(f"{HA_URL}{path}", headers=headers, json=payload, timeout=10)
|
||||
if response.status_code not in (200, 201, 400, 404):
|
||||
fail(f"Home Assistant onboarding step {path} failed: {response.status_code} {response.text}")
|
||||
|
||||
STATE.mkdir(parents=True, exist_ok=True)
|
||||
TOKEN_FILE.write_text(token, encoding="utf-8")
|
||||
return token
|
||||
|
||||
|
||||
class HAWebSocket:
|
||||
def __init__(self, token):
|
||||
scheme = "wss" if HA_URL.startswith("https://") else "ws"
|
||||
address = HA_URL.split("://", 1)[1]
|
||||
self._socket = websocket.create_connection(f"{scheme}://{address}/api/websocket", timeout=15)
|
||||
required = json.loads(self._socket.recv())
|
||||
if required.get("type") != "auth_required":
|
||||
fail(f"unexpected Home Assistant WebSocket greeting: {required}")
|
||||
self._socket.send(json.dumps({"type": "auth", "access_token": token}))
|
||||
authenticated = json.loads(self._socket.recv())
|
||||
if authenticated.get("type") != "auth_ok":
|
||||
fail(f"Home Assistant WebSocket authentication failed: {authenticated}")
|
||||
self._next_id = 1
|
||||
|
||||
def close(self):
|
||||
self._socket.close()
|
||||
|
||||
def call(self, message_type, **kwargs):
|
||||
message_id = self._next_id
|
||||
self._next_id += 1
|
||||
self._socket.send(json.dumps({"id": message_id, "type": message_type, **kwargs}))
|
||||
while True:
|
||||
result = json.loads(self._socket.recv())
|
||||
if result.get("id") != message_id:
|
||||
continue
|
||||
if not result.get("success"):
|
||||
fail(f"WebSocket {message_type} failed: {result}")
|
||||
return result.get("result")
|
||||
|
||||
def registry_entries(self):
|
||||
return self.call("config/entity_registry/list")
|
||||
|
||||
|
||||
def configure_mqtt_integration(token):
|
||||
"""Create HA's MQTT config entry; broker YAML options are no longer accepted."""
|
||||
headers = {"Authorization": f"Bearer {token}"}
|
||||
response = requests.post(
|
||||
f"{HA_URL}/api/config/config_entries/flow",
|
||||
headers=headers,
|
||||
json={"handler": "mqtt"},
|
||||
timeout=15,
|
||||
)
|
||||
flow = response_json(response, "MQTT config-entry flow creation")
|
||||
if flow.get("type") == "create_entry":
|
||||
return
|
||||
if flow.get("type") != "form" or not flow.get("flow_id"):
|
||||
fail(f"unexpected MQTT config-entry flow result: {flow}")
|
||||
|
||||
user_input = {"broker": MQTT_HOST, "port": MQTT_PORT}
|
||||
# Newer HA versions group TLS/transport values in a required section.
|
||||
# Older supported versions do not expose the section and must not receive
|
||||
# unknown fields, so negotiate from the returned data schema.
|
||||
schema_names = {
|
||||
field.get("name")
|
||||
for field in flow.get("data_schema", [])
|
||||
if isinstance(field, dict)
|
||||
}
|
||||
if "other_settings" in schema_names:
|
||||
user_input["other_settings"] = {
|
||||
"set_client_cert": False,
|
||||
"set_ca_cert": "off",
|
||||
"transport": "tcp",
|
||||
}
|
||||
|
||||
response = requests.post(
|
||||
f"{HA_URL}/api/config/config_entries/flow/{flow['flow_id']}",
|
||||
headers=headers,
|
||||
json=user_input,
|
||||
timeout=15,
|
||||
)
|
||||
configured = response_json(response, "MQTT config-entry flow configuration")
|
||||
if configured.get("type") != "create_entry":
|
||||
fail(f"MQTT config-entry flow did not create an entry: {configured}")
|
||||
|
||||
# The successful flow result is returned before the integration has had a
|
||||
# chance to subscribe to retained discovery topics.
|
||||
time.sleep(3)
|
||||
|
||||
|
||||
class RetainedPublisher:
|
||||
def __init__(self):
|
||||
self._client = mqtt.Client(mqtt.CallbackAPIVersion.VERSION2, client_id=f"contract-{uuid.uuid4()}")
|
||||
self._client.connect(MQTT_HOST, MQTT_PORT, keepalive=15)
|
||||
self._client.loop_start()
|
||||
|
||||
def close(self):
|
||||
self._client.loop_stop()
|
||||
self._client.disconnect()
|
||||
|
||||
def publish(self, topic, payload):
|
||||
info = self._client.publish(topic, payload, qos=1, retain=True)
|
||||
info.wait_for_publish(timeout=10)
|
||||
if not info.is_published():
|
||||
fail(f"MQTT publish timed out for {topic}")
|
||||
|
||||
def retained_payload(self, topic):
|
||||
received = []
|
||||
|
||||
def on_message(_client, _userdata, message):
|
||||
if message.topic == topic:
|
||||
received.append(message.payload.decode("utf-8"))
|
||||
|
||||
self._client.on_message = on_message
|
||||
self._client.subscribe(topic, qos=1)
|
||||
value = wait_until(f"retained MQTT message {topic}", lambda: received[0] if received else None)
|
||||
self._client.unsubscribe(topic)
|
||||
self._client.on_message = None
|
||||
return value
|
||||
raise ContractError(message)
|
||||
|
||||
|
||||
def legacy_topic(object_id, device_id=DEVICE_ID):
|
||||
@@ -275,10 +84,15 @@ def wait_for_entry(ws, unique_id):
|
||||
|
||||
|
||||
def migration_contract():
|
||||
publisher = RetainedPublisher()
|
||||
token = onboarding_token()
|
||||
configure_mqtt_integration(token)
|
||||
ws = HAWebSocket(token)
|
||||
publisher = SharedRetainedPublisher(MQTT_HOST, MQTT_PORT)
|
||||
ws = HomeAssistantClient.bootstrap(
|
||||
HA_URL,
|
||||
STATE,
|
||||
owner_name="ArduinoHA Contract Owner",
|
||||
username="arduinoha-contract",
|
||||
password="arduinoha-contract-password",
|
||||
)
|
||||
ws.configure_mqtt(MQTT_HOST, MQTT_PORT)
|
||||
try:
|
||||
# Existing single-component entity and a user-owned registry customization.
|
||||
unique = f"{DEVICE_ID}_temperature"
|
||||
@@ -437,9 +251,8 @@ def migration_contract():
|
||||
def retained_restart_contract():
|
||||
if not (STATE / "migration-complete").exists():
|
||||
fail("retained-restart mode requires the migration contract to run first")
|
||||
publisher = RetainedPublisher()
|
||||
token = onboarding_token()
|
||||
ws = HAWebSocket(token)
|
||||
publisher = SharedRetainedPublisher(MQTT_HOST, MQTT_PORT)
|
||||
ws = HomeAssistantClient.from_state(HA_URL, STATE)
|
||||
try:
|
||||
migrated = wait_for_entry(ws, f"{DEVICE_ID}_temperature")
|
||||
if migrated["entity_id"] != "sensor.contract_temperature_user_name":
|
||||
@@ -453,7 +266,7 @@ def retained_restart_contract():
|
||||
|
||||
|
||||
def main():
|
||||
wait_for_home_assistant()
|
||||
HomeAssistantClient.wait_until_ready(HA_URL)
|
||||
if MODE == "migration":
|
||||
migration_contract()
|
||||
elif MODE == "retained-restart":
|
||||
|
||||
Reference in New Issue
Block a user