#!/usr/bin/env python3
"""Recheck the A04 closed-form models, using 80 decimal digits.

Requires mpmath. Run:
  python foundations-periodic-ode-certificates.py --output result.json
Assertions test implementations, not universal interval certificates. Universal
bounds and hypotheses are proved in the accompanying entries. No ODE solver,
eigenvalue routine, finite difference step, or sampled residual is certified here.
"""
from pathlib import Path
from collections import Counter
import argparse, json
import mpmath as m
m.mp.dps = 80
counts = Counter()

def check(group, truth):
    counts[group] += 1
    if not truth:
        raise AssertionError((group, counts[group]))

def entries(x):
    return list(x) if isinstance(x, m.matrix) else [x]

def close(group, x, y, tol=m.mpf('1e-65')):
    if isinstance(x, m.matrix) or isinstance(y, m.matrix):
        check(group, isinstance(x, m.matrix) and isinstance(y, m.matrix)
              and x.rows == y.rows and x.cols == y.cols)
    xv, yv = entries(x), entries(y)
    check(group, len(xv) == len(yv))
    check(group, max(abs(a-b) for a,b in zip(xv,yv)) <=
          tol*(1+max(abs(a) for a in xv)+max(abs(b) for b in yv)))

def mdiff(fun,t):
    sample=fun(t)
    return m.matrix([[m.diff(lambda s:fun(s)[i,j],t)
                      for j in range(sample.cols)] for i in range(sample.rows)])

def mint(fun,a,b):
    sample=fun((a+b)/2)
    return m.matrix([[m.quad(lambda s:fun(s)[i,j],[a,b])
                      for j in range(sample.cols)] for i in range(sample.rows)])

I=m.eye(2); J=m.matrix([[0,-1],[1,0]]); D=m.diag([1,-3]); Ps=m.diag([0,1])
R=lambda t:m.matrix([[m.cos(t),-m.sin(t)],[m.sin(t),m.cos(t)]])
A=lambda t:3*J+R(3*t)*D*R(3*t).T
X=lambda t:R(3*t)*m.diag([m.exp(t),m.exp(-3*t)])
U=lambda t,s:R(3*t)*m.diag([m.exp(t-s),m.exp(-3*(t-s))])*R(3*s).T
P=lambda t:R(3*t)*Ps*R(3*t).T
g=lambda t:R(3*t)*m.matrix([1,1])
xs=lambda t:R(3*t)*m.matrix([-1,m.mpf(1)/3])
T=2*m.pi; T0=m.pi/3

# Nonlinear variation, analytic derivatives cross-checked with mpmath.diff.
for a in [m.mpf('.2'),m.mpf('.5'),m.mpf('1')]:
    for lam in [-m.mpf('.4'),m.mpf('.1'),m.mpf('.6')]:
        for t in [m.mpf(k)/20 for k in range(21)]:
            phi=lambda aa,ll,tt:aa/(1-ll*aa*tt)
            Y=lambda tt:(1-lam*a*tt)**-2
            Z=lambda tt:a*a*tt/(1-lam*a*tt)**2
            close('nonlinear_variation',m.diff(lambda aa:phi(aa,lam,t),a),Y(t))
            close('nonlinear_variation',m.diff(lambda ll:phi(a,ll,t),lam),Z(t))
            close('nonlinear_variation',m.diff(Y,t),2*lam*phi(a,lam,t)*Y(t))
            close('nonlinear_variation',m.diff(Z,t),2*lam*phi(a,lam,t)*Z(t)+phi(a,lam,t)**2)
# Explicit quadratic Taylor remainder in a compact positive tube for x'=x².
for h in [m.mpf(k)/1000 for k in range(-20,21)]:
    t=m.mpf('.25');a=m.mpf('.5');L=2;H=2
    actual=abs((a+h)/(1-(a+h)*t)-a/(1-a*t)-h/(1-a*t)**2)
    bound=H*t*m.exp(3*L*t)*h*h/2
    check('variation_remainder_bound',actual<=bound)
for t in [-2,-1,1,2]:
    check('lipschitz_not_C1_boundary',m.exp(t)!=m.exp(-t))
