#!/usr/bin/env python3
"""Exact finite-chain and affine-torus certificates; Python standard library only.
No files are written unless --output is supplied. Checks survive python -O.
"""
from fractions import Fraction as Q
from itertools import product, combinations, permutations
from math import gcd, lcm
from pathlib import Path
import argparse, json

CHECKS=0

def check(value, label):
    global CHECKS
    CHECKS+=1
    if not value: raise AssertionError(label)

def eye(n): return [[int(i==j) for j in range(n)] for i in range(n)]
def transpose(A): return [list(x) for x in zip(*A)]
def mm(A,B):
    return [[sum(a*b for a,b in zip(row,col)) for col in zip(*B)] for row in A]
def mv(A,v): return [sum(a*b for a,b in zip(row,v)) for row in A]
def minus(A,B): return [[a-b for a,b in zip(x,y)] for x,y in zip(A,B)]
def trace(A): return sum(A[i][i] for i in range(len(A)))
def power(A,n):
    R=eye(len(A))
    while n:
        if n&1:R=mm(R,A)
        A=mm(A,A);n//=2
    return R

def determinant(A):
    if not A:return 1
    if len(A)==1:return A[0][0]
    return sum((-1)**j*A[0][j]*determinant([row[:j]+row[j+1:] for row in A[1:]]) for j in range(len(A)))

