#!/usr/bin/env python3
"""Exact rational certificates for the A05 worked problems.

No third-party package. Every arithmetic endpoint is a Fraction. exp(t) is
bounded by a proved Taylor/geometric tail; floats are used only for display.
Use --output FILE for the summary and --certificate FILE for a replayable tree.
The finite tree proves the fixed F and D below, not arbitrary input programs.
"""
from fractions import Fraction as Q
from dataclasses import dataclass
from functools import lru_cache
import argparse, json

ASSERTIONS = 0
def check(test, message='certificate assertion'):
    global ASSERTIONS
    ASSERTIONS += 1
    if not test: raise ValueError(message)

@dataclass(frozen=True)
class I:
    lo: Q
    hi: Q
    def __post_init__(self):
        object.__setattr__(self,'lo',Q(self.lo)); object.__setattr__(self,'hi',Q(self.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); v=[self.lo*b.lo,self.lo*b.hi,self.hi*b.lo,self.hi*b.hi];return I(min(v),max(v))
    __rmul__=__mul__
    def reciprocal(self):
        if self.lo<=0<=self.hi: raise ValueError('division through zero')
        return I(1/self.hi,1/self.lo)
    def __truediv__(self,b): return self*iv(b).reciprocal()
    def square(self):
        return I(0 if self.lo<=0<=self.hi else min(self.lo*self.lo,self.hi*self.hi),max(self.lo*self.lo,self.hi*self.hi))
    def mag(self): return max(abs(self.lo),abs(self.hi))
    def midpoint(self): return (self.lo+self.hi)/2
    def radius(self): return (self.hi-self.lo)/2
    def subset(self,b,strict=False):
        return b.lo<self.lo and self.hi<b.hi if strict else b.lo<=self.lo and self.hi<=b.hi
    def record(self): return [str(self.lo),str(self.hi)]
def iv(x): return x if isinstance(x,I) else I(x,x)
def box(values): return [I(Q(a),Q(b)) for a,b in values]
def boxrec(X): return [x.record() for x in X]

@lru_cache(maxsize=None)
def exp_point(t,N=18):
    t=Q(t)
    if abs(t)>1: raise ValueError('this exponential enclosure supports |t|<=1')
    if t<0: return exp_point(-t,N).reciprocal()
    term=Q(1); S=term
    for k in range(1,N+1):
        term=term*t/k; S+=term
    first=term*t/(N+1)
    tail=first/(1-t/Q(N+2))
    return I(S,S+tail)
def exp_interval(x,N=18):
    return I(exp_point(x.lo,N).lo,exp_point(x.hi,N).hi)
def F(X):
    x,y=X
    return [x.square()+y.square()-1,y-exp_interval(x)/2]
def J(X):
    x,y=X
    return [[2*x,2*y],[-exp_interval(x)/2,iv(1)]]
def matvec(A,v): return [sum((a*b for a,b in zip(row,v)),iv(0)) for row in A]
def krawczyk(X,C):
    c=[x.midpoint() for x in X]; radii=[x.radius() for x in X]
    f=F([iv(x) for x in c]); j=J(X)
    z=[-a for a in matvec(C,f)]
    E=[[iv(int(i==k))-sum((C[i][l]*j[l][k] for l in range(2)),iv(0)) for k in range(2)] for i in range(2)]
    K=[iv(c[i])+z[i]+sum((E[i][k]*I(-radii[k],radii[k]) for k in range(2)),iv(0)) for i in range(2)]
    q=max(sum(E[i][k].mag()*radii[k] for k in range(2))/radii[i] for i in range(2))
    return K,E,q

DOMAIN=box([('-1','1'),('0','2')])
ROOT_BOXES=[box([('-983/1000','-981/1000'),('186/1000','188/1000')]),box([('528/1000','530/1000'),('848/1000','850/1000')])]
PRECONDITIONERS=[[[Q('-528/1000'),Q('198/1000')],[Q('-99/1000'),Q('1037/1000')]],[[Q('400/1000'),Q('-680/1000')],[Q('340/1000'),Q('423/1000')]]]

def split(X,axis):
    m=X[axis].midpoint();left=X.copy();right=X.copy()
    left[axis]=I(X[axis].lo,m);right[axis]=I(m,X[axis].hi)
    return left,right

def build_tree(X,roots,budget):
    if budget[0]<=0: return {'kind':'unresolved'}
    budget[0]-=1
    for k,B in enumerate(roots):
        if all(x.subset(b,True) for x,b in zip(X,B)):
            return {'kind':'covered','root':k}
    for k,v in enumerate(F(X)):
        if not v.lo<=0<=v.hi:
            return {'kind':'excluded','component':k,'bound':v.record()}
    axis=max(range(2),key=lambda k:X[k].hi-X[k].lo)
    L,R=split(X,axis)
    return {'kind':'split','axis':axis,'left':build_tree(L,roots,budget),'right':build_tree(R,roots,budget)}

def make_certificate(limit=10000):
    tree=build_tree(DOMAIN,ROOT_BOXES,[limit])
    return {'version':1,'problem':'circle-exp-half','domain':boxrec(DOMAIN),'root_boxes':[boxrec(X) for X in ROOT_BOXES], 'preconditioners':[[[str(a) for a in row] for row in C] for C in PRECONDITIONERS],'tree':tree}

def verify_certificate(cert,allow_unresolved=False):
    check(cert['version']==1 and cert['problem']=='circle-exp-half','unsupported problem')
    check(cert['domain']==boxrec(DOMAIN),'changed domain')
    Bs=[box(x) for x in cert['root_boxes']]
    Cs=[[[Q(a) for a in row] for row in C] for C in cert['preconditioners']]
    check(len(Bs)==len(Cs),'box/preconditioner count')
    root_data=[]
    for B,C in zip(Bs,Cs):
        check(len(B)==2 and all(b.radius()>0 and b.subset(d) for b,d in zip(B,DOMAIN)),'root box domain')
        check(len(C)==2 and all(len(row)==2 for row in C),'matrix shape')
        K,E,q=krawczyk(B,C)
        check(all(k.subset(b,True) for k,b in zip(K,B)),'strict Krawczyk inclusion')
        check(q<1,'weighted contraction')
        root_data.append({'K':boxrec(K),'K_display':[[float(k.lo),float(k.hi)] for k in K],'q':str(q),'q_display':float(q)})
    for i in range(len(Bs)):
        for j in range(i):
            check(any(a.hi<b.lo or b.hi<a.lo for a,b in zip(Bs[i],Bs[j])),'overlapping certified boxes')
    counts={'split':0,'excluded':0,'covered':0,'unresolved':0,'max_depth':0}
    def visit(node,X,depth):
        kind=node['kind'];check(kind in counts and kind!='max_depth','unknown node')
        counts[kind]+=1;counts['max_depth']=max(depth,counts['max_depth'])
        if kind=='split':
            a=node['axis'];check(type(a) is int and a in (0,1),'split axis')
            check(a==max(range(2),key=lambda k:X[k].hi-X[k].lo),'longest side split')
            L,R=split(X,a);visit(node['left'],L,depth+1);visit(node['right'],R,depth+1)
        elif kind=='covered':
            k=node['root'];check(type(k) is int and 0<=k<len(Bs),'root index')
            check(all(x.subset(b,True) for x,b in zip(X,Bs[k])),'false coverage')
        elif kind=='excluded':
            k=node['component'];check(type(k) is int and k in (0,1),'component index')
            v=F(X)[k];check(v.record()==node['bound'],'altered arithmetic witness')
            check(not v.lo<=0<=v.hi,'false exclusion')
        else: check(allow_unresolved,'unresolved leaf cannot certify completeness')
    visit(cert['tree'],DOMAIN,0)
    check(counts['split']+1==counts['excluded']+counts['covered']+counts['unresolved'],'full binary tree')
    return {'status':'complete' if not counts['unresolved'] else 'unresolved','root_count':len(Bs) if not counts['unresolved'] else None,'counts':counts,'root_certificates':root_data}


def tests():
    global ASSERTIONS
    ASSERTIONS = 0
    import copy
    # Interval endpoint algebra: exact rational grid, including zero crossings.
    intervals=[I(Q(a,4),Q(b,4)) for a in range(-4,5) for b in range(a,5)]
    for X in intervals:
        for Y in intervals:
            for a in (X.lo,X.midpoint(),X.hi):
                for b in (Y.lo,Y.midpoint(),Y.hi):
                    check((X+Y).lo<=a+b<=(X+Y).hi)
                    check((X*Y).lo<=a*b<=(X*Y).hi)
            for a in (X.lo,X.midpoint(),X.hi):check(X.square().lo<=a*a<=X.square().hi)
    # Exponential nested enclosures, without relying on floating exp.
    for k in range(-100,101):
        t=Q(k,100);a=exp_point(t,18);b=exp_point(t,25)
        check(b.subset(a),'Taylor degree refinement')
        prod=a*exp_point(-t,18);check(prod.lo<=1<=prod.hi,'exp reciprocal')
    # Oettli-Prager membership: construct admissible data from any admitted x.
    Ac=[[Q(1),Q(0)],[Q(0),Q(1)]];D=[[Q(0),Q(1,5)],[Q(1,10),Q(0)]];bc=[Q(1),Q(1)]
    admitted=0
    for i in range(70,131):
        for j in range(80,121):
            x=[Q(i,100),Q(j,100)];res=[x[k]-bc[k] for k in range(2)]
            s=[sum(D[k][l]*abs(x[l]) for l in range(2)) for k in range(2)]
            if all(abs(res[k])<=s[k] for k in range(2)):
                E=[[-res[k]/s[k]*D[k][l]*(1 if x[l]>=0 else -1) for l in range(2)] for k in range(2)]
                check(all(abs(E[k][l])<=D[k][l] for k in range(2) for l in range(2)))
                check(all(sum((Ac[k][l]+E[k][l])*x[l] for l in range(2))==bc[k] for k in range(2)))
                check(abs(x[0]-1)<=Q(11,49) and abs(x[1]-1)<=Q(6,49))
                admitted+=1
    # Scalar interval Newton, root containment via rational signs.
    X=I(Q(1),Q(2));newton=[]
    for _ in range(4):
        c=X.midpoint();N=iv(c)-iv(c*c-2)/(2*X)
        check(N.subset(X,True));check(N.lo*N.lo<2<N.hi*N.hi)
        newton.append(N.record());X=N
    check(newton[0]==['11/8','23/16'])
    # Kantorovich majorant for x^2-2 from 3/2, including exact track identity.
    t=Q(0);x=Q(3,2);beta=Q(1,12);L=Q(2,3);majorants=[]
    for _ in range(5):
        p=L*t*t/2-t+beta;dt=p/(1-L*t)
        check(p>0 and dt>0 and L*t<1)
        xn=x-(x*x-2)/(2*x);tn=t+dt
        check(xn==Q(3,2)-tn);check(xn*xn>2)
        majorants.append({'t':str(tn),'x':str(xn)});x,t=xn,tn
    # Critical majorant has only linear convergence; supercritical has no real root.
    for k in range(10):
        x=1-Q(1,2**k);f=x-x*x/2-Q(1,2)
        check(x-f/(1-x)==1-Q(1,2**(k+1)))
    cert=make_certificate();report=verify_certificate(cert)
    display_boxes=[box([('-0.9823206','-0.9823149'),('0.1872201','0.1872219')]),box([('0.5290003','0.5290069'),('0.8486169','0.8486226')])]
    for B,C,shown,limit in zip(ROOT_BOXES,PRECONDITIONERS,display_boxes,[Q('0.002767'),Q('0.003237')]):
        K,E,q=krawczyk(B,C);check(all(k.subset(d) for k,d in zip(K,shown)),'printed decimal envelope');check(q<limit,'printed contraction bound')
    # Mutations are rejected; budget exhaustion remains an explicit non-certificate.
    mutations=[]
    c=copy.deepcopy(cert);c['domain'][0][0]='-2';mutations.append(c)
    c=copy.deepcopy(cert);del c['tree']['right'];mutations.append(c)
    c=copy.deepcopy(cert);c['root_boxes'][0]=boxrec(box([('0','1/1000'),('0','1/1000')]));mutations.append(c)
    c=copy.deepcopy(cert);c['preconditioners'][0]=[['0','0'],['0','0']];mutations.append(c)
    c=copy.deepcopy(cert);c['tree']={'kind':'covered','root':0};mutations.append(c)
    c=copy.deepcopy(cert);c['tree']={'kind':'excluded','component':0,'bound':['1','2']};mutations.append(c)
    c=copy.deepcopy(cert);c['tree']={'kind':'unresolved'};mutations.append(c)
    c=copy.deepcopy(cert);c['root_boxes'].append(c['root_boxes'][0]);c['preconditioners'].append(c['preconditioners'][0]);mutations.append(c)
    rejected=0
    for c in mutations:
        try: verify_certificate(c)
        except (ValueError,KeyError,TypeError): rejected+=1
    check(rejected==len(mutations),'mutation accepted')
    partial=make_certificate(5);p=verify_certificate(partial,True);check(p['status']=='unresolved' and p['root_count'] is None)
    report.update(assertions=ASSERTIONS,arithmetic='Python standard-library Fraction; exp Taylor degree 18 with rational tail',oettli_prager_grid_members=admitted,scalar_interval_newton=newton,kantorovich_majorants=majorants,rejected_mutations=rejected,limited_run=p['counts'])
    return cert,report

def main():
    global ASSERTIONS
    ASSERTIONS = 0
    p=argparse.ArgumentParser();p.add_argument('--output');p.add_argument('--certificate');p.add_argument('--verify');a=p.parse_args()
    if a.verify:
        out=verify_certificate(json.load(open(a.verify)));out['assertions']=ASSERTIONS
    else:
        cert,out=tests()
        if a.certificate:
            with open(a.certificate,'w') as f:json.dump(cert,f,ensure_ascii=False,indent=2);f.write('\n')
    text=json.dumps(out,ensure_ascii=False,indent=2)
    if a.output:
        with open(a.output,'w') as f:f.write(text+'\n')
    print(text)
if __name__=='__main__':main()