close('lipschitz_not_C1_boundary',m.exp(0),m.exp(-0))

# Frozen matrices have trace -2 and determinant 6, yet exact solutions grow.
for t in [m.mpf(k)/10 for k in range(-30,51)]:
    close('evolution_derivative',mdiff(X,t),A(t)*X(t))
    close('frozen_spectrum',A(t)[0,0]+A(t)[1,1],-2)
    close('frozen_spectrum',m.det(A(t)),6)
    close('minimal_period',A(t+T0),A(t))
    close('minimal_period',R(3*(t+T0)),-R(3*t))
    close('minimal_period',X(t+T0),X(t)*(-m.diag([m.exp(T0),m.exp(-3*T0)])))
    close('projection',P(t)*P(t),P(t))
    close('forced_solution',mdiff(xs,t),A(t)*xs(t)+g(t))
    close('forced_solution',xs(t+T),xs(t))
    close('growth',m.norm(X(t)*m.matrix([1,0])),m.exp(t))
    close('growth',m.norm(X(t)*m.matrix([0,1])),m.exp(-3*t))
close('average_matrix',mint(A,0,T)/T,3*J-I)
close('monodromy',X(T),m.diag([m.exp(T),m.exp(-3*T)]))
check('minimal_period',m.norm(A(0)-A(T0/2))>1)
# Complex Floquet factor at the minimal coefficient period.
B=m.diag([1+3j,-3+3j])
Pc=lambda t:X(t)*m.diag([m.exp(-(1+3j)*t),m.exp(-(-3+3j)*t)])
for t in [m.mpf(k)/7 for k in range(-8,9)]:
    close('complex_log_branch',Pc(t+T0),Pc(t))
    close('complex_log_branch',Pc(t)*m.diag([m.exp((1+3j)*t),m.exp((-3+3j)*t)]),X(t))
for t in [m.mpf(k)/5 for k in range(-10,11)]:
    for s in [m.mpf(k)/5 for k in range(-10,11)]:
        close('evolution_composition',U(t,s)*U(s,m.mpf('.37')),U(t,m.mpf('.37')))
        close('projection_transport',U(t,s)*P(s),P(t)*U(t,s))
        if t>=s:
            close('dichotomy_rank_one_norm',m.norm(U(t,s)*P(s)),m.exp(-3*(t-s)))
            check('dichotomy_bound',m.norm(U(t,s)*P(s))<=m.exp(-(t-s))*(1+m.mpf('1e-65')))
        else:
            close('dichotomy_rank_one_norm',m.norm(U(t,s)*(I-P(s))),m.exp(t-s))
# Real Floquet obstruction example has doubled real factor period.
XX=lambda t:R(m.pi*t)*m.diag([m.power(2,t),m.power(2,-t)])
AA=lambda t:m.pi*J+R(m.pi*t)*m.diag([m.log(2),-m.log(2)])*R(m.pi*t).T
close('real_log_obstruction',XX(1),m.diag([-2,-m.mpf('.5')]))
for t in [m.mpf(k)/9 for k in range(-9,10)]:
    close('real_log_obstruction',mdiff(XX,t),AA(t)*XX(t))
    close('real_log_obstruction',AA(t+1),AA(t))
    close('real_log_obstruction',R(m.pi*(t+2)),R(m.pi*t))
# Unit-circle Jordan growth: exact integer matrix powers, not spectral-only logic.
N=m.matrix([[0,1],[0,0]])
for k in range(31):close('jordan_boundary',(I+N)**k,I+k*N)

# Periodic forcing, left-null resonance, and an all-window residual bound.
r=mint(lambda s:U(T,s)*g(s),0,T)
close('periodic_boundary',(I-X(T))*xs(0),r)
Bosc=m.matrix([[0,1],[-1,0]])
V=lambda t:R(-t)
for freq in [1,2,3,4]:
    rr=mint(lambda s:V(T-s)*m.matrix([0,m.cos(freq*s)]),0,T)
    close('resonance_increment',rr,m.matrix([0,m.pi]) if freq==1 else m.zeros(2,1))
