#!/usr/bin/env python3
"""Exact arithmetic checks for the stated history-update certificates.

Standard library only. Finite parameter tests are not proofs of general convergence.
Sine enclosures use Taylor's explicit global derivative remainder; all other values
are rational. Checks stay enabled with python -O. No file is written without --output.
"""
from fractions import Fraction as F
from itertools import product
from math import factorial
from pathlib import Path
import argparse,json
counts={}
def check(group,condition):
    counts[group]=counts.get(group,0)+1
    if not condition:raise RuntimeError('failed '+group)
def add(x,y):return tuple(a+b for a,b in zip(x,y))
def sub(x,y):return tuple(a-b for a,b in zip(x,y))
def scale(c,x):return tuple(c*a for a in x)
def dot(x,y):return sum((a*b for a,b in zip(x,y)),F(0))
def norm2(x):return dot(x,x)
def infnorm(x):return max(map(abs,x))
def mv(A,x):return tuple(dot(row,x)for row in A)
def mm(A,B):return tuple(tuple(sum(A[i][k]*B[k][j]for k in range(len(B)))for j in range(len(B[0])))for i in range(len(A)))
def objective(x,a):return norm2(sub(x,a))/2

def sine_bounds(x,terms=45):
    # Taylor through degree 2*terms-1; next-order Lagrange remainder uses
    # |d^(2*terms) sin| <= 1 everywhere, including arguments outside a small disk.
    p=sum(((-1)**k*x**(2*k+1)/factorial(2*k+1)for k in range(terms)),F(0))
    err=abs(x)**(2*terms)/factorial(2*terms)
    return max(F(-1),p-err),min(F(1),p+err)
def square_bounds(lo,hi):
    return (F(0)if lo<=0<=hi else min(lo*lo,hi*hi),max(lo*lo,hi*hi))
def hb_gradient(x):return 25*x if x<1 else x+24 if x<2 else 25*x-24
def hb_objective(x):return F(25,2)*x*x if x<1 else x*x/2+24*x-12 if x<2 else F(25,2)*x*x-24*x+36

def two_history(g,x0,x1):
    r0=sub(g(x0),x0);r1=sub(g(x1),x1);d=sub(r0,r1)
    # Minimum-norm scalar theta for a rank-zero difference is zero.
    t=-dot(r1,d)/norm2(d)if norm2(d)else F(0)
    y=add(scale(t,g(x0)),scale(1-t,g(x1)))
    return t,y,r0,r1

def afw_step(vertices,w,a):
    x=tuple(sum(wi*v[j]for wi,v in zip(w,vertices))for j in range(len(a)))
    grad=sub(x,a);prices=[dot(grad,v)for v in vertices];active=[i for i,u in enumerate(w)if u>0]
    s=min(range(len(vertices)),key=lambda i:(prices[i],i));v=max(active,key=lambda i:(prices[i],-i))
    G=dot(grad,x)-prices[s];away=prices[v]-dot(grad,x)
    if G==0:return {'x':x,'w':w,'gap':G,'stop':True}
    isaway=away>G;d=sub(x,vertices[v])if isaway else sub(vertices[s],x)
    cap=w[v]/(1-w[v])if isaway else F(1);gd=away if isaway else G
    gamma=min(gd/norm2(d),cap);new=[(1+gamma)*u if isaway else(1-gamma)*u for u in w]
    new[v if isaway else s]+= -gamma if isaway else gamma
    return {'x':add(x,scale(gamma,d)),'w':tuple(new),'gap':G,'away_gap':away,'gamma':gamma,'cap':cap,'direction':d,'direction_gap':gd,'drop':isaway and gamma==cap,'oldx':x,'stop':False}

def simplex_projection(a):
    vals=sorted(a,reverse=True);total=F(0);tau=None
    for j,v in enumerate(vals,1):
        total+=v;t=(total-1)/j
        if v>t:tau=t
    return tuple(max(F(0),v-tau)for v in a)

