#!/usr/bin/env python3
"""
Build the page-shaped metrics.json consumed by coding-vs-orchestration.html.

    python3 model-behavior-metrics.py mb-*.jsonl mc-*.jsonl > metrics.json

Inputs (merge as many machines as you like):
  mb-<host>.jsonl   from model-behavior.py       (turn + session records)
  mc-<host>.jsonl   from model-behavior-code.py  (coding records)

Output shape:
  { "<family>": [ {key,label,unit,better,note,v:{model:value}} ... ],
    "index": { "code":{..}, "persist":{..}, ... } }

Edit MODELS below to compare a different set.
"""
import json
import re
import statistics as st
import sys
from collections import defaultdict

MODELS = ["claude-fable-5", "claude-opus-4-8", "claude-opus-5"]
SHORT = {m: m.replace("claude-", "") for m in MODELS}


def load(paths):
    turns, sess, code, attn, cost = [], [], [], [], []
    for p in paths:
        for line in open(p):
            r = json.loads(line)
            k = r.get("k")
            if k == "turn" and r.get("model"):
                turns.append(r)
            elif k == "session" and r.get("model"):
                sess.append(r)
            elif k == "code" and r.get("model"):
                code.append(r)
            elif k == "attn" and r.get("model"):
                attn.append(r)
            elif k == "cost" and r.get("orch_model"):
                cost.append(r)
    return turns, sess, code, attn, cost


def grp(rows):
    g = defaultdict(list)
    for r in rows:
        if r["model"] in MODELS:
            g[r["model"]].append(r)
    return g


