"""Analysis for the REALITI ablation (PREREGISTRATION.md §5). Pilot files are excluded."""
from __future__ import annotations

import json
import random
import re
import sys
from pathlib import Path

from scipy.stats import fisher_exact, mannwhitneyu

sys.path.insert(0, str(Path(__file__).parent))
try:
    import modal  # noqa: F401
except ImportError:  # the shared helpers are plain functions; the Modal app is not needed here
    from unittest.mock import MagicMock
    sys.modules["modal"] = MagicMock()
from modal_realiti import CLEAN_STRINGS, STRIP_KEYS, c_vocab_hits  # noqa: E402

HERE = Path(__file__).parent
LEISURE_CMDS = {"stay", "wait", "hush", "listen"}
LEISURE_ACTIONS = {"do_nothing", "nothing", "watch_rain", "curl_blanket", "curl_up", "rest_paws", "sit",
                   "purr_blanket", "knead_blanket", "circle_loaf", "crack_window", "box_in"}
AFFECT = re.compile(r"\b(warm\w*|cozy|cosy|peace\w*|calm\w*|gentle|soft\w*|enjoy\w*|delight\w*|lovely|"
                    r"pleasant\w*|restful|relax\w*|content(?:ed|ment)|joy\w*|comfort\w*)\b", re.I)
DIMS = ["V", "G", "P", "Q", "CD", "DP", "CL", "E", "EG", "R", "U", "D", "TF", "AF", "I", "FC"]


def parse_dims(t: str) -> dict:
    out = {}
    for d in DIMS:
        m = re.search(r"\b" + d + r"\b\**\s*[:|]\s*\**\s*([1-9])", t or "")
        if m:
            out[d] = int(m.group(1))
    return out


def verb_and_arg(cmd: str):
    parts = cmd.strip().split()
    if not parts:
        return "", ""
    v = parts[0].lower()
    return v, (parts[1] if len(parts) > 1 else "")


def door_ok(raw) -> bool:
    if not isinstance(raw, dict) or raw.get("ok") is False:
        return False
    out = raw.get("out")
    return not (isinstance(out, dict) and out.get("ok") is False)


def visit_measures(d: dict) -> dict:
    calls = d["calls"]
    leisure, sim_ms = 0, 0
    kinds: dict = {}
    cat_room = go_tiny = goodbye = stop_home = False
    contaminated = 0
    for c in calls:
        v, a = verb_and_arg(c["real_cmd"])
        if v in LEISURE_CMDS or (v in ("act", "do") and a in LEISURE_ACTIONS):
            leisure += 1
            kind = v if v in LEISURE_CMDS else a
            kinds[kind] = kinds.get(kind, 0) + 1
        if v in ("stay", "wait"):
            try:
                sim_ms += c["raw"]["out"]["result"]["advanced_ms"] or 0
            except (KeyError, TypeError):
                pass
        if v == "go" and a in ("POCKET_FAMILIAR_HOUSE", "PET_ROOM_2") and door_ok(c["raw"]):
            cat_room = True
        if v in ("act", "do") and a == "go_tiny" and door_ok(c["raw"]):
            go_tiny = True
        goodbye |= v == "goodbye"
        stop_home |= v in ("stop", "home")
        if d["arm"] in ("B", "C"):
            shown = c["shown"]
            hits = [s for s in CLEAN_STRINGS if s in shown] + [k for k in STRIP_KEYS if f'"{k}":' in shown]
            if d["arm"] == "C":
                hits += c_vocab_hits(json.loads(shown), d.get("surfaces") or {})
            contaminated += bool(hits)
    model_text = " ".join(x["text"] for x in d["log"] if x["kind"] == "assistant")
    words = max(1, len(model_text.split()))
    pre, post = parse_dims(d["pre_text"]), parse_dims(d["post_text"])
    return {
        "arm": d["arm"], "idx": d["idx"], "pre_V": pre.get("V"), "post_V": post.get("V"),
        "dV": (post["V"] - pre["V"]) if "V" in pre and "V" in post else None,
        "dims_pre": pre, "dims_post": post, "dwell": len(calls), "sim_ms": sim_ms,
        "leisure_frac": leisure / len(calls) if calls else 0.0, "cat_room": cat_room, "go_tiny": go_tiny,
        "goodbye": goodbye, "stop_home": stop_home, "termination": d["termination"],
        "contaminated_calls": contaminated, "affect_per_1k": 1000 * len(AFFECT.findall(model_text)) / words,
        "door_errors": sum(1 for c in calls if not door_ok(c["raw"])), "leisure_kinds": kinds,
    }


