#!/usr/bin/env python3
"""Exact S04 checks using Python's standard library.

Exhaustive comparison graphs and finite loss tables, exact probability sums,
and Decimal evaluations of the displayed Hoeffding radii. The accompanying
proofs, not finite tests, establish all-distribution guarantees.
"""
from pathlib import Path
from fractions import Fraction as Q
from itertools import product,combinations
from math import comb,ceil,floor
import json,argparse
from decimal import Decimal as D,localcontext
from collections import Counter
def verify(condition,message):
 if not condition:raise AssertionError(message)

counts=Counter()
def check(group,condition):
 counts[group]+=1
 if not condition:raise AssertionError((group,counts[group]))
u=Path(__file__).resolve().parent
mean=lambda a:sum(a,Q())/len(a)
def endpoints(y,alpha,folds=None):
 n=len(y);folds=folds or [[i] for i in range(n)];low=[];high=[];records=[]
 for group in folds:
  f=mean([y[i] for i in range(n) if i not in group])
  for i in group:
   e=abs(y[i]-f);low.append(f-e);high.append(f+e);records.append([i,str(f),str(e),str(f-e),str(f+e)])
 ell=floor(alpha*(n+1));rank=ceil((1-alpha)*(n+1));low.sort();high.sort()
 return (None if ell==0 else low[ell-1],None if rank>n else high[rank-1]),records
res={}
y=list(map(Q,[0,0,2,4]));a=Q(1,5)
res['jackknife_example']={'endpoints':[str(x) for x in endpoints(y,a)[0]],'rows':endpoints(y,a)[1]}
res['cv_example']={'endpoints':[str(x) for x in endpoints(y,a,[[0,1],[2,3]])[0]],'rows':endpoints(y,a,[[0,1],[2,3]])[1]}
# Exact IID coverage over a finite response law, no floating Monte Carlo.
rows=[]
for n in [2,3,4,5]:
 for alpha in [Q(1,10),Q(1,5),Q(1,4),Q(2,5)]:
  covered=0
  for ys in product(range(3),repeat=n+1):
   (lo,hi),_=endpoints(list(map(Q,ys[:-1])),alpha)
   covered+=(lo is None or lo<=ys[-1]) and (hi is None or hi>=ys[-1])
  cov=Q(covered,3**(n+1));verify(cov>=1-2*alpha, 'jackknife exact IID coverage')
  rows.append({'n':n,'alpha':str(alpha),'coverage':str(cov),'bound':str(1-2*alpha)})
res['jackknife_exact_iid_coverages']=rows
rows=[]
for alpha in [Q(1,10),Q(1,5),Q(1,4),Q(2,5)]:
 covered=0;n=4;K=2
 for ys in product(range(3),repeat=n+1):
  (lo,hi),_=endpoints(list(map(Q,ys[:-1])),alpha,[[0,1],[2,3]])
  covered+=(lo is None or lo<=ys[-1]) and (hi is None or hi>=ys[-1])
 cov=Q(covered,3**5);bound=1-2*alpha-Q(1,6);verify(cov>=bound, 'CV exact IID coverage')
 rows.append({'alpha':str(alpha),'coverage':str(cov),'bound':str(bound)})
res['cv_exact_iid_coverages']=rows
# Polynomial helper for exact population subgroup coverage in the CQR example.
def mul(a,b):
 c=[Q(0)]*(len(a)+len(b)-1)
 for i,x in enumerate(a):
  for j,y in enumerate(b):c[i+j]+=x*y
 return c
def pw(a,n):
 x=[Q(1)]
 for _ in range(n):x=mul(x,a)
 return x
def integral(p,a,b):return sum((c*(b**(j+1)-a**(j+1))/Q(j+1) for j,c in enumerate(p)),Q())
# Probability Bin(9,v)<=7, expanded exactly.
poly=[Q(0)]*10
for j in range(8):
 p=mul([Q(0)]*j+[Q(comb(9,j))],pw([Q(1),Q(-1)],9-j))
 for k,t in enumerate(p):poly[k]+=t
