#!/usr/bin/env python3
"""Exact finite models and polynomial certificates; no external dependencies.

The text proves the general theorems. Enumeration checks these finite models,
not an exhaustive search over an infinite rational-function field.
"""
import argparse
from collections import Counter, defaultdict
from functools import lru_cache
from itertools import product
from math import gcd
import json
from pathlib import Path

CHECKS = Counter()

def check(condition, label):
    CHECKS[label] += 1
    if not condition:
        raise ValueError('Certificate failed: ' + label)

def trim(a, p):
    a = [x % p for x in a]
    while a and not a[-1]: a.pop()
    return tuple(a)

def padd(a, b, p, sign=1):
    return trim([(a[i] if i < len(a) else 0) + sign*(b[i] if i < len(b) else 0)
                 for i in range(max(len(a), len(b)))], p)

def pmul(a, b, p):
    out = [0]*max(0, len(a)+len(b)-1)
    for i, x in enumerate(a):
        for j, y in enumerate(b): out[i+j] += x*y
    return trim(out, p)

def prem(a, b, p):
    if not b: raise ZeroDivisionError('zero polynomial divisor')
    a = list(trim(a, p)); inv = pow(b[-1], -1, p)
    while len(a) >= len(b):
        offset = len(a)-len(b); scale = a[-1]*inv % p
        for j, x in enumerate(b): a[offset+j] = (a[offset+j]-scale*x) % p
        while a and not a[-1]: a.pop()
    return tuple(a)

def ppow(a, n, f, p):
    out = (1,)
    while n:
        if n & 1: out = prem(pmul(out, a, p), f, p)
        a = prem(pmul(a, a, p), f, p); n //= 2
    return out

def pgcd(a, b, p):
    while b: a, b = b, prem(a, b, p)
    return trim([x*pow(a[-1], -1, p) for x in a], p) if a else ()

def irreducible(f, p):
    d = len(f)-1; x = prem((0, 1), f, p); current = x
    for k in range(1, d+1):
        current = ppow(current, p, f, p)
        if k <= d//2 and len(pgcd(f, padd(current, x, p, -1), p)) > 1:
            return False
    return current == x

class Field:
    def __init__(self, p, f):
        self.p, self.f, self.d = p, trim(f, p), len(f)-1
        if not irreducible(self.f, p): raise ValueError('reducible modulus')
        self.zero = (0,)*self.d; self.one = (1,)+(0,)*(self.d-1)
        self.elements = tuple(product(range(p), repeat=self.d))
        self.q = p**self.d
        self.x = self.vector(prem((0, 1), self.f, p))
    def vector(self, a): return tuple(a)+(0,)*(self.d-len(a))
    def scalar(self, c): return (c % self.p,)+(0,)*(self.d-1)
    def add(self, a, b): return tuple((x+y) % self.p for x, y in zip(a, b))
    def sub(self, a, b): return tuple((x-y) % self.p for x, y in zip(a, b))
    @lru_cache(None)
    def mul(self, a, b): return self.vector(prem(pmul(a, b, self.p), self.f, self.p))
    @lru_cache(None)
    def power(self, a, n):
        if n < 0:
            if a == self.zero: raise ZeroDivisionError('zero inverse')
            return self.power(self.power(a, self.q-2), -n)
        out = self.one
        while n:
            if n & 1: out = self.mul(out, a)
            a = self.mul(a, a); n //= 2
        return out
    def divide(self, a, b): return self.mul(a, self.power(b, -1))

class Cyclic:
    def __init__(self, field, base_degree=1, generator_power=1):
        self.F = field; self.n = field.d//base_degree; self.q = field.p**base_degree
        if field.d % base_degree or gcd(self.n, generator_power) != 1:
            raise ValueError('not a cyclic generator for this base')
        self.sigma = {a: field.power(a, self.q**generator_power) for a in field.elements}
        self.base = tuple(a for a in field.elements if self.sigma[a] == a)
        self.basis = tuple(field.power(field.x, j) for j in range(self.n))
    def orbit(self, a):
        out = []
        for _ in range(self.n): out.append(a); a = self.sigma[a]
        return out
    def trace(self, a):
        out = self.F.zero
        for x in self.orbit(a): out = self.F.add(out, x)
        return out
    def norm(self, a):
        out = self.F.one
        for x in self.orbit(a): out = self.F.mul(out, x)
        return out
    def weighted(self, u, z):
        F = self.F; c = F.one; out = F.zero
        for v, w in zip(self.orbit(u), self.orbit(z)):
            out = F.add(out, F.mul(c, w)); c = F.mul(c, v)
        return out
    def multiplicative(self, u):
        F = self.F
        if u == F.zero or self.norm(u) != F.one: return None
        for z in self.basis:
            b = self.weighted(u, z)
            if b != F.zero: return b
        raise ValueError('basis search unexpectedly vanished')
    def trace_one(self):
        F = self.F
        for e in self.basis:
            t = self.trace(e)
            if t != F.zero: return F.divide(e, t)
        raise ValueError('trace vanished on basis')
    def additive(self, a):
        F = self.F
        if self.trace(a) != F.zero: return None
        z = self.trace_one(); s = F.zero; out = F.zero
        for v, w in zip(self.orbit(a), self.orbit(z)):
            out = F.add(out, F.mul(s, w)); s = F.add(s, v)
        return out

