Coverage for src/homeassistant_repl/api_objtree.py: 91%

230 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-10-05 13:50 +0000

1"""ApiObjTree - API client mode's read-only `obj[...]`, built entirely from 

2Home Assistant's standard websocket API (get_states, config/*_registry/list) 

3- no ha_repl_server component needed, works against any instance an admin 

4token can reach. Bound to the same name, `obj`, as the live tree (see 

5apirepl.py) so a snippet that only touches `obj` runs unchanged in either 

6mode. Entities only have what that API exposes: state + attributes from 

7get_states, and whatever the entity registry's own partial dict exposes (no 

8live component instance, so no calling methods on it, no fields an 

9integration keeps off the registry/state). 

10 

11Mirrors custom_components/ha_repl_server/objtree.py's shape - __getitem__, 

12.keys()/.values()/.items(), find(), show() all behave the same way modulo the 

13smaller attribute set - but lives in this package rather than that one (a 

14HACS-deployed component with its own packaging boundary), so the handful of 

15small, stable, HA-import-free pieces it needs (the ordered-view mixin, 

16parse_path, the None/empty-collection cleanup) are duplicated here rather than 

17imported across that boundary. Revisit if API client mode sticks and a shared 

18package becomes worth the packaging work. 

19 

20Unlike the live tree, this one is explicitly a *snapshot*: fetching is async 

21(it's a websocket call) but Mapping's `__getitem__`/`__iter__`/`__len__` are 

22not, so refreshing can't happen lazily on access without risking asyncio 

23reentrancy (this is normally used from inside an already-running event loop). 

24Instead the owning REPL loop decides when to call `await cache.refresh()` 

25(e.g. between prompts, when `cache.is_stale()`) - see apirepl.py. `reset_cache()` 

26just marks the cache stale; it does not fetch. 

27""" 

28 

29from __future__ import annotations 

30 

31import re 

32import time 

33from collections.abc import ( 

34 ItemsView, 

35 Iterable, 

36 Iterator, 

37 KeysView, 

38 Mapping, 

39 Sequence, 

40 ValuesView, 

41) 

42from dataclasses import dataclass 

43from datetime import datetime 

44from typing import Any, NamedTuple, Protocol 

45 

46 

47class _ApiClient(Protocol): 

48 """What Cache actually needs from a client - just enough to let tests use 

49 a lightweight fake instead of a real websocket Client.""" 

50 

51 async def call(self, type_: str, **payload: Any) -> Any: ... 

52 

53 

54def parse_path(path: str) -> tuple[str, ...]: 

55 """Split "/integration[/domain[/object_id]]" into 1-3 parts. Duplicated 

56 from custom_components/ha_repl_server/paths.py - see the module 

57 docstring for why.""" 

58 parts = tuple(p for p in path.split("/") if p) 

59 if not parts or len(parts) > 3: 59 ↛ 60line 59 didn't jump to line 60 because the condition on line 59 was never true

60 raise ValueError(f"{path!r} - expected /integration/domain/object_id") 

61 return parts 

62 

63 

64def _as_set(value: str | list[str] | None) -> set[str]: 

65 if value is None: 

66 return set() 

67 return {value} if isinstance(value, str) else set(value) 

68 

69 

70# HA's own slugs (integration/domain/object_id) are always [a-z0-9_] - never any 

71# of these - so a `find()` path containing one is unambiguously a regex, not a 

72# literal path, with no risk of misreading a real path as a pattern. 

73_REGEX_METACHARS = frozenset(".^$*+?{}[]|()\\") 

74 

75 

76def _require_str(value: Any, what: str = "path") -> str: 

77 """A clear TypeError at the boundary beats a confusing one from deep 

78 inside parse_path - duplicated from objtree.py's identical helper, see 

79 that module for the rationale.""" 

80 if not isinstance(value, str): 

81 raise TypeError(f"{what} must be a str, not {type(value).__name__}") 

82 return value 

83 

84 

85class _OrderedView(Iterable[Any]): 

86 """Positional access/slicing on top of a MappingView's existing order - 

87 duplicated from objtree.py's identical mixin, see that module for the 

88 rationale.""" 

89 

90 def __getitem__(self, index: int | slice) -> Any: 

91 return tuple(self)[index] 

92 

93 def __reversed__(self) -> Iterator[Any]: 

94 return reversed(tuple(self)) 

95 

96 

97class OrderedKeysView(_OrderedView, KeysView, Sequence): 

98 pass 

