#!/usr/bin/env python3
"""Exact rational certificates for finite ODE tubes and set propagation.
Only the Python standard library is used. No network or private data access.
--output PATH writes JSON there; without it the JSON is printed.
"""
from fractions import Fraction as F
from dataclasses import dataclass
from collections import Counter
from itertools import product
import argparse
import json

COUNTS = Counter()
def check(ok, label):
    if not ok:
        raise AssertionError(label)
    COUNTS[label] += 1

def frac(x):
    return F(x)

@dataclass(frozen=True)
class Interval:
    lo: F
    hi: F
    def __post_init__(self):
        object.__setattr__(self, 'lo', F(self.lo))
        object.__setattr__(self, 'hi', F(self.hi))
        if self.lo > self.hi:
            raise ValueError('empty interval')
    def __add__(self, other):
        other = 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 + -interval(other)
    def __rsub__(self, other):
        return interval(other) + -self
    def __mul__(self, other):
        other = interval(other)
        ends = [a*b for a in (self.lo, self.hi) for b in (other.lo, other.hi)]
        return Interval(min(ends), max(ends))
    __rmul__ = __mul__
    def __truediv__(self, other):
        other = interval(other)
        if other.lo <= 0 <= other.hi:
            raise ValueError('division by zero-containing interval')
        return self * Interval(1/other.hi, 1/other.lo)
    def __pow__(self, n):
        if not isinstance(n, int) or n < 0:
            raise ValueError('nonnegative integer exponent required')
        if n == 0:
            return Interval(1, 1)
        ends = [self.lo**n, self.hi**n]
        return Interval(0 if n % 2 == 0 and self.lo <= 0 <= self.hi else min(ends), max(ends))
    def contains(self, other):
        other = interval(other)
        return self.lo <= other.lo and other.hi <= self.hi
    def strictly_contains(self, other):
        other = interval(other)
        return self.lo < other.lo and other.hi < self.hi
    def intersect(self, other):
        return Interval(max(self.lo, other.lo), min(self.hi, other.hi))
    def outward_grid(self, den):
        return Interval(F((self.lo * den).__floor__(), den), F((self.hi * den).__ceil__(), den))
    def json(self):
        return [str(self.lo), str(self.hi)]

def interval(x):
    return x if isinstance(x, Interval) else Interval(x, x)

def trim(p):
    p = list(map(F, p))
    while len(p) > 1 and p[-1] == 0:
        p.pop()
    return p

def padd(a, b):
    return trim([(a[i] if i < len(a) else 0) + (b[i] if i < len(b) else 0) for i in range(max(len(a), len(b)))])

def scale(a, c):
    return trim([x*c for x in a])

def pmul(a, b):
    c = [F(0)] * (len(a)+len(b)-1)
    for i,x in enumerate(a):
        for j,y in enumerate(b):
            c[i+j] += x*y
    return trim(c)

def pdiff(a):
    return trim([F(i)*a[i] for i in range(1,len(a))] or [F(0)])

def pvalue(p, x):
    value = F(0)
    for c in reversed(p):
        value = value*x+c
    return value

def exp_bounds(x, n=40):
    """Positive Taylor terms plus a geometric upper tail, all rational."""
    x=F(x)
    if x < 0:
        pos=exp_bounds(-x,n)
        return Interval(1/pos.hi,1/pos.lo)
    if x >= n+2:
        raise ValueError('tail ratio is not below one')
    term=F(1); total=term
    for k in range(1,n+1):
        term *= x/k; total += term
    first = term*x/(n+1)
    return Interval(total, total+first/(1-x/F(n+2)))

def mm(A,B):
    return [[sum((a*b for a,b in zip(row,col)),F(0)) for col in zip(*B)] for row in A]

def mv(A,v):
    return [sum((x*y for x,y in zip(row,v)),F(0)) for row in A]

def transpose(A):
    return [list(x) for x in zip(*A)]

def eye(n):
    return [[F(i==j) for j in range(n)] for i in range(n)]

def absolute(A):
    return [[abs(x) for x in row] for row in A]

def mu_inf(A):
    return max(A[i][i]+sum(abs(x) for j,x in enumerate(row) if i!=j) for i,row in enumerate(A))

def norm_inf(A):
    return max(sum(abs(x) for x in row) for row in A)

def weighted(A,w):
    return [[x*w[j]/w[i] for j,x in enumerate(row)] for i,row in enumerate(A)]

def inv2(A):
    a,b=A[0];c,d=A[1];det=a*d-b*c
    if det==0:raise ValueError('singular coordinate matrix')
    return [[d/det,-b/det],[-c/det,a/det]]

