#!/usr/bin/env python3
"""Exact finite checks for the randomization unit (Python 3 standard library).
Run: python foundations-randomization-exact-check.py [--output PATH]
All probabilities and moments are fractions; no Monte Carlo or dependencies.
The default JSON output is beside this script.
All finite verification gates remain active under python -O.
"""
import argparse
from fractions import Fraction as F
from itertools import combinations,product
from math import comb
from pathlib import Path
import json
def check(condition,message):
 if not condition:
  raise AssertionError(message)

parser=argparse.ArgumentParser(description=__doc__)
parser.add_argument('--output',type=Path,default=Path(__file__).with_name('foundations-randomization-exact-results.json'))
args=parser.parse_args()
mean=lambda xs:sum(xs,F())/len(xs)
var=lambda xs:sum((x-mean(xs))**2 for x in xs)/(len(xs)-1)
diff=lambda y,z:mean([y[i] for i in z])-mean([y[i] for i in range(len(y)) if i not in z])
assign=lambda N,n:list(map(frozenset,combinations(range(N),n)))
obs=lambda y0,y1,z:[y1[i] if i in z else y0[i] for i in range(len(y0))]
def sharp_p(y,z,space):
 t=abs(diff(y,z));return F(sum(abs(diff(y,s))>=t for s in space),len(space))
A=assign(4,2);y0=list(map(F,[0,0,0,0]));y1=list(map(F,[-1,-1,-1,3]));rows=[]
for z in A:
 y=obs(y0,y1,z);est=diff(y,z);vh=var([y[i] for i in z])/2+var([y[i] for i in range(4) if i not in z])/2
 rows.append({'treated':[i+1 for i in z],'estimate':str(est),'conservative_variance_estimate':str(vh),'fisher_sharp_p':str(sharp_p(y,z,A))})
es=[diff(obs(y0,y1,z),z) for z in A];truev=mean([(e-mean(es))**2 for e in es]);vv=var(y1)/2+var(y0)/2-var([b-a for a,b in zip(y0,y1)])/4
check(mean(es)==0 and truev==vv==1, 'sharp-versus-weak design moments')
check(sum(F(row['fisher_sharp_p'])<=F(1,3) for row in rows)==3, 'sharp-null rejection probability')
res={'sharp_vs_weak_null':{'tau':'0','true_design_variance':str(truev),'mean_conservative_variance':'2','fisher_rejection_probability_at_alpha_1_over_3':'1/2','rows':rows}}
# Exact inversion of absolute mean-difference Fisher test.
y=list(map(F,[3,4,1,2]));z=frozenset([0,1]);inversion=[]
for d in list(map(F,[-2,0,1,F(3,2),2,F(5,2),3,4,6])):
 adj=[yy-d*(i in z) for i,yy in enumerate(y)];pv=sharp_p(adj,z,A);inversion.append({'delta':str(d),'p':str(pv),'retained_at_alpha_1_over_3':pv>F(1,3)})
res['constant_effect_inversion']={'observed':[3,4,1,2],'treated':[1,2],'accepted_set':'[1,3] at alpha=1/3; all real numbers at alpha=0.05','grid':inversion}
# Rerandomization accepted support and inclusion probabilities.
x=[-2,-1,1,2];acc=[z for z in A if abs(diff(list(map(F,x)),z))<=1]
check(len(acc)==4, 'restricted support size')
pis=[F(sum(i in z for z in acc),len(acc)) for i in range(4)];check(pis==[F(1,2)]*4, 'restricted inclusion probabilities')
res['rerandomization']={'accepted_treated_sets_1_based':[[i+1 for i in z] for z in acc],'acceptance_probability':str(F(len(acc),len(A))),'inclusion_probabilities':list(map(str,pis)),'expected_draws':str(F(len(A),len(acc)))}
# Blocked design contrasted with complete randomization.
y0=list(map(F,[0,1,2,3,10,11,12,13]));tau=list(map(F,[1,0,2,1,1,2,0,1]));y1=[a+b for a,b in zip(y0,tau)];blocks=[list(range(4)),list(range(4,8))]
bspace=[frozenset(a+b) for a,b in product(list(combinations(range(4),2)),list(combinations(range(4,8),2)))];cspace=assign(8,4)
def design_summary(space):
 es=[diff(obs(y0,y1,z),z) for z in space];return {'assignments':len(space),'mean':str(mean(es)),'variance':str(mean([(e-mean(es))**2 for e in es]))}