99 

100 

101class OrderedValuesView(_OrderedView, ValuesView, Sequence): 

102 pass 

103 

104 

105class OrderedItemsView(_OrderedView, ItemsView, Sequence): # type: ignore[misc] # ty: ignore[invalid-method-override] 

106 # Sequence.__contains__(value) vs ItemsView.__contains__(item: tuple) is a 

107 # real static Liskov mismatch, but fine at runtime - both just delegate to 

108 # tuple(self).__contains__ via _OrderedView/Iterable, deliberately getting 

109 # both Set and Sequence behavior at once (see the module docstring). 

110 pass 

111 

112 

113@dataclass(frozen=True) 

114class ApiEntity: 

115 """One entity's worth of API client mode data - only what get_states and 

116 the entity registry's partial dict expose, not the live component instance. 

117 

118 Named to match the real Entity's own properties where there's a direct 

119 equivalent - state, state_attributes, name - per 

120 https://developers.home-assistant.io/docs/core/entity/. `name` comes from 

121 get_states' `friendly_name` attribute: that's already HA's own fully 

122 resolved result of computing Entity.name (registry override, device name, 

123 has_entity_name, etc.), published into the state machine for exactly this 

124 kind of external consumer - cheaper and more faithful than recomputing it 

125 from the pieces ourselves. `state_attributes` is everything else in that 

126 attributes dict; broader than the live property of the same name (which 

127 excludes things like unit_of_measurement that come from separate Entity 

128 properties), since the state machine doesn't preserve that distinction. 

129 """ 

130 

131 entity_id: str 

132 platform: str 

133 domain: str 

134 object_id: str 

135 state: str | None 

136 name: str | None 

137 state_attributes: dict[str, Any] 

138 area_id: str | None 

139 labels: frozenset[str] 

140 registry: dict[str, Any] 

141 

142 

143class _Found(NamedTuple): 

144 """Internal only - never returned from find()/find_paths()/find_names(), 

145 just the shared (path, entity) pair each of them projects from 

146 (.entity, .path, and .entity_id respectively) so the filtering logic in 

147 _find() lives in exactly one place. Any attribute other than path/entity 

148 falls through to `entity`, so e.g. found.entity_id works without 

149 found.entity.entity_id.""" 

150 

151 path: str 

152 entity: ApiEntity 

153 

154 def __getattr__(self, name: str) -> Any: 

155 return getattr(self.entity, name) 

156 

157 

158# The two registry JSON fields that are Unix-timestamp floats rather than the 

159# ISO datetime strings get_states already returns - see show()/_clean(). 

160_TIMESTAMP_FIELDS = frozenset({"created_at", "modified_at"}) 

161 

162 

163class Cache: 

164 """Fetches and holds one snapshot of states + registries. Refresh is 

165 driven by the REPL loop between commands - see the module docstring.""" 

166 

167 def __init__(self, client: _ApiClient, ttl: float) -> None: 

168 self.client = client 

169 self.ttl = ttl 

170 self._fetched_at: float | None = None 

171 self.entities: dict[str, ApiEntity] = {} 

172 self.areas: dict[str, dict[str, Any]] = {} 

173 self.labels: dict[str, dict[str, Any]] = {} 

174 

175 def is_stale(self) -> bool: 

176 return ( 

177 self._fetched_at is None or (time.monotonic() - self._fetched_at) > self.ttl 

178 ) 

179 

180 def reset(self) -> None: 

181 """Mark the cache stale. Doesn't fetch - see the module docstring.""" 

182 self._fetched_at = None 

183 

184 async def refresh(self) -> None: 

185 states = {s["entity_id"]: s for s in await self.client.call("get_states")} 

186 registry_entries = await self.client.call("config/entity_registry/list") 

187 devices = { 

188 d["id"]: d for d in await self.client.call("config/device_registry/list") 

189 } 

190 self.areas = { 

191 a["area_id"]: a for a in await self.client.call("config/area_registry/list") 

192 } 

193 self.labels = { 

194 l["label_id"]: l 

195 for l in await self.client.call("config/label_registry/list") 

196 } 

197 

198 entities: dict[str, ApiEntity] = {} 

199 for entry in registry_entries: 

200 entity_id = entry["entity_id"] 

201 domain, object_id = entity_id.split(".", 1) 

202 state = states.get(entity_id, {}) 

203 area_id = entry.get("area_id") 

204 if area_id is None and entry.get("device_id"): 

205 device = devices.get(entry["device_id"]) 

