#!/usr/bin/env python3
"""
Finite exact crosscheck for CSM_RH Paper 08.

Verifies:
1. the pinned Vaughan coefficient identity for finite n;
2. the long CSSA block recombination.

This is not evidence for RH.
"""
import math
import csv
import numpy as np

def mobius_sieve(n):
    mu = np.ones(n+1, dtype=int)
    prime = np.ones(n+1, dtype=bool)
    prime[:2] = False
    mu[:] = 1
    is_prime = np.ones(n+1, dtype=bool)
    is_prime[:2] = False
    primes=[]
    lp=np.zeros(n+1,dtype=int)
    mu=np.zeros(n+1,dtype=int)
    mu[1]=1
    for i in range(2,n+1):
        if lp[i]==0:
            lp[i]=i
            primes.append(i)
            mu[i]=-1
        for p in primes:
            if p>lp[i] or i*p>n:
                break
            lp[i*p]=p
            if p==lp[i]:
                mu[i*p]=0
            else:
                mu[i*p]=-mu[i]
    return mu

def von_mangoldt(n):
    lam=np.zeros(n+1,float)
    for p in range(2,n+1):
        isprime=True
        for d in range(2,int(p**0.5)+1):
            if p%d==0:
                isprime=False
                break
        if not isprime:
            continue
        q=p
        lp=math.log(p)
        while q<=n:
            lam[q]=lp
            if q>n//p:
                break
            q*=p
    return lam

def conv(a,b,nmax):
    c=np.zeros(nmax+1,float)
    for d in range(1,nmax+1):
        if a[d]==0:
            continue
        for m in range(1,nmax//d+1):
            if b[m]!=0:
                c[d*m]+=a[d]*b[m]
    return c

def wN(N,n):
    if 1<=n<=N:
        return float(N)
    if N<n<2*N:
        return float(2*N-n)
    return 0.0

def run(N):
    X=2*N
    nmax=X-1
    U=V=max(2,int(X**(1/3)))
    mu=mobius_sieve(nmax)
    lam=von_mangoldt(nmax)
    L=np.zeros(nmax+1,float)
    one=np.zeros(nmax+1,float)
    for n in range(1,nmax+1):
        L[n]=math.log(n)
        one[n]=1.0

    mu_lo=np.zeros_like(L); mu_hi=np.zeros_like(L)
    la_lo=np.zeros_like(L); la_hi=np.zeros_like(L)
    for n in range(1,nmax+1):
        if n<=U: mu_lo[n]=mu[n]
        else: mu_hi[n]=mu[n]
        if n<=V: la_lo[n]=lam[n]
        else: la_hi[n]=lam[n]

    I1=conv(mu_lo,L,nmax)
    tmp=conv(mu_lo,la_lo,nmax)
    I2=conv(tmp,one,nmax)
    tmp2=conv(mu_hi,la_hi,nmax)
    II=conv(tmp2,one,nmax)
    rhs=la_lo+I1-I2+II
    coeff_err=float(np.max(np.abs(rhs[1:]-lam[1:])))

    a=lam-1.0
    a[0]=0.0
    A=np.cumsum(a)

    short=sum(wN(N,n)*a[n]*A[n-1] for n in range(2,min(V,nmax)+1))
    p_long=sum(wN(N,n)*a[n]*A[n-1] for n in range(V+1,X))

    t1=sum(wN(N,n)*A[n-1]*I1[n] for n in range(V+1,X))
    t2=-sum(wN(N,n)*A[n-1]*I2[n] for n in range(V+1,X))
    t3=sum(wN(N,n)*A[n-1]*II[n] for n in range(V+1,X))
    tc=-sum(wN(N,n)*A[n-1] for n in range(V+1,X))
    recomb=t1+t2+t3+tc

    unit_atom=sum(math.log(m)*wN(N,m)*A[m-1] for m in range(V+1,X))

    return [
        N,U,V,coeff_err,p_long,recomb,abs(p_long-recomb),short,
        t1,t2,t3,tc,unit_atom
    ]

rows=[run(N) for N in [40,80,160]]
with open("campaign07_vaughan_exact_crosscheck.csv","w",encoding="utf-8",newline="") as f:
    w=csv.writer(f)
    w.writerow([
        "N","U","V","max_coefficient_identity_error",
        "P_long_direct","P_long_Vaughan","long_recomb_abs_error",
        "P_short","T_I1","T_I2","T_II","T_C","T_I1_unit_atom"
    ])
    w.writerows(rows)

print("PASS")
print("max_coefficient_identity_error",max(r[3] for r in rows))
print("max_long_recomb_abs_error",max(r[6] for r in rows))
