#!/usr/bin/env python3
"""Finite distribution-approximation checks for the Theoryroad S05 unit.
Requires Python 3 and mpmath. Finite laws and identities use Fraction;
Poisson TV uses rigorous rational Taylor/reciprocal intervals; Gaussian
Stein factors use 100-digit diagnostic checks, not a universal proof.
Use --output with your own path; default prints only, never changes pages.
"""
from fractions import Fraction as Q
from itertools import product,combinations
from collections import defaultdict
from math import comb,factorial
from pathlib import Path
import argparse,json
import mpmath as mp
mp.mp.dps=100
def require(condition,message):
 if not condition:raise ValueError(message)

COUNTS={}
def check(t,key):
 COUNTS[key]=COUNTS.get(key,0)+1
 if not t:raise AssertionError(key)
def M(x):return mp.mpf(x.numerator)/x.denominator if isinstance(x,Q) else mp.mpf(x)
def rep(x):return mp.nstr(M(x),60)
def pb(ps):
 v=[Q(1)]+[Q(0)]*len(ps)
 for j,p in enumerate(ps,1):
  for k in range(j,0,-1):v[k]=(1-p)*v[k]+p*v[k-1]
  v[0]*=1-p
 return v
def states(ps):
 for bits in product([0,1],repeat=len(ps)):
  mass=Q(1)
  for b,p in zip(bits,ps):mass*=p if b else 1-p
  if mass:yield bits,mass
def I(a,b=None):return (Q(a),Q(a if b is None else b))
def add(a,b):return (a[0]+b[0],a[1]+b[1])
def neg(a):return (-a[1],-a[0])
def sub(a,b):return add(a,neg(b))
def scale(c,a):return (c*a[0],c*a[1]) if c>=0 else (c*a[1],c*a[0])
def absi(a):return (Q(0) if a[0]<=0<=a[1] else min(abs(v) for v in a),max(abs(v) for v in a))
def expminus(x,N=90):
 require(0<=x<N+2, 'exponential enclosure requires 0 <= x < N+2')
 term=S=Q(1)
 for k in range(1,N+1):term*=x/k;S+=term
 tail=term*x/(N+1)/(1-x/Q(N+2))
 return (1/(S+tail),1/S)
def poisson_tv(v,lam):
 """All mass above finite support is 1-sum(q_k), with outward interval algebra."""
 q=expminus(lam);total=I(0);mass=I(0)
 for k,p in enumerate(v):
  if k:q=scale(lam/k,q)
  total=add(total,absi(sub(I(p),q)));mass=add(mass,q)
 return scale(Q(1,2),add(total,sub(I(1),mass)))
def overlap(m,p):
 d={(0,0):Q(1)}
 for _ in range(m):
  e=defaultdict(Q)
  for (b,k),v in d.items():e[(0,k)]+=v*(1-p);e[(1,k+b)]+=v*p
  d=e
 return [sum(v for (b,j),v in d.items() if j==k) for k in range(m)]
# A density on (a,b) includes exactly atoms at b or to its right. Use an exact midpoint.
def zero_pieces(xs,ps):
 mu=sum(x*p for x,p in zip(xs,ps));ys=[x-mu for x in xs];var=sum(y*y*p for y,p in zip(ys,ps));require(var>0, 'zero-bias construction requires positive variance')
 pcs=[(a,b,sum(y*p for y,p in zip(ys,ps) if y>(a+b)/2)/var) for a,b in zip(ys[:-1],ys[1:])]
 return ys,var,pcs

