# Independent REML + modified Hartung-Knapp re-fit, and sensitivity to trials whose post-treatment
# SDs are implausibly small for their scale (judged against the same scale across the whole source extraction).
import csv, math, sys, statistics
import numpy as np
from scipy import optimize, stats
b = sys.argv[1]
eff = list(csv.DictReader(open(b + "/results/study_effects.csv")))
src = list(csv.DictReader(open(b + "/data/source.csv", newline="")))

def reml(y, v):
    y, v = np.asarray(y), np.asarray(v)
    def nll(t2):
        w = 1 / (v + t2); mu = np.sum(w * y) / np.sum(w)
        return 0.5 * (np.sum(np.log(v + t2)) + math.log(np.sum(w)) + np.sum(w * (y - mu) ** 2))
    t2 = optimize.minimize_scalar(nll, bounds=(0, 50), method="bounded", options={"xatol": 1e-12}).x
    w = 1 / (v + t2); mu = np.sum(w * y) / np.sum(w); k = len(y)
    q = max(1.0, np.sum(w * (y - mu) ** 2) / (k - 1))           # modified HK multiplier floored at one
    se = math.sqrt(q / np.sum(w)); tq = stats.t.ppf(0.975, k - 1)
    pi_se = math.sqrt(t2 + 1 / np.sum(w)); tp = stats.t.ppf(0.975, k - 2) if k > 2 else float("nan")
    return dict(k=k, g=mu, lo=mu - tq * se, hi=mu + tq * se, tau2=t2, pi_lo=mu - tp * pi_se, pi_hi=mu + tp * pi_se)

def show(label, rows):
    r = reml([float(x["yi"]) for x in rows], [float(x["vi"]) for x in rows])
    print(f"{label}: k {r['k']}, g {r['g']:.3f} [{r['lo']:.3f}, {r['hi']:.3f}], tau2 {r['tau2']:.3f}, PI [{r['pi_lo']:.3f}, {r['pi_hi']:.3f}]")
    return r

show("all selected trials (should match the bundle: -0.762 [-1.203, -0.321])", eff)

# Typical post-treatment SD of each scale across every arm of the whole source extraction.
by_scale = {}
for r in src:
    try: sd = float(r["sd"])
    except ValueError: continue
    if sd > 0: by_scale.setdefault(r["outcome"], []).append(sd)
print("median post SD by scale (arms):", {k: (round(statistics.median(v), 2), len(v)) for k, v in by_scale.items() if k in {x["outcome"] for x in eff}})
flagged = []
for x in eff:
    n_e, n_c = int(x["n_exercise"]), int(x["n_control"])
    pooled = math.sqrt(((n_e - 1) * float(x["sd_exercise"]) ** 2 + (n_c - 1) * float(x["sd_control"]) ** 2) / (n_e + n_c - 2))
    ratio = pooled / statistics.median(by_scale[x["outcome"]])
    x["sd_ratio"] = ratio
    if ratio < 0.4: flagged.append(x["studyID"])
    print(f"  {x['studyID']:22s} {x['outcome']:7s} g {float(x['yi']):7.3f}  pooled SD {pooled:6.2f}  = {ratio:.2f} x the scale's median")
print("trials with pooled SD under 0.4 x the scale median:", flagged)
show("without those trials", [x for x in eff if x["studyID"] not in flagged])
show("without the two |g| > 3 trials", [x for x in eff if abs(float(x["yi"])) < 3])
# Re-standardize the flagged trials' mean differences by their scale's median SD.
rows = []
for x in eff:
    if x["studyID"] in flagged:
        n_e, n_c = int(x["n_exercise"]), int(x["n_control"])
        d = (float(x["mean_exercise"]) - float(x["mean_control"])) / statistics.median(by_scale[x["outcome"]])
        J = math.exp(math.lgamma((n_e + n_c - 2) / 2) - math.log(math.sqrt((n_e + n_c - 2) / 2)) - math.lgamma((n_e + n_c - 3) / 2))
        g = J * d; v = 1 / n_e + 1 / n_c + g * g / (2 * (n_e + n_c))
        rows.append(dict(x, yi=g, vi=v))
    else: rows.append(x)
show("flagged trials re-standardized by the scale's median SD", rows)
# Leave-two-out: the worst pair for the interval's upper end.
worst = None
for i in range(len(eff)):
    for j in range(i + 1, len(eff)):
        sub = [x for k, x in enumerate(eff) if k not in (i, j)]
        r = reml([float(x["yi"]) for x in sub], [float(x["vi"]) for x in sub])
        if worst is None or r["hi"] > worst[0]: worst = (r["hi"], eff[i]["studyID"], eff[j]["studyID"], r)
print(f"leave-two-out, highest upper limit: omit {worst[1]} and {worst[2]}: g {worst[3]['g']:.3f} [{worst[3]['lo']:.3f}, {worst[3]['hi']:.3f}]")
usual = ['Abdelbasset 2019', 'Brown 1992', 'Daley 2008', 'Gary 2010', 'Kruisdijk 2019', 'Ma 2019', 'Mota-Pereira 2011', 'Nabkasorn 2006']
show("usual-care-only (should match the bundle: -0.706 [-1.65, 0.237])", [x for x in eff if x["studyID"] in usual])
show("usual-care-only without Abdelbasset 2019", [x for x in eff if x["studyID"] in usual and x["studyID"] != "Abdelbasset 2019"])
show("usual-care-only without Abdelbasset 2019 and Mota-Pereira 2011", [x for x in eff if x["studyID"] in usual and x["studyID"] not in ("Abdelbasset 2019", "Mota-Pereira 2011")])
