Files
manage/backend/src/media_library_viewer_api/clients/ssh.py
T

214 lines
8.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 has_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 = 30,
):
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.
Host keys are trusted on first successful use when a managed
known_hosts path is configured. Subsequent connections stay strict and
reject host-key changes.
"""
if self._client:
return self._client
client = paramiko.SSHClient()
client.load_system_host_keys()
known_hosts_file = Path(self.known_hosts_path) if self.known_hosts_path else None
trusted_before = bool(known_hosts_file and has_known_host(self.host, self.port, known_hosts_file))
if known_hosts_file and known_hosts_file.is_file():
client.load_host_keys(str(known_hosts_file))
client.set_missing_host_key_policy(paramiko.RejectPolicy() if trusted_before else paramiko.AutoAddPolicy())
connect_kwargs: dict[str, Any] = {
"hostname": self.host,
"port": self.port,
"username": self.username,
"password": self.password,
"timeout": self.timeout,
"banner_timeout": self.timeout,
"auth_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
try:
client.connect(**connect_kwargs)
except Exception as exc:
message = str(exc).lower()
if "protocol banner" in message:
raise RuntimeError(
f"SSH banner not received from {self.host}:{self.port}. "
"Confirm the host, port, and firewall; the backend could not complete the SSH handshake."
) from exc
if "no authentication methods available" in message or "authentication failed" in message:
raise RuntimeError(
f"SSH authentication failed for {self.host}:{self.port}. "
"Check the selected key, passphrase, username, or password."
) from exc
raise
if known_hosts_file and not trusted_before:
known_hosts_file.parent.mkdir(parents=True, exist_ok=True)
client.save_host_keys(str(known_hosts_file))
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))