"""One-use resumable transaction controller over concrete sealed ports."""

from __future__ import annotations

import hashlib
import json
from dataclasses import dataclass
from typing import Literal, Protocol, assert_never

from .contract import SOURCE_FILE_DIGEST, SOURCE_ROW_DIGEST, PackageContract, Phase, approval_phrase
from .faults import NO_FAULTS, CrashFaultError, FaultInjector
from .journal import OneUseLedger, PhaseJournal, PhaseRecord
from .observer import ObserverControl, ObserverHead, TimerSubscription, require_pass


class ControllerError(RuntimeError):
    """The authorized continuity transaction failed closed."""


class InputVerifier(Protocol):
    """Exact pre-mutation verifier contract."""

    def verify(self, contract: PackageContract) -> None: ...


class RowExecutor(Protocol):
    """Installed row append/replay contract."""

    def append(self, *, recovery: bool = False) -> None: ...


class ReceiptWriter(Protocol):
    """Durable intent and provenance contract."""

    def write_intent(self, baseline_digest: str) -> None: ...

    def write_provenance(self) -> None: ...


@dataclass(frozen=True, slots=True)
class ControllerPorts:
    """Every concrete transaction effect."""

    verifier: InputVerifier
    ledger: OneUseLedger
    journal: PhaseJournal
    executor: RowExecutor
    receipts: ReceiptWriter
    observer: ObserverControl
    faults: FaultInjector = NO_FAULTS


def execute(
    approval: str,
    contract: PackageContract,
    ports: ControllerPorts,
    *,
    authorized_digest: str | None = None,
) -> str:
    """Reserve and execute the sole approval-bearing launch."""
    digest = contract.package_digest if authorized_digest is None else authorized_digest
    if approval != approval_phrase(digest):
        raise ControllerError("approval phrase")
    ports.verifier.verify(contract)
    ports.ledger.reserve()
    _ = _advance(ports, Phase.RESERVED)
    return _run_boundary(contract, ports, recovery=False)


def recover(contract: PackageContract, ports: ControllerPorts) -> str:
    """Resume only the exact pending reservation without an approval phrase."""
    ports.journal.reload()
    if ports.journal.current is Phase.COMMITTING and ports.ledger.consumed_outcome() == "SUCCEEDED":
        return _advance(ports, Phase.COMMITTED).digest
    ports.ledger.require_pending()
    current = ports.journal.current
    if current is None or _phase_index(current) < _phase_index(Phase.APPEND_INTENT_DURABLE):
        ports.observer.restore_timer()
        ports.ledger.consume("FAILED")
        raise ControllerError("pre-intent interruption consumed without mutation")
    return _run_boundary(contract, ports, recovery=True)


def _run_boundary(
    contract: PackageContract,
    ports: ControllerPorts,
    *,
    recovery: bool,
) -> str:
    try:
        return _resume(contract, ports, recovery=recovery)
    except CrashFaultError:
        raise
    except (OSError, RuntimeError):
        if ports.ledger.pending:
            ports.observer.restore_timer()
            ports.ledger.consume("FAILED")
        raise