def main():
    paths = sys.argv[1:]
    if not paths:
        sys.exit("usage: model-behavior-metrics.py mb-*.jsonl mc-*.jsonl")
    turns, sess, code, attn, cost = load(paths)
    g, gs, gc, ga = grp(turns), grp(sess), grp(code), grp(attn)
    gx = defaultdict(list)
    for r in cost:
        if r["orch_model"] in MODELS:
            gx[r["orch_model"]].append(r)
    M = {}

    def put(fam, key, label, fn, rows, unit="", better="low", note=""):
        M.setdefault(fam, []).append(dict(
            key=key, label=label, unit=unit, better=better, note=note,
            v={SHORT[m]: round(fn(rows[m]), 3) for m in MODELS}))

    STY = lambda r: [t["style"] for t in r if t.get("style")]

    def p1k(r, k):
        s = STY(r)
        return 1000 * sum(x[k] for x in s) / max(1, sum(x["chars"] for x in s))

    # ── corpus ──
    put("corpus", "turns", "Turns analysed", len, g, "turns", "high")
    put("corpus", "sessions", "Sessions", lambda r: len({t["sess"] for t in r}), g, "sessions", "high")
    put("corpus", "projects", "Distinct projects", lambda r: len({t["proj"] for t in r}), g, "projects", "high")

    # ── orchestration ──
    put("orch", "endcold", "Dispatch turns ending cold",
        lambda r: 100 * sum(1 for t in r if t["post_disp"] == 0) /
                  max(1, sum(1 for t in r if t["post_disp"] is not None)),
        g, "%", "low", "Turn ends the moment agents report back")
    put("orch", "postdisp", "Work after last dispatch",
        lambda r: st.mean([t["post_disp"] for t in r if t["post_disp"] is not None]), g, "calls", "high")
    put("orch", "dispatchrate", "Turns containing a dispatch",
        lambda r: 100 * sum(1 for t in r if t["dispatch"]) / len(r), g, "%", "high")
    put("orch", "zerotool", "Turns doing no work at all",
        lambda r: 100 * sum(1 for t in r if not t["tools"]) / len(r), g, "%", "low")
    put("orch", "toolsturn", "Tool calls per turn",
        lambda r: sum(t["tools"] for t in r) / len(r), g, "calls", "high")
    put("orch", "sendmsg", "SendMessage per 100 turns",
        lambda r: 100 * sum(t["tc"].get("SendMessage", 0) for t in r) / len(r), g, "/100", "high",
        "Live-agent feedback rounds")
    put("orch", "agentsperwf", "Agents per Workflow run",
        lambda r: sum(x["wf_agents"] for x in r) / max(1, sum(x["workflows"] for x in r)),
        gs, "agents", "high", "Wave width")
    put("orch", "subpersess", "Subagents per session",
        lambda r: sum(x["wf_agents"] + x["plain_agents"] for x in r) / len(r), gs, "agents", "high")
    put("orch", "subtok", "Worker output tokens/session",
        lambda r: sum(x["sub_tok"] for x in r) / len(r) / 1000, gs, "k tok", "high")

    # ── escalation ──
    put("esc", "askuq", "AskUserQuestion per 100 turns",
        lambda r: 100 * sum(t["tc"].get("AskUserQuestion", 0) for t in r) / len(r), g, "/100", "low")
    put("esc", "endq", "Turns ending in a question",
        lambda r: 100 * sum(t["ends_q"] for t in r) / len(r), g, "%", "low")
    put("esc", "perm", "Permission phrasing", lambda r: p1k(r, "permission"), g, "/1k ch", "low",
        '"want me to", "should I", "let me know"')
    put("esc", "limit", "Limitation phrasing", lambda r: p1k(r, "limitation"), g, "/1k ch", "low",
        '"I didn\'t", "not verified", "out of scope"')
    put("esc", "caveat", "Caveat phrasing", lambda r: p1k(r, "caveat"), g, "/1k ch", "low")
    put("esc", "hedge", "Hedging", lambda r: p1k(r, "hedge"), g, "/1k ch", "low")
    put("esc", "interrupt", "Turns you interrupted",
        lambda r: 100 * sum(t["interrupted"] for t in r) / len(r), g, "%", "low")
    put("esc", "nudge", "Turns needing a nudge",
        lambda r: 100 * sum(t["nudge"] for t in r) / len(r), g, "%", "low")

    # ── prose ──
    put("prose", "finalmed", "Closing message (median)",
        lambda r: st.median([x["chars"] for x in STY(r)]), g, "chars", "low")
    put("prose", "bullets", "Bullets", lambda r: p1k(r, "bullet"), g, "/1k ch", "high")
    put("prose", "headers", "Headers", lambda r: p1k(r, "header"), g, "/1k ch", "high")
    put("prose", "bold", "Bold anchors", lambda r: p1k(r, "bold"), g, "/1k ch", "high")
    put("prose", "tables", "Table rows", lambda r: p1k(r, "table"), g, "/1k ch", "high")
    put("prose", "digits", "Digit density", lambda r: p1k(r, "digit"), g, "/1k ch", "high",
        "Numbers per 1k chars")
    put("prose", "wps", "Words per sentence",
        lambda r: st.mean([x["wps"] for x in STY(r) if x["wps"]]), g, "words", "low")
    put("prose", "uwr", "Unique-word ratio", lambda r: st.mean([x["uwr"] for x in STY(r)]), g, "ratio", "high")
    put("prose", "narr", "Narration per tool call",
        lambda r: sum(t["text_chars"] for t in r) / max(1, sum(t["tools"] for t in r)), g, "chars", "high")
    put("prose", "outtok", "Output tokens per turn",
        lambda r: sum(t["out_tok"] for t in r) / len(r) / 1000, g, "k tok", "low")

    # ── coding: orchestrator + its workers, per session ──
    T = lambda r, k: sum(x["orch"].get(k, 0) + x["work"].get(k, 0) for x in r)
    put("code", "csess", "Sessions", lambda r: len(r), gc, "sessions", "high")
    put("code", "edits", "Edits per session", lambda r: T(r, "edits") / len(r), gc, "edits", "high",
        "Orchestrator + its workers")
    put("code", "editchars", "Code written per session",
        lambda r: T(r, "edit_chars") / len(r) / 1000, gc, "k chars", "high")
    put("code", "editfail", "Edit failure rate",
        lambda r: 100 * T(r, "err_Edit") / max(1, T(r, "edits")), gc, "%", "low",
        "Edit rejected: string not found / file unread")
    put("code", "toolerr", "Tool error rate",
        lambda r: 100 * T(r, "errors") / max(1, T(r, "calls")), gc, "/100 calls", "low")
    put("code", "rework", "Rework ratio",
        lambda r: T(r, "refile") / max(1, T(r, "edits")), gc, "ratio", "low",
        "Repeat edits to the same file")
    put("code", "tests", "Tests per 100 edits",
        lambda r: 100 * T(r, "test_runs") / max(1, T(r, "edits")), gc, "/100", "high")
    put("code", "commits", "Commits per session", lambda r: T(r, "commits") / len(r), gc, "commits", "high")
    put("code", "reverts", "Reverts per session",
        lambda r: sum(x["orch"].get("reverts", 0) for x in r) / len(r), gc, "reverts", "low")
    put("code", "delegation", "Edits done by workers",
        lambda r: 100 * sum(x["work"].get("edits", 0) for x in r) / max(1, T(r, "edits")), gc, "%", "high")

    # ── attention: what it costs the operator, not the token budget ──
    def hrs(r):
        w = [x["wait_s"] for x in r if x["wait_s"] is not None and x["wait_s"] <= 4 * 3600]
        return (sum(x["work_s"] for x in r) + sum(w)) / 3600
    ASK = re.compile(r"(want me to|should i\b|shall i\b|do you want|would you like|"
                     r"let me know|your call|say the word|i can (?:also )?(?:dispatch|run|do|start)|"
                     r"confirm|approve|awaiting|waiting on you)", re.I)
    ENDQ = re.compile(r"\?[\s\"'`*)\]]*$")
    asks = lambda t: bool(t) and (bool(ENDQ.search(t)) or bool(ASK.search(t[-320:])))

    put("attn", "demands_hr", "Handbacks per hour",
        lambda r: len(r) / max(1e-9, hrs(r)), ga, "/hour", "low",
        "Excludes background wakes and subagent traffic")
    put("attn", "cold_hr", "Ends cold — no ask, no resume",
        lambda r: sum(1 for x in r if not asks(x["tail"])) / max(1e-9, hrs(r)), ga, "/hour", "low")
    put("attn", "ask_hr", "Asks a question or offers",
        lambda r: sum(1 for x in r if asks(x["tail"])) / max(1e-9, hrs(r)), ga, "/hour", "low")
    put("attn", "selfresume", "Self-resumes per span",
        lambda r: sum(x["autos"] for x in r) / len(r), ga, "wakes", "high",
        "Background wakes it absorbed instead of asking you")
    put("attn", "run_med", "Median unattended run",
        lambda r: st.median([x["work_s"] for x in r]) / 60, ga, "min", "high")
    put("attn", "colddisp", "Dispatch spans ending cold",
        lambda r: 100 * sum(x["cold"] for x in r if x["dispatch"]) /
                  max(1, sum(1 for x in r if x["dispatch"])), ga, "%", "low",
        "Ended the turn while workers were still out")
    put("attn", "reading_hr", "Reading burden per hour",
        lambda r: sum(x["final_chars"] for x in r) / max(1e-9, hrs(r)) / 1000, ga, "k chars", "low")
    put("attn", "demands_sess", "Handbacks per session",
        lambda r: len(r) / len({x["sess"] for x in r}), ga, "demands", "low")

    # ── cost: API-list spend, actual model mix ──
    RATE = {"claude-fable-5": 10, "claude-opus-5": 5, "claude-opus-4-8": 5,
            "claude-sonnet-5": 3, "claude-haiku-4-5-20251001": 1}

    def dollars(rows):
        t = 0.0
        for x in rows:
            i = RATE.get(x["model"], 5)
            t += (x["input"] * i + x["output"] * 5 * i +
                  x["cache_write"] * 1.25 * i + x["cache_read"] * 0.1 * i) / 1e6
        return t
    nses = lambda r: len({x["sess"] for x in r})
    put("cost", "per_session", "Cost per session",
        lambda r: dollars(r) / nses(r), gx, "$", "low", "API list rates, actual model mix")
    put("cost", "orch_share", "Share spent on the main loop",
        lambda r: 100 * dollars([x for x in r if x["role"] == "orch"]) / max(1e-9, dollars(r)),
        gx, "%", "low")
    put("cost", "cache_read_sess", "Cache-read tokens per session",
        lambda r: sum(x["cache_read"] for x in r) / nses(r) / 1e6, gx, "M tok", "low")

    # ── composite indices ──
    byk = {m["key"]: m for fam in ("orch", "code", "esc", "prose") for m in M[fam]}

    def index(keys):
        tot = {SHORT[m]: 0 for m in MODELS}
        for kk in keys:
            m = byk[kk]
            v = m["v"]
            vals = [v[SHORT[x]] for x in MODELS]
            if m["better"] == "high":
                best = max(vals)
                sc = {SHORT[x]: (100 * v[SHORT[x]] / best if best else 0) for x in MODELS}
            else:
                pos = [q for q in vals if q > 0]
                best = min(pos) if pos else 0
                sc = {SHORT[x]: (100 * best / v[SHORT[x]] if v[SHORT[x]] > 0 else 100) for x in MODELS}
            for x in MODELS:
                tot[SHORT[x]] += sc[SHORT[x]]
        return {k: round(tot[k] / len(keys), 1) for k in tot}

    PERSIST = ["endcold", "postdisp", "agentsperwf", "subpersess"]
    ACTIVITY = ["zerotool", "toolsturn", "subtok", "dispatchrate"]
    CODE = ["edits", "editchars", "editfail", "toolerr", "rework", "tests", "commits", "reverts"]
    M["index"] = {
        "persist": index(PERSIST), "activity": index(ACTIVITY), "code": index(CODE),
        "esc": index([m["key"] for m in M["esc"]]),
        "prose": index([m["key"] for m in M["prose"]]),
        "orch_all": index(PERSIST + ACTIVITY),
        "members": {"persist": PERSIST, "activity": ACTIVITY, "code": CODE},
    }
    json.dump(M, sys.stdout, separators=(",", ":"))


if __name__ == "__main__":
    main()
