#!/usr/bin/env python3
"""
extract_session.py — deterministic extractor for interview-eval.

Turns a Claude Code session transcript (.jsonl) into a normalized "evidence pack"
so the LLM judge scores from clean, quoted evidence instead of re-parsing raw JSONL.

Usage:
    python3 extract_session.py SESSION.jsonl [SESSION2.jsonl ...] [--json] [--out DIR]

Default output: a human/LLM-readable markdown evidence pack per session on stdout.
--json:   also emit the normalized structure as <out>/<sessionId>.session.json
--out DIR: write artifacts to DIR (default: alongside the transcript)

The extractor is purely mechanical. It NEVER scores. It surfaces signals
(prompts, tool usage, modes, errors, auto-detected red flags) for the judge.
"""
import json
import re
import sys
import os
from collections import Counter, OrderedDict
from datetime import datetime

# ---------- transcript record helpers ----------

IDE_TAG = re.compile(r"<ide_[a-z_]+>.*?</ide_[a-z_]+>", re.DOTALL)
CMD_NAME = re.compile(r"<command-name>(.*?)</command-name>", re.DOTALL)
CMD_ARGS = re.compile(r"<command-args>(.*?)</command-args>", re.DOTALL)

# Auto-detected red-flag patterns. These are SIGNALS for the judge, not verdicts.
SECRET_PATTERNS = [
    (re.compile(r"sk-[A-Za-z0-9]{20,}"), "openai-style key"),
    (re.compile(r"ghp_[A-Za-z0-9]{30,}"), "github PAT"),
    (re.compile(r"AKIA[0-9A-Z]{16}"), "aws access key id"),
    (re.compile(r"-----BEGIN (?:RSA |EC )?PRIVATE KEY-----"), "private key"),
    (re.compile(r"(?i)\b(password|passwd|secret|api[_-]?key|token)\b\s*[:=]\s*['\"][^'\"]{6,}"), "inline secret"),
]
DANGER_CMD = [
    (re.compile(r"\brm\s+-rf?\s+[/~]"), "rm -rf on root/home path"),
    (re.compile(r"\bgit\s+push\b.*--force\b|\bgit\s+push\b.*\s-f\b"), "force push"),
    (re.compile(r"\bgit\s+reset\s+--hard\b"), "hard reset"),
    (re.compile(r"\bDROP\s+(TABLE|DATABASE)\b", re.I), "drop table/database"),
    (re.compile(r"\bchmod\s+-R\s+777\b"), "chmod 777 -R"),
    (re.compile(r":\s*\(\)\s*\{\s*:\|:&\s*\}\s*;:"), "fork bomb"),
]
SLASH_CMD = re.compile(r"^\s*/([A-Za-z][\w:-]*)")


def load_records(path):
    recs = []
    with open(path, "r", errors="replace") as f:
        for ln, line in enumerate(f, 1):
            line = line.strip()
            if not line:
                continue
            try:
                recs.append(json.loads(line))
            except json.JSONDecodeError:
                recs.append({"type": "_PARSE_ERROR", "_line": ln})
    return recs


def msg_role(rec):
    return rec.get("message", {}).get("role")


def extract_text(content):
    """Return concatenated text blocks from a message content (str or list)."""
    if isinstance(content, str):
        return content
    if isinstance(content, list):
        parts = [b.get("text", "") for b in content
                 if isinstance(b, dict) and b.get("type") == "text"]
        return "\n".join(p for p in parts if p)
    return ""


def is_tool_result(content):
    return isinstance(content, list) and any(
        isinstance(b, dict) and b.get("type") == "tool_result" for b in content)


def tool_uses(content):
    out = []
    if isinstance(content, list):
        for b in content:
            if isinstance(b, dict) and b.get("type") == "tool_use":
                out.append((b.get("name", "?"), b.get("input", {}) or {}))
    return out


def tool_target(name, inp):
    """One-line human-readable summary of a tool call's target."""
    if not isinstance(inp, dict):
        return ""
    for k in ("file_path", "path", "pattern", "command", "url", "query",
              "prompt", "skill", "notebook_path", "description"):
        if k in inp and inp[k]:
            v = str(inp[k]).replace("\n", " ")
            return (v[:140] + "…") if len(v) > 140 else v
    return ""


def parse_ts(s):
    if not s:
        return None
    try:
        return datetime.fromisoformat(s.replace("Z", "+00:00"))
    except Exception:
        return None


def clean_prompt(text):
    """Strip IDE context tags; flag slash commands and local-command stdout."""
    stripped = IDE_TAG.sub("", text).strip()
    return stripped


# ---------- core normalization ----------

