#!/usr/bin/env python3
"""Exact quaternion certificates over Q and small odd prime fields.
Standard library only. Checks remain enabled under python -O.
The finite tests do not replace the field-independent proofs in the articles.
"""
from fractions import Fraction
from math import isqrt
from itertools import product
from collections import Counter, defaultdict
from pathlib import Path
import argparse
import copy
import json

COUNTS = Counter()

def demand(ok, message):
    if not ok:
        raise ValueError(message)

def check(ok, category):
    demand(ok, category)
    COUNTS[category] += 1

class Field:
    def __init__(self, p=None):
        if p is not None:
            demand(isinstance(p, int) and p > 2 and all(p % d for d in range(2, isqrt(p)+1)),
                   'Only Q or odd prime fields are supported')
        self.p = p
    def __call__(self, x):
        x = Fraction(x)
        if self.p is None:
            return x
        demand(x.denominator % self.p != 0, 'Noninvertible denominator')
        return x.numerator * pow(x.denominator, -1, self.p) % self.p
    def inv(self, x):
        x = self(x)
        demand(x != 0, 'Division by zero')
        return 1/x if self.p is None else pow(x, -1, self.p)
    def div(self, x, y):
        return self(x*self.inv(y))

class Algebra:
    def __init__(self, a, b, p=None):
        self.f = Field(p)
        self.a, self.b = self.f(a), self.f(b)
        demand(self.a != 0 and self.b != 0, 'Quaternion parameters must be nonzero')
        self.basis = [tuple(self.f(int(i == j)) for i in range(4)) for j in range(4)]
        self.one = self.basis[0]
        self.zero = self.q([0]*4)
    def q(self, xs):
        demand(len(xs) == 4, 'A quaternion has four coordinates')
        return tuple(self.f(x) for x in xs)
    def add(self, p, q):
        return self.q([x+y for x, y in zip(p, q)])
    def sub(self, p, q):
        return self.q([x-y for x, y in zip(p, q)])
    def scale(self, c, q):
        return self.q([self.f(c)*x for x in q])
    def mul(self, p, q):
        t, x, y, z = p
        u, v, w, s = q
        a, b = self.a, self.b
        return self.q([t*u+a*x*v+b*y*w-a*b*z*s,
                       t*v+x*u+b*(z*w-y*s),
                       t*w+y*u+a*(x*s-z*v),
                       t*s+z*u+x*w-y*v])
    def bar(self, q):
        t, x, y, z = q
        return self.q([t, -x, -y, -z])
    def norm(self, q):
        t, x, y, z = q
        return self.f(t*t-self.a*x*x-self.b*y*y+self.a*self.b*z*z)
    def trace(self, q):
        return self.f(2*q[0])
    def inv(self, q):
        return self.scale(self.f.inv(self.norm(q)), self.bar(q))
    def scalar(self, q):
        return all(x == 0 for x in q[1:])
    def left_matrix(self, q):
        cols = [self.mul(q, e) for e in self.basis]
        return [list(r) for r in zip(*cols)]
    def intertwiner_matrix(self, alpha, beta):
        cols = [self.sub(self.mul(beta, e), self.mul(e, alpha)) for e in self.basis]
        return [list(r) for r in zip(*cols)]

def rref(f, rows, ncols=None):
    n = ncols if ncols is not None else len(rows[0])
    a = [[f(x) for x in row] for row in rows]
    demand(all(len(row) == n for row in a), 'Matrix row length')
    pivots = []
    r = 0
    for c in range(n):
        pivot = next((i for i in range(r, len(a)) if a[i][c] != 0), None)
        if pivot is None:
            continue
        a[r], a[pivot] = a[pivot], a[r]
        scale = f.inv(a[r][c])
        a[r] = [f(x*scale) for x in a[r]]
        for i in range(len(a)):
            if i != r and a[i][c] != 0:
                scale = a[i][c]
                a[i] = [f(x-scale*y) for x, y in zip(a[i], a[r])]
        pivots.append(c)
        r += 1
        if r == len(a):
            break
    return a, pivots

