#!/usr/bin/env python3
"""Arm A2-id — Implementation A (pandas). Executes PREREG_A2-id_2026-07-25.md.
Usage: python impl_A.py <out.json>   (run from the work/ directory)"""
import json, sys, hashlib
import numpy as np
import pandas as pd

TEST_PQ = "test.parquet"
INST_PQ = "instances.parquet"
SOURCE_HASH = open("source_hash.txt").read().strip()

C3 = {  # publisher-code constant (claim); C2 (README table) is content-identical
    "swe-bench": "Issue Resolution",
    "swe-bench-multimodal": "Frontend",
    "commit0": "Greenfield",
    "swt-bench": "Testing",
    "gaia": "Information Gathering",
}
SLUG = lambda cat: cat.lower().replace(" ", "_")

NO_ANTECEDENT_COLS = (
    ["sdk_version", "openness", "country", "supports_vision", "release_date",
     "average_runtime"]
    + [f"{SLUG(c)}_runtime" for c in C3.values()]
    + [f"{SLUG(c)}_logs_url" for c in C3.values()]
    + [f"{SLUG(c)}_visualization_url" for c in C3.values()]
)
CATEGORICAL_COLS = ["agent_name", "agent_type", "language_model"]

SCORE_T1, SCORE_T2 = 0.05, 0.5
COST_T1, COST_T2 = 0.005, 0.05
SEED, NDRAWS = 20260725, 1000


def r9(x):
    return None if x is None else round(x, 9)


