From c73cfa27c3ef94c0e5eeee5049feea6f967ddb07 Mon Sep 17 00:00:00 2001 From: Alex Blank Date: Thu, 7 May 2026 20:36:37 +0200 Subject: [PATCH] fixes and improvements --- .../media_library_viewer_api/clients/ssh.py | 23 ++++++++++++++----- .../media_library_viewer_api/dependencies.py | 23 ++++++++++++++++++- .../routers/settings.py | 17 ++------------ .../services/known_hosts.py | 10 ++++++++ 4 files changed, 51 insertions(+), 22 deletions(-) diff --git a/backend/src/media_library_viewer_api/clients/ssh.py b/backend/src/media_library_viewer_api/clients/ssh.py index 0dd27ed..1b9adcb 100644 --- a/backend/src/media_library_viewer_api/clients/ssh.py +++ b/backend/src/media_library_viewer_api/clients/ssh.py @@ -20,6 +20,8 @@ from typing import Any import paramiko +from media_library_viewer_api.services.known_hosts import has_known_host + logger = logging.getLogger(__name__) @@ -64,17 +66,23 @@ class RemoteSSHClient: def connect(self) -> paramiko.SSHClient: """Create or reuse the Paramiko connection. - Host keys are expected to be managed ahead of time by startup synthesis - or explicit validation flows. Runtime connections only load the managed - known_hosts file and then let Paramiko enforce strict checking. + 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() - 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()) + 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, @@ -103,6 +111,9 @@ class RemoteSSHClient: "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 diff --git a/backend/src/media_library_viewer_api/dependencies.py b/backend/src/media_library_viewer_api/dependencies.py index c8c21c6..86d556d 100644 --- a/backend/src/media_library_viewer_api/dependencies.py +++ b/backend/src/media_library_viewer_api/dependencies.py @@ -12,7 +12,7 @@ import logging from functools import lru_cache from typing import Any -from fastapi import Request +from fastapi import HTTPException, Request from media_library_viewer_api.clients.jellyfin import JellyfinClient from media_library_viewer_api.clients.jellyseerr import JellyseerrClient @@ -82,6 +82,27 @@ def _ssh_client_for(cache_key: tuple[str, str, str, int, str, str | None, str | ) try: client.connect() + except RuntimeError as exc: + message = str(exc) + lowered = message.lower() + logger.exception("Failed to establish SSH connection to %s", host or "") + if "banner" in lowered: + raise HTTPException( + status_code=502, + detail=( + f"SSH banner not received from {host}:{port}. " + "Confirm the host, port, and firewall; the backend could not complete the SSH handshake." + ), + ) from exc + if "authentication failed" in lowered or "no authentication methods available" in lowered: + raise HTTPException( + status_code=401, + detail=( + f"SSH authentication failed for {host}:{port}. " + "Check the selected key, passphrase, username, or password." + ), + ) from exc + raise HTTPException(status_code=502, detail=message) from exc except Exception: logger.exception("Failed to establish SSH connection to %s", host or "") raise diff --git a/backend/src/media_library_viewer_api/routers/settings.py b/backend/src/media_library_viewer_api/routers/settings.py index b02772e..a8b65ff 100644 --- a/backend/src/media_library_viewer_api/routers/settings.py +++ b/backend/src/media_library_viewer_api/routers/settings.py @@ -13,7 +13,7 @@ from media_library_viewer_api.clients.ssh import RemoteSSHClient from media_library_viewer_api.config import get_settings from media_library_viewer_api.dependencies import get_monitoring_poller, get_settings_store from media_library_viewer_api.services.db_maintenance import remove_sqlite_database -from media_library_viewer_api.services.known_hosts import ensure_known_host +from media_library_viewer_api.services.known_hosts import has_known_host from media_library_viewer_api.services.media_index import MediaIndex from media_library_viewer_api.services.settings_store import SettingsStore @@ -98,20 +98,7 @@ def test_machine_ssh( client, host, port = _resolve_ssh_client(machine, store) settings = get_settings() - try: - known_hosts_updated = ensure_known_host(host, port, settings.ssh_known_hosts_file, strict=True) - except RuntimeError as exc: - message = str(exc) - lowered = message.lower() - if "protocol banner" in lowered: - raise HTTPException( - status_code=status.HTTP_502_BAD_GATEWAY, - detail=( - f"SSH banner not received from {host}:{port}; the backend could not record the host key. " - "Confirm the SSH service, host, and port are reachable." - ), - ) from exc - raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail=message) from exc + known_hosts_updated = not has_known_host(host, port, settings.ssh_known_hosts_file) try: client.connect() diff --git a/backend/src/media_library_viewer_api/services/known_hosts.py b/backend/src/media_library_viewer_api/services/known_hosts.py index 5cc594b..5b0e3a9 100644 --- a/backend/src/media_library_viewer_api/services/known_hosts.py +++ b/backend/src/media_library_viewer_api/services/known_hosts.py @@ -41,6 +41,16 @@ def _fetch_server_key(host: str, port: int, timeout: int = 30) -> paramiko.PKey: sock.close() +def has_known_host(host: str, port: int, known_hosts_path: Path) -> bool: + """Return whether the given host/port is already present in known_hosts.""" + if not host or not known_hosts_path.exists(): + return False + host_alias = _host_alias(host, port) + host_keys = paramiko.HostKeys() + host_keys.load(str(known_hosts_path)) + return host_keys.lookup(host_alias) is not None + + def ensure_known_host(host: str, port: int, known_hosts_path: Path, *, strict: bool = True) -> bool: """Ensure a host key entry exists for the given host/port.