#!/usr/bin/env python3
"""Receive validated router backup files on a restricted SSH account."""

import argparse
import hashlib
import json
import os
import re
import shutil
import sys
import tempfile
import uuid
from datetime import datetime, timedelta
from pathlib import Path, PurePath


MAX_HEADER_SIZE = 65536
MAX_TOTAL_SIZE = int(os.environ.get("BACKUP_MAX_BYTES", 50 * 1024 * 1024))
SAFE_NAME = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]*$")
SNAPSHOT_NAME = re.compile(r"^\d{8}-\d{6}-\d{6}-[0-9a-f]{8}$")
SHA256 = re.compile(r"^[0-9a-f]{64}$")


def read_to_file(stream, destination, size, expected_digest):
    digest = hashlib.sha256()
    remaining = size
    with destination.open("xb") as output:
        while remaining:
            chunk = stream.read(min(1024 * 1024, remaining))
            if not chunk:
                raise ValueError("Transfer ended before all bytes were received")
            output.write(chunk)
            digest.update(chunk)
            remaining -= len(chunk)
        output.flush()
        os.fsync(output.fileno())
    if digest.hexdigest() != expected_digest:
        raise ValueError(f"Checksum mismatch for {destination.name}")
    destination.chmod(0o600)


def validate_metadata(metadata):
    if not isinstance(metadata, dict):
        raise ValueError("Metadata must be an object")
    router = metadata.get("router")
    backup_time_value = metadata.get("backup_time")
    files = metadata.get("files")
    if not isinstance(router, str) or not SAFE_NAME.fullmatch(router):
        raise ValueError("Invalid router name")
    if not isinstance(backup_time_value, str):
        raise ValueError("Backup time is missing")
    try:
        backup_time = datetime.fromisoformat(backup_time_value)
    except ValueError as error:
        raise ValueError("Invalid backup time") from error
    if backup_time.tzinfo is None or backup_time.utcoffset() is None:
        raise ValueError("Backup time must include a time zone")
    backup_time = backup_time.astimezone()
    now = datetime.now().astimezone()
    if backup_time.year < 2000 or backup_time > now + timedelta(days=1):
        raise ValueError("Backup time is not plausible")
    if not isinstance(files, list) or len(files) != 2:
        raise ValueError("Exactly two files are required")

    extensions = set()
    total_size = 0
    for item in files:
        if not isinstance(item, dict):
            raise ValueError("Invalid file metadata")
        name = item.get("name")
        size = item.get("size")
        checksum = item.get("sha256")
        if not isinstance(name, str) or PurePath(name).name != name or not SAFE_NAME.fullmatch(name):
            raise ValueError("Invalid file name")
        extension = name.rsplit(".", 1)[-1].lower() if "." in name else ""
        if extension not in {"backup", "rsc"} or extension in extensions:
            raise ValueError("A .backup and a .rsc file are required")
        if name != f"{router}.{extension}":
            raise ValueError("File name does not match the router name")
        if not isinstance(size, int) or size < 0:
            raise ValueError("Invalid file size")
        if not isinstance(checksum, str) or not SHA256.fullmatch(checksum):
            raise ValueError("Invalid SHA-256 checksum")
        total_size += size
        if total_size > MAX_TOTAL_SIZE:
            raise ValueError("Files exceed the configured size limit")
        extensions.add(extension)
    return router, backup_time, files


def remove_snapshot(path):
    if path.is_dir() and SNAPSHOT_NAME.fullmatch(path.name):
        shutil.rmtree(path)


def prune_router(router_dir, now):
    snapshots = []
    for snapshot_dir in router_dir.glob("*/*/*"):
        try:
            year = int(snapshot_dir.parent.parent.name)
            month = int(snapshot_dir.parent.name)
        except ValueError:
            continue
        if snapshot_dir.is_dir() and SNAPSHOT_NAME.fullmatch(snapshot_dir.name):
            snapshots.append((year, month, snapshot_dir))

    for year in sorted({item[0] for item in snapshots}):
        year_items = [item for item in snapshots if item[0] == year]
        if year < now.year:
            keep = max(year_items, key=lambda item: item[2].name)[2]
            for _, _, path in year_items:
                if path != keep:
                    remove_snapshot(path)
        elif year == now.year:
            for month in sorted({item[1] for item in year_items if item[1] < now.month}):
                month_items = [item for item in year_items if item[1] == month]
                keep = max(month_items, key=lambda item: item[2].name)[2]
                for _, _, path in month_items:
                    if path != keep:
                        remove_snapshot(path)

    for directory in sorted(router_dir.glob("*/*"), reverse=True):
        if directory.is_dir() and not any(directory.iterdir()):
            directory.rmdir()
    for directory in sorted(router_dir.glob("*"), reverse=True):
        if directory.is_dir() and not any(directory.iterdir()):
            directory.rmdir()


def prune_all(archive_dir, now=None):
    now = now or datetime.now().astimezone()
    archive_dir.mkdir(parents=True, exist_ok=True, mode=0o700)
    for router_dir in archive_dir.iterdir():
        if router_dir.is_dir() and SAFE_NAME.fullmatch(router_dir.name):
            prune_router(router_dir, now)


def receive(base_dir):
    os.umask(0o077)
    incoming_dir = base_dir / "incoming"
    archive_dir = base_dir / "archive"
    header = sys.stdin.buffer.readline(MAX_HEADER_SIZE + 1)
    if not header.endswith(b"\n") or len(header) > MAX_HEADER_SIZE:
        raise ValueError("Metadata header is invalid or too large")
    metadata = json.loads(header.decode("ascii"))
    router, backup_time, files = validate_metadata(metadata)

    snapshot = backup_time.strftime("%Y%m%d-%H%M%S-%f-") + uuid.uuid4().hex[:8]
    incoming_dir.mkdir(parents=True, exist_ok=True, mode=0o700)
    staging = Path(tempfile.mkdtemp(prefix=f".{router}-", dir=incoming_dir))
    try:
        for item in files:
            read_to_file(sys.stdin.buffer, staging / item["name"], item["size"], item["sha256"])
        if sys.stdin.buffer.read(1):
            raise ValueError("Unexpected data follows the second file")

        archive_parent = archive_dir / router / backup_time.strftime("%Y") / backup_time.strftime("%m")
        archive_parent.mkdir(parents=True, exist_ok=True, mode=0o700)
        archived = archive_parent / snapshot
        os.replace(staging, archived)
        prune_router(archive_dir / router, datetime.now().astimezone())
        print(json.dumps({"router": router, "archive": str(archived)}, separators=(",", ":")))
    except Exception:
        if staging.exists():
            shutil.rmtree(staging)
        raise


def parse_args():
    parser = argparse.ArgumentParser()
    parser.add_argument("--base-dir", type=Path, default=Path("/srv/router-backups"))
    parser.add_argument("--prune", action="store_true")
    return parser.parse_args()


def main():
    args = parse_args()
    try:
        if args.prune:
            prune_all(args.base_dir / "archive")
        else:
            receive(args.base_dir)
        return os.EX_OK
    except Exception as error:
        print(f"Backup receiver failed: {error}", file=sys.stderr)
        return os.EX_DATAERR


if __name__ == "__main__":
    raise SystemExit(main())