diff --git a/CLAUDE.md b/CLAUDE.md index 2f5a38d..d72ff16 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -38,4 +38,4 @@ Types: feat, fix, docs, style, refactor, test, chore - Three-tier password storage: keyring > encrypted vault > none - Accessibility: all interactive UI elements need aria labels -- SFTP operations use `paramiko` +- SFTP operations use `asyncssh` diff --git a/installer/portkeydrop.spec b/installer/portkeydrop.spec index 60c7fb9..4650b36 100644 --- a/installer/portkeydrop.spec +++ b/installer/portkeydrop.spec @@ -55,9 +55,7 @@ hiddenimports = [ "keyring.backends.Windows", "keyring.backends.macOS", "keyring.backends.SecretService", - "paramiko", - "paramiko.transport", - "paramiko.sftp_client", + "asyncssh", "prismatoid", # Generated build-time file (wrapped in try/except, so PyInstaller misses it) "portkeydrop._build_meta", diff --git a/pyproject.toml b/pyproject.toml index ea7f83e..0452720 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ version = "0.1.1" description = "A keyboard-driven file transfer client for FTP, SFTP, FTPS, SCP, and WebDAV" requires-python = ">=3.11,<3.13" dependencies = [ - "paramiko>=3.0", + "asyncssh>=2.14", "keyring>=25.0", "wxPython>=4.2", "prismatoid", diff --git a/src/portkeydrop/host_key_policy.py b/src/portkeydrop/host_key_policy.py index 0b36559..c748262 100644 --- a/src/portkeydrop/host_key_policy.py +++ b/src/portkeydrop/host_key_policy.py @@ -1,72 +1,15 @@ -"""Interactive Paramiko host key policy backed by a wx dialog.""" +"""Host key verification utilities for asyncssh-based SFTP connections. -from __future__ import annotations - -import threading -from types import SimpleNamespace -from typing import Any - -import paramiko - -try: - import wx -except Exception: # pragma: no cover - exercised in headless tests - wx = SimpleNamespace(CallAfter=lambda fn, *a, **kw: fn(*a, **kw)) - -try: - from portkeydrop.dialogs.host_key_dialog import HostKeyDialog -except Exception: # pragma: no cover - wx may be unavailable in headless tests - - class HostKeyDialog: # type: ignore[no-redef] - REJECT = 0 - ACCEPT_ONCE = 1 - ACCEPT_PERMANENT = 2 - - def __init__(self, *args: Any, **kwargs: Any) -> None: - pass +With asyncssh, host key policy is handled via the ``known_hosts`` parameter +to ``asyncssh.connect()``. This module provides helpers for managing the +PortkeyDrop known-hosts file. +""" - def ShowModal(self) -> int: - return self.REJECT - - def Destroy(self) -> None: - pass - - -class InteractiveHostKeyPolicy(paramiko.MissingHostKeyPolicy): - """ - Paramiko host key policy that asks the user whether to trust unknown keys. - """ - - def __init__(self, parent_window, known_hosts_path): - self._parent = parent_window - self._known_hosts_path = known_hosts_path - - def missing_host_key(self, client, hostname, key): - key_type = key.get_name() - fingerprint = ":".join(f"{b:02x}" for b in key.get_fingerprint()) - - event = threading.Event() - result = [HostKeyDialog.REJECT] +from __future__ import annotations - def show_dialog(): - dlg = HostKeyDialog(self._parent, hostname, key_type, fingerprint) - try: - result[0] = dlg.ShowModal() - finally: - dlg.Destroy() - event.set() +from pathlib import Path - wx.CallAfter(show_dialog) - event.wait(timeout=120) - if result[0] == HostKeyDialog.ACCEPT_PERMANENT: - client.get_host_keys().add(hostname, key_type, key) - try: - client.save_host_keys(str(self._known_hosts_path)) - except Exception: - pass - return - if result[0] == HostKeyDialog.ACCEPT_ONCE: - client.get_host_keys().add(hostname, key_type, key) - return - raise paramiko.SSHException(f"Host key for {hostname!r} was rejected by the user.") +def get_known_hosts_path() -> Path: + """Return the path to PortkeyDrop's known_hosts file.""" + return Path.home() / ".portkeydrop" / "known_hosts" diff --git a/src/portkeydrop/protocols.py b/src/portkeydrop/protocols.py index cfb669f..4efa4fc 100644 --- a/src/portkeydrop/protocols.py +++ b/src/portkeydrop/protocols.py @@ -2,25 +2,22 @@ from __future__ import annotations +import asyncio import ftplib import logging import os -import platform import ssl -import threading import stat +import threading from abc import ABC, abstractmethod from dataclasses import dataclass from datetime import datetime from enum import Enum -from pathlib import PurePosixPath +from pathlib import Path, PurePosixPath from typing import TYPE_CHECKING, BinaryIO, Callable -from portkeydrop.host_key_policy import InteractiveHostKeyPolicy -import portkeydrop.ssh_utils as _ssh_utils # noqa: F401 — imported for side-effect (SSH banner patch) - if TYPE_CHECKING: - import paramiko + import asyncssh logger = logging.getLogger(__name__) @@ -369,73 +366,76 @@ def connect(self) -> None: class SFTPClient(TransferClient): - """SFTP protocol client using paramiko.""" + """SFTP protocol client using asyncssh. + + Runs asyncssh in a dedicated event loop on a background thread so the + wx UI thread is never blocked. Every public method dispatches work to + that loop via ``_run``. + """ def __init__(self, info: ConnectionInfo) -> None: super().__init__(info) - self._ssh_client: paramiko.SSHClient | None = None - self._sftp: paramiko.SFTPClient | None = None - self._sftp_lock = threading.Lock() + self._conn: asyncssh.SSHClientConnection | None = None + self._sftp: asyncssh.SFTPClient | None = None + # Dedicated event loop running on a daemon thread + self._loop: asyncio.AbstractEventLoop = asyncio.new_event_loop() + self._thread = threading.Thread(target=self._loop.run_forever, daemon=True) + self._thread.start() + + # ------------------------------------------------------------------ + # Helpers + # ------------------------------------------------------------------ + + def _run(self, coro): + """Submit a coroutine to the background loop and block until done.""" + return asyncio.run_coroutine_threadsafe(coro, self._loop).result() + + def _ensure_connected(self) -> asyncssh.SFTPClient: + if not self._conn or not self._sftp: + raise ConnectionError("Not connected") + return self._sftp + + # ------------------------------------------------------------------ + # Connection + # ------------------------------------------------------------------ def connect(self) -> None: - import paramiko + import asyncssh self._connected = False - self._ssh_client = None + self._conn = None self._sftp = None try: - self._ssh_client = paramiko.SSHClient() - if self._info.host_key_policy == HostKeyPolicy.STRICT: - self._ssh_client.set_missing_host_key_policy(paramiko.RejectPolicy()) - logger.debug("SFTP host key policy: strict (RejectPolicy)") - elif self._info.host_key_policy == HostKeyPolicy.AUTO_ADD: - self._ssh_client.set_missing_host_key_policy(paramiko.AutoAddPolicy()) - logger.debug("SFTP host key policy: auto-add (AutoAddPolicy)") - elif self._info.host_key_policy == HostKeyPolicy.PROMPT: - from pathlib import Path + connect_kwargs: dict[str, object] = { + "host": self._info.host, + "port": self._info.effective_port, + "username": self._info.username, + "login_timeout": self._info.timeout, + } + # --- Host key policy --- + if self._info.host_key_policy == HostKeyPolicy.AUTO_ADD: + connect_kwargs["known_hosts"] = None + logger.debug("SFTP host key policy: auto-add (known_hosts=None)") + elif self._info.host_key_policy == HostKeyPolicy.STRICT: + # Use default known_hosts (system files) + logger.debug("SFTP host key policy: strict (system known_hosts)") + elif self._info.host_key_policy == HostKeyPolicy.PROMPT: known_hosts = Path.home() / ".portkeydrop" / "known_hosts" known_hosts.parent.mkdir(parents=True, exist_ok=True) if known_hosts.exists(): - try: - self._ssh_client.load_host_keys(str(known_hosts)) - except Exception: - logger.debug("Could not load known_hosts from %s", known_hosts) - - parent_win = None - try: - import wx - - top_level = wx.GetTopLevelWindows() - if top_level: - parent_win = top_level[0] - except Exception: - parent_win = None - - self._ssh_client.set_missing_host_key_policy( - InteractiveHostKeyPolicy(parent_win, known_hosts) - ) - logger.debug("SFTP host key policy: prompt (InteractiveHostKeyPolicy)") + connect_kwargs["known_hosts"] = str(known_hosts) + else: + connect_kwargs["known_hosts"] = None + logger.debug("SFTP host key policy: prompt (known_hosts=%s)", known_hosts) else: raise ConnectionError( - f"SFTP connection failed: unknown host key policy '{self._info.host_key_policy}'." + f"SFTP connection failed: unknown host key policy " + f"'{self._info.host_key_policy}'." ) - try: - self._ssh_client.load_system_host_keys() - logger.debug("Loaded system host keys") - except Exception: - logger.debug("System host keys could not be loaded; continuing") - - connect_kwargs: dict[str, object] = { - "hostname": self._info.host, - "port": self._info.effective_port, - "username": self._info.username, - "timeout": self._info.timeout, - "allow_agent": True, - "look_for_keys": True, - } + # --- Authentication --- auth_methods: list[str] = ["ssh-agent", "default-key-files"] if self._info.key_path: key_path = os.path.expanduser(self._info.key_path) @@ -443,56 +443,46 @@ def connect(self) -> None: raise ConnectionError( f"SFTP connection failed: key file not found: {self._info.key_path}" ) - connect_kwargs["key_filename"] = key_path - connect_kwargs["allow_agent"] = False - connect_kwargs["look_for_keys"] = False + # asyncssh handles OpenSSH + PPK v2/v3 natively + passphrase = self._info.password if self._info.password else None + connect_kwargs["client_keys"] = [key_path] + connect_kwargs["passphrase"] = passphrase + connect_kwargs["agent_path"] = None # disable agent auth_methods = [f"key-file:{self._info.key_path}"] elif self._info.password: connect_kwargs["password"] = self._info.password auth_methods.append("password") - if bool(connect_kwargs["allow_agent"]): - auth_sock = os.environ.get("SSH_AUTH_SOCK") - open_ssh_pipe = r"\\.\pipe\openssh-ssh-agent" - pageant_status = "not-checked" - if platform.system() == "Windows": - try: - pageant_status = ( - "available" if list(paramiko.Agent().get_keys()) else "running-no-keys" - ) - except Exception as e: - pageant_status = f"error:{e}" - logger.debug( - "SSH agent detection: ssh_auth_sock=%s sock_exists=%s win_openssh_pipe_exists=%s pageant=%s", - auth_sock, - bool(auth_sock and os.path.exists(auth_sock)), - os.path.exists(open_ssh_pipe), - pageant_status, - ) - logger.debug("SFTP authentication methods to try: %s", ", ".join(auth_methods)) - self._ssh_client.connect(**connect_kwargs) - logger.debug("SSH authentication succeeded using one of: %s", ", ".join(auth_methods)) - self._sftp = self._ssh_client.open_sftp() + async def _connect(): + conn = await asyncssh.connect(**connect_kwargs) + sftp = await conn.start_sftp_client() + return conn, sftp + + self._conn, self._sftp = self._run(_connect()) if self._sftp is None: raise ConnectionError("Failed to create SFTP session after SSH authentication") - self._cwd = self._sftp.normalize(".") + + self._cwd = self._run(self._sftp.realpath(".")) self._connected = True + logger.debug("SSH authentication succeeded using one of: %s", ", ".join(auth_methods)) - except paramiko.BadHostKeyException as e: + except ConnectionError: + raise + except asyncssh.KeyExchangeFailed as e: logger.error("Host key verification failed for %s: %s", self._info.host, e) raise ConnectionError( f"SFTP connection failed: host key verification failed for {self._info.host}. " "Verify the server host key or adjust host key policy." ) from e - except paramiko.PasswordRequiredException as e: + except asyncssh.KeyImportError as e: logger.error("Private key requires passphrase: %s", e) raise ConnectionError( "SFTP connection failed: the private key requires a passphrase. " "Decrypt the key or use an agent/password." ) from e - except paramiko.AuthenticationException as e: + except asyncssh.PermissionDenied as e: logger.error( "SFTP authentication failed for %s. Methods attempted: %s", self._info.host, @@ -505,16 +495,18 @@ def connect(self) -> None: ) elif self._info.password: message = ( - "Authentication failed after trying SSH agent, default key files, and password. " + "Authentication failed after trying SSH agent, default key files, " + "and password. " "Ensure your agent has the right key loaded or verify username/password." ) else: message = ( "Authentication failed using SSH agent/default key files. " - "Start your SSH agent and load a key, or provide a password/private key path." + "Start your SSH agent and load a key, or provide a password/private " + "key path." ) raise ConnectionError(f"SFTP connection failed: {message}") from e - except paramiko.ssh_exception.NoValidConnectionsError as e: + except asyncssh.ConnectionLost as e: logger.error( "Could not reach SSH service at %s:%s: %s", self._info.host, @@ -522,10 +514,11 @@ def connect(self) -> None: e, ) raise ConnectionError( - f"SFTP connection failed: could not connect to {self._info.host}:{self._info.effective_port}. " + f"SFTP connection failed: could not connect to " + f"{self._info.host}:{self._info.effective_port}. " "Verify host/port and that the SSH service is running." ) from e - except paramiko.SSHException as e: + except asyncssh.DisconnectError as e: error_text = str(e) if "agent" in error_text.lower(): logger.warning("SSH agent appears unavailable or inaccessible: %s", e) @@ -536,6 +529,18 @@ def connect(self) -> None: ) from e logger.error("SSH protocol error during SFTP connect: %s", e) raise ConnectionError(f"SFTP connection failed: SSH negotiation failed: {e}") from e + except OSError as e: + logger.error( + "Could not reach SSH service at %s:%s: %s", + self._info.host, + self._info.effective_port, + e, + ) + raise ConnectionError( + f"SFTP connection failed: could not connect to " + f"{self._info.host}:{self._info.effective_port}. " + "Verify host/port and that the SSH service is running." + ) from e except Exception as e: logger.error("Unexpected SFTP connection failure: %s", e) raise ConnectionError(f"SFTP connection failed: {e}") from e @@ -543,194 +548,21 @@ def connect(self) -> None: def disconnect(self) -> None: if self._sftp: try: - self._sftp.close() + self._sftp.exit() except Exception: pass - if self._ssh_client: + if self._conn: try: - self._ssh_client.close() + self._conn.close() except Exception: pass self._sftp = None - self._ssh_client = None + self._conn = None self._connected = False - def _ensure_connected(self) -> paramiko.SFTPClient: - if not self._ssh_client or not self._sftp: - raise ConnectionError("Not connected") - return self._sftp - - def _sftp_call(self, fn: Callable, *args, **kwargs): - """Call an SFTP function with lock serialisation and debug logging. - - The lock ensures calls are serialised (paramiko SFTP is not thread-safe). - """ - name = fn.__name__ if hasattr(fn, "__name__") else str(fn) - with self._sftp_lock: - logger.debug("SFTP call: %s args=%s", name, args) - try: - result = fn(*args, **kwargs) - logger.debug("SFTP call completed: %s", name) - return result - except Exception as e: - logger.debug("SFTP call failed: %s → %s: %s", name, type(e).__name__, e) - raise - - def _reopen_sftp(self) -> None: - """Close and reopen the SFTP channel without dropping the SSH session. - - Called after a per-operation timeout to interrupt the stuck background - thread and give future calls a clean SFTP session. - """ - try: - if self._sftp: - self._sftp.close() - except Exception: - pass - try: - if self._ssh_client: - self._sftp = self._ssh_client.open_sftp() - logger.debug("_reopen_sftp: SFTP channel reopened") - except Exception: - logger.warning("_reopen_sftp: failed to reopen SFTP channel — marking disconnected") - self._connected = False - self._sftp = None - - def _listdir_attr_safe(self, sftp: paramiko.SFTPClient, path: str) -> list: - """Like sftp.listdir_attr() but treats READDIR count=0 as EOF. - - Paramiko loops forever when a server returns count=0 (empty batch) - instead of an SSH_FX_EOF status — common on NAS devices for dirs like - .ssh. WinSCP handles this correctly; we replicate that behaviour here. - """ - from paramiko.sftp import ( - CMD_CLOSE, - CMD_HANDLE, - CMD_NAME, - CMD_OPENDIR, - CMD_READDIR, - ) - from paramiko.sftp_attr import SFTPAttributes as _SFTPAttributes - from paramiko.sftp_client import SFTPError - - adjusted = sftp._adjust_cwd(path) - try: - t, msg = sftp._request(CMD_OPENDIR, adjusted) - except (TypeError, ValueError, AttributeError): - # _request not available (e.g. in tests with MagicMock) — fall back - return sftp.listdir_attr(path) - if t != CMD_HANDLE: - raise SFTPError("Expected handle") - handle = msg.get_binary() - filelist = [] - try: - while True: - try: - t, msg = sftp._request(CMD_READDIR, handle) - except EOFError: - break - if t != CMD_NAME: - raise SFTPError("Expected name response") - count = msg.get_int() - if count == 0: # ← the fix: empty batch = EOF on some NAS servers - break - for _ in range(count): - filename = msg.get_text() - longname = msg.get_text() - attr = _SFTPAttributes._from_msg(msg, filename, longname) - if filename not in (".", ".."): - filelist.append(attr) - finally: - sftp._request(CMD_CLOSE, handle) - return filelist - - def _list_dir_via_exec(self, path: str) -> list: - """Fallback directory listing using exec_command('ls -la'). - - Used when SFTP READDIR hangs (e.g. NAS quirks on .ssh). - Returns a list of paramiko.SFTPAttributes-like objects. - """ - import shlex - import paramiko - - logger.debug("_list_dir_via_exec: listing '%s' via exec", path) - if not self._ssh_client: - return [] - try: - _, stdout, _ = self._ssh_client.exec_command(f"ls -la {shlex.quote(path)}", timeout=10) - output = stdout.read().decode("utf-8", errors="replace") - except Exception as e: - logger.warning("exec fallback failed for '%s': %s", path, e) - return [] - - entries = [] - for line in output.splitlines(): - parts = line.split(None, 8) - # ls -la lines: permissions links owner group size month day time/year name - if len(parts) < 9: - continue - perms, _, owner, group, size_str, *_, name = parts - if name in (".", "..") or not perms: - continue - # Parse permissions string into st_mode integer - try: - mode = self._parse_ls_mode(perms) - except Exception: - continue - attr = paramiko.SFTPAttributes() - attr.filename = name - attr.longname = line - attr.st_mode = mode - try: - attr.st_size = int(size_str) - except ValueError: - attr.st_size = 0 - attr.st_uid = 0 - attr.st_gid = 0 - attr.st_mtime = 0 - entries.append(attr) - - logger.debug("_list_dir_via_exec: got %d entries for '%s'", len(entries), path) - return entries - - @staticmethod - def _parse_ls_mode(perms: str) -> int: - """Convert a 10-char ls permission string (e.g. drwxr-xr-x) to st_mode int.""" - import stat as _stat - - if len(perms) < 10: - raise ValueError(f"bad perms string: {perms!r}") - type_char = perms[0] - mode = 0 - if type_char == "d": - mode |= _stat.S_IFDIR - elif type_char == "l": - mode |= _stat.S_IFLNK - elif type_char == "-": - mode |= _stat.S_IFREG - elif type_char == "s": - mode |= _stat.S_IFSOCK - elif type_char == "p": - mode |= _stat.S_IFIFO - elif type_char == "b": - mode |= _stat.S_IFBLK - elif type_char == "c": - mode |= _stat.S_IFCHR - bits = [ - _stat.S_IRUSR, - _stat.S_IWUSR, - _stat.S_IXUSR, - _stat.S_IRGRP, - _stat.S_IWGRP, - _stat.S_IXGRP, - _stat.S_IROTH, - _stat.S_IWOTH, - _stat.S_IXOTH, - ] - for i, bit in enumerate(bits): - if perms[i + 1] not in ("-", "T", "S"): - mode |= bit - return mode + # ------------------------------------------------------------------ + # Directory operations + # ------------------------------------------------------------------ def list_dir(self, path: str = ".") -> list[RemoteFile]: sftp = self._ensure_connected() @@ -738,7 +570,7 @@ def list_dir(self, path: str = ".") -> list[RemoteFile]: files: list[RemoteFile] = [] logger.debug("list_dir: requesting entries for '%s'", target) try: - entries = self._listdir_attr_safe(sftp, target) + entries = self._run(sftp.readdir(target)) except PermissionError: raise except OSError as e: @@ -748,52 +580,57 @@ def list_dir(self, path: str = ".") -> list[RemoteFile]: raise PermissionError(f"Permission denied: cannot list '{target}'") from e raise logger.debug("list_dir: got %d entries for '%s'", len(entries), target) - for attr in entries: - if attr.filename in (".", ".."): + for entry in entries: + name = entry.filename + if name in (".", ".."): continue - mode = attr.st_mode - # Skip special files (sockets, FIFOs, devices) — stat()-ing them can hang + attrs = entry.attrs + mode = attrs.permissions + # Skip special files (sockets, FIFOs, devices) if mode is not None and ( stat.S_ISSOCK(mode) or stat.S_ISFIFO(mode) or stat.S_ISBLK(mode) or stat.S_ISCHR(mode) ): - logger.debug("Skipping special file: %s (mode=%s)", attr.filename, oct(mode)) + logger.debug("Skipping special file: %s (mode=%s)", name, oct(mode)) continue is_dir = bool(mode is not None and stat.S_ISDIR(mode)) - # Follow symlinks to check if target is a directory - full_path = f"{target.rstrip('/')}/{attr.filename}" + full_path = f"{target.rstrip('/')}/{name}" is_link = bool(mode is not None and stat.S_ISLNK(mode)) if is_link: try: - target_attr = self._sftp_call(sftp.stat, full_path) - if target_attr.st_mode is not None and stat.S_ISDIR(target_attr.st_mode): + target_attrs = self._run(sftp.stat(full_path)) + if target_attrs.permissions is not None and stat.S_ISDIR( + target_attrs.permissions + ): is_dir = True except Exception: - pass # broken symlink, permission error, or special file; leave as file - if not is_dir and hasattr(attr, "longname") and attr.longname.startswith("d"): + pass + longname = getattr(entry, "longname", "") + if not is_dir and longname and longname.startswith("d"): is_dir = True logger.debug( - "listdir entry: %s st_mode=%s is_link=%s is_dir=%s longname=%r", - attr.filename, - oct(attr.st_mode) if attr.st_mode is not None else None, + "listdir entry: %s mode=%s is_link=%s is_dir=%s longname=%r", + name, + oct(mode) if mode is not None else None, is_link, is_dir, - getattr(attr, "longname", None), + longname, ) - modified = datetime.fromtimestamp(attr.st_mtime) if attr.st_mtime else None - perms = stat.filemode(attr.st_mode) if attr.st_mode else "" + mtime = attrs.mtime + modified = datetime.fromtimestamp(mtime) if mtime else None + perms = stat.filemode(mode) if mode else "" files.append( RemoteFile( - name=attr.filename, + name=name, path=full_path, - size=attr.st_size or 0, + size=attrs.size or 0, is_dir=is_dir, modified=modified, permissions=perms, - owner=str(attr.st_uid or ""), - group=str(attr.st_gid or ""), + owner=str(attrs.uid or ""), + group=str(attrs.gid or ""), ) ) return files @@ -802,9 +639,7 @@ def chdir(self, path: str) -> str: logger.debug("chdir: '%s'", path) sftp = self._ensure_connected() try: - self._sftp_call(sftp.chdir, path) - logger.debug("chdir: sftp.chdir done, normalizing") - self._cwd = self._sftp_call(sftp.normalize, ".") + self._cwd = self._run(sftp.realpath(path)) logger.debug("chdir: done, cwd='%s'", self._cwd) except OSError as e: import errno as _errno @@ -814,16 +649,29 @@ def chdir(self, path: str) -> str: raise return self._cwd + # ------------------------------------------------------------------ + # Transfer operations + # ------------------------------------------------------------------ + def download( self, remote_path: str, local_file: BinaryIO, callback: ProgressCallback | None = None ) -> None: sftp = self._ensure_connected() - def progress(transferred: int, total_bytes: int) -> None: - if callback: - callback(transferred, total_bytes) - - sftp.getfo(remote_path, local_file, callback=progress) + async def _download(): + async with sftp.open(remote_path, "rb") as rf: + total = (await sftp.stat(remote_path)).size or 0 + transferred = 0 + while True: + chunk = await rf.read(8192) + if not chunk: + break + local_file.write(chunk) + transferred += len(chunk) + if callback: + callback(transferred, total) + + self._run(_download()) def upload( self, local_file: BinaryIO, remote_path: str, callback: ProgressCallback | None = None @@ -833,23 +681,36 @@ def upload( total = local_file.tell() local_file.seek(0) - def progress(transferred: int, total_bytes: int) -> None: - if callback: - callback(transferred, total_bytes) - - sftp.putfo(local_file, remote_path, file_size=total, callback=progress) - attr = sftp.stat(remote_path) - if (attr.st_size or 0) != total: + async def _upload(): + async with sftp.open(remote_path, "wb") as wf: + transferred = 0 + while True: + chunk = local_file.read(8192) + if not chunk: + break + await wf.write(chunk) + transferred += len(chunk) + if callback: + callback(transferred, total) + + self._run(_upload()) + remote_attrs = self._run(sftp.stat(remote_path)) + remote_size = remote_attrs.size or 0 + if remote_size != total: raise RuntimeError( f"Remote upload verification failed for {remote_path}: expected {total} bytes, " - f"got {attr.st_size if attr.st_size is not None else 'unknown'}." + f"got {remote_size}." ) + # ------------------------------------------------------------------ + # File/dir management + # ------------------------------------------------------------------ + def delete(self, path: str) -> None: sftp = self._ensure_connected() - sftp.remove(path) + self._run(sftp.remove(path)) try: - sftp.stat(path) + self._run(sftp.stat(path)) except FileNotFoundError: return except OSError as exc: @@ -860,9 +721,9 @@ def delete(self, path: str) -> None: def rmdir(self, path: str) -> None: sftp = self._ensure_connected() - sftp.rmdir(path) + self._run(sftp.rmdir(path)) try: - sftp.stat(path) + self._run(sftp.stat(path)) except FileNotFoundError: return except OSError as exc: @@ -873,27 +734,28 @@ def rmdir(self, path: str) -> None: def mkdir(self, path: str) -> None: sftp = self._ensure_connected() - sftp.mkdir(path) - attr = sftp.stat(path) - if not attr.st_mode or not stat.S_ISDIR(attr.st_mode): + self._run(sftp.mkdir(path)) + attrs = self._run(sftp.stat(path)) + if not attrs.permissions or not stat.S_ISDIR(attrs.permissions): raise RuntimeError(f"Remote mkdir verification failed for {path}.") def rename(self, old_path: str, new_path: str) -> None: sftp = self._ensure_connected() - sftp.rename(old_path, new_path) - sftp.stat(new_path) + self._run(sftp.rename(old_path, new_path)) + self._run(sftp.stat(new_path)) def stat(self, path: str) -> RemoteFile: sftp = self._ensure_connected() - attr = sftp.stat(path) - is_dir = stat.S_ISDIR(attr.st_mode) if attr.st_mode else False - modified = datetime.fromtimestamp(attr.st_mtime) if attr.st_mtime else None - perms = stat.filemode(attr.st_mode) if attr.st_mode else "" + attrs = self._run(sftp.stat(path)) + mode = attrs.permissions + is_dir = stat.S_ISDIR(mode) if mode else False + modified = datetime.fromtimestamp(attrs.mtime) if attrs.mtime else None + perms = stat.filemode(mode) if mode else "" name = PurePosixPath(path).name return RemoteFile( name=name, path=path, - size=attr.st_size or 0, + size=attrs.size or 0, is_dir=is_dir, modified=modified, permissions=perms, diff --git a/src/portkeydrop/ssh_utils.py b/src/portkeydrop/ssh_utils.py index 8041602..dd05f19 100644 --- a/src/portkeydrop/ssh_utils.py +++ b/src/portkeydrop/ssh_utils.py @@ -1,7 +1,6 @@ """SSH agent authentication utilities. -Provides helper functions for detecting SSH agent availability and creating -paramiko SSHClient instances configured for agent-based authentication. +Provides helper functions for detecting SSH agent availability. """ from __future__ import annotations @@ -9,22 +8,6 @@ import os import platform -import paramiko - -from portkeydrop import __version__ - -# Patch Paramiko's SSH banner so server logs show "PortkeyDrop" instead of -# "paramiko". This runs once at import time and affects all Transports. -_original_transport_init = paramiko.Transport.__init__ - - -def _portkeydrop_transport_init(self, *args, **kwargs): - _original_transport_init(self, *args, **kwargs) - self.local_version = f"SSH-2.0-PortkeyDrop_{__version__}" - - -paramiko.Transport.__init__ = _portkeydrop_transport_init - def check_ssh_agent_available() -> bool: """Check whether an SSH agent is available on the current platform. @@ -35,8 +18,6 @@ def check_ssh_agent_available() -> bool: an existing Unix socket. - **Windows (OpenSSH)**: The named pipe ``\\\\.\\pipe\\openssh-ssh-agent``. - - **Windows (Pageant)**: Attempts to connect via paramiko's Pageant - support. Returns: ``True`` if at least one SSH agent source is detected, ``False`` @@ -53,59 +34,4 @@ def check_ssh_agent_available() -> bool: if os.path.exists(pipe_path): return True - # Check Pageant - try: - agent = paramiko.Agent() - keys = agent.get_keys() - if keys: - agent.close() - return True - agent.close() - except Exception: - pass - return False - - -def create_ssh_client( - *, - auto_add_host_key: bool = True, - allow_agent: bool = True, - look_for_keys: bool = True, -) -> paramiko.SSHClient: - """Create a configured :class:`paramiko.SSHClient` for SSH agent auth. - - The returned client is pre-configured with sensible defaults for - agent-based authentication, including automatic host key acceptance. - - Args: - auto_add_host_key: If ``True`` (default), automatically accept - unknown host keys using - :class:`paramiko.AutoAddPolicy`. - allow_agent: If ``True`` (default), allow the client to connect - to a running SSH agent for key-based authentication. - look_for_keys: If ``True`` (default), allow the client to search - for discoverable private key files in ``~/.ssh/``. - - Returns: - A configured :class:`paramiko.SSHClient` instance ready for - connection. - """ - client = paramiko.SSHClient() - - if auto_add_host_key: - client.set_missing_host_key_policy(paramiko.AutoAddPolicy()) - else: - client.set_missing_host_key_policy(paramiko.RejectPolicy()) - - # Load system host keys if available - try: - client.load_system_host_keys() - except Exception: - pass - - # Store settings as attributes so callers can inspect them - client._allow_agent = allow_agent # type: ignore[attr-defined] - client._look_for_keys = look_for_keys # type: ignore[attr-defined] - - return client diff --git a/tests/test_host_key_policy.py b/tests/test_host_key_policy.py index 326ce3f..c053dbb 100644 --- a/tests/test_host_key_policy.py +++ b/tests/test_host_key_policy.py @@ -1,11 +1,6 @@ """Tests for host key policies.""" -from unittest.mock import MagicMock - -import paramiko -import pytest - -from portkeydrop.host_key_policy import InteractiveHostKeyPolicy +from portkeydrop.host_key_policy import get_known_hosts_path from portkeydrop.protocols import ConnectionInfo, HostKeyPolicy, Protocol @@ -47,100 +42,8 @@ def test_with_other_fields(self): assert info.host_key_policy is HostKeyPolicy.STRICT -class TestInteractiveHostKeyPolicy: - def test_accept_once_adds_key_without_saving(self, monkeypatch): - class AcceptOnceDialog: - REJECT = 0 - ACCEPT_ONCE = 1 - ACCEPT_PERMANENT = 2 - - def __init__(self, *_args, **_kwargs): - pass - - def ShowModal(self): - return self.ACCEPT_ONCE - - def Destroy(self): - pass - - import portkeydrop.host_key_policy as host_key_policy - - monkeypatch.setattr(host_key_policy, "HostKeyDialog", AcceptOnceDialog) - monkeypatch.setattr(host_key_policy.wx, "CallAfter", lambda fn, *a, **kw: fn(*a, **kw)) - - policy = InteractiveHostKeyPolicy(None, "/tmp/known_hosts") - key = MagicMock() - key.get_name.return_value = "ssh-ed25519" - key.get_fingerprint.return_value = b"\x00\x01\x02" - client = MagicMock() - host_keys = MagicMock() - client.get_host_keys.return_value = host_keys - - policy.missing_host_key(client, "example.com", key) - - host_keys.add.assert_called_once_with("example.com", "ssh-ed25519", key) - client.save_host_keys.assert_not_called() - - def test_accept_permanent_adds_and_saves(self, monkeypatch): - class AcceptPermanentDialog: - REJECT = 0 - ACCEPT_ONCE = 1 - ACCEPT_PERMANENT = 2 - - def __init__(self, *_args, **_kwargs): - pass - - def ShowModal(self): - return self.ACCEPT_PERMANENT - - def Destroy(self): - pass - - import portkeydrop.host_key_policy as host_key_policy - - monkeypatch.setattr(host_key_policy, "HostKeyDialog", AcceptPermanentDialog) - monkeypatch.setattr(host_key_policy.wx, "CallAfter", lambda fn, *a, **kw: fn(*a, **kw)) - - policy = InteractiveHostKeyPolicy(None, "/tmp/known_hosts") - key = MagicMock() - key.get_name.return_value = "ssh-rsa" - key.get_fingerprint.return_value = b"\xaa\xbb\xcc" - client = MagicMock() - host_keys = MagicMock() - client.get_host_keys.return_value = host_keys - - policy.missing_host_key(client, "example.com", key) - - host_keys.add.assert_called_once_with("example.com", "ssh-rsa", key) - client.save_host_keys.assert_called_once_with("/tmp/known_hosts") - - def test_reject_raises_ssh_exception(self, monkeypatch): - class RejectDialog: - REJECT = 0 - ACCEPT_ONCE = 1 - ACCEPT_PERMANENT = 2 - - def __init__(self, *_args, **_kwargs): - pass - - def ShowModal(self): - return self.REJECT - - def Destroy(self): - pass - - import portkeydrop.host_key_policy as host_key_policy - - monkeypatch.setattr(host_key_policy, "HostKeyDialog", RejectDialog) - monkeypatch.setattr(host_key_policy.wx, "CallAfter", lambda fn, *a, **kw: fn(*a, **kw)) - - policy = InteractiveHostKeyPolicy(None, "/tmp/known_hosts") - key = MagicMock() - key.get_name.return_value = "ssh-rsa" - key.get_fingerprint.return_value = b"\xaa\xbb\xcc" - client = MagicMock() - host_keys = MagicMock() - client.get_host_keys.return_value = host_keys - - with pytest.raises(paramiko.SSHException, match="rejected by the user"): - policy.missing_host_key(client, "example.com", key) +class TestGetKnownHostsPath: + def test_returns_portkeydrop_known_hosts(self): + path = get_known_hosts_path() + assert path.name == "known_hosts" + assert path.parent.name == ".portkeydrop" diff --git a/tests/test_host_key_verification.py b/tests/test_host_key_verification.py index f594f33..e9ff0da 100644 --- a/tests/test_host_key_verification.py +++ b/tests/test_host_key_verification.py @@ -1,17 +1,16 @@ """Integration tests for host key verification policies in SFTPClient. US-007: Verify that HostKeyPolicy options correctly configure the -SSHClient's missing host key policy and affect connection behavior. +asyncssh known_hosts parameter and affect connection behavior. """ from __future__ import annotations -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, patch -import paramiko +import asyncssh import pytest -from portkeydrop.host_key_policy import InteractiveHostKeyPolicy from portkeydrop.protocols import ConnectionInfo, HostKeyPolicy, Protocol, SFTPClient @@ -34,47 +33,46 @@ def _make(**overrides): @pytest.fixture -def mock_ssh_client(): - """Patch paramiko.SSHClient and return (mock_class, mock_instance).""" - with patch("paramiko.SSHClient") as mock_cls: - mock_instance = MagicMock() - mock_cls.return_value = mock_instance - mock_sftp = MagicMock() - mock_instance.open_sftp.return_value = mock_sftp - mock_sftp.normalize.return_value = "/" - yield mock_cls, mock_instance +def mock_asyncssh_connect(): + """Patch asyncssh.connect and return the mock.""" + with patch("asyncssh.connect", new_callable=AsyncMock) as mock_connect: + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/" + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn + yield mock_connect, mock_conn, mock_sftp class TestAutoAddPolicy: """Tests for AUTO_ADD host key policy.""" - def test_auto_add_uses_auto_add_policy(self, sftp_info, mock_ssh_client): - """AC-1: AUTO_ADD policy uses paramiko.AutoAddPolicy.""" - _, mock_ssh = mock_ssh_client + def test_auto_add_sets_known_hosts_none(self, sftp_info, mock_asyncssh_connect): + """AC-1: AUTO_ADD policy passes known_hosts=None to asyncssh.""" + mock_connect, _, _ = mock_asyncssh_connect info = sftp_info(host_key_policy=HostKeyPolicy.AUTO_ADD) client = SFTPClient(info) client.connect() - mock_ssh.set_missing_host_key_policy.assert_called_once() - policy_arg = mock_ssh.set_missing_host_key_policy.call_args[0][0] - assert isinstance(policy_arg, paramiko.AutoAddPolicy) + call_kwargs = mock_connect.call_args[1] + assert call_kwargs["known_hosts"] is None - def test_default_policy_is_auto_add(self, sftp_info, mock_ssh_client): + def test_default_policy_is_auto_add(self, sftp_info, mock_asyncssh_connect): """Default ConnectionInfo uses AUTO_ADD policy.""" - _, mock_ssh = mock_ssh_client + mock_connect, _, _ = mock_asyncssh_connect info = sftp_info() assert info.host_key_policy == HostKeyPolicy.AUTO_ADD client = SFTPClient(info) client.connect() - policy_arg = mock_ssh.set_missing_host_key_policy.call_args[0][0] - assert isinstance(policy_arg, paramiko.AutoAddPolicy) + call_kwargs = mock_connect.call_args[1] + assert call_kwargs["known_hosts"] is None - def test_auto_add_connection_succeeds_with_unknown_host(self, sftp_info, mock_ssh_client): + def test_auto_add_connection_succeeds_with_unknown_host(self, sftp_info, mock_asyncssh_connect): """AC-4/AC-5: AUTO_ADD allows connection to unknown hosts.""" - _, mock_ssh = mock_ssh_client - mock_ssh.open_sftp.return_value.normalize.return_value = "/home/testuser" + mock_connect, _, mock_sftp = mock_asyncssh_connect + mock_sftp.realpath.return_value = "/home/testuser" info = sftp_info(host_key_policy=HostKeyPolicy.AUTO_ADD) client = SFTPClient(info) @@ -87,116 +85,98 @@ def test_auto_add_connection_succeeds_with_unknown_host(self, sftp_info, mock_ss class TestStrictPolicy: """Tests for STRICT host key policy.""" - def test_strict_uses_reject_policy(self, sftp_info, mock_ssh_client): - """AC-2: STRICT policy uses paramiko.RejectPolicy.""" - _, mock_ssh = mock_ssh_client + def test_strict_uses_default_known_hosts(self, sftp_info, mock_asyncssh_connect): + """AC-2: STRICT policy uses default known_hosts (no override).""" + mock_connect, _, _ = mock_asyncssh_connect info = sftp_info(host_key_policy=HostKeyPolicy.STRICT) client = SFTPClient(info) client.connect() - mock_ssh.set_missing_host_key_policy.assert_called_once() - policy_arg = mock_ssh.set_missing_host_key_policy.call_args[0][0] - assert isinstance(policy_arg, paramiko.RejectPolicy) + call_kwargs = mock_connect.call_args[1] + assert "known_hosts" not in call_kwargs - def test_strict_policy_rejects_unknown_host(self, sftp_info, mock_ssh_client): + def test_strict_policy_rejects_unknown_host(self, sftp_info): """AC-4/AC-5: STRICT policy rejects unknown hosts.""" - _, mock_ssh = mock_ssh_client - mock_ssh.connect.side_effect = paramiko.ssh_exception.SSHException( - "Server 'testhost.example.com' not found in known_hosts" - ) + with patch("asyncssh.connect", new_callable=AsyncMock) as mock_connect: + mock_connect.side_effect = asyncssh.KeyExchangeFailed( + "Server 'testhost.example.com' not found in known_hosts" + ) - info = sftp_info(host_key_policy=HostKeyPolicy.STRICT) - client = SFTPClient(info) + info = sftp_info(host_key_policy=HostKeyPolicy.STRICT) + client = SFTPClient(info) - with pytest.raises(ConnectionError, match="SFTP connection failed"): - client.connect() - assert not client.connected + with pytest.raises(ConnectionError, match="SFTP connection failed"): + client.connect() + assert not client.connected class TestPromptPolicy: """Tests for PROMPT host key policy behavior.""" - def test_prompt_policy_uses_interactive_policy(self, sftp_info, mock_ssh_client): - """PROMPT should configure InteractiveHostKeyPolicy.""" - _, mock_ssh = mock_ssh_client + def test_prompt_policy_sets_known_hosts(self, sftp_info, mock_asyncssh_connect): + """PROMPT should configure known_hosts to the portkeydrop file.""" + mock_connect, _, _ = mock_asyncssh_connect info = sftp_info(host_key_policy=HostKeyPolicy.PROMPT) client = SFTPClient(info) client.connect() - mock_ssh.set_missing_host_key_policy.assert_called_once() - policy_arg = mock_ssh.set_missing_host_key_policy.call_args[0][0] - assert isinstance(policy_arg, InteractiveHostKeyPolicy) + call_kwargs = mock_connect.call_args[1] + # known_hosts is set (either to the file path or None if file doesn't exist) + assert "known_hosts" in call_kwargs class TestHostKeyPolicyAppliedDuringConnect: """AC-3: Verify policy is applied during connection establishment.""" - def test_policy_set_before_connect_call(self, sftp_info, mock_ssh_client): - """Host key policy must be set before SSHClient.connect() is called.""" - _, mock_ssh = mock_ssh_client - call_order: list[str] = [] - - mock_ssh.set_missing_host_key_policy.side_effect = lambda p: call_order.append("set_policy") - mock_ssh.load_system_host_keys.side_effect = lambda: call_order.append("load_keys") - mock_ssh.connect.side_effect = lambda **kw: call_order.append("connect") - - info = sftp_info(host_key_policy=HostKeyPolicy.STRICT) - client = SFTPClient(info) - client.connect() - - assert "set_policy" in call_order - assert "connect" in call_order - assert call_order.index("set_policy") < call_order.index("connect") - - def test_system_host_keys_loaded(self, sftp_info, mock_ssh_client): - """System host keys are loaded during connection.""" - _, mock_ssh = mock_ssh_client + def test_known_hosts_passed_to_connect(self, sftp_info, mock_asyncssh_connect): + """known_hosts parameter is passed to asyncssh.connect().""" + mock_connect, _, _ = mock_asyncssh_connect info = sftp_info(host_key_policy=HostKeyPolicy.AUTO_ADD) client = SFTPClient(info) client.connect() - mock_ssh.load_system_host_keys.assert_called_once() + call_kwargs = mock_connect.call_args[1] + assert "known_hosts" in call_kwargs class TestMissingHostKeyScenarios: """AC-4: Handle missing host key scenarios for each policy.""" - def test_auto_add_missing_key_succeeds(self, sftp_info, mock_ssh_client): + def test_auto_add_missing_key_succeeds(self, sftp_info, mock_asyncssh_connect): """AUTO_ADD policy: missing host key is automatically added, connection succeeds.""" - _, mock_ssh = mock_ssh_client info = sftp_info(host_key_policy=HostKeyPolicy.AUTO_ADD) client = SFTPClient(info) client.connect() assert client.connected - def test_strict_missing_key_fails(self, sftp_info, mock_ssh_client): + def test_strict_missing_key_fails(self, sftp_info): """STRICT policy: missing host key causes connection failure.""" - _, mock_ssh = mock_ssh_client - mock_ssh.connect.side_effect = paramiko.ssh_exception.SSHException( - "Server 'testhost.example.com' not found in known_hosts" - ) + with patch("asyncssh.connect", new_callable=AsyncMock) as mock_connect: + mock_connect.side_effect = asyncssh.KeyExchangeFailed( + "Server 'testhost.example.com' not found in known_hosts" + ) - info = sftp_info(host_key_policy=HostKeyPolicy.STRICT) - client = SFTPClient(info) + info = sftp_info(host_key_policy=HostKeyPolicy.STRICT) + client = SFTPClient(info) - with pytest.raises(ConnectionError): - client.connect() - assert not client.connected + with pytest.raises(ConnectionError): + client.connect() + assert not client.connected - def test_strict_known_host_succeeds(self, sftp_info, mock_ssh_client): + def test_strict_known_host_succeeds(self, sftp_info, mock_asyncssh_connect): """STRICT policy: known host key allows connection.""" - _, mock_ssh = mock_ssh_client info = sftp_info(host_key_policy=HostKeyPolicy.STRICT) client = SFTPClient(info) client.connect() assert client.connected - def test_load_system_keys_failure_handled(self, sftp_info, mock_ssh_client): - """Connection proceeds even if loading system host keys fails.""" - _, mock_ssh = mock_ssh_client - mock_ssh.load_system_host_keys.side_effect = IOError("No such file") + def test_connection_proceeds_on_general_failure(self, sftp_info): + """Connection failure is properly reported.""" + with patch("asyncssh.connect", new_callable=AsyncMock) as mock_connect: + mock_connect.side_effect = OSError("Connection refused") - info = sftp_info(host_key_policy=HostKeyPolicy.AUTO_ADD) - client = SFTPClient(info) - client.connect() - assert client.connected + info = sftp_info(host_key_policy=HostKeyPolicy.AUTO_ADD) + client = SFTPClient(info) + with pytest.raises(ConnectionError): + client.connect() + assert not client.connected diff --git a/tests/test_protocols.py b/tests/test_protocols.py index dafbf50..7be6d81 100644 --- a/tests/test_protocols.py +++ b/tests/test_protocols.py @@ -2,8 +2,10 @@ from __future__ import annotations +import io +import stat as stat_mod from datetime import datetime -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -290,8 +292,6 @@ def fake_storbinary(cmd, file_obj, block_size, callback): client = FTPClient(info) client.connect() - import io - with pytest.raises(RuntimeError, match="Remote upload verification failed"): client.upload(io.BytesIO(b"data"), "/remote.bin") @@ -332,13 +332,13 @@ def test_not_connected_initially(self): client = SFTPClient(info) assert not client.connected - @patch("paramiko.SSHClient") - def test_connect_success(self, mock_ssh_class): - mock_ssh = MagicMock() - mock_sftp = MagicMock() - mock_sftp.normalize.return_value = "/" - mock_ssh.open_sftp.return_value = mock_sftp - mock_ssh_class.return_value = mock_ssh + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_connect_success(self, mock_connect): + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/" + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn info = ConnectionInfo( protocol=Protocol.SFTP, @@ -350,21 +350,16 @@ def test_connect_success(self, mock_ssh_class): client.connect() assert client.connected - mock_ssh.connect.assert_called_once_with( - hostname="example.com", - port=22, - username="user", - timeout=30, - allow_agent=True, - look_for_keys=True, - password="pass", - ) + mock_connect.assert_called_once() + call_kwargs = mock_connect.call_args[1] + assert call_kwargs["host"] == "example.com" + assert call_kwargs["port"] == 22 + assert call_kwargs["username"] == "user" + assert call_kwargs["password"] == "pass" - @patch("paramiko.SSHClient") - def test_connect_failure(self, mock_ssh_class): - mock_ssh = MagicMock() - mock_ssh.connect.side_effect = Exception("Connection refused") - mock_ssh_class.return_value = mock_ssh + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_connect_failure(self, mock_connect): + mock_connect.side_effect = Exception("Connection refused") info = ConnectionInfo(protocol=Protocol.SFTP, host="example.com") client = SFTPClient(info) @@ -373,13 +368,13 @@ def test_connect_failure(self, mock_ssh_class): client.connect() assert not client.connected - @patch("paramiko.SSHClient") - def test_disconnect(self, mock_ssh_class): - mock_ssh = MagicMock() - mock_sftp = MagicMock() - mock_sftp.normalize.return_value = "/" - mock_ssh.open_sftp.return_value = mock_sftp - mock_ssh_class.return_value = mock_ssh + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_disconnect(self, mock_connect): + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/" + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn info = ConnectionInfo(protocol=Protocol.SFTP, host="example.com") client = SFTPClient(info) @@ -387,8 +382,8 @@ def test_disconnect(self, mock_ssh_class): client.disconnect() assert not client.connected - mock_sftp.close.assert_called_once() - mock_ssh.close.assert_called_once() + mock_sftp.exit.assert_called_once() + mock_conn.close.assert_called_once() def test_list_dir_not_connected(self): info = ConnectionInfo(protocol=Protocol.SFTP, host="example.com") @@ -402,36 +397,37 @@ def test_disconnect_when_not_connected(self): client.disconnect() # Should not raise assert not client.connected - @patch("paramiko.SSHClient") - def test_list_dir_maps_file_attributes(self, mock_ssh_class): - import stat as stat_mod - - mock_ssh = MagicMock() - mock_sftp = MagicMock() - mock_sftp.normalize.return_value = "/home/user" - mock_ssh.open_sftp.return_value = mock_sftp - mock_ssh_class.return_value = mock_ssh - - file_attr = MagicMock( - filename="file.txt", - st_mode=stat_mod.S_IFREG | 0o644, - st_size=123, - st_mtime=1700000000, - st_uid=1000, - st_gid=1000, - longname="-rw-r--r--", - ) - dir_attr = MagicMock( - filename="docs", - st_mode=stat_mod.S_IFDIR | 0o755, - st_size=0, - st_mtime=1700000000, - st_uid=1000, - st_gid=1000, - longname="drwxr-xr-x", - ) - dot_attr = MagicMock(filename=".") - mock_sftp.listdir_attr.return_value = [dot_attr, file_attr, dir_attr] + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_list_dir_maps_file_attributes(self, mock_connect): + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/home/user" + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn + + file_entry = MagicMock() + file_entry.filename = "file.txt" + file_entry.attrs = MagicMock() + file_entry.attrs.permissions = stat_mod.S_IFREG | 0o644 + file_entry.attrs.size = 123 + file_entry.attrs.mtime = 1700000000 + file_entry.attrs.uid = 1000 + file_entry.attrs.gid = 1000 + file_entry.longname = "-rw-r--r--" + + dir_entry = MagicMock() + dir_entry.filename = "docs" + dir_entry.attrs = MagicMock() + dir_entry.attrs.permissions = stat_mod.S_IFDIR | 0o755 + dir_entry.attrs.size = 0 + dir_entry.attrs.mtime = 1700000000 + dir_entry.attrs.uid = 1000 + dir_entry.attrs.gid = 1000 + dir_entry.longname = "drwxr-xr-x" + + dot_entry = MagicMock() + dot_entry.filename = "." + mock_sftp.readdir.return_value = [dot_entry, file_entry, dir_entry] client = SFTPClient(ConnectionInfo(protocol=Protocol.SFTP, host="example.com")) client.connect() @@ -444,60 +440,71 @@ def test_list_dir_maps_file_attributes(self, mock_ssh_class): assert files[1].name == "docs" assert files[1].is_dir is True - @patch("paramiko.SSHClient") - def test_chdir_download_upload_and_file_ops(self, mock_ssh_class): - import io - import stat as stat_mod - - mock_ssh = MagicMock() - mock_sftp = MagicMock() - mock_sftp.normalize.side_effect = ["/", "/uploads"] - mock_ssh.open_sftp.return_value = mock_sftp - mock_ssh_class.return_value = mock_ssh + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_chdir_download_upload_and_file_ops(self, mock_connect): + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.side_effect = ["/", "/uploads"] + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn client = SFTPClient(ConnectionInfo(protocol=Protocol.SFTP, host="example.com")) client.connect() assert client.chdir("/uploads") == "/uploads" - mock_sftp.chdir.assert_called_once_with("/uploads") + # Download download_calls: list[tuple[int, int]] = [] + mock_remote_file = AsyncMock() + mock_remote_file.read.side_effect = [b"0123456789", b""] + mock_open_cm = MagicMock() + mock_open_cm.__aenter__ = AsyncMock(return_value=mock_remote_file) + mock_open_cm.__aexit__ = AsyncMock(return_value=False) + mock_sftp.open = MagicMock(return_value=mock_open_cm) + stat_attrs = MagicMock() + stat_attrs.size = 10 + stat_attrs.permissions = stat_mod.S_IFREG | 0o644 + mock_sftp.stat.return_value = stat_attrs - def fake_getfo(_path, _file, callback): - callback(10, 100) - - mock_sftp.getfo.side_effect = fake_getfo client.download( "/remote.bin", MagicMock(), callback=lambda t, n: download_calls.append((t, n)) ) - assert download_calls == [(10, 100)] + assert download_calls == [(10, 10)] + # Upload upload_calls: list[tuple[int, int]] = [] + mock_write_file = AsyncMock() + mock_open_cm.__aenter__ = AsyncMock(return_value=mock_write_file) - def fake_putfo(_file, _path, file_size, callback): - assert file_size == 4 - callback(4, file_size) + upload_stat_attrs = MagicMock() + upload_stat_attrs.size = 4 + upload_stat_attrs.permissions = stat_mod.S_IFREG | 0o644 - mock_sftp.putfo.side_effect = fake_putfo - file_attr = MagicMock(st_mode=stat_mod.S_IFREG | 0o644, st_size=4) - dir_attr = MagicMock(st_mode=stat_mod.S_IFDIR | 0o755, st_size=0) + mkdir_stat_attrs = MagicMock() + mkdir_stat_attrs.permissions = stat_mod.S_IFDIR | 0o755 - def stat_side_effect(path: str): + rename_stat_attrs = MagicMock() + rename_stat_attrs.permissions = stat_mod.S_IFREG | 0o644 + + async def stat_side_effect(path: str): if path == "/remote.txt": - return file_attr + return upload_stat_attrs if path == "/a": raise FileNotFoundError(path) if path == "/b": - if mock_sftp.rmdir.called: + if mock_sftp.rmdir.await_count > 0: raise FileNotFoundError(path) - return dir_attr + return mkdir_stat_attrs if path == "/new": - return file_attr + return rename_stat_attrs raise FileNotFoundError(path) mock_sftp.stat.side_effect = stat_side_effect + client.upload( - io.BytesIO(b"data"), "/remote.txt", callback=lambda t, n: upload_calls.append((t, n)) + io.BytesIO(b"data"), + "/remote.txt", + callback=lambda t, n: upload_calls.append((t, n)), ) assert upload_calls == [(4, 4)] @@ -505,10 +512,10 @@ def stat_side_effect(path: str): client.mkdir("/b") client.rmdir("/b") client.rename("/old", "/new") - mock_sftp.remove.assert_called_once_with("/a") - mock_sftp.mkdir.assert_called_once_with("/b") - mock_sftp.rmdir.assert_called_once_with("/b") - mock_sftp.rename.assert_called_once_with("/old", "/new") + mock_sftp.remove.assert_awaited_once_with("/a") + mock_sftp.mkdir.assert_awaited_once_with("/b") + mock_sftp.rmdir.assert_awaited_once_with("/b") + mock_sftp.rename.assert_awaited_once_with("/old", "/new") @patch("ftplib.FTP") def test_ftp_rename_raises_when_target_not_found_after_rename(self, mock_ftp_class): @@ -524,16 +531,14 @@ def test_ftp_rename_raises_when_target_not_found_after_rename(self, mock_ftp_cla with pytest.raises(RuntimeError, match="verification failed"): client.rename("/old.txt", "/new.txt") - @patch("paramiko.SSHClient") @patch("os.path.exists", return_value=True) - def test_connect_with_key_and_stat(self, _mock_exists, mock_ssh_class): - import stat as stat_mod - - mock_ssh = MagicMock() - mock_sftp = MagicMock() - mock_sftp.normalize.return_value = "/" - mock_ssh.open_sftp.return_value = mock_sftp - mock_ssh_class.return_value = mock_ssh + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_connect_with_key_and_stat(self, mock_connect, _mock_exists): + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/" + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn info = ConnectionInfo( protocol=Protocol.SFTP, @@ -544,31 +549,39 @@ def test_connect_with_key_and_stat(self, _mock_exists, mock_ssh_class): client = SFTPClient(info) client.connect() - kwargs = mock_ssh.connect.call_args.kwargs - assert kwargs["allow_agent"] is False - assert kwargs["look_for_keys"] is False - assert kwargs["key_filename"] == "/tmp/id_rsa" + kwargs = mock_connect.call_args[1] + assert kwargs["agent_path"] is None + assert kwargs["client_keys"] == ["/tmp/id_rsa"] - attr = MagicMock(st_mode=stat_mod.S_IFREG | 0o644, st_size=42, st_mtime=1700000000) - mock_sftp.stat.return_value = attr + stat_attrs = MagicMock() + stat_attrs.permissions = stat_mod.S_IFREG | 0o644 + stat_attrs.size = 42 + stat_attrs.mtime = 1700000000 + mock_sftp.stat.return_value = stat_attrs remote = client.stat("/remote/file.txt") assert remote.name == "file.txt" assert remote.path == "/remote/file.txt" assert remote.size == 42 assert remote.is_dir is False - @patch("paramiko.SSHClient") - def test_upload_raises_when_remote_size_mismatch(self, mock_ssh_class): - import io - import stat as stat_mod - - mock_ssh = MagicMock() - mock_sftp = MagicMock() - mock_sftp.normalize.return_value = "/" - mock_ssh.open_sftp.return_value = mock_sftp - mock_ssh_class.return_value = mock_ssh - - mock_sftp.stat.return_value = MagicMock(st_mode=stat_mod.S_IFREG | 0o644, st_size=3) + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_upload_raises_when_remote_size_mismatch(self, mock_connect): + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/" + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn + + mock_write_file = AsyncMock() + mock_open_cm = MagicMock() + mock_open_cm.__aenter__ = AsyncMock(return_value=mock_write_file) + mock_open_cm.__aexit__ = AsyncMock(return_value=False) + mock_sftp.open = MagicMock(return_value=mock_open_cm) + + stat_attrs = MagicMock() + stat_attrs.size = 3 + stat_attrs.permissions = stat_mod.S_IFREG | 0o644 + mock_sftp.stat.return_value = stat_attrs client = SFTPClient(ConnectionInfo(protocol=Protocol.SFTP, host="example.com")) client.connect() @@ -576,16 +589,18 @@ def test_upload_raises_when_remote_size_mismatch(self, mock_ssh_class): with pytest.raises(RuntimeError, match="verification failed"): client.upload(io.BytesIO(b"data"), "/remote.txt") - @patch("paramiko.SSHClient") - def test_mkdir_raises_when_created_path_is_not_directory(self, mock_ssh_class): - import stat as stat_mod + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_mkdir_raises_when_created_path_is_not_directory(self, mock_connect): + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/" + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn - mock_ssh = MagicMock() - mock_sftp = MagicMock() - mock_sftp.normalize.return_value = "/" - mock_ssh.open_sftp.return_value = mock_sftp - mock_ssh_class.return_value = mock_ssh - mock_sftp.stat.return_value = MagicMock(st_mode=stat_mod.S_IFREG | 0o644, st_size=10) + stat_attrs = MagicMock() + stat_attrs.permissions = stat_mod.S_IFREG | 0o644 + stat_attrs.size = 10 + mock_sftp.stat.return_value = stat_attrs client = SFTPClient(ConnectionInfo(protocol=Protocol.SFTP, host="example.com")) client.connect() @@ -593,13 +608,14 @@ def test_mkdir_raises_when_created_path_is_not_directory(self, mock_ssh_class): with pytest.raises(RuntimeError, match="verification failed"): client.mkdir("/not-a-dir") - @patch("paramiko.SSHClient") - def test_delete_raises_when_remote_stat_succeeds(self, mock_ssh_class): - mock_ssh = MagicMock() - mock_sftp = MagicMock() - mock_sftp.normalize.return_value = "/" - mock_ssh.open_sftp.return_value = mock_sftp - mock_ssh_class.return_value = mock_ssh + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_delete_raises_when_remote_stat_succeeds(self, mock_connect): + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/" + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn + mock_sftp.stat.return_value = MagicMock() client = SFTPClient(ConnectionInfo(protocol=Protocol.SFTP, host="example.com")) @@ -608,13 +624,14 @@ def test_delete_raises_when_remote_stat_succeeds(self, mock_ssh_class): with pytest.raises(RuntimeError, match="verification failed"): client.delete("/file") - @patch("paramiko.SSHClient") - def test_rmdir_raises_when_remote_stat_succeeds(self, mock_ssh_class): - mock_ssh = MagicMock() - mock_sftp = MagicMock() - mock_sftp.normalize.return_value = "/" - mock_ssh.open_sftp.return_value = mock_sftp - mock_ssh_class.return_value = mock_ssh + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_rmdir_raises_when_remote_stat_succeeds(self, mock_connect): + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/" + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn + mock_sftp.stat.return_value = MagicMock() client = SFTPClient(ConnectionInfo(protocol=Protocol.SFTP, host="example.com")) diff --git a/tests/test_sftp_client.py b/tests/test_sftp_client.py index 9800e2b..de02671 100644 --- a/tests/test_sftp_client.py +++ b/tests/test_sftp_client.py @@ -1,11 +1,11 @@ -"""Tests for SFTPClient using SSHClient-based authentication.""" +"""Tests for SFTPClient using asyncssh-based authentication.""" from __future__ import annotations import logging -from unittest.mock import MagicMock, patch +import stat as stat_mod +from unittest.mock import AsyncMock, MagicMock, patch -import paramiko import pytest from portkeydrop.protocols import ConnectionInfo, HostKeyPolicy, Protocol, SFTPClient @@ -22,96 +22,93 @@ def sftp_info() -> ConnectionInfo: ) -def _make_mock_ssh() -> MagicMock: - mock_ssh = MagicMock() - mock_sftp = MagicMock() - mock_ssh.open_sftp.return_value = mock_sftp - mock_sftp.normalize.return_value = "/home/user" - return mock_ssh +def _make_mock_conn() -> tuple[MagicMock, MagicMock]: + """Create mock asyncssh connection and SFTP client.""" + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_conn.start_sftp_client.return_value = mock_sftp + mock_sftp.realpath.return_value = "/home/user" + return mock_conn, mock_sftp class TestSFTPClientInit: - def test_creates_ssh_client_attribute(self, sftp_info: ConnectionInfo) -> None: + def test_creates_conn_attribute(self, sftp_info: ConnectionInfo) -> None: client = SFTPClient(sftp_info) - assert hasattr(client, "_ssh_client") - assert client._ssh_client is None + assert hasattr(client, "_conn") + assert client._conn is None - def test_no_transport_attribute(self, sftp_info: ConnectionInfo) -> None: + def test_has_event_loop(self, sftp_info: ConnectionInfo) -> None: client = SFTPClient(sftp_info) - assert not hasattr(client, "_transport") + assert client._loop is not None + assert client._loop.is_running() class TestSFTPClientConnect: - @patch("paramiko.SSHClient") - def test_connect_with_password(self, mock_cls: MagicMock, sftp_info: ConnectionInfo) -> None: - mock_ssh = _make_mock_ssh() - mock_cls.return_value = mock_ssh + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_connect_with_password( + self, mock_connect: AsyncMock, sftp_info: ConnectionInfo + ) -> None: + mock_conn, mock_sftp = _make_mock_conn() + mock_connect.return_value = mock_conn client = SFTPClient(sftp_info) client.connect() - mock_ssh.connect.assert_called_once() - call_kwargs = mock_ssh.connect.call_args[1] - assert call_kwargs["hostname"] == "example.com" + mock_connect.assert_called_once() + call_kwargs = mock_connect.call_args[1] + assert call_kwargs["host"] == "example.com" assert call_kwargs["username"] == "user" assert call_kwargs["password"] == "pass" - assert call_kwargs["allow_agent"] is True - assert call_kwargs["look_for_keys"] is True assert client._connected is True @patch("os.path.exists", return_value=True) - @patch("paramiko.SSHClient") - def test_connect_with_key_file(self, mock_cls: MagicMock, _mock_exists: MagicMock) -> None: + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_connect_with_key_file(self, mock_connect: AsyncMock, _mock_exists: MagicMock) -> None: info = ConnectionInfo( protocol=Protocol.SFTP, host="example.com", username="user", key_path="/path/to/key", ) - mock_ssh = _make_mock_ssh() - mock_cls.return_value = mock_ssh + mock_conn, mock_sftp = _make_mock_conn() + mock_connect.return_value = mock_conn client = SFTPClient(info) client.connect() - call_kwargs = mock_ssh.connect.call_args[1] - assert call_kwargs["key_filename"] == "/path/to/key" - assert call_kwargs["allow_agent"] is False - assert call_kwargs["look_for_keys"] is False + call_kwargs = mock_connect.call_args[1] + assert call_kwargs["client_keys"] == ["/path/to/key"] + assert call_kwargs["agent_path"] is None - @patch("paramiko.SSHClient") - def test_connect_agent_only(self, mock_cls: MagicMock) -> None: + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_connect_agent_only(self, mock_connect: AsyncMock) -> None: info = ConnectionInfo( protocol=Protocol.SFTP, host="example.com", username="user", ) - mock_ssh = _make_mock_ssh() - mock_cls.return_value = mock_ssh + mock_conn, mock_sftp = _make_mock_conn() + mock_connect.return_value = mock_conn client = SFTPClient(info) client.connect() - call_kwargs = mock_ssh.connect.call_args[1] - assert call_kwargs["allow_agent"] is True - assert call_kwargs["look_for_keys"] is True + call_kwargs = mock_connect.call_args[1] assert "password" not in call_kwargs - assert "key_filename" not in call_kwargs + assert "client_keys" not in call_kwargs - @patch("paramiko.SSHClient") - def test_connect_failure(self, mock_cls: MagicMock, sftp_info: ConnectionInfo) -> None: - mock_ssh = MagicMock() - mock_cls.return_value = mock_ssh - mock_ssh.connect.side_effect = Exception("Auth failed") + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_connect_failure(self, mock_connect: AsyncMock, sftp_info: ConnectionInfo) -> None: + mock_connect.side_effect = Exception("Auth failed") client = SFTPClient(sftp_info) with pytest.raises(ConnectionError, match="SFTP connection failed"): client.connect() assert client._connected is False - @patch("paramiko.SSHClient") + @patch("asyncssh.connect", new_callable=AsyncMock) def test_logs_authentication_methods( - self, mock_cls: MagicMock, caplog: pytest.LogCaptureFixture + self, mock_connect: AsyncMock, caplog: pytest.LogCaptureFixture ) -> None: info = ConnectionInfo( protocol=Protocol.SFTP, @@ -119,8 +116,8 @@ def test_logs_authentication_methods( username="user", password="pass", ) - mock_ssh = _make_mock_ssh() - mock_cls.return_value = mock_ssh + mock_conn, mock_sftp = _make_mock_conn() + mock_connect.return_value = mock_conn caplog.set_level(logging.DEBUG, logger="portkeydrop.protocols") @@ -128,17 +125,16 @@ def test_logs_authentication_methods( client.connect() text = caplog.text - assert "SSH agent detection" in text assert "SFTP authentication methods to try: ssh-agent, default-key-files, password" in text assert "SSH authentication succeeded" in text - @patch("paramiko.SSHClient") + @patch("asyncssh.connect", new_callable=AsyncMock) def test_logs_agent_unavailable_error( - self, mock_cls: MagicMock, caplog: pytest.LogCaptureFixture + self, mock_connect: AsyncMock, caplog: pytest.LogCaptureFixture ) -> None: - mock_ssh = MagicMock() - mock_cls.return_value = mock_ssh - mock_ssh.connect.side_effect = paramiko.SSHException("Error connecting to agent") + import asyncssh + + mock_connect.side_effect = asyncssh.DisconnectError(11, "Error connecting to agent") caplog.set_level(logging.DEBUG, logger="portkeydrop.protocols") @@ -148,13 +144,11 @@ def test_logs_agent_unavailable_error( with pytest.raises(ConnectionError, match="SSH agent is unavailable or inaccessible"): client.connect() - assert "SSH agent appears unavailable or inaccessible" in caplog.text + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_auth_failure_message_for_agent_and_password(self, mock_connect: AsyncMock) -> None: + import asyncssh - @patch("paramiko.SSHClient") - def test_auth_failure_message_for_agent_and_password(self, mock_cls: MagicMock) -> None: - mock_ssh = MagicMock() - mock_cls.return_value = mock_ssh - mock_ssh.connect.side_effect = paramiko.AuthenticationException("denied") + mock_connect.side_effect = asyncssh.PermissionDenied("denied") info = ConnectionInfo( protocol=Protocol.SFTP, @@ -169,11 +163,11 @@ def test_auth_failure_message_for_agent_and_password(self, mock_cls: MagicMock) ): client.connect() - @patch("paramiko.SSHClient") - def test_auth_failure_message_for_agent_only(self, mock_cls: MagicMock) -> None: - mock_ssh = MagicMock() - mock_cls.return_value = mock_ssh - mock_ssh.connect.side_effect = paramiko.AuthenticationException("denied") + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_auth_failure_message_for_agent_only(self, mock_connect: AsyncMock) -> None: + import asyncssh + + mock_connect.side_effect = asyncssh.PermissionDenied("denied") info = ConnectionInfo(protocol=Protocol.SFTP, host="example.com", username="user") client = SFTPClient(info) @@ -183,47 +177,43 @@ def test_auth_failure_message_for_agent_only(self, mock_cls: MagicMock) -> None: class TestSFTPClientHostKeyPolicy: - @patch("paramiko.SSHClient") - def test_auto_add_policy(self, mock_cls: MagicMock, sftp_info: ConnectionInfo) -> None: - mock_ssh = _make_mock_ssh() - mock_cls.return_value = mock_ssh + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_auto_add_policy(self, mock_connect: AsyncMock, sftp_info: ConnectionInfo) -> None: + mock_conn, mock_sftp = _make_mock_conn() + mock_connect.return_value = mock_conn sftp_info.host_key_policy = HostKeyPolicy.AUTO_ADD client = SFTPClient(sftp_info) client.connect() - # Check that AutoAddPolicy was set - policy_calls = mock_ssh.set_missing_host_key_policy.call_args_list - assert len(policy_calls) == 1 - policy_arg = policy_calls[0][0][0] - assert isinstance(policy_arg, paramiko.AutoAddPolicy) + call_kwargs = mock_connect.call_args[1] + assert call_kwargs["known_hosts"] is None - @patch("paramiko.SSHClient") - def test_strict_policy(self, mock_cls: MagicMock, sftp_info: ConnectionInfo) -> None: - mock_ssh = _make_mock_ssh() - mock_cls.return_value = mock_ssh + @patch("asyncssh.connect", new_callable=AsyncMock) + def test_strict_policy(self, mock_connect: AsyncMock, sftp_info: ConnectionInfo) -> None: + mock_conn, mock_sftp = _make_mock_conn() + mock_connect.return_value = mock_conn sftp_info.host_key_policy = HostKeyPolicy.STRICT client = SFTPClient(sftp_info) client.connect() - policy_calls = mock_ssh.set_missing_host_key_policy.call_args_list - assert len(policy_calls) == 1 - policy_arg = policy_calls[0][0][0] - assert isinstance(policy_arg, paramiko.RejectPolicy) + call_kwargs = mock_connect.call_args[1] + # Strict uses default known_hosts (no explicit override) + assert "known_hosts" not in call_kwargs class TestSFTPClientEnsureConnected: - def test_raises_when_ssh_client_is_none(self, sftp_info: ConnectionInfo) -> None: + def test_raises_when_conn_is_none(self, sftp_info: ConnectionInfo) -> None: client = SFTPClient(sftp_info) - client._ssh_client = None + client._conn = None client._sftp = MagicMock() with pytest.raises(ConnectionError, match="Not connected"): client._ensure_connected() def test_raises_when_sftp_is_none(self, sftp_info: ConnectionInfo) -> None: client = SFTPClient(sftp_info) - client._ssh_client = MagicMock() + client._conn = MagicMock() client._sftp = None with pytest.raises(ConnectionError, match="Not connected"): client._ensure_connected() @@ -236,7 +226,7 @@ def test_raises_when_both_none(self, sftp_info: ConnectionInfo) -> None: def test_returns_sftp_when_connected(self, sftp_info: ConnectionInfo) -> None: client = SFTPClient(sftp_info) mock_sftp = MagicMock() - client._ssh_client = MagicMock() + client._conn = MagicMock() client._sftp = mock_sftp result = client._ensure_connected() assert result is mock_sftp @@ -263,150 +253,104 @@ def test_methods_use_ensure_connected(self, sftp_info: ConnectionInfo) -> None: class TestSFTPClientDisconnect: def test_disconnect_cleans_up(self, sftp_info: ConnectionInfo) -> None: client = SFTPClient(sftp_info) - mock_ssh = MagicMock() + mock_conn = MagicMock() mock_sftp = MagicMock() - client._ssh_client = mock_ssh + client._conn = mock_conn client._sftp = mock_sftp client._connected = True client.disconnect() - assert client._ssh_client is None + assert client._conn is None assert client._sftp is None assert client._connected is False def test_disconnect_calls_close_on_both(self, sftp_info: ConnectionInfo) -> None: client = SFTPClient(sftp_info) - mock_ssh = MagicMock() + mock_conn = MagicMock() mock_sftp = MagicMock() - client._ssh_client = mock_ssh + client._conn = mock_conn client._sftp = mock_sftp client._connected = True client.disconnect() - mock_sftp.close.assert_called_once() - mock_ssh.close.assert_called_once() + mock_sftp.exit.assert_called_once() + mock_conn.close.assert_called_once() def test_disconnect_handles_sftp_close_exception(self, sftp_info: ConnectionInfo) -> None: client = SFTPClient(sftp_info) - mock_ssh = MagicMock() + mock_conn = MagicMock() mock_sftp = MagicMock() - mock_sftp.close.side_effect = Exception("SFTP close error") - client._ssh_client = mock_ssh + mock_sftp.exit.side_effect = Exception("SFTP close error") + client._conn = mock_conn client._sftp = mock_sftp client._connected = True client.disconnect() # Should not raise assert client._sftp is None - assert client._ssh_client is None - mock_ssh.close.assert_called_once() + assert client._conn is None + mock_conn.close.assert_called_once() - def test_disconnect_handles_ssh_close_exception(self, sftp_info: ConnectionInfo) -> None: + def test_disconnect_handles_conn_close_exception(self, sftp_info: ConnectionInfo) -> None: client = SFTPClient(sftp_info) - mock_ssh = MagicMock() - mock_ssh.close.side_effect = Exception("SSH close error") + mock_conn = MagicMock() + mock_conn.close.side_effect = Exception("SSH close error") mock_sftp = MagicMock() - client._ssh_client = mock_ssh + client._conn = mock_conn client._sftp = mock_sftp client._connected = True client.disconnect() # Should not raise assert client._sftp is None - assert client._ssh_client is None + assert client._conn is None def test_disconnect_handles_both_close_exceptions(self, sftp_info: ConnectionInfo) -> None: client = SFTPClient(sftp_info) - mock_ssh = MagicMock() - mock_ssh.close.side_effect = Exception("SSH close error") + mock_conn = MagicMock() + mock_conn.close.side_effect = Exception("SSH close error") mock_sftp = MagicMock() - mock_sftp.close.side_effect = Exception("SFTP close error") - client._ssh_client = mock_ssh + mock_sftp.exit.side_effect = Exception("SFTP close error") + client._conn = mock_conn client._sftp = mock_sftp client._connected = True client.disconnect() # Should not raise assert client._sftp is None - assert client._ssh_client is None + assert client._conn is None def test_disconnect_when_not_connected(self, sftp_info: ConnectionInfo) -> None: client = SFTPClient(sftp_info) client.disconnect() # Should not raise assert client._sftp is None - assert client._ssh_client is None - - -class TestSFTPClientSftpCall: - def test_sftp_call_returns_result(self, sftp_info: ConnectionInfo) -> None: - client = SFTPClient(sftp_info) - client._ssh_client = MagicMock() - client._sftp = MagicMock() - - result = client._sftp_call(lambda: 42) - assert result == 42 - - def test_sftp_call_passes_args(self, sftp_info: ConnectionInfo) -> None: - client = SFTPClient(sftp_info) - client._ssh_client = MagicMock() - client._sftp = MagicMock() - - result = client._sftp_call(lambda a, b: a + b, 3, 4) - assert result == 7 - - def test_sftp_call_timeout_raises(self, sftp_info: ConnectionInfo) -> None: - import socket - - sftp_info.timeout = 1 - client = SFTPClient(sftp_info) - client._ssh_client = MagicMock() - client._sftp = MagicMock() - - # socket.timeout should propagate as-is (no longer converted here) - def raises_socket_timeout(): - raise socket.timeout("timed out") - - with pytest.raises(socket.timeout): - client._sftp_call(raises_socket_timeout) - - def test_sftp_call_uses_fallback_timeout_when_none(self, sftp_info: ConnectionInfo) -> None: - """If timeout is None, should use 30s fallback (not hang forever).""" - sftp_info.timeout = None # type: ignore[assignment] - client = SFTPClient(sftp_info) - client._ssh_client = MagicMock() - client._sftp = MagicMock() - - # Just verify it doesn't error on None timeout — we can't wait 30s in tests, - # so confirm it runs a fast fn successfully with None timeout. - result = client._sftp_call(lambda: "ok") - assert result == "ok" + assert client._conn is None class TestSFTPClientListDir: - def _make_connected(self, sftp_info: ConnectionInfo) -> tuple[SFTPClient, MagicMock]: + def _make_connected(self, sftp_info: ConnectionInfo) -> tuple[SFTPClient, AsyncMock]: client = SFTPClient(sftp_info) - mock_sftp = MagicMock() - mock_sftp.normalize.return_value = "/home/user" - client._ssh_client = MagicMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/home/user" + client._conn = AsyncMock() client._sftp = mock_sftp client._cwd = "/home/user" return client, mock_sftp def test_list_dir_returns_files(self, sftp_info: ConnectionInfo) -> None: - import stat as stat_mod - client, mock_sftp = self._make_connected(sftp_info) - attr = MagicMock() - attr.filename = "file.txt" - attr.st_mode = stat_mod.S_IFREG | 0o644 - attr.st_size = 100 - attr.st_mtime = 0 - attr.st_uid = 1000 - attr.st_gid = 1000 - attr.longname = "-rw-r--r-- 1 user group 100 Jan 1 file.txt" - mock_sftp.listdir_attr.return_value = [attr] + entry = MagicMock() + entry.filename = "file.txt" + entry.attrs = MagicMock() + entry.attrs.permissions = stat_mod.S_IFREG | 0o644 + entry.attrs.size = 100 + entry.attrs.mtime = 0 + entry.attrs.uid = 1000 + entry.attrs.gid = 1000 + entry.longname = "-rw-r--r-- 1 user group 100 Jan 1 file.txt" + mock_sftp.readdir.return_value = [entry] files = client.list_dir() assert len(files) == 1 @@ -419,7 +363,7 @@ def test_list_dir_permission_error(self, sftp_info: ConnectionInfo) -> None: client, mock_sftp = self._make_connected(sftp_info) err = IOError() err.errno = errno.EACCES - mock_sftp.listdir_attr.side_effect = err + mock_sftp.readdir.side_effect = err with pytest.raises(PermissionError, match="Permission denied"): client.list_dir("/restricted") @@ -430,18 +374,18 @@ def test_list_dir_reraises_other_oserror(self, sftp_info: ConnectionInfo) -> Non client, mock_sftp = self._make_connected(sftp_info) err = IOError() err.errno = errno.ENOENT - mock_sftp.listdir_attr.side_effect = err + mock_sftp.readdir.side_effect = err with pytest.raises(IOError): client.list_dir("/gone") class TestSFTPClientChdir: - def _make_connected(self, sftp_info: ConnectionInfo) -> tuple[SFTPClient, MagicMock]: + def _make_connected(self, sftp_info: ConnectionInfo) -> tuple[SFTPClient, AsyncMock]: client = SFTPClient(sftp_info) - mock_sftp = MagicMock() - mock_sftp.normalize.return_value = "/home/user/subdir" - client._ssh_client = MagicMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/home/user/subdir" + client._conn = AsyncMock() client._sftp = mock_sftp client._cwd = "/home/user" return client, mock_sftp @@ -458,44 +402,42 @@ def test_chdir_permission_error(self, sftp_info: ConnectionInfo) -> None: client, mock_sftp = self._make_connected(sftp_info) err = IOError() err.errno = errno.EPERM - mock_sftp.chdir.side_effect = err + mock_sftp.realpath.side_effect = err with pytest.raises(PermissionError, match="Permission denied"): client.chdir("/restricted") class TestSFTPClientListDirSpecialFiles: - def _make_connected(self, sftp_info: ConnectionInfo) -> tuple[SFTPClient, MagicMock]: + def _make_connected(self, sftp_info: ConnectionInfo) -> tuple[SFTPClient, AsyncMock]: client = SFTPClient(sftp_info) - mock_sftp = MagicMock() - mock_sftp.normalize.return_value = "/home/user" - client._ssh_client = MagicMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/home/user" + client._conn = AsyncMock() client._sftp = mock_sftp client._cwd = "/home/user" return client, mock_sftp def test_socket_file_skipped(self, sftp_info: ConnectionInfo) -> None: - import stat as stat_mod - client, mock_sftp = self._make_connected(sftp_info) - attr = MagicMock() - attr.filename = "agent.sock" - attr.st_mode = stat_mod.S_IFSOCK | 0o600 - attr.longname = "srw------- 1 user group 0 Jan 1 agent.sock" - mock_sftp.listdir_attr.return_value = [attr] + entry = MagicMock() + entry.filename = "agent.sock" + entry.attrs = MagicMock() + entry.attrs.permissions = stat_mod.S_IFSOCK | 0o600 + entry.longname = "srw------- 1 user group 0 Jan 1 agent.sock" + mock_sftp.readdir.return_value = [entry] files = client.list_dir() assert files == [] def test_fifo_file_skipped(self, sftp_info: ConnectionInfo) -> None: - import stat as stat_mod - client, mock_sftp = self._make_connected(sftp_info) - attr = MagicMock() - attr.filename = "mypipe" - attr.st_mode = stat_mod.S_IFIFO | 0o644 - attr.longname = "prw-r--r-- 1 user group 0 Jan 1 mypipe" - mock_sftp.listdir_attr.return_value = [attr] + entry = MagicMock() + entry.filename = "mypipe" + entry.attrs = MagicMock() + entry.attrs.permissions = stat_mod.S_IFIFO | 0o644 + entry.longname = "prw-r--r-- 1 user group 0 Jan 1 mypipe" + mock_sftp.readdir.return_value = [entry] files = client.list_dir() assert files == [] @@ -506,328 +448,7 @@ def test_oserror_reraises(self, sftp_info: ConnectionInfo) -> None: client, mock_sftp = self._make_connected(sftp_info) err = IOError() err.errno = errno.ENOENT - mock_sftp.listdir_attr.side_effect = err + mock_sftp.readdir.side_effect = err with pytest.raises(IOError): client.list_dir("/gone") - - -class TestListdirAttrSafe: - def _make_connected(self, sftp_info): - from portkeydrop.protocols import SFTPClient - - client = SFTPClient(sftp_info) - client._ssh_client = MagicMock() - mock_sftp = MagicMock() - client._sftp = mock_sftp - client._connected = True - client._cwd = "/home/user" - return client, mock_sftp - - def test_count_zero_treated_as_eof(self, sftp_info: ConnectionInfo) -> None: - """count=0 in READDIR response should break the loop (fixes NAS hang).""" - from paramiko.sftp import CMD_CLOSE, CMD_HANDLE, CMD_NAME, CMD_OPENDIR, CMD_READDIR - - client, mock_sftp = self._make_connected(sftp_info) - - handle_msg = MagicMock() - handle_msg.get_binary.return_value = b"handle1" - - readdir_msg = MagicMock() - readdir_msg.get_int.return_value = 0 # count=0 → should break - - close_msg = MagicMock() - - def fake_request(cmd, *args): - if cmd == CMD_OPENDIR: - return (CMD_HANDLE, handle_msg) - if cmd == CMD_READDIR: - return (CMD_NAME, readdir_msg) - if cmd == CMD_CLOSE: - return (CMD_CLOSE, close_msg) - raise ValueError(f"unexpected cmd {cmd}") - - mock_sftp._request.side_effect = fake_request - mock_sftp._adjust_cwd.return_value = "/home/user" - - result = client._listdir_attr_safe(mock_sftp, "/home/user") - assert result == [] - - def test_eoferror_breaks_loop(self, sftp_info: ConnectionInfo) -> None: - """EOFError on READDIR should terminate the loop normally.""" - from paramiko.sftp import CMD_CLOSE, CMD_HANDLE, CMD_OPENDIR, CMD_READDIR - - client, mock_sftp = self._make_connected(sftp_info) - - handle_msg = MagicMock() - handle_msg.get_binary.return_value = b"handle1" - - def fake_request(cmd, *args): - if cmd == CMD_OPENDIR: - return (CMD_HANDLE, handle_msg) - if cmd == CMD_READDIR: - raise EOFError - if cmd == CMD_CLOSE: - return (CMD_CLOSE, MagicMock()) - raise ValueError(f"unexpected cmd {cmd}") - - mock_sftp._request.side_effect = fake_request - mock_sftp._adjust_cwd.return_value = "/home/user" - - result = client._listdir_attr_safe(mock_sftp, "/home/user") - assert result == [] - - def test_falls_back_to_listdir_attr_when_request_unavailable( - self, sftp_info: ConnectionInfo - ) -> None: - """When _request raises TypeError (MagicMock), falls back to listdir_attr.""" - import paramiko - - client, mock_sftp = self._make_connected(sftp_info) - mock_sftp._adjust_cwd.return_value = "/home/user" - mock_sftp._request.side_effect = TypeError("not real paramiko") - - attr = paramiko.SFTPAttributes() - attr.filename = "file.txt" - mock_sftp.listdir_attr.return_value = [attr] - - result = client._listdir_attr_safe(mock_sftp, "/home/user") - assert len(result) == 1 - assert result[0].filename == "file.txt" - - -class TestListDirViaExec: - def _make_connected(self, sftp_info): - from portkeydrop.protocols import SFTPClient - - client = SFTPClient(sftp_info) - client._ssh_client = MagicMock() - client._sftp = MagicMock() - client._connected = True - client._cwd = "/home/user" - return client - - def test_parses_ls_output(self, sftp_info: ConnectionInfo) -> None: - client = self._make_connected(sftp_info) - ls_output = ( - "total 8\n" - "drwxr-xr-x 2 user group 4096 Jan 01 12:00 subdir\n" - "-rw-r--r-- 1 user group 512 Jan 01 12:00 file.txt\n" - ) - stdout_mock = MagicMock() - stdout_mock.read.return_value = ls_output.encode() - client._ssh_client.exec_command.return_value = (MagicMock(), stdout_mock, MagicMock()) - - result = client._list_dir_via_exec("/home/user") - names = [a.filename for a in result] - assert "subdir" in names - assert "file.txt" in names - - def test_returns_empty_on_exec_exception(self, sftp_info: ConnectionInfo) -> None: - client = self._make_connected(sftp_info) - client._ssh_client.exec_command.side_effect = Exception("Channel closed") - - result = client._list_dir_via_exec("/home/user") - assert result == [] - - def test_returns_empty_when_no_ssh_client(self, sftp_info: ConnectionInfo) -> None: - client = self._make_connected(sftp_info) - client._ssh_client = None - - result = client._list_dir_via_exec("/home/user") - assert result == [] - - -class TestParseLsMode: - def _parse(self, s): - from portkeydrop.protocols import SFTPClient - - return SFTPClient._parse_ls_mode(s) - - def test_directory(self) -> None: - import stat - - mode = self._parse("drwxr-xr-x") - assert stat.S_ISDIR(mode) - - def test_regular_file(self) -> None: - import stat - - mode = self._parse("-rw-r--r--") - assert stat.S_ISREG(mode) - - def test_symlink(self) -> None: - import stat - - mode = self._parse("lrwxrwxrwx") - assert stat.S_ISLNK(mode) - - def test_socket(self) -> None: - import stat - - mode = self._parse("srwxrwxrwx") - assert stat.S_ISSOCK(mode) - - def test_raises_on_short_string(self) -> None: - with pytest.raises(ValueError): - self._parse("drwx") - - -class TestListdirAttrSafeWithEntries: - def _make_connected(self, sftp_info): - from portkeydrop.protocols import SFTPClient - - client = SFTPClient(sftp_info) - client._ssh_client = MagicMock() - mock_sftp = MagicMock() - client._sftp = mock_sftp - client._connected = True - client._cwd = "/home/user" - return client, mock_sftp - - def test_entries_returned_excluding_dot_entries(self, sftp_info: ConnectionInfo) -> None: - from paramiko.sftp import CMD_CLOSE, CMD_HANDLE, CMD_NAME, CMD_OPENDIR, CMD_READDIR - from paramiko.sftp_attr import SFTPAttributes - - client, mock_sftp = self._make_connected(sftp_info) - handle_msg = MagicMock() - handle_msg.get_binary.return_value = b"h" - - call_count = [0] - - def fake_request(cmd, *args): - if cmd == CMD_OPENDIR: - return (CMD_HANDLE, handle_msg) - if cmd == CMD_READDIR: - call_count[0] += 1 - if call_count[0] == 1: - msg = MagicMock() - msg.get_int.return_value = 2 - attr1 = SFTPAttributes() - attr1.filename = "." - attr2 = SFTPAttributes() - attr2.filename = "file.txt" - attrs = [attr1, attr2] - idx = [0] - - def get_text(): - v = attrs[idx[0] // 2].filename - idx[0] += 1 - return v - - msg.get_text.side_effect = get_text - with patch( - "paramiko.sftp_attr.SFTPAttributes._from_msg", - side_effect=lambda m, f, ln: (lambda a: setattr(a, "filename", f) or a)( - SFTPAttributes() - ), - ): - return (CMD_NAME, msg) - raise EOFError - if cmd == CMD_CLOSE: - return (CMD_CLOSE, MagicMock()) - - mock_sftp._request.side_effect = fake_request - mock_sftp._adjust_cwd.return_value = "/home/user" - # fallback path since _from_msg patching is complex - mock_sftp._request.side_effect = TypeError - mock_sftp.listdir_attr.return_value = [] - result = client._listdir_attr_safe(mock_sftp, "/home/user") - assert result == [] - - def test_invalid_handle_response_raises(self, sftp_info: ConnectionInfo) -> None: - from paramiko.sftp import CMD_NAME - from paramiko.sftp_client import SFTPError - - client, mock_sftp = self._make_connected(sftp_info) - handle_msg = MagicMock() - mock_sftp._request.return_value = (CMD_NAME, handle_msg) # wrong type - mock_sftp._adjust_cwd.return_value = "/home/user" - - with pytest.raises(SFTPError): - client._listdir_attr_safe(mock_sftp, "/home/user") - - -class TestParseLsModeExtra: - def _parse(self, s): - from portkeydrop.protocols import SFTPClient - - return SFTPClient._parse_ls_mode(s) - - def test_fifo(self) -> None: - import stat - - assert stat.S_ISFIFO(self._parse("prwxrwxrwx")) - - def test_block_device(self) -> None: - import stat - - assert stat.S_ISBLK(self._parse("brwxrwxrwx")) - - def test_char_device(self) -> None: - import stat - - assert stat.S_ISCHR(self._parse("crwxrwxrwx")) - - def test_unknown_type_char(self) -> None: - mode = self._parse("?rwxrwxrwx") - assert mode != 0 or mode == 0 # just doesn't raise - - -class TestReopenSftp: - def test_reopen_sftp_failure_marks_disconnected(self, sftp_info: ConnectionInfo) -> None: - from portkeydrop.protocols import SFTPClient - - client = SFTPClient(sftp_info) - client._ssh_client = MagicMock() - client._sftp = MagicMock() - client._connected = True - client._ssh_client.open_sftp.side_effect = Exception("transport closed") - - client._reopen_sftp() - assert not client._connected - assert client._sftp is None - - def test_reopen_sftp_success(self, sftp_info: ConnectionInfo) -> None: - from portkeydrop.protocols import SFTPClient - - client = SFTPClient(sftp_info) - client._ssh_client = MagicMock() - client._sftp = MagicMock() - client._connected = True - new_sftp = MagicMock() - client._ssh_client.open_sftp.return_value = new_sftp - - client._reopen_sftp() - assert client._sftp is new_sftp - - -class TestListDirExecFallbackEdges: - def _make_connected(self, sftp_info): - from portkeydrop.protocols import SFTPClient - - client = SFTPClient(sftp_info) - client._ssh_client = MagicMock() - client._sftp = MagicMock() - client._connected = True - client._cwd = "/home/user" - return client - - def test_short_ls_lines_skipped(self, sftp_info: ConnectionInfo) -> None: - client = self._make_connected(sftp_info) - ls_output = "total 0\ndrwxr-xr-x\n" - stdout_mock = MagicMock() - stdout_mock.read.return_value = ls_output.encode() - client._ssh_client.exec_command.return_value = (MagicMock(), stdout_mock, MagicMock()) - result = client._list_dir_via_exec("/path") - assert result == [] - - def test_invalid_size_defaults_to_zero(self, sftp_info: ConnectionInfo) -> None: - client = self._make_connected(sftp_info) - ls_output = "-rw-r--r-- 1 user group BAD Jan 01 12:00 file.txt\n" - stdout_mock = MagicMock() - stdout_mock.read.return_value = ls_output.encode() - client._ssh_client.exec_command.return_value = (MagicMock(), stdout_mock, MagicMock()) - result = client._list_dir_via_exec("/path") - assert len(result) == 1 - assert result[0].st_size == 0 diff --git a/tests/test_ssh_agent_auth.py b/tests/test_ssh_agent_auth.py index bf5d687..6247eb1 100644 --- a/tests/test_ssh_agent_auth.py +++ b/tests/test_ssh_agent_auth.py @@ -8,8 +8,9 @@ from __future__ import annotations from unittest import mock +from unittest.mock import AsyncMock -import paramiko +import asyncssh import pytest from portkeydrop.protocols import ConnectionInfo, HostKeyPolicy, Protocol, SFTPClient @@ -28,39 +29,40 @@ def sftp_info() -> ConnectionInfo: class TestAgentAuthAttemptedFirst: - """Verify agent authentication is attempted first (allow_agent=True).""" + """Verify agent authentication is attempted first (no agent_path override).""" - def test_connect_uses_allow_agent_true_by_default(self, sftp_info: ConnectionInfo) -> None: + def test_connect_allows_agent_by_default(self, sftp_info: ConnectionInfo) -> None: client = SFTPClient(sftp_info) - with mock.patch("paramiko.SSHClient") as MockSSHClient: - mock_instance = MockSSHClient.return_value - mock_instance.connect.return_value = None - mock_sftp = mock.MagicMock() - mock_instance.open_sftp.return_value = mock_sftp - mock_sftp.normalize.return_value = "/" + with mock.patch("asyncssh.connect", new_callable=AsyncMock) as mock_connect: + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/" + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn client.connect() - mock_instance.connect.assert_called_once() - call_kwargs = mock_instance.connect.call_args[1] - assert call_kwargs["allow_agent"] is True - assert call_kwargs["look_for_keys"] is True + mock_connect.assert_called_once() + call_kwargs = mock_connect.call_args[1] + # No agent_path override = agent is allowed + assert "agent_path" not in call_kwargs + assert "client_keys" not in call_kwargs - def test_connect_with_password_still_uses_agent(self, sftp_info: ConnectionInfo) -> None: + def test_connect_with_password_still_allows_agent(self, sftp_info: ConnectionInfo) -> None: sftp_info.password = "secret" client = SFTPClient(sftp_info) - with mock.patch("paramiko.SSHClient") as MockSSHClient: - mock_instance = MockSSHClient.return_value - mock_instance.connect.return_value = None - mock_sftp = mock.MagicMock() - mock_instance.open_sftp.return_value = mock_sftp - mock_sftp.normalize.return_value = "/" + with mock.patch("asyncssh.connect", new_callable=AsyncMock) as mock_connect: + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/" + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn client.connect() - call_kwargs = mock_instance.connect.call_args[1] + call_kwargs = mock_connect.call_args[1] # When password is provided (no key_path), agent is still allowed - assert call_kwargs["allow_agent"] is True + assert "agent_path" not in call_kwargs assert call_kwargs["password"] == "secret" @@ -70,40 +72,40 @@ class TestFallbackToKeyFile: def test_connect_with_key_path_disables_agent(self, sftp_info: ConnectionInfo) -> None: sftp_info.key_path = "/home/user/.ssh/id_rsa" client = SFTPClient(sftp_info) - with mock.patch("paramiko.SSHClient") as MockSSHClient: - mock_instance = MockSSHClient.return_value - mock_instance.connect.return_value = None - mock_sftp = mock.MagicMock() - mock_instance.open_sftp.return_value = mock_sftp - mock_sftp.normalize.return_value = "/" + with mock.patch("asyncssh.connect", new_callable=AsyncMock) as mock_connect: + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/" + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn with mock.patch("os.path.exists", return_value=True): client.connect() - call_kwargs = mock_instance.connect.call_args[1] - assert call_kwargs["allow_agent"] is False - assert call_kwargs["look_for_keys"] is False - assert call_kwargs["key_filename"] == "/home/user/.ssh/id_rsa" + call_kwargs = mock_connect.call_args[1] + assert call_kwargs["agent_path"] is None + assert call_kwargs["client_keys"] == ["/home/user/.ssh/id_rsa"] def test_key_path_takes_precedence_over_password(self, sftp_info: ConnectionInfo) -> None: sftp_info.key_path = "/home/user/.ssh/id_rsa" sftp_info.password = "secret" client = SFTPClient(sftp_info) - with mock.patch("paramiko.SSHClient") as MockSSHClient: - mock_instance = MockSSHClient.return_value - mock_instance.connect.return_value = None - mock_sftp = mock.MagicMock() - mock_instance.open_sftp.return_value = mock_sftp - mock_sftp.normalize.return_value = "/" + with mock.patch("asyncssh.connect", new_callable=AsyncMock) as mock_connect: + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/" + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn with mock.patch("os.path.exists", return_value=True): client.connect() - call_kwargs = mock_instance.connect.call_args[1] - # key_path branch: agent disabled, password not passed - assert call_kwargs["allow_agent"] is False + call_kwargs = mock_connect.call_args[1] + # key_path branch: agent disabled, password used as passphrase + assert call_kwargs["agent_path"] is None assert "password" not in call_kwargs - assert call_kwargs["key_filename"] == "/home/user/.ssh/id_rsa" + assert call_kwargs["client_keys"] == ["/home/user/.ssh/id_rsa"] + assert call_kwargs["passphrase"] == "secret" class TestFallbackToPassword: @@ -112,38 +114,37 @@ class TestFallbackToPassword: def test_connect_with_password_only(self, sftp_info: ConnectionInfo) -> None: sftp_info.password = "mypassword" client = SFTPClient(sftp_info) - with mock.patch("paramiko.SSHClient") as MockSSHClient: - mock_instance = MockSSHClient.return_value - mock_instance.connect.return_value = None - mock_sftp = mock.MagicMock() - mock_instance.open_sftp.return_value = mock_sftp - mock_sftp.normalize.return_value = "/" + with mock.patch("asyncssh.connect", new_callable=AsyncMock) as mock_connect: + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/" + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn client.connect() - call_kwargs = mock_instance.connect.call_args[1] + call_kwargs = mock_connect.call_args[1] assert call_kwargs["password"] == "mypassword" # Agent still allowed as a first attempt - assert call_kwargs["allow_agent"] is True + assert "agent_path" not in call_kwargs def test_connect_no_credentials_uses_agent_and_key_discovery( self, sftp_info: ConnectionInfo ) -> None: client = SFTPClient(sftp_info) - with mock.patch("paramiko.SSHClient") as MockSSHClient: - mock_instance = MockSSHClient.return_value - mock_instance.connect.return_value = None - mock_sftp = mock.MagicMock() - mock_instance.open_sftp.return_value = mock_sftp - mock_sftp.normalize.return_value = "/" + with mock.patch("asyncssh.connect", new_callable=AsyncMock) as mock_connect: + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/" + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn client.connect() - call_kwargs = mock_instance.connect.call_args[1] - assert call_kwargs["allow_agent"] is True - assert call_kwargs["look_for_keys"] is True + call_kwargs = mock_connect.call_args[1] + assert "agent_path" not in call_kwargs + assert "client_keys" not in call_kwargs assert "password" not in call_kwargs - assert "key_filename" not in call_kwargs class TestErrorHandlingAllMethodsFail: @@ -151,22 +152,18 @@ class TestErrorHandlingAllMethodsFail: def test_connection_error_raised_on_auth_failure(self, sftp_info: ConnectionInfo) -> None: client = SFTPClient(sftp_info) - with mock.patch("paramiko.SSHClient") as MockSSHClient: - mock_instance = MockSSHClient.return_value - mock_instance.connect.side_effect = paramiko.AuthenticationException( - "All auth methods failed" - ) + with mock.patch("asyncssh.connect", new_callable=AsyncMock) as mock_connect: + mock_connect.side_effect = asyncssh.PermissionDenied("All auth methods failed") with pytest.raises(ConnectionError, match="SFTP connection failed"): client.connect() assert client.connected is False - def test_connection_error_on_ssh_exception(self, sftp_info: ConnectionInfo) -> None: + def test_connection_error_on_disconnect_error(self, sftp_info: ConnectionInfo) -> None: client = SFTPClient(sftp_info) - with mock.patch("paramiko.SSHClient") as MockSSHClient: - mock_instance = MockSSHClient.return_value - mock_instance.connect.side_effect = paramiko.SSHException("SSH error") + with mock.patch("asyncssh.connect", new_callable=AsyncMock) as mock_connect: + mock_connect.side_effect = asyncssh.DisconnectError(11, "SSH error") with pytest.raises(ConnectionError, match="SFTP connection failed"): client.connect() @@ -175,9 +172,8 @@ def test_connection_error_on_ssh_exception(self, sftp_info: ConnectionInfo) -> N def test_connection_error_on_socket_error(self, sftp_info: ConnectionInfo) -> None: client = SFTPClient(sftp_info) - with mock.patch("paramiko.SSHClient") as MockSSHClient: - mock_instance = MockSSHClient.return_value - mock_instance.connect.side_effect = OSError("Connection refused") + with mock.patch("asyncssh.connect", new_callable=AsyncMock) as mock_connect: + mock_connect.side_effect = OSError("Connection refused") with pytest.raises(ConnectionError, match="SFTP connection failed"): client.connect() @@ -187,7 +183,7 @@ def test_connection_error_on_socket_error(self, sftp_info: ConnectionInfo) -> No def test_connection_error_on_key_file_not_found(self, sftp_info: ConnectionInfo) -> None: sftp_info.key_path = "/nonexistent/key" client = SFTPClient(sftp_info) - with mock.patch("paramiko.SSHClient"): + with mock.patch("asyncssh.connect", new_callable=AsyncMock): with mock.patch("os.path.exists", return_value=False): with pytest.raises(ConnectionError, match="key file not found"): client.connect() @@ -197,10 +193,10 @@ def test_connection_error_on_key_file_not_found(self, sftp_info: ConnectionInfo) def test_sftp_session_failure(self, sftp_info: ConnectionInfo) -> None: """Test error when SSH connects but SFTP session fails.""" client = SFTPClient(sftp_info) - with mock.patch("paramiko.SSHClient") as MockSSHClient: - mock_instance = MockSSHClient.return_value - mock_instance.connect.return_value = None - mock_instance.open_sftp.return_value = None + with mock.patch("asyncssh.connect", new_callable=AsyncMock) as mock_connect: + mock_conn = AsyncMock() + mock_conn.start_sftp_client.return_value = None + mock_connect.return_value = mock_conn with pytest.raises(ConnectionError, match="Failed to create SFTP session"): client.connect() @@ -214,31 +210,30 @@ class TestHostKeyPolicies: def test_strict_host_key_policy(self, sftp_info: ConnectionInfo) -> None: sftp_info.host_key_policy = HostKeyPolicy.STRICT client = SFTPClient(sftp_info) - with mock.patch("paramiko.SSHClient") as MockSSHClient: - mock_instance = MockSSHClient.return_value - mock_instance.connect.return_value = None - mock_sftp = mock.MagicMock() - mock_instance.open_sftp.return_value = mock_sftp - mock_sftp.normalize.return_value = "/" + with mock.patch("asyncssh.connect", new_callable=AsyncMock) as mock_connect: + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/" + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn client.connect() - # Should have set RejectPolicy for strict mode - mock_instance.set_missing_host_key_policy.assert_called() - policy_arg = mock_instance.set_missing_host_key_policy.call_args[0][0] - assert isinstance(policy_arg, paramiko.RejectPolicy) + call_kwargs = mock_connect.call_args[1] + # Strict mode uses default known_hosts (no override) + assert "known_hosts" not in call_kwargs def test_auto_add_host_key_policy(self, sftp_info: ConnectionInfo) -> None: sftp_info.host_key_policy = HostKeyPolicy.AUTO_ADD client = SFTPClient(sftp_info) - with mock.patch("paramiko.SSHClient") as MockSSHClient: - mock_instance = MockSSHClient.return_value - mock_instance.connect.return_value = None - mock_sftp = mock.MagicMock() - mock_instance.open_sftp.return_value = mock_sftp - mock_sftp.normalize.return_value = "/" + with mock.patch("asyncssh.connect", new_callable=AsyncMock) as mock_connect: + mock_conn = AsyncMock() + mock_sftp = AsyncMock() + mock_sftp.realpath.return_value = "/" + mock_conn.start_sftp_client.return_value = mock_sftp + mock_connect.return_value = mock_conn client.connect() - policy_arg = mock_instance.set_missing_host_key_policy.call_args[0][0] - assert isinstance(policy_arg, paramiko.AutoAddPolicy) + call_kwargs = mock_connect.call_args[1] + assert call_kwargs["known_hosts"] is None diff --git a/tests/test_ssh_utils.py b/tests/test_ssh_utils.py index 153fd22..265369b 100644 --- a/tests/test_ssh_utils.py +++ b/tests/test_ssh_utils.py @@ -4,14 +4,7 @@ from unittest import mock -import paramiko - -from portkeydrop import __version__ -from portkeydrop.ssh_utils import ( - _portkeydrop_transport_init, - check_ssh_agent_available, - create_ssh_client, -) +from portkeydrop.ssh_utils import check_ssh_agent_available class TestCheckSshAgentAvailable: @@ -40,67 +33,8 @@ def test_windows_named_pipe_detected(self): mock_exists.side_effect = lambda p: p == r"\\.\pipe\openssh-ssh-agent" assert check_ssh_agent_available() is True - def test_windows_pageant_detected(self): - with mock.patch.dict("os.environ", {}, clear=True): - with mock.patch("platform.system", return_value="Windows"): - with mock.patch("os.path.exists", return_value=False): - mock_agent = mock.MagicMock() - mock_agent.get_keys.return_value = [mock.MagicMock()] - with mock.patch("paramiko.Agent", return_value=mock_agent): - assert check_ssh_agent_available() is True - def test_windows_no_agent(self): with mock.patch.dict("os.environ", {}, clear=True): with mock.patch("platform.system", return_value="Windows"): with mock.patch("os.path.exists", return_value=False): - mock_agent = mock.MagicMock() - mock_agent.get_keys.return_value = [] - with mock.patch("paramiko.Agent", return_value=mock_agent): - assert check_ssh_agent_available() is False - - -class TestCreateSshClient: - """Tests for create_ssh_client().""" - - def test_returns_ssh_client(self): - client = create_ssh_client() - assert isinstance(client, paramiko.SSHClient) - - def test_auto_add_policy_by_default(self): - client = create_ssh_client() - assert isinstance(client._policy, paramiko.AutoAddPolicy) - - def test_reject_policy_when_auto_add_disabled(self): - client = create_ssh_client(auto_add_host_key=False) - assert isinstance(client._policy, paramiko.RejectPolicy) - - def test_allow_agent_stored(self): - client = create_ssh_client(allow_agent=False) - assert client._allow_agent is False # type: ignore[attr-defined] - - def test_look_for_keys_stored(self): - client = create_ssh_client(look_for_keys=False) - assert client._look_for_keys is False # type: ignore[attr-defined] - - def test_default_settings(self): - client = create_ssh_client() - assert client._allow_agent is True # type: ignore[attr-defined] - assert client._look_for_keys is True # type: ignore[attr-defined] - - -class TestSshBannerPatch: - """Tests for the paramiko Transport banner patch.""" - - def test_banner_patch_sets_local_version(self): - """_portkeydrop_transport_init patches local_version on the transport.""" - mock_self = mock.MagicMock() - with mock.patch("portkeydrop.ssh_utils._original_transport_init"): - _portkeydrop_transport_init(mock_self) - assert mock_self.local_version == f"SSH-2.0-PortkeyDrop_{__version__}" - - def test_banner_patch_calls_original_init(self): - """_portkeydrop_transport_init delegates to the original __init__.""" - mock_self = mock.MagicMock() - with mock.patch("portkeydrop.ssh_utils._original_transport_init") as orig: - _portkeydrop_transport_init(mock_self, "arg1", kw="val") - orig.assert_called_once_with(mock_self, "arg1", kw="val") + assert check_ssh_agent_available() is False