def _resume(contract: PackageContract, ports: ControllerPorts, *, recovery: bool) -> str:
    baseline = _head(ports.journal.records, "baseline")
    manual = _head(ports.journal.records, "manual")
    timer = _head(ports.journal.records, "timer")
    subscription: TimerSubscription | None = None
    try:
        if _before(ports, Phase.OBSERVER_TIMER_FENCED):
            ports.faults.hit("timer-fence:before-effect")
            baseline = ports.observer.fence_timer_and_capture_failed_head()
            if baseline.status != "FAIL" or baseline.namespace != "observer-r71":
                raise ControllerError("failed observer baseline")
            ports.faults.hit("timer-fence:after-effect")
            _ = _advance(ports, Phase.OBSERVER_TIMER_FENCED, baseline=baseline.digest)
        if baseline is None:
            raise ControllerError("missing observer baseline")
        if _before(ports, Phase.INPUTS_LOCKED_AND_VERIFIED):
            ports.verifier.verify(contract)
            _ = _advance(ports, Phase.INPUTS_LOCKED_AND_VERIFIED, baseline=baseline.digest)
        if _before(ports, Phase.APPEND_INTENT_DURABLE):
            ports.receipts.write_intent(baseline.digest)
            _ = _advance(ports, Phase.APPEND_INTENT_DURABLE, baseline=baseline.digest)
        if _before(ports, Phase.ROW_DURABLE):
            ports.faults.hit("row-append:before-effect")
            ports.executor.append(recovery=recovery)
            ports.faults.hit("row-append:after-effect")
            _ = _advance(ports, Phase.ROW_DURABLE, baseline=baseline.digest)
        if _before(ports, Phase.PROVENANCE_DURABLE):
            ports.receipts.write_provenance()
            _ = _advance(ports, Phase.PROVENANCE_DURABLE, baseline=baseline.digest)
        if _before(ports, Phase.MANUAL_PASS_DURABLE):
            manual = ports.observer.run_manual(baseline)
            require_pass(manual, baseline, natural=False)
            _ = _advance(
                ports, Phase.MANUAL_PASS_DURABLE, baseline=baseline.digest, manual=manual.digest
            )
        if manual is None:
            raise ControllerError("missing manual head")
        if _before(ports, Phase.TIMER_WATCH_ARMED):
            subscription = ports.observer.subscribe_timer(manual)
            _ = _advance(
                ports,
                Phase.TIMER_WATCH_ARMED,
                baseline=baseline.digest,
                manual=manual.digest,
            )
        if _before(ports, Phase.TIMER_PASS_DURABLE):
            if subscription is None:
                subscription = ports.observer.subscribe_timer(manual)
            ports.faults.hit("timer-restore:before-effect")
            ports.observer.restore_timer()
            ports.faults.hit("timer-restore:after-effect")
            timer = subscription.await_pass()
            require_pass(timer, manual, natural=True)
            _ = _advance(
                ports,
                Phase.TIMER_PASS_DURABLE,
                baseline=baseline.digest,
                manual=manual.digest,
                timer=timer.digest,
            )
        if timer is None:
            raise ControllerError("missing timer head")
        if _before(ports, Phase.COMMITTING):
            _ = _advance(
                ports,
                Phase.COMMITTING,
                baseline=baseline.digest,
                manual=manual.digest,
                timer=timer.digest,
            )
        if _before(ports, Phase.COMMITTED):
            ports.ledger.consume("SUCCEEDED")
            return _advance(
                ports,
                Phase.COMMITTED,
                baseline=baseline.digest,
                manual=manual.digest,
                timer=timer.digest,
            ).digest
        return ports.journal.records[-1].digest
    except ControllerError:
        if ports.ledger.pending:
            ports.observer.restore_timer()
            ports.ledger.consume("FAILED")
        raise


def _advance(ports: ControllerPorts, phase: Phase, **heads: str) -> PhaseRecord:
    ports.faults.hit(f"{phase.value}:before-journal")
    record = ports.journal.advance(
        phase,
        baseline=heads.get("baseline"),
        manual=heads.get("manual"),
        timer=heads.get("timer"),
    )
    ports.faults.hit(f"{phase.value}:after-journal")
    return record


def _before(ports: ControllerPorts, phase: Phase) -> bool:
    current = ports.journal.current
    return current is None or _phase_index(current) < _phase_index(phase)


def _phase_index(phase: Phase) -> int:
    return tuple(Phase).index(phase)


def _head(
    records: tuple[PhaseRecord, ...],
    kind: Literal["baseline", "manual", "timer"],
) -> ObserverHead | None:
    for record in reversed(records):
        match kind:
            case "baseline":
                digest = record.observer_baseline_head
            case "manual":
                digest = record.observer_manual_head
            case "timer":
                digest = record.observer_timer_head
            case unreachable:
                assert_never(unreachable)
        if digest is not None:
            status = "FAIL" if kind == "baseline" else "PASS"
            namespace = "observer-r71" if kind == "baseline" else "observer-r71-continuity-r1"
            invocation = "recovered" if kind == "timer" else None
            return ObserverHead(digest, status, "0" * 64, (), namespace, invocation)
    return None


def provenance_payload(contract: PackageContract, intent_digest: str) -> bytes:
    """Build the external authenticated cross-root provenance receipt."""
    payload = {
        "expected_frame_sha256": SOURCE_FILE_DIGEST,
        "intent_digest": intent_digest,
        "method": "canonical_reissue",
        "package_digest": contract.package_digest,
        "source_authority": str(contract.source_root.path),
        "source_row_digest": SOURCE_ROW_DIGEST,
        "target_authority": str(contract.target_root.path),
        "target_row_digest": SOURCE_ROW_DIGEST,
    }
    body = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode()
    receipt = {**payload, "receipt_digest": hashlib.sha256(body).hexdigest()}
    return json.dumps(receipt, sort_keys=True, separators=(",", ":")).encode() + b"\n"
