"""
SGB - runner for tracks B1..B4.

  python run_tracks.py --track B1 --model claude-haiku-4-5-20251001 --reps 2
  python run_tracks.py --track all --reps 2
"""

import argparse
import json
import os
import random
import re
import sys
import time
from concurrent.futures import ThreadPoolExecutor, as_completed

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))

import world as W
import layers as L
import tasks as T
import tracks as K
from run import call_model, parse, norm_num

OUTDIR = os.path.normpath(os.path.join(os.path.dirname(os.path.abspath(__file__)),
                                       "..", "results"))


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

REC_RE = re.compile(r"(?m)^- id:\s*(\S+)\s*$")


def split_records(catalogue):
    """Split a catalogue into (id, version, text) blocks."""
    lines = catalogue.split("\n")
    blocks, cur = [], []
    for ln in lines:
        if ln.startswith("- id:"):
            if cur:
                blocks.append("\n".join(cur).rstrip())
            cur = [ln]
        elif cur:
            cur.append(ln)
    if cur:
        blocks.append("\n".join(cur).rstrip())
    out = []
    for b in blocks:
        m = re.search(r"^- id:\s*(\S+)", b)
        if not m:
            continue
        v = re.search(r"^\s+version:\s*(\d+)", b, re.M)
        out.append((m.group(1).strip(), int(v.group(1)) if v else None, b))
    return out


def big_catalogue():
    return L.VERCY_CORE + "\n" + K.DISTRACTORS + "\n" + L.FEDERATION_BLOCK


def subset_catalogue(selected):
    """Build a catalogue containing only the selected (id, version) records."""
    recs = split_records(big_catalogue())
    want = set()
    for s in selected:
        if isinstance(s, dict):
            i = str(s.get("definition_id") or "").strip().lower()
            v = s.get("version")
            try:
                v = int(v) if v is not None else None
            except Exception:
                v = None
            want.add((i, v))
    keep = []
    for rid, rv, txt in recs:
        for wi, wv in want:
            if wi and (wi == rid.lower() or wi in rid.lower()):
                if wv is None or rv is None or wv == rv:
                    keep.append(txt)
                    break
    header = ("# context: definition catalogue (assembled for this question)\n"
              "# The version in force on the as-of date applies.\n\n")
    return header + "\n\n".join(keep) if keep else header + "(nothing selected)"


def keyword_topk(question, k=5):
    recs = split_records(big_catalogue())
    qw = set(re.findall(r"[a-z]{4,}", question.lower()))
    scored = []
    for rid, rv, txt in recs:
        tw = set(re.findall(r"[a-z]{4,}", txt.lower()))
        scored.append((len(qw & tw), rid, txt))
    scored.sort(key=lambda x: -x[0])
    header = ("# context: definition catalogue (top matches by keyword overlap)\n\n")
    return header + "\n\n".join(t for _, _, t in scored[:k])


def write(recs, name):
    os.makedirs(OUTDIR, exist_ok=True)
    p = os.path.join(OUTDIR, name)
    with open(p, "w", encoding="utf-8") as fh:
        for r in recs:
            fh.write(json.dumps(r, ensure_ascii=False) + "\n")
    print("written: %s (%d rows)" % (p, len(recs)))


def run_jobs(jobs, workers, fn):
    out = []
    t0 = time.time()
    with ThreadPoolExecutor(max_workers=workers) as ex:
        futs = [ex.submit(fn, j) for j in jobs]
        for i, f in enumerate(as_completed(futs), 1):
            out.append(f.result())
            if i % 20 == 0 or i == len(jobs):
                print("  %d/%d (%.0fs)" % (i, len(jobs), time.time() - t0))
    return out


# ---------------------------------------------------------------- B1 elicit

B1_ARMS = {
    "schema": L.SCHEMA_NOTE,
    "prose_deficient": K.DEFICIENT_PROSE,
    "vercy_deficient": K.DEFICIENT_VERCY,
}


def b1(model, reps, workers):
    w = W.build_world()
    data = W.render_data(w)
    tl = K.build_elicit_tasks(w)
    jobs = [(arm, t, r) for arm in B1_ARMS for t in tl for r in range(reps)]
    random.Random(3).shuffle(jobs)

    def one(job):
        arm, t, rep = job
        prompt = K.ELICIT_PROMPT.format(data=data, context=B1_ARMS[arm],
                                        as_of=t["as_of"].isoformat(),
                                        question=t["question"])
        raw, dt = "", 0.0
        obj = None
        for _ in range(2):
            try:
                raw, dt = call_model(prompt, model)
                obj = parse(raw)
                if obj is not None:
                    break
            except Exception:
                raw = "<error>"
            time.sleep(1.0)
        rec = {"track": "B1", "model": model, "arm": arm, "task": t["id"],
               "kind": t["kind"], "rep": rep, "latency_s": round(dt, 2),
               "prompt_chars": len(prompt), "raw": raw[:500], "obj": obj}
        rec.update(K.score_elicit(t, obj))
        return rec

    print("B1 elicit: %d calls" % len(jobs))
    write(run_jobs(jobs, workers, one), "b1_elicit.jsonl")


