Files
manage/backend/src/media_library_viewer_api/clients/ssh.py
T
2026-05-07 15:43:22 +02:00

195 lines
7.6 KiB
Python

"""SSH client helpers for remote filesystem and media inspection.
All command execution goes through ``/bin/sh -c`` and all paths inserted into
commands are shell-quoted by callers. This is important for two reasons:
1. The remote login shell may be fish/csh/etc.; internal commands are POSIX sh.
2. Media paths frequently contain spaces and punctuation.
"""
from __future__ import annotations
import json
import logging
import posixpath
import shlex
from dataclasses import dataclass
from io import StringIO
from pathlib import Path
from typing import Any
import paramiko
from media_library_viewer_api.services.known_hosts import ensure_known_host
logger = logging.getLogger(__name__)
@dataclass
class CommandResult:
"""Plain result object returned by remote command execution."""
command: str
exit_status: int
stdout: str
stderr: str
class RemoteSSHClient:
"""SSH helper for read-only inspection plus explicit job execution."""
def __init__(
self,
host: str,
username: str,
port: int = 22,
key_filename: str | None = None,
private_key: str | None = None,
private_key_passphrase: str | None = None,
password: str | None = None,
known_hosts_path: str | None = None,
timeout: int = 20,
):
if not host or not username:
raise ValueError("SSH host and username are required")
self.host = host
self.username = username
self.port = port
self.key_filename = key_filename or None
self.private_key = private_key or None
self.private_key_passphrase = private_key_passphrase or None
self.password = password or None
self.known_hosts_path = known_hosts_path or None
self.timeout = timeout
self._client: paramiko.SSHClient | None = None
def connect(self) -> paramiko.SSHClient:
"""Create or reuse the Paramiko connection.
Unknown host keys are recorded on first contact in the managed
known_hosts file when one is configured. After that, strict checking
remains in effect so host key changes are still rejected.
"""
if self._client:
return self._client
if self.known_hosts_path:
ensure_known_host(self.host, self.port, Path(self.known_hosts_path), strict=True)
client = paramiko.SSHClient()
client.load_system_host_keys()
if self.known_hosts_path and Path(self.known_hosts_path).is_file():
client.load_host_keys(self.known_hosts_path)
client.set_missing_host_key_policy(paramiko.RejectPolicy())
connect_kwargs: dict[str, Any] = {
"hostname": self.host,
"port": self.port,
"username": self.username,
"password": self.password,
"timeout": self.timeout,
}
if self.private_key:
connect_kwargs["pkey"] = self._load_private_key(self.private_key, self.private_key_passphrase)
else:
connect_kwargs["key_filename"] = self.key_filename
client.connect(**connect_kwargs)
self._client = client
return client
@staticmethod
def _load_private_key(private_key: str, passphrase: str | None = None) -> paramiko.PKey:
key_classes = [paramiko.Ed25519Key, paramiko.RSAKey, paramiko.ECDSAKey]
last_error: Exception | None = None
for key_class in key_classes:
try:
return key_class.from_private_key(StringIO(private_key), password=passphrase or None)
except Exception as exc: # pragma: no cover - try multiple algorithms
last_error = exc
raise RuntimeError("Unable to load SSH private key") from last_error
def close(self) -> None:
if self._client:
self._client.close()
self._client = None
def run(self, command: str, timeout: int | None = None) -> CommandResult:
"""Run a command through POSIX sh, independent of the user's login shell.
Paramiko asks the SSH server to execute a command using the account's
default shell. If that shell is fish/csh/etc., POSIX snippets containing
`if ...; then`, pipes, redirects, or heredocs can fail. All internal app
commands and job templates are written for POSIX shell, so explicitly
dispatch through `/bin/sh -c`.
"""
client = self.connect()
shell_command = f"/bin/sh -c {shlex.quote(command)}"
logger.debug("SSH run host=%s timeout=%s command=%s", self.host, timeout or self.timeout, command)
stdin, stdout, stderr = client.exec_command(shell_command, timeout=timeout or self.timeout)
exit_status = stdout.channel.recv_exit_status()
result = CommandResult(
command=command,
exit_status=exit_status,
stdout=stdout.read().decode(errors="replace"),
stderr=stderr.read().decode(errors="replace"),
)
if result.exit_status == 0:
logger.debug("SSH command ok host=%s exit_status=%s", self.host, result.exit_status)
else:
logger.warning(
"SSH command failed host=%s exit_status=%s stderr=%s",
self.host,
result.exit_status,
result.stderr.strip() or result.stdout.strip(),
)
return result
def list_dir(self, path: str) -> CommandResult:
"""List one remote directory as JSON.
The command first verifies that ``path`` is a directory. Without that
guard, running ``find`` on a file can look like an empty directory, which
was a source of file-browser confusion. Output is NUL-delimited before
Python serializes it, making spaces in filenames safe.
"""
quoted = shlex.quote(path)
not_dir_message = shlex.quote(f"Not a directory: {path}")
command = (
f"test -d {quoted} || "
f"{{ echo {not_dir_message} >&2; exit 20; }}; "
f"find {quoted} -maxdepth 1 -mindepth 1 -printf "
"'%y\\t%s\\t%T@\\t%f\\0' | python3 -c "
+ shlex.quote(
"import sys,json; data=sys.stdin.buffer.read().split(b'\\0'); "
"rows=[]\n"
"for row in data:\n"
" if not row: continue\n"
" t,s,m,n=row.decode('utf-8','replace').split('\\t',3)\n"
" rows.append({'type':t,'size':int(s),'mtime':float(m),'name':n})\n"
"print(json.dumps(rows))"
)
)
result = self.run(command)
logger.info("SSH list_dir path=%s exit_status=%s", path, result.exit_status)
return result
def stat_path(self, path: str) -> CommandResult:
"""Run stat for a remote file or directory path."""
quoted = shlex.quote(path)
result = self.run(f"stat --printf='%F\\n%s bytes\\n%y\\n%n\\n' {quoted}")
logger.info("SSH stat path=%s exit_status=%s", path, result.exit_status)
return result
def ffprobe_json(self, path: str) -> dict[str, Any]:
"""Run ffprobe and parse JSON output for a remote media file."""
quoted = shlex.quote(path)
result = self.run(
"ffprobe -v error -show_format -show_streams -print_format json " + quoted,
timeout=60,
)
logger.info("SSH ffprobe path=%s exit_status=%s", path, result.exit_status)
if result.exit_status != 0:
raise RuntimeError(result.stderr or result.stdout or "ffprobe failed")
return json.loads(result.stdout)
@staticmethod
def join(parent: str, child: str) -> str:
return posixpath.normpath(posixpath.join(parent, child))