"""Descriptor-walked no-follow IO and durable publication."""

from __future__ import annotations

import errno
import hashlib
import os
import stat
from collections.abc import Generator
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from typing import Final, Protocol

from .contract import FilePin, RootPin

_DIRECTORY_FLAGS: Final = os.O_RDONLY | os.O_DIRECTORY | os.O_CLOEXEC | os.O_NOFOLLOW
_FILE_FLAGS: Final = os.O_CLOEXEC | os.O_NOFOLLOW | os.O_NONBLOCK
CONTROL_MODE: Final = 0o700
PRIVATE_MODE: Final = 0o600
SEALED_MODE: Final = 0o400


class SecureIoError(RuntimeError):
    """A filesystem boundary no longer matches its sealed identity."""


class Writer(Protocol):
    """Injectable write syscall."""

    def __call__(self, descriptor: int, payload: bytes, /) -> int: ...


@dataclass(frozen=True, slots=True)
class OpenedRoot:
    """Pinned root descriptor."""

    descriptor: int
    pin: RootPin


def identity(info: os.stat_result) -> tuple[int, int, int, int, int, int]:
    """Project exact filesystem identity fields."""
    return (
        info.st_dev,
        info.st_ino,
        info.st_uid,
        info.st_gid,
        stat.S_IMODE(info.st_mode),
        info.st_nlink,
    )


def write_all(descriptor: int, payload: bytes, writer: Writer = os.write) -> None:
    """Write all bytes with short-write and EINTR handling."""
    offset = 0
    while offset < len(payload):
        try:
            written = writer(descriptor, payload[offset:])
        except InterruptedError:
            continue
        if written <= 0:
            raise SecureIoError("zero write")
        offset += written


def fsync_retry(descriptor: int) -> None:
    """Fsync despite EINTR."""
    while True:
        try:
            os.fsync(descriptor)
            return
        except OSError as error:
            if error.errno != errno.EINTR:
                raise


def open_directory_walk(path: Path) -> int:
    """Open an absolute directory component-by-component without symlinks."""
    if not path.is_absolute():
        raise SecureIoError("directory must be absolute")
    descriptor = os.open(path.anchor, _DIRECTORY_FLAGS)
    try:
        for component in path.parts[1:]:
            child = os.open(component, _DIRECTORY_FLAGS, dir_fd=descriptor)
            os.close(descriptor)
            descriptor = child
    except OSError:
        os.close(descriptor)
        raise
    return descriptor


def verify_directory(descriptor: int, *, mode: int, nlink: int | None = None) -> os.stat_result:
    """Require a real exact-owner directory."""
    info = os.fstat(descriptor)
    valid = (
        stat.S_ISDIR(info.st_mode)
        and info.st_uid == os.geteuid()
        and info.st_gid == os.getegid()
        and stat.S_IMODE(info.st_mode) == mode
        and (nlink is None or info.st_nlink == nlink)
    )
    if not valid:
        raise SecureIoError("directory identity")
    return info


@contextmanager
def open_root(pin: RootPin) -> Generator[OpenedRoot]:
    """Open and verify one exact sealed root pin."""
    descriptor = open_directory_walk(pin.path)
    try:
        info = verify_directory(descriptor, mode=pin.mode, nlink=pin.nlink)
        expected = (pin.device, pin.inode, pin.uid, pin.gid, pin.mode, pin.nlink)
        if identity(info) != expected:
            raise SecureIoError("root identity")
        yield OpenedRoot(descriptor, pin)
    finally:
        os.close(descriptor)


def create_control_root(path: Path) -> None:
    """Create one absent 0700 root beneath a descriptor-walked parent."""
    parent = open_directory_walk(path.parent)
    try:
        _ = verify_directory(parent, mode=CONTROL_MODE)
        os.mkdir(path.name, CONTROL_MODE, dir_fd=parent)
        fsync_retry(parent)
    finally:
        os.close(parent)
    with open_control_root(path):
        pass


@contextmanager
def open_control_root(path: Path) -> Generator[int]:
    """Open a private 0700 control root reached by descriptor walk."""
    descriptor = open_directory_walk(path)
    try:
        _ = verify_directory(descriptor, mode=CONTROL_MODE, nlink=2)
        yield descriptor
    finally:
        os.close(descriptor)


def _verify_open_file(directory: int, name: str, opened: os.stat_result, pin: FilePin) -> None:
    named = os.stat(name, dir_fd=directory, follow_symlinks=False)
    expected = (pin.device, pin.inode, pin.uid, pin.gid, pin.mode, pin.nlink)
    if (
        not stat.S_ISREG(opened.st_mode)
        or identity(opened) != expected
        or identity(named) != expected
        or opened.st_size != pin.size
    ):
        raise SecureIoError("file identity")


def read_pinned(root: OpenedRoot, pin: FilePin) -> bytes:
    """Read and rebind one direct child against exact metadata and digest."""
    if pin.path.parent != root.pin.path:
        raise SecureIoError("file outside root")
    descriptor = os.open(pin.path.name, os.O_RDONLY | _FILE_FLAGS, dir_fd=root.descriptor)
    try:
        _verify_open_file(root.descriptor, pin.path.name, os.fstat(descriptor), pin)
        chunks: list[bytes] = []
        while chunk := os.read(descriptor, 65_536):
            chunks.append(chunk)
        payload = b"".join(chunks)
        if hashlib.sha256(payload).hexdigest() != pin.sha256:
            raise SecureIoError("file bytes")
        return payload
    finally:
        os.close(descriptor)


def read_pinned_absolute(pin: FilePin) -> bytes:
    """Descriptor-walk every absolute parent before opening a pinned leaf."""
    parent = open_directory_walk(pin.path.parent)
    try:
        descriptor = os.open(pin.path.name, os.O_RDONLY | _FILE_FLAGS, dir_fd=parent)
        try:
            _verify_open_file(parent, pin.path.name, os.fstat(descriptor), pin)
            chunks: list[bytes] = []
            while chunk := os.read(descriptor, 65_536):
                chunks.append(chunk)
            payload = b"".join(chunks)
            if hashlib.sha256(payload).hexdigest() != pin.sha256:
                raise SecureIoError("file bytes")
            return payload
        finally:
            os.close(descriptor)
    finally:
        os.close(parent)


def complete_prefix(root: OpenedRoot, pin: FilePin, expected: bytes) -> None:
    """Read, authenticate, and append only a missing exact frame suffix."""
    if pin.path.parent != root.pin.path or pin.mode != PRIVATE_MODE:
        raise SecureIoError("prefix binding")
    descriptor = os.open(
        pin.path.name, os.O_RDWR | os.O_APPEND | _FILE_FLAGS, dir_fd=root.descriptor
    )
    try:
        opened = os.fstat(descriptor)
        _verify_open_file(root.descriptor, pin.path.name, opened, pin)
        current = os.read(descriptor, len(expected) + 1)
        if hashlib.sha256(current).hexdigest() != pin.sha256 or not expected.startswith(current):
            raise SecureIoError("prefix bytes")
        if not current or len(current) >= len(expected):
            raise SecureIoError("not a strict prefix")
        _ = os.lseek(descriptor, 0, os.SEEK_END)
        write_all(descriptor, expected[len(current) :])
        fsync_retry(descriptor)
        _ = os.lseek(descriptor, 0, os.SEEK_SET)
        if os.read(descriptor, len(expected) + 1) != expected:
            raise SecureIoError("prefix readback")
    finally:
        os.close(descriptor)