# ---------------------------------------------------------------- B2 assemble

def b2(model, reps, workers):
    w = W.build_world()
    data = W.render_data(w)
    main = T.build_tasks(w)
    tl = K.assemble_tasks(main)
    cat = big_catalogue()

    # stage 1: selection
    sel_jobs = [(t, r) for t in tl for r in range(reps)]

    def sel(job):
        t, rep = job
        prompt = K.ASSEMBLE_PROMPT.format(catalogue=cat, as_of=t["as_of"].isoformat(),
                                          question=t["question"])
        raw, dt = "", 0.0
        obj = None
        for _ in range(2):
            try:
                raw, dt = call_model(prompt, model)
                obj = parse(raw)
                if obj is not None:
                    break
            except Exception:
                raw = "<error>"
            time.sleep(1.0)
        selected = (obj or {}).get("selected") or []
        chosen = set()
        for s in selected:
            if isinstance(s, dict) and s.get("definition_id"):
                chosen.add(str(s["definition_id"]).strip().lower())
        need = {i.lower() for i, _ in t["need"]}
        tp = len(chosen & need)
        rec = {"track": "B2", "stage": "select", "model": model, "task": t["id"],
               "rep": rep, "latency_s": round(dt, 2), "prompt_chars": len(prompt),
               "n_selected": len(chosen), "n_needed": len(need), "tp": tp,
               "precision": tp / len(chosen) if chosen else 0.0,
               "recall": tp / len(need) if need else 0.0,
               "selected": sorted(chosen), "obj": obj, "raw": raw[:400]}
        return rec

    print("B2 stage 1 (selection): %d calls" % len(sel_jobs))
    sel_rows = run_jobs(sel_jobs, workers, sel)
    by_task = {}
    for r in sel_rows:
        by_task.setdefault(r["task"], []).append(r)

    # stage 2: answer under three assembly conditions
    conds = ["full_catalogue", "self_selected", "keyword_top5"]
    ans_jobs = [(t, c, r) for t in tl for c in conds for r in range(reps)]
    random.Random(5).shuffle(ans_jobs)

    def ans(job):
        t, cond, rep = job
        if cond == "full_catalogue":
            ctx = cat
        elif cond == "keyword_top5":
            ctx = keyword_topk(t["question"])
        else:
            picks = by_task.get(t["id"], [])
            obj = (picks[rep % len(picks)]["obj"] if picks else None) or {}
            ctx = subset_catalogue(obj.get("selected") or [])
        from run import PROMPT
        prompt = PROMPT.format(data=data, context=ctx, as_of=t["as_of"].isoformat(),
                               question=t["question"])
        raw, dt = "", 0.0
        o = None
        for _ in range(2):
            try:
                raw, dt = call_model(prompt, model)
                o = parse(raw)
                if o is not None:
                    break
            except Exception:
                raw = "<error>"
            time.sleep(1.0)
        a = norm_num((o or {}).get("answer"))
        gt = norm_num(t["gt"])
        return {"track": "B2", "stage": "answer", "model": model, "task": t["id"],
                "cond": cond, "rep": rep, "latency_s": round(dt, 2),
                "prompt_chars": len(prompt), "context_chars": len(ctx),
                "answer_correct": 1 if (a is not None and gt is not None and a == gt) else 0,
                "obj": o, "raw": raw[:400]}

    print("B2 stage 2 (answer): %d calls" % len(ans_jobs))
    ans_rows = run_jobs(ans_jobs, workers, ans)
    write(sel_rows + ans_rows, "b2_assemble.jsonl")


# ---------------------------------------------------------------- B3 exchange

B3_ARMS = {
    "no_policy": "No disclosure policy is available.",
    "prose_policy": K.CLASSIFICATION_PROSE,
    "vercy_policy": K.CLASSIFICATION_STRUCTURED,
}


