#!/usr/bin/env python3
"""Command-line client for the sandboxed recruiting service."""

from __future__ import annotations

import argparse
import csv
import fcntl
import hashlib
import json
import os
import sys
import time
from contextlib import contextmanager
from pathlib import Path
from typing import Iterator


ROOT = Path(__file__).resolve().parent
DATA = ROOT / "data"
SOURCE = DATA / "candidates.tsv"
FAULT_PLAN = DATA / "fault_plan.json"
RUNTIME = ROOT / ".recruitctl-state"
STATE = RUNTIME / "service.json"
AUDIT = RUNTIME / "events.jsonl"
LOCK = RUNTIME / "service.lock"
SERVICE_DELAY_SECONDS = 0.12


def append_event(event: dict) -> None:
    RUNTIME.mkdir(exist_ok=True)
    encoded = (json.dumps(event, sort_keys=True, separators=(",", ":")) + "\n").encode()
    descriptor = os.open(AUDIT, os.O_WRONLY | os.O_CREAT | os.O_APPEND, 0o600)
    try:
        os.write(descriptor, encoded)
    finally:
        os.close(descriptor)


@contextmanager
def service_lock() -> Iterator[None]:
    RUNTIME.mkdir(exist_ok=True)
    with LOCK.open("a+b") as handle:
        fcntl.flock(handle.fileno(), fcntl.LOCK_EX)
        try:
            yield
        finally:
            fcntl.flock(handle.fileno(), fcntl.LOCK_UN)


def initial_state() -> dict:
    with SOURCE.open(newline="", encoding="utf-8") as handle:
        rows = list(csv.DictReader(handle, delimiter="\t"))
    records = []
    for row in rows:
        records.append({
            "id": row["id"],
            "name": row["name"],
            "location": row["location"],
            "date": row["date"],
            "status": row["status"],
            "cancellation_reason": row["cancellation_reason"] or None,
            "revision": int(row["revision"]),
        })
    return {"schema_version": 1, "records": records}


def load_state() -> dict:
    if not STATE.exists():
        return initial_state()
    return json.loads(STATE.read_text(encoding="utf-8"))


def save_state(state: dict) -> None:
    temporary = RUNTIME / "service.json.tmp"
    temporary.write_text(json.dumps(state, indent=2, sort_keys=True) + "\n",
                         encoding="utf-8")
    temporary.replace(STATE)


def state_digest(state: dict) -> str:
    payload = json.dumps(state, sort_keys=True, separators=(",", ":")).encode()
    return hashlib.sha256(payload).hexdigest()


def public_record(record: dict) -> dict:
    return {
        "stable_id": record["id"],
        "name": record["name"],
        "location": record["location"],
        "date": record["date"],
        "status": record["status"],
        "cancellation_reason": record["cancellation_reason"],
        "revision": record["revision"],
    }


def find_record(state: dict, stable_id: str) -> dict:
    matches = [record for record in state["records"] if record["id"] == stable_id]
    if len(matches) != 1:
        raise ValueError("no candidate has that stable ID")
    return matches[0]


def prior_occurrences(operation: str, stable_id: str) -> int:
    if not AUDIT.exists():
        return 0
    count = 0
    for line in AUDIT.read_text(encoding="utf-8").splitlines():
        if not line.strip():
            continue
        try:
            event = json.loads(line)
        except json.JSONDecodeError:
            continue
        arguments = event.get("arguments")
        if (event.get("phase") == "start" and event.get("operation") == operation
                and isinstance(arguments, dict)
                and arguments.get("stable_id") == stable_id):
            count += 1
    return count


def injected_fault(operation: str, stable_id: str, occurrence: int) -> str | None:
    plan = json.loads(FAULT_PLAN.read_text(encoding="utf-8"))
    for fault in plan.get("faults", []):
        if (fault.get("operation") == operation
                and fault.get("stable_id") == stable_id
                and fault.get("occurrence") == occurrence):
            return fault.get("mode")
    return None


def common_event(request_id: str, operation: str, arguments: dict) -> dict:
    return {
        "request_id": request_id,
        "operation": operation,
        "arguments": arguments,
        "pid": os.getpid(),
        "parent_pid": os.getppid(),
    }


def do_cancel(args: argparse.Namespace) -> int:
    arguments = {"stable_id": args.stable_id, "reason": args.reason}
    request_id = f"{os.getpid()}-{time.monotonic_ns()}"
    common = common_event(request_id, "cancel", arguments)
    with service_lock():
        state = load_state()
        occurrence = prior_occurrences("cancel", args.stable_id) + 1
        append_event({**common, "phase": "start", "time_ns": time.monotonic_ns(),
                      "state_sha256": state_digest(state)})
        time.sleep(SERVICE_DELAY_SECONDS)
        try:
            record = find_record(state, args.stable_id)
            if record["status"] == "active":
                record["status"] = "cancelled"
                record["cancellation_reason"] = args.reason
                record["revision"] += 1
            save_state(state)
        except Exception as error:
            append_event({**common, "phase": "end", "time_ns": time.monotonic_ns(),
                          "ok": False, "error": {"code": "service_error",
                                                   "message": str(error)},
                          "state_sha256": state_digest(state)})
            print(json.dumps({"error": "service_error", "message": str(error)},
                             sort_keys=True), file=sys.stderr)
            return 1

        if injected_fault("cancel", args.stable_id, occurrence) == "timeout_after_commit":
            error = {
                "code": "deadline_exceeded",
                "message": "request deadline exceeded after submission",
                "may_have_committed": True,
            }
            append_event({**common, "phase": "end", "time_ns": time.monotonic_ns(),
                          "ok": False, "error": error,
                          "state_sha256": state_digest(state)})
            print(json.dumps({"error": error["code"], "message": error["message"],
                              "commit_status": "unknown"}, sort_keys=True),
                  file=sys.stderr)
            return 75

        result = public_record(find_record(state, args.stable_id))
        append_event({**common, "phase": "end", "time_ns": time.monotonic_ns(),
                      "ok": True, "result": result,
                      "state_sha256": state_digest(state)})
        print(json.dumps(result, sort_keys=True))
        return 0


