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

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. 

9 

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""" 

14 

15from __future__ import annotations 

16 

17import ast 

18import builtins 

19import inspect 

20import itertools 

21import linecache 

22import sys 

23import traceback 

24from dataclasses import dataclass 

25from typing import Any 

26 

27from rich.console import Console 

28from rich.pretty import Pretty 

29from rich.traceback import Traceback 

30 

31_cell_counter = itertools.count(1) 

32 

33_console = Console() 

34_error_console = Console(stderr=True) 

35 

36 

37@dataclass 

38class LocalSession: 

39 """A namespace that persists between `run()` calls until the process exits.""" 

40 

41 globals_: dict[str, Any] 

42 auto_await: bool = True 

43 

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) 

49 

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)) 

62 

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) 

79 

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) 

85 

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 ) 

101 

102 

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 

109 

110 

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 

121 

122 

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 """ 

133 

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

135 return node 

136 

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

138 return node 

139 

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 

150 

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) 

164 

165 

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 

175 

176 

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))