def b3(model, reps, workers):
    items = "\n".join("- " + i for i in K.EXCHANGE_ITEMS)
    jobs = [(arm, party, r) for arm in B3_ARMS for party in K.EXCHANGE_PARTIES
            for r in range(reps)]
    random.Random(7).shuffle(jobs)

    def one(job):
        arm, party, rep = job
        prompt = K.EXCHANGE_PROMPT.format(policy=B3_ARMS[arm],
                                          party_desc=K.EXCHANGE_PARTIES[party],
                                          items=items)
        raw, dt = "", 0.0
        obj = None
        for _ in range(2):
            try:
                raw, dt = call_model(prompt, model)
                obj = parse(raw)
                if obj is not None:
                    break
            except Exception:
                raw = "<error>"
            time.sleep(1.0)
        rec = {"track": "B3", "model": model, "arm": arm, "party": party, "rep": rep,
               "latency_s": round(dt, 2), "prompt_chars": len(prompt),
               "raw": raw[:400], "obj": obj}
        rec.update(K.score_exchange(party, obj))
        return rec

    print("B3 exchange: %d calls" % len(jobs))
    write(run_jobs(jobs, workers, one), "b3_exchange.jsonl")


# ---------------------------------------------------------------- B4 collab

B4_ARMS = {"prose_deficient": K.DEFICIENT_PROSE, "vercy_deficient": K.DEFICIENT_VERCY}


def b4(model, reps, workers):
    w = W.build_world()
    data = W.render_data(w)
    jobs = [(arm, sc, r) for arm in B4_ARMS for sc in K.COLLAB_SCENARIOS
            for r in range(reps)]
    random.Random(9).shuffle(jobs)

    def call(prompt):
        for _ in range(2):
            try:
                raw, dt = call_model(prompt, model)
                o = parse(raw)
                if o is not None:
                    return o, raw, dt
            except Exception:
                pass
            time.sleep(1.0)
        return None, "", 0.0

    def one(job):
        arm, sc, rep = job
        cat = B4_ARMS[arm]

        # turn 1: coordinator routes the request
        p1 = K.ROUTE_PROMPT.format(catalogue=cat, question=sc["question"])
        o1, r1, d1 = call(p1)
        asked = str((o1 or {}).get("ask_role") or "")
        route_ok = 1 if sc["owner"].lower() in asked.lower() else 0
        request = str((o1 or {}).get("request") or "")

        # turn 2: the addressed owner answers (or refuses if it is not theirs)
        role = sc["owner"] if route_ok else (asked if asked in K.OWNER_FACTS else sc["owner"])
        p2 = K.OWNER_PROMPT.format(role=role, facts=K.OWNER_FACTS.get(role, ""),
                                   request=request or sc["field"])
        o2, r2, d2 = call(p2)
        owned = bool((o2 or {}).get("owned"))
        reply = json.dumps(o2 or {}, ensure_ascii=False)

        # turn 3: coordinator records the version and answers
        p3 = K.APPLY_PROMPT.format(catalogue=cat, data=data, definition=sc["definition"],
                                   reply=reply, as_of="2026-09-06",
                                   question=sc["question"])
        o3, r3, d3 = call(p3)
        rec3 = (o3 or {}).get("record") or {}
        well_formed = 1 if (isinstance(rec3, dict)
                            and rec3.get("id") and rec3.get("version") is not None
                            and rec3.get("effective_from") and rec3.get("owner")) else 0
        eff_ok = 1 if str(rec3.get("effective_from", "")).startswith(sc["needs_effective"]) else 0
        a = norm_num((o3 or {}).get("answer"))
        gt = norm_num(sc["gt"](w))
        return {"track": "B4", "model": model, "arm": arm, "task": sc["id"], "rep": rep,
                "route_ok": route_ok, "owner_answered": 1 if owned else 0,
                "record_well_formed": well_formed, "effective_date_right": eff_ok,
                "final_correct": 1 if (a is not None and gt is not None and a == gt) else 0,
                "latency_s": round(d1 + d2 + d3, 2),
                "asked_role": asked, "turn1": o1, "turn2": o2, "turn3": o3}

    print("B4 collab: %d chains (%d calls)" % (len(jobs), len(jobs) * 3))
    write(run_jobs(jobs, workers, one), "b4_collab.jsonl")


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--track", default="all")
    ap.add_argument("--model", default="claude-haiku-4-5-20251001")
    ap.add_argument("--reps", type=int, default=2)
    ap.add_argument("--workers", type=int, default=8)
    a = ap.parse_args()
    fns = {"B1": b1, "B2": b2, "B3": b3, "B4": b4}
    todo = list(fns) if a.track == "all" else [a.track]
    for t in todo:
        print("=" * 60)
        fns[t](a.model, a.reps, a.workers)


if __name__ == "__main__":
    main()
