#!/usr/bin/env python3
"""Exact scalar circle-spectrum certificates; standard library only.
Full Laurent factors and finite moment prefixes are distinct inputs.
Supplied factors and roots here are Gaussian-rational; general algebraic root
solving is outside this checker. Automatic recovery covers linear/quadratic kernels.
Use --output to isolate review results.
"""
from dataclasses import dataclass
from fractions import Fraction as F
from itertools import product, combinations
from math import isqrt
from collections import Counter
from pathlib import Path
import argparse,json
CHECKS=Counter()
def require(c,msg):
    if not c: raise ValueError(msg)
def check(c,label):
    CHECKS[label]+=1
    if not c: raise RuntimeError(label)
def rational(x):
    require(type(x) is int or isinstance(x,F),'exact rational input required')
    return F(x)
@dataclass(frozen=True)
class C:
    re:F=F(0)
    im:F=F(0)
    def __post_init__(self):
        object.__setattr__(self,'re',rational(self.re));object.__setattr__(self,'im',rational(self.im))
    def __add__(self,b):
        b=Cv(b);return C(self.re+b.re,self.im+b.im)
    __radd__=__add__
    def __neg__(self):return C(-self.re,-self.im)
    def __sub__(self,b):return self+-Cv(b)
    def __rsub__(self,b):return Cv(b)+-self
    def __mul__(self,b):
        b=Cv(b);return C(self.re*b.re-self.im*b.im,self.re*b.im+self.im*b.re)
    __rmul__=__mul__
    def conjugate(self):return C(self.re,-self.im)
    def norm2(self):return self.re*self.re+self.im*self.im
    def __truediv__(self,b):
        b=Cv(b);q=b.norm2();require(q!=0,'zero denominator');v=self*b.conjugate();return C(v.re/q,v.im/q)
    def __rtruediv__(self,b):return Cv(b)/self
    def __bool__(self):return bool(self.re or self.im)
    def __eq__(self,b):
        try:b=Cv(b)
        except ValueError:return False
        return self.re==b.re and self.im==b.im
    def __hash__(self):return hash((self.re,self.im))
def Cv(x):return x if isinstance(x,C) else C(x)
Z=C();O=C(1);I=C(0,1)
def trim(p):
    p=list(map(Cv,p));require(bool(p),'empty polynomial')
    while len(p)>1 and not p[-1]:p.pop()
    return p
def padd(p,q):return trim([(p[j] if j<len(p) else Z)+(q[j] if j<len(q) else Z) for j in range(max(len(p),len(q)))])
def pscale(p,c):return trim([Cv(c)*x for x in p])
def pmul(p,q):
    r=[Z]*(len(p)+len(q)-1)
    for j,a in enumerate(p):
        for k,b in enumerate(q):r[j+k]+=a*b
    return trim(r)
def pvalue(p,z):
    v=Z
    for c in reversed(p):v=v*z+c
    return v
def pdiv(p,q):
    p=trim(p);q=trim(q);require(any(q),'zero divisor');r=[Z]*max(1,len(p)-len(q)+1)
    while any(p) and len(p)>=len(q):
        j=len(p)-len(q);c=p[-1]/q[-1];r[j]+=c;p=padd(p,pscale([Z]*j+q,-c))
    return trim(r),p
def pgcd(p,q):
    while any(q):p,q=q,pdiv(p,q)[1]
    return pscale(p,1/p[-1])
