#!/usr/bin/env python3
"""Exact, small finite-field certificates; standard library only.
Element integer labels are base-p coefficient vectors, NOT arithmetic modulo |L|.
q-polynomials use coefficient index i for X^(q^i); ordinary polynomials use X^i.
All checks raise explicitly, so python -O performs the same verification.
"""
import argparse
import itertools
import json
from pathlib import Path

CHECKS = 0

def check(ok, why):
    global CHECKS
    CHECKS += 1
    if not ok:
        raise ValueError(why)

def trim(a):
    a = list(a)
    while len(a) > 1 and a[-1] == 0:
        a.pop()
    return a or [0]

class Field:
    def __init__(self, p, modulus, d=1):
        self.p, self.mod = p, modulus
        self.m, self.N = len(modulus)-1, p**(len(modulus)-1)
        if self.m % d or modulus[-1] != 1:
            raise ValueError('bad field presentation')
        self.q, self.n = p**d, self.m//d
        self.digits = [tuple((x//p**i)%p for i in range(self.m)) for x in range(self.N)]
        self.add = [[self.encode([(u+v)%p for u,v in zip(a,b)]) for b in self.digits] for a in self.digits]
        self.neg = [self.encode([(-u)%p for u in a]) for a in self.digits]
        self.mul = [[self.rawmul(a,b) for b in self.digits] for a in self.digits]
        self.inv = [0]+[next((y for y in range(1,self.N) if self.mul[x][y]==1),None) for x in range(1,self.N)]
        if None in self.inv:
            raise ValueError('reducible modulus: a nonzero class is not invertible')
        self.K = [x for x in range(self.N) if self.power(x,self.q)==x]
        check(len(self.K)==self.q, 'fixed field size')
        self.basis = [self.power(p,j) for j in range(self.n)]
        self.coords = {}
        for c in itertools.product(self.K, repeat=self.n):
            self.coords[self.dot(c,self.basis)] = c
        check(len(self.coords)==self.N, 'relative power basis')
        self.dual = [next(x for x in range(self.N) if all(self.trace(self.mul[x][b])==int(i==j) for j,b in enumerate(self.basis))) for i in range(self.n)]
    def encode(self,a):
        return sum(v*self.p**i for i,v in enumerate(a))
    def rawmul(self,a,b):
        c=[0]*(2*self.m-1)
        for i,u in enumerate(a):
            for j,v in enumerate(b): c[i+j]=(c[i+j]+u*v)%self.p
        for i in range(len(c)-1,self.m-1,-1):
            for j in range(self.m): c[i-self.m+j]=(c[i-self.m+j]-c[i]*self.mod[j])%self.p
        return self.encode(c[:self.m])
    def sub(self,a,b): return self.add[a][self.neg[b]]
    def power(self,a,k):
        z=1
        while k:
            if k&1: z=self.mul[z][a]
            a=self.mul[a][a]; k//=2
        return z
    def total(self, seq):
        z=0
        for x in seq: z=self.add[z][x]
        return z
    def dot(self,a,b): return self.total(self.mul[x][y] for x,y in zip(a,b))
    def trace(self,x): return self.total(self.power(x,self.q**i) for i in range(self.n))
    def span(self,basis): return frozenset(self.dot(c,basis) for c in itertools.product(self.K,repeat=len(basis)))
    def basis_of(self,W):
        b=[]; S=frozenset([0])
        for x in sorted(W):
            if x not in S: b.append(x); S=self.span(b)
        if S != frozenset(W): raise ValueError('not a K-subspace')
        return b
    def rank(self,A):
        A=[list(r) for r in A]; r=0
        if not A: return 0
        for j in range(len(A[0])):
            pivot=next((i for i in range(r,len(A)) if A[i][j]),None)
            if pivot is None: continue
            A[r],A[pivot]=A[pivot],A[r]
            s=self.inv[A[r][j]]; A[r]=[self.mul[s][x] for x in A[r]]
            for i in range(len(A)):
                if i != r:
                    s=A[i][j]; A[i]=[self.sub(x,self.mul[s][y]) for x,y in zip(A[i],A[r])]
            r+=1
            if r==len(A): break
        return r
    def det(self,A):
        A=[list(r) for r in A]; z=1
        for j in range(len(A)):
            k=next((i for i in range(j,len(A)) if A[i][j]),None)
            if k is None: return 0
            if k!=j: A[k],A[j]=A[j],A[k]; z=self.neg[z]
            z=self.mul[z][A[j][j]]
            for i in range(j+1,len(A)):
                s=self.mul[A[i][j]][self.inv[A[j][j]]]
                A[i]=[self.sub(x,self.mul[s][y]) for x,y in zip(A[i],A[j])]
        return z
    def matmul(self,A,B): return [[self.dot(r,c) for c in zip(*B)] for r in A]
    def padd(self,a,b):
        return trim([self.add[a[i] if i<len(a) else 0][b[i] if i<len(b) else 0] for i in range(max(len(a),len(b)))])
    def pscale(self,a,c): return trim([self.mul[c][x] for x in a])
    def pmul(self,a,b):
        out=[0]*(len(a)+len(b)-1)
        for i,x in enumerate(a):
            for j,y in enumerate(b): out[i+j]=self.add[out[i+j]][self.mul[x][y]]
        return trim(out)
    def pdiv(self,a,b):
        a,b=trim(a),trim(b)
        if b==[0]: raise ValueError('zero divisor polynomial')
        out=[0]*max(1,len(a)-len(b)+1)
        while a != [0] and len(a)>=len(b):
            j=len(a)-len(b); c=self.mul[a[-1]][self.inv[b[-1]]]; out[j]=c
            for i,x in enumerate(b): a[i+j]=self.sub(a[i+j],self.mul[c][x])
            a=trim(a)
        return trim(out),a
    def gcd(self,a,b):
        while b != [0]: a,b=b,self.pdiv(a,b)[1]
        return self.pscale(a,self.inv[a[-1]]) if a != [0] else [0]
    def qeval(self,a,x): return self.total(self.mul[c][self.power(x,self.q**i)] for i,c in enumerate(a))
    def qcompose(self,a,b,reduce=False):
        out=[0]*(self.n if reduce else len(a)+len(b)-1)
        for i,x in enumerate(a):
            for j,y in enumerate(b):
                k=(i+j)%self.n if reduce else i+j
                out[k]=self.add[out[k]][self.mul[x][self.power(y,self.q**i)]]
        return trim(out)
    def ordinary(self,a):
        out=[0]*(self.q**(len(a)-1)+1)
        for i,c in enumerate(a): out[self.q**i]=c
        return trim(out)
    def subpoly(self,basis):
        a=[1]
        for v in basis:
            c=self.qeval(a,v)
            if c==0: raise ValueError('dependent subspace basis')
            a=self.padd([0]+[self.power(x,self.q) for x in a],self.pscale(a,self.neg[self.power(c,self.q-1)]))
        return a
    def rootpoly(self,W):
        a=[1]
        for w in sorted(W): a=self.pmul(a,[self.neg[w],1])
        return a
    def matrix(self,a): return [list(r) for r in zip(*(self.coords[self.qeval(a,b)] for b in self.basis))]
    def dickson(self,a):
        a=list(a)+[0]*(self.n-len(a))
        return [[self.power(a[(j-i)%self.n],self.q**i) for j in range(self.n)] for i in range(self.n)]
    def interpolate(self,values):
        return trim([self.total(self.mul[v][self.power(d,self.q**i)] for v,d in zip(values,self.dual)) for i in range(self.n)])
    def adjoint(self,a):
        out=[0]*self.n
        for i,c in enumerate(a): out[(-i)%self.n]=self.power(c,self.q**((-i)%self.n))
        return trim(out)

def all_spaces(F):
    spaces={frozenset([0]):[]}; frontier=[frozenset([0])]
    while frontier:
        W=frontier.pop()
        for v in range(F.N):
            if v not in W:
                b=spaces[W]+[v]; U=F.span(b)
                if U not in spaces: spaces[U]=b; frontier.append(U)
    return spaces

def distinct_factors(F,m):
    rest=m; factors=[]
    for d in range(1,len(m)):
        for c in itertools.product(F.K,repeat=d):
            g=list(c)+[1]
            if len(rest)<len(g): continue
            if F.pdiv(rest,g)[1]==[0]:
                factors.append(g)
                while F.pdiv(rest,g)[1]==[0]: rest=F.pdiv(rest,g)[0]
    check(rest==[1], 'full factorization')
    return factors

def normal_checks(F):
    n=F.n; m=[F.neg[1]]+[0]*(n-1)+[1]
    orbit=lambda x:[F.power(x,F.q**i) for i in range(n)]
    normals=[x for x in range(F.N) if F.rank([F.coords[y] for y in orbit(x)])==n]
    check(bool(normals),'normal basis exists in checked field')
    alpha=normals[0]; O=orbit(alpha)
    factors=distinct_factors(F,m); expected=F.N
    for g in factors: expected=expected//(F.q**(len(g)-1))*(F.q**(len(g)-1)-1)
    check(len(normals)==expected,'normal element count')
    for c in itertools.product(F.K,repeat=n):
        theta=F.dot(c,O)
        check((theta in normals)==(F.gcd(trim(c),m)==[1]),'unit/cyclic vector criterion')
    for x in normals: check(F.trace(x)!=0,'normal trace is nonzero')
    k=n
    while k%F.p==0 and k>1: k//=F.p
    if k==1: check(normals==[x for x in range(F.N) if F.trace(x)!=0],'p-power trace criterion')
    dual0=next(x for x in range(F.N) if all(F.trace(F.mul[x][o])==int(j==0) for j,o in enumerate(O)))
    for i,u in enumerate(orbit(dual0)):
        for j,v in enumerate(O): check(F.trace(F.mul[u][v])==int(i==j),'dual normal orbit')
    return {'q':F.q,'n':n,'normal_elements':normals,'chosen_generator':alpha,'orbit':O,'dual_generator':dual0,'distinct_factors':factors,'count':expected}

def operator_checks(F, exhaustive):
    n=F.n; B=[[F.power(b,F.q**i) for b in F.basis] for i in range(n)]
    check(F.det(B)!=0,'Moore basis determinant')
    coeffs=list(itertools.product(range(F.N),repeat=n)) if exhaustive else [tuple(((7*k+3*i*i+5*i) % F.N) for i in range(n)) for k in range(F.N)]+[tuple(int(i==j) for i in range(n)) for j in range(n)]
    rank_counts={}; full=[0]*(F.N+1); full[1]=F.neg[1]; full[F.N]=1
    for raw in coeffs:
        a=trim(raw); A=F.matrix(a); D=F.dickson(a); values=[F.qeval(a,x) for x in range(F.N)]
        ker=[x for x,y in enumerate(values) if y==0]; im=set(values); adj=F.adjoint(a)
        adjker=[x for x in range(F.N) if F.qeval(adj,x)==0]
        rank=F.rank(A); rank_counts[rank]=rank_counts.get(rank,0)+1
        check(F.matmul(D,B)==F.matmul(B,A),'D B = B A')
        check(F.rank(D)==rank and F.det(D)==F.det(A) and F.det(D) in F.K,'Dickson rank/determinant')
        check(F.interpolate([F.qeval(a,b) for b in F.basis])==a,'trace-dual interpolation')
        check(len(ker)==F.q**(n-rank) and len(im)==F.q**rank,'kernel/image dimensions')
        check(F.gcd(F.ordinary(a),full)==F.rootpoly(ker),'kernel ordinary gcd')
        check(F.ordinary(F.subpoly(F.basis_of(im)))==F.rootpoly(im),'image subspace polynomial')
        for x in range(F.N):
            check((x in im)==all(F.trace(F.mul[x][z])==0 for z in adjker),'adjoint image obstruction')
            for y in range(F.N): check(F.trace(F.mul[y][values[x]])==F.trace(F.mul[F.qeval(adj,y)][x]),'trace adjoint')
        b=trim([(3*c+1)%F.N for c in raw]); composed=F.qcompose(a,b,True)
        check(F.matmul(D,F.dickson(b))==F.dickson(composed),'Dickson composition')
        for x in range(F.N): check(F.qeval(composed,x)==F.qeval(a,F.qeval(b,x)),'composition values')
        if rank==n:
            inverse_values=[values.index(b) for b in F.basis]; inverse=F.interpolate(inverse_values)
            check(F.qcompose(a,inverse,True)==[1] and F.qcompose(inverse,a,True)==[1],'two-sided composition inverse')
    return {'operators':len(coeffs),'exhaustive':exhaustive,'rank_counts':rank_counts}

def space_checks(F):
    spaces=all_spaces(F); polys={W:F.subpoly(b) for W,b in spaces.items()}; dimensions={}
    for W,b in spaces.items():
        dimensions[len(b)]=dimensions.get(len(b),0)+1
        P=polys[W]
        check(P[-1]==1 and P[0]!=0 and len(P)==len(b)+1,'monic separable q-polynomial')
        check(F.ordinary(P)==F.rootpoly(W),'recursion equals full root product')
        check({x for x in range(F.N) if F.qeval(P,x)==0}==set(W),'subspace roots')
        for U in spaces:
            V=frozenset(F.add[w][u] for w in W for u in U)
            image=F.span([F.qeval(P,u) for u in spaces[U]])
            check(F.qcompose(F.subpoly(F.basis_of(image)),P)==polys[V],'sum via composition')
            check(F.gcd(F.ordinary(P),F.ordinary(polys[U]))==F.ordinary(polys[W&U]),'intersection via gcd')
    return {'subspaces':len(spaces),'dimensions':dimensions,'ordered_pairs':len(spaces)**2}

def rejects(fn,why):
    try: fn()
    except ValueError: check(True,why)
    else: check(False,why)

def main():
    global CHECKS
    CHECKS=0
    fields=[Field(2,[1,1,0,1]),Field(3,[1,0,1]),Field(2,[1,1,0,0,1]),Field(2,[1,1,0,0,1],2),Field(2,[1,1,0,0,0,0,1],2)]
    results=[]
    for i,F in enumerate(fields):
        r={'p':F.p,'modulus':F.mod,'base_field_labels':F.K,'basis':F.basis,'dual_basis':F.dual,'normal':normal_checks(F)}
        r['operators']=operator_checks(F,i in (0,1,3))
        if i in (1,2,3): r['spaces']=space_checks(F)
        results.append(r)
    F=fields[2]; a=2; alpha=8; ell=[a,0,1]; singular=[a,1]; W=F.span([1,a]); U=F.span([1,4]); P=F.subpoly([1,a]); Q=F.subpoly([1,4])
    inverse=F.interpolate([next(x for x in range(F.N) if F.qeval(ell,x)==b) for b in F.basis])
    sample={'a':a,'alpha':alpha,'alpha_orbit':[F.power(alpha,2**i) for i in range(4)],'a_orbit':[F.power(a,2**i) for i in range(4)],'normal_matrix':[list(r) for r in zip(*(F.coords[F.power(alpha,2**i)] for i in range(4)))], 'ell':ell,'ell_matrix':F.matrix(ell),'ell_Dickson':F.dickson(ell),'ell_inverse':inverse,'singular':singular,'singular_matrix':F.matrix(singular),'singular_adjoint':F.adjoint(singular),'singular_kernel':[x for x in range(16) if F.qeval(singular,x)==0],'adjoint_kernel':[x for x in range(16) if F.qeval(F.adjoint(singular),x)==0],'W':sorted(W),'U':sorted(U),'P_W':P,'P_U':Q,'P_intersection':F.subpoly(F.basis_of(W&U)),'P_sum':F.subpoly(F.basis_of(frozenset(F.add[w][u] for w in W for u in U))),'P_W_U_image':sorted(set(F.qeval(P,u) for u in U))}
    check(F.power(alpha,5)==1 and alpha!=1 and F.power(a,15)==1 and all(F.power(a,k)!=1 for k in (1,3,5)),'normal vs primitive witness')
    check(F.qcompose([0,1],[a],True)!=F.qcompose([a],[0,1],True),'noncommutative composition witness')
    check(F.rank(F.matrix(singular))==3 and F.rank(F.matrix(ell))==4,'main ranks')
    rejects(lambda:F.subpoly([1,1]),'dependent basis rejected')
    rejects(lambda:Field(2,[1,0,1]),'reducible modulus rejected')
    rejects(lambda:F.pdiv([1],[0]),'zero polynomial division rejected')
    fullq=[1,0,0,0,1]
    check(fullq!=[0] and all(F.qeval(fullq,x)==0 for x in range(16)),'formal/function boundary')
    check(F.qcompose(ell,inverse,True)==[1] and F.qcompose(ell,inverse)!=[1],'reduced/formal inverse boundary')
    return {'unit':'algebra-A18','status':'PASS','checks':CHECKS,'encoding':'integer sum c_i p^i for field element sum c_i a^i; ordinary/q exponents explicitly separate','fields':results,'main_F16':sample,'scope':'finite exact certificates, not general theorem proofs; full operator enumeration only for stated exhaustive cases'}

if __name__=='__main__':
    parser=argparse.ArgumentParser(); parser.add_argument('--output',type=Path)
    args=parser.parse_args(); result=main(); data=json.dumps(result,ensure_ascii=False,indent=2,sort_keys=True)+'\n'
    if args.output: args.output.write_text(data,encoding='utf-8')
    print(data,end='')