206 if device is not None: 206 ↛ 212line 206 didn't jump to line 212 because the condition on line 206 was always true

207 area_id = device.get("area_id") 

208 # friendly_name is already HA's own fully-resolved Entity.name, 

209 # published for exactly this kind of external consumer - see 

210 # ApiEntity's docstring. Copy the dict before popping: it's 

211 # `states[entity_id]`'s own attributes dict, not ours to mutate. 

212 attributes = dict(state.get("attributes", {})) 

213 name = attributes.pop("friendly_name", None) 

214 entities[entity_id] = ApiEntity( 

215 entity_id=entity_id, 

216 platform=entry["platform"], 

217 domain=domain, 

218 object_id=object_id, 

219 state=state.get("state"), 

220 name=name, 

221 state_attributes=attributes, 

222 area_id=area_id, 

223 labels=frozenset(entry.get("labels", ())), 

224 registry=entry, 

225 ) 

226 self.entities = entities 

227 self._fetched_at = time.monotonic() 

228 

229 

230def _resolve_from_cache( 

231 values: str | list[str] | None, 

232 by_id: dict[str, dict[str, Any]], 

233 kind: str, 

234 id_key: str, 

235) -> set[str]: 

236 """Resolve each of `values` (an id or a display name) to its canonical id. 

237 

238 Raises KeyError on anything that matches neither - a typo'd area/label 

239 name should fail loudly, not just quietly match nothing in find(). 

240 """ 

241 ids: set[str] = set() 

242 for value in _as_set(values): 

243 if value in by_id: 

244 ids.add(value) 

245 continue 

246 match = next( 

247 ( 

248 k 

249 for k, v in by_id.items() 

250 if v.get("name", "").casefold() == value.casefold() 

251 ), 

252 None, 

253 ) 

254 if match is None: 

255 raise KeyError(f"no such {kind}: {value!r}") 

256 ids.add(match) 

257 return ids 

258 

259 

260def _normalize(value: Any) -> Any: 

261 if isinstance(value, dict): 

262 return _clean(value) 

263 return value 

264 

265 

266def _clean(data: dict[str, Any]) -> dict[str, Any]: 

267 """Recursively drop None/empty-collection values and render the known 

268 timestamp fields as local ISO 8601 strings. 

269 

270 Empty *lists* are dropped here too, not just empty sets: labels/aliases 

271 etc. arrive over the wire as JSON lists (JSON has no set type), so a list 

272 is this mode's wire-equivalent of the live tree's empty-set fields - and 

273 keeping both modes behaving the same way is the point. 

274 """ 

275 cleaned: dict[str, Any] = {} 

276 for key, value in data.items(): 

277 if key in _TIMESTAMP_FIELDS and isinstance(value, (int, float)): 

278 value = datetime.fromtimestamp(value).astimezone().isoformat() 

279 else: 

280 value = _normalize(value) 

281 if value is None or value == set() or value == []: 

282 continue 

283 cleaned[key] = value 

284 return cleaned 

285 

286 

287@dataclass(frozen=True) 

288class ApiObjTree(Mapping[str, "ApiEntity | ApiObjTree"]): 

289 """API client mode's view of the object tree - see the module docstring.""" 

290 

291 cache: Cache 

292 integration: str | None = None 

293 domain: str | None = None 

294 

295 def reset_cache(self) -> None: 

296 self.cache.reset() 

297 

298 def __getitem__(self, key: str) -> ApiEntity | ApiObjTree: 

299 _require_str(key, "key") 

300 if ( 

301 "/" not in key 

302 and "." in key 

303 and self.integration is not None 

304 and self.domain is None 

305 ): 

306 # Courtesy: a full entity_id ("domain.object_id") also works 

307 # scoped to just an integration, not only at the root - translate 

308 # to the equivalent "/"-path so the logic below (which already 

309 # enforces the platform match) handles it the same way. 

310 domain, _, object_id = key.partition(".") 

311 key = f"{domain}/{object_id}" 

312 if "/" not in key and self.integration is None: 

313 entity = self.cache.entities.get(key) 

314 if entity is None: 314 ↛ 315line 314 didn't jump to line 315 because the condition on line 314 was never true

315 raise KeyError(key) 

316 return entity 

317 try: 

318 parts = parse_path(key) 

319 except ValueError as err: 

320 raise KeyError(str(err)) from None 

321 scope = tuple(p for p in (self.integration, self.domain) if p is not None) 

