"""Three-process parent-death supervision for Task22 snapshot workers."""
from __future__ import annotations

import ctypes
import errno
import math
import os
import select
import signal
import time
from collections.abc import Callable
from typing import Protocol
from task22_child_protocol import (
    MAX_RECEIPT, await_child, frame, kill_pidfd, validate_completion,
)
class PidfdOpen(Protocol):
    def __call__(self, pid: int) -> int: ...
def _close(descriptor: int) -> None:
    try:
        os.close(descriptor)
    except OSError as exc:
        if exc.errno != errno.EBADF:
            raise
def _exact_pidfd(pid: int, descriptor: int) -> None:
    with open(f"/proc/self/fdinfo/{descriptor}", encoding="ascii") as stream:
        identities = [line.removeprefix("Pid:").strip() for line in stream
                      if line.startswith("Pid:")]
    if identities != [str(pid)]:
        raise RuntimeError("lifecycle pidfd identity mismatch")
def _pdeathsig(supervisor: int) -> None:
    function = ctypes.CDLL(None, use_errno=True).prctl
    function.argtypes = [ctypes.c_int, ctypes.c_ulong, ctypes.c_ulong,
                         ctypes.c_ulong, ctypes.c_ulong]
    function.restype = ctypes.c_int
    if os.getppid() != supervisor:
        raise RuntimeError("worker supervisor exited before PDEATHSIG binding")
    if function(1, int(signal.SIGKILL), 0, 0, 0) != 0:  # PR_SET_PDEATHSIG
        error = ctypes.get_errno() or errno.EIO
        raise OSError(error, os.strerror(error))
    if os.getppid() != supervisor:
        raise RuntimeError("worker supervisor exited during PDEATHSIG binding")
def _read_exact(descriptor: int, expected: bytes, label: str) -> None:
    if os.read(descriptor, len(expected) + 1) != expected:
        raise RuntimeError(f"{label} lifecycle handshake failed")
def _wait_handshake(pidfd: int, descriptor: int, expected: bytes) -> None:
    poller = select.poll()
    poller.register(pidfd, select.POLLIN)
    poller.register(descriptor, select.POLLIN | select.POLLHUP)
    events = poller.poll(5_000)
    if events != [(descriptor, select.POLLIN)]:
        raise RuntimeError("supervisor lifecycle handshake event mismatch")
    _read_exact(descriptor, expected, "supervisor")

def _empty_receipt(descriptor: int) -> None:
    os.set_blocking(descriptor, False)
    content = bytearray()
    while True:
        try:
            block = os.read(descriptor, MAX_RECEIPT + 5)
        except BlockingIOError:
            break
        if not block:
            break
        content.extend(block)
    if content:
        raise RuntimeError("terminated worker produced an unauthorized receipt")

def _terminate(worker: int, worker_pidfd: int, receipt: int) -> None:
    kill_pidfd(worker, worker_pidfd)
    _empty_receipt(receipt)

def _monitor(
    parent_pidfd: int, control: int, worker: int, worker_pidfd: int,
    gate: int, receipt: int,
) -> tuple[bool, int]:
    poller = select.poll()
    poller.register(parent_pidfd, select.POLLIN)
    poller.register(control, select.POLLIN | select.POLLHUP)
    events = poller.poll(60_000)
    observed = dict(events)
    if observed.get(parent_pidfd):
        if observed[parent_pidfd] != select.POLLIN:
            raise RuntimeError("launcher parent pidfd event mismatch")
        _terminate(worker, worker_pidfd, receipt)
        return True, 0
    if observed.get(control) not in {select.POLLIN, select.POLLIN | select.POLLHUP}:
        _terminate(worker, worker_pidfd, receipt)
        raise RuntimeError("launcher release gate closed or returned an invalid event")
    _read_exact(control, b"R", "launcher release")
    _ = os.write(gate, b"1")
    _close(gate)

    os.set_blocking(receipt, False)
    poller = select.poll()
    poller.register(parent_pidfd, select.POLLIN)
    poller.register(worker_pidfd, select.POLLIN)
    poller.register(receipt, select.POLLIN | select.POLLHUP)
    deadline = time.monotonic() + 60
    chunks: list[bytes] = []
    total, exited, eof = 0, False, False
    while not (exited and eof):
        remaining = max(0, math.ceil((deadline - time.monotonic()) * 1000))
        events = poller.poll(remaining)
        if not events:
            _terminate(worker, worker_pidfd, receipt)
            raise RuntimeError("supervised worker exceeded the bounded execution window")
        observed = dict(events)
        if observed.get(parent_pidfd):
            if observed[parent_pidfd] != select.POLLIN:
                raise RuntimeError("launcher parent pidfd event mismatch")
            _terminate(worker, worker_pidfd, receipt)
            return True, 0
        worker_event = observed.get(worker_pidfd)
        if worker_event is not None:
            if worker_event != select.POLLIN:
                raise RuntimeError("worker pidfd event mismatch")
            exited = True
        receipt_event = observed.get(receipt)
        if receipt_event is not None:
            if receipt_event & ~(select.POLLIN | select.POLLHUP):
                raise RuntimeError("worker receipt event mismatch")
            while True:
                try:
                    block = os.read(receipt, MAX_RECEIPT + 5 - total)
                except BlockingIOError:
                    break
                if not block:
                    eof = True
                    break
                chunks.append(block)
                total += len(block)
                if total > MAX_RECEIPT + 4:
                    _terminate(worker, worker_pidfd, receipt)
                    raise RuntimeError("worker completion receipt exceeds bound")
    waited, status = os.waitpid(worker, 0)
    if waited != worker:
        raise RuntimeError("supervisor reaped the wrong worker")
    return False, validate_completion(b"".join(chunks), status)

