#!/usr/bin/env python3
"""verify_all.py — Electronic Supplementary Material.
Regenerates and asserts every reported quantity in the manuscript from a
single set of definitions. PART A: structural identities and closed forms
(hard assertions). PART B: every number printed in the manuscript,
recomputed and compared."""
import numpy as np
from scipy.integrate import quad
from scipy.optimize import brentq

SIG, MU, B, K0 = 4.0, 0.7, 15.0, 1.0
MUSTAR = (SIG-1)/SIG
rng = np.random.default_rng(42)
A_fail, ledger = [], []

def ok(name, cond):
    print(f"  [{'PASS' if cond else 'FAIL'}] {name}")
    if not cond: A_fail.append(name)
def led(item, printed, recomp, tol):
    v = abs(printed-recomp) <= tol
    ledger.append((item, printed, recomp, v))

# ---------- kernel ----------
def Khat(n, b=B, kb=K0):
    if n == 0: return 2*kb/b*(1-np.exp(-b/(2*kb)))
    return 2*(b/kb)*(1-(-1)**n*np.exp(-b/(2*kb)))/((b/kb)**2+4*np.pi**2*n**2)
def T0f(b=B, kb=K0): return Khat(0,b,kb)
def phif(n, b=B, kb=K0): return Khat(n,b,kb)/T0f(b,kb)

# ---------- linear machinery ----------
Om  = lambda p,s=SIG: p/(s+(s-1)*p)
Acal= lambda p,s=SIG: 1+(2*s-1)/(s-1)*Om(p,s)
Dl  = lambda p,m=MU,s=SIG: m*Acal(p,s)-1
alf = lambda m: m/(1-m)
nuf = lambda m,s=SIG: (alf(m)*(2*s-1)-(s-1)**2)/(s*(s-1))
Wc  = lambda m,s=SIG: m*(2*s-1)/(s-1)*(m*s-(s-1))/((1-m)*(s-1))

print("PART A — machinery assertions")
# closed-form Khat vs quadrature
for n in (0,1,2,5):
    q,_ = quad(lambda y,n=n: np.exp(-B*min(y,1-y))*np.cos(2*np.pi*n*y),0,1,limit=400)
    ok(f"Khat closed form, n={n}", abs(q-Khat(n))<1e-9)
# Delta two forms; nu identity; W lemma algebra
c1=c2=c3=c4=True
for _ in range(3000):
    s=1+3*rng.random(); p=rng.random(); m=0.98*rng.random()+0.01; a=m/(1-m)
    d=m*(1+(2*s-1)/(s-1)*Om(p,s))-1
    c1 &= abs(d-(m*(Om(p,s)*(1-p)+p/(s-1))-(1-m)))<1e-12
    c2 &= abs((1-nuf(m,s)*p)-(-d*(s+(s-1)*p)/(s*(1-m))))<1e-11
    winter=-m*(2*a*(2*s-1)/(s-1)+1)-m*s/(s-1)+(1-m)*a**2*(2*s-1)**2/(s-1)**2
    c3 &= abs(winter-Wc(m,s))<1e-9
    # App-H 3x3 system solve
    kk=0.5+rng.random(); eta=rng.random()-0.5; f=eta*p/kk
    Mx=np.array([[1,-(1-s)*p,-p],[p,s-p,-p],[a/(s-1),a,-1.]],float)
    bb=np.array([f,f,0.]); pp,ww,ll=np.linalg.solve(Mx,bb)
    c4 &= (abs(ww-f/(s*(1-nuf(m,s)*p)))<1e-9*max(1,abs(ww)) and abs(pp-s*ww)<1e-9
           and abs(ll-a*(2*s-1)/(s-1)*ww)<1e-9)
ok("Delta_n: composite == three-force form (3000 draws)",c1)
ok("identity 1-nu*phi == -Delta*(sig+(sig-1)phi)/(sig(1-mu))",c2)
ok("W(mu) closed form == lemma intermediate expression",c3)
ok("App-H 3x3 system: w,p=sigma*w,lambda=alpha(2s-1)/(s-1)w",c4)
ok("nu(mu*) == 1", abs(nuf(MUSTAR)-1)<1e-12)
ok("W(0)=W(mu*)=0; W<0 on (0,mu*)", abs(Wc(1e-14))<1e-9 and abs(Wc(MUSTAR))<1e-12
   and all(Wc(m)<0 for m in np.linspace(.01,MUSTAR-.01,50)))
ok("dW/dmu|mu* == K*sigma", abs((Wc(MUSTAR)-Wc(MUSTAR-1e-7))/1e-7-(2*SIG-1)/(SIG-1)*SIG)<1e-3)

