#!/usr/bin/env python3
"""Reproduce count-observation certificates. Python 3 + mpmath; no author QA files.
Run with --output /path/result.json; otherwise prints JSON only.
All checks remain active under python -O. Numerical checks complement, not replace,
the proofs in the accompanying pages. No data about actual people is used.
"""
import argparse, json, math
from fractions import Fraction as F
from pathlib import Path
import mpmath as mp
mp.mp.dps=90
counts={}
def check(ok,group):
    counts[group]=counts.get(group,0)+1
    if not ok: raise RuntimeError('Check failed: '+group)
def near(a,b,group,tol=mp.mpf('1e-65')):
    check(abs(a-b)<=tol*(1+abs(a)+abs(b)),group)
def val(x): return mp.mpf(x.numerator)/x.denominator if isinstance(x,F) else mp.mpf(x)
def out(x): return mp.nstr(x,55)
def pois(y,l): return mp.exp(-l)*l**y/mp.factorial(y)
def nb(y,mu,k): return mp.rf(k,y)/mp.factorial(y)*(k/(k+mu))**k*(mu/(k+mu))**y
def pos(y,l): return l**y/(mp.factorial(y)*mp.expm1(l))
def meanpos(l): return l/(-mp.expm1(-l))
def varpos(l): return meanpos(l)*(1+l-meanpos(l))
def hurdle(y,q,l): return 1-q if y==0 else q*pos(y,l)
def zipmass(y,p,l): return p+(1-p)*mp.exp(-l) if y==0 else (1-p)*pois(y,l)
def zipll(ys,p,l): return sum(mp.log(zipmass(y,p,l)) for y in ys)
def exp_bounds(x,N=90):
    # Positive x, positive exponential series, then a geometric majorant.
    if x<0 or x>=N+2: raise ValueError('Positive series requires 0 <= x < N+2')
    t=s=F(1)
    for j in range(1,N+1): t*=x/j;s+=t
    nxt=t*x/(N+1)
    return s,s+nxt/(1-x/F(N+2))
def root_numerator_bounds(l,a):
    # m(l)-a has the sign of (l-a)exp(l)+a.
    lo,hi=exp_bounds(l)
    return ((l-a)*hi+a,(l-a)*lo+a) if l<a else ((l-a)*lo+a,(l-a)*hi+a)

# Gamma mixture, normalization, moments, and derivative contracts.
nb_cases=[]
for mu in map(mp.mpf,['0.125','1','3','9']):
  for k in map(mp.mpf,['0.5','1','2','7']):
    # Positive-series moments with an explicit remaining-tail majorant.
    sums=[mp.mpf(0)]*3;p=(k/(k+mu))**k;y=0
    while True:
      for r in range(3): sums[r]+=(mp.mpf(y)**r)*p
      if y>=2:
        rprob=max(mu/(k+mu), (k+y)/(y+1)*mu/(k+mu))
        rho=rprob*((mp.mpf(y+1)/y)**2)
        if rho<1:
          nextp=p*(k+y)/(y+1)*mu/(k+mu)
          bound=(y+1)**2*nextp/(1-rho)
          if bound<mp.mpf('1e-75'):break
      y+=1;p*= (k+y-1)/y*mu/(k+mu)
      if y>30000:raise RuntimeError('Unbounded numerical series')
    near(sums[0],1,'NB normalized positive series')
    near(sums[1],mu,'NB first moment')
    near(sums[2]-sums[1]**2,mu+mu**2/k,'NB second moment')
    for a in range(9):
      near(nb(a+1,mu,k)/nb(a,mu,k),(k+a)/(a+1)*mu/(k+mu),'NB recurrence')
      eta=mp.log(mu)
      f=lambda t:mp.log(nb(a,mp.exp(t),k))
      near(mp.diff(f,eta),k*(a-mu)/(k+mu),'NB beta score')
      near(-mp.diff(f,eta,2),k*mu*(k+a)/(k+mu)**2,'NB observed curvature')
      fshape=lambda kk:mp.log(nb(a,mu,kk))
      u=sum(1/(k+j) for j in range(a))+mp.log(k/(k+mu))+(mu-a)/(k+mu)
      near(mp.diff(fshape,k),u,'NB shape score')
    nb_cases.append({'mean':out(mu),'shape':out(k),'last_y':y,'remaining_second_moment_bound':out(bound)})
# Direct integral checks use k>=1 so endpoint singularities need no special quadrature.
for k in map(mp.mpf,['1','2','3.5']):
 for mu in map(mp.mpf,['0.5','3']):
  for y in range(5):
   integ=mp.quad(lambda u:mp.exp(-mu*u)*(mu*u)**y/mp.factorial(y)*k**k/mp.gamma(k)*u**(k-1)*mp.exp(-k*u),[0,1,mp.inf])
   near(integ,nb(y,mu,k),'Gamma Poisson direct integral',mp.mpf('1e-60'))