def normalize(path):
    recs = load_records(path)
    sid = None
    ai_title = None
    version = None
    cwd = None
    branch = None

    turns = []           # ordered user prompts with following assistant actions
    cur_actions = []     # assistant tool calls / text since last user prompt
    pending_user = None  # the open user turn we are attaching actions to

    tool_counts = Counter()
    perm_modes = Counter()
    perm_transitions = []
    mode_records = []
    interrupts = 0
    sys_errors = []
    parse_errors = 0
    first_ts = None
    last_ts = None
    slash_commands = []
    subagent_calls = []      # Task/Agent tool uses
    askuser_calls = 0
    plan_mode_used = False

    redflags = []            # auto-detected; (kind, where, detail)
    user_prompt_idx = 0

    def close_turn():
        nonlocal pending_user, cur_actions
        if pending_user is not None:
            pending_user["actions"] = cur_actions
            turns.append(pending_user)
        cur_actions = []
        pending_user = None

    for rec in recs:
        t = rec.get("type")
        if t == "_PARSE_ERROR":
            parse_errors += 1
            continue
        ts = parse_ts(rec.get("timestamp"))
        if ts:
            if first_ts is None:
                first_ts = ts
            last_ts = ts

        if t == "mode" or t == "permission-mode":
            m = rec.get("mode")
            if m:
                mode_records.append(m)
                if m in ("plan", "plan-mode"):
                    plan_mode_used = True

        if sid is None and rec.get("sessionId"):
            sid = rec["sessionId"]
        if rec.get("version"):
            version = rec["version"]
        if rec.get("cwd"):
            cwd = rec["cwd"]
        if rec.get("gitBranch"):
            branch = rec["gitBranch"]

        if t == "ai-title" and rec.get("aiTitle"):
            ai_title = rec["aiTitle"]

        if t == "system" and rec.get("level") == "error":
            sys_errors.append({
                "subtype": rec.get("subtype"),
                "error": (str(rec.get("error", ""))[:200]),
                "retryAttempt": rec.get("retryAttempt"),
                "maxRetries": rec.get("maxRetries"),
                "ts": rec.get("timestamp"),
            })

        if t == "user":
            pm = rec.get("permissionMode")
            if pm:
                perm_modes[pm] += 1
                if not perm_transitions or perm_transitions[-1] != pm:
                    perm_transitions.append(pm)
                if pm in ("plan", "plan-mode"):
                    plan_mode_used = True
            content = rec.get("message", {}).get("content")
            if msg_role(rec) != "user":
                continue
            raw = extract_text(content)
            # interrupt markers
            if "interrupted by user" in (raw or ""):
                interrupts += 1
            if is_tool_result(content) and not (raw and raw.strip()):
                continue  # pure tool result, not a human turn
            if not (raw and raw.strip()):
                continue
            text = clean_prompt(raw)
            if not text:
                continue
            # close prior turn, open a new one
            close_turn()
            user_prompt_idx += 1
            # detect slash command
            sm = SLASH_CMD.match(text)
            if sm:
                slash_commands.append(sm.group(1))
            cm = CMD_NAME.search(text)
            if cm:
                slash_commands.append(cm.group(1).strip().lstrip("/"))
            # red-flag scan on the human prompt text
            for pat, label in SECRET_PATTERNS:
                if pat.search(text):
                    redflags.append(("secret-in-prompt", f"U{user_prompt_idx}", label))
            had_ide = bool(IDE_TAG.search(raw))
            pending_user = {
                "idx": user_prompt_idx,
                "ts": rec.get("timestamp"),
                "text": text,
                "len": len(text),
                "had_ide_context": had_ide,
                "permissionMode": pm,
            }

        elif t == "assistant":
            content = rec.get("message", {}).get("content")
            for name, inp in tool_uses(content):
                tool_counts[name] += 1
                tgt = tool_target(name, inp)
                cur_actions.append({"tool": name, "target": tgt})
                if name in ("Task", "Agent"):
                    subagent_calls.append({"target": tgt,
                                           "subagent": inp.get("subagent_type")})
                if name == "AskUserQuestion":
                    askuser_calls += 1
                # danger scan on bash commands the agent ran
                if name == "Bash":
                    cmd = str(inp.get("command", ""))
                    for pat, label in DANGER_CMD:
                        if pat.search(cmd):
                            redflags.append(("danger-cmd", f"after U{user_prompt_idx}", f"{label}: {cmd[:120]}"))
            atext = extract_text(content)
            if atext and atext.strip():
                cur_actions.append({"assistant_text": atext.strip()[:400]})

    close_turn()

    duration_s = None
    if first_ts and last_ts:
        duration_s = int((last_ts - first_ts).total_seconds())

    prompt_lens = [t["len"] for t in turns]
    metrics = {
        "user_prompts": len(turns),
        "assistant_tool_calls": sum(tool_counts.values()),
        "tool_counts": dict(tool_counts.most_common()),
        "permission_modes": dict(perm_modes),
        "permission_mode_transitions": perm_transitions,
        "plan_mode_used": plan_mode_used,
        "subagent_calls": len(subagent_calls),
        "askuser_calls": askuser_calls,
        "slash_commands": slash_commands,
        "interrupts": interrupts,
        "system_errors": len(sys_errors),
        "parse_errors": parse_errors,
        "duration_seconds": duration_s,
        "duration_human": _human_dur(duration_s),
        "avg_prompt_len": round(sum(prompt_lens) / len(prompt_lens), 1) if prompt_lens else 0,
        "median_prompt_len": _median(prompt_lens),
    }

    return {
        "session_id": sid,
        "title": ai_title,
        "version": version,
        "cwd": cwd,
        "git_branch": branch,
        "start": first_ts.isoformat() if first_ts else None,
        "end": last_ts.isoformat() if last_ts else None,
        "metrics": metrics,
        "turns": turns,
        "subagent_calls": subagent_calls,
        "system_errors": sys_errors,
        "auto_redflags": redflags,
        "source_path": os.path.abspath(path),
    }


