#!/usr/bin/env python3
"""judge — swipe through open decisions across every agent session on this machine.

    judge            open the TUI in this terminal
    judge open       open the TUI in a new cmux split (right)
    judge list       print the open queue and exit
    judge gates      print the gate ledger and kill-rate trend
    judge clear      expire every open card

Keys:  →/l approve   ←/h reject   c comment then swipe   s skip
       u undo last   j/k scroll   ?  help      q quit
"""

import curses
import os
import sys
import textwrap
import threading
import time
from datetime import datetime

HERE = os.path.dirname(os.path.abspath(os.path.realpath(__file__)))
sys.path.insert(0, os.path.join(HERE, "..", "scripts"))
import judgegates as G  # noqa: E402
import judgelib as J  # noqa: E402

REFRESH = 1.0


# --------------------------------------------------------------------------- #
# helpers
# --------------------------------------------------------------------------- #

def age(ts):
    s = int(time.time() - ts)
    if s < 60:
        return f"{s}s"
    if s < 3600:
        return f"{s // 60}m"
    if s < 86400:
        return f"{s // 3600}h{(s % 3600) // 60:02d}"
    return f"{s // 86400}d"


def wrap(text, width):
    lines = []
    for para in (text or "").splitlines() or [""]:
        if not para.strip():
            lines.append("")
            continue
        lines.extend(textwrap.wrap(para, width, replace_whitespace=False,
                                   drop_whitespace=False, break_long_words=True) or [""])
    return lines


def heartbeat_loop(stop_evt):
    while not stop_evt.is_set():
        try:
            J.beat()
        except Exception:
            pass
        stop_evt.wait(1.0)


# --------------------------------------------------------------------------- #
# colours
# --------------------------------------------------------------------------- #

C = {}


def init_colours():
    curses.start_color()
    curses.use_default_colors()
    pairs = {
        "dim": 245, "blue": 75, "cyan": 80, "green": 114, "yellow": 179,
        "red": 203, "magenta": 176, "white": 255,
    }
    for i, (name, fg) in enumerate(pairs.items(), start=1):
        curses.init_pair(i, fg, -1)
        C[name] = curses.color_pair(i)
    curses.init_pair(20, 235, 114)
    C["ok_bg"] = curses.color_pair(20)
    curses.init_pair(21, 235, 203)
    C["no_bg"] = curses.color_pair(21)
    curses.init_pair(22, 235, 179)
    C["warn_bg"] = curses.color_pair(22)


# --------------------------------------------------------------------------- #
# TUI
# --------------------------------------------------------------------------- #

