"""Independent thermal ladder verification: python verify.py (NumPy, SciPy).
All units SI. Uses generalized symmetric eigenvalues and adaptive time marching.
"""
import json
import numpy as np
from scipy.linalg import eigh
from scipy.integrate import solve_ivp
from scipy.optimize import brentq
from scipy.stats import gamma


def analyze(n=10, forcing='free', source=5, capacities=None, links=None):
    c = np.array(capacities if capacities is not None else [1/n]*n)
    r = np.array(links if links is not None else [1/n]*n)
    k = np.zeros((n, n))
    k[0, 0] += 1/r[0]
    for i in range(1, n):
        g = 1/r[i]
        k[i, i] += g
        k[i-1, i-1] += g
        k[i, i-1] -= g
        k[i-1, i] -= g
    if forcing in ('both', 'left', 'qboth'):
        k[-1, -1] += n
    b = np.zeros(n)
    if forcing.startswith('q'):
        b[source-1] = 1
    else:
        b[0] += 1/r[0]
        if forcing == 'both':
            b[-1] += n
    x = np.linalg.solve(k, b)
    v = np.linalg.solve(k, c*x)
    w = np.linalg.solve(k, c*v)
    rates, vectors = eigh(k, np.diag(c))
    tau = 1/rates[0]
    coefficients = vectors * (vectors.T @ (c*x))[None, :]
    def response(j, t):
        return 1-np.sum(coefficients[j]*np.exp(-rates*t))/x[j]
    # A separate integration path, not the modal solution.
    def integrate(rtol, atol):
        return solve_ivp(lambda t, y: (b-k@y)/c, (0, 8*tau), np.zeros(n),
                         method='DOP853', rtol=rtol, atol=atol, max_step=0.25/rates[-1], dense_output=True)
    coarse = integrate(1e-8, 1e-10)
    fine = integrate(1e-11, 1e-13)
    assert coarse.success and fine.success
    rows = []
    for j in range(n):
        mu = v[j]/x[j]
        variance = 2*w[j]/x[j]-mu*mu
        modal = [brentq(lambda t: response(j, t)-p, 0, 8*tau) for p in (.1, .5, 1-np.exp(-1), .9)]
        numeric = [brentq(lambda t: fine.sol(t)[j]/x[j]-p, 0, 8*tau) for p in (.1, .5, 1-np.exp(-1), .9)]
        old = [brentq(lambda t: coarse.sol(t)[j]/x[j]-p, 0, 8*tau) for p in (.1, .9)]
        np.testing.assert_allclose(modal, numeric, rtol=2e-7, atol=2e-9)
        np.testing.assert_allclose(numeric[3]-numeric[0], old[1]-old[0], rtol=2e-6)
        shape = mu*mu/variance
        estimate = variance/mu*(gamma.ppf(.9, shape)-gamma.ppf(.1, shape))
        rows.append(dict(mu=mu, variance=variance, t10=modal[0], t50=modal[1],
                         t63=modal[2], t90=modal[3], rise=modal[3]-modal[0],
                         marched_rise=numeric[3]-numeric[0], rcRise=np.log(9)*mu,
                         gammaRise=estimate))
    return dict(n=n, forcing=forcing, tau=tau, x=x.tolist(), stats=rows)


if __name__ == '__main__':
    results = [analyze(n=n, source=1) for n in (1, 5, 10)]
    results += [analyze(forcing=f) for f in ('both', 'left', 'qboth', 'qfree')]
    results += [analyze(n=3, source=1, capacities=[3.5, .8, .2], links=[.5025, 1.2525, 1.3])]
    print(json.dumps(results, indent=2))