for delta in [m.mpf('.01'),m.mpf('.1'),m.mpf('.5'),m.mpf('2')]:
    C=1/(1-m.exp(-delta*T))
    for eta,nu in [(m.mpf('.01'),m.mpf('.02')),(m.mpf('-.03'),m.mpf('.01')),(m.mpf('0'),m.mpf('-.02'))]:
        residual_bound=(1+delta)*abs(eta)+(1/T+delta)*abs(nu)
        error_bound=C*abs(nu)+T*(1+C)*residual_bound
        for k in range(101):
            t=T*k/100
            err=eta*m.sin(t)+nu*t/T
            residual=eta*m.cos(t)+nu/T+delta*err
            check('finite_window_residual',abs(residual)<=residual_bound+m.mpf('1e-70'))
            check('finite_window_error_certificate',abs(err)<=error_bound)
# Green integrals truncated at both ends: an independent numerical integration.
for t in [m.mpf('-.4'),m.mpf('0'),m.mpf('.7')]:
    for L in [m.mpf('.1'),m.mpf('1'),m.mpf('3')]:
        val=mint(lambda s:U(t,s)*P(s)*g(s),t-L,t)-mint(lambda s:U(t,s)*(I-P(s))*g(s),t,t+L)
        exact=R(3*t)*m.matrix([-(1-m.exp(-L)),(1-m.exp(-3*L))/3])
        close('green_quadrature',val,exact)
for k in range(1,201):
    L=m.mpf(k)/10
    exact_tail=m.sqrt(m.exp(-2*L)+m.exp(-6*L)/9)
    check('green_tail_bound',exact_tail<=2*m.sqrt(2)*m.exp(-L))
check('green_L20_tolerance',2*m.sqrt(2)*m.exp(-20)<m.mpf('1e-8'))
for t in [m.mpf(k)/10 for k in range(-20,21)]:
    # Global bounded perturbation z=xs+eta R(3t)(sin t, cos t).
    eta=m.mpf('.01');v=lambda tt:eta*R(3*tt)*m.matrix([m.sin(tt),m.cos(tt)])
    residual=mdiff(v,t)-A(t)*v(t)
    check('all_line_residual',m.norm(residual)<=eta*m.sqrt(20))
    check('all_line_residual_certificate',m.norm(v(t))<=2*eta*m.sqrt(20))

# Exact nonlinear return map, transverse multiplier, and radius bound.
def rad(t,r0,a):return (1/a+(1/r0**2-1/a)*m.exp(-2*a*t))**(-m.mpf('.5'))
def per(lam):return 2*m.pi/(2+lam)
def ret(r,lam):return rad(per(lam),r,1+lam)
for lam in [m.mpf(k)/20 for k in range(-9,10)]:
    a=1+lam;rs=m.sqrt(a);TT=per(lam)
    close('return_fixed_point',ret(rs,lam),rs)
    rho=m.exp(-2*a*TT)
    close('return_multiplier',m.diff(lambda r:ret(r,lam),rs),rho)
    check('return_strict_contraction',0<rho<1)
    for r0 in [m.mpf('.5'),m.mpf('.8'),m.mpf('1.2'),m.mpf('1.5')]:
        for t in [m.mpf('.1'),m.mpf('1'),m.mpf('3')]:
            close('radial_ODE',m.diff(lambda tt:rad(tt,r0,a),t),rad(t,r0,a)*(a-rad(t,r0,a)**2))
for i in range(21):
    r0=m.mpf('.5')+m.mpf(i)/20
    for k in range(61):
        t=m.mpf(k)/10;r=rad(t,r0,m.mpf(1))
        check('invariant_annulus',m.mpf('.5')-m.mpf('1e-70')<=r<=m.mpf('1.5')+m.mpf('1e-70'))
        check('orbital_exp_bound',abs(r-1)<=m.exp(-3*t/4)*abs(r0-1)+m.mpf('1e-70'))