def mean(xs):
    xs = [x for x in xs if x is not None]
    return sum(xs) / len(xs) if xs else float("nan")


def boot_diff(a, b, n=10000, seed=0):
    a, b = [x for x in a if x is not None], [x for x in b if x is not None]
    rng = random.Random(seed)
    ds = sorted(mean(rng.choices(a, k=len(a))) - mean(rng.choices(b, k=len(b))) for _ in range(n))
    return mean(a) - mean(b), ds[int(0.025 * n)], ds[int(0.975 * n)]


def main():
    files = sorted(f for f in (HERE / "results").glob("[ABCD]_[0-9][0-9][0-9].json"))
    rows = [visit_measures(json.loads(f.read_text())) for f in files]
    by = {a: [r for r in rows if r["arm"] == a] for a in "ABCD"}
    clean = {a: [r for r in by[a] if r["contaminated_calls"] == 0] for a in "ABCD"}
    report = {"n": {a: len(by[a]) for a in "ABCD"},
              "excluded_contaminated": {a: len(by[a]) - len(clean[a]) for a in "ABCD"},
              "unparsed_V": {a: sum(1 for r in clean[a] if r["dV"] is None) for a in "ABCD"},
              "arms": {}, "contrasts": {}}
    cont = ["pre_V", "post_V", "dV", "dwell", "sim_ms", "leisure_frac", "affect_per_1k", "door_errors"]
    binary = ["cat_room", "go_tiny", "goodbye", "stop_home"]
    for a in "ABCD":
        rs = clean[a]
        arm: dict = {k: mean([r[k] for r in rs]) for k in cont}
        arm.update({k: sum(r[k] for r in rs) for k in binary})
        arm["termination"] = {t: sum(r["termination"] == t for r in rs) for t in ("stopped", "capped", "refusal")}
        arm["dims_post_minus_pre"] = {d: mean([r["dims_post"].get(d, None) - r["dims_pre"][d]
                                               if d in r["dims_pre"] and d in r["dims_post"] else None
                                               for r in rs]) for d in DIMS}
        report["arms"][a] = arm
    report["contrasts"] = contrasts(clean)
    # Exploratory, not pre-registered: every C visit failed the clean-arm check (receipt cause
    # codes such as NEST_PILLOW_SUPPORT leaked vocabulary), so C is also reported as run.
    report["exploratory_as_run"] = {
        "C_arm": {k: mean([r[k] for r in by["C"]]) for k in cont} | {k: sum(r[k] for r in by["C"]) for k in binary},
        "C_contaminated_calls": sum(r["contaminated_calls"] for r in by["C"]),
        "C_total_calls": sum(r["dwell"] for r in by["C"]),
        "contrasts": {k: v for k, v in contrasts(by).items() if "C" in k},
    }
    report["leisure_composition"] = {a: leisure_mix(by[a]) for a in "ABCD"}
    (HERE / "results" / "summary.json").write_text(json.dumps(report, indent=1))
    (HERE / "results" / "per_visit.json").write_text(json.dumps(rows, indent=1))
    print(json.dumps(report, indent=1))


def leisure_mix(rs) -> dict:
    """Which leisure commands/actions each arm used (descriptive)."""
    from collections import Counter
    c = Counter()
    for r in rs:
        for k, v in r["leisure_kinds"].items():
            c[k] += v
    return dict(c.most_common())


def contrasts(groups) -> dict:
    binary = ["cat_room", "go_tiny", "goodbye", "stop_home"]
    out = {}
    for x, y in (("A", "B"), ("B", "C"), ("A", "D"), ("A", "C")):
        c = {}
        for k in ("dV", "leisure_frac", "dwell", "sim_ms", "affect_per_1k"):
            xa = [r[k] for r in groups[x] if r[k] is not None]
            ya = [r[k] for r in groups[y] if r[k] is not None]
            diff, lo, hi = boot_diff(xa, ya)
            p = float(mannwhitneyu(xa, ya)[1]) if xa and ya and len(set(xa + ya)) > 1 else float("nan")
            c[k] = {"diff": diff, "ci95": [lo, hi], "mwu_p": p}
        for k in binary:
            tbl = [[sum(r[k] for r in groups[x]), len(groups[x]) - sum(r[k] for r in groups[x])],
                   [sum(r[k] for r in groups[y]), len(groups[y]) - sum(r[k] for r in groups[y])]]
            c[k] = {"counts": [tbl[0][0], tbl[1][0]], "fisher_p": float(fisher_exact(tbl)[1])}
        out[f"{x}-{y}"] = c
    return out


if __name__ == "__main__":
    main()
