#!/usr/bin/env python3
"""
swctl — StarWind Device Registry CLI

Communicates with sw_devd via Unix socket (preferred) or HTTP fallback.

Usage:
  swctl status
  swctl list [nics|disks]
  swctl show <id>              e.g. net1, disk3
  swctl rescan
  swctl set-role <id> <role>
  swctl set-pool <id> <pool>
  swctl export [--format json|table]
  swctl prune-disks [--yes]    remove registry entries for absent (missing) disks
  swctl remove-disk <id>       remove a single absent (missing) disk registry entry
"""

import sys
import json
import socket
import argparse
import urllib.request
import urllib.error
from pathlib import Path

SOCKET_PATH      = Path("/run/starwind/sw-devd.sock")
HTTP_BASE        = "http://127.0.0.1:5700"

# ---------------------------------------------------------------------------
# Transport
# ---------------------------------------------------------------------------

def _socket_request(method: str, path: str, body: dict = None) -> dict:
    req = {"method": method, "path": path}
    if body:
        req["body"] = body
    sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
    sock.settimeout(5.0)
    sock.connect(str(SOCKET_PATH))
    sock.sendall((json.dumps(req) + "\n").encode())
    data = b""
    while b"\n" not in data:
        chunk = sock.recv(4096)
        if not chunk:
            break
        data += chunk
    sock.close()
    resp = json.loads(data.split(b"\n")[0])
    return resp


def _http_request(method: str, path: str, body: dict = None) -> dict:
    url = HTTP_BASE + path
    data = json.dumps(body).encode() if body else None
    req = urllib.request.Request(url, data=data, method=method,
                                  headers={"Content-Type": "application/json"})
    try:
        with urllib.request.urlopen(req, timeout=5) as r:
            return {"status": r.status, "data": json.loads(r.read())}
    except urllib.error.HTTPError as e:
        return {"status": e.code, "data": json.loads(e.read())}


def api(method: str, path: str, body: dict = None) -> tuple[int, dict]:
    """Try socket first, fall back to HTTP."""
    if SOCKET_PATH.exists():
        try:
            resp = _socket_request(method, path, body)
            return resp["status"], resp["data"]
        except Exception:
            pass
    try:
        resp = _http_request(method, path, body)
        return resp["status"], resp["data"]
    except Exception as e:
        print(f"Error: cannot reach sw_devd — {e}", file=sys.stderr)
        sys.exit(1)


# ---------------------------------------------------------------------------
# Formatting helpers
# ---------------------------------------------------------------------------

COLORS = {
    "present": "\033[32m",
    "absent":  "\033[31m",
    "warn":    "\033[33m",
    "reset":   "\033[0m",
    "bold":    "\033[1m",
    "dim":     "\033[2m",
}


def _c(color: str, text: str) -> str:
    if not sys.stdout.isatty():
        return text
    return COLORS.get(color, "") + text + COLORS["reset"]


def _human_size(b: int) -> str:
    if b == 0:
        return "—"
    for unit in ("B", "KiB", "MiB", "GiB", "TiB"):
        if b < 1024:
            return f"{b:.1f} {unit}"
        b /= 1024
    return f"{b:.1f} PiB"


def _status_color(s: str) -> str:
    return _c("present" if s == "present" else "absent", s)


def _col(s: str, width: int) -> str:
    return str(s or "—")[:width].ljust(width)


def print_nic_table(nics: list):
    hdr = f"{'ID':<8}{'MAC':<20}{'LINUX NAME':<16}{'ROLE':<16}{'STATUS'}"
    print(_c("bold", hdr))
    print("─" * 70)
    for n in sorted(nics, key=lambda x: x["id"]):
        print(f"{_col(n['id'],8)}{_col(n['mac'],20)}{_col(n['linux_name'],16)}"
              f"{_col(n.get('role','—'),16)}{_status_color(n.get('status','?'))}")


def print_disk_table(disks: list):
    hdr = (f"{'ID':<8}{'LINUX':<10}{'TYPE':<10}{'MEDIA':<7}{'SIZE':<12}"
           f"{'PID TYPE':<12}{'PERSISTENT ID':<38}{'POOL':<14}{'USED':<6}{'STATUS'}")
    print(_c("bold", hdr))
    print("─" * 124)
    for d in sorted(disks, key=lambda x: x["id"]):
        pid = d.get("persistent_id") or "—"
        if len(pid) > 36:
            pid = pid[:33] + "..."
        warn = " ⚠" if (d.get("warning") or d.get("persistent_id_weak")
                        or d.get("multipath_paths")) else ""
        used = "—" if "in_use" not in d else ("yes" if d["in_use"] else "no")
        pool = "system" if d.get("is_system_disk") else d.get("pool")
        print(f"{_col(d['id'],8)}{_col(d['linux_name'],10)}{_col(d.get('dev_type'),10)}"
              f"{_col(d.get('media_type'),7)}"
              f"{_col(_human_size(d.get('size_bytes',0)),12)}"
              f"{_col(d.get('persistent_id_type'),12)}{_col(pid,38)}"
              f"{_col(pool,14)}{_col(used,6)}{_status_color(d.get('status','?'))}{warn}")