322 full = scope + parts 

323 if len(full) > 3: 323 ↛ 324line 323 didn't jump to line 324 because the condition on line 323 was never true

324 raise KeyError(key) 

325 if len(full) < 3: 

326 subtree = ApiObjTree(self.cache, *full) 

327 if not subtree: 

328 raise KeyError(key) 

329 return subtree 

330 integration, domain, object_id = full 

331 entity = self.cache.entities.get(f"{domain}.{object_id}") 

332 if entity is None or entity.platform != integration: 

333 raise KeyError(key) 

334 return entity 

335 

336 def _entries(self, prefix: tuple[str, ...]) -> Iterator[ApiEntity]: 

337 for entity in self.cache.entities.values(): 

338 if len(prefix) >= 1 and entity.platform != prefix[0]: 

339 continue 

340 if len(prefix) >= 2 and entity.domain != prefix[1]: 

341 continue 

342 if len(prefix) >= 3 and entity.object_id != prefix[2]: 342 ↛ 343line 342 didn't jump to line 343 because the condition on line 342 was never true

343 continue 

344 yield entity 

345 

346 def __iter__(self) -> Iterator[str]: 

347 scope = tuple(p for p in (self.integration, self.domain) if p is not None) 

348 children = { 

349 (entity.platform, entity.domain, entity.object_id)[len(scope)] 

350 for entity in self._entries(scope) 

351 } 

352 for child in sorted(children): 

353 yield "/" + child 

354 

355 def __len__(self) -> int: 

356 return sum(1 for _ in self) 

357 

358 def keys(self) -> OrderedKeysView: 

359 return OrderedKeysView(self) 

360 

361 def values(self) -> OrderedValuesView: 

362 return OrderedValuesView(self) 

363 

364 def items(self) -> OrderedItemsView: 

365 return OrderedItemsView(self) 

366 

367 def _find( 

368 self, 

369 path: str, 

370 *, 

371 platform: str | list[str] | None, 

372 domain: str | list[str] | None, 

373 area: str | list[str] | None, 

374 label: str | list[str] | None, 

375 ) -> Iterator[_Found]: 

376 """The real search, shared by find()/find_paths()/find_names() - each 

377 just projects a different field from the (path, entity) pairs this 

378 yields. 

379 

380 `path` is normally an exact /integration/domain/object_id prefix, the 

381 same as indexing - but a path containing a regex metacharacter (e.g. 

382 "/mqtt/binary_sensor/barn.*") is matched as a regular expression 

383 against each candidate's full path instead, anchored at the start 

384 (so it behaves like a prefix match unless you anchor the end 

385 yourself with `$`). That trades the usual early narrowing for a scan 

386 of this view's whole subtree, filtered by the pattern. 

387 

388 `domain` matches the HA domain (e.g. "light") across 

389 integrations, the same way `platform` does for the owning 

390 integration; `area`/`label` each take an id or a display name. Each 

391 of the four ORs within itself when given a list, and they AND 

392 together. No order guarantee; wrap in sorted(...) if you want one. 

393 """ 

394 _require_str(path) 

395 scope = tuple(p for p in (self.integration, self.domain) if p is not None) 

396 pattern = None 

397 if path not in ("", "/") and _REGEX_METACHARS.intersection(path): 

398 pattern = re.compile(path) 

399 prefix = scope 

400 else: 

401 try: 

402 prefix = scope if path in ("", "/") else scope + parse_path(path) 

403 except ValueError as err: 

404 raise KeyError(str(err)) from None 

405 if len(prefix) > 3: 405 ↛ 406line 405 didn't jump to line 406 because the condition on line 405 was never true

406 raise KeyError(path) 

407 

408 platforms = _as_set(platform) 

409 domains = _as_set(domain) 

410 area_ids = _resolve_from_cache(area, self.cache.areas, "area", "area_id") 

411 label_ids = _resolve_from_cache(label, self.cache.labels, "label", "label_id") 

412 

413 def _matches() -> Iterator[_Found]: 

414 for entity in self._entries(prefix): 

415 if platforms and entity.platform not in platforms: 

416 continue 

417 if domains and entity.domain not in domains: 

418 continue 

419 if area_ids and entity.area_id not in area_ids: 

420 continue 

421 if label_ids and not (entity.labels & label_ids): 

422 continue 

423 full_path = "/" + "/".join( 

424 (entity.platform, entity.domain, entity.object_id)[len(scope) :] 

425 ) 

