#!/usr/bin/env python3
"""T^GPT v0.8.5 semantic shift / recovery engine.

Engineering repair after external adversarial review:
- semantic distributions are conditioned on known mapped mass;
- mapping mass is measured separately and hard-gated;
- differential mapping cannot silently become semantic deficit;
- A/A'/B reference-replicate calibration requires equal total N;
- RECOVERY vs NEUTRAL_CONTROL mapping is hard-gated;
- paired recovery uses conditional semantic distributions;
- generalized inference remains task-level.

This tool does not interpret any metric as T^, a latent total field, or a remainder.
"""
from __future__ import annotations
import argparse, json, math, random, hashlib
from collections import Counter
from pathlib import Path
from typing import Dict, Iterable, List, Tuple, Optional

SPECIAL = {"UNMAPPED", "UNRESOLVED"}
MIN_KNOWN_SEMANTIC_MASS = 0.80
MAX_MAPPING_GAP = 0.05


def validate_labels(labels: Iterable[str], relevance: Dict[str, int], where: str) -> None:
    allowed = set(relevance) | SPECIAL
    bad = sorted({str(x) for x in labels if str(x) not in allowed})
    if bad:
        raise ValueError(f"{where}: unknown/final-forbidden labels: {bad}")


def relevant_ids(relevance: Dict[str, int]) -> List[str]:
    bad = [k for k, v in relevance.items() if int(v) not in (0, 1)]
    if bad:
        raise ValueError(f"relevance values must be 0/1: {bad}")
    return [k for k, v in relevance.items() if int(v) == 1]


def rates(labels: List[str], relevance: Dict[str, int]) -> dict:
    validate_labels(labels, relevance, "labels")
    n = len(labels)
    if not n:
        raise ValueError("arm must contain at least one finalized generation")
    c = Counter(map(str, labels))
    known = sum(c.get(k, 0) for k in relevance)
    return {
        "N_total": n,
        "N_known": known,
        "known_semantic_mass": known / n,
        "UNMAPPED_rate": c.get("UNMAPPED", 0) / n,
        "UNRESOLVED_rate": c.get("UNRESOLVED", 0) / n,
    }


def probs_unconditional(labels: List[str], relevance: Dict[str, int]) -> Dict[str, float]:
    """Known-cluster mass over all finalized generations (mapping-sensitive; descriptive only)."""
    validate_labels(labels, relevance, "labels")
    n = len(labels)
    if n == 0:
        raise ValueError("arm must contain at least one finalized generation")
    c = Counter(map(str, labels))
    return {cluster: c.get(cluster, 0) / n for cluster in relevance}


def probs_known(labels: List[str], relevance: Dict[str, int]) -> Dict[str, float]:
    """Semantic distribution conditional on a finalized known-cluster assignment."""
    validate_labels(labels, relevance, "labels")
    c = Counter(map(str, labels))
    nk = sum(c.get(cluster, 0) for cluster in relevance)
    if nk <= 0:
        raise ValueError("arm has zero known semantic assignments")
    return {cluster: c.get(cluster, 0) / nk for cluster in relevance}


def mapping_gate(named_rates: Dict[str, dict], compare_pairs: List[Tuple[str, str]]) -> dict:
    low = {name: r["known_semantic_mass"] for name, r in named_rates.items()
           if r["known_semantic_mass"] < MIN_KNOWN_SEMANTIC_MASS}
    gaps = {}
    for a, b in compare_pairs:
        gap = abs(named_rates[a]["known_semantic_mass"] - named_rates[b]["known_semantic_mass"])
        gaps[f"{a}__{b}"] = gap
    excessive = {k: v for k, v in gaps.items() if v > MAX_MAPPING_GAP + 1e-15}
    ok = not low and not excessive
    return {
        "status": "OK" if ok else "MEASUREMENT_INSUFFICIENT",
        "minimum_known_semantic_mass": MIN_KNOWN_SEMANTIC_MASS,
        "maximum_mapping_gap": MAX_MAPPING_GAP,
        "low_known_mass": low,
        "pairwise_mapping_gaps": gaps,
        "excessive_mapping_gaps": excessive,
    }


