196 lines
5.9 KiB
Python
196 lines
5.9 KiB
Python
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)
|