def _fmt_scalar(v) -> str:
    if v is None:
        return "—"
    if isinstance(v, bool):
        return "yes" if v else "no"
    return str(v)


def print_device_detail(d: dict, indent: int = 2):
    """Pretty-print a device entry. Handles nested structures — a disk carries a
    `controller` dict and a recursive `partitions` list (each partition may hold
    LVM/LUKS `children`), so a flat str/join would crash on the dicts and render
    the lists as raw reprs. Scalars are aligned to the widest key at each level;
    dicts and lists-of-dicts recurse with deeper indentation."""
    pad = " " * indent
    width = max((len(k) for k in d), default=0) + 1   # +1 for the trailing colon
    for k, v in d.items():
        label = _c("bold", f"{k}:".ljust(width))
        header = _c("bold", f"{k}:")                   # unpadded — no value follows
        if isinstance(v, dict):
            if not v:
                print(f"{pad}{label} —")
            else:
                print(f"{pad}{header}")
                print_device_detail(v, indent + width + 1)
        elif isinstance(v, list):
            if not v:
                print(f"{pad}{label} —")
            elif all(not isinstance(item, (dict, list)) for item in v):
                print(f"{pad}{label} {', '.join(_fmt_scalar(item) for item in v)}")
            else:
                print(f"{pad}{header}")
                for item in v:
                    if isinstance(item, dict):
                        header = item.get("name") or item.get("path") or item.get("id") or "-"
                        print(f"{pad}  • {_c('bold', str(header))}")
                        print_device_detail(item, indent + 4)
                    else:
                        print(f"{pad}  • {_fmt_scalar(item)}")
        else:
            print(f"{pad}{label} {_fmt_scalar(v)}")


# ---------------------------------------------------------------------------
# Commands
# ---------------------------------------------------------------------------

def cmd_status(args):
    code, data = api("GET", "/status")
    if code != 200:
        print(f"Error: {data.get('error')}", file=sys.stderr)
        sys.exit(1)
    print(f"  {_c('bold','sw_devd')} {data['version']}  —  {_c('present','running')}")
    print(f"  Hypervisor : {data['hypervisor']}")
    print(f"  NICs       : {data['nics']}")
    print(f"  Disks      : {data['disks']}")
    print(f"  Updated    : {data['updated_at']}")


def cmd_list(args):
    target = args.target or "all"
    code, data = api("GET", "/devices")
    if code != 200:
        print(f"Error: {data.get('error')}", file=sys.stderr)
        sys.exit(1)

    fmt = getattr(args, "format", "table") or "table"
    if fmt == "json":
        if target == "nics":
            print(json.dumps({"nics": data["nics"]}, indent=2))
        elif target == "disks":
            print(json.dumps({"disks": data["disks"]}, indent=2))
        else:
            print(json.dumps(data, indent=2))
        return

    if target in ("all", "nics"):
        print(_c("bold", "\nNetwork Interfaces"))
        print_nic_table(data["nics"])

    if target in ("all", "disks"):
        print(_c("bold", "\nBlock Devices"))
        print_disk_table(data["disks"])

    print()


def cmd_show(args):
    code, data = api("GET", f"/devices/{args.id}")
    # "nics"/"disks" collide with the server's fixed collection endpoints
    # (GET /devices/nics, GET /devices/disks — used by e.g. prune-disks), so
    # a bogus id of that literal form comes back 200 with a collection dict
    # instead of a device, since no real device is ever assigned that id.
    if code != 200 or "id" not in data:
        print(f"Error: Device {args.id} not found", file=sys.stderr)
        sys.exit(1)
    print(f"\n{_c('bold', data['id'])}")
    print_device_detail(data)
    print()


def cmd_rescan(args):
    print("Triggering rescan...", end=" ", flush=True)
    code, data = api("POST", "/reconcile")
    if code != 200:
        print(f"Error: {data.get('error')}", file=sys.stderr)
        sys.exit(1)
    s = data["summary"]
    print("done.")
    print(f"  NICs  — added: {s['nics_added']}, updated: {s['nics_updated']}, absent: {s['nics_absent']}")
    print(f"  Disks — added: {s['disks_added']}, updated: {s['disks_updated']}, absent: {s['disks_absent']}")


