import asyncio import os import shlex import shutil import signal import sys from typing import Any, override import asyncssh # type: ignore[import-not-found, unused-ignore] from asyncssh import SSHReader # type: ignore[import-not-found, unused-ignore] from ...base import Result from ...log import DEBUG, ERR, NOTICE, log from ..SSHClient import SSHClient as Base from .util import join_cmd _USE_DEFAULT_KNOWN_HOSTS = object() class AsyncSSH(Base): def __init__( self, uri: str, *, client_keys: list[str] | None = None, known_hosts: Any = _USE_DEFAULT_KNOWN_HOSTS, term_type: str | None = None, connect_timeout: float | None = 30.0, **kwargs: Any, ) -> None: super().__init__( uri, caps = self.Caps.LogOutput | self.Caps.Wd | self.Caps.Interactive | self.Caps.ModEnv, **kwargs, ) self.__client_keys = client_keys self.__known_hosts = known_hosts self.__term_type = term_type or os.environ.get('TERM', 'xterm') self.__connect_timeout = connect_timeout self.__conn: asyncssh.SSHClientConnection | None = None @override async def _open(self) -> None: await super()._open() await self._conn @override async def _close(self) -> None: if self.__conn is not None: try: self.__conn.close() await self.__conn.wait_closed() except Exception as e: log(DEBUG, f'Failed to close connection ({str(e)}, ignored)') self.__conn = None def _connect_kwargs(self, hide_secrets: bool = False) -> dict[str, Any]: kwargs: dict[str, Any] = { 'host': self.hostname, 'port': self.port, 'username': self.username, 'password': self.password, 'client_keys': self.__client_keys, 'connect_timeout': self.__connect_timeout, } if self.__known_hosts is not _USE_DEFAULT_KNOWN_HOSTS: kwargs['known_hosts'] = self.__known_hosts if hide_secrets and 'password' in kwargs: kwargs['password'] = '' return {k: v for k, v in kwargs.items() if v is not None} @property async def _conn(self) -> asyncssh.SSHClientConnection: if self.__conn is None: try: self.__conn = await asyncssh.connect(**self._connect_kwargs()) except Exception as e: msg = f'-------------------- Failed to connect ({str(e)})' log(ERR, ',', msg) for key, val in self._connect_kwargs(hide_secrets = True).items(): log(ERR, f'| {key:<20} = {val}') log(ERR, '`', msg) raise return self.__conn @staticmethod def _build_remote_command(cmd: list[str], wd: str | None) -> str: inner = f'exec {join_cmd(cmd)}' if wd is not None: inner = f'cd {shlex.quote(wd)} && {inner}' return f'/bin/sh -lc {shlex.quote(inner)}' @staticmethod def _has_local_tty() -> bool: try: return sys.stdin.isatty() and sys.stdout.isatty() except Exception: return False @staticmethod def _get_local_term_size() -> tuple[int, int, int, int]: cols, rows = shutil.get_terminal_size(fallback = (80, 24)) xpixel = ypixel = 0 try: import fcntl import struct import termios packed = fcntl.ioctl(sys.stdout.fileno(), termios.TIOCGWINSZ, b'\0' * 8) rows2, cols2, xpixel, ypixel = struct.unpack('HHHH', packed) if cols2 > 0 and rows2 > 0: cols, rows = cols2, rows2 except Exception: pass return (cols, rows, xpixel, ypixel) async def _read_stream( self, stream: SSHReader[bytes], prio: int, collector: list[bytes], *, verbose: bool, log_prefix: str, log_enc: str, ) -> None: buf = b'' while True: chunk = await stream.read(4096) if not chunk: break collector.append(chunk) if verbose: buf += chunk while b'\n' in buf: line, buf = buf.split(b'\n', 1) log(prio, log_prefix, line.decode(log_enc, errors = 'replace')) if verbose and buf: log(prio, log_prefix, buf.decode(log_enc, errors = 'replace')) async def _run_interactive_on_conn( self, cmd: list[str], wd: str | None, cmd_input: bytes | None, mod_env: dict[str, str] | None, ) -> Result: conn = await self._conn command = self._build_remote_command(cmd, wd) stdout_parts: list[bytes] = [] proc = await conn.create_process( command = command, env = mod_env, stdin = asyncssh.PIPE, stdout = asyncssh.PIPE, stderr = asyncssh.STDOUT, encoding = None, request_pty = 'force', term_type = self.__term_type, term_size = self._get_local_term_size(), ) loop = asyncio.get_running_loop() stdin_fd = sys.stdin.fileno() stdin_queue: asyncio.Queue[bytes | None] = asyncio.Queue() old_tty_state = None old_winch_handler = None stdin_reader_installed = False def _write_local(data: bytes) -> None: try: sys.stdout.buffer.write(data) sys.stdout.buffer.flush() except AttributeError: os.write(sys.stdout.fileno(), data) def _on_stdin_ready() -> None: try: data = os.read(stdin_fd, 4096) except OSError: data = b'' if data: stdin_queue.put_nowait(data) else: try: loop.remove_reader(stdin_fd) except Exception: pass stdin_queue.put_nowait(None) async def _pump_stdin() -> None: if cmd_input is not None and proc.stdin is not None: proc.stdin.write(cmd_input) await proc.stdin.drain() while True: data = await stdin_queue.get() if data is None: if proc.stdin is not None: try: proc.stdin.write_eof() except (BrokenPipeError, OSError): pass return if proc.stdin is not None: proc.stdin.write(data) await proc.stdin.drain() async def _pump_stdout() -> None: while True: chunk = await proc.stdout.read(4096) if not chunk: break stdout_parts.append(chunk) _write_local(chunk) def _on_winch(*_args: Any) -> None: try: proc.change_terminal_size(*self._get_local_term_size()) except Exception: pass try: sys.stdout.flush() sys.stderr.flush() try: import termios import tty old_tty_state = termios.tcgetattr(stdin_fd) tty.setraw(stdin_fd) except Exception: old_tty_state = None try: loop.add_reader(stdin_fd, _on_stdin_ready) stdin_reader_installed = True except (NotImplementedError, RuntimeError): stdin_queue.put_nowait(None) if hasattr(signal, 'SIGWINCH'): try: old_winch_handler = signal.getsignal(signal.SIGWINCH) signal.signal(signal.SIGWINCH, _on_winch) except Exception: old_winch_handler = None stdin_task = asyncio.create_task(_pump_stdin()) stdout_task = asyncio.create_task(_pump_stdout()) completed = await proc.wait(check = False) await stdout_task if not stdin_task.done(): stdin_task.cancel() try: await stdin_task except asyncio.CancelledError: pass exit_code = completed.exit_status if exit_code is None: exit_code = ( completed.returncode if completed.returncode is not None else -1 ) stdout = b''.join(stdout_parts) if stdout_parts else None return Result(stdout, None, exit_code, cmd = cmd) finally: if stdin_reader_installed: try: loop.remove_reader(stdin_fd) except Exception: pass if old_winch_handler is not None and hasattr(signal, 'SIGWINCH'): try: signal.signal(signal.SIGWINCH, old_winch_handler) except Exception: pass if old_tty_state is not None: try: import termios termios.tcsetattr(stdin_fd, termios.TCSADRAIN, old_tty_state) except Exception: pass try: sys.stdout.flush() sys.stderr.flush() except Exception: pass async def _run_captured_pty_on_conn( self, cmd: list[str], wd: str | None, verbose: bool, cmd_input: bytes | None, mod_env: dict[str, str] | None, log_prefix: str, ) -> Result: conn = await self._conn command = self._build_remote_command(cmd, wd) stdout_parts: list[bytes] = [] stdout_log_enc = sys.stdout.encoding or 'utf-8' proc = await conn.create_process( command = command, env = mod_env, stdin = asyncssh.PIPE if cmd_input is not None else asyncssh.DEVNULL, stdout = asyncssh.PIPE, stderr = asyncssh.STDOUT, encoding = None, request_pty = 'force', term_type = self.__term_type, ) task = asyncio.create_task( self._read_stream( proc.stdout, NOTICE, stdout_parts, verbose = verbose, log_prefix = log_prefix, log_enc = stdout_log_enc, ) ) if cmd_input is not None and proc.stdin is not None: proc.stdin.write(cmd_input) await proc.stdin.drain() proc.stdin.write_eof() completed = await proc.wait(check = False) await task exit_code = completed.exit_status if exit_code is None: exit_code = completed.returncode if completed.returncode is not None else -1 stdout = b''.join(stdout_parts) if stdout_parts else None return Result(stdout, None, exit_code, cmd = cmd) @override async def _run_ssh( self, cmd: list[str], wd: str | None, verbose: bool, cmd_input: bytes | None, mod_env: dict[str, str] | None, interactive: bool, log_prefix: str, ) -> Result: try: if interactive: if self._has_local_tty(): return await self._run_interactive_on_conn( cmd = cmd, wd = wd, cmd_input = cmd_input, mod_env = mod_env, ) return await self._run_captured_pty_on_conn( cmd = cmd, wd = wd, verbose = verbose, cmd_input = cmd_input, mod_env = mod_env, log_prefix = log_prefix, ) command = self._build_remote_command(cmd, wd) stdout_parts: list[bytes] = [] stderr_parts: list[bytes] = [] stdout_log_enc = sys.stdout.encoding or 'utf-8' stderr_log_enc = sys.stderr.encoding or 'utf-8' stdin_mode = asyncssh.PIPE if cmd_input is not None else asyncssh.DEVNULL conn = await self._conn proc = await conn.create_process( command = command, env = mod_env, stdin = stdin_mode, stdout = asyncssh.PIPE, stderr = asyncssh.PIPE, encoding = None, request_pty = False, ) tasks = [ asyncio.create_task( self._read_stream( proc.stdout, NOTICE, stdout_parts, verbose = verbose, log_prefix = log_prefix, log_enc = stdout_log_enc, ) ), asyncio.create_task( self._read_stream( proc.stderr, ERR, stderr_parts, verbose = verbose, log_prefix = log_prefix, log_enc = stderr_log_enc, ) ), ] if cmd_input is not None and proc.stdin is not None: proc.stdin.write(cmd_input) await proc.stdin.drain() proc.stdin.write_eof() completed = await proc.wait(check = False) await asyncio.gather(*tasks) stdout = b''.join(stdout_parts) if stdout_parts else None stderr = b''.join(stderr_parts) if stderr_parts else None exit_code = completed.exit_status if exit_code is None: exit_code = ( completed.returncode if completed.returncode is not None else -1 ) return Result(stdout, stderr, exit_code, cmd = cmd) except Exception as e: log(ERR, f'Failed to run command {" ".join(cmd)} ({e})') raise