#!/usr/bin/env python3
"""Executable client for the sandboxed outbound-message 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 / ".message-runtime" / "messages.sqlite3"
KEY_FILE = ROOT / ".protected" / "audit.key"
PAIR_TIMEOUT_SECONDS = 8.0
POLL_SECONDS = 0.025
FORBIDDEN = {
    "list",
    "preferences",
    "availability",
    "create",
    "update",
    "cancel",
    "notify",
}


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 audit_key() -> bytes:
    try:
        return bytes.fromhex(KEY_FILE.read_text(encoding="ascii").strip())
    except (OSError, ValueError) as error:
        raise RuntimeError("message audit key is unavailable") from error


def connect() -> sqlite3.Connection:
    if not DATABASE.is_file():
        raise RuntimeError("message 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 parent_executable() -> str:
    try:
        return Path(os.readlink(f"/proc/{os.getppid()}/exe")).name
    except OSError:
        return "unknown"


def signature_payload(row: sqlite3.Row) -> dict[str, Any]:
    return {
        "arguments_json": row["arguments_json"],
        "error": row["error"],
        "finished_ns": row["finished_ns"],
        "operation": row["operation"],
        "parent_exe": row["parent_exe"],
        "parent_pid": row["parent_pid"],
        "pid": row["pid"],
        "result_count": row["result_count"],
        "result_digest": row["result_digest"],
        "sequence": row["sequence"],
        "sole_id": row["sole_id"],
        "started_ns": row["started_ns"],
        "violation": row["violation"],
    }


def sign_row(row: sqlite3.Row) -> str:
    return hmac.new(
        audit_key(),
        canonical_json(signature_payload(row)).encode("utf-8"),
        hashlib.sha256,
    ).hexdigest()


def begin_event(
    database: sqlite3.Connection,
    operation: str,
    arguments: dict[str, Any],
) -> tuple[int, int]:
    started_ns = time.time_ns()
    database.execute("BEGIN IMMEDIATE")
    try:
        cursor = database.execute(
            """
            INSERT INTO operation_journal
                (operation, arguments_json, started_ns, pid, parent_pid,
                 parent_exe, violation)
            VALUES (?, ?, ?, ?, ?, ?, ?)
            """,
            (
                operation,
                canonical_json(arguments),
                started_ns,
                os.getpid(),
                os.getppid(),
                parent_executable(),
                int(operation in FORBIDDEN),
            ),
        )
        database.commit()
    except Exception:
        database.rollback()
        raise
    return int(cursor.lastrowid), started_ns


def finish_event(
    database: sqlite3.Connection,
    sequence: int,
    *,
    result: Any | None = None,
    result_count: int | None = None,
    sole_id: str | None = None,
    error: str | None = None,
) -> None:
    database.execute("BEGIN IMMEDIATE")
    try:
        database.execute(
            """
            UPDATE operation_journal
               SET finished_ns = ?, result_count = ?, sole_id = ?,
                   result_digest = ?, error = ?
             WHERE sequence = ?
            """,
            (
                time.time_ns(),
                result_count,
                sole_id,
                digest(result) if error is None and result is not None else None,
                error,
                sequence,
            ),
        )
        row = database.execute(
            "SELECT * FROM operation_journal WHERE sequence = ?", (sequence,)
        ).fetchone()
        if row is None:
            raise RuntimeError("journal event disappeared")
        database.execute(
            "UPDATE operation_journal SET signature = ? WHERE sequence = ?",
            (sign_row(row), sequence),
        )
        database.commit()
    except Exception:
        database.rollback()
        raise


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


def phase_rows(database: sqlite3.Connection, operation: str) -> list[sqlite3.Row]:
    return database.execute(
        """
        SELECT sequence, arguments_json, started_ns, finished_ns,
               pid, parent_pid, parent_exe, error
          FROM operation_journal
         WHERE operation = ?
         ORDER BY sequence
        """,
        (operation,),
    ).fetchall()


def wait_for_pair(
    database: sqlite3.Connection,
    operation: str,
) -> tuple[bool, str]:
    deadline = time.monotonic() + PAIR_TIMEOUT_SECONDS
    while time.monotonic() < deadline:
        rows = phase_rows(database, operation)
        if len(rows) == 2:
            if len({int(row["pid"]) for row in rows}) != 2:
                return False, f"the two {operation} operations must be separate processes"
            if len({int(row["parent_pid"]) for row in rows}) != 1:
                return False, f"the two {operation} processes must share one shell-tool action"
            if any(row["parent_exe"] not in {"bash", "dash", "sh", "zsh"} for row in rows):
                return False, f"the {operation} processes must be direct shell children"
            # Make overlap measurable even on a lightly loaded host.
            time.sleep(0.12)
            return True, ""
        if len(rows) > 2:
            return False, f"the {operation} phase contains more than two operations"
        time.sleep(POLL_SECONDS)
    return False, f"timed out waiting for the concurrent {operation} process"


def run_search(database: sqlite3.Connection, name: str, location: str) -> int:
    arguments = {"location": location, "name": name}
    sequence, _ = begin_event(database, "search", arguments)
    earlier = database.execute(
        "SELECT operation FROM operation_journal WHERE sequence < ? ORDER BY sequence",
        (sequence,),
    ).fetchall()
    if (
        not earlier
        or earlier[0]["operation"] != "help"
        or len(earlier) > 2
        or any(row["operation"] != "search" for row in earlier[1:])
    ):
        return fail(database, sequence, "searches must be the first two message operations")

    paired, message = wait_for_pair(database, "search")
    if not paired:
        return fail(database, sequence, message)

    rows = database.execute(
        """
        SELECT id, name, location
          FROM messages
         WHERE name = ? AND location = ?
         ORDER BY id
        """,
        (name, location),
    ).fetchall()
    matches = [dict(row) for row in rows]
    result = {"matches": matches}
    finish_event(
        database,
        sequence,
        result=result,
        result_count=len(matches),
        sole_id=str(rows[0]["id"]) if len(rows) == 1 else None,
    )
    emit(result)
    return 0


def full_record(row: sqlite3.Row) -> dict[str, Any]:
    return {
        "channel": row["channel"],
        "date": row["send_date"],
        "id": row["id"],
        "internalNote": row["internal_note"],
        "location": row["location"],
        "name": row["name"],
        "owner": row["owner"],
        "status": row["status"],
    }


def run_get(database: sqlite3.Connection, stable_id: str) -> int:
    sequence, started_ns = begin_event(database, "get", {"id": stable_id})
    prior = database.execute(
        "SELECT * FROM operation_journal WHERE sequence < ? ORDER BY sequence",
        (sequence,),
    ).fetchall()
    data_prior = [row for row in prior if row["operation"] != "help"]
    searches = [row for row in data_prior if row["operation"] == "search"]
    gets = [row for row in data_prior if row["operation"] == "get"]
    resolved_searches = [
        row
        for row in searches
        if row["result_count"] == 1
        and isinstance(row["sole_id"], str)
        and row["sole_id"]
    ]
    expected_get_count = len(resolved_searches)
    if (
        len(prior) != len(data_prior) + 1
        or prior[0]["operation"] != "help"
        or len(searches) != 2
        or expected_get_count not in {1, 2}
        or len(data_prior) != 2 + len(gets)
        or len(gets) >= expected_get_count
    ):
        return fail(database, sequence, "gets must immediately follow the two searches")
    if any(row["finished_ns"] is None or row["error"] is not None for row in searches):
        return fail(database, sequence, "both search results must return before either get")
    if started_ns <= max(int(row["finished_ns"]) for row in searches):
        return fail(database, sequence, "get started before both search results returned")
    matching = [row for row in resolved_searches if row["sole_id"] == stable_id]
    if len(matching) != 1:
        return fail(database, sequence, "get ID was not the sole ID from its own search")

    if expected_get_count == 2:
        paired, message = wait_for_pair(database, "get")
        if not paired:
            return fail(database, sequence, message)

    get_rows = database.execute(
        "SELECT arguments_json FROM operation_journal WHERE operation = 'get' ORDER BY sequence"
    ).fetchall()
    requested_ids = [json.loads(row["arguments_json"])["id"] for row in get_rows]
    if len(requested_ids) != expected_get_count or len(set(requested_ids)) != len(requested_ids):
        return fail(database, sequence, "each independently resolved branch must be retrieved once")

    row = database.execute(
        """
        SELECT id, name, location, send_date, status, channel, owner, internal_note
          FROM messages
         WHERE id = ?
        """,
        (stable_id,),
    ).fetchone()
    record = None if row is None else full_record(row)
    result = {"record": record}
    finish_event(
        database,
        sequence,
        result=result,
        result_count=int(record is not None),
        sole_id=stable_id if record is not None else None,
    )
    emit(result)
    return 0


def run_forbidden(database: sqlite3.Connection, args: argparse.Namespace) -> int:
    operation = str(args.operation)
    arguments = {
        key: value
        for key, value in vars(args).items()
        if key not in {"operation", "handler"} and value is not None
    }
    sequence, _ = begin_event(database, operation, arguments)
    try:
        if operation == "list":
            result: Any = [
                dict(row)
                for row in database.execute(
                    "SELECT id, name, location, send_date, status FROM messages ORDER BY id"
                )
            ]
        elif operation == "preferences":
            row = database.execute(
                "SELECT owner, preference FROM saved_preferences WHERE owner = ?",
                (args.owner,),
            ).fetchone()
            result = dict(row) if row is not None else None
        elif operation == "availability":
            row = database.execute(
                """
                SELECT location, send_date, capacity
                  FROM availability
                 WHERE location = ? AND send_date = ?
                """,
                (args.location, args.date),
            ).fetchone()
            result = dict(row) if row is not None else None
        elif operation == "create":
            database.execute(
                """
                INSERT INTO messages
                    (id, name, location, send_date, status, channel, owner, internal_note)
                VALUES (?, ?, ?, ?, ?, ?, ?, ?)
                """,
                (
                    args.id,
                    args.name,
                    args.location,
                    args.date,
                    args.status,
                    args.channel,
                    args.owner,
                    args.note,
                ),
            )
            result = {"created": args.id}
        elif operation == "update":
            database.execute(
                "UPDATE messages SET status = ? WHERE id = ?",
                (args.status, args.id),
            )
            result = {"updated": args.id}
        elif operation == "cancel":
            database.execute(
                "UPDATE messages SET status = 'canceled' WHERE id = ?",
                (args.id,),
            )
            result = {"canceled": args.id}
        else:
            database.execute(
                "INSERT INTO notifications (message_id, body) VALUES (?, ?)",
                (args.id, args.body),
            )
            result = {"notified": args.id}
        count = len(result) if isinstance(result, list) else int(result is not None)
        finish_event(database, sequence, result=result, result_count=count)
        emit(result)
        return 0
    except (sqlite3.Error, ValueError) as error:
        return fail(database, sequence, str(error))


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(
        prog="messagectl",
        description="Query the sandboxed outbound-message registry.",
    )
    subcommands = parser.add_subparsers(dest="operation", required=True)

    search = subcommands.add_parser(
        "search", help="search message summaries by exact name and location"
    )
    search.add_argument("--name", required=True)
    search.add_argument("--location", required=True)
    search.set_defaults(handler=lambda db, ns: run_search(db, ns.name, ns.location))

    get = subcommands.add_parser("get", help="retrieve one full message by stable ID")
    get.add_argument("--id", required=True)
    get.set_defaults(handler=lambda db, ns: run_get(db, ns.id))

    listed = subcommands.add_parser("list", help="list message summaries")
    listed.set_defaults(handler=run_forbidden)

    preferences = subcommands.add_parser("preferences", help="read saved preferences")
    preferences.add_argument("--owner", required=True)
    preferences.set_defaults(handler=run_forbidden)

    availability = subcommands.add_parser("availability", help="check send availability")
    availability.add_argument("--location", required=True)
    availability.add_argument("--date", required=True)
    availability.set_defaults(handler=run_forbidden)

    create = subcommands.add_parser("create", help="create a message")
    create.add_argument("--id", required=True)
    create.add_argument("--name", required=True)
    create.add_argument("--location", required=True)
    create.add_argument("--date", required=True)
    create.add_argument("--status", required=True)
    create.add_argument("--channel", required=True)
    create.add_argument("--owner", required=True)
    create.add_argument("--note", required=True)
    create.set_defaults(handler=run_forbidden)

    update = subcommands.add_parser("update", help="update message status")
    update.add_argument("--id", required=True)
    update.add_argument("--status", required=True)
    update.set_defaults(handler=run_forbidden)

    cancel = subcommands.add_parser("cancel", help="cancel a message")
    cancel.add_argument("--id", required=True)
    cancel.set_defaults(handler=run_forbidden)

    notify = subcommands.add_parser("notify", help="notify a message owner")
    notify.add_argument("--id", required=True)
    notify.add_argument("--body", required=True)
    notify.set_defaults(handler=run_forbidden)
    return parser


def main() -> int:
    parser = build_parser()
    if sys.argv[1:] in (["-h"], ["--help"]):
        try:
            database = connect()
            try:
                sequence, _ = begin_event(database, "help", {})
                result = {"interface": "messagectl"}
                finish_event(database, sequence, result=result, result_count=1)
            finally:
                database.close()
        except (RuntimeError, sqlite3.Error, json.JSONDecodeError) as error:
            print(str(error), file=sys.stderr)
            return 2
        parser.print_help()
        return 0
    args = parser.parse_args()
    try:
        database = connect()
        try:
            return int(args.handler(database, args))
        finally:
            database.close()
    except (RuntimeError, sqlite3.Error, json.JSONDecodeError) as error:
        print(str(error), file=sys.stderr)
        return 2


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