"""Focused, preregistered direct-comparison audit of a public trial extraction.

Reproduction is offline. Source outcomes are never repaired or reclassified.
Each study contributes one contrast, with disjoint source arms combined first.
"""
import csv
import hashlib
import json
import math
from collections import defaultdict
from pathlib import Path

import numpy as np
from scipy.optimize import minimize_scalar, brentq
from scipy.stats import t, norm, chi2
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt

ROOT = Path(__file__).resolve().parents[1]
OUT = ROOT / "results"
WALK = "Walking / Jogging"
CONTROLS = {"Educational", "Social", "Social or educational control",
            "Usual care", "Placebo pill", "Stretching"}
RISK = {"Low risk": 0, "Unclear risk": 1, "High risk": 2}
DOMAINS = ["overall", "random_sequence_generation_selection_bias",
           "allocation_concealment_selection_bias", "blinded_outcome_assessor"]

def declared_precision(value):
    """Portable JSON precision, preserving very small nonzero probabilities.

    Aggregated evidence objects must match exactly under the reference harness.
    Statistical summaries are already display-rounded; other numeric values keep
    twelve significant decimal digits, above the precision relevant to inference.
    The unrounded trial contrast CSV and independent R fixtures are also supplied.
    """
    if isinstance(value,float):return float(format(value,'.12g'))
    if isinstance(value,list):return [declared_precision(v) for v in value]
    if isinstance(value,dict):return {k:declared_precision(v) for k,v in value.items()}
    return value

def number(row, field):
    value = float(row[field])
    if not math.isfinite(value):
        raise ValueError("Nonfinite input: " + field)
    return value

def n(row, field="n"):
    value = number(row, field)
    if value <= 1 or value != int(value):
        raise ValueError("Invalid arm sample size")
    return int(value)

def combined(rows, mean="mean", sd="sd", count="n"):
    ns = [n(r, count) for r in rows]
    means, sds = [number(r, mean) for r in rows], [number(r, sd) for r in rows]
    if any(s <= 0 for s in sds):
        raise ValueError("Nonpositive source standard deviation")
    total = sum(ns)
    average = sum(ni*mi for ni,mi in zip(ns,means))/total
    variance = sum((ni-1)*si*si + ni*(mi-average)**2
                   for ni,mi,si in zip(ns,means,sds))/(total-1)
    # Independent second-moment calculation of the combined group variance.
    second = (sum((ni-1)*si*si+ni*mi*mi for ni,mi,si in zip(ns,means,sds))
              - total*average*average)/(total-1)
    assert abs(variance-second) < 1e-8*max(1,variance)
    return total, average, math.sqrt(variance)

def small_sample(df):
    return math.exp(math.lgamma(df/2)-math.lgamma((df-1)/2)-0.5*math.log(df/2))

def effect(exercise, control):
    ne, me, se = combined(exercise)
    nc, mc, sc = combined(control)
    df = ne+nc-2
    pooled = math.sqrt(((ne-1)*se*se+(nc-1)*sc*sc)/df)
    g = small_sample(df)*(me-mc)/pooled
    variance = 1/ne + 1/nc + g*g/(2*(ne+nc))
    assert variance > 0
    return {"yi":g,"vi":variance,"n_exercise":ne,"n_control":nc,
            "mean_exercise":me,"mean_control":mc,"sd_exercise":se,"sd_control":sc}

