643 lines
32 KiB
Python
643 lines
32 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/>.
|
|
"""Community-store bridge: gating + provenance + import/share/catalog orchestration.
|
|
|
|
Pure/near-pure glue between ``store_client`` (network) and ``profile_store`` (local),
|
|
plus the integration-wide account/online flag in ``store_account``. The GitHub
|
|
connection and the online-features switch are device-agnostic (one per HA install);
|
|
brand/model stay per-device. Nothing here runs unless online features are enabled.
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
import math
|
|
from collections.abc import Callable
|
|
from typing import Any
|
|
|
|
from homeassistant.core import HomeAssistant
|
|
|
|
from . import store_account
|
|
from .const import QC_EDITED, QC_MANUAL, QC_RECORDING
|
|
from .store_client import STORE_UNREACHABLE, device_id, get_client, profile_id, trace_hash
|
|
|
|
_LOGGER = logging.getLogger(__name__)
|
|
|
|
|
|
def online_features_enabled(hass: HomeAssistant) -> bool:
|
|
"""True when online store features are enabled integration-wide (default off)."""
|
|
return store_account.online_enabled(hass)
|
|
|
|
|
|
# The community catalog only knows washer/dryer/dishwasher/washer_dryer; HA's
|
|
# washing_machine device type maps to washer. Keep this in sync with the panel's
|
|
# _storeApplianceType() so search, create and share all resolve the same deviceId.
|
|
_STORE_APPLIANCE_TYPE = {"washing_machine": "washer"}
|
|
|
|
|
|
def store_appliance_type(device_type: str) -> str:
|
|
return _STORE_APPLIANCE_TYPE.get(device_type, device_type)
|
|
|
|
|
|
def derive_qc(cycle: dict[str, Any]) -> int:
|
|
"""Derive the obfuscated provenance code for a cycle being uploaded.
|
|
|
|
QC_RECORDING - a recorder capture. Deliberately takes precedence over ``edited``
|
|
(see test_store_provenance.test_recorder_precedence_over_edited): a
|
|
trimmed recording is still classed as a recording, since it began as
|
|
a clean manual capture.
|
|
QC_EDITED - trimmed/edited from a detected cycle.
|
|
QC_MANUAL - a plain detected cycle the user flagged golden by hand.
|
|
Never raises.
|
|
"""
|
|
meta = cycle.get("meta") if isinstance(cycle.get("meta"), dict) else {}
|
|
if meta.get("source") == "recorder" or "original_samples" in meta:
|
|
return QC_RECORDING
|
|
if meta.get("edited"):
|
|
return QC_EDITED
|
|
return QC_MANUAL
|
|
|
|
|
|
def _downsample(points: list[list[float]], max_n: int = 10000) -> list[list[float]]:
|
|
"""Downsample a power trace to at most max_n points using LTTB.
|
|
|
|
LTTB (Largest Triangle Three Buckets) selects the sample in each bucket that
|
|
maximises the triangle area formed by the previously-selected point and the
|
|
centroid of the next bucket. This preserves peaks and troughs (heater pulses,
|
|
pump-out spikes, spin transients) that nearest-index selection can silently drop
|
|
when the step size straddles a narrow transient.
|
|
"""
|
|
n = len(points)
|
|
if n <= max_n:
|
|
return [[float(p[0]), float(p[1])] for p in points]
|
|
if max_n <= 2:
|
|
return [[float(points[0][0]), float(points[0][1])],
|
|
[float(points[-1][0]), float(points[-1][1])]]
|
|
|
|
sampled: list[list[float]] = [[float(points[0][0]), float(points[0][1])]]
|
|
bucket_count = max_n - 2
|
|
bucket_size = (n - 2) / bucket_count
|
|
prev_idx = 0
|
|
|
|
for i in range(bucket_count):
|
|
# Current bucket [a, b)
|
|
a = int(i * bucket_size) + 1
|
|
b = min(int((i + 1) * bucket_size) + 1, n - 1)
|
|
# Next bucket centroid (triangle's third vertex)
|
|
c = b
|
|
d = min(int((i + 2) * bucket_size) + 1, n - 1)
|
|
cnt = d - c
|
|
if cnt > 0:
|
|
avg_x = sum(points[j][0] for j in range(c, d)) / cnt
|
|
avg_y = sum(points[j][1] for j in range(c, d)) / cnt
|
|
else:
|
|
avg_x, avg_y = float(points[-1][0]), float(points[-1][1])
|
|
# Select point in [a, b) with the largest triangle area
|
|
prev = points[prev_idx]
|
|
max_area = -1.0
|
|
max_idx = a
|
|
for j in range(a, b):
|
|
area = abs(
|
|
(prev[0] - avg_x) * (points[j][1] - prev[1]) -
|
|
(prev[0] - points[j][0]) * (avg_y - prev[1])
|
|
) * 0.5
|
|
if area > max_area:
|
|
max_area = area
|
|
max_idx = j
|
|
sampled.append([float(points[max_idx][0]), float(points[max_idx][1])])
|
|
prev_idx = max_idx
|
|
|
|
sampled.append([float(points[-1][0]), float(points[-1][1])])
|
|
return sampled
|
|
|
|
|
|
# Points kept per browse-row trace: the panel's sparkline is 120 px wide.
|
|
_BROWSE_TRACE_POINTS = 200
|
|
|
|
|
|
def _browse_rows(cycles: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
|
"""Browse copies of store cycle rows: no ``importable``, trace downsampled to
|
|
``_BROWSE_TRACE_POINTS`` (see ``StoreBridge.get_cycles``). Executor-safe; never
|
|
raises (a non-numeric point is dropped, store data is untrusted)."""
|
|
out: list[dict[str, Any]] = []
|
|
for cyc in cycles:
|
|
if not isinstance(cyc, dict):
|
|
continue
|
|
row = {k: v for k, v in cyc.items() if k != "importable"}
|
|
trace = cyc.get("trace")
|
|
if isinstance(trace, dict) and isinstance(trace.get("points"), list):
|
|
pts: list[list[float]] = []
|
|
for p in trace["points"]:
|
|
try:
|
|
o, w = float(p[0]), float(p[1])
|
|
except (TypeError, ValueError, IndexError, KeyError, OverflowError):
|
|
continue
|
|
if math.isfinite(o) and math.isfinite(w):
|
|
pts.append([o, w])
|
|
row["trace"] = {**trace, "points": _downsample(pts, _BROWSE_TRACE_POINTS)}
|
|
out.append(row)
|
|
return out
|
|
|
|
|
|
def _cycle_upload_stats(cyc: dict[str, Any], pts: list[list[float]]) -> dict[str, Any]:
|
|
"""Build the community-upload stats for a cycle from its stored metadata + trace.
|
|
|
|
``energy_wh`` is emitted only when it is a known positive value: an older cycle
|
|
or a recording without energy data has no meaningful figure, and sending 0 would
|
|
drag the store's per-program energy average downward. Absent-when-unknown lets the
|
|
aggregate ignore it instead. Shared by share_cycle and share_device so both paths
|
|
serialize a cycle identically.
|
|
"""
|
|
vals = [float(p[1]) for p in pts]
|
|
stats: dict[str, Any] = {
|
|
"duration": float(cyc.get("duration") or (pts[-1][0] - pts[0][0])),
|
|
"peak_w": max(vals) if vals else 0.0,
|
|
"mean_w": (sum(vals) / len(vals)) if vals else 0.0,
|
|
"signature": cyc.get("signature") if isinstance(cyc.get("signature"), dict) else {},
|
|
}
|
|
try:
|
|
energy = float(cyc.get("energy_wh"))
|
|
except (TypeError, ValueError, OverflowError):
|
|
energy = 0.0
|
|
if energy > 0:
|
|
stats["energy_wh"] = energy
|
|
return stats
|
|
|
|
|
|
class StoreBridge:
|
|
"""Orchestrates store browse/import/share/catalog against a ProfileStore.
|
|
|
|
All methods no-op-safe: they return an ``{"error": ...}`` marker rather than raising.
|
|
Callers must gate on ``online_features_enabled`` first. The account/online flag are
|
|
global (via ``store_account``); import/share target this bridge's ProfileStore.
|
|
"""
|
|
|
|
def __init__(self, hass: Any, profile_store: Any) -> None:
|
|
self._hass = hass
|
|
self._ps = profile_store
|
|
# One client for the whole install, not one per appliance: the catalog it reads is
|
|
# public and device-agnostic, so a per-bridge client made an N-appliance install
|
|
# issue N cold-cache copies of the same brand/device queries. See get_client.
|
|
self._client = get_client(hass)
|
|
|
|
def _fire_download_telemetry(self, cycle_ids: list[str]) -> None:
|
|
"""Fire best-effort download/adoption telemetry as a detached background task.
|
|
|
|
Awaiting these store round-trips inline would add their HTTP latency (up to the
|
|
15s per-request timeout when the store is slow/unreachable) to the user-facing
|
|
adopt response. Their outcome does not affect the result, and both bump_* calls
|
|
swallow their own errors, so a detached task can never surface an exception.
|
|
"""
|
|
async def _bump() -> None:
|
|
if cycle_ids:
|
|
await self._client.bump_downloads(cycle_ids)
|
|
await self._client.bump_analytics("downloads", 1)
|
|
|
|
self._hass.async_create_background_task(_bump(), "washdata_store_telemetry")
|
|
|
|
# ── account / status (global) ───────────────────────────────────────────────
|
|
|
|
def status(self) -> dict[str, Any]:
|
|
return {"enabled": store_account.online_enabled(self._hass), **store_account.get_identity(self._hass)}
|
|
|
|
async def connect(self, refresh_token: str, uid: str, name: str | None) -> dict[str, Any]:
|
|
"""Validate the refresh token by exchanging it once, then persist the account.
|
|
|
|
The stored uid is the one the token endpoint returned for that exchange, never
|
|
the caller's ``uid`` (audit STORE-18): a mismatch made every later write 403
|
|
under the store rules' ``uid == request.auth.uid`` checks. ``uid`` is kept in
|
|
the signature for the WS contract; it is only compared, for a log line.
|
|
"""
|
|
# ensure_id_token writes the client-wide _last_error slot on failure, so it
|
|
# takes the same lock the upload paths use: otherwise a sign-in failing here
|
|
# could overwrite the reason a concurrent share is about to report.
|
|
async with self._client.write_lock:
|
|
token = await self._client.ensure_id_token(refresh_token)
|
|
verified = self._client.verified_uid(refresh_token) if token else None
|
|
if not token or not verified:
|
|
return {"error": "token_invalid"}
|
|
if uid and uid != verified:
|
|
_LOGGER.warning("Store connect: ignoring a uid that does not match the token's account")
|
|
await store_account.async_set_account(
|
|
self._hass, {"refresh_token": refresh_token, "uid": verified, "name": name},
|
|
)
|
|
return store_account.get_identity(self._hass)
|
|
|
|
def _uploader_name(self, acct: dict[str, Any]) -> str | None:
|
|
"""The name published on a shared cycle: the account's display name only when
|
|
the user opted in (``share_name``), else None (audit STORE-14)."""
|
|
if not store_account.get_pref(self._hass, "share_name"):
|
|
return None
|
|
name = acct.get("name")
|
|
return name if isinstance(name, str) and name.strip() else None
|
|
|
|
async def disconnect(self) -> dict[str, Any]:
|
|
await store_account.async_clear_account(self._hass)
|
|
return {"connected": False}
|
|
|
|
# ── catalog browse (reads) ───────────────────────────────────────────────────
|
|
|
|
async def list_brands(self, query: str | None = None, *, include_pending: bool = True) -> list[dict[str, Any]]:
|
|
return await self._client.list_brands(query, include_pending=include_pending)
|
|
|
|
async def search_devices(
|
|
self, brand: str | None, appliance_type: str | None,
|
|
model_query: str | None = None, *, include_pending: bool = False,
|
|
) -> list[dict[str, Any]]:
|
|
return await self._client.search_devices(
|
|
brand, appliance_type, model_query=model_query, include_pending=include_pending,
|
|
)
|
|
|
|
async def catalog_entry(self, brand: str, model: str, appliance_type: str) -> dict[str, Any]:
|
|
"""The catalog identity of one appliance -- its brand + device documents, resolved
|
|
by deterministic id, for the settings form's status badges.
|
|
|
|
The badges used to be a by-product of the pickers' full brand/device *lists*, which
|
|
meant simply opening the Settings tab downloaded the catalog. Two point reads answer
|
|
the same question, so the lists are now only fetched when the user actually opens a
|
|
picker. Maps the HA device type to the catalog type first.
|
|
"""
|
|
return await self._client.catalog_entry(brand, model, store_appliance_type(appliance_type))
|
|
|
|
def refresh_catalog(self) -> dict[str, Any]:
|
|
"""Drop the cached catalog so the next read is fresh (panel "refresh" action)."""
|
|
self._client.refresh_catalog()
|
|
return {"ok": True}
|
|
|
|
async def get_profiles(
|
|
self, device_id: str, *, include_pending: bool = True,
|
|
) -> list[dict[str, Any]] | None:
|
|
"""Shared programs for a catalog appliance. Pending-inclusive, like the device
|
|
list this is opened from and like get_cycles below; passed explicitly so the
|
|
browse cannot silently drift back to approved-only (which showed nothing).
|
|
``None`` = the store could not be reached, not "no programs" (audit STORE-09)."""
|
|
return await self._client.get_profiles(device_id, include_pending=include_pending)
|
|
|
|
async def device_profiles(self, brand: str, model: str, appliance_type: str) -> dict[str, Any]:
|
|
"""Profiles for the appliance identified by brand/model/type (for the Share
|
|
dialog's profile picker). Maps the HA device type to the catalog type first."""
|
|
return await self._client.device_profiles(brand, model, store_appliance_type(appliance_type))
|
|
|
|
async def get_cycles(self, profile_id: str) -> list[dict[str, Any]] | None:
|
|
"""A program's shared cycles for the browse list; ``None`` = store unreachable.
|
|
|
|
Slimmed for the wire (audit STORE-20): the panel draws each trace as a 120 px
|
|
sparkline, yet every row carried the full trace twice (``trace.points`` and
|
|
``importable``, up to 7k points each, 50 rows). Import re-fetches the cycle by
|
|
id, so the browse rows drop ``importable`` and carry a peak-preserving
|
|
downsample. The download path reads the client directly and keeps full traces.
|
|
"""
|
|
cycles = await self._client.get_cycles(profile_id)
|
|
if not cycles:
|
|
return cycles
|
|
return await self._hass.async_add_executor_job(_browse_rows, cycles)
|
|
|
|
|
|
# ── community actions (authed writes) ────────────────────────────────────────
|
|
|
|
async def confirm_device(self, device_id: str) -> dict[str, Any]:
|
|
acct = store_account.get_account(self._hass)
|
|
if not acct.get("refresh_token"):
|
|
return {"error": "not_connected"}
|
|
async with self._client.write_lock: # writes _last_error; see connect()
|
|
res = await self._client.confirm_device(acct["refresh_token"], acct.get("uid", ""), device_id)
|
|
return res if res else {"error": "confirm_failed"}
|
|
|
|
async def rate_device(self, device_id: str, rating: int) -> dict[str, Any]:
|
|
acct = store_account.get_account(self._hass)
|
|
if not acct.get("refresh_token"):
|
|
return {"error": "not_connected"}
|
|
async with self._client.write_lock: # writes _last_error; see connect()
|
|
ok = await self._client.rate_device(acct["refresh_token"], acct.get("uid", ""), device_id, rating)
|
|
return {"ok": True} if ok else {"error": "rate_failed"}
|
|
|
|
# ── import / share (target this device's ProfileStore) ───────────────────────
|
|
|
|
async def import_cycle(
|
|
self, cycle_id: str, target_profile: str | None = None, new_profile_name: str | None = None
|
|
) -> dict[str, Any]:
|
|
cyc = await self._client.get_cycle(cycle_id)
|
|
if not cyc:
|
|
return {"error": "not_found"}
|
|
pts = cyc.get("importable")
|
|
if not pts:
|
|
return {"error": "unsupported_schema"}
|
|
# The name comes from the caller (localized) or the store's program label;
|
|
# never fall back to an inline English string. Require a non-empty name.
|
|
raw_profile = new_profile_name or target_profile or cyc.get("program_lc")
|
|
profile = raw_profile.strip() if isinstance(raw_profile, str) else ""
|
|
if not profile:
|
|
return {"error": "profile_name_required"}
|
|
meta = {
|
|
"store_cycle_id": cyc.get("id"),
|
|
"store_uploaded_at": cyc.get("createdAt"),
|
|
"sampling_interval": (cyc.get("trace") or {}).get("sampleIntervalSec"),
|
|
"community": True,
|
|
}
|
|
# Too short / gappy / implausible, or a copy of a trace already stored
|
|
# under any name (audit STORE-03/05): refused before it shapes a profile.
|
|
verdict = self._ps.reference_import_verdict(pts, meta)
|
|
if verdict != "ok":
|
|
return {"error": "invalid_trace" if verdict == "invalid" else verdict}
|
|
local_id = await self._ps.add_reference_cycle(profile, pts, meta)
|
|
if not local_id: # trace failed validation in add_reference_cycle
|
|
return {"error": "invalid_trace"}
|
|
# Credit the download on the source store cycle + record one community-wide
|
|
# adoption for the store's usage dashboard (the real "someone used it" metric).
|
|
# Fired in the background so store latency never delays the adopt response.
|
|
self._fire_download_telemetry([cyc["id"]] if cyc.get("id") else [])
|
|
return {"profile": profile, "cycle_id": local_id}
|
|
|
|
async def share_cycle(
|
|
self, local_cycle_id: str, program: str, brand: str, model: str, appliance_type: str,
|
|
sample_interval_sec: float = 0.0, description: str = "",
|
|
) -> dict[str, Any]:
|
|
acct = store_account.get_account(self._hass)
|
|
if not acct.get("refresh_token"):
|
|
return {"error": "not_connected"}
|
|
pts = self._ps.get_cycle_power_data(local_cycle_id)
|
|
if not pts:
|
|
return {"error": "cycle_not_found"}
|
|
# Look up metadata in BOTH real past_cycles and imported reference_cycles
|
|
# (same by_id behavior as share_device) so an imported/reference cycle keeps
|
|
# its stored duration/energy/signature instead of falling back to trace-derived.
|
|
by_id = {
|
|
c.get("id"): c
|
|
for c in (list(self._ps.get_reference_cycles()) + list(self._ps.get_past_cycles()))
|
|
}
|
|
cyc = by_id.get(local_cycle_id, {})
|
|
stats = _cycle_upload_stats(cyc, pts)
|
|
meta = {
|
|
"applianceType": store_appliance_type(appliance_type), "brand": brand, "model": model, "program": program,
|
|
"sampleIntervalSec": float(sample_interval_sec or cyc.get("sampling_interval") or 0.0),
|
|
"description": description,
|
|
}
|
|
# LTTB downsampling is O(N) pure Python; offload it so a long trace never
|
|
# blocks the event loop while a user shares a cycle.
|
|
downsampled = await self._hass.async_add_executor_job(
|
|
_downsample, [[p[0], p[1]] for p in pts]
|
|
)
|
|
# The client is shared install-wide, and last_error() is a single slot, so the
|
|
# write and the read that interprets it have to be one critical section.
|
|
async with self._client.write_lock:
|
|
new_id = await self._client.upload_reference_cycle(
|
|
acct["refresh_token"], acct.get("uid", ""), self._uploader_name(acct),
|
|
meta, downsampled, stats, derive_qc(cyc),
|
|
)
|
|
if not new_id:
|
|
return {"error": "upload_failed", "detail": self._client.last_error()}
|
|
return {"store_cycle_id": new_id}
|
|
|
|
async def share_device(
|
|
self, brand: str, model: str, appliance_type: str, items: list[dict[str, Any]],
|
|
include_phases: list[str] | None = None, settings: dict[str, Any] | None = None,
|
|
) -> dict[str, Any]:
|
|
"""Share a device bundle. ``items`` = ``[{local_cycle_id, program}]`` (the
|
|
panel's tree selection). Resolves each local cycle's trace + stats and uploads
|
|
the whole set via ``upload_device_bundle``. For each program named in
|
|
``include_phases`` that has local phase ranges, the phase map (seconds) is
|
|
attached to that program's bundle items so it lands on the store profile doc.
|
|
Returns ``{ok, cycle_ids, created, duplicates, errors}`` or ``{error}``.
|
|
"""
|
|
acct = store_account.get_account(self._hass)
|
|
if not acct.get("refresh_token"):
|
|
return {"error": "not_connected"}
|
|
store_type = store_appliance_type(appliance_type)
|
|
want_phases = {str(p).strip() for p in (include_phases or []) if str(p).strip()}
|
|
by_id = {
|
|
c.get("id"): c
|
|
for c in (list(self._ps.get_past_cycles()) + list(self._ps.get_reference_cycles()))
|
|
}
|
|
bundle_items: list[dict[str, Any]] = []
|
|
for it in items or []:
|
|
cid = it.get("local_cycle_id")
|
|
program = str(it.get("program") or "").strip()
|
|
if not cid or not program:
|
|
continue
|
|
pts = self._ps.get_cycle_power_data(cid)
|
|
if not pts:
|
|
continue
|
|
cyc = by_id.get(cid, {})
|
|
# Offload the O(N) LTTB pass per cycle so a large bundle never stalls
|
|
# the event loop while sharing.
|
|
downsampled = await self._hass.async_add_executor_job(
|
|
_downsample, [[p[0], p[1]] for p in pts]
|
|
)
|
|
bundle_items.append({
|
|
"program": program,
|
|
"points": downsampled,
|
|
"stats": _cycle_upload_stats(cyc, pts),
|
|
"qc": derive_qc(cyc),
|
|
"sampleIntervalSec": float(cyc.get("sampling_interval") or 0.0),
|
|
})
|
|
if not bundle_items:
|
|
return {"error": "nothing_to_share"}
|
|
# Stage 2: attach each requested program's phase map to its items. The store
|
|
# cycle id is deterministic (trace_hash), so phaseSourceCycleId can point at
|
|
# the program's first shared cycle for the Stage-4 web editor.
|
|
d_id = device_id(store_type, brand, model)
|
|
for program in want_phases:
|
|
ranges = self._ps.get_profile_phase_ranges(program)
|
|
if not ranges:
|
|
continue
|
|
prog_items = [b for b in bundle_items if b["program"] == program]
|
|
if not prog_items:
|
|
continue
|
|
phases = [{"name": r["name"], "start": r["start"], "end": r["end"]} for r in ranges]
|
|
source_cid = trace_hash(profile_id(d_id, program), prog_items[0]["points"])
|
|
for b in prog_items:
|
|
b["phases"] = phases
|
|
b["phaseSourceCycleId"] = source_cid
|
|
device_meta: dict[str, Any] = {"applianceType": store_type, "brand": brand, "model": model}
|
|
# Stage 3: attach the device's recognition/matching settings (already filtered
|
|
# to the allow-list by the WS layer, which owns entry.options).
|
|
if isinstance(settings, dict) and settings:
|
|
device_meta["settings"] = dict(settings)
|
|
async with self._client.write_lock: # see share_cycle
|
|
res = await self._client.upload_device_bundle(
|
|
acct["refresh_token"], acct.get("uid", ""), self._uploader_name(acct), device_meta, bundle_items,
|
|
)
|
|
# Return the raw bundle result ({ok, cycle_ids, errors}) so the caller can
|
|
# tell a partial upload (some cycle_ids present) from a total failure.
|
|
# Only a pre-flight gate short-circuits with an {"error": ...} marker above.
|
|
if not res.get("ok") and not res.get("cycle_ids"):
|
|
res = {**res, "detail": self._client.last_error()}
|
|
return res
|
|
|
|
async def download_device(
|
|
self,
|
|
device_id_: str,
|
|
device_type: str = "",
|
|
*,
|
|
progress: Callable[[int, int], None] | None = None,
|
|
should_cancel: Callable[[], bool] | None = None,
|
|
) -> dict[str, Any]:
|
|
"""Adopt a whole-device bundle: for each downloaded profile, import its
|
|
reference cycles into ``reference_cycles`` (merge/upsert; real past_cycles are
|
|
never touched) and, when the profile carries a phase map, replace the local
|
|
profile's phase ranges + reconcile any unknown phase labels into the catalog.
|
|
Returns ``{profiles_adopted, cycles_imported, phases_applied, settings}`` where
|
|
``settings`` is the bundle's device settings map (the WS layer applies it to
|
|
entry.options only when the user opts in; the bridge never touches options).
|
|
|
|
Idempotent: a store cycle already imported locally (``meta.source ==
|
|
"store:<id>"``) is skipped, so re-downloading the same device does not
|
|
accumulate duplicate reference cycles.
|
|
|
|
A store that could not be read is reported as ``{"error":
|
|
"store_unreachable"}``, never as an empty adopt (audit STORE-09: offline used
|
|
to read "Nothing new - already on your device"). When only some programs'
|
|
cycles could not be fetched, the result carries ``partial`` +
|
|
``failed_profiles``; if nothing new was imported at all it is that error.
|
|
"""
|
|
bundle = await self._client.get_device_bundle(device_id_)
|
|
if bundle.get("error"):
|
|
return {"error": str(bundle["error"])}
|
|
already = {
|
|
str((c.get("meta") or {}).get("source") or "")
|
|
for c in self._ps.get_reference_cycles()
|
|
}
|
|
# Hashed once for the whole bundle; each import adds its own, so two
|
|
# programs carrying the same recording keep only the first (STORE-05).
|
|
known_hashes = self._ps.stored_trace_hashes()
|
|
profiles_adopted = 0
|
|
cycles_imported = 0
|
|
cycles_skipped = 0
|
|
phases_applied = 0
|
|
imported_store_ids: list[str] = []
|
|
# Batched (audit STORE-10): every cycle used to rebuild its envelope and
|
|
# rewrite the whole store - 40 saves for a 40-cycle bundle, 7 s vs 1.5 s.
|
|
# Now one rebuild per touched program and one save at the end.
|
|
touched: list[str] = []
|
|
profiles_list = [p for p in (bundle.get("profiles", []) or []) if isinstance(p, dict)]
|
|
failed_profiles = sum(1 for p in profiles_list if p.get("cycles_unavailable"))
|
|
total = sum(len(p.get("cycles") or []) for p in profiles_list)
|
|
seen = 0
|
|
cancelled = False
|
|
for prof in profiles_list:
|
|
if cancelled:
|
|
break
|
|
program = str(prof.get("program") or prof.get("program_lc") or "").strip()
|
|
if not program:
|
|
continue
|
|
adopted_any = False
|
|
for cyc in prof.get("cycles", []) or []:
|
|
if should_cancel is not None and should_cancel():
|
|
cancelled = True
|
|
break
|
|
if progress is not None:
|
|
progress(seen, total)
|
|
seen += 1
|
|
pts = cyc.get("importable")
|
|
if not pts:
|
|
continue
|
|
store_cid = cyc.get("id")
|
|
if store_cid and f"store:{store_cid}" in already:
|
|
continue # already imported on a previous download
|
|
local_id = self._ps._add_reference_cycle_nosave(program, pts, { # noqa: SLF001
|
|
"store_cycle_id": store_cid,
|
|
"store_uploaded_at": cyc.get("createdAt"),
|
|
"sampling_interval": (cyc.get("trace") or {}).get("sampleIntervalSec"),
|
|
"community": True,
|
|
}, known_hashes=known_hashes)
|
|
if local_id:
|
|
if program not in touched:
|
|
touched.append(program)
|
|
cycles_imported += 1
|
|
adopted_any = True
|
|
if store_cid:
|
|
imported_store_ids.append(store_cid)
|
|
else:
|
|
cycles_skipped += 1
|
|
if adopted_any:
|
|
profiles_adopted += 1
|
|
# Stage 2: apply the bundled phase map (replace) + reconcile labels. Never
|
|
# raises; a bad/overlapping range set is skipped rather than failing adopt.
|
|
# Run whenever the profile carries phases -- not gated on new cycles -- so a
|
|
# re-download with no new cycles still reconciles updated phase ranges.
|
|
if prof.get("phases") and await self._apply_phases(program, prof.get("phases"), device_type):
|
|
phases_applied += 1
|
|
# Credit a download on every reference cycle actually adopted this run (each
|
|
# underlying object, not just one). Skipped/duplicate cycles aren't re-counted,
|
|
# so re-downloading the same device doesn't inflate the counters. Also record one
|
|
# community-wide "download" (adoption) for the store's usage dashboard -- the real
|
|
# metric of how many people actually pulled this into their integration. Fired in
|
|
# the background so store latency never delays the adopt-bundle response.
|
|
for program in touched:
|
|
await self._ps.async_rebuild_envelope(program)
|
|
if touched:
|
|
await self._ps.async_save()
|
|
if imported_store_ids:
|
|
self._fire_download_telemetry(imported_store_ids)
|
|
if failed_profiles and not cycles_imported and not cancelled:
|
|
# Nothing new arrived and part of the setup could not be read: a failed
|
|
# download, not "already on your device" (audit STORE-09).
|
|
return {
|
|
"error": STORE_UNREACHABLE, "failed_profiles": failed_profiles,
|
|
"cycles_skipped": cycles_skipped, "phases_applied": phases_applied,
|
|
}
|
|
settings = bundle.get("settings") if isinstance(bundle.get("settings"), dict) else {}
|
|
return {
|
|
"profiles_adopted": profiles_adopted,
|
|
"cycles_imported": cycles_imported,
|
|
# Refused by the quality bar or as a duplicate (audit STORE-03/05).
|
|
"cycles_skipped": cycles_skipped,
|
|
**({"cancelled": True} if cancelled else {}),
|
|
**({"partial": True, "failed_profiles": failed_profiles} if failed_profiles else {}),
|
|
"phases_applied": phases_applied,
|
|
"settings": settings,
|
|
}
|
|
|
|
async def _apply_phases(
|
|
self, program: str, phases: Any, device_type: str
|
|
) -> bool:
|
|
"""Replace ``program``'s local phase ranges with the bundled set and merge any
|
|
unknown phase labels into the custom-phase catalog. Returns True when a
|
|
non-empty phase map was applied. Never raises."""
|
|
if not isinstance(phases, list) or not phases:
|
|
return False
|
|
ranges: list[dict[str, Any]] = []
|
|
for p in phases:
|
|
if not isinstance(p, dict):
|
|
continue
|
|
name = str(p.get("name", "")).strip()
|
|
try:
|
|
start, end = float(p.get("start", 0)), float(p.get("end", 0))
|
|
except (TypeError, ValueError, OverflowError):
|
|
continue
|
|
if name and end > start:
|
|
ranges.append({"name": name, "start": start, "end": end})
|
|
if not ranges:
|
|
return False
|
|
try:
|
|
await self._ps.async_set_profile_phase_ranges(program, ranges)
|
|
except Exception as exc: # pylint: disable=broad-exception-caught
|
|
_LOGGER.debug("download_device: could not apply phases for %s: %s", program, exc)
|
|
return False
|
|
# Reconcile labels into the catalog so they carry a name/description in the UI.
|
|
try:
|
|
known = {str(p.get("name", "")).casefold() for p in self._ps.list_phase_catalog(device_type)}
|
|
for r in ranges:
|
|
if r["name"].casefold() not in known:
|
|
try:
|
|
await self._ps.async_create_custom_phase(device_type, r["name"])
|
|
known.add(r["name"].casefold())
|
|
except Exception: # pylint: disable=broad-exception-caught
|
|
pass # duplicate / invalid label -> skip
|
|
except Exception: # pylint: disable=broad-exception-caught
|
|
pass
|
|
return True
|