#!/usr/bin/env python3
from __future__ import annotations

import dataclasses
import json
import logging
import os
import re
import shutil
import signal
import subprocess
import sys
import time
import urllib.parse
import urllib.request
from pathlib import Path
from typing import Callable, NamedTuple


OPTIONS_PATH = Path("/data/options.json")
PROFILE_PATH = Path("/tmp/cloakserve")
CLOAK_ROOT = "http://127.0.0.1:9222/"
CLOAK_VERSION = "http://127.0.0.1:9222/json/version?fingerprint=addon-readiness"
CLOAK_CLOSE = "http://127.0.0.1:9222/fingerprint/addon-readiness/close"
STELLOAUTH_ROOT = "http://127.0.0.1:8080/"
DURATION = re.compile(r"^[1-9][0-9]*(?:ms|s|m|h)$")
CLOAK_COMMAND = [
    "/usr/local/bin/cloakserve",
    "--headless=true",
    "--idle-timeout=30",
    "--data-dir=/tmp/cloakserve",
]
STELLOAUTH_COMMAND = ["/usr/local/bin/stelloauth"]
SHUTDOWN_GRACE_SECONDS = 9.0


class ConfigError(RuntimeError):
    pass


@dataclasses.dataclass(frozen=True)
class Options:
    queue_timeout: str
    rate_limit_count: int
    rate_limit_duration: str


class HttpResponse(NamedTuple):
    status: int
    body: bytes


class ReadinessError(RuntimeError):
    pass


def load_options(path: Path) -> Options:
    try:
        value = json.loads(path.read_text(encoding="utf-8"))
        if not isinstance(value, dict) or set(value) != {
            "queue_timeout",
            "rate_limit_count",
            "rate_limit_duration",
        }:
            raise ValueError

        queue_timeout = value["queue_timeout"]
        rate_limit_count = value["rate_limit_count"]
        rate_limit_duration = value["rate_limit_duration"]
        if not isinstance(queue_timeout, str) or DURATION.fullmatch(queue_timeout) is None:
            raise ValueError
        if (
            isinstance(rate_limit_count, bool)
            or not isinstance(rate_limit_count, int)
            or not 1 <= rate_limit_count <= 20
        ):
            raise ValueError
        if (
            not isinstance(rate_limit_duration, str)
            or DURATION.fullmatch(rate_limit_duration) is None
        ):
            raise ValueError
        return Options(queue_timeout, rate_limit_count, rate_limit_duration)
    except (OSError, UnicodeError, json.JSONDecodeError, KeyError, TypeError, ValueError):
        raise ConfigError("Invalid add-on configuration") from None


def build_environment(options: Options) -> dict[str, str]:
    environment = os.environ.copy()
    environment.update(
        {
            "CLOAK_CDP_URL": "http://127.0.0.1:9222",
            "CLOAK_MAX_SESSIONS": "1",
            "CLOAK_QUEUE_TIMEOUT": options.queue_timeout,
            "RATE_LIMIT_COUNT": str(options.rate_limit_count),
            "RATE_LIMIT_DURATION": options.rate_limit_duration,
            "HTTP_ADDRESS": "0.0.0.0",
            "PORT": "8080",
            "METRICS_ADDRESS": "127.0.0.1",
            "METRICS_PORT": "9090",
        }
    )
    return environment


def http_request(method: str, url: str, timeout: float) -> HttpResponse:
    opener = urllib.request.build_opener(urllib.request.ProxyHandler({}))
    data = b"" if method == "POST" else None
    request = urllib.request.Request(url, data=data, method=method)
    with opener.open(request, timeout=timeout) as response:
        return HttpResponse(response.status, response.read())


def probe_cloak(
    request: Callable[[str, str, float], HttpResponse] = http_request,
) -> None:
    if request("GET", CLOAK_ROOT, 2.0).status != 200:
        raise ReadinessError("CloakBrowser root probe failed")

    version = request("GET", CLOAK_VERSION, 2.0)
    if version.status != 200:
        raise ReadinessError("CloakBrowser CDP probe failed")
    try:
        document = json.loads(version.body)
        if not isinstance(document, dict):
            raise ValueError
        websocket_url = document["webSocketDebuggerUrl"]
        if (
            not isinstance(websocket_url, str)
            or not websocket_url.startswith("ws://127.0.0.1:")
        ):
            raise ValueError
        parsed = urllib.parse.urlsplit(websocket_url)
        if (
            parsed.scheme != "ws"
            or parsed.hostname != "127.0.0.1"
            or parsed.username is not None
            or parsed.password is not None
            or parsed.port is None
        ):
            raise ValueError
    except (json.JSONDecodeError, KeyError, TypeError, UnicodeError, ValueError):
        raise ReadinessError("CloakBrowser CDP response invalid") from None

    if request("POST", CLOAK_CLOSE, 2.0).status != 200:
        raise ReadinessError("CloakBrowser readiness profile close failed")


