Coverage for src/homeassistant_repl/local_session.py: 0%
87 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"""API client mode's exec engine: runs snippets right here in the CLI
2process, not on a remote Home Assistant session - there's no server to send
3code to, since API client mode talks to Home Assistant only through its
4standard API (see api_objtree.py). Deliberately mirrors
5custom_components/ha_repl_server/session.py's exec model (top-level await,
6trailing-expression echo via rich, rich tracebacks) rather than importing it:
7that module lives in a separate HACS-deployed package with its own packaging
8boundary, and duplicating ~100 lines here is simpler than bridging it.
10Simpler than session.py in one respect: output goes straight to the real
11terminal, so there's no print()/help() capturing-and-shipping-back machinery,
12and rich's Console can auto-detect color/width instead of being told.
13"""
15from __future__ import annotations
17import ast
18import builtins
19import inspect
20import itertools
21import linecache
22import sys
23import traceback
24from dataclasses import dataclass
25from typing import Any
27from rich.console import Console
28from rich.pretty import Pretty
29from rich.traceback import Traceback
31_cell_counter = itertools.count(1)
33_console = Console()
34_error_console = Console(stderr=True)
37@dataclass
38class LocalSession:
39 """A namespace that persists between `run()` calls until the process exits."""
41 globals_: dict[str, Any]
42 auto_await: bool = True
44 def __post_init__(self) -> None:
45 self.globals_.setdefault("__name__", "__ha_repl_api__")
46 self.globals_.setdefault("__builtins__", builtins)
47 self.globals_.setdefault("_maybe_await", _maybe_await)
48 self.globals_.setdefault("unawait", unawait)
50 async def run(self, source: str) -> None:
51 """Execute source, printing the trailing expression's value (if any)
52 or a traceback directly to the real terminal - there's no result to
53 ship back over a wire, so this doesn't return one."""
54 try:
55 value = await self._execute(source)
56 except BaseException as err: # noqa: BLE001 - report everything, incl. SystemExit
57 _print_error(err)
58 return
59 if value is not None:
60 self.globals_["_"] = value
61 _console.print(Pretty(value))
63 async def _execute(self, source: str) -> Any:
64 # Not "<ha_repl-N>": rich.traceback refuses to show source for any
65 # filename starting with "<" (treats it like "<stdin>"), no matter what
66 # linecache holds. An absolute-looking path sidesteps that - rich joins a
67 # relative one onto the cwd before the linecache lookup, which would miss.
68 filename = f"/ha_repl/api_cell_{next(_cell_counter)}"
69 linecache.cache[filename] = (
70 len(source),
71 None,
72 source.splitlines(keepends=True),
73 filename,
74 )
75 tree = ast.parse(source, filename, "exec")
76 if self.auto_await:
77 tree = _AutoAwait().visit(tree)
78 ast.fix_missing_locations(tree)
80 last_expr = None
81 last_stmt = tree.body[-1] if tree.body else None
82 if isinstance(last_stmt, ast.Expr):
83 tree.body.pop()
84 last_expr = ast.Expression(last_stmt.value)
86 flags = ast.PyCF_ALLOW_TOP_LEVEL_AWAIT
87 # dont_inherit=True: compile() otherwise inherits this module's own
88 # `from __future__ import annotations`, which would make every
89 # annotation in the user's code a plain string instead of a real value.
90 if tree.body:
91 await _run_code(
92 compile(tree, filename, "exec", flags=flags, dont_inherit=True),
93 self.globals_,
94 )
95 if last_expr is None:
96 return None
97 return await _run_code(
98 compile(last_expr, filename, "eval", flags=flags, dont_inherit=True),
99 self.globals_,
100 )
103async def _maybe_await(value: Any) -> Any:
104 """Finish a coroutine (or other awaitable) the user forgot to `await`;
105 anything else passes straight through. Used both by the AST rewrite below
106 (an unawaited call nested inside a larger expression) and _run_code's own
107 check (a bare reference to one created earlier, e.g. `c = f(); c`)."""
108 return await value if inspect.isawaitable(value) else value
111def unawait(value: Any) -> Any:
112 """Identity function and auto-await escape hatch: `unawait(f())` returns
113 f()'s bare, un-awaited result (a coroutine, if f is async) - for when
114 that's genuinely wanted, e.g. batching into `asyncio.gather(*[unawait(f())
115 for f in ...])`. Recognised by name in the `_AutoAwait` rewrite below,
116 which skips its whole argument rather than calling this at runtime; this
117 plain version only runs if auto-await itself is off (--no-auto-await) or
118 `unawait` is used somewhere the rewrite doesn't reach.
119 """
120 return value
123class _AutoAwait(ast.NodeTransformer):
124 """Rewrites every call not already explicitly awaited to go through
125 _maybe_await() first, so `obj.async_method()` works whether or not the
126 user remembered `await` - including nested inside attribute access
127 (`obj.async_method().attr`), a comprehension, or an argument list.
128 Leaves nested (synchronous) function/lambda bodies alone: `await` there
129 is a SyntaxError, and those calls run later, not as part of this
130 statement anyway. `unawait(expr)` is the escape hatch - its argument is
131 left completely untouched for when the bare coroutine is wanted.
132 """
134 def visit_FunctionDef(self, node: ast.FunctionDef) -> ast.FunctionDef:
135 return node
137 def visit_Lambda(self, node: ast.Lambda) -> ast.Lambda:
138 return node
140 def visit_Await(self, node: ast.Await) -> ast.Await:
141 # Already explicit: don't double-await this call, but still rewrite
142 # anything nested inside its own arguments.
143 if isinstance(node.value, ast.Call):
144 node.value.func = self.visit(node.value.func)
145 node.value.args = [self.visit(a) for a in node.value.args]
146 node.value.keywords = [self.visit(k) for k in node.value.keywords]
147 else:
148 node.value = self.visit(node.value)
149 return node
151 def visit_Call(self, node: ast.Call) -> ast.AST:
152 if (
153 isinstance(node.func, ast.Name)
154 and node.func.id == "unawait"
155 and len(node.args) == 1
156 and not node.keywords
157 ):
158 return node.args[0]
159 self.generic_visit(node)
160 wrapped = ast.Call(
161 func=ast.Name(id="_maybe_await", ctx=ast.Load()), args=[node], keywords=[]
162 )
163 return ast.Await(value=wrapped)
166async def _run_code(code: Any, globals_: dict[str, Any]) -> Any:
167 # No separate "bare coroutine" backstop here: when auto_await is on, the
168 # _AutoAwait rewrite already resolves every call site, and NOT doing so
169 # unconditionally is exactly what `unawait(...)` asks for - a backstop
170 # here would silently defeat it for a trailing `unawait(f())`.
171 result = eval(code, globals_) # nosec B307 - the whole point of a dev shell
172 if code.co_flags & inspect.CO_COROUTINE:
173 result = await result
174 return result
177def _print_error(err: BaseException) -> None:
178 tb = err.__traceback__
179 # Drop the frames belonging to this module so the traceback starts at user code.
180 while tb is not None and tb.tb_frame.f_code.co_filename == __file__:
181 tb = tb.tb_next
182 if isinstance(err, SyntaxError):
183 # No frames worth showing for this one - plain is fine.
184 print(
185 "".join(traceback.format_exception_only(type(err), err)),
186 end="",
187 file=sys.stderr,
188 )
189 return
190 _error_console.print(Traceback.from_exception(type(err), err, tb))