#!/usr/bin/env python3
"""Exact filtered-chain cancellation and persistence certificates.

Standard library. Coefficients are rational numbers or a specified prime field.
The program verifies its complete finite examples; it is not a data-cloud or
Ripser implementation. Explicit checks remain active under python -O.
"""
from fractions import Fraction
from itertools import combinations, permutations
from pathlib import Path
import argparse, json, random
CHECKS=0

def need(b,m):
    if not b: raise ValueError(m)
def check(b,m):
    global CHECKS
    CHECKS+=1
    if not b: raise RuntimeError(m)

class Field:
    def __init__(self,p=0):
        need(isinstance(p,int) and not isinstance(p,bool) and (p==0 or p>=2 and all(p%d for d in range(2,int(p**.5)+1))),'zero or prime characteristic')
        self.p=p
    def value(self,x):
        need(isinstance(x,(int,Fraction)) and not isinstance(x,bool),'exact coefficient')
        if not self.p:return Fraction(x)
        x=Fraction(x);need(x.denominator%self.p!=0,'denominator vanishes')
        return x.numerator*pow(x.denominator,-1,self.p)%self.p
    def inv(self,x):
        x=self.value(x);need(x!=0,'zero pivot')
        return 1/x if not self.p else pow(x,-1,self.p)
    def norm(self,A):return [[self.value(x) for x in row] for row in A]
    def mm(self,A,B,cols=None):
        n=len(B[0]) if B else (0 if cols is None else cols)
        need(all(len(row)==len(B) for row in A),'matrix product dimensions')
        return [[self.value(sum(A[i][k]*B[k][j] for k in range(len(B)))) for j in range(n)] for i in range(len(A))]
    def add(self,A,B,sign=1):return [[self.value(x+sign*y) for x,y in zip(a,b)] for a,b in zip(A,B)]

def identity(n):return [[int(i==j) for j in range(n)] for i in range(n)]
def zero(n,m):return [[0]*m for _ in range(n)]
def col(A,j):return [row[j] for row in A]
def low(v):return next((i for i in range(len(v)-1,-1,-1) if v[i]),None)
def rref(F,A,ncols=None):
    B=F.norm(A);n=len(B[0]) if B else ncols or 0;r=0;piv=[]
    for j in range(n):
        k=next((i for i in range(r,len(B)) if B[i][j]),None)
        if k is None:continue
        B[r],B[k]=B[k],B[r];u=F.inv(B[r][j]);B[r]=[F.value(x*u) for x in B[r]]
        for i in range(len(B)):
            if i!=r:
                t=B[i][j];B[i]=[F.value(x-t*y) for x,y in zip(B[i],B[r])]
        piv.append(j);r+=1
        if r==len(B):break
    return B,piv

def rank(F,A):return len(rref(F,A)[1])
def kernel(F,A,n):
    B,piv=rref(F,A,n);out=[]
    for j in range(n):
        if j in piv:continue
        v=[0]*n;v[j]=1
        for i,k in enumerate(piv):v[k]=F.value(-B[i][j])
        out.append(v)
    return out

def validate(F,D,degrees,grades,ordered=False):
    n=len(D);need(len(degrees)==len(grades)==n and all(len(r)==n for r in D),'square chain metadata')
    need(all(isinstance(k,int) and not isinstance(k,bool) for k in degrees),'integer chain degrees')
    need(all(isinstance(g,(int,Fraction)) and not isinstance(g,bool) for g in grades),'exact grades')
    D=F.norm(D);need(F.mm(D,D)==zero(n,n),'boundary square must vanish')
    for i in range(n):
        for j in range(n):
            if D[i][j]:
                need(degrees[i]==degrees[j]-1 and grades[i]<=grades[j],'homogeneous filtration boundary')
                if ordered:need(i<j,'boundary must precede column')
    if ordered:need(all(grades[i]<=grades[i+1] for i in range(n-1)),'nondecreasing grades')
    return D