# Truncation, derivatives and shifted rather than truncated sampling.
for l in map(mp.mpf,['0.001','0.05','0.5','1','2','4','8']):
 m=meanpos(l);v=varpos(l)
 near(mp.diff(meanpos,l)*l,v,'Truncated derivative identity')
 check(m>1 and 0<v<m,'Truncated strict moment boundaries')
 for y in range(1,15):
  f=lambda eta:mp.log(pos(y,mp.exp(eta)))
  near(mp.diff(f,mp.log(l)),y-m,'Truncated score')
  near(-mp.diff(f,mp.log(l),2),v,'Truncated curvature')
  near(y*pois(y,l)/l,pois(y-1,l),'Size biased Poisson shift')
  near(pos(y,l),pois(y,l)/(1-pois(0,l)),'Truncated normalization factor')
 for j in range(1,10):
  q=mp.mpf(j)/10;mu=q*m;vv=q*v+q*(1-q)*m*m
  near(vv,mu*(1+l-mu),'Hurdle total variance')
  for y in range(1,9):near(hurdle(y,q,l)/q,pos(y,l),'Positive conditioning loses gate')
  near(hurdle(2,q,l)/hurdle(1,q,l),l/2,'Positive ratio identifies lambda')
  p=1-q/(1-mp.exp(-l))
  if p>=0:
   for y in range(15):near(zipmass(y,p,l),hurdle(y,q,l),'ZIP hurdle equality')
  else:check(q>1-mp.exp(-l),'Hurdle excludes ZIP')
# Regression chain-rule derivatives with independent gate and count coefficients.
for x in map(mp.mpf,['-1','0','0.6','2']):
 for beta,gamma in [(mp.mpf('0.7'),mp.mpf('-0.4')),(mp.mpf('-1'),mp.mpf('0.5'))]:
  q=lambda z:1/(1+mp.exp(-gamma*z));lam=lambda z:mp.exp(beta*z)
  hmean=lambda z:q(z)*meanpos(lam(z))
  zmean=lambda z:(1-q(z))*lam(z)
  near(mp.diff(lambda z:mp.log(hmean(z)),x),(1-q(x))*gamma+(1+lam(x)-meanpos(lam(x)))*beta,'Hurdle marginal derivative')
  near(mp.diff(lambda z:mp.log(zmean(z)),x),beta-q(x)*gamma,'ZIP marginal derivative')

# EM Q construction and monotonicity across distinct observed zero patterns.
em_cases=[]
for ys in ([0,0,0,1,2],[0,0,0,0,1,2,3,4],[0,1,1,1,4,8],[0,0,0,0,0,1,1,9]):
 for p0,l0 in [(mp.mpf(1)/3,mp.log(2)),(mp.mpf('0.05'),mp.mpf(4)),(mp.mpf('0.8'),mp.mpf('0.2'))]:
  p,l=p0,l0
  for it in range(60):
   tau=[p/(p+(1-p)*mp.exp(-l)) if y==0 else mp.mpf(0) for y in ys]
   pn=sum(tau)/len(ys);ln=sum(ys)/(len(ys)-sum(tau))
   check(0<=pn<1 and ln>0,'EM admissibility')
   check(zipll(ys,pn,ln)>=zipll(ys,p,l)-mp.mpf('1e-75'),'EM observed monotonicity')
   # Complete-data Q first-order conditions, away from mixture endpoint.
   near(sum(t/pn-(1-t)/(1-pn) for t in tau),0,'EM mixing M-step score')
   near(sum((1-t)*(y/ln-1) for y,t in zip(ys,tau)),0,'EM count M-step score')
   p,l=pn,ln
  em_cases.append({'data':ys,'initial':[out(p0),out(l0)],'after_60':[out(p),out(l)],'loglik':out(zipll(ys,p,l))})

# Exact shared-Gamma count law: verify every allocation on many total-count slices.
shared=[]
for s in range(51):
 masses=[F(math.factorial(s+1),math.factorial(a)*math.factorial(s-a))*F(4*3**a*6**(s-a),11**(s+2)) for a in range(s+1)]
 target=F((s+1)*4*9**s,11**(s+2))
 check(sum(masses)==target,'Shared exact total law')
 check(sum(F(a)*p for a,p in enumerate(masses))==F(s,3)*target,'Shared exact conditional mean')
 check(sum(F(a*(a-1))*p for a,p in enumerate(masses))==F(s*(s-1),9)*target,'Shared exact conditional factorial moment')
 for a,p in enumerate(masses):
  cond=F(math.comb(s,a)*2**(s-a),3**s)
  check(p==target*cond,'Shared exact binomial split')
 shared.append({'total':s,'mass':str(target)})