def kernel(f, rows):
    rr, pivots = rref(f, rows)
    n = len(rows[0])
    out = []
    for c in range(n):
        if c in pivots:
            continue
        v = [f(0)]*n
        v[c] = f(1)
        for r, pivot in enumerate(pivots):
            v[pivot] = f(-rr[r][c])
        out.append(tuple(v))
    return out

def determinant(f, rows):
    n = len(rows)
    demand(all(len(row) == n for row in rows), 'Square determinant')
    a = [[f(x) for x in row] for row in rows]
    det = f(1)
    for c in range(n):
        pivot = next((i for i in range(c, n) if a[i][c] != 0), None)
        if pivot is None:
            return f(0)
        if pivot != c:
            a[pivot], a[c] = a[c], a[pivot]
            det = f(-det)
        d = a[c][c]
        det = f(det*d)
        for i in range(c+1, n):
            scale = f.div(a[i][c], d)
            a[i] = [f(x-scale*y) for x, y in zip(a[i], a[c])]
    return det

def mm(f, a, b):
    return [[f(sum(a[i][k]*b[k][j] for k in range(2))) for j in range(2)] for i in range(2)]

def matrix_add(f, a, b):
    return [[f(a[i][j]+b[i][j]) for j in range(2)] for i in range(2)]

def flat(a):
    return [a[0][0], a[0][1], a[1][0], a[1][1]]

def split_generators(A, u, v):
    f, a, b = A.f, A.a, A.b
    u, v = f(u), f(v)
    demand(f(u*u-a*v*v) == b, 'False quadratic norm witness')
    return [[[f(1), f(0)], [f(0), f(1)]],
            [[f(0), a], [f(1), f(0)]],
            [[u, f(-a*v)], [v, f(-u)]]]

def split_map(A, u, v, q):
    f, a = A.f, A.a
    u, v = f(u), f(v)
    t, x, y, z = q
    return [[f(t+u*y+a*v*z), f(a*(x-v*y-u*z))],
            [f(x+v*y+u*z), f(t-u*y-a*v*z)]]

def split_inverse(A, u, v, M):
    f, a, b = A.f, A.a, A.b
    u, v = f(u), f(v)
    x00, x01, x10, x11 = flat(M)
    t, x = f.div(x00+x11, 2), f.div(f.div(x01, a)+x10, 2)
    d, e = f.div(x00-x11, 2), f.div(x10-f.div(x01, a), 2)
    y = f.div(u*d-a*v*e, b)
    z = f.div(-v*d+u*e, b)
    return A.q([t, x, y, z])

def split_certificate(A, u, v):
    G = split_generators(A, u, v)
    return {'u': str(A.f(u)), 'v': str(A.f(v)),
            'images': [[[str(x) for x in r] for r in M] for M in G]}

def verify_split(A, cert):
    f = A.f
    demand(set(cert) == {'u', 'v', 'images'}, 'Split certificate fields')
    expected = split_generators(A, cert['u'], cert['v'])
    given = [[[f(x) for x in r] for r in M] for M in cert['images']]
    demand(given == expected, 'Split generator image mismatch')
    G = given+[mm(f, given[1], given[2])]
    demand(determinant(f, [flat(M) for M in G]) != 0, 'Split map lacks four independent images')
    for i, e in enumerate(A.basis):
        for j, h in enumerate(A.basis):
            demand(mm(f, G[i], G[j]) == split_map(A, cert['u'], cert['v'], A.mul(e, h)),
                   'Split multiplication table mismatch')
    return True

