Coverage for custom_components/dev_shell_server/session.py: 91%

188 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-10-03 09:02 +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 ) -> ExecResult: 

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

85 async with self._lock: 

86 self.last_used = time.time() 

87 self.executions += 1 

88 out = io.StringIO() 

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

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

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

92 self.globals_.update(self.protected) 

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

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

95 result = ExecResult() 

96 start = time.perf_counter() 

97 try: 

98 coro = self._execute(source) 

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

100 if value is not None: 

101 self.globals_["_"] = value 

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

103 except asyncio.CancelledError: 

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

105 raise 

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

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

108 finally: 

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

110 result.stdout = out.getvalue() 

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

112 text = getattr(result, attr) 

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

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

115 result.truncated = True 

116 return result 

117 

118 async def _execute(self, source: str) -> Any: 

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

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

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

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

123 filename = f"/dev_shell/cell_{next(_cell_counter)}" 

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

125 linecache.cache[filename] = ( 

126 len(source), 

127 None, 

128 source.splitlines(keepends=True), 

129 filename, 

130 ) 

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

132 

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

134 last_expr = None 

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

136 if isinstance(last_stmt, ast.Expr): 

137 tree.body.pop() 

138 last_expr = ast.Expression(last_stmt.value) 

139 

140 flags = ast.PyCF_ALLOW_TOP_LEVEL_AWAIT 

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

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

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

144 # FORWARDREF-based signature rendering in _format_signature below. 

145 if tree.body: 

146 await _run_code( 

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

148 self.globals_, 

149 ) 

150 if last_expr is None: 

151 return None 

152 return await _run_code( 

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

154 self.globals_, 

155 ) 

156 

157 

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

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

160 

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

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

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

164 """ 

165 buf = io.StringIO() 

166 console = Console( 

167 file=buf, 

168 force_terminal=color, 

169 color_system="truecolor" if color else None, 

170 no_color=not color, 

171 highlight=color, 

172 width=width, 

173 ) 

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

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

176 

177 

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

179 """Evaluate code; only await when the code itself used top-level await. 

180 

181 A bare `hass.async_foo()` without await yields a coroutine object, exactly as it 

182 would in component code (strict mode). 

183 """ 

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

185 if code.co_flags & inspect.CO_COROUTINE: 

186 result = await result 

187 return result 

188 

189 

190def _capturing_print(out: io.StringIO): 

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

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

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

194 file = out 

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

196 

197 return shell_print 

198 

199 

200def _capturing_help(out: io.StringIO): 

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

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

203 

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

205 if not args: 

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

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

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

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

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

211 helper.intro() 

212 out.write( 

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

214 ) 

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

216 return 

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

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

219 return 

220 (thing,) = args 

221 if _is_summarisable(thing): 

222 out.write(_class_summary(thing)) 

223 else: 

224 helper(thing) 

225 url = _doc_url(thing) 

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

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

228 

229 return shell_help 

230 

231 

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

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

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

235 

236 

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

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

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

240 return not ( 

241 obj is None 

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

243 or inspect.ismodule(obj) 

244 or inspect.isroutine(obj) 

245 ) 

246 

247 

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

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

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

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

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

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

254 """ 

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

256 header = ( 

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

258 if inspect.isclass(thing) 

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

260 ) 

261 lines = [header, ""] 

262 doc = inspect.getdoc(cls) 

263 if doc: 

264 lines += [doc, ""] 

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

266 lines.append( 

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

268 ) 

269 method_names = sorted( 

270 name 

271 for name in dir(cls) 

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

273 ) 

274 if method_names: 

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

276 for name in method_names: 

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

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

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

280 

281 

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

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

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

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

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

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

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

289 

290 

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

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

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

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

295 

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

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

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

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

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

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

302 

303 

304def _format_signature( 

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

306) -> str: 

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

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

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

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

311 """ 

312 sig = None 

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

314 try: 

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

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

317 sig = None 

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

319 try: 

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

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

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

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

324 try: 

325 sig = inspect.signature(target) 

326 except TypeError, ValueError: 

327 return "(...)" 

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

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

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

331 ret = sig.return_annotation 

332 if ret is not sig.empty and ( 

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

334 ): 

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

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

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

338 

339 

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

341 tb = err.__traceback__ 

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

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

344 tb = tb.tb_next 

345 if isinstance(err, SyntaxError): 

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

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

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

349 else: 

350 text = _render( 

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

352 color=color, 

353 width=width, 

354 ) 

355 return { 

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

357 "message": str(err), 

358 "traceback": text, 

359 } 

360 

361 

362class SessionManager: 

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

364 

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

366 self._bindings = bindings 

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

368 

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

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

371 globals_ = { 

372 "__name__": "__dev_shell__", 

373 "__builtins__": builtins, 

374 **self._bindings, 

375 } 

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

377 name, globals_, dict(self._bindings) 

378 ) 

379 return session 

380 

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

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

383 

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

385 return [ 

386 { 

387 "name": s.name, 

388 "created": s.created, 

389 "last_used": s.last_used, 

390 "executions": s.executions, 

391 "variables": sorted( 

392 k 

393 for k in s.globals_ 

394 if not k.startswith("__") and k not in ("print", "help") 

395 ), 

396 } 

397 for s in self._sessions.values() 

398 ]