#!/usr/bin/env python3
"""Run all C5 examples and the solved capstone.

Requires Python 3, NumPy, SciPy and mpmath (no network or files as inputs).
Run: python verify_matrix_functions.py [--output results.json]
The teaching kernels are below; SciPy and 80-digit mpmath serve as independent
oracles. Schur factorization, linear solves and small dense products are library
primitives, not reimplementations of those prerequisite algorithms.
Operation counters exclude oracle calls and explicitly count multiple-RHS solves.
"""
import argparse
import json
import math
import platform
from fractions import Fraction as F
import numpy as np
import scipy
import scipy.linalg as la
import mpmath as mp

mp.mp.dps = 80
CHECKS = 0

def check(condition, message):
    global CHECKS
    CHECKS += 1
    if not condition:
        raise AssertionError(message)

def close(a, b, tol=2e-12):
    aa, bb = np.asarray(a), np.asarray(b)
    check(la.norm(aa-bb) <= tol * max(1.0, la.norm(bb)), 'numerical comparison failed')

def mp_array(a):
    a = np.asarray(a)
    return mp.matrix([[mp.mpf(float(x)) for x in row] for row in a])

def mp_float(a):
    return np.array(a.tolist(), dtype=float)

def triangular_sylvester(T, S, D):
    """Complex/real UPPER TRIANGULAR AX-XB=C, not real quasi-triangular."""
    m, n = D.shape
    Y = np.zeros((m, n), dtype=np.result_type(T, S, D, float))
    trace = []
    for j in range(n):
        for i in range(m-1, -1, -1):
            rhs = D[i,j] + Y[i,:j] @ S[:j,j] - T[i,i+1:] @ Y[i+1:,j]
            den = T[i,i] - S[j,j]
            if den == 0:
                raise ValueError('common eigenvalues')
            Y[i,j] = rhs / den
            trace.append(dict(row=i, col=j, rhs=float(np.real(rhs)),
                              divisor=float(np.real(den)), value=float(np.real(Y[i,j]))))
    return Y, trace

def bartels_stewart(A, B, C):
    T, Q = la.schur(A, output='complex')
    S, U = la.schur(B, output='complex')
    Y, trace = triangular_sylvester(T, S, Q.conj().T @ C @ U)
    return np.real_if_close(Q @ Y @ U.conj().T), trace

COEFF13 = [64764752532480000,32382376266240000,7771770303897600,
           1187353796428800,129060195264000,10559470521600,670442572800,
           33522128640,1323241920,40840800,960960,16380,182,1]
THETA13 = 5.371920351148152

def expm13(A, forced_s=None):
    """Unshifted fixed [13/13] Padé, norm-based 2005 threshold; not adaptive 2009."""
    A = np.asarray(A, dtype=np.result_type(A, float))
    n = len(A)
    I = np.eye(n, dtype=A.dtype)
    norm = la.norm(A, 1)
    if norm == 0:
        return I, dict(s=0,matrix_multiplications=0,multiple_rhs_solves=0,pade_relative_residual=0.)
    s = max(0, math.ceil(math.log2(norm/THETA13))) if norm else 0
    if forced_s is not None:
        s = forced_s
    X = A / (2.0**s)
    X2 = X@X; X4 = X2@X2; X6 = X2@X4
    b = COEFF13
    U = X@(X6@(b[13]*X6+b[11]*X4+b[9]*X2)+b[7]*X6+b[5]*X4+b[3]*X2+b[1]*I)
    V = X6@(b[12]*X6+b[10]*X4+b[8]*X2)+b[6]*X6+b[4]*X4+b[2]*X2+b[0]*I
    R = la.solve(V-U, V+U)
    core = R.copy()
    for _ in range(s):
        R = R@R
    return R, dict(s=s, matrix_multiplications=6+s, multiple_rhs_solves=1,
                   pade_relative_residual=float(la.norm((V-U)@core-(V+U))/la.norm(V+U)))

