Coverage for custom_components/ha_repl_server/session.py: 90%

215 statements  

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

1"""REPL sessions: persistent namespaces that execute code on the running event loop. 

2 

3Deliberately free of Home Assistant imports so it can be tested standalone. 

4""" 

5 

6from __future__ import annotations 

7 

8import ast 

9import asyncio 

10import builtins 

11import inspect 

12import io 

13import itertools 

14import linecache 

15import pydoc 

16import re 

17import sys 

18import time 

19import traceback 

20import typing 

21from dataclasses import dataclass, field 

22from pathlib import Path 

23from typing import Any 

24 

25import yaml 

26from rich.console import Console 

27from rich.pretty import Pretty 

28from rich.traceback import Traceback 

29 

30MAX_OUTPUT_CHARS = 1_000_000 

31 

32_cell_counter = itertools.count(1) 

33 

34# help() on a hass object documents its *type*; pydoc has no docstrings for core HA 

35# classes written with newcomers in mind, so point at the real docs instead. Keyed by 

36# fully-qualified class name so it applies no matter what name the object is bound to. 

37# This list only grows, and is useful independently of the code around it, hence the 

38# data file rather than a dict literal here. 

39_DOC_URLS: dict[str, str] = yaml.safe_load( 

40 (Path(__file__).parent / "doc_urls.yaml").read_text() 

41) 

42 

43 

44@dataclass 

45class ExecResult: 

46 """Outcome of running one snippet.""" 

47 

48 stdout: str = "" 

49 value: str | None = None 

50 error: dict[str, str] | None = None 

51 duration: float = 0.0 

52 truncated: bool = False 

53 

54 def as_dict(self) -> dict[str, Any]: 

55 return { 

56 "stdout": self.stdout, 

57 "value": self.value, 

58 "error": self.error, 

59 "duration": self.duration, 

60 "truncated": self.truncated, 

61 } 

62 

63 

64@dataclass 

65class Session: 

66 """A named namespace that persists between executions until reset.""" 

67 

68 name: str 

69 globals_: dict[str, Any] 

70 protected: dict[str, Any] = field(default_factory=dict) 

71 created: float = field(default_factory=time.time) 

72 last_used: float = field(default_factory=time.time) 

73 executions: int = 0 

74 _lock: asyncio.Lock = field(default_factory=asyncio.Lock) 

75 

76 async def run( 

77 self, 

78 source: str, 

79 timeout: float | None = None, 

80 *, 

81 color: bool = False, 

82 width: int = 88, 

83 auto_await: bool = True, 

84 ) -> ExecResult: 

85 """Execute source in this session, returning captured output and the last value.""" 

86 async with self._lock: 

87 self.last_used = time.time() 

88 self.executions += 1 

89 out = io.StringIO() 

90 # hass/obj/open are re-seeded every run, not just at session creation, so a 

91 # snippet that does `obj = obj["/mqtt"]["..."]` only shadows them for its own 

92 # run - the next command always starts from the real bindings again. 

93 self.globals_.update(self.protected) 

94 self.globals_["print"] = _capturing_print(out) 

95 self.globals_["help"] = _capturing_help(out) 

96 self.globals_["_maybe_await"] = _maybe_await 

97 self.globals_["unawait"] = unawait 

98 result = ExecResult() 

99 start = time.perf_counter() 

100 try: 

101 coro = self._execute(source, auto_await=auto_await) 

102 value = await (asyncio.wait_for(coro, timeout) if timeout else coro) 

103 if value is not None: 

104 self.globals_["_"] = value 

105 result.value = _render(Pretty(value), color=color, width=width) 

106 except asyncio.CancelledError: 

107 # Cancellation of the caller must propagate, not be reported as a result. 

108 raise 

109 except BaseException as err: # noqa: BLE001 - report everything, incl. SystemExit 

110 result.error = _format_error(err, color=color, width=width) 

111 finally: 

112 result.duration = time.perf_counter() - start 

113 result.stdout = out.getvalue() 

114 for attr in ("stdout", "value"): 