v=Q(171,200);tail=integral(poly,v,Q(1));left=integral(poly,Q(0),v)
cb=Q(200,29)*tail;ca=Q(10,9)*left+Q(10,29)*tail
verify(Q(9,10)*ca+Q(1,10)*cb==Q(4,5), 'CQR marginal coverage identity')
emptyprob=sum((Q(comb(9,j))*v**j*(1-v)**(9-j) for j in [8,9]),Q())
res['cqr_population_example']={'group_A_coverage':str(ca),'group_B_coverage':str(cb),'marginal_coverage':'4/5','probability_B_interval_empty':str(emptyprob),'group_A_decimal':float(ca),'group_B_decimal':float(cb),'empty_decimal':float(emptyprob)}
# CRC expected risk and the probability of bad realized calibration risk.
res['crc_conditional_counterexample']={'n':4,'alpha':'2/5','risky_loss_probability':'1/2','risky_selected_probability':'5/16','expected_test_loss':'5/32','probability_conditional_risk_above_alpha':'5/16'}
verify(Q(1,2)*Q(5,16)==Q(5,32)<Q(2,5), 'CRC conditional-risk example')
res['nonmonotone_crc_failure']={'n':1,'alpha':'1/2','candidate_count':10,'per_candidate_risk':'9/10','expected_selected_risk':str(Q(9,10)*(1-Q(9,10)**10))}
verify(Q(9,10)*(1-Q(9,10)**10)>Q(1,2), 'nonmonotone CRC failure')
# Exact binomial tail for deployment/calibration-conditional threshold planning.
def failure(n,k,alpha):return sum((Q(comb(n,j))*(1-alpha)**j*alpha**(n-j) for j in range(k,n+1)),Q())
plans=[]
for n in [4,9,14,19,29,39,59,99]:
 for alpha in [Q(1,5),Q(1,10)]:
  delta=Q(1,20);k=next((k for k in range(1,n+1) if failure(n,k,alpha)<=delta),n+1)
  plans.append({'n':n,'alpha':str(alpha),'delta':str(delta),'rank':k,'bad_calibration_probability':str(failure(n,k,alpha))})
res['conditional_rank_plans']=plans

# Exhaust every strict comparison outcome, including ties (no winner).
for N in range(2,6):
 pairs=list(combinations(range(N),2))
 for graph in product(range(3),repeat=len(pairs)):
  wins=[0]*N
  for (i,j),v in zip(pairs,graph):
   if v==1:wins[i]+=1
   elif v==2:wins[j]+=1
  for t in range(1,N+1):
   size=sum(x>=t for x in wins)
   check('jackknife_comparison_graph',size==0 or size<=2*(N-t)-1)
