Files
HomeAssistantVS/custom_components/ha_washdata/task_registry.py
T

289 lines
11 KiB
Python

# WashData - Home Assistant integration for appliance cycle monitoring via smart plugs.
# Copyright (C) 2026 Lukas Bandura
# SPDX-License-Identifier: AGPL-3.0-or-later
#
# This program is free software: you can redistribute it and/or modify
# it under the terms of the GNU Affero General Public License as published
# by the Free Software Foundation, either version 3 of the License, or
# (at your option) any later version.
#
# This program is distributed in the hope that it will be useful,
# but WITHOUT ANY WARRANTY; without even the implied warranty of
# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
# GNU Affero General Public License for more details.
#
# You should have received a copy of the GNU Affero General Public License
# along with this program. If not, see <https://www.gnu.org/licenses/>.
"""In-memory registry of long-running background tasks (reprocess, ML training,
Playground history/optimize).
Purpose: keep a task's progress, cancel handle and result on the *server* so they
survive a dropped WebSocket (backgrounded tab), can be cancelled, and can be
re-fetched on reconnect. One registry per ``hass``; each task is tagged with the
``entry_id`` it belongs to. No persistence - results live for the session and the
last few finished tasks are retained for reload.
Pure asyncio + synchronous listener callbacks; the WebSocket layer registers a
listener to push updates and calls :func:`get_registry` to read/kick/cancel.
"""
from __future__ import annotations
import asyncio
import logging
import uuid
from collections import OrderedDict
from dataclasses import dataclass, field
from typing import Any, Callable
from homeassistant.core import HomeAssistant
from homeassistant.util import dt as dt_util
from .const import DOMAIN
# Keep this many *finished* tasks (with their results) around for reload; older
# ones are evicted. Running tasks are never evicted.
_MAX_FINISHED = 30
_REGISTRY_KEY = f"{DOMAIN}_task_registry"
# Task lifecycle states.
STATE_RUNNING = "running"
STATE_DONE = "done"
STATE_ERROR = "error"
STATE_CANCELLED = "cancelled"
@dataclass
class Task:
"""A single tracked background operation."""
id: str
entry_id: str
kind: str # 'reprocess' | 'ml_training' | 'pg_history' | 'pg_sweep'
label: str # English fallback shown only if no label_key resolves
# Panel-localizable label: the pill renders _t(label_key, label_params, label)
# so per-step progress text is translated. When label_key is None the pill
# falls back to a per-kind translated action label.
label_key: str | None = None
label_params: dict[str, Any] = field(default_factory=dict)
total: int = 0
done: int = 0
state: str = STATE_RUNNING
error: str | None = None
started_at: float = field(default_factory=lambda: dt_util.now().timestamp())
updated_at: float = field(default_factory=lambda: dt_util.now().timestamp())
finished_at: float | None = None
result: Any = None
_cancelled: bool = False
@property
def cancel_requested(self) -> bool:
return self._cancelled
def progress(self) -> float | None:
"""Fraction complete in [0, 1], or None when total is unknown."""
if self.total <= 0:
return None
return max(0.0, min(1.0, self.done / self.total))
def eta_s(self) -> float | None:
"""Rough seconds-to-completion from elapsed time and progress."""
p = self.progress()
if not p or p <= 0 or self.state != STATE_RUNNING:
return None
elapsed = self.updated_at - self.started_at
if elapsed <= 0:
return None
return max(0.0, elapsed * (1.0 - p) / p)
def snapshot(self, include_result: bool = False) -> dict[str, Any]:
"""JSON-safe view for the WS layer. ``include_result`` embeds the payload."""
data: dict[str, Any] = {
"id": self.id,
"entry_id": self.entry_id,
"kind": self.kind,
"label": self.label,
"label_key": self.label_key,
"label_params": self.label_params,
"state": self.state,
"done": self.done,
"total": self.total,
"progress": self.progress(),
"eta_s": self.eta_s(),
"started_at": self.started_at,
"updated_at": self.updated_at,
"finished_at": self.finished_at,
"error": self.error,
"has_result": self.result is not None,
}
if include_result:
data["result"] = self.result
return data
class TaskRegistry:
"""Holds active + recently-finished tasks and notifies listeners on change."""
def __init__(self) -> None:
self._tasks: OrderedDict[str, Task] = OrderedDict()
self._listeners: set[Callable[[dict[str, Any]], None]] = set()
# Raw asyncio Task handles linked by ws_api so cancel_entry_tasks can
# actually inject CancelledError, not just set the polling flag.
self._asyncio_tasks: dict[str, asyncio.Task[Any]] = {}
# -- listeners -----------------------------------------------------------
def add_listener(self, cb: Callable[[dict[str, Any]], None]) -> Callable[[], None]:
"""Register a change callback; returns an unsubscribe function."""
self._listeners.add(cb)
return lambda: self._listeners.discard(cb)
def _notify(self, task: Task) -> None:
snap = task.snapshot()
for cb in list(self._listeners):
try:
cb(snap)
except Exception: # pylint: disable=broad-exception-caught
logging.getLogger(__name__).debug("Task registry listener error", exc_info=True)
# -- lifecycle -----------------------------------------------------------
def create(
self,
entry_id: str,
kind: str,
label: str,
total: int = 0,
*,
label_key: str | None = None,
label_params: dict[str, Any] | None = None,
) -> Task:
task = Task(
id=uuid.uuid4().hex[:12],
entry_id=entry_id,
kind=kind,
label=label,
label_key=label_key,
label_params=dict(label_params) if label_params else {},
total=max(0, int(total or 0)),
)
self._tasks[task.id] = task
self._notify(task)
self._evict()
return task
def link_asyncio_task(self, task_id: str, asyncio_task: asyncio.Task[Any]) -> None:
"""Associate the raw asyncio Task with a registry entry.
Call this immediately after hass.async_create_task() so that
cancel_entry_tasks() can inject CancelledError rather than only
setting the polling flag.
"""
self._asyncio_tasks[task_id] = asyncio_task
def update(
self,
task: Task,
*,
done: int | None = None,
total: int | None = None,
label: str | None = None,
label_key: str | None = None,
label_params: dict[str, Any] | None = None,
) -> None:
if done is not None:
task.done = done
if total is not None:
task.total = total
if label is not None:
task.label = label
# A supplied label_key replaces the localized label; passing label without
# label_key (legacy callers) clears any stale key so the fallback shows.
if label_key is not None or label is not None:
task.label_key = label_key
task.label_params = dict(label_params) if label_params else {}
task.updated_at = dt_util.now().timestamp()
self._notify(task)
def finish(
self,
task: Task,
*,
state: str = STATE_DONE,
result: Any = None,
error: str | None = None,
) -> None:
# A cancellation that races a late-arriving normal completion must win:
# once STATE_CANCELLED is set, no subsequent finish() call can downgrade it.
if task.state == STATE_CANCELLED and state != STATE_CANCELLED:
return
self._asyncio_tasks.pop(task.id, None)
task.state = state
task.error = error
if result is not None:
task.result = result
task.finished_at = task.updated_at = dt_util.now().timestamp()
self._notify(task)
self._evict()
def cancel(self, task_id: str) -> bool:
"""Request cancellation of a running task. Consumers poll
:attr:`Task.cancel_requested` between chunks. Returns True if a running
task was flagged."""
task = self._tasks.get(task_id)
if task is not None and task.state == STATE_RUNNING:
task._cancelled = True # noqa: SLF001 - registry owns the flag
return True
return False
def cancel_entry_tasks(self, entry_id: str) -> list[asyncio.Task[Any]]:
"""Cancel all running tasks for an entry and mark them finished.
Sets the polling flag and — when a raw asyncio Task was linked via
link_asyncio_task — injects CancelledError so the coroutine stops
promptly rather than waiting for the next cancel_requested poll.
Returns the list of raw asyncio Tasks that were cancelled so the caller
can await their completion (ensuring finally blocks run and locks are
released) before tearing down shared state.
Called during async_unload_entry.
"""
cancelled: list[asyncio.Task[Any]] = []
for task in list(self._tasks.values()):
if task.entry_id == entry_id and task.state == STATE_RUNNING:
task._cancelled = True # noqa: SLF001
raw = self._asyncio_tasks.pop(task.id, None)
if raw is not None and not raw.done():
raw.cancel()
cancelled.append(raw)
self.finish(task, state=STATE_CANCELLED)
return cancelled
# -- reads ---------------------------------------------------------------
def get(self, task_id: str) -> Task | None:
return self._tasks.get(task_id)
def snapshot(self, entry_id: str | None = None) -> list[dict[str, Any]]:
return [
t.snapshot()
for t in self._tasks.values()
if entry_id is None or t.entry_id == entry_id
]
def _evict(self) -> None:
# Evict per entry_id so that a busy entry cannot displace another entry's
# finished results before the panel reads them via get_task_result.
by_entry: dict[str, list[Task]] = {}
for t in self._tasks.values():
if t.state != STATE_RUNNING:
by_entry.setdefault(t.entry_id, []).append(t)
for tasks in by_entry.values():
tasks.sort(key=lambda t: t.finished_at or 0.0)
while len(tasks) > _MAX_FINISHED:
self._tasks.pop(tasks.pop(0).id, None)
def get_registry(hass: HomeAssistant) -> TaskRegistry:
"""Get (or lazily create) the per-hass task registry."""
reg = hass.data.get(_REGISTRY_KEY)
if not isinstance(reg, TaskRegistry):
reg = TaskRegistry()
hass.data[_REGISTRY_KEY] = reg
return reg