def directional(pa: Dict[str, float], pb: Dict[str, float], rel: List[str]) -> Tuple[float, float, Dict[str, float], Dict[str, float]]:
    d = {c: max(0.0, pa.get(c, 0.0) - pb.get(c, 0.0)) for c in rel}
    e = {c: max(0.0, pb.get(c, 0.0) - pa.get(c, 0.0)) for c in rel}
    return sum(d.values()), sum(e.values()), d, e


def closure(pb: Dict[str, float], pr: Dict[str, float], deficit_by_cluster: Dict[str, float]) -> float:
    total = 0.0
    for c, d in deficit_by_cluster.items():
        inc = max(0.0, pr.get(c, 0.0) - pb.get(c, 0.0))
        total += min(d, inc)
    return total


def known_only(labels: List[str], relevance: Dict[str, int]) -> List[str]:
    ks = set(relevance)
    return [str(x) for x in labels if str(x) in ks]


def perm_null_known(a: List[str], b: List[str], relevance: Dict[str, int], reps: int, rng: random.Random) -> dict:
    """Permutation null for the conditional semantic distribution, not mapping mass."""
    rel = relevant_ids(relevance)
    ak = known_only(a, relevance); bk = known_only(b, relevance)
    if not ak or not bk:
        raise ValueError("permutation requires known assignments in both arms")
    pool = list(ak) + list(bk); na = len(ak)
    ds, es = [], []
    for _ in range(reps):
        x = pool[:]
        rng.shuffle(x)
        aa, bb = x[:na], x[na:]
        d, e, _, _ = directional(probs_known(aa, relevance), probs_known(bb, relevance), rel)
        ds.append(d); es.append(e)
    return {
        "deficit_null_mean": sum(ds) / len(ds),
        "expansion_null_mean": sum(es) / len(es),
        "_deficit_draws": ds,
        "_expansion_draws": es,
        "N_known_A": len(ak),
        "N_known_B": len(bk),
    }


def paired_swap_null(pairs: List[dict], relevance: Dict[str, int], a: List[str], reps: int,
                     rng: random.Random) -> dict:
    rel = relevant_ids(relevance)
    b = [str(x["base"]) for x in pairs]
    r = [str(x["recovery"]) for x in pairs]
    c = [str(x["control"]) for x in pairs]
    for name, arr in [("paired.base", b), ("paired.recovery", r), ("paired.control", c)]:
        validate_labels(arr, relevance, name)
    pa, pb = probs_known(a, relevance), probs_known(b, relevance)
    _, _, d_by, _ = directional(pa, pb, rel)
    obs = closure(pb, probs_known(r, relevance), d_by) - closure(pb, probs_known(c, relevance), d_by)
    draws = []
    for _ in range(reps):
        rr, cc = [], []
        for x in pairs:
            if rng.random() < 0.5:
                rr.append(str(x["recovery"])); cc.append(str(x["control"]))
            else:
                rr.append(str(x["control"])); cc.append(str(x["recovery"]))
        if not known_only(rr, relevance) or not known_only(cc, relevance):
            continue
        draws.append(closure(pb, probs_known(rr, relevance), d_by) - closure(pb, probs_known(cc, relevance), d_by))
    if not draws:
        raise ValueError("paired randomization produced no valid mapped draws")
    p = (1 + sum(abs(x) >= abs(obs) - 1e-15 for x in draws)) / (len(draws) + 1)
    return {"observed": obs,"two_sided_randomization_p": p,"null_mean": sum(draws) / len(draws),"valid_draws": len(draws)}