def probe_stelloauth(
    request: Callable[[str, str, float], HttpResponse] = http_request,
) -> None:
    if request("GET", STELLOAUTH_ROOT, 2.0).status != 200:
        raise ReadinessError("Stelloauth root probe failed")


def wait_until_ready(
    name: str,
    probe: Callable[[], None],
    timeout: float,
    stopping: Callable[[], bool],
    monotonic: Callable[[], float] = time.monotonic,
    sleep: Callable[[float], None] = time.sleep,
) -> None:
    deadline = monotonic() + timeout
    delay = 0.25
    while monotonic() < deadline and not stopping():
        try:
            probe()
            return
        except (OSError, ValueError, ReadinessError):
            sleep(delay)
            delay = min(delay * 2, 2.0)
    raise ReadinessError(f"{name} did not become ready")


def process_group_alive(
    pgid: int,
    killpg: Callable[[int, int], None] = os.killpg,
) -> bool:
    try:
        killpg(pgid, 0)
    except ProcessLookupError:
        return False
    except PermissionError:
        return True
    except OSError:
        return True
    return True


class ProcessManager:
    def __init__(
        self,
        environment: dict[str, str],
        *,
        popen: Callable[..., subprocess.Popen] = subprocess.Popen,
        killpg: Callable[[int, int], None] = os.killpg,
        group_alive: Callable[[int], bool] = process_group_alive,
        cloak_probe: Callable[[], None] = probe_cloak,
        stelloauth_probe: Callable[[], None] = probe_stelloauth,
        monotonic: Callable[[], float] = time.monotonic,
        sleep: Callable[[float], None] = time.sleep,
        profile_path: Path = PROFILE_PATH,
    ) -> None:
        self.environment = environment
        self._popen = popen
        self._killpg = killpg
        self._group_alive = group_alive
        self._cloak_probe = cloak_probe
        self._stelloauth_probe = stelloauth_probe
        self._monotonic = monotonic
        self._sleep = sleep
        self._profile_path = profile_path
        self._children: list[tuple[str, subprocess.Popen]] = []
        self._reaped_children: set[int] = set()
        self._groups: list[int] = []
        self._dead_groups: set[int] = set()
        self._stopping = False
        self._term_sent = False
        self._shutdown_deadline: float | None = None

    def _handle_signal(self, _signum: int, _frame: object) -> None:
        if self._stopping:
            return
        logging.info("Shutdown requested")
        self._stopping = True
        self._begin_shutdown()

    def _clean_profiles(self) -> None:
        logging.info("Cleaning CloakBrowser profiles")
        if self._profile_path.is_symlink():
            self._profile_path.unlink()
        elif self._profile_path.exists():
            shutil.rmtree(self._profile_path)
        self._profile_path.mkdir(parents=True, mode=0o700)
        self._profile_path.chmod(0o700)

    def _spawn(self, name: str, command: list[str]) -> subprocess.Popen:
        process = self._popen(
            command,
            env=self.environment,
            start_new_session=True,
        )
        self._children.append((name, process))
        self._groups.append(process.pid)
        if self._stopping:
            self._signal_groups(signal.SIGTERM, [process.pid])
        return process

    def _first_exited_child(self) -> tuple[str, int] | None:
        for name, process in self._children:
            returncode = process.poll()
            if returncode is not None:
                return name, returncode
        return None

    def _startup_stopping(self) -> bool:
        return self._stopping or self._first_exited_child() is not None

    def _reap_exited_children(self) -> None:
        for _name, process in self._children:
            if process.pid in self._reaped_children or process.poll() is None:
                continue
            try:
                process.wait(timeout=None)
            except OSError:
                logging.error("Child process reap failed")
            else:
                self._reaped_children.add(process.pid)

    def _live_groups(self) -> list[int]:
        self._reap_exited_children()
        live_groups = []
        for pgid in self._groups:
            if pgid in self._dead_groups:
                continue
            if self._group_alive(pgid):
                live_groups.append(pgid)
            else:
                self._dead_groups.add(pgid)
        return live_groups

    def _signal_groups(
        self, sent_signal: int, groups: list[int] | None = None
    ) -> None:
        for pgid in self._live_groups() if groups is None else groups:
            if pgid in self._dead_groups:
                continue
            if groups is not None and not self._group_alive(pgid):
                self._dead_groups.add(pgid)
                continue
            try:
                self._killpg(pgid, sent_signal)
            except ProcessLookupError:
                self._dead_groups.add(pgid)
            except OSError:
                pass

    def _begin_shutdown(self) -> None:
        if self._shutdown_deadline is None:
            self._shutdown_deadline = self._monotonic() + SHUTDOWN_GRACE_SECONDS
        if not self._term_sent:
            self._signal_groups(signal.SIGTERM)
            self._term_sent = True

    def _reap_all(self) -> None:
        for _name, process in self._children:
            if process.pid in self._reaped_children:
                continue
            try:
                process.wait(timeout=None)
            except OSError:
                logging.error("Child process reap failed")
            else:
                self._reaped_children.add(process.pid)

    def _shutdown(self) -> None:
        if not self._children:
            return
        logging.info("Stopping child processes")
        self._begin_shutdown()
        assert self._shutdown_deadline is not None
        live_groups = self._live_groups()
        while live_groups and self._monotonic() < self._shutdown_deadline:
            remaining = self._shutdown_deadline - self._monotonic()
            self._sleep(min(0.25, max(0.0, remaining)))
            live_groups = self._live_groups()

        if live_groups:
            logging.info("Forcing child processes to stop")
            self._signal_groups(signal.SIGKILL, live_groups)
        self._reap_all()
        logging.info("Child processes stopped")

    @staticmethod
    def _failure_status(returncode: int) -> int:
        return returncode if returncode != 0 else 1

    def _startup_failure(self, service: str) -> int:
        exited = self._first_exited_child()
        if exited is not None:
            name, returncode = exited
            logging.error(f"{name} exited unexpectedly")
            status = self._failure_status(returncode)
        else:
            logging.error(f"{service} readiness failed")
            status = 1
        self._shutdown()
        return status

    def _run(self) -> int:
        try:
            self._clean_profiles()
            if self._stopping:
                self._shutdown()
                return 0

            logging.info("Starting CloakBrowser")
            self._spawn("CloakBrowser", CLOAK_COMMAND)
            wait_until_ready(
                "CloakBrowser",
                self._cloak_probe,
                60.0,
                self._startup_stopping,
                self._monotonic,
                self._sleep,
            )
            if self._stopping:
                self._shutdown()
                return 0
            if self._first_exited_child() is not None:
                return self._startup_failure("CloakBrowser")
            logging.info("CloakBrowser ready")

            logging.info("Starting Stelloauth")
            self._spawn("Stelloauth", STELLOAUTH_COMMAND)
            wait_until_ready(
                "Stelloauth",
                self._stelloauth_probe,
                30.0,
                self._startup_stopping,
                self._monotonic,
                self._sleep,
            )
            if self._stopping:
                self._shutdown()
                return 0
            if self._first_exited_child() is not None:
                return self._startup_failure("Stelloauth")
            logging.info("Stelloauth listening on 0.0.0.0:8080")
        except ReadinessError:
            if self._stopping:
                self._shutdown()
                return 0
            service = "Stelloauth" if len(self._children) > 1 else "CloakBrowser"
            return self._startup_failure(service)
        except OSError:
            logging.error("Process startup failed")
            self._shutdown()
            return 1

        while not self._stopping:
            exited = self._first_exited_child()
            if exited is not None:
                name, returncode = exited
                logging.error(f"{name} exited unexpectedly")
                status = self._failure_status(returncode)
                self._shutdown()
                return status
            self._sleep(0.25)

        self._shutdown()
        return 0

    def run(self) -> int:
        previous_handlers = {
            signal.SIGTERM: signal.signal(signal.SIGTERM, self._handle_signal),
            signal.SIGINT: signal.signal(signal.SIGINT, self._handle_signal),
        }
        try:
            return self._run()
        finally:
            for handled_signal, previous_handler in previous_handlers.items():
                signal.signal(handled_signal, previous_handler)


def main() -> int:
    logging.basicConfig(level=logging.INFO, format="%(message)s")
    try:
        options = load_options(OPTIONS_PATH)
    except ConfigError:
        logging.error("Invalid add-on configuration")
        return 2
    return ProcessManager(build_environment(options)).run()


if __name__ == "__main__":
    sys.exit(main())
