#!/usr/bin/env python3
"""Validate a router backup e-mail and forward its attachments over SSH."""

import hashlib
import json
import os
import re
import subprocess
import sys
from email.parser import BytesParser
from email.policy import default
from email.utils import parsedate_to_datetime
from pathlib import PurePath


MAX_TOTAL_SIZE = int(os.environ.get("BACKUP_MAX_BYTES", 50 * 1024 * 1024))
MAX_MESSAGE_SIZE = int(os.environ.get("BACKUP_MAX_MESSAGE_BYTES", 75 * 1024 * 1024))
SUBJECT_PATTERN = re.compile(r"^Router-Backup:\s+([A-Za-z0-9][A-Za-z0-9._-]*)$")
SAFE_NAME = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]*$")


class ConfigurationError(RuntimeError):
    pass


class ValidationError(ValueError):
    pass


def required_setting(name):
    value = os.environ.get(name, "").strip()
    if not value:
        raise ConfigurationError(f"Required environment variable is missing: {name}")
    return value


def extract_backup(raw_message):
    message = BytesParser(policy=default).parsebytes(raw_message)
    subject = str(message.get("Subject", "")).strip()
    match = SUBJECT_PATTERN.fullmatch(subject)
    if not match:
        raise ValidationError("Subject does not match the backup format")

    router = match.group(1)
    try:
        backup_time = parsedate_to_datetime(message.get("Date", ""))
    except (TypeError, ValueError) as error:
        raise ValidationError("Invalid Date header") from error
    if backup_time is None or backup_time.tzinfo is None or backup_time.utcoffset() is None:
        raise ValidationError("Date header must include a time zone")

    attachments = {}
    total_size = 0
    for part in message.walk():
        filename = part.get_filename()
        if filename is None:
            continue
        filename = str(filename)
        if PurePath(filename).name != filename or not SAFE_NAME.fullmatch(filename):
            raise ValidationError(f"Unsafe attachment name: {filename!r}")

        extension = filename.rsplit(".", 1)[-1].lower() if "." in filename else ""
        if extension not in {"backup", "rsc"}:
            raise ValidationError(f"Unexpected attachment: {filename}")
        if extension in attachments:
            raise ValidationError(f"More than one .{extension} attachment")
        if filename != f"{router}.{extension}":
            raise ValidationError(f"Attachment does not belong to {router}: {filename}")

        data = part.get_payload(decode=True)
        if data is None:
            raise ValidationError(f"Cannot decode attachment: {filename}")
        total_size += len(data)
        if total_size > MAX_TOTAL_SIZE:
            raise ValidationError("Attachments exceed the configured size limit")
        attachments[extension] = {"name": filename, "data": data}

    if set(attachments) != {"backup", "rsc"}:
        raise ValidationError("Exactly one .backup and one .rsc attachment are required")
    return router, backup_time.isoformat(), [attachments["backup"], attachments["rsc"]]


def transfer_backup(router, backup_time, attachments):
    nas_host = required_setting("BACKUP_NAS_HOST")
    nas_user = required_setting("BACKUP_NAS_USER")
    ssh_key = required_setting("BACKUP_SSH_KEY")
    known_hosts = required_setting("BACKUP_KNOWN_HOSTS")
    nas_port = os.environ.get("BACKUP_NAS_PORT", "22")

    metadata = {
        "router": router,
        "backup_time": backup_time,
        "files": [
            {
                "name": item["name"],
                "size": len(item["data"]),
                "sha256": hashlib.sha256(item["data"]).hexdigest(),
            }
            for item in attachments
        ],
    }
    payload = json.dumps(metadata, separators=(",", ":")).encode("ascii") + b"\n"
    payload += b"".join(item["data"] for item in attachments)

    command = [
        "ssh",
        "-T",
        "-p",
        nas_port,
        "-i",
        ssh_key,
        "-o",
        "BatchMode=yes",
        "-o",
        "IdentitiesOnly=yes",
        "-o",
        "StrictHostKeyChecking=yes",
        "-o",
        f"UserKnownHostsFile={known_hosts}",
        f"{nas_user}@{nas_host}",
    ]
    result = subprocess.run(
        command,
        input=payload,
        stdout=subprocess.PIPE,
        stderr=subprocess.PIPE,
        timeout=300,
        check=False,
    )
    if result.returncode != 0:
        error = result.stderr.decode("utf-8", "replace").strip()
        raise RuntimeError(error or f"SSH transfer exited with status {result.returncode}")


def main():
    try:
        raw_message = sys.stdin.buffer.read(MAX_MESSAGE_SIZE + 1)
        if len(raw_message) > MAX_MESSAGE_SIZE:
            raise ValidationError("Message exceeds the configured size limit")
        router, backup_time, attachments = extract_backup(raw_message)
        transfer_backup(router, backup_time, attachments)
        print(f"Archived backup for {router}", file=sys.stderr)
        return os.EX_OK
    except (ConfigurationError, subprocess.TimeoutExpired, RuntimeError) as error:
        print(f"Temporary backup processing failure: {error}", file=sys.stderr)
        return os.EX_TEMPFAIL
    except ValidationError as error:
        print(f"Rejected backup message: {error}", file=sys.stderr)
        return os.EX_DATAERR


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