def median(vals):
    s = sorted(vals)
    n = len(s)
    if n == 0:
        return None
    return s[n // 2] if n % 2 else (s[n // 2 - 1] + s[n // 2]) / 2.0


def tier_numeric(rec, pub, dp, t1, t2):
    if rec is None:
        return None
    if round(rec, dp) == pub:
        return 0
    d = abs(rec - pub)
    if d <= t1:
        return 1
    if d <= t2:
        return 2
    return "div"


def main(out_path):
    td = pd.read_parquet(TEST_PQ)
    idf = pd.read_parquet(INST_PQ)
    out = {"meta": {"impl": "A", "prereg": "PREREG_A2-id_2026-07-25.md"}}

    # ---------- P4: join-key search ----------
    test_ids = td["id"].tolist()
    cands = {}
    keyfuncs = {
        "K1_id": (td["id"], idf["id"]),
        "K2_language_model": (td["language_model"], idf["language_model"]),
        "K3_agentname_lm": (td["agent_name"] + "\x1f" + td["language_model"],
                            idf["agent_name"] + "\x1f" + idf["language_model"]),
        "K4_id_lower": (td["id"].str.lower(), idf["id"].str.lower()),
    }
    for name, (tk, ik) in keyfuncs.items():
        groups = set(ik.unique())
        cover = float(np.mean([k in groups for k in tk]))
        # uniqueness: no two test rows share a key, and each matched group hits 1 test row
        unique = bool(tk.is_unique)
        cands[name] = {"coverage": cover, "test_key_unique": unique,
                       "n_instance_groups": int(ik.nunique())}
    order = ["K1_id", "K2_language_model", "K3_agentname_lm", "K4_id_lower"]
    ranked = sorted(order, key=lambda n: (-cands[n]["coverage"], not cands[n]["test_key_unique"], order.index(n)))
    winner, runner_up = ranked[0], ranked[1]
    out["join"] = {"candidates": cands, "winner": winner, "runner_up": runner_up}
    tkey, ikey = keyfuncs[winner]
    td = td.assign(_k=tkey.values)
    idf = idf.assign(_k=ikey.values)

    # ---------- P4: column-family mapping ----------
    c1_pairs = sorted(set(zip(idf["benchmark"], idf["category"])))
    c1_map = {}
    amb = []
    for b, c in c1_pairs:
        if b in c1_map and c1_map[b] != c:
            amb.append(b)
        c1_map.setdefault(b, c)
    agrees_c3 = {b: (c1_map.get(b) == c) for b, c in C3.items()}
    out["colmap"] = {"C1_pairs": [list(p) for p in c1_pairs],
                     "C1_ambiguous_benchmarks": amb,
                     "agrees_with_C2_README_and_C3_code": agrees_c3,
                     "mapping_used": {b: SLUG(c) for b, c in c1_map.items() if c is not None}}
    b2slug = {b: SLUG(c) for b, c in c1_map.items()}
    slugs_in_test = [SLUG(c) for c in C3.values()]

    # instances-side negative space
    inst_groups = set(idf["_k"].unique())
    test_keys = set(td["_k"])
    out["instances_side"] = {
        "n_rows": int(len(idf)),
        "distinct_benchmarks": sorted(idf["benchmark"].unique().tolist()),
        "groups_without_test_row": sorted(inst_groups - test_keys),
        "test_rows_without_group": sorted(test_keys - inst_groups),
        "n_groups": len(inst_groups),
        "resolved_nulls": int(idf["resolved"].isna().sum()),
        "cost_nulls": int(idf["cost"].isna().sum()),
    }

    # ---------- group stats ----------
    gb = {}
    for (k, b), g in idf.groupby(["_k", "benchmark"]):
        res = g["resolved"]
        costs = g["cost"].dropna().astype(float).tolist()
        gb[(k, b)] = {
            "n_all": int(len(g)),
            "n_obs": int(res.notna().sum()),
            "r_true": int((res == True).sum()),  # noqa: E712
            "costs": costs,
        }
    gmodel = {}
    for k, g in idf.groupby("_k"):
        gmodel[k] = {
            "vals": {c: sorted(set(g[c].dropna().tolist())) for c in CATEGORICAL_COLS},
            "modal": {c: (g[c].mode().sort_values().iloc[0] if len(g) else None)
                      for c in CATEGORICAL_COLS},
            "benchmarks": sorted(set(g["benchmark"])),
            "bench_with_obs": sorted(set(g.loc[g["resolved"].notna(), "benchmark"])),
        }

    # ---------- per-column policies ----------
    columns = {}
    rows = td.to_dict("records")

    def score_policies(k, b):
        s = gb.get((k, b))
        if s is None:
            return None
        p = {}
        p["S1_pct_Dall"] = 100.0 * s["r_true"] / s["n_all"] if s["n_all"] else None
        p["S2_pct_Dobs"] = 100.0 * s["r_true"] / s["n_obs"] if s["n_obs"] else None
        p["S1_frac_Dall"] = s["r_true"] / s["n_all"] if s["n_all"] else None
        p["S2_frac_Dobs"] = s["r_true"] / s["n_obs"] if s["n_obs"] else None
        return p

    def cost_policies(k, b):
        s = gb.get((k, b))
        if s is None:
            return None
        c, na = s["costs"], s["n_all"]
        p = {}
        p["C_mean_excl"] = sum(c) / len(c) if c else None
        p["C_mean_zero"] = sum(c) / na if na else None
        p["C_median_excl"] = median(c)
        p["C_median_zero"] = median(c + [0.0] * (na - len(c))) if na else None
        p["C_sum_over_resolved"] = (sum(c) / s["r_true"]) if s["r_true"] else None
        return p

    POLICY_ORDER = {}  # column -> ordered policy names (prereg enumeration order)

    def add_col(col, family, dp, t1, t2, per_row_policies, primary):
        pub = {r["id"]: r[col] for r in rows}
        pols = {}
        pnames = list(next(v for v in per_row_policies.values() if v is not None).keys()) \
            if any(v is not None for v in per_row_policies.values()) else []
        for pn in pnames:
            per_row, cnt = {}, {0: 0, 1: 0, 2: 0, "div": 0, None: 0}
            for r in rows:
                rid, k = r["id"], r["_k"]
                pp = per_row_policies.get(rid)
                rec = None if pp is None else pp.get(pn)
                if family == "count":
                    t = (0 if rec == pub[rid] else "div") if rec is not None else None
                else:
                    t = tier_numeric(rec, pub[rid], dp, t1, t2)
                d = None if rec is None else r9(rec - pub[rid])
                per_row[rid] = [rec, d, t]
                cnt[t if t in (0, 1, 2, "div") else None] += 1
            pols[pn] = {"per_row": per_row, "t0": cnt[0], "t1": cnt[1], "t2": cnt[2],
                        "div": cnt["div"], "undef": cnt[None]}
        vals = pd.Series([pub[r["id"]] for r in rows])
        prim = pols.get(primary, {}).get("per_row", {})
        deltas = {rid: abs(v[1]) for rid, v in prim.items() if v[1] is not None}
        mx = max(deltas.values()) if deltas else None
        columns[col] = {
            "family": family, "dp": dp,
            "published_range": [float(vals.min()), float(vals.max())] if family != "categorical" else None,
            "zero_variance": bool(vals.nunique() == 1),
            "policies": pols, "primary": primary,
            "max_abs_delta": mx,
            "max_abs_rows": sorted([rid for rid, d in deltas.items() if d == mx]) if mx is not None else [],
        }
        POLICY_ORDER[col] = pnames

    # per-category score / cost columns
    slug2b = {v: kk for kk, v in b2slug.items()}
    for slug in slugs_in_test:
        b = slug2b.get(slug)
        sp = {r["id"]: score_policies(r["_k"], b) for r in rows}
        add_col(f"{slug}_score", "score", 1, SCORE_T1, SCORE_T2, sp, "S1_pct_Dall")
        cp = {r["id"]: cost_policies(r["_k"], b) for r in rows}
        add_col(f"{slug}_cost", "cost", 4, COST_T1, COST_T2, cp, "C_mean_excl")

    # average_score
    asp = {}
    for r in rows:
        k = r["_k"]
        s1, s2, na_sum, no_sum, rt_sum = [], [], 0, 0, 0
        for b in C3:
            st = gb.get((k, b))
            if st is None:
                continue
            if st["n_all"]:
                s1.append(100.0 * st["r_true"] / st["n_all"])
            if st["n_obs"]:
                s2.append(100.0 * st["r_true"] / st["n_obs"])
            na_sum += st["n_all"]; no_sum += st["n_obs"]; rt_sum += st["r_true"]
        asp[r["id"]] = {
            "AS_macro_raw_Dall": sum(s1) / len(s1) if s1 else None,
            "AS_macro_raw_Dobs": sum(s2) / len(s2) if s2 else None,
            "AS_macro_pub1dp": sum(round(x, 1) for x in s1) / len(s1) if s1 else None,
            "AS_micro_Dall": 100.0 * rt_sum / na_sum if na_sum else None,
            "AS_micro_Dobs": 100.0 * rt_sum / no_sum if no_sum else None,
            "AS_macro_all5_zero": sum(s1) / 5.0 if s1 else None,
        }
    add_col("average_score", "score", 2, SCORE_T1, SCORE_T2, asp, "AS_macro_raw_Dall")

    # average_cost
    acp = {}
    for r in rows:
        k = r["_k"]
        cm, all_costs, n_all_tot = [], [], 0
        for b in C3:
            st = gb.get((k, b))
            if st is None:
                continue
            if st["costs"]:
                cm.append(sum(st["costs"]) / len(st["costs"]))
            all_costs += st["costs"]; n_all_tot += st["n_all"]
        acp[r["id"]] = {
            "AC_macro_raw": sum(cm) / len(cm) if cm else None,
            "AC_macro_4dp": sum(round(x, 4) for x in cm) / len(cm) if cm else None,
            "AC_micro_excl": sum(all_costs) / len(all_costs) if all_costs else None,
            "AC_micro_zero": sum(all_costs) / n_all_tot if n_all_tot else None,
            "AC_macro_all5_zero": sum(cm) / 5.0 if cm else None,
        }
    add_col("average_cost", "cost", 4, COST_T1, COST_T2, acp, "AC_macro_raw")

    # categories_completed
    ccp = {}
    for r in rows:
        gm = gmodel.get(r["_k"])
        ccp[r["id"]] = None if gm is None else {
            "CC1_any_row": len([b for b in gm["benchmarks"] if b in C3]),
            "CC2_any_obs": len([b for b in gm["bench_with_obs"] if b in C3]),
        }
    add_col("categories_completed", "count", None, None, None, ccp, "CC1_any_row")

    # categorical columns
    for col in CATEGORICAL_COLS:
        pub = {r["id"]: r[col] for r in rows}
        pols = {}
        for pn in ["U_unique", "U_major"]:
            per_row, cnt = {}, {0: 0, "div": 0, None: 0}
            for r in rows:
                gm = gmodel.get(r["_k"])
                if gm is None:
                    per_row[r["id"]] = [None, None, None]; cnt[None] += 1; continue
                if pn == "U_unique":
                    vs = gm["vals"][col]
                    rec = vs[0] if len(vs) == 1 else None
                    if len(vs) != 1:
                        per_row[r["id"]] = [None, f"multiplicity={len(vs)}", "div"]; cnt["div"] += 1; continue
                else:
                    rec = gm["modal"][col]
                t = 0 if rec == pub[r["id"]] else "div"
                per_row[r["id"]] = [rec, None, t]; cnt[t] += 1
            pols[pn] = {"per_row": per_row, "t0": cnt[0], "t1": 0, "t2": 0,
                        "div": cnt["div"], "undef": cnt[None]}
        vals = pd.Series([pub[r["id"]] for r in rows])
        columns[col] = {"family": "categorical", "dp": None, "published_range": None,
                        "zero_variance": bool(vals.nunique() == 1),
                        "policies": pols, "primary": "U_unique",
                        "max_abs_delta": None, "max_abs_rows": []}
        POLICY_ORDER[col] = ["U_unique", "U_major"]

    # id column (join key when K1/K4 wins -> structural self-match)
    columns["id"] = {"family": "join_key", "dp": None, "published_range": None,
                     "zero_variance": False, "policies": {}, "primary": None,
                     "max_abs_delta": None, "max_abs_rows": []}

    # ---------- P12 cell accounting ----------
    ADDRESSED = [c for c in td.columns if c not in NO_ANTECEDENT_COLS and c != "_k"]
    by_col, totals = {}, {"reconciled": 0, "divergent": 0, "no_antecedent": 0, "structural": 0}
    tier_break = {0: 0, 1: 0, 2: 0}
    for col in [c for c in td.columns if c != "_k"]:
        cls = {}
        if col in NO_ANTECEDENT_COLS:
            cls = {r["id"]: "no_antecedent" for r in rows}
        elif col == "id" and winner in ("K1_id", "K4_id_lower"):
            cls = {r["id"]: "structural" for r in rows}
        elif columns[col]["zero_variance"]:
            cls = {r["id"]: "structural" for r in rows}
        else:
            for r in rows:
                rid = r["id"]
                best, bestp = None, None
                for pn in POLICY_ORDER[col]:
                    t = columns[col]["policies"][pn]["per_row"][rid][2]
                    if t in (0, 1, 2) and (best is None or best == "div" or t < best):
                        if best is None or best == "div" or t < best:
                            best, bestp = t, pn
                    elif t == "div" and best is None:
                        best, bestp = "div", pn
                if best in (0, 1, 2):
                    cls[rid] = f"reconciled_T{best}:{bestp}"
                    tier_break[best] += 1
                elif best == "div":
                    cls[rid] = "divergent"
                else:
                    cls[rid] = "no_antecedent(all_policies_undefined)"
        c = {"reconciled": 0, "divergent": 0, "no_antecedent": 0, "structural": 0}
        for v in cls.values():
            if v.startswith("reconciled"):
                c["reconciled"] += 1
            elif v == "divergent":
                c["divergent"] += 1
            elif v.startswith("no_antecedent"):
                c["no_antecedent"] += 1
            else:
                c["structural"] += 1
        by_col[col] = {"counts": c, "cells": cls}
        for kk in totals:
            totals[kk] += c[kk]
    totals["total"] = sum(v for v in totals.values())
    totals["reconciled_by_tier"] = {f"T{k}": v for k, v in tier_break.items()}
    out["cells"] = {"totals": totals, "by_column": by_col}

    # ---------- P9 permutation nulls ----------
    rng = np.random.default_rng(SEED)
    perms = [rng.permutation(len(rows)) for _ in range(NDRAWS)]
    nulls = {}
    for col in ADDRESSED:
        if col == "id" or columns[col].get("family") == "join_key":
            continue
        fam = columns[col]["family"]
        prim = columns[col]["primary"]
        pr = columns[col]["policies"][prim]["per_row"]
        pub = [rows[i][col] for i in range(len(rows))]
        rec = [pr[rows[i]["id"]][0] for i in range(len(rows))]
        if columns[col]["zero_variance"]:
            nulls[col] = {"structural_zero_variance": True,
                          "note": "null distribution degenerate; floor structural, ratios undefined"}
            continue
        dp, t1, t2 = columns[col]["dp"], None, None
        if fam == "score":
            t1, t2 = SCORE_T1, SCORE_T2
        elif fam == "cost":
            t1, t2 = COST_T1, COST_T2
        def stats(perm):
            t0 = 0
            devs = []
            for i in range(len(rows)):
                rv = rec[perm[i]]
                if rv is None:
                    continue
                if fam in ("count", "categorical"):
                    t0 += 1 if rv == pub[i] else 0
                else:
                    if round(rv, dp) == pub[i]:
                        t0 += 1
                    devs.append(abs(rv - pub[i]))
            mad = sum(devs) / len(devs) if devs else None
            return t0, mad
        obs_t0, obs_mad = stats(list(range(len(rows))))
        nt0, nmad = [], []
        for p in perms:
            a, b = stats(p)
            nt0.append(a)
            if b is not None:
                nmad.append(b)
        nulls[col] = {
            "observed_t0": obs_t0, "observed_mad": obs_mad,
            "t0_median": float(np.median(nt0)), "t0_p5": float(np.percentile(nt0, 5)),
            "t0_p95": float(np.percentile(nt0, 95)),
            "mad_median": float(np.median(nmad)) if nmad else None,
            "mad_p5": float(np.percentile(nmad, 5)) if nmad else None,
            "n_draws": NDRAWS, "seed": SEED,
        }
    out["nulls"] = nulls

    # ---------- L4 tie census on rankings ----------
    def tie_census(vals):
        n = len(vals)
        pairs = sum(1 for i in range(n) for j in range(i + 1, n) if vals[i] == vals[j])
        s = sorted(vals, reverse=True)
        adj = sum(1 for i in range(n - 1) if s[i] == s[i + 1])
        return {"n": n, "distinct": len(set(vals)), "tied_pairs": pairs,
                "adjacent_ties_in_desc_sort": adj, "strictly_decidable": pairs == 0}
    pub_as = [r["average_score"] for r in rows]
    prim_pr = columns["average_score"]["policies"]["AS_macro_raw_Dall"]["per_row"]
    rec_as = [prim_pr[r["id"]][0] for r in rows if prim_pr[r["id"]][0] is not None]
    out["tie_census"] = {
        "published_average_score": tie_census(pub_as),
        "recomputed_primary_raw": tie_census(rec_as),
        "recomputed_primary_2dp": tie_census([round(x, 2) for x in rec_as]),
        "published_sort_descending": bool(all(pub_as[i] >= pub_as[i + 1] for i in range(len(pub_as) - 1))),
    }

    # ---------- §3.1 self-identification ----------
    def find_rows(pred):
        return [i for i, r in enumerate(rows) if pred(str(r["language_model"]).lower())]
    a2_exact = find_rows(lambda s: s == "claude-fable-5")
    a2 = a2_exact if a2_exact else find_rows(lambda s: "fable" in s)
    b2 = find_rows(lambda s: s.startswith("gemini"))
    def rowinfo(i):
        r = {k: v for k, v in rows[i].items() if k != "_k"}
        r2 = {k: (None if isinstance(v, float) and pd.isna(v) else v) for k, v in r.items()}
        srt = sorted(pub_as, reverse=True)
        return {"file_order_rank_1based": i + 1,
                "rank_by_average_score": srt.index(rows[i]["average_score"]) + 1,
                "published_row": r2}
    out["self_id"] = {
        "rule": "exact lower(language_model)=='claude-fable-5', fallback substring 'fable'; B2: startswith 'gemini' (pin unknown to this arm)",
        "A2_rows": {str(i): rowinfo(i) for i in a2},
        "A2_match_type": "exact" if a2_exact else ("substring" if a2 else "none"),
        "B2_candidate_rows": {str(i): rowinfo(i) for i in b2},
    }

    # ---------- M-H hash reconstruction (Impl A only, prereg addition A2) ----------
    h = hashlib.sha256()
    for n_, df_ in enumerate([pd.read_parquet(TEST_PQ), pd.read_parquet(INST_PQ)]):
        if n_:
            h.update(b"\n--\n")
        h.update(df_.reindex(sorted(df_.columns), axis=1)
                 .to_csv(index=False, float_format="%.10g").encode("utf-8"))
    out["hash_check"] = {"executed": True, "recomputed": h.hexdigest(),
                         "shipped_source_hash": SOURCE_HASH,
                         "equal": h.hexdigest() == SOURCE_HASH}

    # P10 ranges for addressed numeric columns are in columns[*].published_range
    out["columns"] = columns

    def np_default(o):
        if isinstance(o, np.integer):
            return int(o)
        if isinstance(o, np.bool_):
            return bool(o)
        if isinstance(o, np.floating):
            return float(o)
        raise TypeError(f"not serializable: {type(o)}")

    with open(out_path, "w") as f:
        json.dump(out, f, sort_keys=True, indent=1, default=np_default)
        f.write("\n")
    print(f"impl_A wrote {out_path}; cell totals: {totals}")


if __name__ == "__main__":
    main(sys.argv[1])
