"""Exact ownership for Task22 parent-side fork resources."""
from __future__ import annotations

import fcntl
import hashlib
import os
import stat
from pathlib import Path
from types import TracebackType
from typing import Literal, Protocol


class Closeable(Protocol):
    def close(self) -> None: ...


LockHandoff = tuple[Path, Path, int, dict[str, object]]


def _lock_identity(descriptor: int) -> dict[str, object]:
    info = os.fstat(descriptor)
    content = os.pread(descriptor, info.st_size + 1, 0)
    return {
        "inode": info.st_ino, "mode": stat.S_IMODE(info.st_mode),
        "uid": info.st_uid, "gid": info.st_gid, "size": info.st_size,
        "sha256": hashlib.sha256(content).hexdigest(),
    }


def validate_lock_residue(
    path: Path, descriptor: int, expected: dict[str, object]
) -> None:
    current, pinned = path.lstat(), os.fstat(descriptor)
    writable = fcntl.fcntl(descriptor, fcntl.F_GETFL) & os.O_ACCMODE == os.O_RDWR
    if (
        _lock_identity(descriptor) != expected
        or (current.st_dev, current.st_ino) != (pinned.st_dev, pinned.st_ino)
        or not stat.S_ISREG(pinned.st_mode) or pinned.st_nlink != 1 or not writable
    ):
        raise RuntimeError("trainer-removal transaction lock residue identity changed")


def open_exact_lock_residue(path: Path, expected: dict[str, object]) -> int:
    descriptor = os.open(
        path, os.O_RDWR | os.O_CLOEXEC | getattr(os, "O_NOFOLLOW", 0)
        | getattr(os, "O_NOATIME", 0)
    )
    try:
        validate_lock_residue(path, descriptor, expected)
        return descriptor
    except BaseException:
        os.close(descriptor)
        raise


def revalidate_lock_residue(
    path: Path, descriptor: int, expected: dict[str, object] | None = None
) -> None:
    if expected is not None:
        validate_lock_residue(path, descriptor, expected)
        return
    current, pinned = path.lstat(), os.fstat(descriptor)
    if (current.st_dev, current.st_ino) != (pinned.st_dev, pinned.st_ino):
        raise RuntimeError("trainer-removal transaction lock residue was rebound")


def _under_profile(target: str, root: Path) -> bool:
    value = target.removesuffix(" (deleted)")
    return value == str(root) or value.startswith(f"{root}/")


def assert_parent_profile_free(root: Path) -> None:
    tasks = tuple(Path("/proc/self/task").iterdir())
    if len(tasks) != 1:
        raise RuntimeError("launcher parent is multithreaded during profile handoff")
    if _under_profile(os.readlink("/proc/self/cwd"), root):
        raise RuntimeError("launcher parent cwd references the profile")
    references: list[str] = []
    for entry in Path("/proc/self/fd").iterdir():
        try:
            target = entry.readlink()
        except FileNotFoundError:
            continue
        if _under_profile(str(target), root):
            references.append(f"{entry.name}:{target}")
    if references:
        raise RuntimeError("launcher parent retained profile descriptors: " + ",".join(references))


def await_gate(descriptor: int, label: str) -> None:
    content = bytearray()
    while block := os.read(descriptor, 2 - len(content)):
        content.extend(block)
        if len(content) > 1:
            break
    if bytes(content) != b"1":
        raise RuntimeError(f"{label} gate was not released")
    os.close(descriptor)


class ChildOwnership:
    """Own descriptors, one direct child, and an optional runtime seal."""

    def __init__(self, seal: Closeable | None = None) -> None:
        self._fds: set[int] = set()
        self._gate: int | None = None
        self._child: int | None = None
        self._reaped: bool = False
        self._seal: Closeable | None = seal

    def __enter__(self) -> ChildOwnership:
        return self

    def own(self, descriptor: int) -> int:
        self._fds.add(descriptor)
        return descriptor

    def pipe(self) -> tuple[int, int]:
        read_fd, write_fd = os.pipe2(os.O_CLOEXEC)
        return self.own(read_fd), self.own(write_fd)

    def close(self, descriptor: int) -> None:
        if descriptor not in self._fds:
            raise RuntimeError("fork resource is not owned or was already closed")
        os.close(descriptor)
        self._fds.remove(descriptor)

    def forked(self, pid: int, gate_write: int) -> None:
        if pid <= 0 or self._child is not None or gate_write not in self._fds:
            raise RuntimeError("invalid direct-child ownership transition")
        self._child = pid
        self._gate = gate_write

    def reaped(self) -> None:
        self._reaped = True

    def _reap(self) -> None:
        if self._child is None or self._reaped:
            return
        while True:
            try:
                waited, _status = os.waitpid(self._child, 0)
                break
            except InterruptedError:
                continue
            except ChildProcessError:
                self._reaped = True
                return
        if waited != self._child:
            raise RuntimeError("cleanup wait reaped the wrong child")
        self._reaped = True

    def _cleanup(self) -> list[BaseException]:
        failures: list[BaseException] = []
        if self._gate is not None and self._gate in self._fds:
            try:
                self.close(self._gate)
            except BaseException as exc:
                failures.append(exc)
        if self._seal is not None:
            try:
                self._seal.close()
            except BaseException as exc:
                failures.append(exc)
        try:
            self._reap()
        except BaseException as exc:
            failures.append(exc)
        for descriptor in tuple(self._fds):
            try:
                self.close(descriptor)
            except BaseException as exc:
                failures.append(exc)
        return failures

    def __exit__(
        self,
        exception_type: type[BaseException] | None,
        exception: BaseException | None,
        traceback: TracebackType | None,
    ) -> Literal[False]:
        del exception_type, traceback
        failures = self._cleanup()
        if exception is not None:
            for failure in failures:
                exception.add_note(f"cleanup failure: {type(failure).__name__}: {failure}")
            return False
        if failures:
            raise failures[0]
        return False
