"""Exact integer-weight clustered binary null benchmark. No random numbers."""
import bisect,csv,json,math
from fractions import Fraction
from functools import lru_cache
from pathlib import Path
ROOT=Path(__file__).resolve().parents[1];OUT=ROOT/'results';OUT.mkdir(exist_ok=True)
config=json.loads((ROOT/'data/design.json').read_text())

def encode_fraction(x):return {'numerator':str(x.numerator),'denominator':str(x.denominator),'decimal':float(x)}

def convolution(a,b):
    c=[0]*(len(a)+len(b)-1)
    for i,x in enumerate(a):
        if x:
            for j,y in enumerate(b):
                if y:c[i+j]+=x*y
    return c

@lru_cache(None)
def rejection_tables(n):
    fisher={};oracle_data={}
    for total in range(2*n+1):
        lo=max(0,total-n);hi=min(n,total)
        weights={a:math.comb(n,a)*math.comb(n,total-a) for a in range(lo,hi+1)}
        denominator=math.comb(2*n,total);assert sum(weights.values())==denominator
        byweight={}
        for w in weights.values():byweight[w]=byweight.get(w,0)+w
        cumulative=0;reject={}
        for w in sorted(byweight):
            cumulative+=byweight[w];reject[w]=(20*cumulative<=denominator)
        for a,w in weights.items():fisher[a,total-a]=reject[w]
    return fisher

def donor_weights(k,regime):
    if regime=='independent':return [math.comb(k,x) for x in range(k+1)],2**k,Fraction(0)
    if regime=='perfect_copy':return [1]+[0]*(k-1)+[1],2,Fraction(1)
    a=int(regime)
    weights=[math.comb(x+a-1,x)*math.comb(k-x+a-1,k-x) for x in range(k+1)]
    denominator=math.comb(k+2*a-1,k);assert sum(weights)==denominator
    return weights,denominator,Fraction(1,2*a+1)

def evaluate(m,k,regime):
    q,d,rho=donor_weights(k,regime)
    # Exact validation of the beta-binomial variance and the Bernoulli controls.
    mean=Fraction(sum(i*x for i,x in enumerate(q)),d)
    variance=Fraction(sum(i*i*x for i,x in enumerate(q)),d)-mean*mean
    assert mean==Fraction(k,2)
    assert variance==Fraction(k,4)*(1+(k-1)*rho)
    arm=[1]
    for _ in range(m):arm=convolution(arm,q)
    D=d**m;assert sum(arm)==D
    n=m*k;mask=rejection_tables(n);numerator=0
    absdiff=[0]*(n+1)
    for a,wa in enumerate(arm):
        if wa:
            for b,wb in enumerate(arm):
                if wb:
                    w=wa*wb;absdiff[abs(a-b)]+=w
                    if mask[a,b]:numerator+=w
    den=D*D;assert sum(absdiff)==den
    tail=0;oracle_numerator=0
    for distance in range(n,-1,-1):
        tail+=absdiff[distance]
        if 20*tail<=den:oracle_numerator+=absdiff[distance]
    naive=Fraction(numerator,den);oracle=Fraction(oracle_numerator,den)
    assert oracle<=Fraction(1,20)
    if regime=='independent':assert naive<=Fraction(1,20)
    return {'donors_per_arm':m,'cells_per_donor':k,'regime':str(regime),'rho':encode_fraction(rho),'design_effect':encode_fraction(1+(k-1)*rho),'nominal_alpha':encode_fraction(Fraction(1,20)),'fisher_null_rejection':encode_fraction(naive),'oracle_null_rejection':encode_fraction(oracle)}

rows=[]
for m in config['donors_per_arm']:
    for k in config['cells_per_donor']:
        for regime in ['independent',*config['beta_symmetric_shapes'],'perfect_copy']:
            rows.append(evaluate(m,k,regime))
lookup={(r['donors_per_arm'],r['cells_per_donor'],r['regime']):r for r in rows}
selected=lookup[5,20,'5'];reference=lookup[5,1,'5'];small=lookup[5,20,'50']
R1={'grid_cases':len(rows),'independent_cases':sum(r['regime']=='independent' for r in rows),'oracle_cases_with_size_at_most_nominal':sum(Fraction(int(r['oracle_null_rejection']['numerator']),int(r['oracle_null_rejection']['denominator']))<=Fraction(1,20) for r in rows),'selected_m':5,'selected_k':20,'selected_rho':float(Fraction(1,11)),'selected_naive_type1':selected['fisher_null_rejection']['decimal'],'selected_oracle_type1':selected['oracle_null_rejection']['decimal'],'selected_design_effect':selected['design_effect']['decimal'],'one_cell_type1':reference['fisher_null_rejection']['decimal'],'small_rho_type1':small['fisher_null_rejection']['decimal'],'max_naive_type1':max(r['fisher_null_rejection']['decimal'] for r in rows),'cases_exceeding_nominal':sum(Fraction(int(r['fisher_null_rejection']['numerator']),int(r['fisher_null_rejection']['denominator']))>Fraction(1,20) for r in rows)}
(OUT/'exact-grid.json').write_text(json.dumps(rows,indent=2)+'\n');(OUT/'R1.json').write_text(json.dumps(R1,indent=2)+'\n')
with (OUT/'grid.csv').open('w',newline='') as f:
    columns=['donors_per_arm','cells_per_donor','regime','rho','design_effect','fisher_null_rejection','oracle_null_rejection']
    w=csv.DictWriter(f,fieldnames=columns);w.writeheader()
    for r in rows:w.writerow({c:r[c]['decimal'] if isinstance(r[c],dict) else r[c] for c in columns})
print(json.dumps(R1,indent=2))