115 text = getattr(result, attr) 

116 if text is not None and len(text) > MAX_OUTPUT_CHARS: 116 ↛ 117line 116 didn't jump to line 117 because the condition on line 116 was never true

117 setattr(result, attr, text[:MAX_OUTPUT_CHARS]) 

118 result.truncated = True 

119 return result 

120 

121 async def _execute(self, source: str, *, auto_await: bool = True) -> Any: 

122 # Not "<ha_repl-N>": rich.traceback refuses to show source for any 

123 # filename starting with "<" (treats it like "<stdin>"), no matter what 

124 # linecache holds. An absolute-looking path sidesteps that - rich joins a 

125 # relative one onto the cwd before the linecache lookup, which would miss. 

126 filename = f"/ha_repl/cell_{next(_cell_counter)}" 

127 # Register the source so tracebacks can show the offending lines. 

128 linecache.cache[filename] = ( 

129 len(source), 

130 None, 

131 source.splitlines(keepends=True), 

132 filename, 

133 ) 

134 tree = ast.parse(source, filename, "exec") 

135 if auto_await: 

136 tree = _AutoAwait().visit(tree) 

137 ast.fix_missing_locations(tree) 

138 

139 # Like the interactive interpreter, echo the value of a trailing expression. 

140 last_expr = None 

141 last_stmt = tree.body[-1] if tree.body else None 

142 if isinstance(last_stmt, ast.Expr): 

143 tree.body.pop() 

144 last_expr = ast.Expression(last_stmt.value) 

145 

146 flags = ast.PyCF_ALLOW_TOP_LEVEL_AWAIT 

147 # dont_inherit=True: compile() otherwise inherits this module's own 

148 # `from __future__ import annotations`, which would make every 

149 # annotation in the user's code a plain string - breaking the 

150 # FORWARDREF-based signature rendering in _format_signature below. 

151 if tree.body: 

152 await _run_code( 

153 compile(tree, filename, "exec", flags=flags, dont_inherit=True), 

154 self.globals_, 

155 ) 

156 if last_expr is None: 

157 return None 

158 return await _run_code( 

159 compile(last_expr, filename, "eval", flags=flags, dont_inherit=True), 

160 self.globals_, 

161 ) 

162 

163 

164def _render(renderable: Any, *, color: bool, width: int) -> str: 

165 """Render a Rich renderable (a value's pretty repr, a traceback) to text. 

166 

167 `force_terminal`/`no_color` are set explicitly rather than auto-detected: 

168 the real terminal is on the far end of a websocket call, not this process, 

169 so the caller (which does know) decides via `color`. 

170 """ 

171 buf = io.StringIO() 

172 console = Console( 

173 file=buf, 

174 force_terminal=color, 

175 color_system="truecolor" if color else None, 

176 no_color=not color, 

177 highlight=color, 

178 width=width, 

179 ) 

180 console.print(renderable, end="") 

181 return buf.getvalue().rstrip("\n") 

182 

183 

184async def _maybe_await(value: Any) -> Any: 

185 """Finish a coroutine (or other awaitable) the user forgot to `await`; 

186 anything else passes straight through. Used both by the AST rewrite below 

187 (an unawaited call nested inside a larger expression) and _run_code's own 

188 check (a bare reference to one created earlier, e.g. `c = f(); c`).""" 

189 return await value if inspect.isawaitable(value) else value 

190 

191 

192def unawait(value: Any) -> Any: 

193 """Identity function and auto-await escape hatch: `unawait(f())` returns 

194 f()'s bare, un-awaited result (a coroutine, if f is async) - for when 

195 that's genuinely wanted, e.g. batching into `asyncio.gather(*[unawait(f()) 

196 for f in ...])`. Recognised by name in the `_AutoAwait` rewrite below, 

197 which skips its whole argument rather than calling this at runtime; this 

198 plain version only runs if auto-await itself is off (`auto_await=False`) 

199 or `unawait` is used somewhere the rewrite doesn't reach. 

200 """ 

201 return value 

202 

203 

204class _AutoAwait(ast.NodeTransformer): 

