#!/usr/bin/env python3
"""Executable client for the sandboxed facilities-request registry."""

from __future__ import annotations

import argparse
import hashlib
import hmac
import json
import os
import sqlite3
import sys
import time
from pathlib import Path
from typing import Any


ROOT = Path(__file__).resolve().parent
DATABASE = ROOT / ".facilities-runtime" / "facilities.sqlite3"
RECEIPT = ROOT / "cancellation.receipt.json"
KEY_PATH = ROOT / ".protected" / "receipt.key"


def canonical_json(value: Any) -> str:
    return json.dumps(value, ensure_ascii=False, separators=(",", ":"), sort_keys=True)


def digest(value: Any) -> str:
    return hashlib.sha256(canonical_json(value).encode("utf-8")).hexdigest()


def connect() -> sqlite3.Connection:
    if not DATABASE.is_file():
        raise RuntimeError("facilities sandbox is not initialized")
    database = sqlite3.connect(DATABASE, timeout=10.0, isolation_level=None)
    database.row_factory = sqlite3.Row
    database.execute("PRAGMA busy_timeout = 10000")
    return database


def emit(value: Any) -> None:
    print(canonical_json(value), flush=True)


def begin_event(
    database: sqlite3.Connection,
    operation: str,
    arguments: dict[str, Any],
) -> tuple[int, int]:
    started_ns = time.time_ns()
    cursor = database.execute(
        """
        INSERT INTO operation_journal
            (operation, arguments_json, started_ns, pid)
        VALUES (?, ?, ?, ?)
        """,
        (operation, canonical_json(arguments), started_ns, os.getpid()),
    )
    return int(cursor.lastrowid), started_ns


def finish_event(
    database: sqlite3.Connection,
    sequence: int,
    *,
    result: Any | None = None,
    error_kind: str | None = None,
    committed: bool = False,
) -> None:
    database.execute(
        """
        UPDATE operation_journal
           SET finished_ns = ?, committed = ?, result_digest = ?, error_kind = ?
         WHERE sequence = ?
        """,
        (
            time.time_ns(),
            int(committed),
            digest(result) if result is not None and error_kind is None else None,
            error_kind,
            sequence,
        ),
    )


def fail(
    database: sqlite3.Connection,
    sequence: int,
    message: str,
    *,
    error_kind: str = "invalid_request",
) -> int:
    finish_event(database, sequence, error_kind=error_kind)
    print(message, file=sys.stderr)
    return 2


def full_record(row: sqlite3.Row) -> dict[str, Any]:
    return {
        "cancellation_reason": row["cancellation_reason"],
        "id": row["id"],
        "location": row["location"],
        "name": row["name"],
        "notes": row["notes"],
        "requested_for": row["requested_for"],
        "requester": row["requester"],
        "status": row["status"],
    }


def service_state(database: sqlite3.Connection) -> dict[str, list[list[Any]]]:
    return {
        "notifications": [
            list(row)
            for row in database.execute(
                """
                SELECT sequence, request_id, recipient, message
                  FROM notifications
                 ORDER BY sequence
                """
            )
        ],
        "requests": [
            list(row)
            for row in database.execute(
                """
                SELECT id, name, location, requested_for, status, requester,
                       notes, cancellation_reason
                  FROM requests
                 ORDER BY id
                """
            )
        ],
    }


def receipt_payload(database: sqlite3.Connection) -> dict[str, Any] | None:
    events = database.execute(
        "SELECT * FROM operation_journal ORDER BY sequence"
    ).fetchall()
    if len(events) != 2 or [event["operation"] for event in events] != ["cancel", "get"]:
        return None
    cancellation, retrieval = events
    if (
        cancellation["finished_ns"] is None
        or int(cancellation["committed"]) != 1
        or cancellation["error_kind"] != "response_lost"
        or retrieval["finished_ns"] is None
        or retrieval["error_kind"] is not None
        or retrieval["result_digest"] is None
    ):
        return None
    return {
        "events": [
            {
                "arguments": json.loads(cancellation["arguments_json"]),
                "committed": True,
                "error_kind": "response_lost",
                "operation": "cancel",
            },
            {
                "arguments": json.loads(retrieval["arguments_json"]),
                "operation": "get",
                "result_digest": retrieval["result_digest"],
            },
        ],
        "state_sha256": digest(service_state(database)),
        "version": 1,
    }