def select(rows, controls=CONTROLS, prefer_clinician=True):
    post = defaultdict(list)
    for r in rows:
        if number(r,"weeks_from_end_of_treatment_to_measurement")==0:
            post[r["studyID"]].append(r)
    included, flow = [], []
    for study in sorted({r["studyID"] for r in rows}):
        relevant = [r for r in post[study] if r["trt"]==WALK or r["trt"] in controls]
        walk = [r for r in relevant if r["trt"]==WALK]
        ctrl = [r for r in relevant if r["trt"] in controls]
        if not walk or not ctrl:
            flow.append({"studyID":study,"included":False,"reason":"No immediate target and eligible control arms"})
            continue
        # Validate only records that can enter this direct-comparison audit.
        # Non-target duplicate keys are preserved and reported, never averaged.
        keys=[(r['arm_number'],r['outcome']) for r in relevant]
        if len(keys)!=len(set(keys)):
            raise ValueError("Duplicate eligible study/arm/measure/time")
        arms=defaultdict(set)
        for r in relevant:arms[r["arm_number"]].add(r["outcome"])
        common=set.intersection(*arms.values())
        if not common:
            flow.append({"studyID":study,"included":False,"reason":"No common immediate measure across all eligible arms"})
            continue
        def rank(measure):
            measured=[r for r in relevant if r["outcome"]==measure]
            clinician=all(r["reported"]=="Clinician" for r in measured)
            return (0 if prefer_clinician and clinician else 1,measure)
        measure=min(common,key=rank)
        chosen=[r for r in relevant if r["outcome"]==measure]
        e=[r for r in chosen if r["trt"]==WALK]
        c=[r for r in chosen if r["trt"] in controls]
        try:
            estimate=effect(e,c)
        except ValueError as error:
            flow.append({"studyID":study,"included":False,"reason":str(error)})
            continue
        assert len({r["arm_number"] for r in chosen})==len(chosen)
        risk={domain:max((r[domain] for r in chosen),key=lambda v:RISK[v]) for domain in DOMAINS}
        included.append({"studyID":study,"outcome":measure,**estimate,"risk":risk,
                         "source_row_ids":[r["row_id"] for r in chosen],
                         "source_csv_lines":[r['_source_csv_line'] for r in chosen],
                         "exercise_rows":e,"control_rows":c})
        flow.append({"studyID":study,"included":True,"reason":"Eligible", "outcome":measure,
                     "source_row_ids":[r["row_id"] for r in chosen]})
    assert len({r["studyID"] for r in included})==len(included)
    return included, flow

def reml(y,v):
    def objective(tau):
        w=1/(v+tau)
        mu=np.dot(w,y)/w.sum()
        return float(np.log(v+tau).sum()+math.log(w.sum())+np.dot(w,(y-mu)**2))
    upper=max(1,float(np.var(y))*4)
    for _ in range(20):
        fit=minimize_scalar(objective,bounds=(0,upper),method="bounded",options={"xatol":1e-12})
        if fit.x < upper*0.95:break
        upper*=4
    else:raise ValueError("REML search did not bracket optimum")
    tau=0.0 if objective(0)<=fit.fun else float(fit.x)
    # An independent derivative-root check, not the likelihood minimizer.
    def score(x):
        w=1/(v+x);mu=np.dot(w,y)/w.sum()
        return float(np.dot(w*w,(y-mu)**2)-w.sum()+np.dot(w,w)/w.sum())
    if score(0)<=0:root=0.0
    else:
        end=upper
        while score(end)>0:end*=4
        root=float(brentq(score,0,end,xtol=1e-13))
    assert abs(tau-root)<1e-6
    return root