def run():
    counts.clear()
    # Certified finite evaluations of the nonconvex PL example, not the global proof.
    for j in range(-40,41):
        x=F(j,8);lo,hi=sine_bounds(x);_,s2hi=square_bounds(lo,hi)
        a,b=sine_bounds(2*x);glo=2*x+F(3,2)*a;ghi=2*x+F(3,2)*b
        g2lo,_=square_bounds(glo,ghi);fhi=x*x+F(3,2)*s2hi
        check('certified_finite_pl_evaluations',g2lo/2>=fhi/5)
    k=0
    while F(5,2)*F(24,25)**k>F(1,100):k+=1
    check('pl_minimum_bound_budget',k==136 and F(5,2)*F(24,25)**(k-1)>F(1,100))
    for x,y,g in product([F(j,5)for j in range(-5,6)],repeat=3):
        nx=x-(g+F(1,10))/10;ny=y-(-g+F(1,10))/10
        check('flat_direction_difference_identity',nx-ny==x-y-g/5)
        check('flat_direction_mean_identity',(nx+ny)/2==(x+y)/2-F(1,100))
    check('flat_noise_generic_floor',F(1,50)/(2*F(2,5))==F(1,40))
    check('flat_noise_target_mean',F(1,2)-F(136,100)==-F(43,50))

    # Modal recurrence and optimal endpoint roots over rational square-root inputs.
    for m,l in product(range(1,6),range(2,10)):
        if l<=m:continue
        mu=F(m*m);L=F(l*l);q=F(l-m,l+m);alpha=F(4,(l+m)**2);beta=q*q
        for j in range(11):
            lam=mu+(L-mu)*F(j,10);aa=1+beta-alpha*lam
            check('quadratic_optimal_modal_discriminant',aa*aa<=4*beta)
            check('quadratic_stability_endpoints',alpha*lam>0 and 2*(1+beta)-alpha*lam>0)
        check('quadratic_repeated_endpoints',1+beta-alpha*mu==2*q and 1+beta-alpha*L==-2*q)
    histories=[]
    for lam in [F(1),F(25)]:
        old=now=F(1);row=[now]
        for j in range(1,51):
            new=(1+F(4,9)-lam/9)*now-F(4,9)*old;old,now=now,new;row.append(new)
            exact=(1+F(j,3))*F(2,3)**j if lam==1 else (1+F(5*j,3))*(-F(2,3))**j
            check('heavy_ball_modal_closed_form',new==exact)
        histories.append(row)
    fvals=[(histories[0][j]**2+25*histories[1][j]**2)/2 for j in range(51)]
    check('heavy_ball_initial_overshoot',fvals[0]==13 and fvals[1]==F(3232,81)>13)
    first=next(j for j,v in enumerate(fvals)if v<=F(1,100));check('heavy_ball_actual_budget',first==18)
    cycle=(F(792,1225),-F(2208,1225),F(2592,1225))
    check('nonlinear_cycle_branches',cycle[0]<1 and cycle[1]<1 and cycle[2]>2)
    for j in range(3):
        cur=cycle[j];prev=cycle[(j-1)%3];nxt=F(13,9)*cur-F(4,9)*prev-hb_gradient(cur)/9
        check('nonlinear_cycle_exact_update',nxt==cycle[(j+1)%3])
    sample=[F(j,10)for j in range(-20,41)]
    for x,y in product(sample,repeat=2):
        if x<y:check('nonlinear_strong_monotonicity_finite_pairs',1<=(hb_gradient(y)-hb_gradient(x))/(y-x)<=25)
    check('nonlinear_piece_values',hb_objective(F(1))==F(25,2) and hb_objective(F(2))==38)
    pp=((F(4,3),-F(4,9)),(F(1),F(0)));pm=((-F(4,3),-F(4,9)),(F(1),F(0)));P=mm(pm,pp)
    check('switched_matrix',P==((-F(20,9),F(16,27)),(F(4,3),-F(4,9))))
    tr=P[0][0]+P[1][1];det=P[0][0]*P[1][1]-P[0][1]*P[1][0]
    check('switched_unstable_root_sign',tr==-F(8,3) and det==F(16,81) and 1+tr+det==-F(119,81)<0)
    state=(F(1),F(1));seq=[]
    for j in range(40):
        base=state;state=mv(pm,mv(pp,state));direct=mv(P,base)
        check('switched_direct_product_replay',state==direct)
        seq.append(state)
    check('switched_initial_not_eigenline',mv(P,(F(1),F(1)))[0]!=mv(P,(F(1),F(1)))[1])
    check('switched_first_points',seq[0]==(-F(44,27),F(8,9)) and seq[1]==(F(112,27),-F(208,81)))

    # Standard two-node examples and explicit rejection before a domain call.
    A=((F(1,2),F(1,4)),(F(0),F(1,4)));b=(F(1,4),F(3,4));g=lambda x:add(mv(A,x),b)
    t,y,r0,r1=two_history(g,(F(0),F(0)),b);ry=sub(g(y),y)
    check('two_dimensional_aa_exact',t==-F(11,41) and y==(F(53,82),F(81,82)) and ry==(F(57,328),F(3,328)))
    check('two_dimensional_aa_guard',infnorm(ry)/infnorm(r1)==F(114,205)<F(3,5))
    gp=lambda x:(F(9,10)*x[0]+F(1,10),)if x[0]<=F(1,2)else(x[0]/10+F(1,2),)
    t,y,r0,r1=two_history(gp,(F(0),),(F(1,10),))
    check('piecewise_aa_rejection',t==-9 and y==(F(1),) and abs(gp(y)[0]-y[0])==F(2,5)>F(3,5)*abs(r1[0]))
    check('piecewise_fallback',gp((F(1,10),))==(F(19,100),))
    A3=((F(1,4),F(1),F(0)),(F(0),F(1,4),F(1)),(F(0),F(0),F(1,4)));w=(F(16),F(4),F(1));b3=sub(w,mv(A3,w));gg=lambda x:add(mv(A3,x),b3)
    wn=lambda x:max(abs(v)/wi for v,wi in zip(x,w));inside=lambda x:all(0<=v<=wi for v,wi in zip(x,w))
    check('weighted_contraction_constant',max(sum(abs(A3[i][j])*w[j]/w[i]for j in range(3))for i in range(3))==F(1,2))
    t,yy,_,rr1=two_history(gg,(F(0),)*3,b3);ry=sub(gg(yy),yy)
    check('weighted_aa_coefficients',t==-F(4363,4321) and yy==tuple(F(v,4321)for v in (69304,19497,4869)))
    check('weighted_domain_rejection',not inside(yy) and all(v>wi for v,wi in zip(yy,w)))
    check('weighted_misleading_residual',wn(ry)/wn(rr1)==F(6576,21605)<F(3,5))
    fallback=gg(b3);rf=sub(gg(fallback),fallback)
    check('weighted_fallback_certificate',fallback==(F(12),F(13,4),F(15,16)) and rf==(F(9,4),F(1,2),F(3,64)) and wn(rf)==F(9,64) and wn(sub(fallback,w))==F(1,4)<=2*wn(rf))
    # Different contractive matrices and initial states: verify every actual guard branch.
    for a,c,d in product([F(0),F(1,4),F(1,2)],repeat=3):
        M=((a,c),(F(0),d));q=max(a+c,d)
        if q>=1:continue
        ones=(F(1),F(1));bb=sub(ones,mv(M,ones));fun=lambda x:add(mv(M,x),bb)
        for start in product([F(0),F(1,3),F(1)],repeat=2):
            old=start;cur=fun(old)
            for _ in range(3):
                rk=sub(fun(cur),cur)
                if rk==(F(0),F(0)):break
                _,trial,_,_=two_history(fun,old,cur)
                accept=all(0<=v<=1 for v in trial) and infnorm(sub(fun(trial),trial))<=F(3,5)*infnorm(rk)
                nxt=trial if accept else fun(cur);rn=sub(fun(nxt),nxt)
                check('guarded_actual_residual_contract',infnorm(rn)<=max(q,F(3,5))*infnorm(rk))
                check('guarded_actual_error_certificate',infnorm(sub(nxt,ones))<=infnorm(rn)/(1-q))
                old,cur=cur,nxt

    # AFW uses all vertices and maintains the explicit representation.
    V=tuple(tuple(F(i==j)for j in range(3))for i in range(3));a=(F(4,5),F(1,2),-F(3,10));ww=(F(1,3),)*3
    s1=afw_step(V,ww,a);s2=afw_step(V,s1['w'],a);s3=afw_step(V,s2['w'],a)
    check('simplex_drop_then_good',s1['drop'] and s1['gamma']==F(1,2) and s1['x']==(F(1,2),F(1,2),F(0)) and not s2['drop'] and s2['gamma']==F(3,10))
    check('simplex_exact_optimality',s2['x']==(F(13,20),F(7,20),F(0)) and s3['gap']==0 and objective(s2['x'],a)==F(27,400))
    for center in product([F(-1,2),F(1,3),F(1)],repeat=3):
        optimum=simplex_projection(center);fstar=objective(optimum,center);pg=sub(optimum,center)
        level=min(pg);check('simplex_projection_kkt',sum(optimum)==1 and all(v>=0 for v in optimum) and all(v==0 or pg[j]==level for j,v in enumerate(optimum)))
        for initial in [(F(1),F(0),F(0)),(F(1,3),)*3,(F(1,2),F(1,3),F(1,6))]:
            ww=initial;s0=sum(v>0 for v in ww);good=drop=0;C=None
            for it in range(6):
                step=afw_step(V,ww,center)
                if step['stop']:break
                oldx=step['oldx'];newx=step['x'];h=objective(oldx,center)-fstar;hn=objective(newx,center)-fstar
                if C is None:C=max(step['gap'],F(2))
                check('afw_weight_invariant',all(v>=0 for v in step['w']) and sum(step['w'])==1 and step['w']==newx)
                check('afw_gap_certificate',0<=h<=step['gap'])
                check('afw_descent_payment',h-hn>=step['gamma']*step['direction_gap']-step['gamma']**2*norm2(step['direction'])/2)
                if step['drop']:drop+=1
                else:
                    good+=1;check('afw_good_step_budget',hn<=h-h*h/(2*C))
                check('afw_drop_inventory',drop<=good+s0-1)
                check('afw_accumulated_value_bound',hn<=2*C/(good+2))
                ww=step['w']
    square=((F(0),F(0)),(F(1),F(0)),(F(1),F(1)),(F(0),F(1)));target=(F(1),F(1));rows=[]
    for initial,expected,gamma in [((F(1,10),F(1,5),F(1,2),F(1,5)),F(7,9),F(1,9)),((F(1,5),F(1,10),F(3,5),F(1,10)),F(7,8),F(1,4))]:
        st=afw_step(square,initial,target);nxt=afw_step(square,st['w'],target)
        check('same_point_distinct_state',st['oldx']==(F(7,10),)*2 and st['x']==(expected,)*2 and st['gamma']==gamma and st['drop'])
        check('same_point_post_gap',objective(st['x'],target)==(1-expected)**2 and nxt['gap']==2*(1-expected)**2)
        rows.append({'gamma':str(gamma),'x':[str(v)for v in st['x']],'weights':[str(v)for v in st['w']],'objective':str(objective(st['x'],target)),'gap':str(nxt['gap'])})
    return {'status':'PASS','checks':sum(counts.values()),'groups':dict(counts),'pl_bound_budget':k,'heavy_ball_quadratic_first_target_step':first,'heavy_ball_first_objective':str(fvals[1]),'nonlinear_cycle':[str(v)for v in cycle],'switched_first_two_states':[[str(v)for v in s]for s in seq[:2]],'weighted_anderson':{'theta':str(t),'candidate':[str(v)for v in yy],'domain_accepted':False,'fallback':[str(v)for v in fallback],'weighted_error_bound':'9/32'},'same_point_away_states':rows,'limits':['Finite parameter grids test implementations and named examples, not global PL or convergence theorems.','Sine comparisons use an explicit rational Taylor remainder; no floating-point tolerance is a certificate.','A low residual does not waive the declared domain check.','Good/drop counts follow the page definition; no unconditional half-good claim for arbitrary initial support.']}

def main():
    ap=argparse.ArgumentParser();ap.add_argument('--output',type=Path);args=ap.parse_args();result=run();s=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(s)
    print(s,end='')
if __name__=='__main__':main()