def _human_dur(s):
    if s is None:
        return None
    h, rem = divmod(s, 3600)
    m, sec = divmod(rem, 60)
    if h:
        return f"{h}h{m}m"
    if m:
        return f"{m}m{sec}s"
    return f"{sec}s"


def _median(xs):
    if not xs:
        return 0
    s = sorted(xs)
    n = len(s)
    return s[n // 2] if n % 2 else (s[n // 2 - 1] + s[n // 2]) / 2


# ---------- rendering ----------

def render_md(norm):
    m = norm["metrics"]
    L = []
    L.append(f"# Evidence pack — session `{norm['session_id']}`")
    if norm["title"]:
        L.append(f"**Title:** {norm['title']}")
    L.append("")
    L.append("## Session metrics")
    L.append(f"- User prompts (human turns): **{m['user_prompts']}**")
    L.append(f"- Wall-clock: **{m['duration_human']}** ({m['duration_seconds']}s)  ·  cwd: `{norm['cwd']}`  ·  branch: `{norm['git_branch']}`  ·  cli: {norm['version']}")
    L.append(f"- Tool calls: **{m['assistant_tool_calls']}** — {m['tool_counts']}")
    L.append(f"- Permission modes: {m['permission_modes']}  ·  transitions: {' → '.join(m['permission_mode_transitions']) or '(none)'}")
    L.append(f"- Plan mode used: **{m['plan_mode_used']}**  ·  Subagents: **{m['subagent_calls']}**  ·  AskUserQuestion: **{m['askuser_calls']}**  ·  Slash cmds: {m['slash_commands'] or '(none)'}")
    L.append(f"- Interrupts (user stopped a tool): **{m['interrupts']}**  ·  System/API errors: **{m['system_errors']}**")
    L.append(f"- Prompt length: avg **{m['avg_prompt_len']}** chars, median **{m['median_prompt_len']}**")
    L.append("")
    if norm["auto_redflags"]:
        L.append("## ⚠ Auto-detected red-flag signals (verify before scoring)")
        for kind, where, detail in norm["auto_redflags"]:
            L.append(f"- [{kind}] {where}: {detail}")
        L.append("")
    if norm["system_errors"]:
        L.append("## System/API errors (context noise, not always the candidate's fault)")
        for e in norm["system_errors"][:8]:
            L.append(f"- {e.get('subtype')} attempt {e.get('retryAttempt')}/{e.get('maxRetries')}: {e.get('error')}")
        L.append("")
    L.append("## Turn-by-turn (human prompt → agent actions)")
    for t in norm["turns"]:
        ide = "  ·  [had IDE file context]" if t["had_ide_context"] else ""
        pm = f"  ·  mode={t['permissionMode']}" if t.get("permissionMode") else ""
        L.append("")
        L.append(f"### U{t['idx']}  ({(t['ts'] or '')[:19]}{pm}{ide})")
        L.append("> " + t["text"].replace("\n", "\n> "))
        acts = t.get("actions", [])
        tool_line = _summarize_actions(acts)
        if tool_line:
            L.append("")
            L.append(f"_agent did:_ {tool_line}")
    L.append("")
    return "\n".join(L)


def _summarize_actions(acts):
    bits = []
    tc = Counter()
    for a in acts:
        if "tool" in a:
            tc[a["tool"]] += 1
    for name, n in tc.most_common():
        bits.append(f"{name}×{n}")
    # include first few concrete targets for color
    targets = [a["target"] for a in acts if a.get("tool") in ("Bash", "Write", "Edit", "Read") and a.get("target")][:3]
    s = ", ".join(bits)
    if targets:
        s += "  | e.g. " + " ; ".join(targets)
    return s


def main():
    args = [a for a in sys.argv[1:] if not a.startswith("--")]
    want_json = "--json" in sys.argv
    out_dir = None
    if "--out" in sys.argv:
        i = sys.argv.index("--out")
        if i + 1 < len(sys.argv):
            out_dir = sys.argv[i + 1]
            args = [a for a in args if a != out_dir]
    if not args:
        print(__doc__)
        sys.exit(1)
    for path in args:
        if not os.path.exists(path):
            print(f"!! not found: {path}", file=sys.stderr)
            continue
        norm = normalize(path)
        md = render_md(norm)
        print(md)
        print("\n" + "=" * 80 + "\n")
        base = out_dir or os.path.dirname(os.path.abspath(path))
        if want_json or out_dir:
            os.makedirs(base, exist_ok=True)
            jp = os.path.join(base, f"{norm['session_id']}.session.json")
            with open(jp, "w") as f:
                json.dump(norm, f, indent=2)
            print(f"[wrote {jp}]", file=sys.stderr)


if __name__ == "__main__":
    main()
