from __future__ import annotations

import ast
import ctypes
import errno
import hashlib
import importlib.util
import os
import select
import signal
import sys
from collections.abc import Callable
from pathlib import Path
from types import ModuleType
from typing import Protocol, cast

import pytest

ROOT = Path(__file__).resolve().parent


class _Syscall(Protocol):
    argtypes: list[object]
    restype: object

    def __call__(self, number: int, pid: int, flags: int, /) -> int: ...


def _protocol() -> ModuleType:
    spec = importlib.util.spec_from_file_location(
        "timeout_protocol", ROOT / "task22_child_protocol.py"
    )
    if spec is None or spec.loader is None:
        raise AssertionError("completion protocol loader missing")
    module = importlib.util.module_from_spec(spec)
    spec.loader.exec_module(module)
    return module


def _pidfd_open(pid: int) -> int:
    syscall = cast(_Syscall, cast(object, ctypes.CDLL(None, use_errno=True).syscall))
    syscall.argtypes = [ctypes.c_long, ctypes.c_int, ctypes.c_uint]
    syscall.restype = ctypes.c_long
    descriptor = syscall(434, pid, 0)
    if descriptor < 0:
        error = ctypes.get_errno()
        raise OSError(error, os.strerror(error))
    return descriptor


def _libc_kill(pidfd: int) -> None:
    function = ctypes.CDLL(None, use_errno=True).pidfd_send_signal
    function.argtypes = [ctypes.c_int, ctypes.c_int, ctypes.c_void_p, ctypes.c_uint]
    function.restype = ctypes.c_int
    if function(pidfd, signal.SIGKILL, None, 0) != 0:
        error = ctypes.get_errno()
        raise OSError(error, os.strerror(error))


def test_real_hanging_child_timeout_does_not_require_signal_module_binding(
    monkeypatch: pytest.MonkeyPatch,
) -> None:
    protocol = _protocol()
    await_child = cast(
        Callable[[int, int, int, int], int], getattr(protocol, "await_child")
    )
    ready_read, ready_write = os.pipe2(os.O_CLOEXEC)
    gate_read, gate_write = os.pipe2(os.O_CLOEXEC)
    receipt_read, receipt_write = os.pipe2(os.O_CLOEXEC)
    pid = os.fork()
    if pid == 0:
        os.close(ready_read)
        os.close(gate_write)
        os.close(receipt_read)
        _ = os.write(ready_write, b"1")
        _ = os.read(gate_read, 1)
        os._exit(99)
    os.close(ready_write)
    os.close(gate_read)
    os.close(receipt_write)
    pidfd = _pidfd_open(pid)
    try:
        poller = select.poll()
        poller.register(ready_read, select.POLLIN)
        assert poller.poll(2_000)
        assert os.read(ready_read, 1) == b"1"
        monkeypatch.delattr(signal, "pidfd_send_signal", raising=False)
        with pytest.raises(RuntimeError, match="bounded execution window"):
            _ = await_child(pid, pidfd, receipt_read, 1)
    finally:
        try:
            try:
                waited, _ = os.waitpid(pid, os.WNOHANG)
            except ChildProcessError:
                waited = pid
            if waited == 0:
                _libc_kill(pidfd)
                reaper = select.poll()
                reaper.register(pidfd, select.POLLIN)
                assert reaper.poll(2_000)
                waited, _ = os.waitpid(pid, 0)
            assert waited == pid
        finally:
            os.close(pidfd)
            os.close(receipt_read)
            os.close(gate_write)
            os.close(ready_read)


def test_timeout_tests_use_no_sleep_or_polling_delay() -> None:
    tree = ast.parse(Path(__file__).read_text(encoding="utf-8"))
    calls = [node.func for node in ast.walk(tree) if isinstance(node, ast.Call)]
    assert not any(
        isinstance(call, ast.Attribute) and call.attr == "sleep" for call in calls
    )


class _Seal:
    def __init__(self) -> None:
        self.closed: bool = False

    def close(self) -> None:
        self.closed = True


def _module(path: Path, name: str) -> ModuleType:
    spec = importlib.util.spec_from_file_location(name, path)
    if spec is None or spec.loader is None:
        raise AssertionError(f"cannot load {path}")
    module = importlib.util.module_from_spec(spec)
    spec.loader.exec_module(module)
    return module