res['blocked_vs_complete']={'blocked':design_summary(bspace),'complete':design_summary(cspace)}
# Fisher exact independent-binomial conditional table 1,9;11,3.
N=24;n=10;K=12;xobs=1;mass={x:F(comb(K,x)*comb(N-K,n-x),comb(N,n)) for x in range(max(0,n-(N-K)),min(n,K)+1)}
left=sum(p for x,p in mass.items() if x<=xobs);two=sum(p for p in mass.values() if p<=mass[xobs]);check(sum(mass.values())==1, 'Fisher mass normalization')
res['fisher_exact']={'table':[[1,9],[11,3]],'conditional_mass':{str(k):str(v) for k,v in mass.items()},'left_p':str(left),'probability_ordered_two_sided_p':str(two)}
# Paired binary table: rows before=0,1 and cols after=0,1.
b=9;c=1;D=b+c;exact=2*sum(F(comb(D,j),2**D) for j in range(b,D+1));check(exact==F(11,512), 'McNemar example probability')
res['mcnemar']={'table':[[10,9],[1,10]],'discordant':D,'two_sided_exact_p':str(exact),'pearson_statistic':str(F((b-c)**2,D))}
# Extended finite checks, using rational arithmetic throughout.
def weighted_tail(values,weights):
 return [sum((w for t,w in zip(values,weights) if t>=v),F()) for v in values]
def superuniform(pvals,weights):
 return all(sum((w for p,w in zip(pvals,weights) if p<=u),F())<=u for u in set(pvals))
checks={};count=0
for ints in product(range(1,4),repeat=3):
 ws=[F(x,sum(ints)) for x in ints]
 for vals in product(range(3),repeat=3):
  check(superuniform(weighted_tail(vals,ws),ws), 'weighted tail superuniformity');count+=1
checks['weighted_tail_laws']=count
res['weighted_assignment_example']={'values':[3,1,0],'weights':['1/2','1/3','1/6'],'pvalues':list(map(str,weighted_tail([3,1,0],[F(1,2),F(1,3),F(1,6)])))}
# Neyman identity on every binary potential-outcome table for N=4, all arm sizes.
count=0
for vals in product(range(2),repeat=8):
 a=list(map(F,vals[:4]));b=list(map(F,vals[4:]));tau=[v-u for u,v in zip(a,b)]
 for n1 in (1,2,3):
  sp=assign(4,n1);es=[diff(obs(a,b,z),z) for z in sp];v=mean([(e-mean(es))**2 for e in es]);formula=var(b)/n1+var(a)/(4-n1)-var(tau)/4
  check(mean(es)==mean(tau) and v==formula, 'Neyman finite-population identity')
  if n1==2:
   vh=[var([b[i] for i in z])/2+var([a[i] for i in range(4) if i not in z])/2 for z in sp]
   check(mean(vh)-v==var(tau)/4, 'Neyman conservative variance gap')
  count+=1
checks['neyman_all_binary_tables_and_arm_sizes']=count
# New Neyman page table.
a=list(map(F,[1,2,3,4]));b=list(map(F,[2,4,6,8]));rows=[]
for z in assign(4,2):
 y=obs(a,b,z);rows.append({'treated':[i+1 for i in sorted(z)],'estimate':str(diff(y,z)),'variance_estimate':str(var([y[i] for i in z])/2+var([y[i] for i in range(4) if i not in z])/2)})
check(mean([F(x['estimate']) for x in rows])==F(5,2), 'Neyman example mean')
check(mean([(F(x['estimate'])-F(5,2))**2 for x in rows])==F(15,4), 'Neyman example variance')
check(mean([F(x['variance_estimate']) for x in rows])==F(25,6), 'Neyman example estimated variance')
res['neyman_page_example']={'rows':rows,'true_variance':'15/4','mean_variance_estimate':'25/6'}
# Pair moments, every binary potential table with two pairs.
count=0
for vals in product(range(2),repeat=8):
 # Two pairs, each encoded as Y10,Y11,Y20,Y21.
 plus=[];minus=[]
 for j in range(2):
  q=list(map(F,vals[4*j:4*j+4]));plus.append(q[1]-q[2]);minus.append(q[3]-q[0])
 mus=[(a+b)/2 for a,b in zip(plus,minus)];deltas=[(a-b)/2 for a,b in zip(plus,minus)]
 ds=[list(row) for row in product(*zip(plus,minus))];es=[mean(row) for row in ds];truev=sum(x*x for x in deltas)/4
 check(mean(es)==mean(mus) and mean([(e-mean(mus))**2 for e in es])==truev, 'matched-pair exact moments')
 check(mean([var(row)/2 for row in ds])-truev==sum((m-mean(mus))**2 for m in mus)/2, 'matched-pair conservative variance gap')
 count+=1
