Commands executed by ExecContext and its derived classes don't populate the "cmd" parameter of "Result"'s constructor. Fixing that makes for nicer error messages. Signed-off-by: Jan Lindemann <jan@janware.com>
453 lines
14 KiB
Python
453 lines
14 KiB
Python
import asyncio
|
|
import os
|
|
import shlex
|
|
import shutil
|
|
import signal
|
|
import sys
|
|
|
|
import asyncssh # type: ignore[import-not-found]
|
|
|
|
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 = _USE_DEFAULT_KNOWN_HOSTS,
|
|
term_type: str | None = None,
|
|
connect_timeout: float | None = 30.0,
|
|
**kwargs,
|
|
) -> 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
|
|
|
|
async def _open(self) -> None:
|
|
await super()._open()
|
|
await self._conn
|
|
|
|
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:
|
|
kwargs: dict = {
|
|
'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
|
|
ret = {k: v for k, v in kwargs.items() if v is not None}
|
|
if hide_secrets and 'password' in kwargs:
|
|
kwargs['password'] = '<hidden>'
|
|
return ret
|
|
|
|
@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,
|
|
prio,
|
|
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 None:
|
|
return
|
|
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) -> 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)
|
|
|
|
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
|