def matrix_norm_checks():
    for a,b,c,d in product(range(-2,3),repeat=4):
        A=[[F(a),F(b)],[F(c),F(d)]];h=F(1,8)
        near=[[F(i==j)+h*A[i][j] for j in range(2)]for i in range(2)]
        check((norm_inf(near)-1)/h==mu_inf(A),'log_norm_exact_small_step')
        w=[F(3),F(2)];B=weighted(A,w)
        check(mu_inf(B)==max(A[i][i]+sum(abs(A[i][j])*w[j]/w[i]for j in range(2)if i!=j)for i in range(2)),'weighted_row_formula')
    A=[[F(-2),F(8)],[F(0),F(-2)]]
    check(mu_inf(A)==6 and norm_inf(A)==10,'nonnormal_example')
    check(mu_inf(weighted(A,[F(8),F(1)]))==-1,'scaled_nonnormal_contraction')
    for k in range(81):
        t=F(k,20);em2=exp_bounds(-2*t);em1=exp_bounds(-t)
        check(em2.hi*(1+t)<=em1.lo,'weighted_exact_flow_exponential_bound')
    return {'A':[['-2','8'],['0','-2']],'mu_inf':'6','mu_2':'2','weights':['8','1'],'mu_weighted_inf':'-1'}

def tube_checks():
    # Manufactured candidate is nonconstant; the nonlinear system retains cubic terms.
    z1=[F(0),F(1,10)];z2=[F(0),F(0),F(1,20)]
    q1=padd(padd(pdiff(z1),scale(z1,2)),scale(z2,-8))
    q1=padd(padd(q1,pmul(pmul(z1,z1),z1)),[F(-1,1000)])
    q2=padd(pdiff(z2),scale(z2,2))
    q2=padd(padd(q2,pmul(pmul(z2,z2),z2)),[F(-1,2000)])
    fz1=padd(padd(scale(z1,-2),scale(z2,8)),padd(scale(pmul(pmul(z1,z1),z1),-1),q1))
    fz2=padd(scale(z2,-2),padd(scale(pmul(pmul(z2,z2),z2),-1),q2))
    check(padd(pdiff(z1),scale(fz1,-1))==[F(1,1000)],'tube_residual_polynomial_identity')
    check(padd(pdiff(z2),scale(fz2,-1))==[F(1,2000)],'tube_residual_polynomial_identity')
    for i,j in product(range(-15,16),repeat=2):
        x=F(i,10);y=F(j,10)
        J=[[-2-3*x*x,F(8)],[F(0),-2-3*y*y]]
        check(mu_inf(weighted(J,[F(8),F(1)]))<=-1,'tube_weighted_jacobian_sample')
    r0=F(1,1000);delta=F(1,2000);rho=F(1,500)
    # R(t)=delta+(r0-delta)e^-t <= r0; the strict inequality is exact.
    check(delta<=r0<rho,'tube_uniform_strict_margin')
    lower_exp4=sum((F(4)**k/F(__import__('math').factorial(k))for k in range(8)),F(0))
    check(lower_exp4>50,'tube_endpoint_exp4_rational_certificate')
    Rupper=F(51,100000)
    check(delta+(r0-delta)/50==Rupper,'tube_endpoint_radius')
    resultbox=[Interval(F(2,5)-8*Rupper,F(2,5)+8*Rupper),Interval(F(4,5)-Rupper,F(4,5)+Rupper)]
    check(resultbox[0]==Interval(F(4949,12500),F(5051,12500)),'tube_first_physical_coordinate')
    # Formal-page example residual t^3/1000 and radius .008+.002exp(-t).
    formal_residual=[F(0),F(0),F(0),F(1,1000)]
    for k in range(81):
        t=F(k,40)
        check(F(0)<=pvalue(formal_residual,t)<=F(1,125),'formal_tube_residual_interval')
    return {'candidate_polynomials':[[str(x)for x in z1],[str(x)for x in z2]],'forcing_polynomials':[[str(x)for x in q1],[str(x)for x in q2]],'residual':['1/1000','1/2000'],'weighted_residual_bound':'1/2000','initial_weighted_radius':'1/1000','working_radius':'1/500','m':'-1','T':'4','exp4_lower_partial_sum_degree7':str(lower_exp4),'endpoint_radius_upper':str(Rupper),'endpoint_physical_box':[x.json()for x in resultbox]}

def flow_coefficients(X,p):
    """Normalized series for u'=u^2, v'=u*v using convolution, not exact solution."""
    u=[X[0]];v=[X[1]]
    for k in range(p):
        u.append(sum((u[j]*u[k-j]for j in range(k+1)),interval(0))/(k+1))
        v.append(sum((u[j]*v[k-j]for j in range(k+1)),interval(0))/(k+1))
    return [[u[k],v[k]]for k in range(p+1)]

