from __future__ import annotations import io import stat from dataclasses import dataclass from pathlib import Path import paramiko import pytest from backup_tool.adapters import SourceError from backup_tool.ssh_adapter import SSHAdapter, load_private_key from backup_tool.ssh_source import SSHSourcePublicConfig from tests.conftest import make_settings @dataclass class Attributes: filename: str st_mode: int st_size: int = 0 st_mtime: int = 1 class ServerKey: def __init__(self, algorithm: str = "ssh-ed25519", encoded: str = "AQID") -> None: self.algorithm = algorithm self.encoded = encoded def get_name(self) -> str: return self.algorithm def get_base64(self) -> str: return self.encoded class Transport: def __init__(self, key: ServerKey) -> None: self.key = key self.events: list[str] = [] self.closed = False def start_client(self, *, timeout: float) -> None: self.events.append("start") def get_remote_server_key(self) -> ServerKey: self.events.append("host_key") return self.key def auth_publickey(self, username: str, private_key: object) -> None: self.events.append("auth") def close(self) -> None: self.closed = True class Channel: def __init__(self, events: list[str]) -> None: self.events = events def settimeout(self, timeout: float) -> None: self.events.append("timeout") class Handle: def __init__(self, chunks: list[bytes]) -> None: self.chunks = chunks self.closed = False def read(self, size: int) -> bytes: assert size == 4096 return self.chunks.pop(0) if self.chunks else b"" def close(self) -> None: self.closed = True class SFTP: def __init__(self, events: list[str], entries: list[Attributes]) -> None: self.events = events self.entries = entries self.handle = Handle([b"one", b"two"]) self.closed = False def get_channel(self) -> Channel: return Channel(self.events) def listdir_iter(self, path: str, *, read_aheads: int): self.events.append(f"list:{path}:{read_aheads}") return iter(self.entries) def lstat(self, path: str) -> Attributes: self.events.append(f"lstat:{path}") return Attributes("file", stat.S_IFREG | 0o640, 6, 1) def open(self, path: str, mode: str, bufsize: int) -> Handle: self.events.append(f"open:{path}:{mode}:{bufsize}") return self.handle def close(self) -> None: self.closed = True def config() -> SSHSourcePublicConfig: return SSHSourcePublicConfig( hostname="backup.example.test", port=22, username="backup", host_key="ssh-ed25519 AQID", root="/", ) def adapter( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, transport: Transport, sftp: SFTP ) -> SSHAdapter: monkeypatch.setattr("backup_tool.ssh_adapter.load_private_key", lambda _: object()) settings = make_settings(tmp_path).model_copy(update={"ssh_read_chunk_bytes": 4096}) return SSHAdapter( config(), "private-key-is-never-sent-to-a-log", settings, transport_factory=lambda *_: transport, sftp_factory=lambda _: sftp, ) @pytest.mark.asyncio async def test_host_pin_mismatch_never_authenticates_or_opens_sftp( tmp_path: Path, monkeypatch: pytest.MonkeyPatch ) -> None: transport = Transport(ServerKey(encoded="BAUG")) sftp = SFTP(transport.events, []) reader = adapter(tmp_path, monkeypatch, transport, sftp) with pytest.raises(SourceError, match="host key") as error: await reader.probe() assert error.value.reason_code == "source_trust" assert transport.events == ["start", "host_key"] assert transport.closed assert "timeout" not in sftp.events assert not any(event.startswith("list:") for event in sftp.events) @pytest.mark.asyncio async def test_pinned_transport_authenticates_before_sftp_and_streams_bounded_reads( tmp_path: Path, monkeypatch: pytest.MonkeyPatch ) -> None: transport = Transport(ServerKey()) sftp = SFTP(transport.events, [Attributes("file", stat.S_IFREG | 0o640, 6, 1)]) reader = adapter(tmp_path, monkeypatch, transport, sftp) entries = [entry async for entry in reader.enumerate_entries()] content = b"".join([chunk async for chunk in reader.open_content("file")]) await reader.close() assert entries[0].path == "file" assert content == b"onetwo" assert transport.events.index("host_key") < transport.events.index("auth") assert transport.events.index("auth") < transport.events.index("timeout") assert "list:/:32" in sftp.events assert "open:/file:rb:4096" in sftp.events assert sftp.handle.closed and sftp.closed and transport.closed @pytest.mark.asyncio @pytest.mark.parametrize("mode", [stat.S_IFLNK | 0o777, stat.S_IFIFO | 0o600]) async def test_sftp_rejects_symlinks_and_special_entries( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, mode: int ) -> None: transport = Transport(ServerKey()) sftp = SFTP(transport.events, [Attributes("unsafe", mode)]) reader = adapter(tmp_path, monkeypatch, transport, sftp) with pytest.raises(SourceError, match="symlink|unsupported"): await anext(reader.enumerate_entries()) assert "auth" in transport.events await reader.close() def test_private_key_loader_rejects_short_rsa_and_accepts_strong_rsa() -> None: short = paramiko.RSAKey.generate(2048) strong = paramiko.RSAKey.generate(3072) short_buffer = io.StringIO() strong_buffer = io.StringIO() short.write_private_key(short_buffer) strong.write_private_key(strong_buffer) with pytest.raises(SourceError, match="algorithm") as error: load_private_key(short_buffer.getvalue()) assert error.value.reason_code == "source_auth" loaded = load_private_key(strong_buffer.getvalue()) assert isinstance(loaded, paramiko.RSAKey)