#!/usr/bin/env python3
"""Skava company backup client. Credentials come from SKAVA_BACKUP_TOKEN.

Stores each verified full/diff, never overwrites historical backups, and persists
an idempotency key before admission so interrupted runs can safely continue.
"""

from __future__ import annotations

import argparse
import fcntl
import hashlib
import json
import os
import secrets
import time
import uuid
import zipfile
from pathlib import Path
from typing import Any
from urllib.error import HTTPError
from urllib.parse import urlparse
from urllib.request import HTTPRedirectHandler, Request, build_opener


class BackupClientError(Exception):
    def __init__(
        self, message: str, *, status: int = 0, code: str = "", retry_after: float = 0
    ) -> None:
        super().__init__(message)
        self.status = status
        self.code = code
        self.retry_after = retry_after


class NoRedirect(HTTPRedirectHandler):
    def redirect_request(self, request, fp, code, msg, headers, newurl):
        raise BackupClientError(
            "Unexpected redirect; backup credentials were not forwarded."
        )


class BackupClient:
    def __init__(
        self, server: str, company_id: int, token: str, *, allow_http: bool = False
    ) -> None:
        parsed = urlparse(server)
        if (
            parsed.scheme not in {"https", "http"}
            or not parsed.hostname
            or (parsed.scheme != "https" and not allow_http)
        ):
            raise BackupClientError("Use HTTPS for backup authentication.")
        if (
            parsed.username
            or parsed.password
            or parsed.query
            or parsed.fragment
            or parsed.path not in {"", "/"}
        ):
            raise BackupClientError(
                "Use the server origin without credentials, path or query."
            )
        self.root = server.rstrip("/") + f"/api/v1/companies/{company_id}"
        self.token = token
        self.opener = build_opener(NoRedirect())
        self.identity = hashlib.sha256(self.root.encode()).hexdigest()[:16]

    def _request(self, path: str, payload: dict[str, Any] | None = None):
        raw = json.dumps(payload).encode() if payload is not None else None
        request = Request(
            self.root + path,
            data=raw,
            headers={
                "Authorization": "Bearer " + self.token,
                "Content-Type": "application/json",
                "Accept-Encoding": "identity",
            },
        )
        try:
            return self.opener.open(request, timeout=90)
        except HTTPError as error:
            try:
                body = json.loads(error.read())
                code = body.get("code", "HTTP_ERROR")
            except (ValueError, UnicodeError):
                code = "HTTP_ERROR"
            try:
                retry = max(0, float(error.headers.get("Retry-After", "0")))
            except (TypeError, ValueError):
                retry = 0
            error.close()
            wait = f" Retry after {int(retry) + 1} seconds." if retry else ""
            raise BackupClientError(
                f"Backup API: {code} (HTTP {error.code})." + wait,
                status=error.code,
                code=code,
                retry_after=retry,
            ) from None

    def _json(self, path: str, payload: dict[str, Any] | None = None) -> dict[str, Any]:
        with self._request(path, payload) as response:
            return json.load(response)

    @staticmethod
    def _save(path: Path, value: dict[str, Any]) -> None:
        partial = path.with_suffix(".tmp")
        with partial.open("w") as output:
            os.chmod(partial, 0o600)
            json.dump(value, output)
            output.flush()
            os.fsync(output.fileno())
        os.replace(partial, path)
        BackupClient._sync_directory(path.parent)

    @staticmethod
    def _sync_directory(path: Path) -> None:
        descriptor = os.open(path, os.O_RDONLY)
        try:
            os.fsync(descriptor)
        finally:
            os.close(descriptor)

    def backup(
        self, kind: str, directory: Path | str, *, max_wait: int = 86400
    ) -> Path:
        directory = Path(directory).resolve()
        directory.mkdir(mode=0o700, parents=True, exist_ok=True)
        with (directory / f".skava-{self.identity}.lock").open("a") as lock:
            os.chmod(lock.name, 0o600)
            try:
                fcntl.flock(lock, fcntl.LOCK_EX | fcntl.LOCK_NB)
            except BlockingIOError:
                raise BackupClientError(
                    "Another backup for this company is running in this directory."
                ) from None
            return self._backup(kind, directory, max_wait=max_wait)

    def _confirm(
        self, job_id: str, download_id: str, confirmation: dict[str, Any]
    ) -> None:
        for attempt in range(3):
            try:
                self._json(
                    f"/backups/{job_id}/downloads/{download_id}/confirm", confirmation
                )
                return
            except BackupClientError as error:
                if error.status != 429 or attempt == 2:
                    raise
                time.sleep(max(1, error.retry_after))

    def _record(
        self, state_path: Path, state: dict[str, Any], job: dict[str, Any], target: Path
    ) -> None:
        part = {"id": job["id"], "sha256": job["sha256"], "file": target.name}
        if job["kind"] == "full":
            state["full"] = part
            state["chain"] = [part]
        else:
            state.setdefault("chain", [state["full"]]).append(part)
        state.pop("pending", None)
        self._save(state_path, state)

    def _backup(self, kind: str, directory: Path, *, max_wait: int) -> Path:
        if kind not in {"full", "diff"}:
            raise BackupClientError("Choose full or diff.")
        directory = Path(directory).resolve()
        directory.mkdir(mode=0o700, parents=True, exist_ok=True)
        state_path = directory / f".skava-{self.identity}.json"
        state = json.loads(state_path.read_text()) if state_path.exists() else {}
        pending = state.get("pending")
        if pending is None:
            pending = {"kind": kind, "request_id": str(uuid.uuid4())}
            if kind == "diff":
                full = state.get("full")
                chain = state.get("chain", [full] if full else [])
                if not full or not chain or chain[0] != full:
                    raise BackupClientError(
                        "Keep the entire verified backup chain locally before requesting an incremental."
                    )
                for part in chain:
                    path = directory / part["file"]
                    if not path.is_file():
                        raise BackupClientError(
                            "A backup chain part is missing. Restore it before requesting another incremental."
                        )
                    sha = hashlib.sha256()
                    with path.open("rb") as source:
                        for chunk in iter(lambda: source.read(1024 * 1024), b""):
                            sha.update(chunk)
                    if not secrets.compare_digest(sha.hexdigest(), part["sha256"]):
                        raise BackupClientError(
                            "A local backup chain part is corrupt. Restore the verified file first."
                        )
                pending["base_id"] = full["id"]
                pending["parent_id"] = chain[-1]["id"]
            state["pending"] = pending
            self._save(state_path, state)
        if pending["kind"] != kind:
            raise BackupClientError(
                "Finish the pending backup before requesting another type."
            )
        try:
            job = self._json(
                "/backups",
                {
                    key: pending[key]
                    for key in ("kind", "request_id", "base_id", "parent_id")
                    if key in pending
                },
            )["job"]
        except BackupClientError as error:
            if error.code in {
                "BACKUP_QUOTA",
                "BACKUP_BUSY",
                "BACKUP_QUEUE_FULL",
                "BACKUP_FULL_REQUIRED",
                "BACKUP_BASE_CHANGED",
                "BACKUP_PARENT_CHANGED",
                "BACKUP_CONFIRM_REQUIRED",
                "BACKUP_RETRY_LIMIT",
            }:
                state.pop("pending", None)
                self._save(state_path, state)
            raise
        pending["job_id"] = job["id"]
        self._save(state_path, state)
        deadline = time.monotonic() + max_wait
        while job["status"] in {"queued", "running"}:
            if time.monotonic() >= deadline:
                raise BackupClientError(
                    "Backup still queued. Run the same command later to resume."
                )
            time.sleep(10)
            try:
                job = self._json(f"/backups/{job['id']}")["job"]
            except BackupClientError as error:
                if error.status != 429 and error.status < 500:
                    raise
                time.sleep(
                    min(max(10, error.retry_after), max(0, deadline - time.monotonic()))
                )
        if job["status"] != "ready":
            state.pop("pending", None)
            self._save(state_path, state)
            raise BackupClientError(
                "Backup generation failed; no complete backup is available."
            )
        target = directory / f"skava-{kind}-{job['id']}.zip"
        partial = target.with_suffix(".zip.part")
        if "receipt" in pending and target.is_file():
            # Recheck the durable file before resuming an interrupted confirmation.
            sha = hashlib.sha256()
            with target.open("rb") as source:
                for chunk in iter(lambda: source.read(1024 * 1024), b""):
                    sha.update(chunk)
            if target.stat().st_size == job["bytes"] and secrets.compare_digest(
                sha.hexdigest(), job["sha256"]
            ):
                self._confirm(job["id"], pending["download_id"], pending["receipt"])
                self._record(state_path, state, job, target)
                return target
        if not job["available"]:
            state.pop("pending", None)
            self._save(state_path, state)
            raise BackupClientError("The download expired. Request a new backup.")
        sha = hashlib.sha256()
        size = 0
        try:
            with (
                self._request(f"/backups/{job['id']}/download") as response,
                partial.open("wb") as output,
            ):
                os.chmod(partial, 0o600)
                receipt = response.headers.get("X-Backup-Receipt")
                download_id = response.headers.get("X-Backup-Download-Id")
                if not receipt or not download_id:
                    raise BackupClientError("Missing download receipt.")
                while chunk := response.read(1024 * 1024):
                    size += len(chunk)
                    if size > job["bytes"]:
                        raise BackupClientError("Unexpected backup size.")
                    output.write(chunk)
                    sha.update(chunk)
                output.flush()
                os.fsync(output.fileno())
            if size != job["bytes"] or not secrets.compare_digest(
                sha.hexdigest(), job["sha256"]
            ):
                raise BackupClientError(
                    "Incomplete download or incorrect checksum. No success was confirmed."
                )
            with zipfile.ZipFile(partial) as archive:
                manifest = json.loads(archive.read("manifest.json"))
                if (
                    manifest.get("version") != 2
                    or manifest["job_id"] != job["id"]
                    or manifest["kind"] != kind
                ):
                    raise BackupClientError("Incorrect backup manifest.")
                if kind == "diff" and (
                    manifest["base_id"] != state["full"]["id"]
                    or manifest["base_sha256"] != state["full"]["sha256"]
                    or manifest["parent_id"] != pending["parent_id"]
                    or manifest["parent_sha256"]
                    != state.get("chain", [state["full"]])[-1]["sha256"]
                ):
                    raise BackupClientError(
                        "The incremental does not match the locally saved backup chain."
                    )
                if archive.testzip() is not None:
                    raise BackupClientError("Corrupt backup entry.")
            os.replace(partial, target)
            self._sync_directory(directory)
            confirmation = {
                "receipt": receipt,
                "bytes": size,
                "sha256": sha.hexdigest(),
            }
            # Preserve the receipt across process death after durable local save.
            state["pending"]["receipt"] = confirmation
            state["pending"]["download_id"] = download_id
            self._save(state_path, state)
            self._confirm(job["id"], download_id, confirmation)
            self._record(state_path, state, job, target)
            return target
        finally:
            partial.unlink(missing_ok=True)


def main() -> None:
    parser = argparse.ArgumentParser(
        description="Download and verify Skava company chat backups."
    )
    parser.add_argument("--server", default="https://chat.skava.io")
    parser.add_argument("--company", type=int, required=True)
    parser.add_argument("--kind", choices=["full", "diff"], required=True)
    parser.add_argument("--directory", type=Path, default=Path("skava-backups"))
    args = parser.parse_args()
    token = os.environ.get("SKAVA_BACKUP_TOKEN")
    if not token:
        parser.error("Set SKAVA_BACKUP_TOKEN in the protected process environment.")
    try:
        target = BackupClient(args.server, args.company, token).backup(
            args.kind, args.directory
        )
    except Exception as error:
        message = (
            str(error)
            if isinstance(error, BackupClientError)
            else "Backup failed. Run the same command to resume."
        )
        parser.exit(1, message + "\n")
    print(f"Backup received, verified and confirmed: {target.name}")


if __name__ == "__main__":
    main()