# ---------- H(r), Jensen coefficients ----------
Hr = lambda r,b=B: 2/b**2*(np.exp(-b*min(r,1-r))-np.exp(-b/2))-2/b*(0.5-min(r,1-r))*np.exp(-b/2)
Cc = lambda b=B: (2-(2+b)*np.exp(-b/2))/b**2
def Hhat(n,b=B):
    q,_=quad(lambda r: Hr(r,b)*np.cos(2*np.pi*n*r),0,1,limit=800); return q
ok("arc convention: -2bC+b^2*Hhat0 == -b e^{-b/2}/2",
   abs(-2*B*Cc()+B**2*Hhat(0)-(-B*np.exp(-B/2)/2))<1e-9)

def Jn(n,b=B,kb=K0):
    p=phif(n,b,kb); t0=T0f(b,kb)
    return -2*b*Cc(b)+b**2*Hhat(n,b), 2*t0*p**2/(1-p)
def Lam(n,m,b=B,kb=K0,s=SIG):
    p=phif(n,b,kb); t0=T0f(b,kb); nu=nuf(m,s); jen,_=Jn(n,b,kb)
    return (Wc(m,s)*p**2/(s**2*kb**2*(1-nu*p)**2)
            + m*jen/((s-1)*t0*kb**2) + 2*m*nu*p**2/((s-1)*kb**2*(1-nu*p)))

# structural: Lam at mu* equals J_n/(sigma*T0*K0^2)
t0=T0f(); c=all(abs(Lam(n,MUSTAR)-sum(Jn(n))/(SIG*t0))<1e-10 for n in (1,2,3,5))
ok("Lambda_n(mu*) == J_n/(sigma T0 K0^2)  (thm right-endpoint)",c)
# poles & finiteness
grid=np.linspace(1e-3,MUSTAR-1e-9,40001)
ok("1-nu*phi_1 > 0 on (0,mu*]  (no pole; Lambda_1(mu*-) finite)",
   min(1-nuf(m)*phif(1) for m in grid)>0)
pole=brentq(lambda m:1-nuf(m)*phif(1),MUSTAR,0.9)
ok("Lambda_1 pole at mu = 1/A_1 = mu_1*", abs(pole-1/Acal(phif(1)))<1e-10)
# single crossing
def crossings(n):
    v=np.array([Lam(n,m) for m in grid]); return int(np.sum(np.sign(v[1:])!=np.sign(v[:-1]))), v
n1,v1=crossings(1); n2,v2=crossings(2); n3,v3=crossings(3); n4,v4=crossings(4)
ok("single crossing: Lambda_1 (1), Lambda_2 (1), Lambda_3 (0), Lambda_4 (0)",
   (n1,n2,n3,n4)==(1,1,0,0))
ok("left endpoint: Lambda_n<0 near mu->0+ for n=1..4",
   all(v[0]<0 for v in (v1,v2,v3,v4)))
mss1=brentq(lambda m:Lam(1,m),0.1,MUSTAR-1e-9)
mss2=brentq(lambda m:Lam(2,m),0.1,MUSTAR-1e-9)

print("\nPART B — ledger: printed in manuscript vs recomputed")
led("phi_1 (manuscript prints 0.8517)",0.8517,phif(1),1e-3)
led("Omega_1 wage elasticity (manuscript prints 0.1299)",0.1299,Om(phif(1)),1e-3)
led("A_1 composite (L927 prints as 'Omega_1'=1.302)",1.302,Acal(phif(1)),1.5e-3)
led("Delta_1 (L927: -0.088)",-0.088,Dl(phif(1)),1e-3)
led("referee: mu*Om-(1-mu) with Om=1.302",0.6114,MU*1.302-(1-MU),1e-4)
led("free-trade limit Omega_n -> 1/(2 sigma - 1) (Sec 3.3)",1/(2*SIG-1),Om(1.0),1e-9)
led("Vbar symmetric welfare (manuscript prints 0.6248)",0.6248,t0**(MU/(SIG-1)),1e-3)
led("nu(0.7) (implied by Sec.4)",0.6111,nuf(MU),1e-4)
for n,pj,pc,pJ in [(1,-0.0401,1.303,1.263),(2,-0.1088,0.223,0.114),(3,-0.1621,0.0656,-0.0965),
                   (4,-0.1954,0.0249,-0.170),(5,-0.2159,0.0113,-0.205),(6,-0.2290,0.0058,-0.223),
                   (8,-0.2436,0.0019,-0.242),(10,-0.2510,0.0008,-0.250),(20,-0.2617,5.3e-5,-0.262)]:
    jen,cpl=Jn(n)
    led(f"tab jn_right n={n}: Jensen",pj,jen,7e-4)
    led(f"tab jn_right n={n}: Coupling",pc,cpl,7e-4)
    led(f"tab jn_right n={n}: J_n",pJ,jen+cpl,1.2e-3)