def snf(A):
    """Euclidean integer row/column reduction with explicit U and V."""
    M=[row[:] for row in A];n=len(M);m=len(M[0]);U=eye(n);V=eye(m)
    def rs(i,j):M[i],M[j]=M[j],M[i];U[i],U[j]=U[j],U[i]
    def cs(i,j):
        for C in (M,V):
            for row in C:row[i],row[j]=row[j],row[i]
    def ra(i,j,c):
        M[i]=[a+c*b for a,b in zip(M[i],M[j])];U[i]=[a+c*b for a,b in zip(U[i],U[j])]
    def ca(i,j,c):
        for C in (M,V):
            for row in C:row[i]+=c*row[j]
    for k in range(min(n,m)):
        nz=[(abs(M[i][j]),i,j) for i in range(k,n) for j in range(k,m) if M[i][j]]
        if not nz:break
        _,i,j=min(nz);rs(k,i);cs(k,j)
        while True:
            restart=False
            for i in range(k+1,n):
                if M[i][k]:
                    ra(i,k,-(M[i][k]//M[k][k]))
                    if M[i][k]:rs(i,k)
                    restart=True;break
            if restart:continue
            for j in range(k+1,m):
                if M[k][j]:
                    ca(j,k,-(M[k][j]//M[k][k]))
                    if M[k][j]:cs(j,k)
                    restart=True;break
            if restart:continue
            bad=next(((i,j) for i in range(k+1,n) for j in range(k+1,m) if M[i][j]%M[k][k]),None)
            if bad:ra(k,bad[0],1);continue
            break
        if M[k][k]<0:
            M[k]=[-x for x in M[k]];U[k]=[-x for x in U[k]]
    return U,M,V

def verify_snf(A,U,D,V):
    check(mm(mm(U,A),V)==D,'UNV=D')
    check(abs(determinant(U))==abs(determinant(V))==1,'integer units')
    d=[D[i][i] for i in range(min(len(D),len(D[0])))];nz=[x for x in d if x]
    check(all(D[i][j]==0 for i in range(len(D)) for j in range(len(D[0])) if i!=j),'diagonal')
    check(all(x>0 for x in nz) and d==nz+[0]*(len(d)-len(nz)),'positive then zero')
    check(all(b%a==0 for a,b in zip(nz,nz[1:])),'divisibility')

def fracmod(x):return Q(x)%1

def finite_solutions(N,b):
    """Solve N*x=-b mod Z for nonsingular N and exact integer/rational b."""
    b=[Q(x) for x in b]
    U,D,V=snf(N);verify_snf(N,U,D,V)
    d=[D[i][i] for i in range(len(N))]
    if not all(d):raise ValueError('positive-dimensional or incompatible: use family classification')
    c=[-x for x in mv(U,b)]
    points={tuple(fracmod(x) for x in mv(V,[(c[i]+j[i])/d[i] for i in range(len(d))])) for j in product(*(range(s) for s in d))}
    check(len(points)==abs(determinant(N)),'complete distinct finite solutions')
    check(all(all(fracmod(a+t)==0 for a,t in zip(mv(N,p),b)) for p in points),'all points satisfy original equation')
    return points,(U,D,V)

def affine_iterate(A,b,n):
    d=len(A);P=eye(d);c=[Q(0)]*d
    for _ in range(n):P=mm(A,P);c=[x+y for x,y in zip(mv(A,c),b)]
    return P,c

def cycles_on(points,A,b):
    todo=set(points);out=[]
    while todo:
        p=min(todo);q=p;orbit=[]
        while q not in orbit:
            check(q in todo,'cycle has no pretail/repeated earlier orbit')
            orbit.append(q);todo.remove(q)
            q=tuple(fracmod(x+y) for x,y in zip(mv(A,q),b))
        check(q==p,'orbit closes at first point')
        out.append(orbit)
    return out

def mu(n):
    ans=1;p=2
    while p*p<=n:
        if n%p==0:
            n//=p;ans=-ans
            if n%p==0:return 0
        p+=1
    return -ans if n>1 else ans

def rref(A,ncols=None):
    M=[[Q(x) for x in row] for row in A];m=len(M);n=len(M[0]) if m else ncols or 0;piv=[];r=0
    for c in range(n):
        k=next((i for i in range(r,m) if M[i][c]),None)
        if k is None:continue
        M[r],M[k]=M[k],M[r];a=M[r][c];M[r]=[x/a for x in M[r]]
        for i in range(m):
            if i!=r and M[i][c]:
                a=M[i][c];M[i]=[x-a*y for x,y in zip(M[i],M[r])]
        piv.append(c);r+=1
        if r==m:break
    return M,piv

def kernel(A,n):
    R,piv=rref(A,n);free=[j for j in range(n) if j not in piv];out=[]
    for j in free:
        v=[Q(0)]*n;v[j]=1
        for i,c in enumerate(piv):v[c]=-R[i][j]
        out.append(v)
    return out

def column_basis(A):
    if not A:return []
    _,piv=rref(A);return [[row[j] for row in A] for j in piv]

def rank_vectors(vectors):
    if not vectors:return 0
    return len(rref(transpose(vectors))[1])

def coordinates(v,basis):
    if not basis:
        check(not any(v),'zero vector in empty basis');return []
    R,piv=rref([row+[a] for row,a in zip(transpose(basis),v)])
    check(len(piv)==len(basis) and all(c<len(basis) for c in piv),'unique basis coordinates')
    return [R[i][-1] for i in range(len(basis))]

def faces(facets):
    return sorted({tuple(c) for f in facets for n in range(1,len(f)+1) for c in combinations(sorted(f),n)},key=lambda x:(len(x),x))
def sign_order(xs):
    if len(set(xs))<len(xs):return 0
    return (-1)**sum(xs[i]>xs[j] for i in range(len(xs)) for j in range(i+1,len(xs)))
def chains(K):
    return [[s for s in K if len(s)==q+1] for q in range(max(map(len,K)))]
def boundary(C,q):
    if q==0:return []
    M=[[0]*len(C[q]) for _ in C[q-1]];index={s:i for i,s in enumerate(C[q-1])}
    for j,s in enumerate(C[q]):
        for k in range(len(s)):M[index[s[:k]+s[k+1:]]][j]+=(-1)**k
    return M

def induced(C,D,g,q):
    M=[[0]*len(C[q]) for _ in D[q]] if q<len(D) else []
    ix={s:i for i,s in enumerate(D[q])} if q<len(D) else {}
    for j,s in enumerate(C[q]):
        v=tuple(g[i] for i in s);sgn=sign_order(v)
        if sgn:M[ix[tuple(sorted(v))]][j]=sgn
    return M

def homology_trace(C,F,q):
    n=len(C[q]);Z=kernel(boundary(C,q),n)
    B=column_basis(boundary(C,q+1)) if q+1<len(C) else []
    H=[];basis=B[:]
    for v in Z:
        if rank_vectors(basis+[v])>len(basis):basis.append(v);H.append(v)
    tr=Q(0)
    for j,h in enumerate(H):tr+=coordinates(mv(F,h),basis)[len(B)+j]
    return len(H),tr

def check_simplicial():
    records=[];total_maps=0
    complexes=[faces([(0,1,2)]),faces([(0,1),(1,2),(0,2)]),faces([(0,1),(1,2),(2,3),(0,3)]),faces([(0,1),(1,2),(2,3),(3,4),(0,4)]),faces(list(combinations(range(4),3)))]
    for K in complexes:
        C=chains(K);verts=sorted({v for s in K for v in s});Ks=set(K);count=0
        for vals in product(verts,repeat=len(verts)):
            g=dict(zip(verts,vals))
            if any(tuple(sorted(set(g[v] for v in s))) not in Ks for s in K):continue
            F=[induced(C,C,g,q) for q in range(len(C))]
            for q in range(1,len(C)):check(mm(boundary(C,q),F[q])==mm(F[q-1],boundary(C,q)),'chain-map boundary')
            ht=[homology_trace(C,F[q],q) for q in range(len(C))]
            L=sum((-1)**q*trace(F[q]) for q in range(len(C)))
            check(L==sum((-1)**q*t for q,(_,t) in enumerate(ht)),'Hopf trace independently via kernel quotient')
            fixed_simplex=any(set(g[v] for v in s)==set(s) for s in K)
            check(not L or fixed_simplex,'nonzero Lefschetz implies invariant simplex/barycenter')
            count+=1
        total_maps+=count;records.append({'vertices':len(verts),'simplices':len(K),'self_maps':count})
    # Degree-two map: oriented cycle bases, with wrap edges oriented forward.
    B=[[int(i==(j+1)%12)-int(i==j) for j in range(12)] for i in range(12)]
    T=[[int(i==(j+1)%3)-int(i==j) for j in range(3)] for i in range(3)]
    labels=[(i//2)%3 for i in range(12)]
    G0=[[int(i==labels[j]) for j in range(12)] for i in range(3)]
    G1=[[int(j%2==1 and i==labels[j]) for j in range(12)] for i in range(3)]
    check(mm(T,G1)==mm(G0,B),'degree-two chain-map equation')
    check(mv(G1,[1]*12)==[2]*3,'degree two fundamental cycle')
    # Exact open-star interval endpoints after choosing a lift.
    for i in range(12):
        j=i//2;lo=Q(i-1,6);hi=Q(i+1,6);center=Q(j,3)
        check(lo>=center-Q(1,3) and hi<=center+Q(1,3),'open-star inclusion endpoints')
    # Prism signs for all maps of an edge into a full 2-simplex.
    K=faces([(0,1)]);L=faces([(0,1,2)]);C=chains(K);D=chains(L)
    def prism(g,h,q):
        M=[[0]*len(C[q]) for _ in D[q+1]];idx={s:i for i,s in enumerate(D[q+1])}
        for j,s in enumerate(C[q]):
            for i in range(q+1):
                v=tuple(g[x] for x in s[:i+1])+tuple(h[x] for x in s[i:]);sgn=sign_order(v)
                if sgn:M[idx[tuple(sorted(v))]][j]+=(-1)**i*sgn
        return M
    maps=list(product(range(3),repeat=2))
    for g,h in product(maps,repeat=2):
        P0,P1=prism(g,h,0),prism(g,h,1)
        check(mm(boundary(D,1),P0)==minus(induced(C,D,h,0),induced(C,D,g,0)),'prism q0')
        a=mm(boundary(D,2),P1);b=mm(P0,boundary(C,1))
        check([[x+y for x,y in zip(r,s)] for r,s in zip(a,b)]==minus(induced(C,D,h,1),induced(C,D,g,1)),'prism q1')
    # Actual barycentric subdivision matrices and a last-vertex chain inverse.
    for K in complexes:
        C=chains(K);index={s:i for i,s in enumerate(K)};facets=[s for s in K if not any(set(s)<set(t) for t in K)]
        jf=[]
        for s in facets:
            for perm in permutations(s):jf.append(tuple(index[tuple(sorted(perm[:i+1]))] for i in range(len(s))))
        J=faces(jf);E=chains(J);last={i:max(s) for i,s in enumerate(K)};subs=[]
        for q in range(len(C)):
            ix={t:i for i,t in enumerate(E[q])};S=[[0]*len(C[q]) for _ in E[q]]
            for j,simplex in enumerate(C[q]):
                for perm in permutations(simplex):
                    t=tuple(index[tuple(sorted(perm[:i+1]))] for i in range(q+1))
                    S[ix[t]][j]+=sign_order(perm)
            subs.append(S)
            check(mm(induced(E,C,last,q),S)==eye(len(C[q])),'last-vertex after subdivision chain identity')
            if q:check(mm(boundary(E,q),S)==mm(subs[q-1],boundary(C,q)),'subdivision commutes with boundary')
    # Mesh contraction on every comparable pair of nonempty faces in dimensions 1..5.
    for d in range(1,6):
        fs=[s for k in range(1,d+2) for s in combinations(range(d+1),k)]
        for a in fs:
            for b in fs:
                if not set(a)<set(b):continue
                ba=[Q(int(i in a),len(a)) for i in range(d+1)];bb=[Q(int(i in b),len(b)) for i in range(d+1)]
                check(sum((x-y)**2 for x,y in zip(ba,bb))<=2*Q(d,d+1)**2,'squared barycentric mesh contraction')
    return {'finite_complex_maps':records,'total_self_maps':total_maps,'degree_two_vertex_labels':labels,'prism_pairs':len(maps)**2}

def check_grid_cases():
    models=0
    for vals in product(range(-2,3),repeat=4):
        A=[list(vals[:2]),list(vals[2:])];N=minus(A,eye(2));U,D,V=snf(N);verify_snf(N,U,D,V)
        for m in (2,3,5):
            b=[Q((vals[0]+vals[2])%m,m),Q((vals[1]-vals[3])%m,m)]
            bm=[int(t*m) for t in b];c=[-int(t*m) for t in mv(U,b)]
            raw={z for z in product(range(m),repeat=2) if all((v+t)%m==0 for v,t in zip(mv(N,z),bm))}
            diag=[z for z in product(range(m),repeat=2) if all((D[i][i]*z[i]-c[i])%m==0 for i in range(2))]
            recovered={tuple(v%m for v in mv(V,z)) for z in diag}
            check(raw==recovered,'direct versus Smith finite-grid equation')
            models+=1
    # Higher-dimensional certificates, including zero and rank deficient matrices.
    for k in range(80):
        A=[[(i*7+j*3+k*(i-j+1))%9-4 for j in range(3)] for i in range(3)]
        U,D,V=snf(A);verify_snf(A,U,D,V)
        m=3;c=[k%m,(2*k+1)%m,(k+2)%m]
        raw={z for z in product(range(m),repeat=3) if all((a-b)%m==0 for a,b in zip(mv(A,z),c))}
        uc=mv(U,c)
        recovered={tuple(x%m for x in mv(V,z)) for z in product(range(m),repeat=3) if all((D[i][i]*z[i]-uc[i])%m==0 for i in range(3))}
        check(raw==recovered,'3D Smith affine rank case');models+=1
    return models

def integer_inverse(V):
    R,piv=rref([row+e for row,e in zip(V,eye(len(V)))])
    check(piv[:len(V)]==list(range(len(V))),'unimodular inverse pivots')
    inv=[row[len(V):] for row in R]
    check(all(x.denominator==1 for row in inv for x in row),'inverse integer entries')
    return [[int(x) for x in row] for row in inv]

def quotient_group(relation,n):
    if n==0:return {'free_rank':0,'torsion':[]}
    if not relation or not relation[0]:return {'free_rank':n,'torsion':[]}
    U,D,V=snf(relation);verify_snf(relation,U,D,V)
    ds=[D[i][i] for i in range(min(len(D),len(D[0]))) if D[i][i]]
    return {'free_rank':n-len(ds),'torsion':[x for x in ds if x>1]}

def integer_chain_homology(boundaries,dims,q):
    M=boundaries[q];U,D,V=snf(M);verify_snf(M,U,D,V)
    r=sum(D[i][i]!=0 for i in range(min(len(D),len(D[0]))))
    if q+1==len(dims):rel=[[] for _ in range(dims[q]-r)]
    else:
        nxt=mm(integer_inverse(V),boundaries[q+1])
        check(not any(x for row in nxt[:r] for x in row),'next boundary lies in integer kernel')
        rel=nxt[r:]
    return quotient_group(rel,dims[q]-r)

def torus_mapping_homology(A):
    M=minus(eye(2),A);det=determinant(A);cok=quotient_group(M,2)
    rank=2-cok['free_rank'];a=1-det
    expected=[{'free_rank':1,'torsion':[]},
              {'free_rank':1+cok['free_rank'],'torsion':cok['torsion']},
              {'free_rank':2-rank+int(a==0),'torsion':[abs(a)] if abs(a)>1 else []},
              {'free_rank':int(a==0),'torsion':[]}]
    boundaries=[[[0]],[[0,0,0]],[[0]+M[0],[0]+M[1],[0,0,0]],[[a],[0],[0]]]
    actual=[integer_chain_homology(boundaries,[1,3,3,1],q) for q in range(4)]
    check(actual==expected,'mapping-torus integer cone versus kernel-cokernel')
    return actual

def main():
    global CHECKS
    CHECKS=0
    simplicial=check_simplicial();grid_cases=check_grid_cases()
    A=[[3,2],[1,1]];b=[Q(0),Q(0)];fixed={};records=[]
    hand=[([[0,1],[1,-2]],eye(2),[1,2]),([[1,-2],[2,-5]],[[1,-2],[0,1]],[2,6]),([[-1,3],[3,-8]],eye(2),[5,10]),([[3,-8],[7,-19]],[[1,-2],[0,1]],[8,24])]
    for n in range(1,5):
        P,c=affine_iterate(A,b,n);N=minus(P,eye(2));points,cert=finite_solutions(N,c)
        U,V,ds=hand[n-1];D=[[ds[0],0],[0,ds[1]]];verify_snf(N,U,D,V)
        denominator=lcm(*(x.denominator for p in points for x in p))
        raw={tuple(Q(z,denominator) for z in p) for p in product(range(denominator),repeat=2) if all(v%denominator==0 for v in mv(N,p))}
        check(raw==points,'complete independent denominator grid')
        orbit=cycles_on(points,A,b);fixed[n]=len(points)
        primitive=sum(mu(n//d)*fixed[d] for d in range(1,n+1) if n%d==0)
        check(primitive==sum(len(o) for o in orbit if len(o)==n),'Mobius primitive points')
        check(primitive%n==0,'orbit divisibility')
        L=1-trace(P)+determinant(P);check(L==determinant(minus(eye(2),P))==-len(points),'homology trace determinant sign')
        records.append({'n':n,'A_power':P,'N':N,'U':U,'V':V,'smith_factors':ds,'fixed_points':len(points),'least_period_points':primitive,'least_period_orbits':primitive//n,'Lefschetz':L,'common_denominator':denominator,'all_orbits':[[[str(x) for x in p] for p in o] for o in orbit]})
    # Translation conjugacy for a rational affine perturbation.
    p=[Q(1,3),Q(1,5)];b1=mv(minus(eye(2),A),p)
    for n in range(1,5):
        P,c=affine_iterate(A,b1,n);points,_=finite_solutions(minus(P,eye(2)),c)
        linear,_=finite_solutions(minus(P,eye(2)),[Q(0),Q(0)])
        check(points=={tuple(fracmod(x+y) for x,y in zip(z,p)) for z in linear},'affine conjugacy exact sets')
    # Shear with half-translation: no fixed points; two period-two circles.
    S=[[1,1],[0,1]];bb=[Q(0),Q(1,2)]
    for m in (4,8,12):
        for n in (1,2):
            P,c=affine_iterate(S,bb,n);N=minus(P,eye(2));raw=[]
            for z in product(range(m),repeat=2):
                x=[Q(t,m) for t in z]
                if all(fracmod(v+t)==0 for v,t in zip(mv(N,x),c)):raw.append(tuple(x))
            expected=[] if n==1 else [(Q(j,m),t) for j in range(m) for t in (Q(1,4),Q(3,4))]
            check(set(raw)==set(expected),'empty versus two circles')
    # Mayer-Vietoris two-seam block identity for scalar and matrix endomorphisms.
    for vals in product(range(-2,3),repeat=4):
        F=[list(vals[:2]),list(vals[2:])];I=eye(2);Z=[[0,0],[0,0]]
        Phi=[I[i]+F[i] for i in range(2)]+[[-x for x in I[i]]+[-x for x in I[i]] for i in range(2)]
        R=[Z[i]+[-x for x in I[i]] for i in range(2)]+[I[i]+F[i] for i in range(2)]
        C=[Z[i]+I[i] for i in range(2)]+[I[i]+[-x for x in I[i]] for i in range(2)]
        expect=[I[i]+Z[i] for i in range(2)]+[Z[i]+minus(I,F)[i] for i in range(2)]
        check(mm(mm(R,Phi),C)==expect,'two-seam I minus F block')
        check(abs(determinant(R))==abs(determinant(C))==1,'two-seam transformations units')
    homology_examples={name:torus_mapping_homology(B) for name,B in [('A',A),('A_squared',power(A,2)),('coordinate_swap',[[0,1],[1,0]]),('singular_projection',[[1,0],[0,0]])]}
    for vals in product(range(-2,3),repeat=4):torus_mapping_homology([list(vals[:2]),list(vals[2:])])
    result={'status':'PASS','arithmetic':'integers and fractions; checks are active under -O','checks':CHECKS,'simplicial':simplicial,'affine_finite_grid_models':grid_cases,'torus':records,'mapping_torus_integer_invariants':homology_examples}
    parser=argparse.ArgumentParser();parser.add_argument('--output',type=Path);args=parser.parse_args()
    text=json.dumps(result,ensure_ascii=False,indent=2)+'\n'
    if args.output:args.output.parent.mkdir(parents=True,exist_ok=True);args.output.write_text(text)
    else:print(text,end='')

if __name__=='__main__':main()
