"""Shared rollback boundary for sealed NutriCoach live controllers."""

from __future__ import annotations

from types import TracebackType
from collections.abc import Callable
from typing import Literal, Self, final

from scripts.nutricoach_v150_concrete_host import ConcreteLiveHost
from scripts.nutricoach_v150_phase_journal import PhaseJournal
from scripts.nutricoach_v150_sealed_authority import Snapshot, restore


@final
class _Capture:
    """Record one rollback failure while continuing later recovery."""

    def __init__(self, failures: list[str], stage: str) -> None:
        self.failures = failures
        self.stage = stage

    def __enter__(self) -> Self:
        return self

    def __exit__(
        self,
        error_type: type[BaseException] | None,
        error: BaseException | None,
        traceback: TracebackType | None,
    ) -> Literal[True]:
        del error_type, traceback
        if error is not None:
            self.failures.append(f"{self.stage}:{error}")
        return True


@final
class RollbackGuard:
    """Exact v7-derived stop-through-fence rollback boundary."""

    def __init__(
        self,
        host: ConcreteLiveHost,
        journal: PhaseJournal,
        uncommitted_error: Callable[[str], BaseException],
        predecessor_cron_fence: Callable[[], None] | None = None,
    ) -> None:
        self.host = host
        self.journal = journal
        self._uncommitted_error = uncommitted_error
        self._predecessor_cron_fence = predecessor_cron_fence
        self.snapshot: Snapshot | None = None
        self.committed = False

    def __enter__(self) -> Self:
        return self

    def bind(self, snapshot: Snapshot) -> None:
        self.snapshot = snapshot

    def commit(self) -> None:
        self.committed = True

    def rollback(self, original: BaseException) -> None:
        failures = self.host.rollback_failures
        if self.host.service.running:
            with _Capture(failures, "stop"):
                self.host.service.stop()
        if self.snapshot is not None:
            with _Capture(failures, "authority_restore"):
                self.host.restore_runtime_authority()
            with _Capture(failures, "snapshot_restore"):
                restore(self.snapshot)
            with _Capture(failures, "created_remove"):
                self.host.remove_created()
            restore_failed = any(
                failure.startswith(("authority_restore:", "snapshot_restore:"))
                for failure in failures
            )
            if not restore_failed:
                with _Capture(failures, "systemd_reload"):
                    self.host.service.reload()
        else:
            restore_failed = False
        fence_failed = False
        if not restore_failed and self._predecessor_cron_fence is not None:
            with _Capture(failures, "predecessor_cron_pause"):
                self._predecessor_cron_fence()
            fence_failed = any(
                failure.startswith("predecessor_cron_pause:")
                for failure in failures
            )
        if not restore_failed and not fence_failed:
            if not self.host.service.running:
                with _Capture(failures, "service_restore"):
                    self.host.service.start()
            if not self.host.service.running:
                with _Capture(failures, "service_restore_retry"):
                    self.host.service.start()
            if not self.host.service.running:
                failures.append("service_restore:inactive")
            with _Capture(failures, "protected_restore"):
                self.host.verify_preflight()
        protected_failed = any(
            failure.startswith("protected_restore:") for failure in failures
        )
        safe_paused = self.journal.phase() == "ROLLED_BACK_SAFE_CRON_PAUSED"
        if safe_paused and not restore_failed and not fence_failed and not protected_failed:
            with _Capture(failures, "safe_r70_rollback_verify"):
                _verify_safe_r70_rollback(self.host)
        safe_failed = any(
            failure.startswith("safe_r70_rollback_verify:") for failure in failures
        )
        if restore_failed or fence_failed or protected_failed or not self.host.service.running or safe_failed:
            phase = "RECOVERY_REQUIRED"
        elif safe_paused:
            phase = "ROLLED_BACK_SAFE_CRON_PAUSED"
        else:
            phase = "ROLLED_BACK"
        self.journal.advance(phase)
        if failures:
            original.add_note("rollback_failures=" + ",".join(failures))

    def __exit__(
        self,
        error_type: type[BaseException] | None,
        error: BaseException | None,
        traceback: TracebackType | None,
    ) -> Literal[False]:
        del error_type, traceback
        if error is not None:
            self.rollback(error)
        elif not self.committed:
            missing = self._uncommitted_error("uncommitted")
            self.rollback(missing)
            raise missing
        return False


def _verify_safe_r70_rollback(host: ConcreteLiveHost) -> None:
    """Fence the restarted predecessor from a due provider-capable cron tick."""
    from scripts.nutricoach_v150_r71b_maintenance_transaction import (
        verify_paused_r70_scheduler,
    )

    verify_paused_r70_scheduler(host.paths.profile / "cron/jobs.json")
    state = host.service_state()
    if (
        state.get("ActiveState") != "active"
        or state.get("SubState") != "running"
        or int(state.get("MainPID", "0")) <= 0
        or state.get("NRestarts") != "0"
    ):
        raise RuntimeError("predecessor_service_contract")