# State-dependent return time: fixed-time tangent drift must be projected away.
Fshear=lambda rr:m.matrix([rr*m.cos(2*m.pi*rr),rr*m.sin(2*m.pi*rr)])
close('return_time_projection',mdiff(Fshear,m.mpf(1)),m.matrix([1,2*m.pi]))
close('return_time_projection',m.diff(lambda r:2*m.pi/r,m.mpf(1)),-2*m.pi)
Proj=m.diag([1,0]);close('return_time_projection',Proj*m.matrix([1,2*m.pi]),m.matrix([1,0]))
# Parameter derivative at fixed initial radius and fixed time is not p'(lambda).
def fixed_endpoint(lam):
    radius=rad(m.pi,m.mpf(1),1+lam);angle=(2+lam)*m.pi
    return radius*m.matrix([m.cos(angle),m.sin(angle)])
rho=m.exp(-2*m.pi);w=mdiff(fixed_endpoint,m.mpf(0))
close('parameter_variation',w,m.matrix([(1-rho)/2,m.pi]))
H=m.matrix([[rho-1,0,0],[0,0,2],[0,1,0]])
answer=m.lu_solve(H,m.matrix([-w[0],-w[1],0]))
close('bordered_period_derivative',answer,m.matrix([m.mpf('.5'),0,-m.pi/2]))
close('bordered_period_derivative',m.diff(per,m.mpf(0)),answer[2])
close('bordered_period_derivative',m.diff(lambda lam:m.sqrt(1+lam),m.mpf(0)),answer[0])
q=m.matrix([0,m.mpf('.5')]);close('adjoint_normalization',(q.T*m.matrix([0,2]))[0],1)
close('adjoint_period_derivative',-(q.T*w)[0],answer[2])
# Adjoint integral in physical coordinates, independently from terminal w.
YY=lambda t:R(2*t)*m.diag([m.exp(-2*t),1])
psi=lambda t:YY(t).T**-1*q
source=lambda t:R(2*t)*m.matrix([1,1])
close('adjoint_period_integral',-m.quad(lambda t:(psi(t).T*source(t))[0],[0,m.pi]),-m.pi/2)
for k in range(21):
    t=m.pi*k/20
    # Df on the unit circle is -2 zz^T + 2J.
    jac=-2*(R(2*t)*m.matrix([1,0]))*(R(2*t)*m.matrix([1,0])).T+2*J
    close('adjoint_ODE',mdiff(psi,t),-jac.T*psi(t))
close('adjoint_periodicity',psi(m.pi),psi(0))
# Neutral multiplier cannot decide nonlinear stability; fold has two/zero circles.
for u0 in [m.mpf('.01'),m.mpf('.1'),m.mpf('-.1')]:
    for t in [0,1,10,100]:
        stable=lambda s:u0/m.sqrt(1+2*u0*u0*s)
        close('neutral_cubic_boundary',m.diff(stable,t),-stable(t)**3)
    crossing=(1-(abs(u0)/m.mpf('.2'))**2)/(2*u0*u0)
    close('neutral_cubic_boundary',abs(u0/m.sqrt(1-2*u0*u0*crossing)),m.mpf('.2'))
for lam in [-m.mpf('.01'),-m.mpf('.04')]:
    for sign in [-1,1]:close('fold_boundary',lam+(sign*m.sqrt(-lam))**2,0)
check('fold_boundary',m.mpf('.01')>0)

result={'status':'passed','decimal_precision':m.mp.dps,'assertions':sum(counts.values()),
        'groups':dict(sorted(counts.items())),
        'selected_values':{'T':m.nstr(T,35),'growing_multiplier':m.nstr(m.exp(T),35),
                           'transverse_multiplier':m.nstr(rho,35),'period_derivative':m.nstr(answer[2],35),
                           'L20_general_tail_bound':m.nstr(2*m.sqrt(2)*m.exp(-20),35)},
        'scope':'Closed-form models and derived bounds, cross-checked by differentiation and quadrature; universal claims rely on proofs in the entries.',
        'limits':['Equality-bound comparisons permit 1e-70 rounding slack at 80-digit working precision','No interval-arithmetic integration certificate','Finite samples do not prove global residual bounds','Does not test production site rendering']}
args=argparse.ArgumentParser();args.add_argument('--output',type=Path);args=args.parse_args()
text=json.dumps(result,ensure_ascii=False,indent=2)+'\n'
if args.output:args.output.write_text(text)
print(text,end='')