426 if pattern is not None and not pattern.match(full_path): 

427 continue 

428 yield _Found(full_path, entity) 

429 

430 return _matches() 

431 

432 def find( 

433 self, 

434 path: str = "/", 

435 *, 

436 platform: str | list[str] | None = None, 

437 domain: str | list[str] | None = None, 

438 area: str | list[str] | None = None, 

439 label: str | list[str] | None = None, 

440 raw: bool = False, 

441 ) -> Iterator[ApiEntity] | Iterator[dict[str, Any]]: 

442 """Every entity (or raw dict, if raw=True) matching the filters under 

443 `path` (relative to this view, "/" meaning this view's whole 

444 subtree) - flat, skipping the directory-style one-level-at-a-time 

445 grouping .keys()/indexing give you. find_paths()/find_names() are the 

446 same search with just the tree-path or entity_id strings, if that's 

447 all you want - see _find() for the filters. 

448 

449 `raw=True` yields the object as received instead of the typed, 

450 renamed-to-match-Entity ApiEntity view - exactly what show(path) 

451 would return for that path (the merged, cleaned registry+state 

452 dict). The dict's own "entity_id" key identifies which entity it 

453 came from. 

454 """ 

455 matches = self._find( 

456 path, platform=platform, domain=domain, area=area, label=label 

457 ) 

458 if raw: 

459 return (self.show(found.path) for found in matches) 

460 return (found.entity for found in matches) 

461 

462 def find_paths( 

463 self, 

464 path: str = "/", 

465 *, 

466 platform: str | list[str] | None = None, 

467 domain: str | list[str] | None = None, 

468 area: str | list[str] | None = None, 

469 label: str | list[str] | None = None, 

470 ) -> Iterator[str]: 

471 """Just the tree-path strings from find() (e.g. 

472 "/demo/light/kitchen_lights") - see _find() for the filters.""" 

473 return ( 

474 found.path 

475 for found in self._find( 

476 path, platform=platform, domain=domain, area=area, label=label 

477 ) 

478 ) 

479 

480 def find_names( 

481 self, 

482 path: str = "/", 

483 *, 

484 platform: str | list[str] | None = None, 

485 domain: str | list[str] | None = None, 

486 area: str | list[str] | None = None, 

487 label: str | list[str] | None = None, 

488 ) -> Iterator[str]: 

489 """Just the HA entity_id strings from find() (e.g. 

490 "light.kitchen_lights") - see _find() for the filters. Not the same 

491 as find_paths(): this is the flat entity_id, not this tree's 

492 /integration/domain/object_id path.""" 

493 return ( 

494 found.entity_id 

495 for found in self._find( 

496 path, platform=platform, domain=domain, area=area, label=label 

497 ) 

498 ) 

499 

500 def show(self, path: str) -> dict[str, Any]: 

501 """Same shape as the live tree's show() - the registry entry's own 

502 fields plus `state`/`state_attributes`, cleaned the same way (None/ 

503 empty-collection dropped, timestamps as local ISO 8601).""" 

504 _require_str(path) 

505 scope = tuple(p for p in (self.integration, self.domain) if p is not None) 

506 try: 

507 full = scope + parse_path(path) 

508 except ValueError as err: 

509 raise KeyError(str(err)) from None 

510 if len(full) != 3: 510 ↛ 511line 510 didn't jump to line 511 because the condition on line 510 was never true

511 raise KeyError(f"{path!r} - show() needs a full entity path") 

512 integration, domain, object_id = full 

513 entity = self.cache.entities.get(f"{domain}.{object_id}") 

514 if entity is None or entity.platform != integration: 514 ↛ 515line 514 didn't jump to line 515 because the condition on line 514 was never true

515 raise KeyError(path) 

516 

517 data: dict[str, Any] = dict(entity.registry) 

518 data["name"] = entity.name 

519 data["state"] = entity.state 

520 data["state_attributes"] = entity.state_attributes 

521 return _clean(data) 

522 

523 def __repr__(self) -> str: 

524 # "(api)" is just a display hint for a human reading output - the 

525 # binding name and API are identical to the live tree's on purpose. 

526 scope = "/".join(p for p in (self.integration, self.domain) if p is not None) 

527 label = f"obj:/{scope}" if scope else "obj" 

528 kind = ( 

529 "entities" 

530 if self.domain is not None 

531 else "domains" 

532 if self.integration 

533 else "integrations" 

534 ) 

535 return f"<{label}: {len(self)} {kind} (api)>"