#!/usr/bin/env python3
"""Exact finite certificates for elliptic two-isogeny and height descent.

Only the standard library is used. Polynomial identities are checked by exact
coefficient arithmetic; point and height loops are finite regression checks.
The all-rational-point conclusion additionally uses the proof in the capstone.
"""
import argparse
from collections import Counter
from fractions import Fraction as Q
from math import gcd, isqrt
import json
from pathlib import Path

COUNTS = Counter()

def check(condition, category):
    if not condition:
        raise AssertionError(category)
    COUNTS[category] += 1


def div(x, y, p):
    if not y:
        raise ZeroDivisionError('zero field denominator')
    return x * pow(y % p, -1, p) % p if p else Q(x) / Q(y)


def norm(x, p):
    return x % p if p else Q(x)


def on_curve(P, a, b, p=0):
    return P is None or norm(P[1]**2-P[0]**3-a*P[0]**2-b*P[0], p) == 0


def add(P, R, a, b, p=0):
    if P is None:
        return R
    if R is None:
        return P
    x, y = P
    r, s = R
    if x == r and norm(y+s, p) == 0:
        return None
    m = div(3*x*x+2*a*x+b, 2*y, p) if P == R else div(s-y, r-x, p)
    X = norm(m*m-a-x-r, p)
    return X, norm(m*(x-X)-y, p)


def phi(P, a, b, p=0):
    if P is None or P[0] == 0:
        return None
    x, y = P
    return norm(x+a+div(b, x, p), p), norm(y*(1-div(b, x*x, p)), p)


def dual(P, a, b, p=0):
    if P is None or P[0] == 0:
        return None
    X, Y = P
    B = a*a-4*b
    return div(X-2*a+div(B, X, p), 4, p), div(Y*(1-div(B, X*X, p)), 8, p)


def squareclass(q):
    q = Q(q)
    if not q:
        raise ValueError('zero has no multiplicative square class')
    m = abs(q.numerator*q.denominator)
    d = -1 if q < 0 else 1
    prime = 2
    while prime*prime <= m:
        parity = 0
        while m % prime == 0:
            m //= prime
            parity ^= 1
        if parity:
            d *= prime
        prime += 1
    return d*m