def eye(n):return [[O if i==j else Z for j in range(n)] for i in range(n)]
def star(A):return [[x.conjugate() for x in row] for row in zip(*A)]
def mm(A,B):return [[sum((a*b for a,b in zip(row,col)),Z) for col in zip(*B)] for row in A]
def msub(A,B):return [[a-b for a,b in zip(x,y)] for x,y in zip(A,B)]
def mscale(A,c):return [[c*x for x in row] for row in A]
def quad(A,v):return sum((v[i].conjugate()*A[i][j]*v[j] for i in range(len(v)) for j in range(len(v))),Z)
def psd(A):
    n=len(A);require(all(len(row)==n for row in A) and A==star(A),'Hermitian square matrix required')
    if not n:return dict(psd=True,rank=0,pivots=[])
    for i in range(n):
        if A[i][i].re<0:
            v=[Z]*n;v[i]=O;return dict(psd=False,witness=v,value=A[i][i])
        if not A[i][i]:
            for j in range(n):
                if A[i][j]:
                    v=[Z]*n;v[j]=O;v[i]=-A[i][j]*((A[j][j].re+1)/(2*A[i][j].norm2()))
                    return dict(psd=False,witness=v,value=quad(A,v))
    piv=next((i for i in range(n) if A[i][i]),None)
    if piv is None:return dict(psd=True,rank=0,pivots=[])
    rem=[j for j in range(n) if j!=piv];d=A[piv][piv]
    S=[[A[i][j]-A[i][piv]*A[piv][j]/d for j in rem] for i in rem];r=psd(S)
    if r['psd']:return dict(psd=True,rank=1+r['rank'],pivots=[d]+r['pivots'])
    v=[Z]*n
    for j,x in zip(rem,r['witness']):v[j]=x
    v[piv]=-sum((A[piv][j]*v[j] for j in rem),Z)/d
    return dict(psd=False,witness=v,value=quad(A,v))
def power(z,k):
    z=Cv(z);require(type(k) is int,'integer power')
    if k<0:return O/power(z,-k)
    out=O
    while k:
        if k&1:out*=z
        z*=z;k//=2
    return out

def solve(A,b):
    A=[[Cv(x) for x in row] for row in A];b=list(map(Cv,b));n=len(A)
    require(n>0 and len(b)==n and all(len(row)==n for row in A),'square system')
    B=[row+[v] for row,v in zip(A,b)]
    for j in range(n):
        k=next((i for i in range(j,n) if B[i][j]),None)
        require(k is not None,'singular system')
        B[j],B[k]=B[k],B[j];a=B[j][j];B[j]=[x/a for x in B[j]]
        for i in range(n):
            if i!=j:
                a=B[i][j];B[i]=[x-a*y for x,y in zip(B[i],B[j])]
    return [row[-1] for row in B]

def spectral_coefficients(p):
    p=trim(p)
    return trim([sum((p[j+k]*p[j].conjugate() for j in range(len(p)-k)),Z) for k in range(len(p))])
def laurent_value(q,z):
    z=Cv(z);require(z.norm2()==1,'unit circle evaluation')
    return Cv(q[0])+sum((Cv(q[k])*power(z,k)+Cv(q[k]).conjugate()*power(z,-k) for k in range(1,len(q))),Z)
def roots_polynomial(roots):
    p=[O]
    for z in roots:p=pmul(p,[-Cv(z),O])
    return p

def factor_certificate(q,p,roots=None):
    q=trim(q);p=trim(p);require(q[0].im==0,'real constant Laurent coefficient')
    require(spectral_coefficients(p)==q,'all Laurent coefficients must match')
    if not any(p):return dict(nonnegative=True,zero=True,canonical=False,strict=False)
    result=dict(nonnegative=True,zero=False,canonical=None,strict=None)
    if roots is not None:
        roots=list(map(Cv,roots));require(len(roots)==len(p)-1,'complete root multiplicities')
        require(pscale(roots_polynomial(roots),p[-1])==p,'root factorization identity')
        result.update(canonical=all(z.norm2()>=1 for z in roots) and p[0].im==0 and p[0].re>0,strict=all(z.norm2()!=1 for z in roots))
    return result

def moment_prefix(c):
    c=list(map(Cv,c));require(bool(c) and c[0].im==0,'nonempty moments with real mass');return c
def toeplitz(c):
    c=moment_prefix(c);return [[c[i-j] if i>=j else c[j-i].conjugate() for j in range(len(c))] for i in range(len(c))]