def _runtime() -> ModuleType:
    for helper in ("task22_dependency_closure", "task22_child_protocol", "task22_resource_ownership"):
        sys.modules[helper] = _module(ROOT / f"{helper}.py", f"fault_{helper}")
    return _module(ROOT / "task22_launcher_runtime.py", "fault_runtime")


def _open_fds() -> set[int]:
    return {int(entry) for entry in os.listdir("/proc/self/fd")}


def _assert_fault_cleanup(
    monkeypatch: pytest.MonkeyPatch,
    invoke: Callable[[ModuleType, _Seal], object],
    error: int,
    *,
    seal_expected: bool,
) -> None:
    runtime = _runtime()
    seal = _Seal()
    before = _open_fds()
    children: list[int] = []
    real_fork = os.fork

    def fork() -> int:
        pid = real_fork()
        if pid > 0:
            children.append(pid)
        return pid

    def fail(_pid: int) -> int:
        raise OSError(error, os.strerror(error))

    monkeypatch.setattr(os, "fork", fork)
    monkeypatch.setattr(runtime, "pidfd_open", fail)
    caught: OSError | None = None
    leaked: set[int] = set()
    unreaped = False
    try:
        with pytest.raises(OSError) as raised:
            _ = invoke(runtime, seal)
        caught = raised.value
        leaked = _open_fds() - before
        assert len(children) == 1
        try:
            waited, _ = os.waitpid(children[0], os.WNOHANG)
            unreaped = waited == 0 or waited == children[0]
        except ChildProcessError:
            unreaped = False
    finally:
        for descriptor in leaked:
            try:
                os.close(descriptor)
            except OSError:
                pass
        for pid in children:
            try:
                while True:
                    try:
                        waited, _ = os.waitpid(pid, 0)
                        break
                    except InterruptedError:
                        continue
                assert waited == pid
            except ChildProcessError:
                pass
    assert caught is not None and caught.errno == error
    assert leaked == set()
    assert not unreaped
    assert seal.closed is seal_expected


@pytest.mark.parametrize("error", (errno.ENOSYS, errno.EPERM, errno.EMFILE))
@pytest.mark.parametrize("iteration", range(2))
def test_sealed_source_pidfd_open_failure_has_no_fd_delta_or_orphan(
    monkeypatch: pytest.MonkeyPatch, error: int, iteration: int
) -> None:
    del iteration

    def invoke(runtime: ModuleType, _seal: _Seal) -> object:
        sealed = cast(Callable[[str, bytes], int], getattr(runtime, "sealed_memfd"))
        run = cast(Callable[[int, str, list[str], Path], tuple[int, str, str]], getattr(runtime, "run_sealed_source"))
        descriptor = sealed("pidfd-fault-source", b"pass\n")
        try:
            return run(
                descriptor, hashlib.sha256(b"pass\n").hexdigest(), ["sealed"], ROOT
            )
        finally:
            os.close(descriptor)

    _assert_fault_cleanup(monkeypatch, invoke, error, seal_expected=False)


@pytest.mark.parametrize("error", (errno.ENOSYS, errno.EPERM, errno.EMFILE))
@pytest.mark.parametrize("iteration", range(2))
def test_snapshot_pidfd_open_failure_closes_seal_gate_fds_and_reaps(
    monkeypatch: pytest.MonkeyPatch, error: int, iteration: int
) -> None:
    del iteration

    def invoke(runtime: ModuleType, seal: _Seal) -> object:
        sealed = cast(Callable[[str, bytes], int], getattr(runtime, "sealed_memfd"))
        run = cast(Callable[[dict[str, list[object]], int, str, bytes, Callable[[], _Seal]], int], getattr(runtime, "run_snapshot_child"))
        descriptor = sealed("pidfd-fault-cli", b"pass\n")
        try:
            return run(
                {}, descriptor, hashlib.sha256(b"pass\n").hexdigest(), b"pass\n",
                lambda: seal,
            )
        finally:
            os.close(descriptor)

    _assert_fault_cleanup(monkeypatch, invoke, error, seal_expected=True)
