Files
houseplan-card/custom_components/houseplan/zigbee_topology.py
T
2026-10-05 20:08:37 +03:00

442 lines
18 KiB
Python

"""Bounded, entry-owned Zigbee2MQTT requests; never a polling radio scanner."""
from __future__ import annotations
import asyncio
import json
import logging
import math
import re
from collections.abc import Awaitable, Callable
from dataclasses import dataclass, field
from time import monotonic, time
from typing import Any
from uuid import uuid4
from homeassistant.components import mqtt
from homeassistant.config_entries import ConfigEntry
from homeassistant.core import HomeAssistant, callback
MAX_TOPICS = 8
MAX_LISTENERS = 32
MAX_PAYLOAD_BYTES = 2 * 1024 * 1024
MAX_NODES = 1000
MAX_LINKS = 6000
CANCEL_AFTER_MS = 600_000
TRANSPORT_TIMEOUT = 10.0
INFO_TIMEOUT = 4.0
_LOGGER = logging.getLogger(__name__)
class ZigbeeScanError(Exception):
"""A stable, already localized WebSocket command refusal."""
def __init__(self, code: str) -> None:
super().__init__(code)
self.code = code
def normalize_base_topic(value: object) -> str:
"""Match the frontend's exact-topic normalization, without a saved-topic gate."""
if not isinstance(value, str):
raise ZigbeeScanError("invalid_data")
topic = re.sub(r"/{2,}", "/", value.strip().strip("/"))
# JS counts UTF-16 code units; reject lone surrogates as invalid MQTT UTF-8.
try:
length = len(topic.encode("utf-16-le")) // 2
topic.encode("utf-8")
except UnicodeError as err:
raise ZigbeeScanError("invalid_data") from err
if not topic or length > 180 or re.search(r"[#+\x00-\x1f\x7f]", topic):
raise ZigbeeScanError("invalid_data")
return topic
def _parse_payload(payload: object) -> dict[str, Any] | None:
"""Bound bytes before decoding; uncorrelatable garbage must not end a job."""
try:
if isinstance(payload, str):
payload = payload.encode("utf-8")
if not isinstance(payload, bytes) or len(payload) > MAX_PAYLOAD_BYTES:
return None
value = json.loads(payload)
if not isinstance(value, dict):
return None
pending = [(value, 0)]
inspected = 0
while pending:
item, depth = pending.pop()
inspected += 1
if depth > 32 or inspected > 100_000:
return None
if isinstance(item, dict):
pending.extend((child, depth + 1) for child in item.values())
elif isinstance(item, list):
pending.extend((child, depth + 1) for child in item)
elif isinstance(item, float) and not math.isfinite(item):
return None
return value
except (ValueError, UnicodeError, RecursionError):
return None
def _valid_map(message: dict[str, Any]) -> bool:
"""Check the raw-map shape, not routing evidence (that remains TypeScript)."""
data = message.get("data")
value = data.get("value", data) if isinstance(data, dict) else data
if value is None:
value = message.get("value", message)
if isinstance(value, str):
value = _parse_payload(value)
if not isinstance(value, dict):
return False
nodes, links = value.get("nodes"), value.get("links")
if not isinstance(nodes, list) or len(nodes) > MAX_NODES:
return False
if not isinstance(links, list) or len(links) > MAX_LINKS:
return False
for node in nodes:
if not isinstance(node, dict):
return False
ieee = node.get("ieeeAddr", node.get("ieee_address", node.get("ieee")))
if not isinstance(ieee, str) or not re.fullmatch(
r"(?:0x)?[0-9a-fA-F]{16}|(?:[0-9a-fA-F]{2}[:-]){7}[0-9a-fA-F]{2}", ieee,
):
return False
routes_count = 0
for link in links:
if not isinstance(link, dict):
return False
for side in ("source", "target"):
endpoint = link.get(f"{side}IeeeAddr", link.get(side))
if isinstance(endpoint, bool) or not isinstance(endpoint, (dict, str, int)):
return False
if not endpoint and endpoint != 0:
return False
routes = link.get("routes", [])
if not isinstance(routes, list) or any(not isinstance(route, dict) for route in routes):
return False
routes_count += len(routes)
if routes_count > MAX_LINKS:
return False
return True
@dataclass
class _Job:
topic: str
job_id: str
started_at: int
started: float
info: asyncio.Future[None]
phase: str = "loading"
stage: str = "connecting"
finished: float | None = None
error: str | None = None
stale: bool = False
result: dict[str, Any] | None = None
obtained_at: int | None = None
response_active: bool = False
task: asyncio.Task[None] | None = None
cleanups: list[Callable[[], None]] = field(default_factory=list)
class ZigbeeScanCoordinator:
"""One in-memory slot per topic, shared by all browsers of the HA entry."""
def __init__(self, hass: HomeAssistant, entry: ConfigEntry) -> None:
self.hass = hass
self.entry = entry
self.session_id = str(uuid4())
self.revision = 0
self.closed = False
self._jobs: dict[str, _Job] = {}
self._listeners: set[Callable[[dict[str, Any]], None]] = set()
self._operations: set[asyncio.Task] = set()
def _active(self, job: _Job) -> bool:
return not self.closed and self._jobs.get(job.topic) is job and job.phase == "loading"
def _state(self, job: _Job) -> dict[str, Any]:
provider: dict[str, Any] = {
"topic": job.topic, "job_id": job.job_id, "phase": job.phase,
"stage": job.stage, "started_at": job.started_at,
"elapsed_ms": max(0, int(((job.finished if job.finished is not None else monotonic()) - job.started) * 1000)),
"cancel_after_ms": CANCEL_AFTER_MS,
}
if job.error is not None:
provider["error"] = job.error
if job.result is not None:
provider["result"] = job.result
provider["obtained_at"] = job.obtained_at
provider["stale"] = job.stale
return {"kind": "state", "session_id": self.session_id,
"revision": self.revision, "provider": provider}
def snapshot(self, topic: str) -> dict[str, Any]:
"""Current public state; no MQTT work, no timestamp reset."""
if self.closed:
raise ZigbeeScanError("not_ready")
job = self._jobs.get(normalize_base_topic(topic))
if job is None:
raise ZigbeeScanError("invalid_data")
return self._state(job)
def initial_events(self) -> list[dict[str, Any]]:
"""Bound each event to one map; never aggregate eight maps into one frame."""
return [{"kind": "reset", "session_id": self.session_id,
"revision": self.revision, "topics": list(self._jobs)},
*(self._state(job) for job in self._jobs.values())]
def add_listener(self, listener: Callable[[dict[str, Any]], None]) -> Callable[[], None]:
if self.closed:
raise ZigbeeScanError("not_ready")
if len(self._listeners) >= MAX_LISTENERS:
raise ZigbeeScanError("capacity_exceeded")
self._listeners.add(listener)
return lambda: self._listeners.discard(listener)
def _emit(self, event: dict[str, Any]) -> None:
for listener in tuple(self._listeners):
try:
listener(event)
except Exception: # noqa: BLE001 - one disconnected UI must not interrupt a shared job
self._listeners.discard(listener)
def _changed(self, job: _Job) -> None:
self.revision += 1
self._emit(self._state(job))
def start(self, base_topic: str) -> dict[str, Any]:
"""Reserve before creating any task: concurrent WS clients join one job."""
if self.closed:
raise ZigbeeScanError("not_ready")
topic = normalize_base_topic(base_topic)
previous = self._jobs.get(topic)
if previous is not None and previous.phase == "loading":
return self._state(previous)
# A transport which acknowledges cancellation late still owns a bounded
# slot until its cleanup returns; repeated failed starts cannot grow it.
if len(self._operations) >= MAX_TOPICS:
raise ZigbeeScanError("capacity_exceeded")
if previous is None and len(self._jobs) >= MAX_TOPICS:
candidate = next((item for item in self._jobs.values() if item.phase != "loading"), None)
if candidate is None:
raise ZigbeeScanError("capacity_exceeded")
del self._jobs[candidate.topic]
self.revision += 1
self._emit({"kind": "removed", "session_id": self.session_id,
"revision": self.revision, "topic": candidate.topic})
job = _Job(topic, f"houseplan-{uuid4()}", int(time() * 1000), monotonic(),
self.hass.loop.create_future(),
stale=previous.stale if previous else False,
result=previous.result if previous else None,
obtained_at=previous.obtained_at if previous else None)
self._jobs.pop(topic, None)
self._jobs[topic] = job
job.task = self.entry.async_create_background_task(
self.hass, self._run(job), "House Plan Zigbee map", eager_start=False,
)
self._changed(job)
return self._state(job)
def cancel(self, base_topic: str, job_id: str) -> dict[str, Any]:
if self.closed:
raise ZigbeeScanError("not_ready")
topic = normalize_base_topic(base_topic)
job = self._jobs.get(topic)
if job is None or job.job_id != job_id:
raise ZigbeeScanError("conflict")
# A repeated cancellation or a success/cancel race never rolls back a final state.
if job.phase != "loading":
return self._state(job)
if (monotonic() - job.started) * 1000 < CANCEL_AFTER_MS:
raise ZigbeeScanError("conflict")
self._finish(job, "cancelled")
return self._state(job)
@staticmethod
def _cleanup(cleanup: Callable[[], None]) -> None:
try:
cleanup()
except Exception: # noqa: BLE001 - cleanup remains idempotent across unload and transport failure
_LOGGER.debug("House Plan Zigbee transport cleanup failed")
def _add_cleanup(self, job: _Job, cleanup: Callable[[], None]) -> None:
if self._active(job):
job.cleanups.append(cleanup)
else:
self._cleanup(cleanup)
def _release(self, job: _Job) -> None:
job.response_active = False
cleanups, job.cleanups = job.cleanups, []
for cleanup in cleanups:
self._cleanup(cleanup)
if not job.info.done():
job.info.cancel()
def _finish(self, job: _Job, phase: str, error: str | None = None,
result: dict[str, Any] | None = None) -> None:
if not self._active(job):
return
job.phase, job.error, job.finished = phase, error, monotonic()
if result is not None:
job.result, job.obtained_at = result, int(time() * 1000)
job.stale = False
elif job.result is not None:
job.stale = True
self._release(job)
# Including when called synchronously inside publish: a terminal provider
# reply must not wait another ten seconds for the publish acknowledgement.
if job.task is not None and not job.task.done():
job.task.cancel()
self._changed(job)
async def _operation(self, job: _Job, action: Callable[[], Awaitable],
seconds: float, *, subscription: bool = False):
"""Bound transport even if a cancelled subscribe later returns an unsubscribe."""
async def perform():
result = await action()
if subscription and callable(result):
self._add_cleanup(job, result)
return result
task = self.entry.async_create_background_task(
self.hass, perform(), "House Plan Zigbee transport", eager_start=False,
)
self._operations.add(task)
@callback
def settled(done: asyncio.Task) -> None:
self._operations.discard(done)
if not done.cancelled():
done.exception() # consume a late transport exception after the job ended
task.add_done_callback(settled)
try:
done, _pending = await asyncio.wait({task}, timeout=max(0, seconds))
if not done:
raise TimeoutError
return task.result()
finally:
if not task.done():
task.cancel()
async def _subscribe(self, job: _Job, topic: str, handler: Callable) -> None:
deadline = self.hass.loop.time() + TRANSPORT_TIMEOUT
await self._operation(job, lambda: mqtt.async_subscribe(
self.hass, topic, handler, qos=0, encoding=None,
), TRANSPORT_TIMEOUT, subscription=True)
# Optional on the declared HA floor (2024.6). This acknowledges HA's
# processing of subscriptions, not provider acceptance or radio progress.
processed = getattr(mqtt, "async_on_subscribe_done", None)
if callable(processed) and self._active(job):
ready = self.hass.loop.create_future()
@callback
def on_ready() -> None:
if not ready.done():
ready.set_result(None)
remove = processed(self.hass, topic, 0, on_ready)
self._add_cleanup(job, remove)
async with asyncio.timeout(max(0, deadline - self.hass.loop.time())):
await ready
async def _run(self, job: _Job) -> None:
try:
if not mqtt.mqtt_config_entry_enabled(self.hass):
self._finish(job, "error", "unsupported")
return
if not mqtt.is_connected(self.hass):
self._finish(job, "error", "connection")
return
@callback
def connection_changed(connected: bool) -> None:
if not connected:
self._finish(job, "error", "connection")
self._add_cleanup(job, mqtt.async_subscribe_connection_status(
self.hass, connection_changed,
))
response_topic = f"{job.topic}/bridge/response/networkmap"
info_topic = f"{job.topic}/bridge/info"
@callback
def receive_response(message) -> None:
if (not self._active(job) or not job.response_active
or message.topic != response_topic or message.retain):
return
value = _parse_payload(message.payload)
if value is None:
return
data = value.get("data")
transaction = value.get("transaction", data.get("transaction") if isinstance(data, dict) else None)
if transaction != job.job_id:
return
if value.get("status") not in (None, "ok"):
self._finish(job, "error", "provider")
elif not _valid_map(value):
self._finish(job, "error", "invalid_payload")
else:
self._finish(job, "ready", result=value)
@callback
def receive_info(message) -> None:
if (self._active(job) and message.topic == info_topic and message.retain
and not job.info.done() and _parse_payload(message.payload) is not None):
job.info.set_result(None)
# Response first also improves the old-HA retained-info preflight;
# retained info alone is not claimed to be a strict broker SUBACK.
await self._subscribe(job, response_topic, receive_response)
await self._subscribe(job, info_topic, receive_info)
async with asyncio.timeout(INFO_TIMEOUT):
await job.info
if not self._active(job):
return
job.response_active = True
await self._operation(job, lambda: mqtt.async_publish(
self.hass, f"{job.topic}/bridge/request/networkmap",
json.dumps({"type": "raw", "routes": True, "transaction": job.job_id}),
qos=0, retain=False,
), TRANSPORT_TIMEOUT)
if self._active(job):
job.stage = "waiting"
self._changed(job)
# Only MQTT, explicit cancel, transport loss or entry teardown
# ends this wait. 600 seconds is a cancel threshold, not a deadline.
await self.hass.loop.create_future()
except asyncio.CancelledError:
if self._active(job):
self._finish(job, "error", "connection")
raise
except TimeoutError:
self._finish(job, "error", "timeout")
except Exception: # noqa: BLE001 - contain optional MQTT errors without logging network payloads
self._finish(job, "error", "connection")
finally:
self._release(job)
async def async_close(self) -> None:
"""Stop entry-owned work, clear volatile maps, never restart a radio scan."""
if self.closed:
return
self.closed = True
# Reload does not disconnect HA WebSockets. Tell their observers that
# this session and its volatile cache are gone before dropping them.
self.revision += 1
self._emit({"kind": "closed", "session_id": self.session_id,
"revision": self.revision})
tasks = [job.task for job in self._jobs.values() if job.task is not None]
for job in self._jobs.values():
self._release(job)
for task in (*tasks, *self._operations):
if not task.done():
task.cancel()
if tasks:
await asyncio.gather(*tasks, return_exceptions=True)
self._jobs.clear()
self._listeners.clear()