class Judge:
    def __init__(self, scr):
        self.scr = scr
        self.queue = []
        self.idx = 0
        self.scroll = 0
        self.last_load = 0
        self.flash = ("", 0)
        self.history = []  # decided card ids this run
        self.skipped = set()
        self.drafts = {}  # card id → comment typed but not yet swiped

    # ---- data ----
    def load(self, force=False):
        if not force and time.time() - self.last_load < REFRESH:
            return
        self.last_load = time.time()
        J.gc()
        cards = J.open_cards()
        # keep skipped cards at the back
        cards.sort(key=lambda c: c["id"] in self.skipped)
        cur_id = self.current()["id"] if self.current() else None
        self.queue = cards
        if cur_id:
            for i, c in enumerate(self.queue):
                if c["id"] == cur_id:
                    self.idx = i
                    break
            else:
                self.idx = 0
                self.scroll = 0
        self.idx = min(self.idx, max(0, len(self.queue) - 1))

    def current(self):
        return self.queue[self.idx] if self.queue else None

    def say(self, msg, ttl=2.5):
        self.flash = (msg, time.time() + ttl)

    # ---- drawing ----
    def put(self, y, x, s, attr=0):
        h, w = self.scr.getmaxyx()
        if 0 <= y < h and x < w:
            try:
                self.scr.addnstr(y, x, s, max(0, w - x - 1), attr)
            except curses.error:
                pass

    def draw(self, x_off=0):
        scr = self.scr
        scr.erase()
        h, w = scr.getmaxyx()
        perms = sum(1 for c in self.queue if c["kind"] in J.BLOCKING_KINDS)
        sessions = len({c["session_id"] for c in self.queue})

        # header
        self.put(0, 1, "judge", curses.A_BOLD | C["blue"])
        stat = f"  {len(self.queue)} open"
        if perms:
            stat += f"  ·  {perms} blocking"
        stat += f"  ·  {sessions} session{'s' if sessions != 1 else ''}"
        self.put(0, 7, stat, C["dim"])
        clock = datetime.now().strftime("%H:%M")
        self.put(0, w - len(clock) - 1, clock, C["dim"])

        card = self.current()
        if not card:
            msg = "nothing to judge — sessions will feed cards here as they work"
            self.put(h // 2, max(1, (w - len(msg)) // 2), msg, C["dim"])
            self.footer()
            return

        # card frame
        cw = min(w - 4, 100)
        cx = max(1, (w - cw) // 2) + x_off
        top, bottom = 2, h - 3
        inner = cw - 4

        conclusion = (m0 := card.get("meta") or {}).get("verify", {}).get("conclusion", "")
        if card["kind"] == "permission":
            kind_attr, kind_txt = C["warn_bg"], " PERMISSION "
        elif card["kind"] == "gate":
            kind_attr = C["no_bg"] if conclusion in ("not-a-gate", "rubber-stamp") else C["warn_bg"]
            kind_txt = f" GATE · {m0.get('stage', '?').upper()} "
        else:
            kind_attr, kind_txt = C["cyan"] | curses.A_REVERSE, " REVIEW "
        self.put(top, cx, "╭" + "─" * (cw - 2) + "╮", C["dim"])
        self.put(top, cx + 2, kind_txt, kind_attr)
        pos = f" {self.idx + 1}/{len(self.queue)} · {age(card['created'])} "
        self.put(top, cx + cw - len(pos) - 2, pos, C["dim"])
        for y in range(top + 1, bottom):
            self.put(y, cx, "│", C["dim"])
            self.put(y, cx + cw - 1, "│", C["dim"])
        self.put(bottom, cx, "╰" + "─" * (cw - 2) + "╯", C["dim"])

        m = card.get("meta", {})
        y = top + 1
        # where
        where = J.short_dir(card.get("cwd", ""))
        tag = []
        if m.get("branch"):
            tag.append(m["branch"][:40])
        if m.get("slug"):
            tag.append(m["slug"])
        if m.get("model"):
            tag.append(m["model"].replace("claude-", ""))
        self.put(y, cx + 2, where, curses.A_BOLD | C["white"])
        self.put(y, cx + 3 + len(where), "  ".join(tag), C["dim"])
        y += 1
        if card["kind"] == "gate":
            title_attr = C["red"] if conclusion in ("not-a-gate", "rubber-stamp") else (
                C["yellow"] if conclusion in ("weak-red", "partial", "timeout", "unrunnable") else C["green"])
        else:
            title_attr = C["yellow"] if card["kind"] == "permission" else C["green"]
        self.put(y, cx + 2, card["title"], curses.A_BOLD | title_attr)
        y += 1
        if m.get("prompt"):
            first = m["prompt"].splitlines()[0][:inner - 6]
            self.put(y, cx + 2, "you » " + first, C["dim"])
            y += 1
        self.put(y, cx + 1, "─" * (cw - 2), C["dim"])
        y += 1

        # body
        body = []
        if card["kind"] == "gate":
            body += self.gate_body(card, inner)
        elif card["kind"] == "review":
            if m.get("files"):
                body.append(("files", C["magenta"]))
                body += [("  " + J.short_dir(f), C["dim"]) for f in m["files"][:12]]
                if len(m["files"]) > 12:
                    body.append((f"  … +{len(m['files']) - 12} more", C["dim"]))
            if m.get("bash"):
                body.append(("shell", C["magenta"]))
                body += [("  $ " + b, C["dim"]) for b in m["bash"]]
            if m.get("tools"):
                tools = "  ".join(f"{k}×{v}" for k, v in sorted(m["tools"].items(), key=lambda kv: -kv[1]))
                body.append(("tools  " + tools, C["dim"]))
            if body:
                body.append(("", 0))
            body.append(("claude said", C["magenta"]))
            body += [(ln, 0) for ln in wrap(card["body"], inner)]
        else:
            body += [(ln, 0) for ln in wrap(card["body"], inner)]

        avail = bottom - y
        max_scroll = max(0, len(body) - avail)
        self.scroll = min(self.scroll, max_scroll)
        for i, (ln, attr) in enumerate(body[self.scroll:self.scroll + avail]):
            self.put(y + i, cx + 2, ln, attr)
        if max_scroll:
            frac = self.scroll / max_scroll
            self.put(bottom, cx + cw - 12, f" {int(frac * 100):3d}% j/k ", C["dim"])

        # comment preview
        if self.drafts.get(card["id"]):
            self.put(bottom - 1, cx + 2, ("✎ " + self.drafts[card["id"]])[:inner], C["yellow"])

        self.footer(card)

    def gate_body(self, card, inner):
        """The evidence, then the test. The evidence is why the human is here."""
        m = card.get("meta", {})
        v = m.get("verify") or {}
        rows = []

        if m.get("claims"):
            rows.append(("claim", C["magenta"]))
            rows += [("  " + c[:inner - 4], 0) for c in m["claims"][:6]]
            if len(m["claims"]) > 6:
                rows.append((f"  … +{len(m['claims']) - 6} more", C["dim"]))
            rows.append(("", 0))

        rows.append(("evidence", C["magenta"]))
        status = v.get("status")
        if status == "running":
            rows.append(("  replaying the test against the code as it stands…", C["yellow"]))
        elif status == "error":
            rows.append((f"  could not replay: {v.get('error', '')}"[:inner], C["red"]))

        for run in v.get("runs", []):
            label = {"assertion": "fails on an assertion", "error": "fails on an import/compile error",
                     "pass": "PASSES without the implementation", "timeout": "timed out",
                     "unrunnable": "no runnable test command"}.get(run["verdict"], run["verdict"])
            attr = C["green"] if run["verdict"] == "assertion" else (
                C["red"] if run["verdict"] == "pass" else C["yellow"])
            rows.append((f"  {J.short_dir(run['file'])}", C["dim"]))
            rows.append((f"    {label}", attr))
            if run.get("command"):
                rows.append((f"    $ {run['command']}"[:inner], C["dim"]))

        if v.get("stage") == "green":
            scored, killed = v.get("scored", 0), v.get("killed", 0)
            attr = C["green"] if scored and killed == scored else (C["red"] if killed == 0 else C["yellow"])
            rows.append((f"  mutations caught   {killed}/{scored}", attr))
            for mutant in v.get("mutants", []):
                if mutant.get("invalid"):
                    continue
                mark, mattr = ("killed", C["dim"]) if mutant["killed"] else ("SURVIVED", C["red"])
                rows.append((f"  {mark:>8}  {mutant['file']}:{mutant['line']}  {mutant['op']}"[:inner], mattr))
                if not mutant["killed"]:
                    rows.append((f"            still passed with: {mutant['after']}"[:inner], C["dim"]))
        if v.get("note"):
            rows.append(("  " + v["note"], C["dim"]))

        if m.get("smells"):
            rows.append(("", 0))
            rows.append(("smells", C["magenta"]))
            for smell in m["smells"]:
                rows.append((f"  {smell['code']}", C["red"]))
                rows += [("    " + ln, C["dim"]) for ln in wrap(smell["note"], inner - 6)]

        if m.get("next_edit"):
            rows.append(("", 0))
            rows.append((f"held before editing  {m['next_edit']}"[:inner], C["dim"]))

        rows.append(("", 0))
        rows.append(("the test", C["magenta"]))
        rows += [(ln, 0) for ln in wrap(card["body"], inner)]
        return rows

    def footer(self, card=None):
        h, w = self.scr.getmaxyx()
        if card and card["kind"] == "permission":
            left, right = "← deny", "allow →"
        elif card and card["kind"] == "gate":
            left, right = "← weak gate", "real gate →"
        else:
            left, right = "← reject", "approve →"
        keys = [(left, C["red"]), ("   c comment", C["dim"]), ("   s skip", C["dim"]), ("   u undo", C["dim"]),
                ("   ? help", C["dim"]), ("   q quit", C["dim"])]
        x = 1
        for txt, attr in keys:
            self.put(h - 1, x, txt, attr | curses.A_BOLD if attr == C["red"] else attr)
            x += len(txt)
        self.put(h - 1, w - len(right) - 1, right, C["green"] | curses.A_BOLD)
        msg, until = self.flash
        if msg and time.time() < until:
            self.put(h - 2, max(1, (w - len(msg)) // 2), msg, C["yellow"] | curses.A_BOLD)

    def swipe_anim(self, direction, label, attr):
        h, w = self.scr.getmaxyx()
        for step in range(1, 7):
            self.draw(x_off=direction * step * max(2, w // 24))
            self.put(h // 2, max(1, (w - len(label)) // 2), f"  {label}  ", attr | curses.A_BOLD)
            self.scr.refresh()
            time.sleep(0.03)

    # ---- input ----
    def read_line(self, prompt, initial=""):
        h, w = self.scr.getmaxyx()
        curses.curs_set(1)
        self.scr.timeout(-1)  # block while typing; the 1s refresh timeout would raise "no input"
        buf = list(initial)
        while True:
            self.put(h - 2, 0, " " * (w - 1))
            self.put(h - 1, 0, " " * (w - 1))
            self.put(h - 2, 1, prompt, C["yellow"] | curses.A_BOLD)
            shown = "".join(buf)[-(w - 3):]
            self.put(h - 1, 1, shown)
            self.scr.move(h - 1, min(w - 2, 1 + len(shown)))
            self.scr.refresh()
            try:
                ch = self.scr.get_wch()
            except curses.error:
                continue
            if ch in ("\n", "\r", curses.KEY_ENTER):
                break
            if ch == "\x1b":
                buf = None
                break
            if ch in (curses.KEY_BACKSPACE, "\x7f", "\b"):
                if buf:
                    buf.pop()
            elif ch == "\x15":  # ctrl-u
                buf = []
            elif isinstance(ch, str) and ch.isprintable():
                buf.append(ch)
        curses.curs_set(0)
        self.scr.timeout(int(REFRESH * 1000))
        return None if buf is None else "".join(buf).strip()

    def decide(self, verdict):
        card = self.current()
        if not card:
            return
        comment = self.drafts.pop(card["id"], "")
        labels = {"permission": ("ALLOW", "DENY"), "gate": ("REAL GATE", "WEAK GATE")}
        yes, no = labels.get(card["kind"], ("APPROVE", "REJECT"))
        if verdict == "approve":
            self.swipe_anim(+1, yes, C["ok_bg"])
        else:
            self.swipe_anim(-1, no, C["no_bg"])
        J.decide(card, verdict, comment)
        self.history.append(card["id"])
        if card["kind"] == "gate":
            # The waiting hook writes the ledger for gates it decides itself; a
            # human decision is recorded here, at the moment it is made.
            G.append_ledger(J.load_card(card["id"]) or card)
        note = "sent to session" if card["kind"] == "review" and (verdict == "reject" or comment) else ""
        self.say(f"{verdict}{' · ' + note if note else ''}", 1.5)
        self.scroll = 0
        self.load(force=True)

    def comment(self):
        card = self.current()
        if not card:
            return
        text = self.read_line("comment (enter to keep, esc to cancel) — then swipe ← / →", self.drafts.get(card["id"], ""))
        if text is None:
            return
        self.drafts[card["id"]] = text
        self.say("comment set — now swipe", 3)

    def undo(self):
        if not self.history:
            self.say("nothing to undo")
            return
        card = J.load_card(self.history[-1])
        if not card:
            self.history.pop()
            self.say("card gone")
            return
        ok, note = J.undo(card)
        if ok:
            self.history.pop()
            self.load(force=True)
            for i, c in enumerate(self.queue):
                if c["id"] == card["id"]:
                    self.idx = i
                    break
        self.say(note, 4)

    def skip(self):
        card = self.current()
        if not card:
            return
        self.skipped.add(card["id"])
        self.scroll = 0
        self.load(force=True)
        self.say("skipped to back", 1)

    def help(self):
        h, w = self.scr.getmaxyx()
        lines = [
            "judge — human judgment queue",
            "",
            "GATE cards are the point: a test the session just wrote, replayed",
            "against the code to show whether it can actually fail. RED means the",
            "implementation is still unwritten; GREEN means mutations were run.",
            "  → / l  real gate                   ← / h  weak gate → the session is",
            "         held and must strengthen the test before writing the code",
            "PERMISSION cards block a session until you swipe.",
            "  → / l  allow the tool call        ← / h  deny (your comment is the reason)",
            "REVIEW cards summarise a finished turn.",
            "  → / l  approve (comment optional)  ← / h  reject → comment lands in that",
            "         session on its next prompt/turn, which must address it",
            "",
            "  c  write a comment, then swipe      s  skip card to the back",
            "  u  undo last decision (review cards; pulls back or retracts feedback)",
            "  j/k  scroll body                    r  reload      q  quit",
            "",
            "`judge gates` prints the kill-rate trend from every decided gate.",
            "Closing judge hands pending prompts and gates back to their sessions.",
            "any key to close",
        ]
        self.scr.erase()
        for i, ln in enumerate(lines):
            self.put(h // 2 - len(lines) // 2 + i, max(1, (w - 78) // 2), ln, curses.A_BOLD if i == 0 else 0)
        self.scr.refresh()
        self.scr.timeout(-1)
        self.scr.getch()
        self.scr.timeout(int(REFRESH * 1000))

    def run(self):
        curses.curs_set(0)
        self.scr.keypad(True)
        self.scr.timeout(int(REFRESH * 1000))
        self.load(force=True)
        while True:
            self.load()
            self.draw()
            self.scr.refresh()
            try:
                ch = self.scr.getch()
            except KeyboardInterrupt:
                return
            if ch == -1:
                continue
            if ch in (ord("q"), 27):
                return
            if ch in (curses.KEY_RIGHT, ord("l"), ord("L")):
                self.decide("approve")
            elif ch in (curses.KEY_LEFT, ord("h"), ord("H")):
                self.decide("reject")
            elif ch == ord("c"):
                self.comment()
            elif ch == ord("s"):
                self.skip()
            elif ch == ord("u"):
                self.undo()
            elif ch in (ord("j"), curses.KEY_DOWN):
                self.scroll += 1
            elif ch in (ord("k"), curses.KEY_UP):
                self.scroll = max(0, self.scroll - 1)
            elif ch == ord(" ") or ch == curses.KEY_NPAGE:
                self.scroll += 10
            elif ch == curses.KEY_PPAGE:
                self.scroll = max(0, self.scroll - 10)
            elif ch == ord("r"):
                self.load(force=True)
            elif ch == ord("?"):
                self.help()
            elif ch == curses.KEY_RESIZE:
                pass


def tui():
    sys.stdout.write("\033]0;judge\007")  # pane / tab title
    sys.stdout.flush()
    stop_evt = threading.Event()
    J.beat()
    t = threading.Thread(target=heartbeat_loop, args=(stop_evt,), daemon=True)
    t.start()
    try:
        curses.wrapper(lambda scr: (init_colours(), Judge(scr).run()))
    finally:
        stop_evt.set()
        try:
            J.HEARTBEAT.unlink()
        except OSError:
            pass


def cmd_open():
    """Pop a cmux split to the right and start the TUI in it."""
    import re
    import shutil
    import subprocess
    cmux = shutil.which("cmux") or "/Applications/cmux.app/Contents/Resources/bin/cmux"
    if not os.path.exists(cmux):
        print("judge: cmux CLI not found — run `judge` in another terminal pane instead")
        return 1
    if J.tui_alive():
        print("judge: already running")
        return 0
    res = subprocess.run([cmux, "new-split", "right", "--focus", "true"], capture_output=True, text=True)
    m = re.search(r"surface:\d+", res.stdout + res.stderr)
    if res.returncode != 0 or not m:
        print(f"judge: cmux new-split failed: {(res.stdout + res.stderr).strip()}")
        return 1
    subprocess.run([cmux, "send", "--surface", m.group(0), f"clear; {os.path.realpath(__file__)}\n"], check=False)
    print(f"judge: opened in cmux {m.group(0)}")
    return 0


def cmd_list():
    cards = J.open_cards()
    if not cards:
        print("judge: queue empty")
        return
    for c in cards:
        print(f"{c['kind']:<10} {age(c['created']):>5}  {J.short_dir(c['cwd']):<28} {c['title']}")


def cmd_gates(limit=25):
    """The gate ledger: what was judged, what the machine found, how it went."""
    records = G.read_ledger()
    if not records:
        print("judge: no gates judged yet")
        return
    human = [r for r in records if not r.get("auto")]
    weak = [r for r in records if r["verdict"] == "reject"]
    scored = [r for r in records if r.get("scored")]
    killed = sum(r.get("killed") or 0 for r in scored)
    total = sum(r.get("scored") or 0 for r in scored)
    print(f"judge gates — {len(records)} judged ({len(human)} by hand, "
          f"{len(records) - len(human)} auto-denied), {len(weak)} called weak")
    if total:
        print(f"mutation kill rate across {len(scored)} scored gate(s): {killed}/{total}"
              f" ({100 * killed // total}%)")
    smells = {}
    for record in records:
        for code in record.get("smells", []):
            smells[code] = smells.get(code, 0) + 1
    if smells:
        print("smells: " + "  ".join(f"{k}×{v}" for k, v in sorted(smells.items(), key=lambda kv: -kv[1])))
    print()
    for record in records[-limit:]:
        when = datetime.fromtimestamp(record["ts"]).strftime("%m-%d %H:%M")
        mark = "weak" if record["verdict"] == "reject" else "real"
        score = f"{record['killed']}/{record['scored']}" if record.get("scored") else "—"
        claim = (record["claims"] or record["files"] or [""])[0]
        print(f"{when}  {record['stage']:<5} {mark:<4} {score:>5}  "
              f"{record.get('conclusion', ''):<12} {record['repo']:<16} {claim[:44]}")


def cmd_clear():
    n = 0
    for c in J.open_cards():
        J.expire(c)
        n += 1
    print(f"judge: expired {n} card(s)")


if __name__ == "__main__":
    arg = sys.argv[1] if len(sys.argv) > 1 else ""
    if arg in ("-h", "--help", "help"):
        print(__doc__)
    elif arg == "open":
        sys.exit(cmd_open())
    elif arg == "list":
        cmd_list()
    elif arg == "clear":
        cmd_clear()
    elif arg == "gates":
        cmd_gates()
    else:
        tui()
