#!/usr/bin/env python3
"""Finite, exact certificates; general theorems are proved in the accompanying text."""
import argparse
from fractions import Fraction as Q
import json
from pathlib import Path

CHECKS = 0

def check(condition, label):
    global CHECKS
    CHECKS += 1
    if not condition:
        raise ArithmeticError(label)

def vp(x, p):
    x = Q(x)
    if not x:
        return float("inf")
    a, b, value = abs(x.numerator), x.denominator, 0
    while a % p == 0:
        a //= p
        value += 1
    while b % p == 0:
        b //= p
        value -= 1
    return value

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

def add(a, b):
    return trim([(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 scale(a, c):
    return trim([c*x for x in a])

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

def mod(a, p):
    return trim([x % p for x in a])

def divmodp(a, b, p):
    a, b = mod(a,p), mod(b,p)
    if b == [0]:
        raise ZeroDivisionError
    q = [0] * max(1, len(a)-len(b)+1)
    while a != [0] and len(a) >= len(b):
        k = len(a)-len(b)
        c = a[-1]*pow(b[-1],-1,p) % p
        q[k] = c
        a = mod(add(a,[0]*k+scale(b,-c)),p)
    return trim(q),a

def bezoutp(a, b, p):
    r0,r1 = mod(a,p),mod(b,p)
    s0,s1,t0,t1 = [1],[0],[0],[1]
    while r1 != [0]:
        q,r = divmodp(r0,r1,p)
        r0,r1 = r1,r
        s0,s1 = s1,mod(add(s0,scale(mul(q,s1),-1)),p)
        t0,t1 = t1,mod(add(t0,scale(mul(q,t1),-1)),p)
    c = pow(r0[-1],-1,p)
    return mod(scale(r0,c),p),mod(scale(s0,c),p),mod(scale(t0,c),p)

def evaluate(a, x):
    y = 0
    for c in reversed(a):
        y = y*x+c
    return y

def factor_lift(f, g0, h0, p, precision):
    g0,h0 = mod(g0,p),mod(h0,p)
    if precision < 1 or f[-1] != 1 or g0[-1] != 1 or h0[-1] != 1:
        raise ValueError("positive precision and monic inputs required")
    d,A,B = bezoutp(h0,g0,p)
    if d != [1] or mod(add(f,scale(mul(g0,h0),-1)),p) != [0]:
        raise ValueError("input factors must be coprime and multiply to f mod p")
    g,h = g0[:],h0[:]
    rows = [{"n":1,"g":g[:],"h":h[:]}]
    for n in range(1,precision):
        modulus = p**n
        difference = add(f,scale(mul(g,h),-1))
        check(all(c % modulus == 0 for c in difference),"existing factor precision")
        e = mod([c//modulus for c in difference],p)
        _,u = divmodp(mul(A,e),g0,p)
        v,remainder = divmodp(add(e,scale(mul(u,h0),-1)),g0,p)
        check(remainder == [0],"exact correction division")
        check((u == [0] or len(u)<len(g0)) and
              (v == [0] or len(v)<len(h0)),"correction degree")
        check(mod(add(mul(u,h0),mul(v,g0)),p) == e,"correction equation")
        g,h = add(g,scale(u,modulus)),add(h,scale(v,modulus))
        check(mod(add(f,scale(mul(g,h),-1)),p**(n+1)) == [0],"lifted product")
        rows.append({"n":n+1,"g":g[:],"h":h[:]})
    return rows

def polygon(a, p):
    points = [(i,vp(c,p)) for i,c in enumerate(a) if c]
    if not points:
        raise ValueError("zero polynomial")
    hull = []
    for point in points:
        while len(hull) >= 2:
            x,y = hull[-2],hull[-1]
            if (y[1]-x[1])*(point[0]-y[0]) < (point[1]-y[1])*(y[0]-x[0]):
                break
            hull.pop()
        hull.append(point)
    slopes = []
    for a,b in zip(hull,hull[1:]):
        slopes += [Q(b[1]-a[1],b[0]-a[0])]*(b[0]-a[0])
    return hull,slopes

def weight(a, p, t):
    scores = [(vp(c,p)+i*t,i) for i,c in enumerate(a) if c]
    w = min(x[0] for x in scores)
    indices = [i for x,i in scores if x == w]
    return w,min(indices),max(indices)

def simple_digits(f, p, initial, precision):
    a = initial % p
    check(f(a) % p == 0,"initial root")
    rows = [a]
    for n in range(1,precision):
        candidates = [a+t*p**n for t in range(p) if f(a+t*p**n) % p**(n+1) == 0]
        check(len(candidates)==1,"simple root unique next digit")
        a=candidates[0]
        rows.append(a)
    return rows

def infinite_value(x, precision):
    # n^2 >= precision proves the omitted terms vanish, for every x in Z_3.
    modulus = 3**precision
    answer = x**3-x*x-3*x+30
    n = 4
    while n*n < precision:
        answer += 3**(n*n)*pow(x,n,modulus)
        n += 1
    return answer % modulus

def main():
    global CHECKS
    CHECKS = 0
    f = [30,-3,-1,1]
    rows = factor_lift(f,[0,0,1],[2,1],3,8)
    g,h = rows[-1]["g"],rows[-1]["h"]
    check(g == [375,1836,1] and h == [4724,1],"3^8 factors")
    certificate = [c//3**8 for c in add(f,scale(mul(g,h),-1))]
    check(certificate == [-270,-1322,-1],"full product certificate")
    check(polygon(f,3)[1] == [Q(-1,2),Q(-1,2),Q(0)],"cubic slopes")
    for p in (2,3,5,7):
        for seed in range(1,41):
            a = [((seed*(i+2)+i*i)%17-8)*p**((seed+i)%4)
                 for i in range(1+seed%5)]
            b = [((seed*(i+5)+2*i*i)%13-6)*p**((2*seed+i)%3)
                 for i in range(1+(seed+2)%4)]
            a[0] = a[0] or 1
            b[0] = b[0] or p
            a[-1] = a[-1] or 1
            b[-1] = b[-1] or 1
            product = mul(a,b)
            check(polygon(product,p)[1] == sorted(polygon(a,p)[1]+polygon(b,p)[1]),
                  "full product slopes")
            for t in (Q(i,j) for j in (1,2,3) for i in range(-4,5)):
                x,y,z = weight(a,p,t),weight(b,p,t),weight(product,p,t)
                check(z == tuple(x[k]+y[k] for k in range(3)),"weighted endpoints")
    # Different degrees, primes, and non-squarefree individual blocks.
    models = 0
    for p in (2,3,5,7):
        for r in range(1,5):
            for s in range(1,5):
                g0 = [0]*r+[1]
                h0 = [1]
                for _ in range(s):
                    h0 = mod(mul(h0,[-1,1]),p)
                target = mul(g0,h0)
                target = add(target,[p*((i+2)*(r+s)%11-5) for i in range(r+s)])
                result = factor_lift(target,g0,h0,p,5)
                check(result[-1]["g"][-1] == result[-1]["h"][-1] == 1,
                      "generic leading coefficients")
                models += 1
    # Constant-block edge case and explicit rejection of noncoprime inputs.
    constant = factor_lift(f,[1],mod(f,3),3,8)
    check(constant[-1]["g"] == [1] and
          constant[-1]["h"] == mod(f,3**8),"constant factor")
    try:
        factor_lift([-3,0,1],[0,1],[0,1],3,3)
    except ValueError:
        check(True,"noncoprime rejection")
    else:
        check(False,"noncoprime accepted")
    x=Q(1)
    newton=[]
    for n in range(6):
        m,s=vp(x*x-17,2),vp(2*x,2)
        check(s==1 and m==2+2**(n+1),"exact Newton precision")
        newton.append({"n":n,"a":str(x),"residual_v":m,"error_v":m-s})
        x=(x+17/x)/2
    alpha4096 = 1889*pow(441,-1,4096)%4096
    check(alpha4096==1769,"2-adic representative")
    root_tables = []
    for n in range(3,13):
        modulus=2**n
        roots=[r for r in range(modulus) if (r*r-17)%modulus==0]
        true=sorted({1889*pow(441,-1,modulus)%modulus,
                     -1889*pow(441,-1,modulus)%modulus})
        check(len(roots)==4,"four finite roots")
        check(roots==sorted({t for r in true for t in (r,(r+modulus//2)%modulus)}),
              "four-class formula")
        next_roots=[r for r in range(2*modulus) if (r*r-17)%(2*modulus)==0]
        for r in roots:
            lifts=[x for x in next_roots if x%modulus==r]
            check(len(lifts)==(2 if r in true else 0),"dead versus continuing classes")
        root_tables.append({"n":n,"all_roots":roots,"infinite_branch_residues":true})
    # Many different 2-adic units, independently enumerate both adjacent finite levels.
    for u in range(1,258,8):
        for n in range(3,10):
            roots=[r for r in range(2**n) if (r*r-u)%2**n==0]
            next_roots=[r for r in range(2**(n+1)) if (r*r-u)%2**(n+1)==0]
            check(len(roots)==len(next_roots)==4,"general unit square classes")
            check(sorted(len([s for s in next_roots if s%2**n==r]) for r in roots)
                  ==[0,0,2,2],"finite branch pattern")
    rho = simple_digits(lambda x:evaluate(f,x),3,1,18)
    beta=[1]
    for n in range(1,18):
        candidates=[beta[-1]+t*3**n for t in range(3)
                    if infinite_value(beta[-1]+t*3**n,n+1)==0]
        check(len(candidates)==1,"analytic unique next digit")
        beta.append(candidates[0])
    check(rho[-1]==244635283 and beta[-1]==72448399,"infinite series root certificate")
    check(beta[-1]-rho[-1]==-4*3**16,"first changed digit")
    check(rho[15]==beta[15]==29401678,"shared sixteen digits")
    for n in range(1,9):
        for a in range(3**n):
            if n <= 6:
                check(infinite_value(a,n)==evaluate(f,a)%3**n,"certified tail truncation")
            if a%3 != 1:
                check(infinite_value(a,2)!=0,"excluded residue classes")
    # Exact finite polynomial division is a separate check of the Strassmann mechanism.
    divisions=0
    for p in (2,3,5):
        for alpha in range(-3,4):
            for seed in range(1,31):
                quotient=[p**((i+seed)%4)*((seed+i)%7-3) for i in range(6)]
                quotient[-1]=quotient[-1] or 1
                dividend=mul([-alpha,1],quotient)
                N=lambda a:max(i for i,c in enumerate(a) if c and vp(c,p)==min(vp(d,p) for d in a if d))
                recovered=[sum(dividend[k]*alpha**(k-n-1)
                               for k in range(n+1,len(dividend)))
                           for n in range(len(dividend)-1)]
                check(trim(recovered)==trim(quotient),"coefficient tail division")
                check(N(dividend)==N(quotient)+1,"last maximum descends once")
                divisions+=1
    check(vp(9*9-324,3)==5 and vp(18,3)==2,"strong input")
    check(vp(-18-9,3)==3 and vp(18-9,3)==2,"correct unique ball")
    result={"status":"PASS","checks":CHECKS,"factor_models":models,
            "strassmann_finite_divisions":divisions,"cubic_factor_rows":rows,
            "cubic_product_quotient":certificate,"newton_binary":newton,
            "binary_root_tables":root_tables,"rho_mod_3_18":rho[-1],
            "beta_mod_3_18":beta[-1],"common_mod_3_16":rho[15],
            "difference_div_3_16":-4,
            "scope":"Exact finite certificates and tail-controlled truncations; general proofs are in the text."}
    return result

if __name__=="__main__":
    parser=argparse.ArgumentParser()
    parser.add_argument("--output",type=Path)
    args=parser.parse_args()
    result=main()
    serialized=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(serialized)
    else:
        print(serialized,end="")