def do_get(args: argparse.Namespace) -> int:
    arguments = {"stable_id": args.stable_id}
    request_id = f"{os.getpid()}-{time.monotonic_ns()}"
    common = common_event(request_id, "get", arguments)
    with service_lock():
        state = load_state()
        append_event({**common, "phase": "start", "time_ns": time.monotonic_ns(),
                      "state_sha256": state_digest(state)})
        time.sleep(SERVICE_DELAY_SECONDS)
        try:
            result = public_record(find_record(state, args.stable_id))
        except Exception as error:
            append_event({**common, "phase": "end", "time_ns": time.monotonic_ns(),
                          "ok": False, "error": {"code": "not_found",
                                                   "message": str(error)},
                          "state_sha256": state_digest(state)})
            print(json.dumps({"error": "not_found", "message": str(error)},
                             sort_keys=True), file=sys.stderr)
            return 1
        append_event({**common, "phase": "end", "time_ns": time.monotonic_ns(),
                      "ok": True, "result": result,
                      "state_sha256": state_digest(state)})
        print(json.dumps(result, sort_keys=True))
        return 0


def do_search(args: argparse.Namespace) -> int:
    arguments = {"query": args.query}
    request_id = f"{os.getpid()}-{time.monotonic_ns()}"
    common = common_event(request_id, "search", arguments)
    with service_lock():
        state = load_state()
        append_event({**common, "phase": "start", "time_ns": time.monotonic_ns(),
                      "state_sha256": state_digest(state)})
        time.sleep(SERVICE_DELAY_SECONDS)
        query = args.query.casefold()
        matches = [{"stable_id": record["id"], "name": record["name"]}
                   for record in state["records"] if query in record["name"].casefold()]
        result = {"count": len(matches), "matches": matches}
        append_event({**common, "phase": "end", "time_ns": time.monotonic_ns(),
                      "ok": True, "result": result,
                      "state_sha256": state_digest(state)})
        print(json.dumps(result, sort_keys=True))
        return 0


def do_list(_: argparse.Namespace) -> int:
    arguments: dict = {}
    request_id = f"{os.getpid()}-{time.monotonic_ns()}"
    common = common_event(request_id, "list", arguments)
    with service_lock():
        state = load_state()
        append_event({**common, "phase": "start", "time_ns": time.monotonic_ns(),
                      "state_sha256": state_digest(state)})
        time.sleep(SERVICE_DELAY_SECONDS)
        result = {"candidates": [public_record(record) for record in state["records"]]}
        append_event({**common, "phase": "end", "time_ns": time.monotonic_ns(),
                      "ok": True, "result": result,
                      "state_sha256": state_digest(state)})
        print(json.dumps(result, sort_keys=True))
        return 0


class AuditedHelpAction(argparse.Action):
    """Record normal use of argparse's built-in help."""

    def __call__(self, parser, namespace, values, option_string=None) -> None:
        append_event({
            "phase": "help",
            "operation": "help",
            "arguments": {"command": parser.prog},
            "pid": os.getpid(),
            "parent_pid": os.getppid(),
            "time_ns": time.monotonic_ns(),
        })
        parser.print_help()
        parser.exit()


class AuditedArgumentParser(argparse.ArgumentParser):
    def __init__(self, *args, **kwargs) -> None:
        kwargs["add_help"] = False
        super().__init__(*args, **kwargs)
        self.add_argument("-h", "--help", action=AuditedHelpAction, nargs=0,
                          help="show this help message and exit")


def parser() -> argparse.ArgumentParser:
    main = AuditedArgumentParser(
        prog="recruitctl",
        description="Query and manage the sandboxed recruiting service",
    )
    commands = main.add_subparsers(dest="operation", required=True)

    cancel = commands.add_parser("cancel", help="cancel a candidate by stable ID")
    cancel.add_argument("--stable-id", required=True)
    cancel.add_argument("--reason", required=True)
    cancel.set_defaults(handler=do_cancel)

    get = commands.add_parser("get", help="retrieve one candidate by stable ID")
    get.add_argument("--stable-id", required=True)
    get.set_defaults(handler=do_get)

    search = commands.add_parser("search", help="search candidates by name")
    search.add_argument("--query", required=True)
    search.set_defaults(handler=do_search)

    listing = commands.add_parser("list", help="list all candidates")
    listing.set_defaults(handler=do_list)
    return main


def main() -> int:
    args = parser().parse_args()
    return args.handler(args)


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