def atom_moments(nodes,weights,n):
    nodes=list(map(Cv,nodes));weights=list(map(rational,weights))
    require(type(n) is int and n>=0 and len(nodes)==len(weights),'atom dimensions/order')
    require(all(z.norm2()==1 for z in nodes) and all(w>0 for w in weights),'unit nodes and strictly positive masses')
    require(len(set(nodes))==len(nodes),'combine duplicate nodes first')
    return [sum((w*power(z,-k) for z,w in zip(nodes,weights)),Z) for k in range(n+1)]
def verify_atoms(c,nodes,weights):
    c=moment_prefix(c);require(atom_moments(nodes,weights,len(c)-1)==c,'all known moments must match');return True

def monic_kernel(c):
    c=moment_prefix(c);state=psd(toeplitz(c));require(state['psd'],'positive moment matrix required')
    if state['rank']==0:return [O]
    require(state['rank']<len(c),'singular matrix required')
    r=state['rank'];T=toeplitz(c[:r+1]);a=solve([row[:r] for row in T[:r]],[-row[r] for row in T[:r]])+[O]
    require(all(sum(T[i][j]*a[j] for j in range(r+1))==0 for i in range(r+1)),'exact kernel identity')
    return a

def sqrt_fraction(x):
    x=rational(x);require(x>=0,'nonnegative rational square root')
    a,b=isqrt(x.numerator),isqrt(x.denominator);require(a*a==x.numerator and b*b==x.denominator,'square root leaves rationals');return F(a,b)
def sqrt_gaussian(z):
    z=Cv(z);r=sqrt_fraction(z.norm2());u=sqrt_fraction((r+z.re)/2);v=sqrt_fraction((r-z.re)/2)
    if z.im<0:v=-v
    out=C(u,v);require(out*out==z,'Gaussian square root');return out
def small_roots(p):
    p=trim(p)
    if len(p)==2:return [-p[0]/p[1]]
    require(len(p)==3,'supply certified roots for degree above two')
    d=sqrt_gaussian(p[1]*p[1]-4*p[0]*p[2]);return [(-p[1]+d)/(2*p[2]),(-p[1]-d)/(2*p[2])]
def recover_singular(c,roots=None):
    c=moment_prefix(c);p=monic_kernel(c)
    if len(p)==1:verify_atoms(c,[],[]);return dict(nodes=[],weights=[],kernel=p)
    roots=small_roots(p) if roots is None else list(map(Cv,roots))
    require(roots_polynomial(roots)==p,'complete annihilator roots')
    require(len(set(roots))==len(roots) and all(z.norm2()==1 for z in roots),'distinct unit annihilator roots')
    r=len(roots);w=solve([[power(z,-k) for z in roots] for k in range(r)],c[:r]);require(all(x.im==0 and x.re>0 for x in w),'positive real recovered masses');w=[x.re for x in w]
    verify_atoms(c,roots,w);return dict(nodes=roots,weights=w,kernel=p)

def maximal_atom(c,z):
    c=moment_prefix(c);z=Cv(z);require(z.norm2()==1,'prescribed unit node');T=toeplitz(c);state=psd(T);require(state['psd'] and state['rank']==len(c),'positive definite prefix required')
    v=[power(z,-k) for k in range(len(c))];u=solve(T,v);d=sum((a.conjugate()*b for a,b in zip(v,u)),Z);require(d.im==0 and d.re>0,'positive evaluation norm');w=1/d.re
    residual=[ck-w*power(z,-k) for k,ck in enumerate(c)];s=psd(toeplitz(residual));require(s['psd'] and s['rank']==len(c)-1,'rank-one boundary certificate')
    return dict(node=z,weight=w,residual_moments=residual,kernel=monic_kernel(residual))

def rejected(fn):
    try:fn()
    except ValueError:return True
    return False