205 """Rewrites every call not already explicitly awaited to go through 

206 _maybe_await() first, so `hass.async_foo()` works whether or not the 

207 user remembered `await` - including nested inside attribute access 

208 (`hass.async_foo().attr`), a comprehension, or an argument list. Leaves 

209 nested (synchronous) function/lambda bodies alone: `await` there is a 

210 SyntaxError, and those calls run later, not as part of this statement 

211 anyway. `unawait(expr)` is the escape hatch - its argument is left 

212 completely untouched for when the bare coroutine is wanted. Opt out 

213 entirely with `auto_await=False` on the exec call (strict mode: 

214 forgetting await behaves exactly as in component code). 

215 """ 

216 

217 def visit_FunctionDef(self, node: ast.FunctionDef) -> ast.FunctionDef: 

218 return node 

219 

220 def visit_Lambda(self, node: ast.Lambda) -> ast.Lambda: 

221 return node 

222 

223 def visit_Await(self, node: ast.Await) -> ast.Await: 

224 # Already explicit: don't double-await this call, but still rewrite 

225 # anything nested inside its own arguments. 

226 if isinstance(node.value, ast.Call): 226 ↛ 231line 226 didn't jump to line 231 because the condition on line 226 was always true

227 node.value.func = self.visit(node.value.func) 

228 node.value.args = [self.visit(a) for a in node.value.args] 

229 node.value.keywords = [self.visit(k) for k in node.value.keywords] 

230 else: 

231 node.value = self.visit(node.value) 

232 return node 

233 

234 def visit_Call(self, node: ast.Call) -> ast.AST: 

235 if ( 

236 isinstance(node.func, ast.Name) 

237 and node.func.id == "unawait" 

238 and len(node.args) == 1 

239 and not node.keywords 

240 ): 

241 return node.args[0] 

242 self.generic_visit(node) 

243 wrapped = ast.Call( 

244 func=ast.Name(id="_maybe_await", ctx=ast.Load()), args=[node], keywords=[] 

245 ) 

246 return ast.Await(value=wrapped) 

247 

248 

249async def _run_code(code: Any, globals_: dict[str, Any]) -> Any: 

250 # No separate "bare coroutine" backstop here: when auto_await is on, the 

251 # _AutoAwait rewrite already resolves every call site, and NOT doing so 

252 # unconditionally is exactly what `unawait(...)` asks for - a backstop 

253 # here would silently defeat it for a trailing `unawait(f())`. 

254 result = eval(code, globals_) # nosec B307 - the whole point of a dev shell 

255 if code.co_flags & inspect.CO_COROUTINE: 

256 result = await result 

257 return result 

258 

259 

260def _capturing_print(out: io.StringIO): 

261 def shell_print(*args: Any, file: Any = None, **kwargs: Any) -> None: 

262 # stdout and stderr both come back to the client; other files are honoured. 

263 if file is None or file is sys.stdout or file is sys.stderr: 263 ↛ 265line 263 didn't jump to line 265 because the condition on line 263 was always true

264 file = out 

265 builtins.print(*args, file=file, **kwargs) 

266 

267 return shell_print 

268 

269 

270def _capturing_help(out: io.StringIO): 

271 # A fresh Helper per call, output redirected into the same buffer as print(). 

272 helper = pydoc.Helper(input=io.StringIO(), output=out) 

273 

274 def shell_help(*args: Any) -> None: 

275 if not args: 

276 # help()'s real interactive loop (reading "help> " commands) spins the 

277 # CPU forever on this Python version when its input isn't a real tty - 

278 # verified in isolation, nothing project-specific about it - and there's no 

279 # live stdin to browse with anyway over this request/response API. Show 

280 # the same intro banner and stop there instead of entering interact(). 

281 helper.intro() 

282 out.write( 

283 "\nGet help on any object, with links for known Home Assistant classes\n" 

284 ) 

285 out.write("\ne.g. help(hass) or help(obj['/sun/sun']).\n") 

286 return 

