jw-pkg/src/python/jw/pkg/lib/ec/ssh/AsyncSSH.py
Jan Lindemann 994c264689
All checks were successful
CI / Packaging - Kali Linux (pull_request) Successful in 4m18s
CI / Packaging - OpenSUSE Tumbleweed (pull_request) Successful in 4m20s
CI / Packaging test (pull_request) Successful in 0s
CI / Packaging - Kali Linux (push) Successful in 3m50s
CI / Packaging - OpenSUSE Tumbleweed (push) Successful in 4m20s
CI / Packaging test (push) Successful in 0s
lib.ec.ssh.AsyncSSH: Actually hide password
_connect_kwargs(hide_secrets = True) is used to log the connection
parameters when a connection fails, without leaking the password. The
filtered dictionary is built before the password is replaced with
'<hidden>', and the replacement is applied to the local kwargs
dictionary afterwards, after the filtered copy has already been made.
The dictionary that ends up in the log therefore still contains the
real password.

Hide the password before building the filtered dictionary.

Assisted-by: unsloth/Qwen3.8-27B-GGUF:Q4_K_M with pi.dev v0.84.2
Signed-off-by: Jan Lindemann <jan@janware.com>
2026-09-05 20:53:33 +02:00

458 lines
14 KiB
Python

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'] = '<hidden>'
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