def exp_taylor_block(T, tol=2e-16):
    """Entire exp only. Stop by a norm remainder bound, not one small term."""
    n = len(T)
    sigma = np.trace(T)/n
    M = T-sigma*np.eye(n)
    q = la.norm(M, 1)
    P = np.eye(n); Fv = P.copy()
    tail = math.exp(q)*q
    k = 0
    while math.exp(float(sigma))*tail > tol*max(1.0, la.norm(math.exp(float(sigma))*Fv,1)):
        k += 1
        P = P@M/k
        Fv = Fv+P
        tail = math.exp(q)*q**(k+1)/math.factorial(k+1)
        if k > 200:
            raise RuntimeError('Taylor cap')
    return math.exp(float(sigma))*Fv, k, math.exp(float(sigma))*tail

def parlett_cluster_exp(T):
    """Worked 3x3 partition {0,1}|{2}; no general clustering claimed."""
    F11, k, bound = exp_taylor_block(T[:2,:2])
    f33 = math.exp(T[2,2])
    rhs = F11@T[:2,2] - T[:2,2]*f33
    f13 = la.solve(T[:2,:2]-T[2,2]*np.eye(2), rhs)
    R = np.zeros((3,3)); R[:2,:2] = F11; R[:2,2] = f13; R[2,2] = f33
    return R, dict(taylor_products=k, diagonal_block_truncation_bound=bound,
                   off_diagonal_solves=1)

def triangular_sqrt(T):
    """Positive-real-spectrum upper triangular teaching kernel."""
    n = len(T)
    R = np.zeros_like(T, dtype=float)
    if np.any(np.diag(T) <= 0):
        raise ValueError('this teaching kernel requires positive real diagonal')
    for i in range(n):
        R[i,i] = math.sqrt(T[i,i])
    for gap in range(1,n):
        for i in range(n-gap):
            j = i+gap
            R[i,j] = (T[i,j]-R[i,i+1:j]@R[i+1:j,j])/(R[i,i]+R[j,j])
    return R

def log_inverse_ss(T, s=None, m=10):
    """Illustrative triangular inverse scaling, Gauss-Legendre partial fractions.
    Auto roots until ||X||_1 <= 1/4. Bound is exact-arithmetic truncation ONLY.
    Does not implement all 2012 cancellation corrections or adaptive degrees.
    """
    R = np.array(T, dtype=float)
    n = len(R); I = np.eye(n); roots = []
    k = 0
    while (k < s if s is not None else la.norm(R-I,1) > .25):
        R = triangular_sqrt(R); roots.append(R.copy()); k += 1
        if k > 60:
            raise RuntimeError('root cap')
    X = R-I
    nodes, weights = np.polynomial.legendre.leggauss(m)
    nodes = (nodes+1)/2; weights = weights/2
    L = np.zeros_like(T, dtype=float)
    for node, weight in zip(nodes, weights):
        L += weight*la.solve(I+node*X, X)
    L *= 2**k
    r = la.norm(X,1)
    bound = 2**k * 2*r**(2*m+1)/(1-r) if r < 1 else None
    return L, dict(s=k, m=m, roots=roots, radius_1=float(r),
                   multiple_rhs_solves=m, truncation_bound_1=bound)

def denman_beavers(A, steps=7):
    Y = np.array(A, dtype=float); Z = np.eye(len(A)); I = Z.copy()
    trace = []
    for k in range(steps+1):
        trace.append(dict(k=k, Y=Y.copy(), Z=Z.copy(),
                          square_residual=float(la.norm(Y@Y-A)),
                          inverse_residual=float(la.norm(Y@Z-I))))
        if k < steps:
            # Both updates use OLD Y,Z; never overwrite one before forming the other.
            Y, Z = (Y+la.solve(Z,I))/2, (Z+la.solve(Y,I))/2
    return Y, Z, trace

def polar_newton(A, steps=7):
    X = np.array(A, dtype=float); I = np.eye(len(A)); trace = []
    for k in range(steps+1):
        trace.append(dict(k=k, X=X.copy(), singular_values=la.svdvals(X),
                          orthogonality_residual=float(la.norm(X.T@X-I))))
        if k < steps:
            X = (X+la.solve(X.T,I))/2
    return X, trace

