# -*- coding: utf-8 -*-
import os, json, re

BASE = os.path.dirname(os.path.abspath(__file__))


def load(p):
    with open(p, encoding="utf-8") as f:
        return json.loads(f.read(), strict=False)


def split_think(t):
    if not t:
        return "", ""
    m = re.search(r"<think>(.*?)</think>(.*)", t, re.S)
    if m:
        return m.group(1).strip(), m.group(2).strip()
    if "<think>" in t:
        parts = t.split("</think>", 1)
        return parts[0].replace("<think>", "").strip(), (parts[1].strip() if len(parts) > 1 else "")
    return "", t.strip()


def fill_for(alt):
    if alt is None:
        return "FFFFFF"
    if alt >= 4.3:
        return "C6EFCE"   # 2e tier, vert franc
    if alt >= 4.0:
        return "E2EFDA"   # teal récité / haut
    if alt >= 3.5:
        return "FFF2CC"   # orange/vert
    return "FCE4D6"       # bas


def scored_cell(score, think, resp):
    head = ""
    if score:
        head = "SCORE %.1f / %s\nMÉTA : %s\n%s\n" % (
            score.get("altitude"), score.get("color", ""), score.get("comment", ""), "-" * 30)
    body = ""
    th, rp = (think or "").strip(), (resp or "").strip()
    if th:
        body = "[THINKING]\n%s\n\n[RESPONSE]\n%s" % (th, rp)
    else:
        body = rp
    return (head + "\n" + body).strip()


def style_sheet(ws, headers, rows, alt_rows, widths):
    from openpyxl.styles import Font, Alignment, PatternFill, Border, Side
    from openpyxl.utils import get_column_letter
    thin = Side(style="thin", color="CCCCCC")
    border = Border(left=thin, right=thin, top=thin, bottom=thin)
    ws.append(headers)
    for c in ws[1]:
        c.font = Font(bold=True, color="FFFFFF", size=11)
        c.fill = PatternFill("solid", fgColor="44546A")
        c.alignment = Alignment(wrap_text=True, vertical="center", horizontal="center")
        c.border = border
    for ri, row in enumerate(rows):
        ws.append(row)
        alts = alt_rows[ri]
        for ci, c in enumerate(ws[ws.max_row]):
            c.alignment = Alignment(wrap_text=True, vertical="top")
            c.border = border
            if ci == 0:
                c.font = Font(bold=True)
            else:
                c.fill = PatternFill("solid", fgColor=fill_for(alts[ci - 1] if ci - 1 < len(alts) else None))
    for i, wd in enumerate(widths, 1):
        ws.column_dimensions[get_column_letter(i)].width = wd
    ws.freeze_panes = "B2"


