"""Narrow catchable-termination deferral; SIGKILL is not recoverable."""

from __future__ import annotations

import signal
from enum import StrEnum
from threading import current_thread, get_ident, main_thread
from types import FrameType, TracebackType
from collections.abc import Callable
from typing import Final, Self, TypedDict, cast

from checkin_cli.weekly_operations import WeeklyOperationsCorruption, WeeklyOperationsError, WeeklyOperationsInputError
from checkin_cli.weekly_operations_parent import WeeklyOperationsAuthorityBinding

_CATCHABLE_TERMINATION: Final = frozenset((signal.SIGTERM, signal.SIGINT))


class DeferredTerminationSection:
    """Defer catchable termination only until a created inode is registered."""

    __slots__: Final = ("_previous_mask",)

    def __init__(self) -> None:
        self._previous_mask: set[int | signal.Signals] | None = None

    def __enter__(self) -> Self:
        if current_thread() is main_thread():
            self._previous_mask = signal.pthread_sigmask(signal.SIG_BLOCK, _CATCHABLE_TERMINATION)
        return self

    def __exit__(self, _kind: type[BaseException] | None, error: BaseException | None, traceback: TracebackType | None) -> None:
        if self._previous_mask is None:
            return
        try:
            _ = signal.pthread_sigmask(signal.SIG_SETMASK, self._previous_mask)
        except (WeeklyOperationsCorruption, KeyboardInterrupt, SystemExit) as cancellation:
            if error is not None:
                raise error.with_traceback(traceback) from cancellation
            raise


class InitializationCommittedEvidence(TypedDict):
    schema_version: str
    exit_context: str
    cancellation_kind: str | None
    original_context_type: str | None
    binding: dict[str, str | int]


class InitializationCancellationKind(StrEnum):
    SIGTERM = "sigterm"
    SIGINT = "sigint"


class InitializationCommitContext(StrEnum):
    MISSING_ACKNOWLEDGEMENT = "missing_acknowledgement"
    BODY_EXCEPTION = "body_exception"
    CANCELLATION = "cancellation"


class WeeklyOperationsInitializationCommitted(WeeklyOperationsError):
    """A committed binding was not acknowledged before transaction exit."""

    __slots__: tuple[str, ...] = ("binding", "exit_context", "cancellation_kind", "original_context")

    def __init__(
        self,
        binding: WeeklyOperationsAuthorityBinding,
        exit_context: InitializationCommitContext,
        cancellation_kind: InitializationCancellationKind | None = None,
        original_context: BaseException | None = None,
    ) -> None:
        super().__init__(f"initialization committed during {exit_context.value}")
        self.binding: WeeklyOperationsAuthorityBinding = binding
        self.exit_context: InitializationCommitContext = exit_context
        self.cancellation_kind: InitializationCancellationKind | None = cancellation_kind
        self.original_context: BaseException | None = original_context

    def to_evidence(self) -> InitializationCommittedEvidence:
        binding = {field: getattr(self.binding, field) for field in self.binding.__dataclass_fields__}
        return {
            "schema_version": "nutricoach-weekly-operations-initialization-committed-v2",
            "exit_context": self.exit_context.value,
            "cancellation_kind": None if self.cancellation_kind is None else self.cancellation_kind.value,
            "original_context_type": None if self.original_context is None else type(self.original_context).__name__,
            "binding": binding,
        }


class InitializationCancellation(WeeklyOperationsCorruption):
    __slots__: tuple[str, ...] = ("cancellation_kind",)

    def __init__(self, cancellation_kind: InitializationCancellationKind) -> None:
        super().__init__(f"authority initialization cancelled by {cancellation_kind.value}")
        self.cancellation_kind: InitializationCancellationKind = cancellation_kind


def _cancel_initialization(signal_number: int, _frame: FrameType | None) -> None:
    kind = InitializationCancellationKind.SIGTERM if signal_number == signal.SIGTERM else InitializationCancellationKind.SIGINT
    raise InitializationCancellation(kind)


class InitializationSignalGuard:
    """Retain and exactly restore one creator thread's signal state."""

    __slots__: Final = ("_creator", "_active", "_previous_mask", "_previous_sigterm", "_previous_sigint", "_handlers_installed")

    def __init__(self) -> None:
        self._creator: int = get_ident()
        self._active: bool = False
        self._previous_mask: set[int | signal.Signals] | None = None
        self._previous_sigterm: signal.Handlers | int | Callable[[int, FrameType | None], None] | None = None
        self._previous_sigint: signal.Handlers | int | Callable[[int, FrameType | None], None] | None = None
        self._handlers_installed: bool = False

    def _require_creator(self) -> None:
        if get_ident() != self._creator:
            raise WeeklyOperationsInputError("initialization transaction thread mismatch")

    def block(self) -> None:
        self._require_creator()
        if self._active:
            raise WeeklyOperationsInputError("initialization signal guard already active")
        self._previous_mask = signal.pthread_sigmask(signal.SIG_BLOCK, _CATCHABLE_TERMINATION)
        self._active = True
        if current_thread() is main_thread():
            self._previous_sigterm = cast(signal.Handlers | int | Callable[[int, FrameType | None], None], signal.getsignal(signal.SIGTERM))
            self._previous_sigint = cast(signal.Handlers | int | Callable[[int, FrameType | None], None], signal.getsignal(signal.SIGINT))
            _ = signal.signal(signal.SIGTERM, _cancel_initialization)
            _ = signal.signal(signal.SIGINT, _cancel_initialization)
            self._handlers_installed = True

    def restore(self, primary: BaseException | None = None) -> InitializationCancellation | None:
        self._require_creator()
        if not self._active or self._previous_mask is None:
            raise WeeklyOperationsInputError("initialization signal guard is inactive")
        cancellation: InitializationCancellation | None = None
        try:
            try:
                _ = signal.pthread_sigmask(signal.SIG_SETMASK, self._previous_mask)
            except InitializationCancellation as error:
                cancellation = error
        finally:
            if self._handlers_installed:
                _ = signal.signal(signal.SIGINT, self._previous_sigint)
                _ = signal.signal(signal.SIGTERM, self._previous_sigterm)
            self._active = False
        if primary is not None:
            if cancellation is not None:
                raise primary from cancellation
            raise primary
        return cancellation