def reduce_boundary(F,D,degrees,grades):
    D=validate(F,D,degrees,grades,True);n=len(D);R=[r[:] for r in D];V=F.norm(identity(n));owner={}
    for j in range(n):
        while (i:=low(col(R,j))) is not None and i in owner:
            k=owner[i];a=F.value(R[i][j]*F.inv(R[i][k]))
            for row in range(n):
                R[row][j]=F.value(R[row][j]-a*R[row][k]);V[row][j]=F.value(V[row][j]-a*V[row][k])
        i=low(col(R,j))
        if i is not None:owner[i]=j
    W=[r[:] for r in V];B=zero(n,n)
    for i,j in owner.items():
        for k in range(n):W[k][i]=R[k][j]
        B[i][j]=1
    pairs=sorted(owner.items());essential=[i for i in range(n) if low(col(R,i)) is None and i not in owner]
    bars=sorted([(degrees[i],grades[i],grades[j]) for i,j in pairs if grades[i]<grades[j]]+[(degrees[i],grades[i],None) for i in essential],key=lambda z:(z[0],z[1],float('inf') if z[2] is None else z[2]))
    return {'R':R,'V':V,'W':W,'B':B,'pairs':pairs,'essential':essential,'bars':bars}

def certify(F,D,degrees,grades,c):
    D=validate(F,D,degrees,grades,True);n=len(D);R,V,W,B=(c[k] for k in ['R','V','W','B'])
    check(F.mm(D,V)==R,'R=DV');check(F.mm(D,W)==F.mm(W,B),'DW=WB')
    check(F.mm(B,B)==zero(n,n),'paired boundary square')
    for A,unit in [(V,True),(W,False)]:
        check(len(A)==n and all(len(row)==n for row in A),'basis dimensions')
        check(all(A[i][j]==0 for i in range(n) for j in range(i)),'upper basis')
        check(all(A[i][i]==1 if unit else A[i][i]!=0 for i in range(n)),'basis diagonal')
        check(all(not A[i][j] or degrees[i]==degrees[j] and grades[i]<=grades[j] for i in range(n) for j in range(n)),'graded filtered basis')
    piv=[low(col(R,j)) for j in range(n) if low(col(R,j)) is not None]
    check(len(piv)==len(set(piv)),'unique pivots')
    check(all(low(col(R,i)) is None for i in piv),'paired birth column is zero')
    # Verify the submitted combinatorial output as well as its basis identities.
    pairs=sorted((low(col(R,j)),j) for j in range(n) if low(col(R,j)) is not None)
    births={i for i,j in pairs}
    essential=[i for i in range(n) if low(col(R,i)) is None and i not in births]
    expected_W=[row[:] for row in V];expected_B=zero(n,n)
    for i,j in pairs:
        for k in range(n):expected_W[k][i]=R[k][j]
        expected_B[i][j]=1
    bars=sorted([(degrees[i],grades[i],grades[j]) for i,j in pairs if grades[i]<grades[j]]+[(degrees[i],grades[i],None) for i in essential],key=lambda z:(z[0],z[1],float('inf') if z[2] is None else z[2]))
    need(c['pairs']==pairs,'certificate pairs disagree with reduced lows')
    need(c['essential']==essential,'certificate essential indices disagree with reduced lows')
    need(W==expected_W,'certificate W disagrees with paired-column construction')
    need(B==expected_B,'certificate B disagrees with paired-column construction')
    need(c['bars']==bars,'certificate bars disagree with pairs and actual grades')
    return c

def persistent_rank(F,D,degrees,grades,k,s,t):
    """Independent cycle-plus-boundary rank oracle, without reduced columns."""
    n=len(D);early=[j for j in range(n) if degrees[j]==k and grades[j]<=s]
    Z=kernel(F,[[D[i][j] for j in early] for i in range(n)],len(early))
    cycles=[]
    for v in Z:
        z=[0]*n
        for j,a in zip(early,v):z[j]=a
        cycles.append(z)
    boundaries=[col(D,j) for j in range(n) if degrees[j]==k+1 and grades[j]<=t]
    mat=lambda vs:[[v[i] for v in vs] for i in range(n)]
    return rank(F,mat(boundaries+cycles))-rank(F,mat(boundaries))