check(F(15,2)+24+2*9==F(99,2),'Shared total variance')
check(F(15,2)+24==F(63,2),'Independent total variance')
check(F(4,25)*F(1,16)==F(1,100),'Independent joint zero')
check(F(4,121)>F(1,100),'Shared joint zero differs')

# Exact rational exponential enclosures for the capstone root.
a=F(5,2);L=F(223161188402,10**11);H=F(223161188403,10**11)
loL,hiL=root_numerator_bounds(L,a);loH,hiH=root_numerator_bounds(H,a)
check(hiL<0 and loH>0,'Exact exponential root bracket')
check(H-L==F(1,10**11),'Exact bracket width')
lhat=mp.findroot(lambda x:meanpos(x)-val(a),(mp.mpf(2),mp.mpf(3)))
check(val(L)<lhat<val(H),'Root within exact bracket')
pihat=1-mp.mpf(5)/(4*lhat);qhat=mp.mpf('0.5')
for y in range(50):near(zipmass(y,pihat,lhat),hurdle(y,qhat,lhat),'Capstone full-law equivalence')
ys=[0,0,0,0,1,2,3,4];empmean=F(sum(ys),len(ys));empvar=sum(F(y*y,len(ys)) for y in ys)-empmean**2
k=empmean**2/(empvar-empmean)
check(empmean==F(5,4) and empvar==F(35,16) and k==F(5,3),'Capstone NB moments')
lnb0=nb(0,val(empmean),val(k));near(lnb0,(mp.mpf(4)/7)**(mp.mpf(5)/3),'Capstone NB zero')
under_l=mp.log(2);under_q=mp.mpf(9)/10;under_mu=under_q*meanpos(under_l);under_var=under_mu*(1+under_l-under_mu)
near(1-under_q/(1-mp.exp(-under_l)),-mp.mpf(4)/5,'Capstone negative mixing weight')
check(0<under_var<under_mu,'Capstone underdispersion')
near(under_mu,mp.mpf(9)/5*mp.log(2),'Capstone hurdle mean')
check(F(sum(y-1 for y in [1,2,3,4]),4)==F(3,2),'Capstone event weighted MLE')
near(mp.diff(lambda t:sum(mp.log(pois(y-1,mp.exp(t))) for y in [1,2,3,4]),mp.log(mp.mpf('1.5'))),0,'Event sampling likelihood score')
# Page-level examples.
near(nb(0,mp.mpf(3),mp.mpf(2)),mp.mpf(4)/25,'Page NB zero')
near(nb(1,mp.mpf(3),mp.mpf(2)),mp.mpf(24)/125,'Page NB one')
near(nb(2,mp.mpf(3),mp.mpf(2)),mp.mpf(108)/625,'Page NB two')
check(sum(F(2*(y-1),3) for y in [0,1,5])==2,'Page NB scoring numerator')
check(3*F(2,3)==2,'Page NB scoring information')
l2=mp.findroot(lambda x:meanpos(x)-2,(1,2));near(meanpos(l2),2,'Page truncated root')

result={'status':'PASS','precision_decimal_digits':mp.mp.dps,'checks':sum(counts.values()),'check_groups':counts,
 'scope':'Independent finite computations and exact rational brackets, supplementing analytic proofs; no simulated evidence or confidence claims.',
 'nb_series_cases':nb_cases,'em_cases':em_cases,
 'capstone':{'lambda':out(lhat),'pi':out(pihat),'q':'1/2','maximum_log_likelihood':out(zipll(ys,pihat,lhat)),
 'lambda_bracket':[str(L),str(H)],'bracket_width':str(H-L),'low_endpoint_numerator_upper':str(hiL),'high_endpoint_numerator_lower':str(loH),
 'em_start_loglik':out(zipll(ys,mp.mpf(1)/3,mp.log(2))),'em_first_pi':'1/4','em_first_lambda':'5/3','em_first_loglik':out(zipll(ys,mp.mpf(1)/4,mp.mpf(5)/3)),
 'NB_moment_shape':'5/3','NB_zero':out(lnb0),'underdispersed_mean':out(under_mu),'underdispersed_variance':out(under_var),
 'shared_covariance':'9','shared_total_variance':'99/2','independent_total_variance':'63/2','shared_joint_zero':'4/121','independent_joint_zero':'1/100','event_sample_lambda':'3/2'},
 'page_truncated_mean_2_root':out(l2),'exact_shared_total_slices':shared}
ap=argparse.ArgumentParser();ap.add_argument('--output',type=Path);args=ap.parse_args();text=json.dumps(result,ensure_ascii=False,indent=2)+'\n'
if args.output:args.output.parent.mkdir(parents=True,exist_ok=True);args.output.write_text(text)
else:print(text,end='')
