#!/usr/bin/env python3
"""Exact causal-polynomial and fractional-memory certificates, standard library only.
The 80-place interval primitives are reused from the A15 tomography certificate.
Finite tests validate algebra and interval operations; the pages prove continuous results.
"""
import argparse
from dataclasses import dataclass
from fractions import Fraction as Q
from math import factorial, isqrt, comb
from pathlib import Path
import json

DIGITS=80
SCALE=10**DIGITS
CHECKS=0

def check(ok,label):
    global CHECKS
    CHECKS+=1
    if not ok:raise ArithmeticError(label)

def floorq(q):return q.numerator//q.denominator

def ceilq(q):return -((-q.numerator)//q.denominator)

@dataclass(frozen=True)
class Interval:
    lo:int
    hi:int
    def __post_init__(self):
        if self.lo>self.hi:raise ValueError('reversed interval')
    @staticmethod
    def rational(x):
        q=Q(x)*SCALE
        return Interval(floorq(q),ceilq(q))
    def __add__(self,other):
        other=as_interval(other)
        return Interval(self.lo+other.lo,self.hi+other.hi)
    __radd__=__add__
    def __neg__(self):return Interval(-self.hi,-self.lo)
    def __sub__(self,other):return self+-as_interval(other)
    def __rsub__(self,other):return as_interval(other)+-self
    def __mul__(self,other):
        other=as_interval(other)
        p=[a*b for a in [self.lo,self.hi] for b in [other.lo,other.hi]]
        return Interval(min(p)//SCALE,-((-max(p))//SCALE))
    __rmul__=__mul__
    def reciprocal(self):
        if self.lo<=0<=self.hi:raise ZeroDivisionError('interval contains zero')
        return Interval(floorq(Q(SCALE*SCALE,self.hi)),ceilq(Q(SCALE*SCALE,self.lo)))
    def __truediv__(self,other):return self*as_interval(other).reciprocal()
    def __rtruediv__(self,other):return as_interval(other)*self.reciprocal()
    def __pow__(self,n):
        if not isinstance(n,int) or n<0:raise ValueError('nonnegative integer exponent')
        result=Interval.rational(1)
        for _ in range(n):result=result*self
        return result
    def abs_upper(self):return max(abs(self.lo),abs(self.hi))
    def absolute(self):
        if self.lo>=0:return self
        if self.hi<=0:return -self
        return Interval(0,self.abs_upper())
    def contains(self,q):return self.lo<=Q(q)*SCALE<=self.hi
    def intersects(self,other):return self.lo<=other.hi and other.lo<=self.hi
    def record(self):return {'lower':decimal(self.lo),'upper':decimal(self.hi)}

def as_interval(x):return x if isinstance(x,Interval) else Interval.rational(x)
def decimal(n):
    sign='-' if n<0 else '';a,b=divmod(abs(n),SCALE)
    return sign+str(a)+'.'+str(b).zfill(DIGITS)

def atan_reciprocal(m,terms):
    # Alternating series, decreasing magnitudes; adjacent sums bracket atan(1/m).
    a=sum((Q(1,(2*j+1)*m**(2*j+1)) if j%2==0 else -Q(1,(2*j+1)*m**(2*j+1))) for j in range(terms))
    next_term=Q((-1)**terms,(2*terms+1)*m**(2*terms+1))
    return Interval(floorq(min(a,a+next_term)*SCALE),ceilq(max(a,a+next_term)*SCALE))

PI=16*atan_reciprocal(5,100)-4*atan_reciprocal(239,24)
ROOT2=Interval(isqrt(2*SCALE*SCALE),isqrt(2*SCALE*SCALE)+1)

def sqrt_interval(x):
    x=as_interval(x)
    if x.lo<0:raise ValueError('nonnegative square root input required')
    a=isqrt(x.lo*SCALE);b=isqrt(x.hi*SCALE)
    if b*b<x.hi*SCALE:b+=1
    return Interval(a,b)

def exp_negative(x):
    """exp(-x), x>=0. Positive Taylor bracket after range reduction."""
    x=as_interval(x)
    if x.lo<0:raise ValueError('nonnegative argument required')
    if x.hi==0:return as_interval(1)
    y=x;n=0
    while y.hi>SCALE//2:y=y/2;n+=1
    term=as_interval(1);s=term
    for j in range(1,81):term=term*y/j;s=s+term
    # For all 0<=y<=1/2, ratios after term 81 are <=1/(2*82).
    tail=Q(1,2)**81/factorial(81)/(1-Q(1,164))
    s=Interval(s.lo,s.hi+ceilq(tail*SCALE));ans=s.reciprocal()
    for _ in range(n):ans=ans*ans
    return ans


def natural(n,name):
    if isinstance(n,bool) or not isinstance(n,int) or n<0:raise ValueError(name)
    return n

def clean(p):return {Q(e):Q(c) for e,c in p.items() if Q(c)}

def add(a,b):
    out=clean(a)
    for e,c in b.items():out[Q(e)]=out.get(Q(e),Q(0))+Q(c)
    return clean(out)

def scale(a,c):return clean({e:v*Q(c) for e,v in a.items()})

def ordinary(p):
    q=clean(p)
    if any(e.denominator!=1 or e<0 for e in q):raise ValueError('ordinary polynomial exponent')
    return {int(e):c for e,c in q.items()}

def multiply(a,b):
    a,b=ordinary(a),ordinary(b);out={}
    for i,c in a.items():
        for j,d in b.items():out[i+j]=out.get(i+j,Q(0))+c*d
    return ordinary(out)

def power(p,n):
    natural(n,'polynomial power');out={0:Q(1)}
    for _ in range(n):out=multiply(out,p)
    return out

def evaluate(p,t):return sum((c*Q(t)**e for e,c in ordinary(p).items()),Q(0))

def kernel_checked(k):
    out={}
    for (i,j),c in k.items():
        natural(i,'kernel t exponent');natural(j,'kernel s exponent')
        if Q(c):out[i,j]=Q(c)
    return out

def apply_kernel(k,p):
    out={};p=ordinary(p)
    for (i,j),c in kernel_checked(k).items():
        for q,b in p.items():
            e=i+j+q+1;out[e]=out.get(e,Q(0))+c*b/Q(j+q+1)
    return ordinary(out)

def kernel_bound(k,T):
    k=kernel_checked(k);T=Q(T)
    if T<=0:raise ValueError('positive finite time window')
    if all(i+j<=1 for i,j in k):
        # An affine function on the triangle is a convex combination of vertex values.
        return max(abs(sum((c*t**i*s**j for (i,j),c in k.items()),Q(0)))
                   for t,s in [(Q(0),Q(0)),(T,Q(0)),(T,T)])
    return sum((abs(c)*T**(i+j) for (i,j),c in k.items()),Q(0))

def tail_budget(B,x,N):
    B,x=Q(B),Q(x);natural(N,'Picard index')
    if B<0 or x<0 or x>=N+2:raise ValueError('tail requires B,x>=0 and x<N+2')
    return B*x**(N+1)/factorial(N+1)/(1-x/Q(N+2))

def model_bounds(k,f,lam,T):
    T=Q(T);M=kernel_bound(k,T);f=ordinary(f)
    B=sum((abs(c)*T**e for e,c in f.items()),Q(0))
    return M,B,abs(Q(lam))*M*T

def make_picard(k,f,lam,T,N):
    natural(N,'Picard index');f=ordinary(f);k=kernel_checked(k);lam=Q(lam);T=Q(T)
    M,B,x=model_bounds(k,f,lam,T);log=[f]
    for _ in range(N):log.append(ordinary(add(f,scale(apply_kernel(k,log[-1]),lam))))
    residual=ordinary(add(add(log[-1],scale(f,-1)),scale(apply_kernel(k,log[-1]),-lam)))
    return dict(kernel=k,f=f,lam=lam,T=T,N=N,M=M,B=B,x=x,log=log,residual=residual,
                error_upper=tail_budget(B,x,N))

def verify_picard(c):
    k,f,lam,T,N=c['kernel'],ordinary(c['f']),Q(c['lam']),Q(c['T']),c['N']
    natural(N,'Picard index');M,B,x=model_bounds(k,f,lam,T)
    if len(c['log'])!=N+1 or c['log'][0]!=f:raise ValueError('initial iterate or log length')
    for n in range(N):
        expected=ordinary(add(f,scale(apply_kernel(k,c['log'][n]),lam)))
        if ordinary(c['log'][n+1])!=expected:raise ValueError('Picard transition coefficients')
    r=ordinary(add(add(c['log'][-1],scale(f,-1)),scale(apply_kernel(k,c['log'][-1]),-lam)))
    if r!=ordinary(c['residual']):raise ValueError('residual coefficients')
    if (c['M'],c['B'],c['x'])!=(M,B,x):raise ValueError('whole-triangle model bounds')
    if Q(c['error_upper'])!=tail_budget(B,x,N):raise ValueError('uncertified tail budget')
    return True

def fractional_integral(p,a):
    a=Q(a);p=clean(p)
    if a<0 or any(e<=-1 for e in p):raise ValueError('integral order/exponents')
    return clean({e+a:c for e,c in p.items()})

def caputo(p,a):
    a=Q(a);p=clean(p)
    if not 0<a<1 or any(e<0 for e in p):raise ValueError('AC power sum and 0<alpha<1')
    return clean({e-a:c for e,c in p.items() if e>0})

def verify_caputo(p,a,u0,lam,f):
    p=clean(p);f=clean(f)
    if p.get(Q(0),Q(0))!=Q(u0):raise ValueError('initial trace')
    residual=add(add(caputo(p,a),scale(p,-Q(lam))),scale(f,-1))
    if residual:raise ValueError('fractional equation coefficient mismatch')
    return True

def sqrtpi():return sqrt_interval(PI)

def half_basis(j,t):
    natural(j,'half exponent index');t=Q(t)
    if t<0:raise ValueError('nonnegative time')
    if j%2==0:return as_interval(t**(j//2)/factorial(j//2))
    k=(j+1)//2
    return sqrt_interval(t)*t**(k-1)*Q(4**k*factorial(k),factorial(2*k))/sqrtpi()

def power_sum_value(p,t):
    ans=as_interval(0)
    for e,c in clean(p).items():
        j=2*e
        if j<0 or j.denominator!=1:raise ValueError('numeric evaluator needs nonnegative half-integers')
        ans+=c*half_basis(int(j),t)
    return ans

def half_ml(x,m,positive=False):
    """Enclose E_(1/2)(+/-x). Two positive subseries and geometric tail checks."""
    x=Q(x);natural(m,'paired truncation index')
    if x<0 or x*x>=m+2:raise ValueError('x>=0 and x^2<m+2 required')
    A=Q(1);B=2*x/sqrtpi();even=as_interval(1);odd=B
    for j in range(m):
        A*=x*x/Q(j+1);B=B*x*x/Q(2*j+3,2);even+=A;odd+=B
    A_next=A*x*x/Q(m+1);B_next=B*x*x/Q(2*m+3,2)
    tailA=as_interval(A_next/(1-x*x/Q(m+2)))
    tailB=B_next/(1-x*x/Q(2*m+5,2))
    if positive:
        finite=even+odd;interval=Interval(finite.lo,finite.hi+tailA.hi+tailB.hi)
    else:
        finite=even-odd;interval=Interval(finite.lo-tailB.hi,finite.hi+tailA.hi)
    return dict(interval=interval,finite=finite,tail_even=tailA,tail_odd=tailB,m=m,x=x)

def long_tail(x):
    x=Q(x)
    if x<=0:raise ValueError('positive tail argument')
    upper=1/(sqrtpi()*x);lower=upper*max(Q(0),1-Q(1,2*x*x))
    return Interval(lower.lo,upper.hi)

def reject(call,label):
    try:call()
    except (ValueError,ZeroDivisionError):check(True,label);return
    raise ArithmeticError('invalid certificate accepted: '+label)

def poly_record(p):return [{'exponent':str(e),'coefficient':str(c)} for e,c in sorted(p.items())]

def interval_record(z):return z.record()

def fibonacci_solution(t,N):
    """Independent Taylor certificate for u''=u'+u, u(0)=u'(0)=1."""
    t=Q(t);natural(N,'Taylor index')
    if t<0 or 2*t>=N+2:raise ValueError('Taylor tail ratio')
    a,b=1,1;total=Q(0)
    for n in range(N+1):
        total+=a*t**n/Q(factorial(n));a,b=b,a+b
    first=a*t**(N+1)/Q(factorial(N+1))
    tail=first/(1-2*t/Q(N+2))
    return Interval(floorq(total*SCALE),ceilq((total+tail)*SCALE))

def run():
    global CHECKS
    CHECKS = 0
    # Independent integer/rational containment checks for the primitive operations.
    check(3*SCALE<PI.lo<PI.hi<4*SCALE,'Machin pi bracket')
    for n in range(1,55):
        for d in [1,3,11]:
            q=Q(n,d);s=sqrt_interval(q)
            check(Q(s.lo,SCALE)**2<=q<=Q(s.hi,SCALE)**2,'integer-square radical enclosure')
    for a in [Q(-7,3),Q(-1,7),Q(0),Q(3,11),Q(23,4)]:
        for b in [Q(-2),Q(1,13),Q(2,3)]:
            A,B=as_interval(a),as_interval(b)
            check((A+B).contains(a+b) and (A*B).contains(a*b),'outward arithmetic')
            check((A/B).contains(a/b),'outward rational division')
    k={(0,0):Q(1),(1,0):Q(1),(0,1):Q(-1)};f={0:Q(1)}
    eps=Q(1,10**5);N=1
    while tail_budget(1,2,N)>eps:N+=1
    main=make_picard(k,f,1,1,N);check(verify_picard(main),'main Picard certificate')
    check(main['log'][2]=={0:Q(1),1:Q(1),2:Q(1),3:Q(1,3),4:Q(1,24)},'main second iterate')
    check(main['error_upper']<=eps and tail_budget(1,2,N-1)>eps,'first adequate a-priori bound')
    true_one=fibonacci_solution(1,100);actual= true_one-evaluate(main['log'][-1],1)
    check(actual.lo>0 and actual.hi<main['error_upper']*SCALE,'independent ODE Taylor solution validates Picard error')
    # Three exact families verify the integral engine without using its recurrence formula.
    for n in range(13):
        for c in [Q(-2),Q(0),Q(1,3),Q(1),Q(3)]:
            for typ in ['constant','2s','difference']:
                kern=({(0,0):c} if typ=='constant' else {(0,1):2*c} if typ=='2s' else {(1,0):c,(0,1):-c})
                p={0:Q(1)}
                for _ in range(n):p=apply_kernel(kern,p)
                if typ=='constant':expected=ordinary({n:c**n/Q(factorial(n))})
                elif typ=='2s':expected=ordinary({2*n:c**n/Q(factorial(n))})
                else:expected=ordinary({2*n:c**n/Q(factorial(2*n))})
                check(p==expected,'closed-form iterated kernel family')
    # General signed nonconvolution kernels and changing forcing/time/feedback.
    models=0
    for i in range(1,19):
        kern={(0,0):Q(i-7,5),(1,0):Q(2-i,7),(0,1):Q(i,9),(1,1):Q((-1)**i,4)}
        forcing={0:Q(i,3),1:Q(-1,i+1),2:Q(1,5)}
        lam=Q((-1)**i,i+1);T=Q(i%4+1,4)
        cert=make_picard(kern,forcing,lam,T,10);check(verify_picard(cert),'varying polynomial model')
        # Linear residual identity from a separately propagated omitted term.
        term=forcing
        for _ in range(11):term=scale(apply_kernel(kern,term),lam)
        check(cert['residual']==ordinary(scale(term,-1)),'exact omitted-term residual')
        for j in range(17):
            t=T*j/16
            check(abs(evaluate(cert['residual'],t))<=cert['B']*cert['x']**11/factorial(11),'diagnostic residual bound')
        models+=1
    from copy import deepcopy
    for field in ['log','residual','M','error_upper']:
        c=deepcopy(main)
        if field=='log':c['log'][2][2]+=1
        elif field=='residual':c['residual'][0]=1
        else:c[field]+=1
        reject(lambda c=c:verify_picard(c),'corrupt '+field)
    reject(lambda:tail_budget(1,4,2),'noncontracting geometric tail')
    reject(lambda:make_picard(k,f,1,0,12),'zero time window')
    reject(lambda:make_picard(k,f,1,1,True),'boolean iteration index')
    # Fractional operations use normalized phi_p coordinates, never guessed Gamma samples.
    frac_models=0
    for a in [Q(1,5),Q(1,3),Q(1,2),Q(2,3),Q(4,5)]:
        for b in [Q(1,7),Q(1,3),Q(1,2),Q(5,4)]:
            for j in range(1,17):
                p={Q(0):Q(j,3),Q(j,6):Q(-2),Q(j+1,3):Q(3)}
                check(fractional_integral(fractional_integral(p,a),b)==fractional_integral(p,a+b),'integral order semigroup')
                restored=fractional_integral(caputo(p,a),a)
                check(restored==add(p,{Q(0):-p.get(Q(0),0)}),'Caputo integral retains initial constant')
                shifted=fractional_integral(p,a)
                check(caputo(shifted,a)==clean(p),'left inverse with AC finite-power input')
                frac_models+=1
    u={Q(0):Q(2),Q(1,2):Q(1),Q(1):Q(-2),Q(3,2):Q(3)}
    force={Q(0):Q(5),Q(1):Q(-1),Q(3,2):Q(6)}
    check(verify_caputo(u,Q(1,2),2,-2,force),'manufactured half-order equation')
    reject(lambda:verify_caputo(u,Q(1,2),0,-2,force),'incorrect initial trace')
    broken=add(u,{Q(1):Q(1,100)})
    reject(lambda:verify_caputo(broken,Q(1,2),2,-2,force),'changed solution coefficient')
    reject(lambda:caputo({Q(-1,3):1},Q(1,2)),'non-AC singular solution class')
    reject(lambda:caputo(u,1),'order outside declared interface')
    reject(lambda:fractional_integral({Q(-1):1},Q(1,2)),'nonintegrable power input')
    nonsg=[]
    for a,b in [(Q(1,3),Q(1,3)),(Q(1,4),Q(1,2)),(Q(2,5),Q(1,5))]:
        p={b:Q(1)};successive=caputo(caputo(p,b),a);direct=caputo(p,a+b)
        check(not successive and direct=={-a:Q(1)},'valid-domain Caputo nonsemigroup')
        nonsg.append(dict(a=str(a),b=str(b),successive=poly_record(successive),direct=poly_record(direct)))
    # Every numeric comparison below is made on outward integer endpoints.
    ml_rows=[]
    for x in [Q(0),Q(1,8),Q(1,4),Q(1,2),Q(1),Q(2),Q(3),Q(4)]:
        fine=half_ml(x,120);coarse=half_ml(x,40);finer=half_ml(x,150)
        check(coarse['interval'].intersects(fine['interval']),'coarse and fine certified intervals agree')
        check(fine['interval'].hi-fine['interval'].lo<10**30,'fine width below 1e-50')
        check(fine['interval'].intersects(finer['interval']),'separate truncation interval agreement')
        check(fine['interval'].lo>0 and fine['interval'].hi<=SCALE if x else fine['interval'].contains(1),'positive and at most one')
        if x:
            bounds=long_tail(x)
            check(bounds.lo<=fine['interval'].lo and fine['interval'].hi<=bounds.hi,'positive-integral long-tail enclosure')
        ml_rows.append(dict(x=str(x),interval=fine['interval'].record(),paired_terms=121))
    reject(lambda:half_ml(4,14),'tail ratio boundary x^2=m+2')
    reject(lambda:half_ml(-1,20),'negative x interface')
    reject(lambda:half_ml(1,False),'boolean truncation index')
    forced=Q(3,2)+half_ml(2,120)['interval']/2
    check(forced.lo>Q(162769783,10**8)*SCALE and forced.hi<Q(162769785,10**8)*SCALE,'forced response displayed decimals')
    # Same time subdivision but a different dynamical continuation.
    S1=half_ml(1,120)['interval'];S4=half_ml(2,120)['interval']
    check((S1**4).hi<S4.lo,'four one-unit restarts differ from full four-unit memory')
    old_memory=(4-2*sqrt_interval(3))/sqrtpi()
    full=4/sqrtpi();restart=2*sqrt_interval(3)/sqrtpi()
    check(old_memory.lo>0 and (old_memory+restart).intersects(full),'Caputo restart memory at t=4,a=1')
    # Three equivalent exact model transformations, beyond one fixed forcing.
    transfers=[]
    for a in [Q(1,3),Q(1,2),Q(2,3)]:
        for lam in [Q(-2),Q(0),Q(3,4)]:
            for c in [Q(-1),Q(0),Q(3,2)]:
                p={Q(0):c,a:Q(1),2*a:Q(-2),3*a:Q(3)}
                f2=add(caputo(p,a),scale(p,-lam))
                check(verify_caputo(p,a,c,lam,f2),'changed order feedback and initial trace')
                check(add(fractional_integral(add(scale(p,lam),f2),a),{Q(0):c})==clean(p),'integral model matches derivative model')
                transfers.append(dict(alpha=str(a),lambda_=str(lam),initial=str(c),forcing=poly_record(f2)))
    check((half_ml(2,120,positive=True)['interval']+S4).intersects(2/exp_negative(4)),'positive-negative sum equals twice exp(4)')
    # A small perturbation of a known exact solution under lambda=-1/2.
    eta=Q(1,100)*(2/sqrtpi()+Q(1,2));growth=half_ml(Q(1,2),120,positive=True)['interval']
    perturbation_bound=2*eta*(growth-1)
    check(perturbation_bound.lo>Q(1,100)*SCALE,'comparison bound covers exact t endpoint error')
    check(perturbation_bound.hi<Q(1,25)*SCALE,'comparison budget below 0.04')
    result=dict(status='PASS',checks=CHECKS,polynomial_transfer_models=models,fractional_identity_models=frac_models,
        pi_interval=PI.record(),main_picard=dict(N=N,M=str(main['M']),x=str(main['x']),error_upper=str(main['error_upper']),
            previous_bound=str(tail_budget(1,2,N-1)),tolerance=str(eps),
            iterates=[poly_record(p) for p in main['log']],residual=poly_record(main['residual']),value_at_one=str(evaluate(main['log'][-1],1)),true_value_at_one=true_one.record(),actual_error_at_one=actual.record()),
        manufactured=dict(solution=poly_record(u),forcing=poly_record(force),values={str(t):power_sum_value(u,t).record() for t in [Q(0),Q(1,4),Q(1),Q(4)]}),
        caputo_nonsemigroup=nonsg,mittag_leffler=ml_rows,forced_at_one=forced.record(),
        memory_at_four=dict(full=full.record(),restarted=restart.record(),old_memory=old_memory.record(),
            full_relaxation_at_four=S4.record(),four_restarts=(S1**4).record()),
        long_tail_at_x20=long_tail(20).record(),perturbed_model=dict(lambda_='-1/2',perturbation='phi_1/100',
            residual_bound=eta.record(),error_bound_at_one=perturbation_bound.record(),actual_error_at_one='1/100'),
        model_transfers=transfers,
        scope='Exact finite algebra and outward interval arithmetic. Continuous convergence, AC regularity and whole-interval bounds are proved in the pages; no finite sample proves them.')
    return result

if __name__=='__main__':
    parser=argparse.ArgumentParser();parser.add_argument('--output',type=Path,required=True)
    args=parser.parse_args();result=run();args.output.parent.mkdir(parents=True,exist_ok=True)
    args.output.write_text(json.dumps(result,ensure_ascii=False,indent=2)+'\n')
    print(json.dumps({k:result[k] for k in ['status','checks','polynomial_transfer_models','fractional_identity_models']},ensure_ascii=False))