def check_ranks(F,D,degrees,grades,bars):
    levels=sorted(set(grades))
    for s in levels:
        for t in levels:
            if s>t:continue
            for k in sorted(set(degrees)):
                count=sum(dim==k and birth<=s and (death is None or t<death) for dim,birth,death in bars)
                check(count==persistent_rank(F,D,degrees,grades,k,s,t),'barcode versus direct persistent image rank')

def cancel(F,D,degrees,grades,a,b,filtered=True,integer_units=False):
    D=validate(F,D,degrees,grades);n=len(D);need(isinstance(a,int) and isinstance(b,int) and 0<=a<n and 0<=b<n,'pivot indices')
    need(degrees[a]==degrees[b]+1 and D[b][a]!=0,'adjacent nonzero pivot')
    if integer_units:need(F.p==0 and all(x.denominator==1 for row in D for x in row) and abs(D[b][a])==1,'integer pivot must be a unit')
    if filtered:need(grades[a]==grades[b],'equal filtration level required')
    h=zero(n,n);h[a][b]=F.inv(D[b][a]);P=F.add(F.add(identity(n),F.mm(D,h),-1),F.mm(h,D),-1)
    keep=[j for j in range(n) if j not in [a,b]];m=len(keep)
    inc=[[P[i][j] for j in keep] for i in range(n)];proj=[P[i][:] for i in keep]
    d=F.mm(F.mm(proj,D),[[int(i==j) for j in keep] for i in range(n)])
    check(F.mm(proj,inc)==identity(m),'pi=1')
    check(F.mm(inc,proj,n)==P,'ip=1-dh-hd')
    check(F.mm(D,inc)==F.mm(inc,d),'di=idprime')
    check(F.mm(proj,D)==F.mm(d,proj,n),'pd=dprimep')
    check(F.mm(d,d)==zero(m,m),'reduced differential')
    check(F.mm(h,h)==zero(n,n) and F.mm(h,inc)==zero(n,m) and F.mm(proj,h)==zero(m,n),'contraction side conditions')
    if filtered:
        check(all(not inc[i][j] or grades[i]<=grades[keep[j]] for i in range(n) for j in range(m)),'filtered inclusion')
        check(all(not proj[i][j] or grades[keep[i]]<=grades[j] for i in range(m) for j in range(n)),'filtered projection')
    return dict(D=d,degrees=[degrees[j] for j in keep],grades=[grades[j] for j in keep],keep=keep,i=inc,p=proj,h=h)

def simplicial(F,assignment):
    cells=sorted(assignment,key=lambda c:(assignment[c],len(c),c));n=len(cells);idx={c:j for j,c in enumerate(cells)};D=zero(n,n)
    for j,c in enumerate(cells):
        need(tuple(sorted(set(c)))==c and c,'ordered nonempty simplex')
        if len(c)>1:
            for k in range(len(c)):
                face=c[:k]+c[k+1:];need(face in idx,'missing face');D[idx[face]][j]=F.value((-1)**k)
    deg=[len(c)-1 for c in cells];g=[assignment[c] for c in cells];return cells,validate(F,D,deg,g,True),deg,g

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