287 if len(args) > 1: 287 ↛ 288line 287 didn't jump to line 288 because the condition on line 287 was never true

288 helper(*args) # raises the same TypeError real help() would 

289 return 

290 (thing,) = args 

291 if _is_summarisable(thing): 

292 out.write(_class_summary(thing)) 

293 else: 

294 helper(thing) 

295 url = _doc_url(thing) 

296 if url: 296 ↛ 297line 296 didn't jump to line 297 because the condition on line 296 was never true

297 out.write(f"\nSee also: {url}\n") 

298 

299 return shell_help 

300 

301 

302def _doc_url(obj: Any) -> str | None: 

303 cls = obj if inspect.isclass(obj) else type(obj) 

304 return _DOC_URLS.get(f"{cls.__module__}.{cls.__qualname__}") 

305 

306 

307def _is_summarisable(obj: Any) -> bool: 

308 """Classes and plain instances get the condensed summary below; modules, 

309 functions/methods, and primitives go through pydoc as usual (already short).""" 

310 return not ( 

311 obj is None 

312 or isinstance(obj, (bool, int, float, complex, str, bytes)) 

313 or inspect.ismodule(obj) 

314 or inspect.isroutine(obj) 

315 ) 

316 

317 

318def _class_summary(thing: Any) -> str: 

319 """A class's full pydoc page repeats its docstring once per method, which 

320 balloons for a class like HomeAssistant with hundreds of methods. A method's 

321 own docstring is one `help(hass.the_method)` away, so here we only show the 

322 class docstring, the constructor signature, and the methods' names and 

323 signatures (no live attribute values either - that's what `hass.foo` is for). 

324 """ 

325 cls = thing if inspect.isclass(thing) else type(thing) 

326 header = ( 

327 f"Help on class {cls.__qualname__} in module {cls.__module__}:" 

328 if inspect.isclass(thing) 

329 else f"Help on {cls.__qualname__} object in module {cls.__module__}:" 

330 ) 

331 lines = [header, ""] 

332 doc = inspect.getdoc(cls) 

333 if doc: 

334 lines += [doc, ""] 

335 # A constructor "returning Self" is implied, not useful to state. 

336 lines.append( 

337 f"{cls.__qualname__}{_format_signature(cls, drop_return=(typing.Self,))}" 

338 ) 

339 method_names = sorted( 

340 name 

341 for name in dir(cls) 

342 if not name.startswith("_") and _is_plain_method(getattr(cls, name, None)) 

343 ) 

344 if method_names: 

345 lines += ["", "Methods:"] 

346 for name in method_names: 

347 sig = _format_signature(getattr(cls, name), drop_self=True) 

348 lines.append(f" {name}{sig}") 

349 return "\n".join(lines) + "\n" 

350 

351 

352def _is_plain_method(obj: Any) -> bool: 

353 # inspect.isroutine() also matches any non-data descriptor (defines __get__ but 

354 # not __set__) via its ismethoddescriptor() check - which, alongside genuine 

355 # methods, catches cached_property (both functools's and propcache's, used all 

356 # over Home Assistant's own entity classes), listing a read-only property as a 

357 # fake "name(...)" method. isfunction/ismethod is the properly narrow check. 

358 return inspect.isfunction(obj) or inspect.ismethod(obj) 

359 

360 

361# Matches a dotted path ending in an identifier, e.g. "homeassistant.core.State" 

362# or the "collections.abc.Coroutine" inside a bigger type expression - used to 

363# shorten type annotations down to just the class name. 

364_DOTTED_NAME = re.compile(r"\b(?:[A-Za-z_]\w*\.)+([A-Za-z_]\w*)\b") 

365 

366# Python 3.14+ (PEP 649): __annotations__ access eagerly evaluates every annotation 

367# by default, which raises for names only imported under TYPE_CHECKING - common in 

368# HA's own codebase (e.g. Entity methods taking an "EntityPlatform"). FORWARDREF 

369# format resolves what it can and leaves the rest as an inert placeholder instead 

370# of raising. Not available before 3.14, hence the getattr dance. 