def norm_from_zero(A, q):
    demand(q != A.zero and A.norm(q) == 0, 'Nonzero zero-norm input required')
    t, x, y, z = q
    f, a = A.f, A.a
    denominator = f(y*y-a*z*z)
    demand(denominator != 0, 'Zero denominator: use the square-parameter branch')
    u = f.div(t*y-a*x*z, denominator)
    v = f.div(x*y-t*z, denominator)
    demand(f(u*u-a*v*v) == A.b, 'Recovered norm identity')
    return u, v

def solve_conjugacy(A, alpha, beta, basis=None):
    demand(not A.scalar(alpha) and not A.scalar(beta), 'This solver requires two noncentral inputs')
    demand((A.trace(alpha), A.norm(alpha)) == (A.trace(beta), A.norm(beta)), 'Trace/norm mismatch')
    W = kernel(A.f, A.intertwiner_matrix(alpha, beta)) if basis is None else [A.q(v) for v in basis]
    demand(len(W) == 2, 'Expected two-dimensional intertwiner space')
    for r, s in product([0, 1, -1], repeat=2):
        c = A.add(A.scale(r, W[0]), A.scale(s, W[1]))
        if A.norm(c) != 0:
            cert = {'basis': [[str(x) for x in v] for v in W],
                    'weights': [str(A.f(r)), str(A.f(s))], 'conjugator': [str(x) for x in c]}
            verify_conjugacy(A, alpha, beta, cert)
            return cert
    raise ValueError('No invertible point in certified grid')

def verify_conjugacy(A, alpha, beta, cert):
    demand(set(cert) == {'basis', 'weights', 'conjugator'}, 'Conjugacy certificate fields')
    demand(not A.scalar(alpha) and not A.scalar(beta), 'Noncentral inputs required')
    demand((A.trace(alpha), A.norm(alpha)) == (A.trace(beta), A.norm(beta)), 'Input invariants differ')
    W = [A.q(v) for v in cert['basis']]
    demand(len(W) == 2 and len(rref(A.f, W)[1]) == 2, 'Independent two-vector basis required')
    actual = kernel(A.f, A.intertwiner_matrix(alpha, beta))
    demand(len(actual) == 2, 'Original intertwiner nullity')
    for w in W:
        demand(A.mul(beta, w) == A.mul(w, alpha), 'False kernel vector')
    demand(len(cert['weights']) == 2, 'Two weights required')
    r, s = [A.f(x) for x in cert['weights']]
    demand(all(x in {A.f(0), A.f(1), A.f(-1)} for x in [r, s]), 'Weights outside three-point grid')
    c = A.q(cert['conjugator'])
    demand(c == A.add(A.scale(r, W[0]), A.scale(s, W[1])), 'Wrong recorded linear combination')
    demand(A.norm(c) != 0, 'Singular intertwiner rejected')
    demand(A.mul(A.mul(c, alpha), A.inv(c)) == beta, 'Conjugation fails on original input')
    return True

def automorphism_matrix(A, ip, jp):
    images = [A.one, ip, jp, A.mul(ip, jp)]
    demand(A.mul(ip, ip) == A.scale(A.a, A.one), 'First generator square')
    demand(A.mul(jp, jp) == A.scale(A.b, A.one), 'Second generator square')
    demand(A.mul(jp, ip) == A.scale(-1, A.mul(ip, jp)), 'Generator anticommutation')
    demand(determinant(A.f, images) != 0, 'Generator-image rank')
    return images

def solve_automorphism(A, ip, jp):
    images = automorphism_matrix(A, ip, jp)
    rows = A.intertwiner_matrix(A.basis[1], ip)+A.intertwiner_matrix(A.basis[2], jp)
    W = kernel(A.f, rows)
    demand(len(W) == 1 and A.norm(W[0]) != 0, 'Automorphism joint kernel')
    c = W[0]
    for e, target in zip(A.basis, images):
        demand(A.mul(A.mul(c, e), A.inv(c)) == target, 'Common conjugator check')
    return c