def run():
    global CHECKS
    CHECKS=0;rng=random.Random(202620);F=Field();main={(0,):0,(1,):0,(2,):0,(3,):0,(0,1):1,(1,2):1,(0,3):2,(2,3):2,(0,2):3,(0,1,2):3,(0,2,3):5}
    cells,D,deg,g=simplicial(F,main);c=certify(F,D,deg,g,reduce_boundary(F,D,deg,g));check_ranks(F,D,deg,g,c['bars'])
    expected=[(0,0,1),(0,0,1),(0,0,2),(0,0,None),(1,2,5)]
    check(c['bars']==expected,'main barcode');check(c['pairs']==[(1,4),(2,5),(3,6),(7,10),(8,9)],'main index pairs')
    check(F.mm(c['R'],c['R'])!=zero(len(D),len(D)),'R is not a differential')
    check(col(c['R'],10)==[0,0,0,0,1,1,-1,1,0,0,0],'dying square cycle')
    check(col(c['V'],10)==[0]*9+[1,1],'filling chain')
    contraction=cancel(F,D,deg,g,9,8,integer_units=True);small=certify(F,contraction['D'],contraction['degrees'],contraction['grades'],reduce_boundary(F,contraction['D'],contraction['degrees'],contraction['grades']))
    check(small['bars']==expected,'equal-level cancellation barcode');check_ranks(F,contraction['D'],contraction['degrees'],contraction['grades'],small['bars'])
    variant=dict(main);variant[(0,1,2)]=4;_,Dv,dv,gv=simplicial(F,variant);cv=certify(F,Dv,dv,gv,reduce_boundary(F,Dv,dv,gv));check_ranks(F,Dv,dv,gv,cv['bars'])
    check((1,3,4) in cv['bars'] and len(cv['bars'])==6,'delayed filling adds a bar');check(reject(lambda:cancel(F,Dv,dv,gv,9,8)),'unequal-level refusal')
    # Same pointwise Betti numbers, different actual maps.
    same_betti=[]
    for killed in [1,2]:
        A=zero(4,4);A[killed][3]=1;de=[0,1,1,2];gr=[0,0,1,1];cc=certify(F,A,de,gr,reduce_boundary(F,A,de,gr));check_ranks(F,A,de,gr,cc['bars']);same_betti.append(cc['bars'])
        check(persistent_rank(F,A,de,gr,1,0,0)==persistent_rank(F,A,de,gr,1,1,1)==1,'pointwise Betti same')
        check(persistent_rank(F,A,de,gr,1,0,1)==(killed==2),'cross-time map differs')
    # Coefficient dependence and the integer-unit boundary.
    A=[[0,2],[0,0]];check(reject(lambda:cancel(F,A,[0,1],[0,0],1,0,integer_units=True)),'Z nonunit refusal')
    coeff={}
    for p in [0,2,3,5]:
        K=Field(p);A=K.norm([[0,2],[0,0]]);cc=certify(K,A,[0,1],[0,1],reduce_boundary(K,A,[0,1],[0,1]));check_ranks(K,A,[0,1],[0,1],cc['bars']);coeff[str(p)]=cc['bars']
    check(coeff['2']==[(0,0,None),(1,1,None)] and coeff['0']==[(0,0,1)],'field changes bars')
    # All face-closed complexes on four fixed vertices; several prime fields.
    optional=[c for size in range(2,5) for c in combinations(range(4),size)];families=0
    for mask in range(1<<len(optional)):
        faces={(i,):0 for i in range(4)};chosen=[c for j,c in enumerate(optional) if mask>>j&1]
        if any(any(c[:i]+c[i+1:] not in faces and c[:i]+c[i+1:] not in chosen for i in range(len(c))) for c in chosen):continue
        for cell in sorted(chosen,key=lambda z:(len(z),z)):
            faces[cell]=max(faces[cell[:i]+cell[i+1:]] for i in range(len(cell)))+rng.randrange(3)
        for p in [0,2,3,5]:
            K=Field(p);_,A,de,gr=simplicial(K,faces);cc=certify(K,A,de,gr,reduce_boundary(K,A,de,gr));check_ranks(K,A,de,gr,cc['bars'])
            # Every available same-level pivot, not just a geometric free face.
            for a in range(len(A)):
                for b in range(len(A)):
                    if A[b][a] and gr[a]==gr[b]:
                        z=cancel(K,A,de,gr,a,b);dd=reduce_boundary(K,z['D'],z['degrees'],z['grades']);check(dd['bars']==cc['bars'],'all equal-level pivot barcodes')
            families+=1
    # Reorder equal-level homogeneous cells: bars independent of tie order.
    ties=0
    for order0 in permutations(range(4)):
        for order1 in [(4,5),(5,4)]:
            for order2 in [(6,7),(7,6)]:
                order=list(order0)+list(order1)+list(order2)+[8,9,10]
                for p in [0,2,3]:
                    K=Field(p);A=K.norm([[D[i][j] for j in order] for i in order]);de=[deg[i] for i in order];gr=[g[i] for i in order]
                    cc=certify(K,A,de,gr,reduce_boundary(K,A,de,gr));check(cc['bars']==expected,'equal-level ordering invariant');ties+=1
    # Compose successive contractions: h_total=h_old+i_old h_new p_old.
    all_zero={cell:0 for size in range(1,5) for cell in combinations(range(4),size)}
    _,A,de,gr=simplicial(F,all_zero);original=[row[:] for row in A];n=len(A)
    it=F.norm(identity(n));pt=F.norm(identity(n));ht=zero(n,n);steps=0
    while any(any(row) for row in A):
        a,b=next((j,i) for j in range(len(A)) for i in range(len(A)) if A[i][j])
        z=cancel(F,A,de,gr,a,b,integer_units=True)
        ht=F.add(ht,F.mm(F.mm(it,z['h']),pt))
        it=F.mm(it,z['i']);pt=F.mm(z['p'],pt)
        A,de,gr=z['D'],z['degrees'],z['grades'];steps+=1
        check(F.mm(pt,it)==identity(len(A)),'composed pi')
        check(F.mm(it,pt,n)==F.add(F.add(identity(n),F.mm(original,ht),-1),F.mm(ht,original),-1),'composed chain homotopy')
        check(F.mm(ht,ht)==zero(n,n) and F.mm(ht,it)==zero(n,len(A)) and F.mm(pt,ht)==zero(len(A),n),'composed side conditions')
    check(steps==7 and de==[0],'tetrahedron contracts to one vertex')
    # Empty input and a fully cancellable atom retain meaningful dimensions.
    ce=certify(F,[],[],[],reduce_boundary(F,[],[],[]));check(ce['bars']==[],'empty barcode')
    cancel(F,[[0,1],[0,0]],[0,1],[0,0],1,0,integer_units=True)
    check(reject(lambda:Field(4)),'composite characteristic refused');check(reject(lambda:reduce_boundary(F,[[0,1],[0,0]],[0,1],[1,0])),'bad grade refused')
    check(reject(lambda:reduce_boundary(F,[[0,1,0],[0,0,1],[0,0,0]],[0,1,2],[0,1,2])),'D square nonzero refused')
    check(reject(lambda:reduce_boundary(F,[[0,1],[0,0]],[0,0],[0,1])),'wrong degrees refused')
    check(reject(lambda:reduce_boundary(F,[[0.0]], [0],[0])),'float coefficient refused')
    return dict(status='PASS',checks=CHECKS,main_cells=cells,main_grades=g,main=c,contraction=contraction,delayed_bars=cv['bars'],same_betti_different_maps=same_betti,coefficient_bars=coeff,face_closed_field_cases=families,equal_level_order_cases=ties,composed_cancellations=steps,scope='Exact finite examples, contraction identities and independent persistent-image rank oracle; no general-purpose point-cloud filtration or performance claim.')

def serial(x):
    if isinstance(x,Fraction):return str(x)
    if isinstance(x,(tuple,list)):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__':
    p=argparse.ArgumentParser(description=__doc__);p.add_argument('--output',type=Path);a=p.parse_args();result=json.dumps(serial(run()),ensure_ascii=False,indent=2)+'\n'
    if a.output:a.output.parent.mkdir(parents=True,exist_ok=True);a.output.write_text(result);print('PASS',CHECKS,'checks;',a.output)
    else:print(result,end='')