def fit(items, method="REML", interval="modified_HK"):
    k=len(items)
    if k<2:return {"k":k,"status":"insufficient studies; no pooled estimate"}
    y=np.array([r["yi"] for r in items]);v=np.array([r["vi"] for r in items])
    assert np.isfinite(y).all() and np.isfinite(v).all() and (v>0).all()
    fixed=1/v;fixed_mu=np.dot(fixed,y)/fixed.sum()
    q=float(np.dot(fixed,(y-fixed_mu)**2))
    if method=="REML":tau=reml(y,v)
    elif method=="DL":tau=max(0,(q-(k-1))/(fixed.sum()-np.dot(fixed,fixed)/fixed.sum()))
    elif method=="fixed":tau=0
    else:raise ValueError("Unknown model")
    w=1/(v+tau);mu=float(np.dot(w,y)/w.sum())
    hk=float(np.dot(w,(y-mu)**2)/(k-1))
    multiplier=max(1,hk) if interval=="modified_HK" else hk if interval=="HK" else 1
    variance=float(multiplier/w.sum());se=math.sqrt(variance)
    dist=norm if interval=="normal" else t(k-1)
    quantile=float(dist.ppf(.975));p=float(2*dist.sf(abs(mu/se)))
    prediction_quantile=float(t(k-2).ppf(.975)) if k>2 else None
    prediction_width=prediction_quantile*math.sqrt(tau+variance) if k>2 else None
    return {"status":"estimated","k":k,"model":method,"interval_method":interval,
            "n_exercise":sum(r["n_exercise"] for r in items),"n_control":sum(r["n_control"] for r in items),
            "n_total":sum(r["n_exercise"]+r["n_control"] for r in items),
            "g":round(mu,3),"se_g":round(se,3),
            "ci95_low_g":round(mu-quantile*se,3),"ci95_high_g":round(mu+quantile*se,3),
            "p_two_sided":p,"p_display":"<0.001" if p<.001 else format(p,'.3g'),"tau2":round(float(tau),6),"Q":round(q,6),
            "Q_p_value":float(chi2.sf(q,k-1)),
            "Q_p_display":"<0.001" if chi2.sf(q,k-1)<.001 else format(chi2.sf(q,k-1),'.3g'),
            "I2_percent":round(max(0,(q-k+1)/q)*100,2) if q else 0,
            "hk_multiplier_raw":round(hk,6),
            "prediction95_low_g":round(mu-prediction_width,3) if k>2 else None,
            "prediction95_high_g":round(mu+prediction_width,3) if k>2 else None,
            "interval_excludes_zero":bool(mu+quantile*se<0 or mu-quantile*se>0)}

def changed_contrasts(items, correlation=None):
    results=[]
    for item in items:
        estimate=dict(item);means=[];variances=[]
        for name in ["exercise_rows","control_rows"]:
            group=item[name];counts=[n(r) for r in group];weights=np.array(counts)/sum(counts)
            gs=np.array([number(r,"smd") for r in group])
            variances_arm=[number(r,"se_smd")**2 if correlation is None
                           else 2*(1-correlation)/n(r)+number(r,"smd")**2/(2*n(r)) for r in group]
            means.append(float(np.dot(weights,gs)))
            variances.append(float(np.dot(weights*weights,variances_arm)))
        estimate["yi"]=means[0]-means[1];estimate["vi"]=sum(variances)
        results.append(estimate)
    return results

def validate_reference(selected,primary,sensitivity,leave):
    reference=json.loads((ROOT/'data/metafor-reference.json').read_text())
    expected=list(reference['studies'].values())
    assert [(r['studyID'],r['outcome']) for r in selected]==[(r['studyID'],r['outcome']) for r in expected]
    effects=max(abs(a['yi']-b['yi']) for a,b in zip(selected,expected))
    variances=max(abs(a['vi']-b['vi']) for a,b in zip(selected,expected))
    assert effects<1e-10 and variances<1e-10
    estimates=[(primary,reference['primary'])]
    estimates.extend((sensitivity[key],value) for key,value in reference['sensitivity'].items())
    estimates.extend((a,b) for a,b in zip(leave,reference['leave_one_out']))
    for actual,expected in estimates:
        assert actual['k']==expected['k']
        for field in ['g','se_g','ci95_low_g','ci95_high_g']:
            assert abs(actual[field]-expected[field])<=0.000501
        assert abs(actual['tau2']-expected['tau2'])<=0.000001
    return {'status':'matched within declared display precision',
            'software':reference['software'],'fits_compared':len(estimates),
            'max_study_g_difference':effects,'max_study_variance_difference':variances}