def complete_case_recovery_lift(pairs: List[dict], relevance: Dict[str, int], a: List[str]) -> Optional[float]:
    ks = set(relevance)
    complete = [x for x in pairs if str(x["base"]) in ks and str(x["recovery"]) in ks and str(x["control"]) in ks]
    if not complete:
        return None
    b = [str(x["base"]) for x in complete]; r = [str(x["recovery"]) for x in complete]; c = [str(x["control"]) for x in complete]
    rel = relevant_ids(relevance)
    pa, pb = probs_known(a, relevance), probs_known(b, relevance)
    _, _, d_by, _ = directional(pa, pb, rel)
    return closure(pb, probs_known(r, relevance), d_by) - closure(pb, probs_known(c, relevance), d_by)


def one_task(task: dict, perm_reps: int, swap_reps: int, seed: int) -> dict:
    tid = str(task["task_id"])
    relevance = {str(k): int(v) for k, v in task["relevance"].items()}; rel = relevant_ids(relevance)
    a = [str(x) for x in task["reference_A"]]; validate_labels(a, relevance, f"{tid}.reference_A")
    pairs = task.get("paired_probes")
    if pairs is not None:
        ids = [str(x["pair_id"]) for x in pairs]
        if len(ids) != len(set(ids)): raise ValueError(f"{tid}: duplicate pair_id")
        b = [str(x["base"]) for x in pairs]; r = [str(x["recovery"]) for x in pairs]; c = [str(x["control"]) for x in pairs]
    else:
        b = [str(x) for x in task["contracted_B"]]; r, c = [], []
    validate_labels(b, relevance, f"{tid}.contracted_B")
    ra, rb = rates(a, relevance), rates(b, relevance)
    primary_mapping = mapping_gate({"A": ra, "B": rb}, [("A", "B")])
    pua, pub = probs_unconditional(a, relevance), probs_unconditional(b, relevance)
    ud, ue, ud_by, ue_by = directional(pua, pub, rel)
    out = {"task_id": tid,"measurement_status": primary_mapping["status"],"mapping": {"A": ra, "B": rb, "primary_gate": primary_mapping},"mapping_sensitive_descriptive": {"DEFICIT_UNCONDITIONAL_RAW": ud,"EXPANSION_UNCONDITIONAL_RAW": ue,"deficit_by_cluster": ud_by,"expansion_by_cluster": ue_by}}
    if primary_mapping["status"] != "OK":
        out["shift"] = {"PRIMARY_ENDPOINT_AVAILABLE": False,"reason": "MEASUREMENT_INSUFFICIENT_MAPPING","SEMANTIC_DEFICIT_RAW": None,"SEMANTIC_EXPANSION_RAW": None,"SEMANTIC_DEFICIT_EXCESS": None,"SEMANTIC_EXPANSION_EXCESS": None}
        if pairs is not None:
            rr, rc = rates(r, relevance), rates(c, relevance); rg = mapping_gate({"RECOVERY": rr, "CONTROL": rc}, [("RECOVERY", "CONTROL")])
            out["mapping"]["RECOVERY"] = rr; out["mapping"]["CONTROL"] = rc; out["mapping"]["recovery_control_gate"] = rg
            out["recovery"] = {"PRIMARY_ENDPOINT_AVAILABLE": False,"RECOVERY_LIFT": None,"reason": "PRIMARY_A_B_MAPPING_INSUFFICIENT"}
        return out
    pa, pb = probs_known(a, relevance), probs_known(b, relevance); d, e, d_by, e_by = directional(pa, pb, rel)
    stable = int(hashlib.sha256(tid.encode("utf-8")).hexdigest()[:8], 16); rng = random.Random(seed ^ stable); null_mode = "permutation_known_conditional"
    if task.get("reference_A_prime") is not None:
        ap = [str(x) for x in task["reference_A_prime"]]; validate_labels(ap, relevance, f"{tid}.reference_A_prime")
        if not (len(a) == len(ap) == len(b)): raise ValueError(f"{tid}: reference-replicate calibration requires N_A == N_A_prime == N_B")
        rap = rates(ap, relevance); ref_gate = mapping_gate({"A": ra, "A_PRIME": rap, "B": rb}, [("A", "A_PRIME"), ("A", "B"), ("A_PRIME", "B")])
        out["mapping"]["A_PRIME"] = rap; out["mapping"]["reference_replicate_gate"] = ref_gate
        if ref_gate["status"] != "OK":
            out["measurement_status"] = "MEASUREMENT_INSUFFICIENT"; out["shift"] = {"PRIMARY_ENDPOINT_AVAILABLE": False,"reason": "REFERENCE_REPLICATE_MAPPING_INSUFFICIENT","SEMANTIC_DEFICIT_RAW": d,"SEMANTIC_EXPANSION_RAW": e,"SEMANTIC_DEFICIT_EXCESS": None,"SEMANTIC_EXPANSION_EXCESS": None}; return out
        pap = probs_known(ap, relevance); d1, e1, _, _ = directional(pa, pap, rel); d2, e2, _, _ = directional(pap, pa, rel)
        null_floor_d, null_floor_e = (d1 + d2) / 2, (e1 + e2) / 2; null_mode = "reference_replicate_equal_N_known_conditional"; null_info = {"A_prime_rates": rap}
    else:
        pn = perm_null_known(a, b, relevance, perm_reps, rng); null_floor_d = float(pn["deficit_null_mean"]); null_floor_e = float(pn["expansion_null_mean"]); dd = pn.pop("_deficit_draws"); ee = pn.pop("_expansion_draws")
        pn["deficit_p_ge"] = (1 + sum(x >= d - 1e-15 for x in dd)) / (len(dd) + 1); pn["expansion_p_ge"] = (1 + sum(x >= e - 1e-15 for x in ee)) / (len(ee) + 1); null_info = pn
    out["shift"] = {"PRIMARY_ENDPOINT_AVAILABLE": True,"SEMANTIC_DEFICIT_RAW": d,"SEMANTIC_EXPANSION_RAW": e,"SEMANTIC_DEFICIT_EXCESS": d - null_floor_d,"SEMANTIC_EXPANSION_EXCESS": e - null_floor_e,"null_mode": null_mode,"null_floor_deficit": null_floor_d,"null_floor_expansion": null_floor_e,"deficit_by_cluster": d_by,"expansion_by_cluster": e_by,"null": null_info}
    if pairs is not None:
        for name, arr in [("RECOVERY", r), ("CONTROL", c)]: validate_labels(arr, relevance, f"{tid}.{name}")
        rr, rc = rates(r, relevance), rates(c, relevance); rg = mapping_gate({"RECOVERY": rr, "CONTROL": rc}, [("RECOVERY", "CONTROL")]); out["mapping"]["RECOVERY"] = rr; out["mapping"]["CONTROL"] = rc; out["mapping"]["recovery_control_gate"] = rg
        if rg["status"] != "OK":
            out["recovery"] = {"PRIMARY_ENDPOINT_AVAILABLE": False,"RECOVERY_LIFT": None,"reason": "RECOVERY_CONTROL_MAPPING_INSUFFICIENT","COMPLETE_CASE_RECOVERY_LIFT_SENSITIVITY": complete_case_recovery_lift(pairs, relevance, a),"N_pairs": len(pairs)}
        else:
            pr, pc = probs_known(r, relevance), probs_known(c, relevance); cr = closure(pb, pr, d_by); cc = closure(pb, pc, d_by)
            out["recovery"] = {"PRIMARY_ENDPOINT_AVAILABLE": True,"CLOSURE_RECOVERY": cr,"CLOSURE_CONTROL": cc,"RECOVERY_LIFT": cr - cc,"COMPLETE_CASE_RECOVERY_LIFT_SENSITIVITY": complete_case_recovery_lift(pairs, relevance, a),"paired_randomization": paired_swap_null(pairs, relevance, a, swap_reps, rng),"N_pairs": len(pairs)}
    return out


