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
« 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.
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 ) -> 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
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")
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)
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 )
158def _render(renderable: Any, *, color: bool, width: int) -> str:
159 """Render a Rich renderable (a value's pretty repr, a traceback) to text.
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")
178async def _run_code(code: Any, globals_: dict[str, Any]) -> Any:
179 """Evaluate code; only await when the code itself used top-level await.
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
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)
197 return shell_print
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)
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")
229 return shell_help
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__}")
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 )
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"
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)
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")
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)
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)
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 }
362class SessionManager:
363 """Holds sessions by name; a session lives until reset or process restart."""
365 def __init__(self, bindings: dict[str, Any]) -> None:
366 self._bindings = bindings
367 self._sessions: dict[str, Session] = {}
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
381 def reset(self, name: str) -> bool:
382 return self._sessions.pop(name, None) is not None
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 ]