"""Independent unequal-capacity/resistance examples (NumPy + SciPy).
Run: python verify-boundary-cases.py --output unequal-reference.json
All nodes start at 25 C. Temperatures computed below are changes from 25 C.
"""
import json,argparse
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

CAPACITIES=[.08,.12,.20,.35,.50,.25,.15,.10,.15,.10]
LINKS=[.10,.15,.25,.30,.20,.40,.15,.10,.20,.15]
CASES=[
 dict(key='ambient-free',label='A. Ambient step; far end insulated',forcing='free',amplitude=10,extraHeat=0,rightR=None,ambientStep=10,heat=0),
 dict(key='heat-free',label='B. Fixed ambient + middle heat; far end insulated',forcing='qfree',amplitude=2,extraHeat=0,rightR=None,ambientStep=0,heat=2),
 dict(key='ambient-fixed',label='C. Ambient step; opposite reservoir fixed',forcing='left',amplitude=10,extraHeat=0,rightR=.25,ambientStep=10,heat=0),
 dict(key='heat-fixed',label='D. Both reservoirs fixed + middle heat',forcing='qboth',amplitude=2,extraHeat=0,rightR=.25,ambientStep=0,heat=2),
 dict(key='combined-fixed',label='E. Ambient step + middle heat; opposite reservoir fixed',forcing='left',amplitude=10,extraHeat=2,rightR=.25,ambientStep=10,heat=2)
]

def analyze(case):
    n=10;c=np.array(CAPACITIES);r=np.array(LINKS)
    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 case['rightR']:k[-1,-1]+=1/case['rightR']
    b=np.zeros(n);b[0]=case['ambientStep']/r[0];b[4]+=case['heat']
    x=np.linalg.solve(k,b);v=np.linalg.solve(k,c*x);w=np.linalg.solve(k,c*v)
    rates,vec=eigh(k,np.diag(c));co=vec*(vec.T@(c*x))[None,:]/x[:,None]
    tau=1/rates[0];end=8*tau
    response=lambda j,t:1-np.sum(co[j]*np.exp(-rates*t))
    solutions=[solve_ivp(lambda t,y:(b-k@y)/c,(0,end),np.zeros(n),method='DOP853',rtol=tol,atol=tol/100,max_step=.25/rates[-1],dense_output=True) for tol in [1e-8,1e-11]]
    assert all(sol.success for sol in solutions)
    rows=[]
    for j in range(n):
        mu=v[j]/x[j];variance=2*w[j]/x[j]-mu*mu;shape=mu*mu/variance;scale=variance/mu
        crossings=[brentq(lambda t:response(j,t)-p,0,end) for p in [.1,.9]]
        marches=[[brentq(lambda t:sol.sol(t)[j]/x[j]-p,0,end) for p in [.1,.9]] for sol in solutions]
        np.testing.assert_allclose(crossings,marches[1],atol=2e-9,rtol=2e-8)
        np.testing.assert_allclose(marches[0],marches[1],atol=2e-8,rtol=2e-7)
        rows.append(dict(node=j+1,steadyChange=x[j],finalTemperature=25+x[j],mu=mu,wNormalized=w[j]/x[j],variance=variance,shape=shape,scale=scale,t10=crossings[0],t90=crossings[1],rise=crossings[1]-crossings[0],oneMoment=np.log(9)*mu,twoMoment=gamma.ppf(.9,shape,scale=scale)-gamma.ppf(.1,shape,scale=scale),dop853Rise=marches[1][1]-marches[1][0]))
    sample_times=np.linspace(0,max(r['t90'] for r in rows)*1.18,150)
    samples=[dict(time=float(t),temperatures=(25+solutions[1].sol(t)).tolist()) for t in sample_times]
    return dict(**case,capacities=CAPACITIES,links=LINKS,K=k.tolist(),b=b.tolist(),x=x.tolist(),v=v.tolist(),w=w.tolist(),tau=tau,rows=rows,samples=samples)

if __name__=='__main__':
    parser=argparse.ArgumentParser();parser.add_argument('--output');args=parser.parse_args()
    data=[analyze(c) for c in CASES]
    # Combined forcing must superpose in unnormalized temperature, not in normalized timing.
    np.testing.assert_allclose(data[4]['x'],np.array(data[2]['x'])+data[3]['x'],rtol=1e-12)
    encoded=json.dumps(data,indent=2)
    if args.output:
        from pathlib import Path
        Path(args.output).write_text(encoded,encoding='utf-8')
    else:print(encoded)
