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
« 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.
3Deliberately free of Home Assistant imports so it can be tested standalone.
4"""
6from __future__ import annotations
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
25import yaml
26from rich.console import Console
27from rich.pretty import Pretty
28from rich.traceback import Traceback
30MAX_OUTPUT_CHARS = 1_000_000
32_cell_counter = itertools.count(1)
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)
44@dataclass
45class ExecResult:
46 """Outcome of running one snippet."""
48 stdout: str = ""
49 value: str | None = None
50 error: dict[str, str] | None = None
51 duration: float = 0.0
52 truncated: bool = False
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 }
64@dataclass
65class Session:
66 """A named namespace that persists between executions until reset."""
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)
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
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)
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)
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 )
164def _render(renderable: Any, *, color: bool, width: int) -> str:
165 """Render a Rich renderable (a value's pretty repr, a traceback) to text.
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")
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
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
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 """
217 def visit_FunctionDef(self, node: ast.FunctionDef) -> ast.FunctionDef:
218 return node
220 def visit_Lambda(self, node: ast.Lambda) -> ast.Lambda:
221 return node
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
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)
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
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)
267 return shell_print
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)
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")
299 return shell_help
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__}")
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 )
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"
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)
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")
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)
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)
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 }
432class SessionManager:
433 """Holds sessions by name; a session lives until reset or process restart."""
435 def __init__(self, bindings: dict[str, Any]) -> None:
436 self._bindings = bindings
437 self._sessions: dict[str, Session] = {}
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
451 def reset(self, name: str) -> bool:
452 return self._sessions.pop(name, None) is not None
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 ]