from pathlib import Path
import re, math, random

HERE=Path(__file__).resolve().parent
PAPER=HERE/"CSM_RH_Paper_54_Seeded_MLEPG_Amplification_and_Critical_Locking_Barrier_v0.1_2026-09-08.md"

def Phi(k,a,d):
    return min(a,d,2-a*(2-k))

def phi_star(k):
    return 2/(3-k)

def check_amplification_gate():
    random.seed(54)
    for _ in range(10000):
        k=random.uniform(0.01,0.99)
        a=random.uniform(0.01,0.999)
        d=random.uniform(0.01,1.5)
        amp=Phi(k,a,d)>k+1e-12
        gate=(a>k and d>k)
        assert amp==gate, (k,a,d,Phi(k,a,d),amp,gate)
    return True

def check_optimal_map():
    worst=0.0
    for k in [0.05,0.2,0.4,0.6,0.8,0.95]:
        target=phi_star(k)
        vals=[]
        for j in range(1,20000):
            a=j/20000
            vals.append(min(a,2-a*(2-k)))
        num=max(vals)
        worst=max(worst,abs(num-target))
        assert abs(num-target)<1e-4, (k,num,target)
        assert target>k and target<1
    return worst

def check_iteration_formula():
    worst=0.0
    for k0 in [0.05,0.2,0.5,0.8]:
        k=k0
        for j in range(1,12):
            k=phi_star(k)
            e_formula=1/(2**j*(1+1/(1-k0))-1)
            worst=max(worst,abs((1-k)-e_formula))
            assert abs((1-k)-e_formula)<1e-12
    return worst

def check_power_model():
    # Compare finite sums with the predicted exponents; only a numerical sanity check.
    k=0.4
    beta=1-k/2
    alpha=0.7
    ratios=[]
    for N in [2000,4000,8000,16000]:
        H=max(1,int(N**alpha))
        S=sum(((n+H)**beta-n**beta)**2 for n in range(N,2*N))
        scale=N*(H**2)*(N**(-k))
        ratios.append(S/scale)
    assert min(ratios)>0.01 and max(ratios)<10
    return ratios

def check_source():
    s=PAPER.read_text(encoding="utf-8")
    forbidden=[
        r"(?<!\\)\\\(",
        r"(?<!\\)\\\)",
        r"(?<!\\)\\\[",
        r"(?<!\\)\\\]",
    ]
    for pat in forbidden:
        assert re.search(pat,s) is None, pat
    assert s.count("$$")%2==0
    tmp=re.sub(r"\$\$.*?\$\$","",s,flags=re.S)
    assert len(re.findall(r"(?<!\\)\$",tmp))%2==0
    return True

if __name__=="__main__":
    print("amplification_gate",check_amplification_gate())
    print("optimal_map_grid_error",check_optimal_map())
    print("iteration_formula_error",check_iteration_formula())
    print("power_model_ratios",check_power_model())
    print("source_delimiters",check_source())
    print("PASS")