led("Lambda_1(0.7) (+0.290)",0.290,Lam(1,MU),1.5e-3)
led("Lambda_2(0.7) (-0.056)",-0.056,Lam(2,MU),1.5e-3)
led("Lambda_3(0.7) (-0.234)",-0.234,Lam(3,MU),1.5e-3)
led("mu**_1 (0.6377)",0.6377,mss1,2e-4)
led("mu**_2 (0.7176)",0.7176,mss2,2e-4)
led("band width mu*-mu**_1 (0.1123)",0.1123,MUSTAR-mss1,2e-4)
led("band %% of mu* (14.97)",14.97,100*(MUSTAR-mss1)/MUSTAR,0.03)
bstar=brentq(lambda b: sum(Jn(1,b)),3.5,7.5)
led("b* (5.24)",5.24,bstar,0.02)
led("phi_1(b*) (0.47)",0.47,phif(1,bstar),0.005)
m200=brentq(lambda m:Lam(1,m,200.0),0.55,MUSTAR-1e-9)
led("mu**_1 at b=200 -> 3/5 claim",0.600,m200,0.004)



# exact-rational verification of the single-crossing polynomial identity
from fractions import Fraction as _F
import random as _rnd; _rnd.seed(11)
def _nuF(a,s): return (a*(2*s-1)-(s-1)**2)/(s*(s-1))
_okid=True
for _ in range(8):
    _s=_F(_rnd.randint(2,6)); _p=_F(_rnd.randint(1,9),10); _J=_F(_rnd.randint(-9,-1),10); _T=_F(_rnd.randint(1,9),10); _a=_F(_rnd.randint(1,40),10)
    _mu=_a/(1+_a); _nu=_nuF(_a,_s)
    _W=_mu*((2*_s-1)/(_s-1))*(_mu*_s-(_s-1))/((1-_mu)*(_s-1))
    _L=_W*_p**2/(_s**2*(1-_nu*_p)**2)+_mu*_J/((_s-1)*_T)+2*_mu*_nu*_p**2/((_s-1)*(1-_nu*_p))
    _Q=((2*_s-1)*_p**2/((_s-1)**2*_s**2))*(_a-(_s-1))+(_J/((_s-1)*_T))*(1-_nu*_p)**2+(2*_p**2/(_s-1))*_nu*(1-_nu*_p)
    _okid &= (_L*(1+_a)*(1-_nu*_p)**2 == _a*_Q)
ok("single-crossing identity L*(1+a)(1-nu*phi)^2 == a*Q (exact rationals)", _okid)
# closed-form mu**_n from the single-crossing proposition (quadratic in alpha)
def mu2_closed(n):
    p=phif(n); jen=-2*B*Cc()+B**2*Hhat(n); s=SIG
    A2=(2*s-1)*p**2/((s-1)**2*s**2); B2=jen/((s-1)*t0); C2=2*p**2/(s-1)
    e=p*(2*s-1)/(s*(s-1)); g=(s+(s-1)*p)/s
    c2=B2*e*e - C2*e*(2*s-1)/(s*(s-1))
    c1=A2 - 2*B2*g*e + C2*((2*s-1)*g+e*(s-1)**2)/(s*(s-1))
    c0=-A2*(s-1)+B2*g*g - C2*g*(s-1)/s
    r=[x.real for x in np.roots([c2,c1,c0]) if abs(x.imag)<1e-12 and 0<x.real<s-1]
    assert len(r)==1
    return r[0]/(1+r[0])
ok("closed-form mu**_1 == brentq root", abs(mu2_closed(1)-mss1)<1e-9)
ok("closed-form mu**_2 == brentq root", abs(mu2_closed(2)-mss2)<1e-9)

w=max(len(x[0]) for x in ledger)
for it,pv,rv,v in ledger:
    print(f"  [{'OK ' if v else 'DIFF'}] {it:<{w}}  printed={pv:<10.4g} recomputed={rv:.5g}")
print(f"\nPART A: {'ALL PASS' if not A_fail else 'FAILURES: '+str(A_fail)}")
print(f"PART B: {sum(1 for *_,v in ledger if v)}/{len(ledger)} match.")