371_FORWARDREF_FORMAT = getattr(getattr(inspect, "Format", None), "FORWARDREF", None) 

372 

373 

374def _format_signature( 

375 target: Any, *, drop_self: bool = False, drop_return: tuple = () 

376) -> str: 

377 """Render target's signature, trimmed down for a quick overview: 

378 self dropped (when `drop_self`), a void or otherwise uninformative return 

379 annotation dropped (`drop_return`), type annotations shortened to their bare 

380 class name, and no padding around `=` for default values. 

381 """ 

382 sig = None 

383 if _FORWARDREF_FORMAT is not None: 383 ↛ 388line 383 didn't jump to line 388 because the condition on line 383 was always true

384 try: 

385 sig = inspect.signature(target, annotation_format=_FORWARDREF_FORMAT) 

386 except Exception: # noqa: BLE001 - fall through to the attempts below 

387 sig = None 

388 if sig is None: 388 ↛ 389line 388 didn't jump to line 389 because the condition on line 388 was never true

389 try: 

390 sig = inspect.signature(target, eval_str=True) 

391 except Exception: # noqa: BLE001 - unresolvable forward refs can raise almost 

392 # anything (NameError, AttributeError, ...); a signature with raw/partial 

393 # annotations is still better than losing the method entirely. 

394 try: 

395 sig = inspect.signature(target) 

396 except TypeError, ValueError: 

397 return "(...)" 

398 params = list(sig.parameters.values()) 

399 if drop_self and params and params[0].name == "self": 

400 sig = sig.replace(parameters=params[1:]) 

401 ret = sig.return_annotation 

402 if ret is not sig.empty and ( 

403 ret is None or ret is type(None) or any(ret is d for d in drop_return) 

404 ): 

405 sig = sig.replace(return_annotation=sig.empty) 

406 text = _DOTTED_NAME.sub(r"\1", str(sig)) 

407 return re.sub(r"\s*=\s*", "=", text) 

408 

409 

410def _format_error(err: BaseException, *, color: bool, width: int) -> dict[str, str]: 

411 tb = err.__traceback__ 

412 # Drop the frames belonging to this module so the traceback starts at user code. 

413 while tb is not None and tb.tb_frame.f_code.co_filename == __file__: 

414 tb = tb.tb_next 

415 if isinstance(err, SyntaxError): 

416 # No frames worth showing for this one, just the offending line and caret - 

417 # a plain rendering already does that job, so it skips the Rich treatment. 

418 text = "".join(traceback.format_exception_only(type(err), err)) 

419 else: 

420 text = _render( 

421 Traceback.from_exception(type(err), err, tb, width=width), 

422 color=color, 

423 width=width, 

424 ) 

425 return { 

426 "type": type(err).__name__, 

427 "message": str(err), 

428 "traceback": text, 

429 } 

430 

431 

432class SessionManager: 

433 """Holds sessions by name; a session lives until reset or process restart.""" 

434 

435 def __init__(self, bindings: dict[str, Any]) -> None: 

436 self._bindings = bindings 

437 self._sessions: dict[str, Session] = {} 

438 

439 def get(self, name: str) -> Session: 

440 if (session := self._sessions.get(name)) is None: 

441 globals_ = { 

442 "__name__": "__ha_repl__", 

443 "__builtins__": builtins, 

444 **self._bindings, 

445 } 

446 session = self._sessions[name] = Session( 

447 name, globals_, dict(self._bindings) 

448 ) 

449 return session 

450 

451 def reset(self, name: str) -> bool: 

452 return self._sessions.pop(name, None) is not None 

453 

454 def describe(self) -> list[dict[str, Any]]: 

455 return [ 

456 { 

457 "name": s.name, 

458 "created": s.created, 

459 "last_used": s.last_used, 

460 "executions": s.executions, 

461 "variables": sorted( 

462 k 

463 for k in s.globals_ 

464 if not k.startswith("__") 

465 and k not in ("print", "help", "_maybe_await", "unawait") 

466 ), 

467 } 

468 for s in self._sessions.values() 

469 ]