def alpha(P, b, p=0):
    value = 1 if P is None else b if P[0] == 0 else P[0]
    if p:
        return 1 if pow(value % p, (p-1)//2, p) == 1 else -1
    return squareclass(value)


def height(P):
    if P is None:
        return 1
    x = Q(P[0])
    return max(abs(x.numerator), x.denominator)


def rational_sqrt(q):
    q = Q(q)
    if q < 0:
        return None
    a, b = isqrt(q.numerator), isqrt(q.denominator)
    return Q(a, b) if a*a == q.numerator and b*b == q.denominator else None


def rational_points(a, b, bound):
    points = {None}
    for v in range(1, bound+1):
        for u in range(-bound, bound+1):
            if gcd(u, v) != 1:
                continue
            x = Q(u, v)
            y = rational_sqrt(x*x*x+a*x*x+b*x)
            if y is not None:
                points.update([(x, y), (x, -y)])
    return sorted(points, key=lambda P: (P is not None, P or (0, 0)))


# Multivariate integer polynomials with exponent tuples in three variables.
def poly(c=0, exponent=(0, 0, 0)):
    return {exponent: c} if c else {}


def padd(*terms):
    out = Counter()
    for term in terms:
        for exponent, coefficient in term.items():
            out[exponent] += coefficient
    return {e:c for e,c in out.items() if c}


def pmul(*terms):
    out = poly(1)
    for term in terms:
        nxt = Counter()
        for e,c in out.items():
            for f,d in term.items():
                nxt[tuple(x+y for x,y in zip(e, f))] += c*d
        out = {e:c for e,c in nxt.items() if c}
    return out


def ppow(p, n):
    return pmul(*([p]*n))


def scale(p, c):
    return {e:c*v for e,v in p.items() if c*v}


def symbolic_checks():
    x, a, b = [poly(1, tuple(int(j == i) for j in range(3))) for i in range(3)]
    x2 = ppow(x, 2)
    D = padd(x2, pmul(a, x), b)
    B = padd(ppow(a, 2), scale(b, -4))
    L = padd(x2, scale(b, -1))
    check(padd(ppow(D, 2), scale(pmul(a, D, x), -2), pmul(B, x2), scale(ppow(L, 2), -1)) == {}, 'symbolic_target_equation')
    N = padd(ppow(x, 6), scale(pmul(a, ppow(x, 5)), 2), scale(pmul(b, ppow(x, 4)), 5), scale(pmul(ppow(b, 2), x2), -5), scale(pmul(a, ppow(b, 2), x), -2), scale(ppow(b, 3), -1))
    comp = pmul(L, padd(ppow(D, 2), scale(pmul(B, x2), -1)))
    check(padd(comp, scale(N, -1)) == {}, 'symbolic_dual_composite_y')
    derivative = padd(scale(x2, 3), scale(pmul(a, x), 2), b)
    double = padd(pmul(derivative, padd(scale(pmul(x2, D), 4), scale(ppow(L, 2), -1))), scale(pmul(x2, ppow(D, 2)), -8))
    check(padd(double, scale(N, -1)) == {}, 'symbolic_doubling_y')
    # Rename the three indeterminates to x,r,k with f(t)=t^3-k*t.
    r, k = a, b
    A = pmul(padd(x, r), padd(pmul(x, r), scale(k, -1)))
    fx, fr = padd(ppow(x, 3), scale(pmul(k, x), -1)), padd(ppow(r, 3), scale(pmul(k, r), -1))
    product = padd(ppow(A, 2), scale(pmul(fx, fr), -4), scale(pmul(ppow(padd(pmul(x, r), k), 2), ppow(padd(x, scale(r, -1)), 2)), -1))
    check(product == {}, 'symbolic_paired_addition_product')
    pair_sum = padd(fx, fr, scale(pmul(padd(x, r), ppow(padd(x, scale(r, -1)), 2)), -1), scale(A, -1))
    check(pair_sum == {}, 'symbolic_paired_addition_sum')


def finite_field_checks():
    models = 0
    corrections = Counter()
    for p in [3, 5, 7, 11, 13]:
        for a in range(p):
            for b in range(1, p):
                B = (a*a-4*b) % p
                if not B:
                    continue
                models += 1
                ap = -2*a % p
                G = [None]+[(x,y) for x in range(p) for y in range(p) if on_curve((x,y), a,b,p)]
                H = [None]+[(x,y) for x in range(p) for y in range(p) if on_curve((x,y), ap,B,p)]
                image = {phi(P,a,b,p) for P in G}
                back = {dual(R,a,b,p) for R in H}
                double = {add(P,P,a,b,p) for P in G}
                check(image <= set(H) and back <= set(G), 'finite_field_image_equations')
                check({P for P in G if phi(P,a,b,p) is None} == {None,(0,0)}, 'finite_field_first_kernel')
                check({R for R in H if dual(R,a,b,p) is None} == {None,(0,0)}, 'finite_field_dual_kernel')
                for P in G:
                    check(dual(phi(P,a,b,p),a,b,p) == add(P,P,a,b,p), 'finite_field_first_composite')
                    check((alpha(P,b,p) == 1) == (P in back), 'finite_field_first_squareclass_kernel')
                    for Qp in G:
                        S = add(P,Qp,a,b,p)
                        check(phi(S,a,b,p) == add(phi(P,a,b,p),phi(Qp,a,b,p),ap,B,p), 'finite_field_first_homomorphism')
                        check(alpha(S,b,p) == alpha(P,b,p)*alpha(Qp,b,p), 'finite_field_squareclass_homomorphism')
                for R in H:
                    check(phi(dual(R,a,b,p),a,b,p) == add(R,R,ap,B,p), 'finite_field_second_composite')
                    check((alpha(R,B,p) == 1) == (R in image), 'finite_field_second_squareclass_kernel')
                kappa = 2//len({None,(0,0)} & image)
                corrections[kappa] += 1
                check(kappa == (1 if pow(B,(p-1)//2,p) == 1 else 2), 'finite_field_correction_square_test')
                check(len(G)*len(image)*len(back)*kappa == len(G)*len(H)*len(double), 'finite_field_exact_index_formula')
                # A pair of cardinalities would give the wrong index for kappa=2.
                if kappa == 2:
                    check(len(G)*len(H)*len(double) != len(G)*len(image)*len(back), 'reject_missing_kernel_correction')
    return {'models':models, 'kernel_correction_model_counts':dict(sorted(corrections.items()))}


def rational_checks():
    models = [(-3,2), (1,-2), (2,-3), (0,-1), (0,-9), (0,-25), (0,4), (0,36)]
    checked = 0
    for a,b in models:
        B = a*a-4*b
        if b*B == 0:
            continue
        G = rational_points(a,b,36)
        for P in G:
            checked += 1
            R = phi(P,a,b)
            check(on_curve(R,-2*a,B), 'rational_first_image')
            check(dual(R,a,b) == add(P,P,a,b), 'rational_composite')
            check(b % abs(alpha(P,b)) == 0, 'rational_squarefree_support')
            for Qp in G:
                S = add(P,Qp,a,b)
                check(phi(S,a,b) == add(phi(P,a,b),phi(Qp,a,b),-2*a,B), 'rational_homomorphism')
                check(alpha(S,b) == squareclass(alpha(P,b)*alpha(Qp,b)), 'rational_squareclass_homomorphism')
            if P is not None and P[0]:
                x,y=P
                d=alpha(P,b)
                z=rational_sqrt(x/d)
                u,v=z.numerator,z.denominator
                N=y*v**3/(d*u)
                check(N.denominator == 1 and N*N == d*u**4+a*u*u*v*v+(b//d)*v**4, 'rational_quartic_forward')
                if d == 1:
                    X=2*x+a+2*y/z
                    R=(X,2*z*X)
                    check(on_curve(R,-2*a,B) and dual(R,a,b)==P, 'rational_dual_preimage')
        for d in [d for d in range(-abs(b),abs(b)+1) if d and b%d == 0 and squareclass(d)==d]:
            for u in range(-5,6):
                for v in range(0,6):
                    if gcd(u,v)!=1:
                        continue
                    N=rational_sqrt(d*u**4+a*u*u*v*v+(b//d)*v**4)
                    if N is None:
                        continue
                    if v == 0:
                        check(d==1, 'quartic_zero_denominator')
                    elif u == 0:
                        check(d==squareclass(b), 'quartic_zero_numerator')
                    else:
                        P=(Q(d*u*u,v*v),Q(d*u*N,v**3))
                        check(on_curve(P,a,b) and alpha(P,b)==d, 'rational_quartic_reverse')
    # The three local impossibility certificates for E_3'.
    for d,modulus in [(2,27),(3,9),(6,9)]:
        squares={z*z%modulus for z in range(modulus)}
        for u in range(modulus):
            for v in range(modulus):
                if u%3 == v%3 == 0:
                    continue
                check((d*u**4+(36//d)*v**4)%modulus not in squares, 'E3_local_quartic_obstruction')
    return {'rational_models':len(models), 'rational_points_in_height_36_boxes':checked}


def height_checks():
    pairs=0
    for n in range(1,26):
        for u in range(-50,51):
            for v in range(0,41):
                if gcd(u,v)!=1:
                    continue
                F=(u*u+n*n*v*v)**2
                G=4*u*v*(u*u-n*n*v*v)
                g=gcd(F,abs(G))
                H=max(abs(u),v)
                HH=max(F,abs(G))//g
                check((4*n**4)%g == 0, 'height_gcd_divisibility')
                check(4*n**4*HH >= H**4 and HH <= 2*(n*n+1)**2*H**4, 'height_doubling_bounds')
                pairs += 1
    for n in [1,2,3,5,6,7,10]:
        pts=rational_points(0,-n*n,40)
        for P in pts:
            for R in pts:
                check(height(add(P,R,0,-n*n))*height(add(P,None if R is None else (R[0],-R[1]),0,-n*n)) <= 4*(n*n+1)**2*height(P)**2*height(R)**2, 'height_paired_addition_bound')
    for a in range(-8,9):
        for b in range(-8,9):
            if gcd(a,b)!=1:
                continue
            for c in range(-8,9):
                for d in range(-8,9):
                    if gcd(c,d)!=1:
                        continue
                    W=(a*c,a*d+b*c,b*d)
                    check(gcd(gcd(abs(W[0]),abs(W[1])),abs(W[2]))==1 and 2*max(map(abs,W))>=max(abs(a),abs(b))*max(abs(c),abs(d)), 'height_symmetric_pair_lemma')
    return {'primitive_height_inputs':pairs}


def capstone_checks():
    T={None,(Q(0),Q(0)),(Q(1),Q(0)),(Q(-1),Q(0))}
    check(set(rational_points(0,-1,2))==T, 'E1_complete_height_two_box')
    check(set(rational_points(0,-1,80))==T, 'E1_larger_box_regression_only')
    check({alpha(P,-1) for P in T}=={-1,1}, 'E1_first_squareclasses')
    for R in [(Q(2),Q(4)),(Q(2),Q(-4))]:
        check(on_curve(R,0,4) and dual(R,0,-1)==(Q(1),Q(0)) and alpha(R,4)==2, 'E1_second_layer_non_double')
    for u in range(-80,81):
        for v in range(0,61):
            if gcd(u,v)!=1:
                continue
            h=max(abs(u),v)
            for A,B in [(-v,u),(u+v,u-v),(v-u,u+v)]:
                check(max(abs(A),abs(B))//gcd(A,B)<=2*h, 'E1_torsion_translation_height')
    P=(Q(25,4),Q(75,8))
    check(on_curve(P,0,-25) and height(P)**3>2500, 'E5_non_torsion_threshold')
    two=add(P,P,0,-25)
    check(two==(Q(1681,144),Q(-62279,1728)), 'E5_first_double')
    heights=[]
    for _ in range(5):
        heights.append(str(height(P)))
        R=add(P,P,0,-25)
        check(height(R)>height(P) and 2500*height(R)>=height(P)**4, 'E5_repeated_height_growth')
        P=R
    for R in [(Q(20),Q(100)),(Q(5),Q(-25))]:
        check(on_curve(R,0,100) and dual(R,0,-25)==(Q(25,4),Q(75,8)) and alpha(R,100)==5, 'E5_two_preimages')
    # Concrete broken choices are rejected without changing the valid formulas.
    check((Q(8),Q(0)) != dual((Q(2),Q(4)),0,-1), 'reject_missing_dual_scaling')
    P=(Q(25,4),Q(75,8))
    image=phi(P,0,-25)
    check(dual((image[0],-image[1]),0,-25)==(two[0],-two[1]) and two[1]!=0, 'reject_one_sided_y_sign')
    return {'E1_rational_points_proved_in_text':['O',[0,0],[1,0],[-1,0]], 'E5_first_five_doubling_heights':heights, 'finite_search_scope':'Only bounded-height boxes are enumerated; completeness uses the text descent proof.'}


def main():
    COUNTS.clear()
    parser=argparse.ArgumentParser()
    parser.add_argument('--output',type=Path,required=True)
    args=parser.parse_args()
    symbolic_checks()
    details={}
    details.update(finite_field_checks())
    details.update(rational_checks())
    details.update(height_checks())
    details.update(capstone_checks())
    out={'status':'PASS','assertions':sum(COUNTS.values()),'categories':dict(sorted(COUNTS.items())),'details':details}
    args.output.parent.mkdir(parents=True,exist_ok=True)
    args.output.write_text(json.dumps(out,ensure_ascii=False,indent=2)+'\n')
    print(json.dumps({'status':out['status'],'assertions':out['assertions']},ensure_ascii=False))

if __name__=='__main__':
    main()
