#!/usr/bin/env python3
"""Exact finite certificates. General population theorems are proved in the text."""
import argparse,json
from collections import Counter,defaultdict
from fractions import Fraction as F
from itertools import product
from math import comb
from pathlib import Path
CHECKS=Counter()
def check(condition,label):
 CHECKS[label]+=1
 if not condition:raise ValueError('Failed certificate: '+label)
def encode(x):
 if isinstance(x,F):return str(x)
 if isinstance(x,dict):return {str(k):encode(v)for k,v in x.items()}
 if isinstance(x,(list,tuple)):return [encode(v)for v in x]
 return x

def missing_bounds(observed,lower,upper):
 if not observed or len(observed)!=len(lower) or len(lower)!=len(upper):raise ValueError('nonempty equal lengths required')
 L=[];U=[]
 for value,a,b in zip(observed,lower,upper):
  a,b=F(a),F(b)
  if a>b:raise ValueError('reversed support')
  if value is None:L.append(a);U.append(b)
  else:
   value=F(value)
   if not a<=value<=b:raise ValueError('observed outcome outside known support')
   L.append(value);U.append(value)
 return sum(L,F())/len(L),sum(U,F())/len(U)

def trim_distribution(values,masses,keep):
 values=tuple(map(F,values));masses=tuple(map(F,masses));keep=F(keep)
 if len(values)!=len(masses) or not values:raise ValueError('nonempty paired support required')
 if any(m<0 for m in masses)or sum(masses)!=1 or not 0<keep<=1:raise ValueError('invalid mass or retained fraction')
 def side(reverse):
  amount=[F()for _ in values];left=keep
  for j in sorted(range(len(values)),key=lambda j:values[j],reverse=reverse):
   take=min(left,masses[j]);amount[j]=take;left-=take
  if left:raise ValueError('insufficient mass')
  return sum((y*r for y,r in zip(values,amount)),F())/keep,tuple(amount)
 return side(False),side(True)

def simplex_vertices(masses,keep):
 """Independent box/hyperplane vertices: all but one coordinate at endpoints."""
 masses=tuple(map(F,masses));keep=F(keep);out=set();k=len(masses)
 for pivot in range(k):
  other=[j for j in range(k)if j!=pivot]
  for bits in product([0,1],repeat=k-1):
   v=[F()for _ in masses]
   for j,b in zip(other,bits):v[j]=masses[j]*b
   v[pivot]=keep-sum(v)
   if 0<=v[pivot]<=masses[pivot]:out.add(tuple(v))
 return sorted(out)

def weighted_distribution(weights,success):
 if any(not isinstance(q,int)or q<0 for q in weights):raise ValueError('nonnegative integer weights required')
 ps=[F(success)]*len(weights) if not isinstance(success,(list,tuple)) else list(map(F,success))
 if len(ps)!=len(weights)or any(not 0<=p<=1 for p in ps):raise ValueError('invalid success probabilities')
 out={0:F(1)}
 for q,p in zip(weights,ps):
  nxt=defaultdict(F)
  for s,mass in out.items():nxt[s]+=mass*(1-p);nxt[s+q]+=mass*p
  out=dict(nxt)
 return out

def enumerated_distribution(weights,ps):
 out=defaultdict(F)
 for bits in product([0,1],repeat=len(weights)):
  probability=F(1)
  for b,p in zip(bits,ps):probability*=p if b else 1-p
  out[sum(q*b for q,b in zip(weights,bits))]+=probability
 return dict(out)
def tail(dist,t):return sum((v for s,v in dist.items()if s>=t),F())
def gamma_bounds(weights,t,gamma):
 gamma=F(gamma)
 if gamma<1:raise ValueError('Gamma must be at least one')
 l=1/(1+gamma);h=gamma/(1+gamma)
 return tail(weighted_distribution(weights,l),t),tail(weighted_distribution(weights,h),t)
def sign_candidate(differences,delta,gamma):
 adjusted=[F(d)-F(delta)for d in differences];nonzero=[d for d in adjusted if d]
 m=len(nonzero);k=sum(d>0 for d in nonzero)
 return {'nonzero':m,'positive':k,'upper':gamma_bounds([1]*m,k,gamma)[1]}
def compositions(total,k):
 if k==1:yield (total,);return
 for first in range(total+1):
  for rest in compositions(total-first,k-1):yield (first,)+rest