def arnoldi_action(A, b, m, t=1.):
    n = len(b); beta = la.norm(b)
    if beta == 0:
        return np.zeros(n), dict(matvecs=0, happy=True, defect=0.)
    V = np.zeros((n,m+1)); H = np.zeros((m+1,m)); V[:,0] = b/beta
    happy = False
    for j in range(m):
        w = A@V[:,j]
        for _ in range(2):
            for i in range(j+1):
                h = V[:,i]@w; H[i,j] += h; w -= h*V[:,i]
        H[j+1,j] = la.norm(w)
        if H[j+1,j] < 1e-14*max(1.,la.norm(A)):
            m = j+1; happy = True; break
        V[:,j+1] = w/H[j+1,j]
    K = H[:m,:m]; z = la.expm(t*K)[:,0]*beta
    y = V[:,:m]@z
    defect = 0. if happy else abs(H[m,m-1]*z[-1])
    return y, dict(matvecs=m, happy=happy, defect=float(defect), H=K,
                   h_next=float(H[m,m-1]), basis=V[:,:m])

def run():
    out = {'environment': dict(python=platform.python_version(), numpy=np.__version__,
                              scipy=scipy.__version__, mpmath=mp.__version__, mp_dps=mp.mp.dps)}
    A = np.array([[2,1,0],[0,2,0],[0,0,5.]])
    p = 6/5*np.eye(3)-9/20*A+A@A/20
    q = .7*np.eye(3)-A/10
    close(A@p,np.eye(3)); close((A@q)[0,1],.3)
    out['primary'] = dict(inverse=p, value_only_interpolant=q, wrong_product=A@q)
    out['separation'] = []
    for M in [0,10,100,10000]:
        A = np.array([[1.,M],[0,2.]])
        x = la.solve(A,[0,1]); sep = la.svdvals(A)[-1]
        formula = 2/math.sqrt(((M*M+5)+math.sqrt((M*M+5)**2-16))/2)
        close(sep,formula); close(x,[-M/2,.5])
        out['separation'].append(dict(M=M, gap=1, sep=float(sep), solution=x))
    T = np.array([[1,2,-1],[0,2,3],[0,0,4.]])
    S = np.array([[-1,1,2],[0,-2,-1],[0,0,-3.]])
    D = np.array([[-2,7,5],[3,5,10],[10,-2,-11.]])
    Y, trace = triangular_sylvester(T,S,D)
    close(Y,[[1,2,0],[-1,1,2],[2,0,-1]]); close(T@Y-Y@S,D)
    Q = np.array([[.6,-.8,0],[.8,.6,0],[0,0,1]])
    U = np.eye(3)[:,[2,0,1]]
    A=Q@T@Q.T; B=U@S@U.T; C=Q@D@U.T
    X,_ = bartels_stewart(A,B,C)
    close(X,Q@Y@U.T); close(A@X-X@B,C)
    Ar=np.array([[0.,-1],[1,0]]); Br=np.array([[2.]]); Cr=np.array([[1.],[0]])
    Yr,scale,info = la.lapack.dtrsyl(Ar,Br,Cr,isgn=-1)
    close(Yr,[[-.4],[-.2]]); close(Ar@Yr-Yr@Br,scale*Cr); check(info==0,'DTRSYL status')
    out['bartels_stewart']=dict(T=T,S=S,D=D,Y=Y,trace=trace,original_solution=X,
                                original_residual=float(la.norm(A@X-X@B-C)),
                                real_block=dict(Y=Yr,scale=scale,info=info))
    rng=np.random.default_rng(503)
    for n in range(1,6):
        for _ in range(5):
            A=rng.normal(size=(n,n))+5*np.eye(n); B=rng.normal(size=(n,n))-5*np.eye(n)
            C=rng.normal(size=(n,n)); X,_=bartels_stewart(A,B,C)
            close(A@X-X@B,C,2e-11); close(X,la.solve_sylvester(A,-B,C),2e-11)
    A=np.diag([0.,1.]); E=np.array([[0.,1.],[1,0]])
    K=np.block([[A,E],[np.zeros_like(A),A]])
    L=la.expm(K)[:2,2:]; close(L,math.expm1(1)*E)
    close(L,la.expm_frechet(A,E,compute_expm=False))
    out['frechet']=dict(A=A,E=E,L=L,wrong=la.expm(A)@E,
                        relative_condition_F=math.e/math.sqrt(1+math.e**2))
    A=np.array([[-1.,20],[0,-2]])
    exact=np.array([[math.exp(-1),20*(math.exp(-1)-math.exp(-2))],[0,math.exp(-2)]])
    R,counts=expm13(A); close(R,exact)
    X=A/2; P=la.solve(np.eye(2)-X/2,np.eye(2)+X/2)
    close(P,[[.6,16/3],[0,1/3]]); close(P@P,[[9/25,224/45],[0,1/9]])
    out['exponential']=dict(A=A,exact=exact,low_order_P=P,low_order_square=P@P,
                            pade13=R,counts=counts)
    N=np.array([[0.,1e8],[0,0]])
    RN,cn=expm13(N); RN0,cn0=expm13(N,0)
    close(RN,np.eye(2)+N,1e-6); close(RN0,np.eye(2)+N)
    check(cn['s']==25 and cn0['s']==0,'overscaling operation count')
    out['nilpotent_overscaling']=dict(norm_based=cn,no_scaling=cn0,
        norm_based_relative_error=float(la.norm(RN-(np.eye(2)+N))/la.norm(np.eye(2)+N)),
        no_scaling_relative_error=float(la.norm(RN0-(np.eye(2)+N))/la.norm(np.eye(2)+N)),
        note='Observed floating-point error is reported, not asserted to be a platform-independent constant.')
    for n in [2,3,5]:
        for _ in range(8):
            A=rng.normal(size=(n,n))*rng.choice([.01,1.,5.])
            R,_=expm13(A); ref=mp_float(mp.expm(mp_array(A)))
            close(R,ref,2e-11)
    out['near_defective_family']=[]
    for eps in [.1,1e-6,1e-12,0.]:
        T=np.array([[1.,2,1],[0,1+eps,3],[0,0,4]])
        exp_ref=mp_float(mp.expm(mp_array(T)))
        sqrt_ref=mp_float(mp.sqrtm(mp_array(T)))
        log_ref=mp_float(mp.logm(mp_array(T)))
        P,pc=parlett_cluster_exp(T); X,xc=expm13(T)
        Sr=triangular_sqrt(T); L,lc=log_inverse_ss(T)
        close(P,exp_ref); close(X,exp_ref); close(Sr,sqrt_ref); close(L,log_ref)
        close(Sr@Sr,T); close(la.expm(L),T)
        check(np.min(np.real(la.eigvals(Sr)))>0,'principal square-root branch')
        check(np.max(np.abs(np.imag(la.eigvals(L))))<math.pi,'principal logarithm strip')
        check(lc['truncation_bound_1']<4.11e-13,'log truncation bound')
        E=np.zeros((3,3)); E[1,0]=1
        K=np.block([[T,E],[np.zeros((3,3)),T]])
        Lexp=expm13(K)[0][:3,3:]
        Lref=mp_float(mp.expm(mp_array(K)))[:3,3:]
        close(Lexp,Lref)
        hrows=[]
        for h in [1e-2,1e-3,1e-4,1e-5,1e-6,1e-7]:
            fd=(expm13(T+h*E)[0]-expm13(T-h*E)[0])/(2*h)
            hrows.append(dict(h=h,error=float(la.norm(fd-Lref))))
        Kry,kr=arnoldi_action(T,np.array([0.,0.,1]),3)
        close(Kry,exp_ref[:,2])
        actual_eps=T[1,1]-1
        naive12=(math.exp(T[1,1])-math.e)*2/actual_eps if actual_eps else None
        out['near_defective_family'].append(dict(epsilon=eps,input=T,exp=exp_ref,sqrt=Sr,log=L,
            cluster_counts=pc,expm_counts=xc,log_counts=lc,frechet_E=E,frechet=Lexp,
            finite_difference=hrows,naive_f12=naive12,krylov=Kry,krylov_counts=kr))
    T0=out['near_defective_family'][-1]['input']
    S0=np.array([[1.,1,0],[0,1,1],[0,0,2]])
    L0=np.array([[0.,2,math.log(4)-2],[0,0,math.log(4)],[0,0,math.log(4)]])
    F0=np.array([[math.e,2*math.e,math.exp(4)-3*math.e],[0,math.e,math.exp(4)-math.e],[0,0,math.exp(4)]])
    close(S0@S0,T0); close(S0,out['near_defective_family'][-1]['sqrt'])
    close(L0,out['near_defective_family'][-1]['log']); close(F0,out['near_defective_family'][-1]['exp'])
    LF0=np.array([[math.e,2*math.e/3,(2*math.exp(4)-23*math.e)/9],
                   [math.e,math.e,(math.exp(4)-7*math.e)/3],[0,0,0]])
    close(LF0,out['near_defective_family'][-1]['frechet'])
    close(expm13(np.zeros((3,3)))[0],np.eye(3))
    E=np.zeros((3,3)); E[1,0]=1
    sep=float(la.svdvals(np.kron(np.eye(3),S0)+np.kron(S0.T,np.eye(3)))[-1])
    perturb=1e-6*E; root_res=(S0+perturb)@(S0+perturb)-T0
    close(root_res,S0@perturb+perturb@S0,1e-14)
    check(la.norm(perturb)<=la.norm(root_res)/sep,'root error certificate for nilpotent perturbation')
    out['root_separation']=dict(sep=sep,perturbation=perturb,residual=root_res,
         error=float(la.norm(perturb)),residual_norm=float(la.norm(root_res)),bound=float(la.norm(root_res)/sep))
    A=np.diag([-2.]*4)+np.diag([1.]*3,1)+np.diag([1.]*3,-1)
    ref=la.expm(A)[:,0]; rows=[]
    for m in range(1,5):
        y,details=arnoldi_action(A,np.eye(4)[:,0],m)
        H=details['H']; z=la.expm(H)[:,0]
        bound=0. if details['happy'] else details['h_next']*la.solve(H,z-np.eye(m)[:,0])[-1]
        err=la.norm(y-ref)
        check(err<=bound+1e-14,'contractive defect integral bound')
        rows.append(dict(m=m,y=y,error=float(err),defect=details['defect'],integral_bound=float(bound),details=details))
    out['krylov_diffusion']=dict(A=A,exact=ref,rows=rows)
    A=np.array([[4.,6],[0,9]])
    L,details=log_inverse_ss(A,s=2,m=3)
    exact=np.array([[math.log(4),6*math.log(9/4)/5],[0,math.log(9)]])
    LA,da=log_inverse_ss(A); close(LA,exact)
    Y,Z,db=denman_beavers(A)
    close(Y,[[2,1.2],[0,3]]); close(Y@Z,np.eye(2)); close(Y@Y,A)
    for state in db:
        close(state['Y'],A@state['Z'])
        check(np.min(np.real(la.eigvals(state['Y'])))>0,'DB root branch')
    out['logarithm']=dict(A=A,exact=exact,two_roots_m3=L,demonstration_counts=details,
                         demonstration_error_1=float(la.norm(L-exact,1)),auto=LA,auto_counts=da)
    out['denman_beavers']=dict(A=A,trace=db,multiple_rhs_solves=14,negative_root_square=(-Y)@(-Y))
    A=np.array([[2.,1],[0,1]])
    U,pol=polar_newton(A); H=U.T@A
    close(U,np.array([[3,1],[-1,3]])/math.sqrt(10)); close(U@H,A); close(H,H.T)
    check(np.min(la.eigvalsh(H))>0,'polar positive stretch')
    for state in pol[1:]:
        check(la.norm(state['X']-U)<=state['orthogonality_residual']/2+1e-14,'polar orthogonality error bound')
    out['polar']=dict(A=A,U=U,H=H,trace=pol,multiple_rhs_solves=7,
                      closest_distance=float(la.norm(A-U)),QR_distance=float(la.norm(A-np.eye(2))))
    out['checks_passed']=CHECKS
    return out

def convert(x):
    if isinstance(x,np.ndarray): return x.tolist()
    if isinstance(x,np.generic): return x.item()
    raise TypeError(type(x).__name__)

if __name__=='__main__':
    parser=argparse.ArgumentParser(); parser.add_argument('--output')
    args=parser.parse_args(); result=run()
    text=json.dumps(result,ensure_ascii=False,indent=2,default=convert,allow_nan=False)+'\n'
    if args.output:
        with open(args.output,'w') as f: f.write(text)
        print(f"{result['checks_passed']} checks passed; wrote {args.output}")
    else:
        print(text,end='')