def wp(a, p):
    out = [0]*(p*max(0, len(a)-1)+1)
    for j, c in enumerate(a): out[p*j] += c; out[j] -= c
    return trim(out, p)

def reduce_pole(a, p):
    """Coefficients in s=1/t. Return a = remainder + wp(correction)."""
    a = list(trim(a, p)); correction = [0]*len(a)
    for j in range(len(a)-1, 0, -1):
        if j % p == 0:
            c = a[j]; a[j] = 0; a[j//p] = (a[j//p]+c) % p
            correction[j//p] = (correction[j//p]+c) % p
    return trim(a, p), trim(correction, p)

def same_line(r, s, p):
    if not r or not s: return None
    for c in range(1, p):
        if trim([c*x for x in r], p) == s: return c
    return None

def run():
    CHECKS.clear(); models = []
    specs = [(2,(1,1),1),(2,(1,1,1),1),(2,(1,1,0,1),1),
             (2,(1,1,0,0,1),1),(2,(1,1,0,0,1),2),
             (2,(1,0,1,0,0,1),1),(3,(2,2,0,1),1),
             (5,(3,0,1),1),(7,(5,0,0,1),1),(13,(11,0,1),1)]
    for p, f, r in specs:
        F = Field(p, f)
        for k in range(1, max(2,F.d//r)):
            if gcd(k,F.d//r) != 1: continue
            C = Cyclic(F,r,k); mult = defaultdict(set); add = defaultdict(set)
            for b in F.elements:
                add[F.sub(b,C.sigma[b])].add(b)
                if b != F.zero: mult[F.divide(b,C.sigma[b])].add(b)
            check(len(C.base)==C.q, 'cyclic_base_size')
            check(C.trace(C.trace_one())==F.one,'trace_one')
            for a in F.elements:
                b=C.multiplicative(a)
                check((b is not None)==(a in mult),'multiplicative_solvability')
                if b is not None:
                    check(F.divide(b,C.sigma[b])==a,'multiplicative_substitution')
                    check({F.mul(b,c) for c in C.base if c!=F.zero}==mult[a], 'multiplicative_full_fibre')
                b=C.additive(a)
                check((b is not None)==(a in add),'additive_solvability')
                if b is not None:
                    check(F.sub(b,C.sigma[b])==a,'additive_substitution')
                    check({F.add(b,c) for c in C.base}==add[a],'additive_full_fibre')
            check(len(mult)==(F.q-1)//(C.q-1),'norm_kernel_size')
            check(len(add)==F.q//C.q,'trace_kernel_size')
            models.append({'p':p,'modulus':f,'base_degree':r,'generator_power':k,
                           'norm_one':len(mult),'trace_zero':len(add)})
    F=Field(7,(5,0,0,1));C=Cyclic(F);alpha=F.x
    attempts=[C.weighted(F.scalar(4),z) for z in C.basis]
    check(attempts==[(0,0,0),(0,0,0),(0,0,3)],'main_failed_then_successful_trials')
    u=(4,6,4);b=(1,1,0)
    check(F.divide(b,C.sigma[b])==u and C.weighted(u,F.one)==b,'main_nonconstant_norm_one')
    check(C.norm(b)==F.scalar(3),'main_nonunit_norm_of_preimage')
    check(C.sigma[u]==(4,3,1) and F.mul(u,C.sigma[u])==(3,2,3),'main_weighted_prefix_coordinates')
    beta=F.power(alpha,2)
    check(F.power(beta,3)==F.scalar(4) and F.mul(F.scalar(4),F.power(beta,2))==alpha,'main_kummer_inverse')
    G=Field(3,(2,2,0,1));D=Cyclic(G)
    check([D.trace(z) for z in D.basis]==[G.zero,G.zero,G.scalar(2)],'wild_characteristic_trace')
    check(D.trace_one()==(0,0,2) and D.additive(G.scalar(2))==G.x,'wild_characteristic_constructor')
    binomials=[]
    for p in [3,5,7,13,17,19]:
        for n in range(2,7):
            if (p-1)%n: continue
            powers={pow(c,n,p) for c in range(1,p)}
            for a in range(1,p):
                d=next(j for j in range(1,n+1) if pow(a,j,p) in powers)
                f=trim([-a]+[0]*(n-1)+[1],p);x=(0,1);current=x
                for j in range(1,n+1):
                    current=ppow(current,p,f,p)
                    degree=len(pgcd(f,padd(current,x,p,-1),p))-1
                    check(degree==(n if j%d==0 else 0),'kummer_frobenius_factor_degree')
                check(irreducible(f,p)==(d==n),'kummer_irreducible_class_order')
                binomials.append({'p':p,'n':n,'a':a,'degree':d})
    check({pow(x,4,13)for x in range(1,13)}=={1,3,9},'quartic_power_classes')
    H=Field(13,(11,0,1))
    check(H.power(H.x,4)==H.scalar(4),'quartic_degree_two_root')
    check(pmul((11,0,1),(2,0,1),13)==(9,0,0,0,1),'quartic_factorization')
    pole_models=[]
    for p,M in [(2,8),(3,6),(5,5)]:
        fibres=defaultdict(int)
        for raw in product(range(p),repeat=M+1):
            a=trim(raw,p);r,B=reduce_pole(a,p)
            check(padd(r,wp(B,p),p)==a,'pole_reconstruction')
            check(all(not c or j==0 or j%p for j,c in enumerate(r)),'pole_reduced_support')
            check(reduce_pole(r,p)==(r,()),'pole_idempotence')
            fibres[r]+=1
        check(set(fibres.values())=={p**(M//p)},'pole_complete_quotient_fibres')
        pole_models.append({'p':p,'max_pole':M,'inputs':p**(M+1),'classes':len(fibres),'fibre':p**(M//p)})
    a=(0,1,0,2,1);r,B=reduce_pole(a,3)
    check(r==(0,0,0,0,1) and B==(0,2),'main_pole_reduction')
    check(same_line(r,(0,0,0,0,2),3)==2,'main_same_extension')
    check(same_line(r,(0,1),3) is None,'main_different_extension')
    for p in [2,3,5]:
        for raw in product(range(p),repeat=4):
            a=trim(raw,p);r,A=reduce_pole(a,p)
            for c in range(1,p):
                B=(1,c,1);a2=padd(trim([c*x for x in a],p),wp(B,p),p);r2,A2=reduce_pole(a2,p)
                check(r2==trim([c*x for x in r],p),'affine_parameter_class')
                if r:
                    recovered=same_line(r,r2,p)
                    check(recovered==c,'affine_nonzero_line_scalar')
    # Damage certificates: use the wrong sign, wrong base, or wrong modulus.
    check(padd((0,0,0,0,1),wp((0,1),3),3)!=(0,1,0,2,1),'damaged_pole_record_rejected')
    try: Cyclic(Field(2,(1,1,0,0,1)),1,2)
    except ValueError: check(True,'nongenerator_rejected')
    else: check(False,'nongenerator_rejected')
    try: Field(13,(9,0,0,0,1))
    except ValueError: check(True,'reducible_modulus_rejected')
    else: check(False,'reducible_modulus_rejected')
    check(C.multiplicative(F.zero) is None,'zero_multiplicative_input_rejected')
    return {'status':'PASS','checks':sum(CHECKS.values()),'groups':dict(sorted(CHECKS.items())),
            'cyclic_models':models,'kummer_binomial_models':binomials,'pole_models':pole_models,
            'main':{'norm_one_input':u,'preimage':b,'basis_trials_for_4':attempts,
                    'artin_schreier_remainder':[0,0,0,0,1],'artin_schreier_correction':[0,2]},
            'scope':'Finite exact implementation checks; general cyclic and rational-function claims proved in the text.'}

if __name__=='__main__':
    ap=argparse.ArgumentParser();ap.add_argument('--output',type=Path);args=ap.parse_args()
    result=run();data=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(data)
    else: print(data,end='')
