#!/usr/bin/env python3
"""Rational enclosures for sampling and tanh-sinh quadrature.

All certificate endpoints and analytic error budgets use Fraction. Decimal is
used only for readable output. The mathematical strip constants are proved in
the accompanying pages; finite checks do not prove an arbitrary function's
analytic extension or replace those proofs.
"""
from fractions import Fraction as Q
from decimal import Decimal,localcontext
from collections import Counter
from functools import lru_cache
from math import isqrt
from pathlib import Path
import argparse,json,random

BITS=640
GRID=1<<BITS
COUNT=Counter()
def check(value,category):
 if not value:raise AssertionError(category)
 COUNT[category]+=1

def down(x):
 x=Q(x);return Q(x.numerator*GRID//x.denominator,GRID)
def up(x):return -down(-Q(x))
class I:
 def __init__(self,lo,hi=None):
  self.lo=down(lo);self.hi=up(lo if hi is None else hi)
  if self.lo>self.hi:raise ValueError('reversed interval')
 def __add__(self,b):
  b=iv(b);return I(self.lo+b.lo,self.hi+b.hi)
 __radd__=__add__
 def __neg__(self):return I(-self.hi,-self.lo)
 def __sub__(self,b):return self+-iv(b)
 def __rsub__(self,b):return iv(b)+-self
 def __mul__(self,b):
  b=iv(b);p=[x*y for x in [self.lo,self.hi] for y in [b.lo,b.hi]]
  return I(min(p),max(p))
 __rmul__=__mul__
 def inv(self):
  if self.lo<=0<=self.hi:raise ZeroDivisionError('interval contains zero')
  return I(1/self.hi,1/self.lo)
 def __truediv__(self,b):return self*iv(b).inv()
 def __rtruediv__(self,b):return iv(b)*self.inv()
 def square(self):return I(0 if self.lo<=0<=self.hi else min(self.lo**2,self.hi**2),max(self.lo**2,self.hi**2))
 def contains(self,x):return self.lo<=x<=self.hi
 def expand(self,e):return I(self.lo-e,self.hi+e)
 def width(self):return self.hi-self.lo
 def pair(self):return [str(self.lo),str(self.hi)]
def iv(x):return x if isinstance(x,I) else I(x)

@lru_cache(maxsize=None)
def exp_positive(q):
 q=Q(q)
 if q<0:raise ValueError('nonnegative argument required')
 if not q:return I(1)
 k=0
 while q>Q(1,2):q/=2;k+=1
 term=I(1);total=I(1);N=120
 for j in range(1,N+1):
  term=term*q/j;total=total+term
 # Every later term ratio is at most q/(N+2).
 remainder=term.hi*q/(N+1)/(1-q/(N+2))
 result=I(total.lo,total.hi+remainder)
 for _ in range(k):result=result.square()
 return result

def exp_point(q):return exp_positive(q) if q>=0 else exp_positive(-q).inv()
def exp(x):
 x=iv(x);return I(exp_point(x.lo).lo,exp_point(x.hi).hi)
def sinh(x):return (exp(x)-exp(-iv(x)))/2
def cosh(x):return (exp(x)+exp(-iv(x)))/2

def atan_series(q,N=150):
 q=Q(q);s=Q(0)
 for j in range(N):s+=(-1)**j*q**(2*j+1)/(2*j+1)
 t=(-1)**N*q**(2*N+1)/(2*N+1)
 return I(min(s,s+t),max(s,s+t))
def pi_interval():
 t=Q(1,5);t=2*t/(1-t*t);t=2*t/(1-t*t)
 check((t-Q(1,239))/(1+t/239)==1,'Machin_tangent_identity')
 return 16*atan_series(Q(1,5))-4*atan_series(Q(1,239))
PI=pi_interval()

def display(x,digits=50):
 with localcontext() as ctx:
  ctx.prec=digits
  return str(Decimal(x.numerator)/Decimal(x.denominator))

def interval_tests():
 check(PI.lo>3 and PI.hi<4,'pi_coarse_bounds')
 ee=exp(I(1));check(ee.lo>Q(8,3) and ee.hi<Q(11,4),'e_coarse_bounds')
 rng=random.Random(11108)
 for _ in range(250):
  a,b,c,d=sorted([Q(rng.randrange(-1000,1001),rng.randrange(1,50)) for j in range(4)])
  A=I(a,b);B=I(c,d)
  for x,y in [(a,c),(a,d),(b,c),(b,d),((a+b)/2,(c+d)/2)]:
   check((A+B).contains(x+y) and (A*B).contains(x*y),'rounded_interval_arithmetic')
   if not c<=0<=d:check((A/B).contains(x/y),'rounded_interval_division')
 for q in [Q(0),Q(1,7),Q(1),Q(4),Q(40),Q(320)]:
  # Independent identities check enclosures, not a black-box exp evaluation.
  E=exp(I(q));F=exp(I(-q));check((E*F).contains(1),'exp_reciprocal_enclosures')
  E2=exp(I(2*q));sq=E.square();check(max(E2.lo,sq.lo)<=min(E2.hi,sq.hi),'exp_doubling_enclosures')
 try:I(-1,1).inv()
 except ZeroDivisionError:check(True,'reject_zero_denominator_interval')
 else:raise AssertionError('failed denominator rejection')

def periodic_checks():
 for N in range(1,31):
  for k in range(-4*N,4*N+1):
   # Sum of N roots of unity equals N precisely when the frequency is 0 mod N.
   # Checking the integer exponent labels verifies the exact filter mechanism.
   residues=[k*j%N for j in range(N)]
   check((len(set(residues))==1)==(k%N==0),'periodic_frequency_filter_labels')
  r=Q(1,2);rho=r**N
  unshifted=(1+rho)/(1-rho);shifted=(1-rho)/(1+rho)
  check(unshifted-1==2*rho/(1-rho) and shifted-1==-2*rho/(1+rho),'Poisson_kernel_exact_errors')
 eps=Q(1,1000);N=next(n for n in range(1,50) if Q(2,2**n-1)<=eps)
 check(N==11 and Q(2,2**10-1)>eps,'periodic_minimal_N')
 for n in [3,8,17]:
  for size in [n,2*n]:
   check(all((2*n*j)%size==0 for j in range(size)),'identical_mesh_values_do_not_certify_integral')
 return {'r':'1/2','tolerance':'1/1000','minimal_N':11,'N10_error':str(Q(2,1023)),'N11_error':str(Q(2,2047)),'half_shift_N11_error':str(-Q(2,2049))}

def gaussian_tail(a,b):
 # Sum_{j>=0} exp(-a (b+j)^2), b>0, using j²>=j.
 return exp(-a*b*b)/(1-exp(-a*(2*b+1)))
def gaussian_checks():
 out=[];a=PI/100
 for shift in [Q(0),Q(1,3)]:
  direct=I(0);N=100
  for k in range(-N,N+1):direct=direct+exp(-a*(Q(k)+shift)**2)
  tail=gaussian_tail(a,N+1+shift)+gaussian_tail(a,N+1-shift)
  direct=direct.expand(tail.hi)
  q=exp(-100*PI)
  # Keep frequency m=1; m>=2 is bounded by 2 exp(-400π)/(1-exp(-500π)).
  dual=10*(1+(2 if shift==0 else -1)*q)
  rest=20*exp(-400*PI)/(1-exp(-500*PI));dual=dual.expand(rest.hi)
  check(max(direct.lo,dual.lo)<=min(direct.hi,dual.hi),'Gaussian_two_scales_enclosures_overlap')
  check(direct.lo>10 if shift==0 else direct.hi<10,'Gaussian_shift_changes_correction_sign')
  check(dual.lo>10 if shift==0 else dual.hi<10,'Gaussian_dual_signed_correction')
  coarse=Q(20)*(Q(3,8)**300)/(1-Q(3,8)**900)
  check(direct.lo>10-coarse and direct.hi<10+coarse,'Gaussian_coarse_rational_tolerance')
  out.append({'shift':str(shift),'spatial_nodes':201,'direct_interval':direct.pair(),'dual_interval':dual.pair(),'signed_offset_display':display((dual.lo+dual.hi)/2-10),'coarse_error_radius':str(coarse)})
 return out

def sqrt_positive(q):
 q=Q(q)
 if q<0:raise ValueError('nonnegative square root required')
 n=isqrt(q.numerator*GRID*GRID//q.denominator)
 lo=Q(n,GRID);hi=lo if lo*lo==q else Q(n+1,GRID)
 check(lo*lo<=q<=hi*hi,'rational_square_root_enclosure')
 return I(lo,hi)

def de_value(t,model):
 t=I(t);a=PI/2;u=a*sinh(t);den=cosh(u);base=a*cosh(t)/den
 if model==0:return base
 mapped=sinh(u)/den
 return base*mapped.square() if model==2 else base/(5-mapped)

def de_certificate(model,steps,N,tolerance=Q(1,10**10)):
 contracts={0:(10,40,30,64,Q(1)),2:(12,48,36,1024,Q(1)),3:(10,40,30,64,Q(1,4))}
 if model not in contracts:raise ValueError('only three proved input models')
 m,n,exponent,numerator,tail_factor=contracts[model]
 if (steps,N)!=(m,n):raise ValueError('unproved step/window contract')
 h=Q(1,steps);L=N*h
 total=I(0);positive=I(0)
 for k in range(-N,N+1):
  term=h*de_value(Q(k,steps),model);total=total+term
  if k>0:positive=positive+term
 disc=Q(numerator)/(Q(8,3)**exponent-1);tail=tail_factor*22*Q(3,8)**37
 check(disc+tail<Q(1,10**10),'DE_analytic_error_budget')
 check(total.width()<Q(1,10**60),'DE_finite_evaluation_width')
 result=total.expand(disc+tail)
 reference=PI if model==0 else PI/2 if model==2 else PI/sqrt_positive(24)
 check(result.lo<reference.lo<=reference.hi<result.hi,'DE_certificate_contains_closed_form')
 if model==3:
  wrong=(2*positive+h*de_value(Q(0),model)).expand(disc+tail)
  check(wrong.lo>reference.hi,'reject_false_even_symmetry_on_asymmetric_input')
 radius=result.width()/2
 check(radius<tolerance,'DE_total_radius_under_tolerance')
 labels={0:'(1-x^2)^(-1/2)',2:'x^2(1-x^2)^(-1/2)',3:'1/((5-x)sqrt(1-x^2))'}
 return {'model':labels[model],'h':str(h),'N':N,'nodes':2*N+1,'strip_half_width':'pi/6','one_boundary_L1_bound':512 if model==2 else 32,'finite_sum_interval':total.pair(),'discretization_bound':str(disc),'tail_bound':str(tail),'integral_interval':result.pair(),'midpoint_display':display((result.lo+result.hi)/2),'certified_radius_display':display(radius,20),'finite_sum_width_display':display(total.width(),20),'status':'CERTIFIED'}

def main():
 COUNT.clear()
 pi_interval()  # Replay the import-time Machin check for this run.
 p=argparse.ArgumentParser();p.add_argument('--output',type=Path,required=True);args=p.parse_args()
 interval_tests();periodic=periodic_checks();gaussian=gaussian_checks()
 de=[de_certificate(0,10,40),de_certificate(2,12,48),de_certificate(3,10,40)]
 for inputs in [(0,10,4),(2,10,40),(4,10,40)]:
  try:de_certificate(*inputs)
  except ValueError:check(True,'reject_unproved_DE_contract')
  else:raise AssertionError('unproved contract accepted')
 out={'status':'PASS','assertions':sum(COUNT.values()),'categories':dict(sorted(COUNT.items())),'fixed_point_bits':BITS,'pi_interval':PI.pair(),'periodic':periodic,'Gaussian':gaussian,'double_exponential':de,'scope':'Fraction interval endpoints certify the three stated models using the proved analytic constants. Decimal strings are display only; finite tests do not certify arbitrary functions.'}
 args.output.parent.mkdir(parents=True,exist_ok=True);args.output.write_text(json.dumps(out,ensure_ascii=False,indent=2)+'\n');print(json.dumps({'status':'PASS','assertions':out['assertions'],'DE_radii':[r['certified_radius_display'] for r in de]}))
if __name__=='__main__':main()