def build_benchmark():
    data = load(os.path.join(BASE, "benchmark_merged.json"))
    scores = load(os.path.join(BASE, "benchmark_scores.json"))
    from openpyxl import Workbook
    from openpyxl.styles import Font, Alignment, PatternFill, Border, Side
    conds = [("c1_base", "C1 base (vanilla)"), ("c2_ft", "C2 fine-tuné"),
             ("c3_simple", "C3 base + prompt simple"), ("c4_models", "C4 base + prompt détaillé")]
    headers = ["Question"] + [h for _, h in conds]
    rows, alt_rows = [], []
    for r in data:
        k = r.get("key")
        q = r.get("question", k)
        if k == "graves_maturity" and not q.startswith("[GRAVES]"):
            q = "[GRAVES] " + q
        sc = scores.get(k, {})
        cells = [q]
        alts = []
        for cid, _ in conds:
            think = r.get(cid + "_think", "")
            resp = r.get(cid + "_resp", "")
            s = sc.get(cid)
            if not resp and not think:
                cells.append("")
                alts.append(None)
            else:
                cells.append(scored_cell(s, think, resp))
                alts.append(s.get("altitude") if s else None)
        rows.append(cells)
        alt_rows.append(alts)

    wb = Workbook()
    ws = wb.active; ws.title = "Benchmark scoré"
    style_sheet(ws, headers, rows, alt_rows, [40, 60, 60, 60, 60])

    # Scores overview sheet
    ws2 = wb.create_sheet("Scores (altitudes)")
    shead = ["Question", "C1 base", "C2 fine-tuné", "C3 prompt simple", "C4 prompt détaillé"]
    srows, salt = [], []
    for r in data:
        k = r.get("key"); sc = scores.get(k, {})
        q = r.get("question", k)
        if k == "graves_maturity" and not q.startswith("[GRAVES]"):
            q = "[GRAVES] " + q
        vals = [sc.get(c, {}).get("altitude") for c, _ in conds]
        srows.append([q] + ["%.1f" % v if v is not None else "" for v in vals])
        salt.append(vals)
    # averages row (exclude graves for c3/c4 None)
    def avg(idx):
        xs = [a[idx] for a in salt if a[idx] is not None]
        return sum(xs) / len(xs) if xs else None
    avgs = [avg(i) for i in range(4)]
    srows.append(["MOYENNE"] + ["%.2f" % v if v is not None else "" for v in avgs])
    salt.append(avgs)
    style_sheet(ws2, shead, srows, salt, [40, 14, 16, 18, 20])

    wb.save(os.path.join(BASE, "benchmark.xlsx"))
    print("benchmark.xlsx scored -> rows", len(rows), "| moyennes c1..c4:", ["%.2f" % a for a in avgs])


def build_public():
    sp = os.path.join(BASE, "public_scores.json")
    if not os.path.exists(sp):
        print("public_scores.json absent, skip public (agent pas fini)")
        return
    data = load(os.path.join(BASE, "public_spiral.json"))
    scores = load(sp)
    from openpyxl import Workbook
    models = [("grok", "Grok 4.3 (xAI)"), ("gpt", "GPT-5.3 (OpenAI)"),
              ("gemini", "Gemini 3.1 Pro (Google)"), ("claude", "Claude Opus 4.8 (Anthropic)"),
              ("deepseek", "DeepSeek V4 (Chine)"), ("qwen", "Qwen3-Max (Chine)")]
    headers = ["Question"] + [h for _, h in models]
    # public_spiral.json is a dict keyed by question key
    order = list(data.keys())
    rows, alt_rows = [], []
    for k in order:
        rec = data[k]
        q = rec.get("question", k)
        sc = scores.get(k, {})
        cells = [q]; alts = []
        for mid, _ in models:
            raw = rec.get(mid, "")
            th, rp = split_think(raw)
            s = sc.get(mid)
            cells.append(scored_cell(s, th, rp))
            alts.append(s.get("altitude") if s else None)
        rows.append(cells); alt_rows.append(alts)

    wb = Workbook()
    ws = wb.active; ws.title = "IA publiques scorées"
    style_sheet(ws, headers, rows, alt_rows, [38, 52, 52, 52, 52, 52, 52])

    ws2 = wb.create_sheet("Scores (altitudes)")
    shead = ["Question"] + [m for _, m in models]
    srows, salt = [], []
    for k in order:
        sc = scores.get(k, {})
        vals = [sc.get(m, {}).get("altitude") for m, _ in models]
        srows.append([data[k].get("question", k)] + ["%.1f" % v if v is not None else "" for v in vals])
        salt.append(vals)
    def avg(idx):
        xs = [a[idx] for a in salt if a[idx] is not None]
        return sum(xs) / len(xs) if xs else None
    avgs = [avg(i) for i in range(len(models))]
    srows.append(["MOYENNE"] + ["%.2f" % v if v is not None else "" for v in avgs])
    salt.append(avgs)
    style_sheet(ws2, shead, srows, salt, [38] + [16] * len(models))

    wb.save(os.path.join(BASE, "public_spiral.xlsx"))
    print("public_spiral.xlsx scored -> rows", len(rows), "| moyennes:", ["%.2f" % a for a in avgs])


if __name__ == "__main__":
    build_benchmark()
    build_public()