def run():
    CHECKS.clear()
    p=list(map(Cv,[1,F(-3,2),F(1,2)]));q=spectral_coefficients(p)
    check(q==list(map(Cv,[F(7,2),F(-9,4),F(1,2)])),'terminal boundary Laurent coefficients')
    check(factor_certificate(q,p,[1,2])==dict(nonnegative=True,zero=False,canonical=True,strict=False),'terminal boundary canonical factor')
    reflected=list(map(Cv,[F(-1,2),F(3,2),-1]));check(factor_certificate(q,reflected,[1,F(1,2)])['canonical'] is False,'root reflection preserves spectrum but not canonicality')
    strict=pmul([O,C(F(-1,2))],[O,I/3]);sq=spectral_coefficients(strict)
    check(sq==[C(F(25,18)),C(F(-5,9),F(5,12)),-I/6],'terminal complex factor coefficients')
    check(factor_certificate(sq,strict,[2,3*I])['strict'],'strict factor roots outside circle')
    check([laurent_value(sq,z) for z in [O,-O,I]]==[C(F(5,18)),C(F(5,2)),C(F(5,9))],'terminal exact spectral readings')
    c=[O,C(F(3,4))];check(psd(toeplitz(c))['rank']==2,'legal finite prefix')
    wrong=c+[Z];v=list(map(Cv,[1,F(-3,2),1]));check(quad(toeplitz(wrong),v)==F(-1,4) and laurent_value(c,-O)==F(-1,2),'zero-extension and full-polynomial rejection')
    complex_c=[O,C(F(2,3),F(-1,3)),C(F(1,3))];rec=recover_singular(complex_c)
    check(dict(zip(rec['nodes'],rec['weights']))=={O:F(2,3),I:F(1,3)} and rec['kernel']==[I,-O-I,O],'terminal unique complex moment recovery')
    check(atom_moments(rec['nodes'],rec['weights'],3)[3]==C(F(2,3),F(1,3)),'forced next moment sign')
    for eps in [F(1,1000),F(1,10),F(2)]:
        perturbed=complex_c[:];perturbed[2]+=I*eps;check(quad(toeplitz(perturbed),[I,-O-I,O])==-2*eps,'actual outward perturbation negative direction')
    extremes=[]
    for z,expected in [(O,F(7,8)),(I,F(7,32))]:
        e=maximal_atom(c,z);rest=recover_singular(e['residual_moments']);check(e['weight']==expected,'terminal maximal prescribed atom');verify_atoms(c,[z]+rest['nodes'],[expected]+rest['weights']);extremes.append(dict(**e,remainder=rest))
    check(extremes[1]['remainder']['nodes']==[C(F(24,25),F(-7,25))],'terminal second extremal node')
    # General exact root and coefficient identities. Sampling is diagnostic only.
    circle=sorted({C((1-t*t)/(1+t*t),2*t/(1+t*t)) for t in [F(k,5) for k in range(-8,9)]}|{-O},key=lambda z:(z.re,z.im))
    roots_pool=[Z,O,-O,I,-I,C(F(1,2)),C(2),3*I,C(F(3,5),F(4,5))]
    for roots in product(roots_pool,repeat=2):
        a=pscale(roots_polynomial(roots),C(F(2,3),F(1,3)));s=spectral_coefficients(a);cert=factor_certificate(s,a,roots)
        check(cert['nonnegative'],'all root families yield verified nonnegative spectra')
        for z in circle[::3]:check(laurent_value(s,z)==pvalue(a,z).norm2(),'coefficient identity at rational circle diagnostics')
        check(s[0].re==sum(x.norm2() for x in a),'exact Parseval constant coefficient')
        # Even with circle zeros, nonzero density gives every finite block positive definite.
        for n in [1,2,4]:
            padded=s+[Z]*max(0,n+1-len(s));state=psd(toeplitz(padded[:n+1]));check(state['psd'] and state['rank']==n+1,'nonzero polynomial square has all finite Toeplitz blocks PD')
    # Independent data generation from exact unit atoms, covering ranks and zeros.
    atom_pool=[O,-O,I,-I,C(F(3,5),F(4,5))]
    for r in range(1,5):
        for nodes in combinations(atom_pool,r):
            for raw in product([1,2],repeat=r):
                w=[F(x,sum(raw)) for x in raw]
                for n in [0,r-1,r,r+1]:
                    data=atom_moments(nodes,w,n);state=psd(toeplitz(data));check(state['psd'] and state['rank']==min(r,n+1),'generated atoms exact Toeplitz rank')
                    check(verify_atoms(data,nodes,w),'full moment witness verified')
                    if n>=r:
                        out=recover_singular(data,nodes);check(out['weights']==w,'singular reconstruction and unique weights')
                        check(all(pvalue(out['kernel'],z)==0 for z in nodes),'all recovered spectral lines annihilated')
                        for phase in [I,C(F(3,5),F(4,5))]:
                            shifted=[phase*z for z in nodes];rot=[power(phase,-k)*ck for k,ck in enumerate(data)]
                            check(verify_atoms(rot,shifted,w),'rotation with fixed negative exponent moment convention')
    # Exact first-lag completion disk, including singular boundary and infeasible exterior.
    for c0 in [F(1),F(3,2),F(2)]:
        for x,y in product([F(-1,2),F(0),F(1,3)],repeat=2):
            c1=C(x,y)
            if c1.norm2()>=c0*c0:continue
            radius=c0-c1.norm2()/c0;center=c1*c1/c0
            for h in [Z,O,-O,I,C(F(3,5),F(4,5)),C(F(1,2),F(1,3)),C(F(6,5)),C(1,1)]:
                data=[C(c0),c1,center+radius*h];state=psd(toeplitz(data));check(state['psd']==(h.norm2()<=1),'exact full complex extension disk')
                if state['psd']:check(state['rank']==(2 if h.norm2()==1 else 3),'extension boundary rank')
                else:check(quad(toeplitz(data),state['witness'])==state['value'] and state['value'].re<0,'negative PSD witness reproduced')
            for z in circle[::4]:
                e=maximal_atom([C(c0),c1],z);rest=recover_singular(e['residual_moments']);check(verify_atoms([C(c0),c1],[z]+rest['nodes'],[e['weight']]+rest['weights']),'all rational prescribed nodes attain bound')
                larger=e['weight']+F(1,100);bad=[C(c0)-larger,c1-larger/z];state=psd(toeplitz(bad));check(not state['psd'],'increased prescribed mass rejected')
    for n in range(5):check(recover_singular([Z]*(n+1))['weights']==[],'zero moment table has zero measure')
    check(factor_certificate([0],[0])['zero'],'zero Laurent polynomial separated')
    for fn in [lambda:factor_certificate([1,F(3,4)],[1,1]),lambda:factor_certificate(q,p,[1,3]),lambda:atom_moments([0],[1],1),lambda:atom_moments([1,1],[F(1,2),F(1,2)],1),lambda:atom_moments([1],[-1],1),lambda:toeplitz([]),lambda:toeplitz([I]),lambda:recover_singular([1,0]),lambda:maximal_atom([1,1],1),lambda:C(.5)]:
        check(rejected(fn),'invalid or false certificate rejected')
    return dict(status='PASS',checks=sum(CHECKS.values()),groups=dict(CHECKS),boundary_factor=p,boundary_coefficients=q,strict_factor=strict,strict_coefficients=sq,unique_moments=complex_c,unique_recovery=rec,prescribed_atom_extremizers=extremes,scope='Exact supplied-factor and finite-moment certificates; general algebraic root isolation, noisy rank decisions and entropy selection are outside this checker')
def serial(x):
    if isinstance(x,C):return {'re':str(x.re),'im':str(x.im)}
    if isinstance(x,F):return str(x)
    if isinstance(x,(list,tuple)):return [serial(v) for v in x]
    if isinstance(x,dict):return {k:serial(v) for k,v in x.items()}
    return x
if __name__=='__main__':
    ap=argparse.ArgumentParser(description=__doc__);ap.add_argument('--output',type=Path);args=ap.parse_args();data=json.dumps(serial(run()),ensure_ascii=False,indent=2)+'\n'
    if args.output:args.output.parent.mkdir(parents=True,exist_ok=True);args.output.write_text(data);print('PASS',sum(CHECKS.values()),'checks;',args.output)
    else:print(data,end='')
