"""Secure openat and atomic-write primitives for Topic-59 maintenance."""

from __future__ import annotations

import os
import stat
from dataclasses import dataclass
from pathlib import Path
from typing import Final, TypeVar

from pydantic import BaseModel

from .nutrition_weekly_maintenance_contract import (
    Topic59MaintenanceAuthorityV1,
    Topic59MaintenanceContractError,
    canonical_document,
)

_CREDENTIAL: Final = "nutricoach-topic59-maintenance-r71b.json"
_DIRECTORY_FLAGS: Final = os.O_RDONLY | os.O_DIRECTORY | os.O_CLOEXEC | os.O_NOFOLLOW
_FILE_FLAGS: Final = os.O_RDONLY | os.O_CLOEXEC | os.O_NOFOLLOW | os.O_NONBLOCK
_MAX_BYTES: Final = 8192
_DocumentT = TypeVar("_DocumentT", bound=BaseModel)


@dataclass(frozen=True, slots=True)
class OwnedMaintenanceDirectory:
    descriptor: int
    uid: int
    gid: int


def read_credential(root: Path) -> Topic59MaintenanceAuthorityV1:
    root_fd = os.open(root, _DIRECTORY_FLAGS)
    try:
        raw = read_regular_at(root_fd, _CREDENTIAL, 0o400, os.geteuid(), None)
        return parse_document(raw, Topic59MaintenanceAuthorityV1)
    finally:
        os.close(root_fd)


def open_owned_directory(profile_root: Path) -> OwnedMaintenanceDirectory:
    profile_fd = os.open(profile_root, _DIRECTORY_FLAGS)
    profile_info = os.fstat(profile_fd)
    if profile_info.st_uid != os.geteuid() or profile_info.st_gid != os.getegid():
        os.close(profile_fd)
        raise PermissionError
    try:
        data_fd = os.open("data", _DIRECTORY_FLAGS, dir_fd=profile_fd)
    finally:
        os.close(profile_fd)
    data_info = os.fstat(data_fd)
    if data_info.st_uid != profile_info.st_uid or data_info.st_gid != profile_info.st_gid:
        os.close(data_fd)
        raise PermissionError
    try:
        directory_fd = os.open("topic59-maintenance-r71b", _DIRECTORY_FLAGS, dir_fd=data_fd)
    finally:
        os.close(data_fd)
    info = os.fstat(directory_fd)
    if (
        not stat.S_ISDIR(info.st_mode) or stat.S_IMODE(info.st_mode) != 0o700
        or info.st_uid != profile_info.st_uid or info.st_gid != profile_info.st_gid
    ):
        os.close(directory_fd)
        raise PermissionError
    return OwnedMaintenanceDirectory(
        descriptor=directory_fd, uid=profile_info.st_uid, gid=profile_info.st_gid,
    )


def read_regular_at(
    directory_fd: int, name: str, mode: int, uid: int, gid: int | None,
) -> bytes:
    descriptor = open_file(directory_fd, name, mode, uid, gid)
    try:
        raw = os.read(descriptor, _MAX_BYTES + 1)
        if not raw or len(raw) > _MAX_BYTES or os.read(descriptor, 1):
            raise Topic59MaintenanceContractError()
        return raw
    finally:
        os.close(descriptor)


def open_file(
    directory_fd: int, name: str, mode: int, uid: int | None = None, gid: int | None = None,
) -> int:
    descriptor = os.open(name, _FILE_FLAGS, dir_fd=directory_fd)
    info = os.fstat(descriptor)
    expected_uid = os.geteuid() if uid is None else uid
    if (
        not stat.S_ISREG(info.st_mode) or stat.S_IMODE(info.st_mode) != mode
        or info.st_uid != expected_uid or (gid is not None and info.st_gid != gid)
        or info.st_nlink != 1
    ):
        os.close(descriptor)
        raise PermissionError
    return descriptor


def parse_document(raw: bytes, model: type[_DocumentT]) -> _DocumentT:
    parsed = model.model_validate_json(raw)
    if canonical_document(parsed) != raw:
        raise Topic59MaintenanceContractError()
    return parsed


def atomic_write_at(directory_fd: int, name: str, raw: bytes) -> None:
    temporary = f".{name}.tmp"
    descriptor = os.open(
        temporary, os.O_WRONLY | os.O_CREAT | os.O_EXCL | os.O_CLOEXEC | os.O_NOFOLLOW,
        0o600, dir_fd=directory_fd,
    )
    try:
        view = memoryview(raw)
        while view:
            view = view[os.write(descriptor, view):]
        os.fsync(descriptor)
    finally:
        os.close(descriptor)
    os.rename(temporary, name, src_dir_fd=directory_fd, dst_dir_fd=directory_fd)
    fsync_directory(directory_fd)
    if read_regular_at(directory_fd, name, 0o600, os.geteuid(), os.getegid()) != raw:
        raise OSError


def fsync_directory(directory_fd: int) -> None:
    os.fsync(directory_fd)