checks['matched_pair_binary_tables']=count
plus=list(map(F,[-1,1,0]));minus=list(map(F,[3,3,2]));ds=list(product(*zip(plus,minus)));es=[mean(row) for row in ds]
check(mean(es)==F(4,3) and mean([(e-F(4,3))**2 for e in es])==F(2,3), 'pair example true moments')
check(mean([var(row)/3 for row in ds])==F(7,9), 'pair example estimated variance')
res['pair_page_example']={'estimates':list(map(str,sorted(es))),'true_variance':'2/3','mean_variance_estimate':'7/9','sharp_p_for_3_3_2':str(F(sum(abs(sum(s*d for s,d in zip(signs,[3,3,2])))>=8 for signs in product([-1,1],repeat=3)),8))}
# Restricted-design support and variance counterexamples.
sp=assign(4,2);acc=[z for z in sp if abs(diff(list(map(F,[-2,-1,1,2])),z))<=1];y=list(map(F,[0,1,2,6]));z=frozenset([0,2])
check(sharp_p(y,z,acc)==F(1,2) and sharp_p(y,z,sp)==F(2,3), 'restricted-support p-values')
y=list(map(F,[0,1,0,1]));check(mean([diff(y,z)**2 for z in acc])==F(1,2) and mean([diff(y,z)**2 for z in sp])==F(1,3), 'restricted-design variance comparison')
res['restricted_design_checks']={'conditional_p':'1/2','wrong_original_support_p':'2/3','outcome_variance_restricted':'1/2','outcome_variance_complete':'1/3'}
# Inversion coverage at a true constant effect, all binary baselines for 4 units.
count=0
for base in product(range(2),repeat=4):
 for delta in map(F,[-2,0,2]):
  ps=[]
  for z in sp:
   y=[F(x)+delta*(i in z) for i,x in enumerate(base)];adjusted=[yy-delta*(i in z) for i,yy in enumerate(y)];ps.append(sharp_p(adjusted,z,sp))
  check(superuniform(ps,[F(1,6)]*6), 'constant-effect inversion coverage');count+=1
checks['constant_effect_coverage_tables']=count
# Enumerate every exact breakpoint and representative interval for the inversion example.
y=list(map(F,[3,4,1,2]));zo=frozenset([0,1]);coeff=[]
for z in sp:
 a=diff(y,z);b=diff([yy-F(i in zo) for i,yy in enumerate(y)],z)-a;coeff.append((a,b))
ao,bo=coeff[sp.index(zo)];roots=set()
for a,b in coeff:
 for s in [-1,1]:
  if b-s*bo:roots.add((s*ao-a)/(b-s*bo))
roots=sorted(roots);check(roots==[F(1),F(2),F(3)], 'exact inversion breakpoints')
cells=[]
for i in range(len(roots)+1):
 left=roots[i-1] if i else None;right=roots[i] if i<len(roots) else None
 rep=(left+right)/2 if left is not None and right is not None else (right-1 if left is None else left+1)
 vals=[abs(a+b*rep) for a,b in coeff];pv=weighted_tail(vals,[F(1,6)]*6)[sp.index(zo)]
 cells.append({'kind':'open_interval','left':str(left),'right':str(right),'representative':str(rep),'p':str(pv),'retained':pv>F(1,3)})
 if right is not None:
  vals=[abs(a+b*right) for a,b in coeff];pv=weighted_tail(vals,[F(1,6)]*6)[sp.index(zo)]
  cells.append({'kind':'point','delta':str(right),'p':str(pv),'retained':pv>F(1,3)})