def field(B):return [B[0]**2,B[0]*B[1]]

def certified_step(X,h,B,p):
    if h<=0 or p<1:raise ValueError('invalid step')
    fi=field(B);Q=[x+Interval(0,h)*f for x,f in zip(X,fi)]
    if not all(b.strictly_contains(q)for b,q in zip(B,Q)):
        raise ValueError('unverified full-step box')
    coeff=flow_coefficients(X,p);remainder=flow_coefficients(B,p+1)[p+1]
    Y=[]
    for j in range(2):
        y=sum((h**k*coeff[k][j]for k in range(p+1)),interval(0))+h**(p+1)*remainder[j]
        Y.append(B[j].intersect(y))
    return Y,Q,remainder

def make_certificate():
    X=[Interval(1,F(101,100)),Interval(2,F(201,100))];h=F(1,16);p=6;t=F(0);steps=[]
    for _ in range(8):
        B=[Interval(F(9,10)*x.lo,F(3,2)*x.hi)for x in X]
        Y,Q,E=certified_step(X,h,B,p);out=[y.outward_grid(10**12)for y in Y]
        steps.append({'t':str(t),'h':str(h),'start':[x.json()for x in X],'full_step':[b.json()for b in B],'picard_image':[q.json()for q in Q],'remainder_coefficient':[e.json()for e in E],'output':[x.json()for x in out]})
        t+=h;X=out
    return {'model':'u_prime=u^2;v_prime=u*v','order':p,'grid_denominator':10**12,'initial':[['1','101/100'],['2','201/100']],'target':'1/2','steps':steps}

def parsebox(rows):
    if len(rows)!=2:
        raise ValueError('this model requires exactly two coordinates')
    return [Interval(F(a),F(b))for a,b in rows]

def verify_certificate(cert):
    if cert['model']!='u_prime=u^2;v_prime=u*v':raise ValueError('wrong model')
    X=parsebox(cert['initial']);t=F(0);p=cert['order']
    for row in cert['steps']:
        if F(row['t'])!=t or parsebox(row['start'])!=X:raise ValueError('broken time or state chain')
        h=F(row['h']);B=parsebox(row['full_step']);Y,Q,E=certified_step(X,h,B,p)
        if Q!=parsebox(row['picard_image']) or E!=parsebox(row['remainder_coefficient']):raise ValueError('claimed bounds not reproduced')
        output=parsebox(row['output'])
        if not all(a.contains(b)for a,b in zip(output,Y)):raise ValueError('lost endpoint containment')
        X=output;t+=h
    if t!=F(cert['target']):raise ValueError('target not covered')
    return X

def taylor_checks():
    cert=make_certificate();last=verify_certificate(cert)
    X0=parsebox(cert['initial'])
    for n,row in enumerate(cert['steps'],1):
        t=F(n,16);output=parsebox(row['output'])
        exact=[Interval(X0[0].lo/(1-t*X0[0].lo),X0[0].hi/(1-t*X0[0].hi)),Interval(X0[1].lo/(1-t*X0[0].lo),X0[1].hi/(1-t*X0[0].hi))]
        check(all(a.contains(b)for a,b in zip(output,exact)),'coupled_taylor_exact_solution_crosscheck')
        for a,b in product(range(11),repeat=2):
            u0=1+F(a,1000);v0=2+F(b,1000)
            check(output[0].contains(u0/(1-t*u0))and output[1].contains(v0/(1-t*u0)),'coupled_initial_grid_diagnostic')
    # Independent coefficient identity at varying point/box inputs.
    for a in range(1,11):
        for b in range(1,7):
            X=[Interval(F(a,10),F(a+1,10)),Interval(F(b,10),F(b+1,10))]
            C=flow_coefficients(X,8)
            for k,row in enumerate(C):
                check(row[0]==X[0]**(k+1) and row[1]==X[1]*X[0]**k,'coupled_normalized_coefficient_identity')
    import copy
    corruptions=[]
    for name,edit in [('missing_last_step',lambda c:c['steps'].pop()),('time_gap',lambda c:c['steps'][2].update(t='7/32')),('too_small_full_step_box',lambda c:c['steps'][0]['full_step'][0].__setitem__(1,'21/20')),('false_remainder',lambda c:c['steps'][0]['remainder_coefficient'].__setitem__(0,['0','0'])),('false_endpoint',lambda c:c['steps'][-1]['output'][0].__setitem__(1,'2')),('extra_coordinate',lambda c:c['steps'][-1]['output'].append(['0','0']))]:
        bad=copy.deepcopy(cert);edit(bad)
        try:verify_certificate(bad)
        except (ValueError,KeyError):corruptions.append(name)
        else:raise AssertionError('damaged certificate accepted: '+name)
        check(True,'damaged_certificate_rejected')
    # Formal-page scalar p=4 step, recomputed separately.
    x=Interval(1,F(101,100));B=Interval(F(9,10),F(6,5));h=F(1,8)
    y=sum((h**k*x**(k+1)for k in range(5)),interval(0))+h**5*B**6
    check(y==Interval(F(37448531441,32768000000),F(47349395541301,40960000000000)),'formal_scalar_taylor_exact_bounds')
    check(y.contains(Interval(F(8,7),F(808,699))),'formal_scalar_exact_flow')
    return {'certificate':cert,'final_box':[x.json()for x in last],'rejected_corruptions':corruptions,'formal_scalar_box':y.json()}