def _supervisor(
    parent: int, opener: PidfdOpen, setup: int, ready: int, result: int, outcome: int,
    worker_body: Callable[[int, int], None], close_references: Callable[[], None],
) -> None:
    parent_pidfd = worker_pidfd = gate_write = receipt_read = -1
    worker, supervisor_pid = 0, os.getpid()
    try:
        parent_pidfd = opener(parent)
        _exact_pidfd(parent, parent_pidfd)
        _ = os.write(ready, b"B")
        _read_exact(setup, b"A", "parent pidfd acknowledgement")
        gate_read, gate_write = os.pipe2(os.O_CLOEXEC)
        receipt_read, receipt_write = os.pipe2(os.O_CLOEXEC)
        worker = os.fork()
        if worker == 0:
            _close(gate_write)
            _close(receipt_read)
            _close(parent_pidfd)
            _close(setup)
            _close(ready)
            _close(result)
            _close(outcome)
            _pdeathsig(supervisor_pid)
            worker_body(gate_read, receipt_write)
            os._exit(1)
        _close(gate_read)
        _close(receipt_write)
        worker_pidfd = opener(worker)
        _exact_pidfd(worker, worker_pidfd)
        close_references()
        _ = os.write(ready, b"W")
        parent_died, code = _monitor(
            parent_pidfd, setup, worker, worker_pidfd, gate_write, receipt_read
        )
        if parent_died:
            for descriptor in (worker_pidfd, receipt_read, setup, ready, result, outcome,
                               parent_pidfd):
                _close(descriptor)
            os._exit(0)
        _ = os.write(outcome, b"S")
        _ = os.write(result, frame("exit", code))
        os._exit(code)
    except BaseException as exc:
        if worker > 0:
            try:
                if worker_pidfd >= 0:
                    _terminate(worker, worker_pidfd, receipt_read)
                else:
                    os.kill(worker, signal.SIGKILL)
                    waited, _ = os.waitpid(worker, 0)
                    if waited != worker:
                        raise RuntimeError("supervisor fallback reaped wrong worker")
            except (ChildProcessError, OSError, RuntimeError):
                pass
        try:
            _ = os.write(outcome, b"E")
            _ = os.write(result, frame("error", 1, f"{type(exc).__name__}: {exc}"))
        except OSError:
            pass
        os._exit(1)

def supervise(
    opener: PidfdOpen, worker_body: Callable[[int, int], None],
    close_supervisor_references: Callable[[], None],
    parent_release: Callable[[], None],
) -> int:
    setup_read, setup_write = os.pipe2(os.O_CLOEXEC)
    ready_read, ready_write = os.pipe2(os.O_CLOEXEC)
    result_read, result_write = os.pipe2(os.O_CLOEXEC)
    outcome_read, outcome_write = os.pipe2(os.O_CLOEXEC)
    supervisor = os.fork()
    if supervisor == 0:
        _close(setup_write)
        _close(ready_read)
        _close(result_read)
        _close(outcome_read)
        _supervisor(os.getppid(), opener, setup_read, ready_write, result_write, outcome_write,
                    worker_body, close_supervisor_references)
    _close(setup_read)
    _close(ready_write)
    _close(result_write)
    _close(outcome_write)
    supervisor_pidfd = -1
    try:
        supervisor_pidfd = opener(supervisor)
        _wait_handshake(supervisor_pidfd, ready_read, b"B")
        _ = os.write(setup_write, b"A")
        _wait_handshake(supervisor_pidfd, ready_read, b"W")
        parent_release()
        _ = os.write(setup_write, b"R")
        _close(setup_write)
        setup_write = -1
        code = await_child(supervisor, supervisor_pidfd, result_read, 65_000)
        marker = os.read(outcome_read, 2)
        if marker != b"S":
            raise RuntimeError("supervisor completion failed closed")
        return code
    finally:
        for descriptor in (setup_write, ready_read, result_read, outcome_read, supervisor_pidfd):
            if descriptor >= 0:
                _close(descriptor)
        try:
            waited, _ = os.waitpid(supervisor, os.WNOHANG)
            # A setup failure closes the acknowledgement pipe before worker fork.
            if waited == 0:
                waited, _ = os.waitpid(supervisor, 0)
            if waited not in {0, supervisor}:
                raise RuntimeError("launcher reaped the wrong supervisor")
        except ChildProcessError:
            pass