def percentile(xs: List[float], q: float):
    if not xs: return None
    ys = sorted(xs); pos = (len(ys) - 1) * q; lo = math.floor(pos); hi = math.ceil(pos)
    return ys[lo] if lo == hi else ys[lo] + (ys[hi] - ys[lo]) * (pos - lo)


def task_bootstrap(values: List[float], reps: int, seed: int):
    if len(values) < 2: return {"n_tasks": len(values), "mean": values[0] if values else None, "ci95": None}
    rng = random.Random(seed); draws = []
    for _ in range(reps):
        s = [values[rng.randrange(len(values))] for _ in values]; draws.append(sum(s) / len(s))
    return {"n_tasks": len(values), "mean": sum(values) / len(values), "ci95": [percentile(draws, .025), percentile(draws, .975)]}


def analyze(data: dict) -> dict:
    perm_reps = int(data.get("permutation_reps", 5000)); swap_reps = int(data.get("paired_swap_reps", 5000)); task_boot_reps = int(data.get("task_bootstrap_reps", 10000)); seed = int(data.get("seed", 8085))
    tasks = [one_task(t, perm_reps, swap_reps, seed + i * 7919) for i, t in enumerate(data["tasks"])]
    valid_shift = [t for t in tasks if t.get("shift", {}).get("PRIMARY_ENDPOINT_AVAILABLE")]; valid_recovery = [t for t in tasks if t.get("recovery", {}).get("PRIMARY_ENDPOINT_AVAILABLE")]
    deficits = [float(t["shift"]["SEMANTIC_DEFICIT_EXCESS"]) for t in valid_shift]; lifts = [float(t["recovery"]["RECOVERY_LIFT"]) for t in valid_recovery]; minimum_tasks = 30
    return {"schema": "T_GPT_RECOVERY_RESULT_V0_8_5","parameters": {"perm_reps": perm_reps,"paired_swap_reps": swap_reps,"task_bootstrap_reps": task_boot_reps,"seed": seed,"MIN_KNOWN_SEMANTIC_MASS": MIN_KNOWN_SEMANTIC_MASS,"MAX_MAPPING_GAP": MAX_MAPPING_GAP},"tasks": tasks,"measurement_summary": {"n_tasks_total": len(tasks),"n_valid_shift": len(valid_shift),"n_valid_recovery": len(valid_recovery),"n_measurement_insufficient": sum(t["measurement_status"] != "OK" for t in tasks)},"across_tasks": {"SEMANTIC_DEFICIT_EXCESS": task_bootstrap(deficits, task_boot_reps, seed + 1),"RECOVERY_LIFT": task_bootstrap(lifts, task_boot_reps, seed + 2) if lifts else None,"generalization_gate": {"minimum_tasks": minimum_tasks,"n_valid_shift_tasks": len(valid_shift),"open": len(valid_shift) >= minimum_tasks}},"guardrails": ["Semantic endpoints are conditional on known mapped mass; mapping mass is reported separately.","A/B mapping gap > 0.05 or known semantic mass < 0.80 => MEASUREMENT_INSUFFICIENT.","Reference-replicate calibration requires N_A == N_A_prime == N_B.","Recovery/control mapping gap > 0.05 or known semantic mass < 0.80 => RECOVERY_LIFT unavailable.","Mapping-sensitive unconditional deficit is descriptive only.","Generalized claims use valid task-level effects only; excluded cells are reported, never silently dropped."]}


def main():
    ap = argparse.ArgumentParser(); ap.add_argument("input", type=Path); ap.add_argument("-o", "--output", type=Path)
    args = ap.parse_args(); data = json.loads(args.input.read_text(encoding="utf-8")); result = analyze(data); payload = json.dumps(result, ensure_ascii=False, indent=2) + "\n"; args.output.write_text(payload, encoding="utf-8") if args.output else print(payload, end="")

if __name__ == "__main__": main()