res['inversion_exact_cells']=cells
# Conditional Fisher probabilities, both hypergeometric representations, recurrence and calibration.
count=0;ratios=0
for n1 in range(1,13):
 for n0 in range(1,13):
  N=n1+n0
  for K in range(N+1):
   L=max(0,K-n0);U=min(n1,K);xs=list(range(L,U+1));mass=[F(comb(n1,x)*comb(n0,K-x),comb(N,K)) for x in xs]
   check(sum(mass)==1, 'hypergeometric normalization')
   check(mass==[F(comb(K,x)*comb(N-K,n1-x),comb(N,n1)) for x in xs], 'hypergeometric representations')
   for i,x in enumerate(xs[:-1]):
    check(mass[i+1]/mass[i]==F((K-x)*(n1-x),(x+1)*(N-K-n1+x+1)), 'hypergeometric adjacent-mass ratio');ratios+=1
   pp=[sum((m for m in mass if m<=mi),F()) for mi in mass]
   pl=[sum(mass[:i+1]) for i in range(len(mass))];pr=[sum(mass[i:]) for i in range(len(mass))];pd=[min(F(1),2*min(a,b)) for a,b in zip(pl,pr)]
   for ps in [pp,pl,pr,pd]:check(superuniform(ps,mass), 'conditional Fisher p-value calibration')
   count+=1
checks['fisher_conditional_strata']=count;checks['fisher_mass_ratio_equalities']=ratios
res['asymmetric_fisher_table']={'table':[[0,3],[3,2]],'mass':['10/56','30/56','15/56','1/56'],'probability_ordered_p':'11/56','doubled_min_tail_p':'5/14'}
# McNemar every discordance count through 40, with exact double-tail equivalence.
count=0
for D in range(41):
 mass=[F(comb(D,k),2**D) for k in range(D+1)];ps=[]
 for b in range(D+1):
  p=sum((mass[k] for k in range(D+1) if abs(2*k-D)>=abs(2*b-D)),F());other=min(F(1),2*sum(mass[:min(b,D-b)+1]));check(p==other, 'McNemar double-tail equivalence');ps.append(p)
 check(superuniform(ps,mass), 'McNemar conditional calibration');count+=1
checks['mcnemar_conditional_strata']=count
res['dependent_pair_failure']={'D':10,'p_both_possible_outcomes':'1/512','rejection_probability_at_0_05':'1'}
res['checks']=checks;res['status']='PASS'
# Additional transfer task: a different four-unit constant-effect table.
a=list(map(F,[0,1,2,3]));b=[x+1 for x in a];sp=assign(4,2);zo=sp[0];y=obs(a,b,zo)
es=[diff(obs(a,b,z),z) for z in sp]
check(mean(es)==1 and mean([(e-1)**2 for e in es])==F(5,3), 'transfer design moments')
check(y==[1,2,2,3] and sharp_p(y,zo,sp)==F(2,3), 'transfer sharp-null p-value')
for delta,expected in [(-3,F(1,3)),(-2,F(2,3)),(-1,F(1)),(0,F(2,3)),(1,F(1,3))]:
 adjusted=[v-delta*(i in zo) for i,v in enumerate(y)]
 check(sharp_p(adjusted,zo,sp)==expected, 'transfer inversion grid')
true_ps=[sharp_p([v-(i in z) for i,v in enumerate(obs(a,b,z))],z,sp) for z in sp]
check(true_ps==[F(1,3),F(2,3),F(1),F(1),F(2,3),F(1,3)], 'transfer true-effect p-value vector')
check(F(sum(p>F(1,3) for p in true_ps),len(sp))==F(2,3), 'transfer confidence coverage')
check(F(2,comb(8,4))==F(1,35), 'independent table minimum p-value')
res['capstone_transfer']={'potential_y0':[0,1,2,3],'potential_y1':[1,2,3,4],'observed_y':[1,2,2,3],'true_tau':'1','design_variance':'5/3','fisher_zero_p':'2/3','accepted_at_alpha_1_over_3':'[-2,0]','paired_four_discordances_p':'1/8','independent_four_vs_four_p':'1/35'}
args.output.parent.mkdir(parents=True,exist_ok=True)
args.output.write_text(json.dumps(res,ensure_ascii=False,indent=2)+'\n')
print(json.dumps({'status':'PASS','checks':checks},ensure_ascii=False,indent=2))