def rejection(thunk, category):
    try:
        thunk()
    except (ValueError, ZeroDivisionError):
        COUNTS[category] += 1
        return
    raise ValueError('Broken data accepted: '+category)

def stringify(x):
    if isinstance(x, (Fraction,)):
        return str(x)
    if isinstance(x, dict):
        return {k: stringify(v) for k, v in x.items()}
    if isinstance(x, (list, tuple)):
        return [stringify(v) for v in x]
    return x

def run():
    COUNTS.clear()
    finite_summaries = []
    A = Algebra(2, -1)
    q, r = A.q([1, 1, 1, 0]), A.q([1, 1, 1, 1])
    uv = norm_from_zero(A, q)
    check(uv == (1, 1), 'Rational zero divisor to norm witness')
    M = [[Fraction(3), Fraction(4)], [Fraction(5), Fraction(6)]]
    invM = split_inverse(A, *uv, M)
    check(invM == A.q(['9/2', '7/2', '9/2', -3]), 'Target matrix exact inverse coordinates')
    check(split_map(A, *uv, invM) == M, 'Target matrix round trip')
    good_split = split_certificate(A, *uv)
    check(verify_split(A, good_split), 'Complete split certificate')
    alpha, beta = A.basis[1], A.scale(-1, A.basis[3])
    conj = solve_conjugacy(A, alpha, beta)
    auto = solve_automorphism(A, beta, A.basis[2])
    check(auto == A.q([1, 0, 1, 0]), 'Joint kernel recovers rational automorphism')
    for t, x in product(range(-3, 4), repeat=2):
        c = A.mul(A.q([1, 0, 1, 0]), A.q([t, x, 0, 0]))
        check(A.mul(beta, c) == A.mul(c, alpha), 'Complete rational conjugator family')
        check(A.norm(c) == 2*(t*t-2*x*x), 'Family norm identity')
    for aa, bb in [(2, -1), (-1, -1), (2, 3), (1, 3), (3, -2)]:
        T = Algebra(aa, bb)
        es = T.basis
        for x, y, z in product(es, repeat=3):
            check(T.mul(T.mul(x, y), z) == T.mul(x, T.mul(y, z)), 'Rational basis associativity')
        values = [T.q(v) for v in product([-1, 0, 1], repeat=4)]
        for n, x in enumerate(values):
            y = values[(17*n+8) % len(values)]
            check(T.bar(T.mul(x, y)) == T.mul(T.bar(y), T.bar(x)), 'Rational conjugation reverses product')
            check(T.norm(T.mul(x, y)) == T.norm(x)*T.norm(y), 'Rational norm multiplication')
            L = T.left_matrix(x)
            check(determinant(T.f, L) == T.norm(x)**2, 'Four-dimensional left determinant')
            check(sum(L[i][i] for i in range(4)) == 2*T.trace(x), 'Four-dimensional left trace')
            check(T.sub(T.add(T.mul(x, x), T.scale(T.norm(x), T.one)), T.scale(T.trace(x), x)) == T.zero,
                  'Reduced quadratic identity')
            if T.norm(x):
                check(T.mul(x, T.inv(x)) == T.one and T.mul(T.inv(x), x) == T.one, 'Both inverse directions')
            elif x != T.zero:
                check(T.bar(x) != T.zero and T.mul(x, T.bar(x)) == T.zero and T.mul(T.bar(x), x) == T.zero,
                      'Nonzero two-sided annihilator')
    # An intentionally bad basis: neither individual vector is invertible.
    T = Algebra(1, 1)
    bad_basis = [T.q([1, 1, 0, 0]), T.q([1, -1, 0, 0])]
    check(all(T.norm(v) == 0 for v in bad_basis), 'Two singular basis vectors')
    grid_cert = solve_conjugacy(T, T.basis[1], T.basis[1], bad_basis)
    check(verify_conjugacy(T, T.basis[1], T.basis[1], grid_cert), 'Three-point grid succeeds despite bad basis')
    N = [[Fraction(0), Fraction(1)], [Fraction(0), Fraction(0)]]
    NT = [[Fraction(0), Fraction(0)], [Fraction(1), Fraction(0)]]
    nq, ntq = split_inverse(T, 1, 0, N), split_inverse(T, 1, 0, NT)
    nilpotent_cert = solve_conjugacy(T, nq, ntq)
    check(verify_conjugacy(T, nq, ntq, nilpotent_cert), 'Noncentral repeated-root conjugacy')
    rejection(lambda: solve_conjugacy(T, T.zero, nq), 'Scalar and nilpotent branch rejected')
    square_model = Algebra(1, 3)
    infinity_zero = square_model.q([1, 1, 0, 0])
    rejection(lambda: norm_from_zero(square_model, infinity_zero), 'Zero denominator branch rejected')
    check(verify_split(square_model, split_certificate(square_model, 2, 1)), 'Square-parameter fallback gives affine witness')
    # Exhaust all elements of every quaternion algebra over F3 and F5.
    for p in [3, 5]:
        for aa, bb in product(range(1, p), repeat=2):
            T = Algebra(aa, bb, p)
            f, es = T.f, T.basis
            uvp = next((u, v) for u, v in product(range(p), repeat=2) if f(u*u-aa*v*v) == bb)
            check(verify_split(T, split_certificate(T, *uvp)), 'Finite-field split certificate')
            for x, y, z in product(es, repeat=3):
                check(T.mul(T.mul(x, y), z) == T.mul(x, T.mul(y, z)), 'Finite basis associativity')
            values = [T.q(v) for v in product(range(p), repeat=4)]
            units = [x for x in values if T.norm(x)]
            by_invariants = defaultdict(list)
            for index, x in enumerate(values):
                mx = split_map(T, *uvp, x)
                check(split_inverse(T, *uvp, mx) == x, 'All finite elements matrix round trip')
                check(determinant(f, mx) == T.norm(x), 'All finite matrix determinants equal reduced norm')
                y = values[(31*index+7) % len(values)]
                check(mm(f, mx, split_map(T, *uvp, y)) == split_map(T, *uvp, T.mul(x, y)),
                      'Finite matrix multiplication matches quaternion multiplication')
                if not T.scalar(x):
                    by_invariants[(T.trace(x), T.norm(x))].append(x)
            for key, group in sorted(by_invariants.items()):
                representative = group[0]
                discriminant = f(key[0]*key[0]-4*key[1])
                if discriminant == 0:
                    expected_size, expected_centralizer = p*p-1, p*(p-1)
                elif pow(discriminant, (p-1)//2, p) == 1:
                    expected_size, expected_centralizer = p*(p+1), (p-1)**2
                else:
                    expected_size, expected_centralizer = p*(p-1), p*p-1
                check(len(group) == expected_size, 'Class size by discriminant type')
                orbit = {T.mul(T.mul(c, representative), T.inv(c)) for c in units}
                check(orbit == set(group), 'Independent full unit orbit equals invariant class')
                centralizer = [c for c in units if T.mul(c, representative) == T.mul(representative, c)]
                check(len(orbit)*len(centralizer) == len(units), 'Orbit centralizer exact count')
                check(len(centralizer) == expected_centralizer, 'Centralizer size by quadratic algebra type')
                for target in group:
                    certificate = solve_conjugacy(T, representative, target)
                    check(verify_conjugacy(T, representative, target, certificate), 'All finite noncentral targets certified')
                # The returned complete kernel is also checked against a full field enumeration.
                cert = solve_conjugacy(T, representative, group[-1])
                W = [T.q(v) for v in cert['basis']]
                span = {T.add(T.scale(r, W[0]), T.scale(s, W[1])) for r, s in product(range(p), repeat=2)}
                actual = {c for c in values if T.mul(group[-1], c) == T.mul(c, representative)}
                check(span == actual, 'Full finite intertwiner space enumerated independently')
            # Every inner conjugation is supplied only through its two generator images.
            seen_maps = set()
            for c in units:
                inv = T.inv(c)
                ip = T.mul(T.mul(c, es[1]), inv)
                jp = T.mul(T.mul(c, es[2]), inv)
                if (ip, jp) in seen_maps:
                    continue
                seen_maps.add((ip, jp))
                recovered = solve_automorphism(T, ip, jp)
                quotient = T.mul(T.inv(recovered), c)
                check(T.scalar(quotient) and quotient != T.zero, 'All finite inner maps recover unique scalar class')
            check(len(seen_maps)*(p-1) == len(units), 'Finite automorphism scalar-fiber count')
            finite_summaries.append({'p': p, 'a': aa, 'b': bb, 'norm_witness': list(uvp),
                                     'elements': len(values), 'units': len(units),
                                     'noncentral_classes': len(by_invariants), 'automorphisms_recovered': len(seen_maps)})
    # Corrupt fields independently, while keeping the original input fixed.
    x = copy.deepcopy(good_split); x['u'] = '2'
    rejection(lambda: verify_split(A, x), 'False norm witness rejected')
    x = copy.deepcopy(good_split); x['images'][2][0][1] = '2'
    rejection(lambda: verify_split(A, x), 'Wrong matrix sign rejected')
    x = copy.deepcopy(conj); x['conjugator'][0] = str(Fraction(x['conjugator'][0])+1)
    rejection(lambda: verify_conjugacy(A, alpha, beta, x), 'Altered conjugator rejected')
    x = copy.deepcopy(conj); x['basis'][1] = x['basis'][0][:]
    rejection(lambda: verify_conjugacy(A, alpha, beta, x), 'Dependent kernel basis rejected')
    x = copy.deepcopy(grid_cert); x['weights'] = ['1', '0']; x['conjugator'] = ['1', '1', '0', '0']
    R = Algebra(1, 1)
    rejection(lambda: verify_conjugacy(R, R.basis[1], R.basis[1], x), 'Nonzero singular intertwiner rejected')
    rejection(lambda: solve_automorphism(A, A.basis[1], A.basis[1]), 'False automorphism images rejected')
    rejection(lambda: Algebra(0, 1), 'Zero parameter rejected')
    rejection(lambda: Algebra(1, 1, 2), 'Characteristic two rejected')
    rejection(lambda: Algebra(1, 1, 9), 'Composite coefficient modulus rejected')
    return stringify({'status': 'PASS', 'checks': sum(COUNTS.values()), 'categories': dict(COUNTS),
        'rational_capstone': {'a': 2, 'b': -1, 'zero_divisor': q, 'annihilator': A.bar(q),
          'invertible_element': r, 'inverse': A.inv(r), 'norm_witness': uv,
          'zero_divisor_matrix': split_map(A, *uv, q), 'target_matrix': M,
          'target_inverse_coordinates': invM, 'split_certificate': good_split,
          'alpha': alpha, 'beta': beta, 'conjugacy_certificate': conj,
          'joint_automorphism_conjugator': auto},
        'singular_basis_grid_certificate': grid_cert,
        'nilpotent_matrix_conjugacy': {'alpha': nq, 'beta': ntq, 'certificate': nilpotent_cert},
        'finite_models': finite_summaries,
        'scope': 'Exact finite algebra and complete F3/F5 element, unit-orbit and intertwiner checks; not a proof of arbitrary-field theorems or a decision algorithm for rational conic solvability.'})

if __name__ == '__main__':
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--output', type=Path)
    args = parser.parse_args()
    result = run()
    encoded = json.dumps(result, ensure_ascii=False, indent=2)+'\n'
    if args.output:
        args.output.write_text(encoded)
    print(json.dumps({'status': result['status'], 'checks': result['checks'], 'finite_models': len(result['finite_models'])}))