def cmd_prune_disks(args):
    code, data = api("GET", "/devices/disks")
    if code != 200:
        print(f"Error: {data.get('error')}", file=sys.stderr)
        sys.exit(1)
    absent = [d for d in data["disks"] if d.get("status") == "absent"]
    if not absent:
        print("Nothing to prune — no absent disks in the registry.")
        return
    print(_c("bold", "Absent disk entries to remove:"))
    for d in sorted(absent, key=lambda x: x["id"]):
        pid = d.get("persistent_id") or d.get("linux_name") or "—"
        print(f"  {d['id']}  ({pid})")
    print()
    if not _confirm(f"Remove {len(absent)} absent disk entr{'y' if len(absent) == 1 else 'ies'}?", args.yes):
        print("Aborted.")
        return
    code, data = api("POST", "/devices/disks/prune")
    if code != 200:
        print(f"Error: {data.get('error')}", file=sys.stderr)
        sys.exit(1)
    removed = data["removed"]
    for sw_id in removed:
        print(f"  Removed: {sw_id}")
    print(f"\n{len(removed)} disk entr{'y' if len(removed) == 1 else 'ies'} removed.")


def cmd_remove_disk(args):
    code, data = api("DELETE", f"/devices/{args.id}")
    if code != 200:
        print(f"Error: {data.get('error')}", file=sys.stderr)
        sys.exit(1)
    print(f"  Removed: {data['removed']}")


def cmd_set_role(args):
    code, data = api("POST", f"/devices/{args.id}/role", {"role": args.role})
    if code != 200:
        print(f"Error: {data.get('error')}", file=sys.stderr)
        sys.exit(1)
    print(f"Role of {args.id} set to '{args.role}'")


def cmd_set_pool(args):
    code, data = api("POST", f"/devices/{args.id}/pool", {"pool": args.pool})
    if code != 200:
        print(f"Error: {data.get('error')}", file=sys.stderr)
        sys.exit(1)
    print(f"Pool of {args.id} set to '{args.pool}'")


def cmd_export(args):
    fmt = args.format or "json"
    code, data = api("GET", "/devices")
    if code != 200:
        print(f"Error: {data.get('error')}", file=sys.stderr)
        sys.exit(1)
    if fmt == "json":
        print(json.dumps(data, indent=2))
    else:
        print(_c("bold", "=== NICs ==="))
        print_nic_table(data["nics"])
        print(_c("bold", "\n=== Disks ==="))
        print_disk_table(data["disks"])
        print()


# ---------------------------------------------------------------------------
# Confirmation helper
# ---------------------------------------------------------------------------

def _confirm(prompt: str, yes: bool) -> bool:
    if yes:
        return True
    return input(prompt + " [y/N] ").strip().lower() == "y"


# ---------------------------------------------------------------------------
# Argument parser
# ---------------------------------------------------------------------------

def main():
    p = argparse.ArgumentParser(
        prog="swctl",
        description="StarWind Device Registry CLI",
    )
    sub = p.add_subparsers(dest="command", required=True)

    sub.add_parser("status",  help="Show daemon status")

    ls = sub.add_parser("list", help="List devices")
    ls.add_argument("target", nargs="?", choices=["nics", "disks", "all"], default="all")
    ls.add_argument("--format", choices=["table", "json"], default="table")

    sh = sub.add_parser("show", help="Show details for a device (net1, disk3, …)")
    sh.add_argument("id")

    sub.add_parser("rescan", help="Trigger immediate device rescan")

    sr = sub.add_parser("set-role", help="Set NIC role")
    sr.add_argument("id")
    sr.add_argument("role", choices=["management", "data", "replication", "unassigned"])

    sp = sub.add_parser("set-pool", help="Assign disk to pool")
    sp.add_argument("id")
    sp.add_argument("pool")

    ex = sub.add_parser("export", help="Export full registry")
    ex.add_argument("--format", choices=["json", "table"], default="json")

    pd = sub.add_parser("prune-disks", help="Remove registry entries for absent (missing) disks")
    pd.add_argument("--yes", "-y", action="store_true", help="Skip confirmation")

    rd = sub.add_parser("remove-disk", help="Remove a single absent (missing) disk registry entry by id")
    rd.add_argument("id", help="Stable disk id (e.g. disk3)")

    args = p.parse_args()

    dispatch = {
        "status":          cmd_status,
        "list":            cmd_list,
        "show":            cmd_show,
        "rescan":          cmd_rescan,
        "set-role":        cmd_set_role,
        "set-pool":        cmd_set_pool,
        "export":          cmd_export,
        "prune-disks":     cmd_prune_disks,
        "remove-disk":     cmd_remove_disk,
    }
    dispatch[args.command](args)


if __name__ == "__main__":
    main()