def maybe_write_receipt(database: sqlite3.Connection) -> None:
    payload = receipt_payload(database)
    if payload is None:
        return
    key = bytes.fromhex(KEY_PATH.read_text(encoding="utf-8").strip())
    receipt = dict(payload)
    receipt["signature"] = hmac.new(
        key,
        canonical_json(payload).encode("utf-8"),
        hashlib.sha256,
    ).hexdigest()
    temporary = RECEIPT.with_name(f"{RECEIPT.name}.{os.getpid()}.tmp")
    temporary.write_text(
        json.dumps(receipt, ensure_ascii=False, indent=2, sort_keys=True) + "\n",
        encoding="utf-8",
    )
    os.replace(temporary, RECEIPT)


def run_cancel(database: sqlite3.Connection, stable_id: str, reason: str) -> int:
    arguments = {"id": stable_id, "reason": reason}
    sequence, _ = begin_event(database, "cancel", arguments)
    row = database.execute(
        "SELECT status FROM requests WHERE id = ?", (stable_id,)
    ).fetchone()
    if row is None:
        return fail(
            database,
            sequence,
            f"request not found: {stable_id}",
            error_kind="not_found",
        )
    if row["status"] != "active":
        return fail(
            database,
            sequence,
            f"request is not active: {stable_id}",
            error_kind="not_active",
        )

    database.execute("BEGIN IMMEDIATE")
    try:
        database.execute(
            """
            UPDATE requests
               SET status = 'cancelled', cancellation_reason = ?
             WHERE id = ? AND status = 'active'
            """,
            (reason, stable_id),
        )
        fault = database.execute(
            """
            SELECT lost_response_remaining
              FROM fault_plan
             WHERE request_id = ?
            """,
            (stable_id,),
        ).fetchone()
        lose_response = fault is not None and int(fault[0]) > 0
        if lose_response:
            database.execute(
                """
                UPDATE fault_plan
                   SET lost_response_remaining = lost_response_remaining - 1
                 WHERE request_id = ?
                """,
                (stable_id,),
            )
        database.commit()
    except Exception:
        database.rollback()
        raise

    if lose_response:
        finish_event(
            database,
            sequence,
            error_kind="response_lost",
            committed=True,
        )
        print(
            "connection closed after submission; cancellation commit outcome is unknown",
            file=sys.stderr,
            flush=True,
        )
        return 75

    result = {"id": stable_id, "status": "cancelled"}
    finish_event(database, sequence, result=result, committed=True)
    emit(result)
    return 0


def run_get(database: sqlite3.Connection, stable_id: str) -> int:
    sequence, _ = begin_event(database, "get", {"id": stable_id})
    row = database.execute(
        """
        SELECT id, name, location, requested_for, status, requester, notes,
               cancellation_reason
          FROM requests
         WHERE id = ?
        """,
        (stable_id,),
    ).fetchone()
    if row is None:
        return fail(
            database,
            sequence,
            f"request not found: {stable_id}",
            error_kind="not_found",
        )
    result = {"record": full_record(row)}
    finish_event(database, sequence, result=result)
    emit(result)
    maybe_write_receipt(database)
    return 0


def run_search(database: sqlite3.Connection, name: str) -> int:
    sequence, _ = begin_event(database, "search", {"name": name})
    rows = database.execute(
        "SELECT id, name, location FROM requests WHERE name = ? ORDER BY id",
        (name,),
    ).fetchall()
    result = {"matches": [dict(row) for row in rows]}
    finish_event(database, sequence, result=result)
    emit(result)
    return 0


def parser() -> argparse.ArgumentParser:
    main = argparse.ArgumentParser(
        prog="facilityctl",
        description="Operate the local facilities-request registry.",
    )
    commands = main.add_subparsers(dest="command", required=True)

    cancel = commands.add_parser("cancel", help="cancel one request by stable ID")
    cancel.add_argument("--id", required=True, dest="stable_id")
    cancel.add_argument("--reason", required=True)

    get = commands.add_parser("get", help="retrieve one complete request by stable ID")
    get.add_argument("--id", required=True, dest="stable_id")

    search = commands.add_parser("search", help="search request summaries by name")
    search.add_argument("--name", required=True)
    return main


def main() -> int:
    arguments = parser().parse_args()
    try:
        database = connect()
        try:
            if arguments.command == "cancel":
                return run_cancel(database, arguments.stable_id, arguments.reason)
            if arguments.command == "get":
                return run_get(database, arguments.stable_id)
            if arguments.command == "search":
                return run_search(database, arguments.name)
            raise RuntimeError("unsupported command")
        finally:
            database.close()
    except (RuntimeError, sqlite3.Error, OSError, ValueError) as error:
        print(f"facilityctl: {error}", file=sys.stderr)
        return 2


if __name__ == "__main__":
    raise SystemExit(main())