def plot(primary, selected, leave_one_out):
    plt.rcParams.update({"font.family":"DejaVu Sans","font.size":10})
    fig,ax=plt.subplots(figsize=(10,11),layout="constrained")
    ys=np.arange(len(selected),0,-1)
    for pos,item in zip(ys,selected):
        g=item["yi"];width=norm.ppf(.975)*math.sqrt(item["vi"])
        ax.errorbar(g,pos,xerr=width,fmt="o",color="#166676",capsize=3,markersize=4)
    ax.errorbar(primary["g"],0,xerr=[[primary["g"]-primary["ci95_low_g"]],[primary["ci95_high_g"]-primary["g"]]],fmt="D",color="#7c2d55",capsize=5)
    ax.axvline(0,color="#777",linewidth=1)
    ax.set_yticks([*ys,0],[f"{r['studyID']} ({r['outcome']}; n={r['n_exercise']+r['n_control']})" for r in selected]+["Pooled REML, modified Hartung-Knapp"])
    ax.set_xlabel("Hedges g at post-treatment: walking/jogging minus active control\nNegative values favor walking/jogging")
    ax.set_title("Direct comparisons from the fixed Noetel et al. dataset\nStudy intervals: normal; pooled interval: modified Hartung-Knapp")
    ax.grid(axis="x",alpha=.2);fig.savefig(OUT/'forest.png',dpi=160);plt.close(fig)
    fig,ax=plt.subplots(figsize=(9,10),layout="constrained")
    for pos,r in zip(ys,leave_one_out):
        ax.errorbar(r["g"],pos,xerr=[[r["g"]-r["ci95_low_g"]],[r["ci95_high_g"]-r["g"]]],fmt="o",color="#166676",capsize=3)
    ax.set_yticks(ys,[r["omitted"] for r in leave_one_out]);ax.axvline(0,color="#777")
    ax.axvline(primary["g"],color="#7c2d55",linestyle="--",label="All-study estimate")
    ax.set_xlabel("Pooled Hedges g after omitting the named study\nModified Hartung-Knapp intervals")
    ax.set_title("Every prespecified leave-one-study-out analysis");ax.legend();ax.grid(axis="x",alpha=.2)
    fig.savefig(OUT/'leave_one_out.png',dpi=160);plt.close(fig)