# CV augmentation: 3 equal folds of size 2; all 3^12 cross-fold outcomes.
N=6;m=2;n=4
pairs=[(i,j) for i,j in combinations(range(N),2) if i//m!=j//m]
for graph in product(range(3),repeat=len(pairs)):
 wins=[0]*N
 for (i,j),v in zip(pairs,graph):
  if v==1:wins[i]+=1
  elif v==2:wins[j]+=1
 for t in range(1,n+2):
  size=sum(x>=t for x in wins)
  check('cv_fold_comparison_graph',size==0 or size<=2*(n-t)+m)
# Exchangeable finite loss tables: integer units of 1/2, B=1.
# For every held-out position calculate the actual CRC choice, then average
# held-out losses. This covers all uniform permutations of the fixed table.
lossrows=[(a,b,0) for a in range(3) for b in range(a+1)]
for N in range(2,6):
 for rows in product(lossrows,repeat=N):
  sums=[sum(row[j] for row in rows) for j in range(3)]
  for alpha in [Q(1,4),Q(1,3),Q(1,2),Q(3,4)]:
   A,B=alpha.numerator,alpha.denominator
   total=0
   oracle=next(j for j in range(3) if sums[j]*B<=2*N*A)
   for held,row in enumerate(rows):
    j=next((j for j in range(3) if (sums[j]-row[j]+2)*B<=2*N*A),2)
    check('crc_oracle_order',j>=oracle and row[j]<=row[oracle])
    total+=row[j]
   check('crc_exchangeable_mean_budget',total*B<=2*N*A)
# RCPS pointwise confidence may not be searched as an arbitrary menu.
for size in range(1,11):
 for delta in [Q(1,20),Q(1,10),Q(1,5)]:
  anyfail=Q();suffixfail=Q()
  for flags in product([0,1],repeat=size):
   mass=delta**sum(flags)*(1-delta)**(size-sum(flags))
   if any(flags):anyfail+=mass
   U=[0 if v else 1 for v in flags]+[0]
   selected=next(j for j in range(size+1) if max(U[j:])<=Q(1,2))
   if selected<size:suffixfail+=mass
  check('rcps_suffix_exact_probability',suffixfail==delta)
  check('rcps_arbitrary_search_failure',anyfail==1-(1-delta)**size)
  if size>1:check('rcps_arbitrary_search_exceeds_delta',anyfail>delta)
# Exact Beta integrals vs binomial tail; factorial identity is evaluated
# independently as polynomial integration, not by calling a Beta routine.
from math import factorial
for n in range(1,21):
 for k in range(1,n+1):
  density=[Q(0)]*(k-1)+[Q(factorial(n),factorial(k-1)*factorial(n-k))*x for x in pw([Q(1),Q(-1)],n-k)]
  check('beta_density_integral',integral(density,Q(0),Q(1))==1)
  check('beta_mean_integral',integral([Q(0)]+density,Q(0),Q(1))==Q(k,n+1))
  check('beta_second_moment_integral',integral([Q(0),Q(0)]+density,Q(0),Q(1))==Q(k*(k+1),(n+1)*(n+2)))
  for alpha in [Q(1,10),Q(1,5),Q(1,3),Q(1,2)]:
   check('beta_binomial_identity',integral(density,Q(0),1-alpha)==failure(n,k,alpha))
# Discrete-score PIT bound: enumerate all Bernoulli calibration samples.
for n in range(1,9):
 for pzero in [Q(1,5),Q(1,2),Q(4,5)]:
  for k in range(1,n+1):
   for alpha in [Q(1,10),Q(1,5),Q(1,2)]:
    prob=Q()
    for zeros in range(n+1):
     coverage=pzero if zeros>=k else Q(1)
     if coverage<1-alpha:prob+=Q(comb(n,zeros))*pzero**zeros*(1-pzero)**(n-zeros)
    check('discrete_score_conservative_tail',prob<=failure(n,k,alpha))
# Independent Decimal evaluation of Hoeffding upper bounds for IID binary
# losses; the probability of accepting a truly unsafe action is exact.
with localcontext() as ctx:
 ctx.prec=80
 dec=lambda x:D(x.numerator)/D(x.denominator)
 numerical=[]
 for n in [4,100]:
  radius=(D(20).ln()/(2*n)).sqrt()
  numerical.append({'n':n,'radius':str(radius),'middle_upper':str(D(1)/8+radius)})
  check('hoeffding_displayed_choice', (D(1)/8+radius<=D('0.4'))==(n==100))
 for n in range(1,41):
  for delta in [Q(1,100),Q(1,20),Q(1,5)]:
   radius=((1/dec(delta)).ln()/(2*n)).sqrt()
   for p in [Q(1,5),Q(1,2),Q(4,5)]:
    for alpha in [Q(1,10),Q(1,3),Q(3,4)]:
     if p<=alpha:continue
     prob=sum((Q(comb(n,j))*p**j*(1-p)**(n-j) for j in range(n+1) if D(j)/n+radius<=dec(alpha)),Q())
     check('hoeffding_exact_unsafe_selection',prob<=delta)
res['hoeffding_examples']=numerical
# Hand interval inversion and loss table, including empty CQR output.
scores=[max(-y,y-4) for y in map(Q,[2,1,3,Q(3,2),Q(1,2),Q(7,2),Q(1,5),Q(19,5),5])]
q=sorted(scores)[7]
check('cqr_negative_threshold',q==Q(-1,5))
check('cqr_empty_output',1-q>Q(6,5)+q)
check('calibration_80percent_counterexample',failure(4,4,Q(1,5))==Q(256,625))
check('minimum_sample_80',failure(13,13,Q(1,5))>Q(1,20)>=failure(14,14,Q(1,5)))
check('minimum_sample_90',failure(28,28,Q(1,10))>Q(1,20)>=failure(29,29,Q(1,10)))
# Old split-conformal page deepening: exact label-conditional coverage with
# random class sizes, unequal class masses, ties and unseen-class fallback.
classes=[(0,0,Q(1,2)),(0,1,Q(1,4)),(1,0,Q(1,8)),(1,1,Q(1,8))]
for n in range(1,6):
 for alpha in [Q(1,5),Q(1,4),Q(2,5)]:
  cover=[Q(),Q()];mass=[Q(),Q()]
  for data in product(classes,repeat=n+1):
   label,score,_=data[-1];w=Q(1)
   for _,_,weight in data:w*=weight
   vals=sorted(s for lab,s,_ in data[:-1] if lab==label)
   k=ceil((len(vals)+1)*(1-alpha));q=None if k>len(vals) else vals[k-1]
   mass[label]+=w
   if q is None or score<=q:cover[label]+=w
  for label in [0,1]:check('label_conditional_exact_coverage',cover[label]>=mass[label]*(1-alpha))
check('label_conditional_empty_group',ceil(Q(4,5))==1)
check('label_conditional_hand_set',[name for name,score,q in [('A',Q(9,10),Q(4,5)),('B',Q(1,5),None),('C',Q(9,10),None)] if q is None or score<=q]==['B','C'])
res['capstone_99_calibration']={str(k):str(failure(99,k,Q(1,10))) for k in [90,94,95]}
check('capstone_99_minimum_rank',failure(99,94,Q(1,10))>Q(1,20)>=failure(99,95,Q(1,10)))
res['status']='PASS';res['assertions']=sum(counts.values());res['groups']=dict(counts)
res['limits']=['Exact finite enumerations check implementations, while distribution-free statements use the published proofs in the entries.','Decimal logarithms are high-precision numerical calculations, not interval certificates.','No production-site rendering test is claimed.']
args=argparse.ArgumentParser();args.add_argument('--output',type=Path);args=args.parse_args()
text=json.dumps(res,ensure_ascii=False,indent=2)+'\n'
if args.output:args.output.write_text(text)
print(text,end='')