def run():
 CHECKS.clear();missing_models=0;ate_models=0;trim_models=0;tail_models=0
 # Finite labelled missing records, including all missing and all observed.
 for n in range(1,7):
  for observed in product([None,0,1],repeat=n):
   lo,hi=missing_bounds(observed,[0]*n,[1]*n);m=sum(x is None for x in observed)
   values=[]
   for fill in product([0,1],repeat=m):
    it=iter(fill);v=[next(it)if x is None else x for x in observed];values.append(F(sum(v),n))
   check(min(values)==lo and max(values)==hi,'missing_complete_extrema')
   check(hi-lo==F(m,n),'missing_identification_width');missing_models+=1
 # ATE endpoints simultaneously attained on the same potential table.
 for n in range(1,8):
  for counts in compositions(n,4):
   obs=[(a,y)for (a,y),c in zip([(0,0),(0,1),(1,0),(1,1)],counts)for _ in range(c)]
   p=F(sum(a for a,y in obs),n);m1=F(sum(a*y for a,y in obs),n);m0=F(sum((1-a)*y for a,y in obs),n)
   theory=(m1-m0-p,m1-m0+1-p);effects=[]
   for unseen in product([0,1],repeat=n):
    effects.append(F(sum((y-u)if a else(u-y)for (a,y),u in zip(obs,unseen)),n))
   check((min(effects),max(effects))==theory,'ate_joint_complete_extrema')
   check(theory[1]-theory[0]==1 and theory[0]<=0<=theory[1],'ate_width_and_zero');ate_models+=1
 observed=[1,1,None,None]+[1]*6+[None,None]
 check(missing_bounds(observed,[0]*12,[1]*4+[4]*8)==(F(2,3),F(3,2)),'terminal_heterogeneous_support')
 check(missing_bounds(observed,[0]*12,[4]*12)==(F(2,3),F(2)),'terminal_coarsened_support')
 # True box-simplex extrema, compared with sorting for atoms and negative supports.
 for k in range(2,6):
  values=[-3,-1,2,4,9][:k]
  for raw in compositions(6,k):
   if not all(raw):continue
   masses=tuple(F(c,6)for c in raw)
   for keep in [F(j,7)for j in range(1,8)]:
    low,high=trim_distribution(values,masses,keep);vertices=simplex_vertices(masses,keep)
    objectives=[sum((F(y)*r for y,r in zip(values,v)),F())/keep for v in vertices]
    check(low[0]==min(objectives)and high[0]==max(objectives),'trim_independent_vertex_extrema')
    for mean,allocation in [low,high]:
     check(sum(allocation)==keep and all(0<=r<=f for r,f in zip(allocation,masses)),'trim_mass_certificate')
     check(mean==sum((y*r for y,r in zip(values,allocation)),F())/keep,'trim_objective_certificate')
     # The selected and induced distributions reconstruct the same full observed law.
     if keep<1:
      fa=tuple(r/keep for r in allocation);fi=tuple((f-r)/(1-keep)for f,r in zip(masses,allocation))
      check(sum(fa)==sum(fi)==1,'lee_component_normalization')
      check(tuple(keep*x+(1-keep)*y for x,y in zip(fa,fi))==masses,'lee_observation_reconstruction')
    trim_models+=1
 low,high=trim_distribution([0,2,6],[F(1,4),F(1,2),F(1,4)],F(2,3))
 check((low[0],high[0])==(F(5,4),F(7,2)),'main_lee_means')
 check(low[1]==(F(1,4),F(5,12),F())and high[1]==(F(),F(5,12),F(1,4)),'main_fractional_atom')
 check((low[0]-1,high[0]-1)==(F(1,4),F(5,2)),'main_lee_effect')
 check((4-high[0],4-low[0])==(F(1,2),F(11,4)),'terminal_reverse_selection')
 # Explicit population type tables carry unconditional mass.
 tables=[[(1,1,0,1,F(3,16)),(1,1,2,1,F(5,16)),(0,1,2,0,F(1,16)),(0,1,6,0,F(3,16)),(0,0,0,0,F(1,4))],
         [(1,1,6,1,F(3,16)),(1,1,2,1,F(5,16)),(0,1,2,0,F(1,16)),(0,1,0,0,F(3,16)),(0,0,0,0,F(1,4))]]
 for table,target in zip(tables,[F(1,4),F(5,2)]):
  treatment=defaultdict(F);control=defaultdict(F);mass=F();effect=F()
  for s0,s1,y1,y0,w in table:
   check(s1>=s0,'lee_population_monotonicity')
   if s1:treatment[y1]+=w
   if s0:control[y0]+=w
   if s0 and s1:mass+=w;effect+=w*(y1-y0)
  check(dict(treatment)=={0:F(3,16),2:F(6,16),6:F(3,16)}and dict(control)=={1:F(1,2)},'lee_full_table_observation')
  check(effect/mass==target,'lee_full_table_effect')
 pooled=trim_distribution([0,2,10],[F(1,4),F(1,4),F(1,2)],F(3,4))
 check((pooled[0][0]-6,pooled[1][0]-6)==(F(-2),F(4,3)),'terminal_pooled_bounds')
 stratified=(F(1,3)*0+F(2,3),F(1,3)*2+F(2,3))
 check(stratified==(F(2,3),F(4,3)),'terminal_target_stratum_weights')
 # Duplicate support entries do not lose probability mass.
 check(trim_distribution([2,0,2,6],[F(1,4)]*4,F(2,3))[0][0]==low[0],'repeated_support_entries')
 # Complete product distributions, all cube corners and interior probabilities.
 for weights in [(),(0,),(1,),(1,2),(0,1,1),(1,2,2,3),(1,2,3,4),(0,2,2,3,4),(1,1,1,1,1,1)]:
  I=len(weights)
  for gamma in [F(1),F(3,2),F(2),F(3)]:
   l=1/(1+gamma);h=gamma/(1+gamma)
   endpoint=weighted_distribution(weights,h)
   check(endpoint==enumerated_distribution(weights,[h]*I),'convolution_matches_enumeration')
   check(sum(endpoint.values())==1,'distribution_normalization')
   tests=sorted(set(range(sum(weights)+2)))
   extrema={t:[]for t in tests}
   for ps in product([l,h],repeat=I):
    d=enumerated_distribution(weights,ps)
    for t in tests:extrema[t].append(tail(d,t))
    # Exact rejection probability for every achievable p-value level.
    upper={t:tail(endpoint,t)for t in d}
    for alpha in sorted(set(upper.values())|{F(1,20),F(1,4)}):
     rejected=sum((mass for t,mass in d.items()if upper[t]<=alpha),F())
     check(rejected<=alpha,'gamma_uniform_pvalue_validity')
   for t in tests:
    bounds=gamma_bounds(weights,t,gamma)
    check(bounds==(min(extrema[t]),max(extrema[t])),'gamma_attained_corner_extrema')
   if I:
    ps=[l+(h-l)*F(j+1,I+1)for j in range(I)];d=enumerated_distribution(weights,ps)
    for t in tests:
     a,b=gamma_bounds(weights,t,gamma);check(a<=tail(d,t)<=b,'gamma_interior_tail_bounds')
   tail_models+=1
 check(gamma_bounds([1]*6,6,1)[1]==F(1,64),'six_positive_gamma_one')
 check(gamma_bounds([1]*6,6,F(3,2))[1]==F(729,15625),'six_positive_gamma_three_halves')
 check(gamma_bounds([1]*6,6,2)[1]==F(64,729),'six_positive_gamma_two')
 check(gamma_bounds([1,2,2,3],7,2)==(F(1,27),F(8,27)),'weighted_main_tail')
 double=[]
 for ps in product([F(1,3),F(2,3)],repeat=2):
  d=enumerated_distribution([1,2],ps);double.append(d[0]+d[3])
 check(max(double)==F(5,9)and 2*gamma_bounds([1,2],3,2)[1]==F(8,9),'two_sided_bound_not_sharp')
 deltas=[-1,0,F(1,2),1,F(3,2),2,3];candidates=[sign_candidate([0,1,1,2],delta,2)for delta in deltas]
 check([v['upper']for v in candidates]==[F(16,81),F(8,27),F(16,27),F(8,9),F(80,81),F(1),F(1)],'terminal_candidate_breakpoints')
 check([v['upper']>F(1,4)for v in candidates]==[False,True,True,True,True,True,True],'terminal_acceptance_closed_endpoint')
 for ps in product([F(1,3),F(2,3)],repeat=4):
  cover=F()
  for bits in product([0,1],repeat=4):
   prob=F(1)
   for bit,p in zip(bits,ps):prob*=p if bit else 1-p
   ds=[a*(2*b-1)for a,b in zip([0,1,1,2],bits)]
   if sign_candidate(ds,0,2)['upper']>F(1,4):cover+=prob
  check(cover==1,'terminal_true_table_complete_coverage')
 for fn,args in [(trim_distribution,([0,1],[F(1,2)]*2,0)),(trim_distribution,([0,1],[1,1],F(1,2))),
                 (missing_bounds,([2],[0],[1])),(weighted_distribution,([-1],F(1,2))),(gamma_bounds,([1],1,F(1,2)))]:
  try:fn(*args)
  except ValueError:check(True,'invalid_contract_rejected')
  else:check(False,'invalid_contract_rejected')
 return encode({'status':'PASS','checks':sum(CHECKS.values()),'groups':dict(sorted(CHECKS.items())),
  'models':{'missing_tables':missing_models,'ate_observation_histograms':ate_models,'trimming_models':trim_models,'weighted_tail_models':tail_models},
  'main':{'heterogeneous_support_bounds':(F(2,3),F(3,2)),'lee_low':low,'lee_high':high,'stratified_lee':stratified,'pooled_lee':(F(-2),F(4,3)),'weighted_gamma_two':(F(1,27),F(8,27)),'candidate_deltas':deltas,'candidate_results':candidates},
  'scope':'Exact finite certificates; general sharpness, arbitrary distributions, and coverage theorems are proved in the text. No empirical interval is promoted to a confidence interval.'})
if __name__=='__main__':
 parser=argparse.ArgumentParser();parser.add_argument('--output',type=Path);args=parser.parse_args();res=run();data=json.dumps(res,ensure_ascii=False,indent=2)+'\n'
 if args.output:args.output.parent.mkdir(parents=True,exist_ok=True);args.output.write_text(data)
 else:print(data,end='')