def main():
    OUT.mkdir(exist_ok=True)
    rows=list(csv.DictReader((ROOT/'data/source.csv').open()))
    for line,r in enumerate(rows,2):r['_source_csv_line']=line
    selected,flow=select(rows)
    primary=fit(selected)
    sensitivity={"REML_normal":fit(selected,interval="normal"),
                 "REML_unmodified_HK":fit(selected,interval="HK"),
                 "fixed_normal":fit(selected,method="fixed",interval="normal"),
                 "DL_normal":fit(selected,method="DL",interval="normal")}
    restrictions={"low_randomization_and_allocation":lambda r:r['risk']['random_sequence_generation_selection_bias']=='Low risk' and r['risk']['allocation_concealment_selection_bias']=='Low risk',
                  "low_blinded_assessor":lambda r:r['risk']['blinded_outcome_assessor']=='Low risk',
                  "overall_low_risk":lambda r:r['risk']['overall']=='Low risk'}
    for label,keep in restrictions.items():
        subset=[r for r in selected if keep(r)]
        sensitivity[label]={**fit(subset),"studies":[r['studyID'] for r in subset]}
    alternative,_=select(rows,prefer_clinician=False)
    sensitivity['lexicographic_measure']=fit(alternative)
    sensitivity['lexicographic_measure']['changed_measures']=[{'studyID':r['studyID'],'outcome':r['outcome']} for r in alternative if next(s for s in selected if s['studyID']==r['studyID'])['outcome']!=r['outcome']]
    usual,_=select(rows,controls={'Usual care'})
    sensitivity['usual_care_only']={**fit(usual),"studies":[r['studyID'] for r in usual]}
    sensitivity['published_arm_change']=fit(changed_contrasts(selected))
    for rho in [0,.5,.8]:sensitivity['arm_change_r_'+str(rho).replace('.','p')]=fit(changed_contrasts(selected,rho))
    leave=[{"omitted":r['studyID'],**fit([s for s in selected if s['studyID']!=r['studyID']])} for r in selected]
    summaries={"minimum_pooled_g":min(r['g'] for r in leave),"maximum_pooled_g":max(r['g'] for r in leave),
               "intervals_excluding_zero":sum(r['interval_excludes_zero'] for r in leave),"comparisons":len(leave),
               "largest_absolute_point_shift_study":max(leave,key=lambda r:abs(r['g']-primary['g']))['omitted'],
               "maximum_absolute_point_shift_g":round(max(abs(r['g']-primary['g']) for r in leave),3)}
    # Compare the documented esc small-sample implementation to source arm changes.
    g_errors=[];se_errors=[];source_discrepancies=[]
    for item in selected:
        for r in item['exercise_rows']+item['control_rows']:
            size=n(r)
            corrected=(number(r,'mean')-number(r,'pre_mean'))/number(r,'pre_sd')*(1-3/(4*size-9))
            g_errors.append(abs(corrected-number(r,'smd')))
            if abs(corrected-number(r,'smd')) > 1e-10:
                source_discrepancies.append({'studyID':r['studyID'],'arm_number':r['arm_number'],
                                            'source_csv_line':r['_source_csv_line'],
                                            'stored_g':number(r,'smd'),'documented_formula_g':corrected,
                                            'stored_mean_diff':number(r,'mean_diff'),
                                            'post_minus_baseline_mean':round(number(r,'mean')-number(r,'pre_mean'),6),
                                            'absolute_difference':abs(corrected-number(r,'smd'))})
            rho=.18 if r['reported']=='Clinician' else .25
            recalculated=math.sqrt(2*(1-rho)/size+number(r,'smd')**2/(2*size))
            se_errors.append(abs(recalculated-number(r,'se_smd')))
    public=[{k:v for k,v in r.items() if k not in ['exercise_rows','control_rows']} for r in selected]
    source_keys=defaultdict(list)
    for r in rows:
        source_keys[(r['studyID'],r['arm_number'],r['outcome'],r['weeks_from_end_of_treatment_to_measurement'])].append(r['_source_csv_line'])
    duplicate_keys=[{'studyID':key[0],'arm_number':key[1],'outcome':key[2],'time':key[3],'csv_lines':lines}
                    for key,lines in source_keys.items() if len(lines)>1]
    result={"primary":primary,"sensitivity":sensitivity,"leave_one_out":leave,"leave_one_out_summary":summaries,
            "interval_level_percent":95,
            "selected_studies":public,"flow":flow,"source_rows":len(rows),
            "source_studies":len({r['studyID'] for r in rows}),
            "source_overall_risk":{risk.lower().replace(' ','_'):sum(r['risk']['overall']==risk for r in selected) for risk in RISK},
            "checks":{"one_contrast_per_study":True,"unique_eligible_source_keys":True,
                      "noneligible_source_duplicate_keys":duplicate_keys,
                      "combined_variance_identity":True,"REML_score_matches_likelihood":True,
                      "max_source_arm_g_discrepancy":round(max(g_errors),10),
                      "max_source_arm_g_discrepancy_display":round(max(g_errors),5),
                      "source_arm_g_discrepancies":source_discrepancies,
                      "source_arm_g_discrepancy_count":len(source_discrepancies),
                      "max_source_arm_se_discrepancy":round(max(se_errors),10)},
            "source_sha256":hashlib.sha256((ROOT/'data/source.csv').read_bytes()).hexdigest()}
    result['checks']['independent_metafor_reference']=validate_reference(selected,primary,sensitivity,leave)
    result=declared_precision(result)
    (OUT/'R1.json').write_text(json.dumps(result,ensure_ascii=False,indent=2,allow_nan=False)+'\n')
    with (OUT/'study_effects.csv').open('w',newline='') as f:
        names=['studyID','outcome','yi','vi','n_exercise','n_control','mean_exercise','mean_control','sd_exercise','sd_control']
        writer=csv.DictWriter(f,fieldnames=names);writer.writeheader();writer.writerows({k:r[k] for k in names} for r in selected)
    (OUT/'selection.json').write_text(json.dumps(flow,ensure_ascii=False,indent=2)+'\n')
    plot(primary,selected,leave)
    print(json.dumps({'primary':primary,'sensitivity':sensitivity,'leave_one_out':summaries,'checks':result['checks']},indent=2))

if __name__=='__main__':main()