def run():
 COUNTS.clear()
 # PB exact DP versus independently enumerated bit vectors, including endpoints.
 for n in range(6):
  for ps in product([Q(0),Q(1,4),Q(1,2),Q(1)],repeat=n):
   v=pb(ps);brute=[Q(0)]*(n+1)
   for bits,p in states(ps):brute[sum(bits)]+=p
   check(v==brute,'PB_full_enumeration');check(sum(v)==1,'PB_mass')
   mu=sum(ps);check(sum(k*p for k,p in enumerate(v))==mu,'PB_mean')
   check(sum((k-mu)**2*p for k,p in enumerate(v))==sum(p*(1-p) for p in ps),'PB_variance')
   check(v==pb(tuple(reversed(ps))),'PB_order_invariance')
 # Size bias of arbitrary dependent three-bit distributions. Integer weights /3.
 for marks in combinations(range(10),7):
  # Stars-and-bars representation of eight weights with total 3.
  cuts=(-1,)+marks+(10,);weights=[cuts[i+1]-cuts[i]-1 for i in range(8)]
  law={b:Q(w,3) for b,w in zip(product([0,1],repeat=3),weights) if w}
  lam=sum(sum(b)*p for b,p in law.items())
  if not lam:continue
  ps=[sum(p*b[i] for b,p in law.items()) for i in range(3)];ws=defaultdict(Q);direct=defaultdict(Q)
  for b,p in law.items():direct[sum(b)]+=sum(b)*p/lam
  for i,pi in enumerate(ps):
   if not pi:continue
   for b,p in law.items():
    if b[i]:ws[sum(b)]+=(pi/lam)*(p/pi)
  check(all(ws[k]==direct[k] for k in range(4)),'dependent_size_bias_law')
  for k in range(4):check(sum(sum(b)**(k+1)*p for b,p in law.items())==lam*sum(w**k*p for w,p in ws.items()),'dependent_size_bias_moment')
 check(Q(1,2)!=Q(1,4),'pairwise_independence_insufficient')
 # Zero-bias density and identities for arbitrary centered three-point laws.
 raw=[Q(-2),Q(0),Q(3)]
 for a in range(1,10):
  for b in range(1,10-a):
   ps=[Q(a,10),Q(b,10),Q(10-a-b,10)];xs,var,pcs=zero_pieces(raw,ps)
   check(all(d>=0 for _,_,d in pcs),'zero_density_nonnegative')
   check(sum((v-u)*d for u,v,d in pcs)==1,'zero_density_normalized')
   for k in range(1,7):
    lhs=sum(x**(k+1)*p for x,p in zip(xs,ps));rhs=var*sum(d*(v**k-u**k) for u,v,d in pcs)
    check(lhs==rhs,'zero_polynomial_identity')
   absfirst=sum(d*((v*abs(v)-u*abs(u))/2) for u,v,d in pcs)
   check(absfirst==sum(abs(x)**3*p for x,p in zip(xs,ps))/(2*var),'zero_absolute_first_moment')
 for j in range(1,20):
  p=Q(j,20);xs,var,pcs=zero_pieces([Q(0),Q(1)],[1-p,p]);check(pcs==[(-p,1-p,Q(1))],'Bernoulli_zero_uniform')
  check((1-p)*Q(1,2)+p*Q(1,2)==Q(1,2),'Bernoulli_coupling_half_distance')
 # Heterogeneous Bernoulli independent sum: variance-weighted zero-bias identity.
 for n in range(1,6):
  ps=[Q(i+1,n+2) for i in range(n)];mu=sum(ps);var=sum(p*(1-p) for p in ps)
  for k in range(1,7):
   lhs=sum((sum(b)-mu)**(k+1)*q for b,q in states(ps));rhs=Q(0)
   for i,p in enumerate(ps):
    others=ps[:i]+ps[i+1:];other_mu=sum(others)
    for b,q in states(others):
     t=sum(b)-other_mu;lo=t-p;hi=t+1-p
     rhs+=p*(1-p)*q*(hi**k-lo**k)
   check(lhs==rhs,'zero_independent_sum_identity')
 # Independent-coordinate exchange pairs; all conditional identities exact.
 for n in range(1,6):
  ps=[Q(i+1,n+2) for i in range(n)];mu=sum(ps);var=sum(p*(1-p) for p in ps);lam=Q(1,n);joint=defaultdict(Q)
  for bits,mass in states(ps):
   S=sum(bits)-mu;ed=ed2=Q(0)
   for i,p in enumerate(ps):
    for z,r in [(0,1-p),(1,p)]:
     d=z-bits[i];q=r/n;ed+=q*d;ed2+=q*d*d;joint[(S,S+d)]+=mass*q
   check(ed==-lam*S,'exchange_independent_regression')
   check(ed2==(var+sum((b-p)**2 for b,p in zip(bits,ps)))/n,'exchange_independent_conditional_second')
  check(all(v==joint[(b,a)] for (a,b),v in list(joint.items())),'exchange_independent_symmetry')
 # General finite-population values; do not just test the balanced +/-1 example.
 for N in range(3,10):
  vals=[Q(i)-Q(N-1,2) for i in range(N)];v=sum(x*x for x in vals)/N
  for m in range(1,N):
   sig2=Q(m*(N-m),N-1)*v;lam=Q(N,m*(N-m));joint=defaultdict(Q);moment=Q(0);den=comb(N,m)
   for S0 in combinations(range(N),m):
    S=set(S0);T=sum(vals[i] for i in S);moment+=T*T/den;ed=ed2=Q(0)
    for i in S:
     for j in set(range(N))-S:
      d=vals[j]-vals[i];q=Q(1,m*(N-m));ed+=q*d;ed2+=q*d*d;joint[(T,T+d)]+=q/den
    check(ed==-lam*T,'without_replacement_regression')
    raw2=sum((vals[j]-vals[i])**2 for i in S for j in set(range(N))-S)/Q(m*(N-m))
    check(ed2==raw2,'without_replacement_conditional_second')
   check(moment==sig2,'without_replacement_variance')
   check(all(q==joint[(b,a)] for (a,b),q in list(joint.items())),'without_replacement_symmetry')
   check(sum((b-a)**2*q for (a,b),q in joint.items())==2*lam*sig2,'without_replacement_jump_second')
 # Overlap dynamic programming against independent bit enumeration.
 for m in range(3,11):
  for p in [Q(1,10),Q(3,10),Q(1,2)]:
   exact=[Q(0)]*m
   for bits,q in states([p]*m):exact[sum(a*b for a,b in zip(bits,bits[1:]))]+=q
   v=overlap(m,p);check(v==exact,'overlap_DP_enumeration')
   lam=(m-1)*p*p;b1=(3*m-5)*p**4;b2=2*(m-2)*p**3;tv=poisson_tv(v,lam)
   check(tv[1]<=b1+b2,'overlap_rigorous_Poisson_TV')
 # Discrete Stein equation and difference factors: finite A, untruncated target probabilities.
 for lam in map(mp.mpf,['.1','.5','1','3','7']):
  for mask in range(64):
   A={i for i in range(6) if mask>>i&1};pa=sum(mp.exp(-lam)*lam**i/mp.factorial(i) for i in A)
   f=[mp.mpf(0)]*20;f[1]=(int(0 in A)-pa)/lam;f[0]=f[1]
   for k in range(1,19):f[k+1]=(k*f[k]+int(k in A)-pa)/lam
   for k in range(19):
    check(abs(lam*f[k+1]-k*f[k]-(int(k in A)-pa))<mp.mpf('1e-70'),'Poisson_Stein_equation')
    check(abs(f[k])<=1+mp.mpf('1e-60') and abs(f[k+1]-f[k])<=1+mp.mpf('1e-60'),'Poisson_Stein_factors')
 # Explicit survival/immigration kernel backward identity.
 def ph(k,t,lam,A):
  q=mp.exp(-t);b=lam*(1-q)
  return sum(mp.binomial(k,j)*q**j*(1-q)**(k-j)*sum(mp.exp(-b)*b**(a-j)/mp.factorial(a-j) for a in A if a>=j) for j in range(k+1))
 for k,ls,ts,A in product(range(5),['.3','2'],['.1','.7','2'],[{0},{1,3},{0,2,4}]):
  lam=mp.mpf(ls);t=mp.mpf(ts);der=mp.diff(lambda u:ph(k,u,lam,A),t)
  gen=lam*(ph(k+1,t,lam,A)-ph(k,t,lam,A))+(k*(ph(k-1,t,lam,A)-ph(k,t,lam,A)) if k else 0)
  check(abs(der-gen)<mp.mpf('1e-85'),'Poisson_kernel_backward_identity')
 # Gaussian Lipschitz factors for h(x)=|x-a|, with explicit piecewise CDF kernels.
 phi=lambda x:mp.exp(-x*x/2)/mp.sqrt(2*mp.pi)
 Phi=lambda x:(1+mp.erf(x/mp.sqrt(2)))/2
 A0=lambda x:x*Phi(x)+phi(x)
 B0=lambda x:phi(x)-x*(1-Phi(x))
 for aa in [-3,-1,0,1,3]:
  a=mp.mpf(aa);eh=2*phi(a)+a*(2*Phi(a)-1)
  for j in range(-60,61):
   x=mp.mpf(j)/10;am=(1-Phi(x))/phi(x);bm=Phi(x)/phi(x);A=-A0(x) if x<=a else A0(x)-2*A0(a);B=B0(x) if x>=a else 2*B0(a)-B0(x)
   f=-am*A-bm*B;fp=-(x*am-1)*A-(x*bm+1)*B
   check(abs(f)<=1+mp.mpf('1e-80'),'normal_Stein_bounded_solution')
   check(abs(fp)<=mp.sqrt(2/mp.pi)+mp.mpf('1e-80'),'normal_Stein_first_derivative')
   check(abs(fp-x*f-(abs(x-a)-eh))<mp.mpf('1e-80'),'normal_Stein_equation')
   check(abs(am*A0(x)+bm*B0(x)-1)<mp.mpf('1e-80'),'normal_kernel_weight_identity')
   a2=(1+x*x)*am-x;b2=(1+x*x)*bm+x
   check(a2>=0 and b2>=0 and abs(a2*A0(x)+b2*B0(x)-1)<mp.mpf('1e-80'),'normal_kernel_second_weight')
   if x!=a:check(abs(mp.sign(x-a)-a2*A-b2*B)<=2+mp.mpf('1e-80'),'normal_Stein_second_derivative')
 # Four capstone outputs.
 ps=list(map(Q,['.02','.05','.08','.15']));iv=pb(ps);lam=sum(ps);itv=poisson_tv(iv,lam);budget=sum(p*p for p in ps);check(itv[1]<budget,'capstone_independent_certificate')
 m=20;p=Q(1,10);ov=overlap(m,p);ol=(m-1)*p*p;otv=poisson_tv(ov,ol);ob=(3*m-5)*p**4+2*(m-2)*p**3;check(otv[1]<ob,'capstone_overlap_certificate')
 s=100;probs=[Q(comb(s,k)**2,comb(2*s,s)) for k in range(s+1)];sig2=Q(s*s,2*s-1);sig=mp.sqrt(M(sig2));lam=Q(2,s)
 cond=[Q(4)*Q(k*k+(s-k)**2,s*s)/sig2 for k in range(s+1)];defect=sum(q*abs(1-a/(2*lam)) for q,a in zip(probs,cond))
 check(sum(probs)==1 and sum((2*k-s)*q for k,q in enumerate(probs))==0,'capstone_hypergeometric_law')
 check(sum((2*k-s)**2*q for k,q in enumerate(probs))==sig2,'capstone_hypergeometric_variance')
 check(sum(a*q for a,q in zip(cond,probs))==2*lam,'capstone_conditional_second_average')
 third=2/sig;bound=mp.sqrt(2/mp.pi)*M(defect)+third;pay=sum(M(q)*max(M(2*k-s)/sig,0) for k,q in enumerate(probs));normal=1/mp.sqrt(2*mp.pi)
 n=400;zero_pay=sum(mp.mpf(comb(n,k))*max(mp.mpf(k-200)/10,0)/mp.mpf(2)**n for k in range(n+1))
 check(abs(zero_pay-normal)<mp.mpf('.1'),'capstone_zero_bias_payoff')
 result={'status':'PASS','assertions':sum(COUNTS.values()),'checks':dict(COUNTS),'precision_digits':mp.mp.dps,'rational_TV_enclosure_degree':90,'independent':{'pmf':list(map(str,iv)),'lambda':str(sum(ps)),'budget':str(budget),'TV_interval_rational':list(map(str,itv)),'TV_decimal_approx':rep((itv[0]+itv[1])/2)},'overlap':{'bit_count':m,'indicator_count':m-1,'pmf':list(map(str,ov)),'lambda':str(ol),'b1':str((3*m-5)*p**4),'b2':str(2*(m-2)*p**3),'budget':str(ob),'P0':str(ov[0]),'TV_interval_rational':list(map(str,otv)),'TV_decimal_approx':rep((otv[0]+otv[1])/2)},'without_replacement':{'population_size':2*s,'sample_size':s,'sigma_squared':str(sig2),'lambda':str(lam),'conditional_variance_defect':str(defect),'W1_budget':rep(bound),'positive_part':rep(pay),'actual_positive_error':rep(abs(pay-normal))},'zero_bias':{'n':n,'sigma':10,'W1_budget':'1/10','positive_part':rep(zero_pay),'actual_positive_error':rep(abs(zero_pay-normal))},'normal_positive_part':rep(normal),'limitations':'Finite enumeration validates the specified constructions; universal conclusions require the proofs in the formal pages. Gaussian numerical checks are 100-digit diagnostics, not directed-rounding certificates.'}
 return result
if __name__=='__main__':
 ap=argparse.ArgumentParser();ap.add_argument('--output');a=ap.parse_args();res=run();text=json.dumps(res,ensure_ascii=False,indent=2)
 if a.output:Path(a.output).write_text(text+'\n')
 print(text)