def wrapping_checks():
    R=[[F(3,5),F(-4,5)],[F(4,5),F(3,5)]];I=eye(2);eps=F(1,100);delta=F(1,10000)
    check(mm(transpose(R),R)==I,'rational_rotation_orthogonal')
    Q=I;rad=[eps,eps];naive=eps;noise_sum=[F(0),F(0)];powers=[I];terminal=None
    for n in range(1,41):
        oldQ=Q;Q=mm(R,Q);powers.append(Q)
        inv=inv2(Q);check(mm(inv,Q)==I and mm(Q,inv)==I,'coordinate_inverse_exact')
        check(mm(inv,mm(R,oldQ))==I,'retained_linear_coordinates')
        extra=mv(absolute(inv),[delta,delta]);rad=[x+y for x,y in zip(rad,extra)]
        naive=F(7,5)*naive+delta
        check(naive==F(7,5)**n*eps+F(5,2)*delta*(F(7,5)**n-1),'naive_wrapping_geometric_sum')
        for row in Q:
            check(sum(abs(x)for x in row)**2<=2,'rotation_row_bound_squared')
        check(all(x<=eps+F(3,2)*n*delta for x in rad),'moving_radius_linear_bound')
        physical=mv(absolute(Q),rad)
        noise=mv(absolute(powers[n-1]),[delta,delta]);noise_sum=[x+y for x,y in zip(noise_sum,noise)]
        exact_hull=[x+y for x,y in zip(mv(absolute(Q),[eps,eps]),noise_sum)]
        check(all(a>=b for a,b in zip(physical,exact_hull)),'exact_zonotope_support_contained')
        if n==20:
            check(all(x<F(1,50)for x in physical),'twenty_steps_physical_tolerance')
            check(naive>8,'twenty_steps_naive_width')
            terminal={'n':n,'epsilon':str(eps),'delta':str(delta),'coordinate_radius':[str(x)for x in rad],'physical_radius':[str(x)for x in physical],'exact_axis_hull_radius':[str(x)for x in exact_hull],'naive_axis_radius':str(naive),'simple_coordinate_radius_upper':'13/1000','simple_physical_radius_upper':'39/2000'}
    # Nonlinear shear: check sampled rational grid points against the shifted box.
    for k in range(1,9):
        e=F(k,10);center=e*e/2;newr=e+e*e/2
        for i,j in product(range(-10,11),repeat=2):
            x=e*F(i,10);y=e*F(j,10)
            check(abs(x+y*y-center)<=newr and abs(y)<=e,'nonlinear_shear_centered_remainder')
    try:inv2([[F(1),F(2)],[F(2),F(4)]])
    except ValueError:check(True,'singular_coordinate_rejected')
    else:raise AssertionError('singular inverse accepted')
    return terminal

def run():
    COUNTS.clear()
    norm=matrix_norm_checks();tube=tube_checks();taylor=taylor_checks();wrapping=wrapping_checks()
    return {'status':'PASS','checks':sum(COUNTS.values()),'categories':dict(sorted(COUNTS.items())),'logarithmic_norm':norm,'nonlinear_weighted_tube':tube,'validated_coupled_taylor':taylor,'wrapping':wrapping,'arithmetic':'Fraction throughout; exponential bounds use positive Taylor partial sums and rational geometric tails. No floating-point comparisons or exact-solution data are used to accept a Taylor step. Finite diagnostics supplement the general proofs.'}

if __name__=='__main__':
    ap=argparse.ArgumentParser();ap.add_argument('--output');args=ap.parse_args();text=json.dumps(run(),ensure_ascii=False,indent=2)+'\n'
    if args.output:
        from pathlib import Path
        Path(args.output).write_text(text)
    else:print(text,end='')
