#!/usr/bin/env python3
"""Exact certificates for the Theoryroad real-root unit.
Python 3 standard library only; no floating-point root finding or CAS calls.
Polynomial coefficients are in increasing degree order. Run with optional
--output result.json to save the report. This is an executable mathematical
check, not a formal proof-assistant verification or a production-site test.
Certificate checks and input validation remain active under python -O.
"""
from fractions import Fraction as Q
from itertools import product
from pathlib import Path
import argparse
import json


def verify(condition, message):
    """Execute a certificate check in both ordinary and optimized Python."""
    if not condition:
        raise AssertionError(message)


def poly(a):
    a = [Q(v) for v in a]
    while a and not a[-1]:
        a.pop()
    return tuple(a)


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


def sub(a, b):
    return add(a, scale(b, -1))


def mul(a, b):
    if not a or not b:
        return ()
    c = [Q(0)]*(len(a)+len(b)-1)
    for i, u in enumerate(a):
        for j, v in enumerate(b):
            c[i+j] += u*v
    return poly(c)


def divrem(a, b):
    if not b:
        raise ZeroDivisionError("zero polynomial divisor")
    r = poly(a)
    c = [Q(0)]*max(0, len(a)-len(b)+1)
    while r and len(r) >= len(b):
        k = len(r)-len(b)
        v = r[-1]/b[-1]
        c[k] += v
        r = sub(r, (Q(0),)*k+scale(b, v))
    return poly(c), r


def exactdiv(a, b):
    q, r = divrem(a, b)
    if r:
        raise ValueError("division was not exact")
    return q


def monic(a):
    return scale(a, 1/a[-1]) if a else ()


def gcd(a, b):
    while b:
        a, b = b, divrem(a, b)[1]
    return monic(a)


def derivative(a):
    return poly([i*a[i] for i in range(1, len(a))])


def ev(a, x):
    v = Q(0)
    for c in reversed(a):
        v = v*x+c
    return v


def sign(v):
    return (v > 0)-(v < 0)


def squarefree(a):
    if not a:
        raise ValueError("zero polynomial has no finite root isolation")
    return monic(exactdiv(a, gcd(a, derivative(a))))


def squarefree_blocks(a):
    a = monic(a)
    g = gcd(a, derivative(a))
    w = exactdiv(a, g)
    k, out = 1, []
    while len(w) > 1:
        y = gcd(w, g)
        z = exactdiv(w, y)
        if len(z) > 1:
            out.append((z, k))
        w, g, k = y, exactdiv(g, y), k+1
    verify(len(g) == 1, 'certificate or invariant failed: len(g) == 1')
    return out


def signed_chain(a, b):
    if not a:
        raise ValueError("signed chain requires a nonzero first polynomial")
    if not b:
        return [a]
    out = [a, b]
    while True:
        r = scale(divrem(out[-2], out[-1])[1], -1)
        if not r:
            return out
        out.append(r)


def variations(seq, x):
    signs = [sign(ev(p, x)) for p in seq if ev(p, x)]
    return sum(u != v for u, v in zip(signs, signs[1:]))


def count(p, a, b):
    if not p or not a < b:
        raise ValueError("nonzero polynomial and ordered endpoints required")
    if not ev(p, a) or not ev(p, b):
        raise ValueError("open-interval certificate requires non-root endpoints")
    f = squarefree(p)
    if len(f) == 1:
        return 0
    chain = signed_chain(f, derivative(f))
    return variations(chain, a)-variations(chain, b)


def query(p, q, a, b):
    if not p or not a < b or not ev(p, a) or not ev(p, b):
        raise ValueError("invalid query or root endpoint")
    p = squarefree(p)
    if len(p) == 1:
        return 0
    r = divrem(mul(derivative(p), q), p)[1]
    if not r:
        return 0
    d = gcd(p, r)
    chain = signed_chain(exactdiv(p, d), exactdiv(r, d))
    return variations(chain, a)-variations(chain, b)


def bound(p):
    if len(p) < 2:
        raise ValueError("root bound requires a nonconstant polynomial")
    return 1+max(abs(c/p[-1]) for c in p[:-1])


def isolate(p, epsilon=None):
    f = squarefree(p)
    if len(f) == 1:
        return [], 0
    B = bound(f)
    pending, out, fallback = [(-B, B)], [], 0
    while pending:
        a, b = pending.pop()
        n = count(f, a, b)
        if n == 0:
            continue
        if n == 1 and (epsilon is None or b-a <= epsilon):
            out.append((a, b))
            continue
        m = (a+b)/2
        if ev(f, m) == 0:
            fallback += 1
            degree = len(f)-1
            for j in range(1, degree+2):
                m = a+(Q(1, 4)+Q(j, 2*(degree+2)))*(b-a)
                if ev(f, m):
                    break
            else:
                raise AssertionError("n+1 distinct candidates cannot all be roots")
        verify(a < m < b and max(m-a, b-m) <= Q(3, 4)*(b-a), 'certificate or invariant failed: a < m < b and max(m-a, b-m) <= Q(3, 4)*(b-a)')
        verify(count(f, a, m)+count(f, m, b) == n, 'certificate or invariant failed: count(f, a, m)+count(f, m, b) == n')
        pending.extend([(a, m), (m, b)])
    return sorted(out), fallback


def equal_alg(p, I, q, J):
    if count(p, *I) != 1 or count(q, *J) != 1:
        raise ValueError("each algebraic record must isolate exactly one real root")
    a, b = max(I[0], J[0]), min(I[1], J[1])
    return a < b and count(gcd(p, q), a, b) == 1


def transpose(A):
    return [list(row) for row in zip(*A)]


def mmul(A, B):
    return [[sum(u*v for u, v in zip(row, col)) for col in zip(*B)] for row in A]


def eye(n):
    return [[Q(i == j) for j in range(n)] for i in range(n)]


def madd(A, B):
    return [[u+v for u, v in zip(r, s)] for r, s in zip(A, B)]


def mscale(A, c):
    return [[c*v for v in row] for row in A]


def trace(A):
    return sum(A[i][i] for i in range(len(A)))


def mpow(A, n):
    R = eye(len(A))
    while n:
        if n & 1:
            R = mmul(R, A)
        A, n = mmul(A, A), n//2
    return R


def matrix_poly(q, C):
    n = len(C)
    R = [[Q(0)]*n for _ in range(n)]
    for c in reversed(q):
        R = madd(mmul(R, C), mscale(eye(n), c))
    return R


def companion(p):
    p = monic(p)
    n = len(p)-1
    C = [[Q(0)]*n for _ in range(n)]
    for j in range(n-1):
        C[j+1][j] = Q(1)
    for i in range(n):
        C[i][-1] = -p[i]
    return C


def hermite(p, q):
    p = squarefree(p)
    C = companion(p)
    n = len(C)
    qC = matrix_poly(q, C)
    values = [trace(mmul(qC, mpow(C, k))) for k in range(2*n-1)]
    return [[values[i+j] for j in range(n)] for i in range(n)]


def inertia(A):
    """Exact symmetric elimination, with 2x2 pivots if all diagonals vanish."""
    n = len(A)
    if not n:
        return (0, 0, 0)
    if any(len(row) != n for row in A) or A != transpose(A):
        raise ValueError("inertia requires a square symmetric matrix")
    pivot = next((i for i in range(n) if A[i][i]), None)
    if pivot is not None:
        order = [pivot]+[i for i in range(n) if i != pivot]
        B = [[A[i][j] for j in order] for i in order]
        v = B[0][0]
        S = [[B[i][j]-B[i][0]*B[0][j]/v for j in range(1, n)] for i in range(1, n)]
        p, m, z = inertia(S)
        return (p+(v > 0), m+(v < 0), z)
    pair = next(((i, j) for i in range(n) for j in range(i+1, n) if A[i][j]), None)
    if pair is None:
        return (0, 0, n)
    i, j = pair
    order = [i, j]+[k for k in range(n) if k not in pair]
    B = [[A[i][j] for j in order] for i in order]
    v = B[0][1]
    S = [[B[i][j]-(B[i][0]*B[1][j]+B[i][1]*B[0][j])/v
          for j in range(2, n)] for i in range(2, n)]
    p, m, z = inertia(S)
    return (p+1, m+1, z)


def sig(A):
    p, m, _ = inertia(A)
    return p-m


def diagonal(v):
    return [[Q(v[i]) if i == j else Q(0) for j in range(len(v))] for i in range(len(v))]


def run():
    p, q, linear = poly([-1, -1, 0, 1]), poly([-2, 0, 1]), poly([-1, 1])
    P = mul(mul(p, q), linear)
    F = mul(P, linear)
    verify(F == poly([2, -2, -3, 1, 5, -2, -2, 1]), 'certificate or invariant failed: F == poly([2, -2, -3, 1, 5, -2, -2, 1])')
    verify(P == poly([-2, 0, 3, 2, -3, -1, 1]), 'certificate or invariant failed: P == poly([-2, 0, 3, 2, -3, -1, 1])')
    verify(gcd(F, derivative(F)) == linear, 'certificate or invariant failed: gcd(F, derivative(F)) == linear')
    verify(squarefree(F) == P, 'certificate or invariant failed: squarefree(F) == P')
    verify(squarefree_blocks(F) == [(mul(p, q), 1), (linear, 2)], 'certificate or invariant failed: squarefree_blocks(F) == [(mul(p, q), 1), (linear, 2)]')
    verify(bound(p) == 2 and bound(P) == 4, 'certificate or invariant failed: bound(p) == 2 and bound(P) == 4')
    I = [(Q(-3, 2), Q(-11, 8)), (Q(3, 4), Q(5, 4)),
         (Q(5, 4), Q(4, 3)), (Q(11, 8), Q(3, 2))]
    verify(all(I[i][1] <= I[i+1][0] for i in range(3)), 'certificate or invariant failed: all(I[i][1] <= I[i+1][0] for i in range(3))')
    verify([count(P, a, b) for a, b in I] == [1]*4, 'certificate or invariant failed: [count(P, a, b) for a, b in I] == [1]*4')
    verify(count(P, -4, 4) == count(F, -4, 4) == 4, 'certificate or invariant failed: count(P, -4, 4) == count(F, -4, 4) == 4')
    multiplicities = [next(k for block, k in squarefree_blocks(F) if count(block, a, b)) for a, b in I]
    verify(multiplicities == [1, 2, 1, 1], 'certificate or invariant failed: multiplicities == [1, 2, 1, 1]')
    verify([query(P, q, a, b) for a, b in I] == [0, -1, -1, 0], 'certificate or invariant failed: [query(P, q, a, b) for a, b in I] == [0, -1, -1, 0]')
    aggregate = [query(P, z, -4, 4) for z in [poly([1]), q, mul(q, q)]]
    verify(aggregate == [4, -2, 2], 'certificate or invariant failed: aggregate == [4, -2, 2]')
    verify([query(p, z, Q(5, 4), Q(4, 3)) for z in [poly([1]), q, mul(q, q)]] == [1, -1, 1], 'certificate or invariant failed: [query(p, z, Q(5, 4), Q(4, 3)) for z in [poly([1]), q, mul(q, q)]] == [1, -1, 1]')
    sturm = signed_chain(p, derivative(p))
    verify(sturm == [p, poly([-1, 0, 3]), poly([1, Q(2, 3)]), poly([Q(-23, 4)])], 'certificate or invariant failed: sturm == [p, poly([-1, 0, 3]), poly([1, Q(2, 3)]), poly([Q(-23, 4)])]')
    r = divrem(mul(derivative(p), q), p)[1]
    weighted_chain = signed_chain(p, r)
    verify(weighted_chain == [p, poly([2, 3, -4]), poly([Q(5, 8), Q(-1, 16)]), poly([368])], 'certificate or invariant failed: weighted_chain == [p, poly([2, 3, -4]), poly([Q(5, 8), Q(-1, 16)]), poly([368])]')
    C = companion(p)
    verify(mpow(C, 3) == madd(C, eye(3)), 'certificate or invariant failed: mpow(C, 3) == madd(C, eye(3))')
    verify([trace(mpow(C, i)) for i in range(9)] == [3, 0, 2, 3, 2, 5, 5, 7, 10], 'certificate or invariant failed: [trace(mpow(C, i)) for i in range(9)] == [3, 0, 2, 3, 2, 5, 5, 7, 10]')
    Ls = [[[1, 0, 0], [0, 1, 0], [Q(2, 3), Q(3, 2), 1]],
          [[1, 0, 0], [Q(-3, 4), 1, 0], [Q(1, 2), -10, 1]],
          [[1, 0, 0], [Q(-7, 6), 1, 0], [Q(5, 6), Q(-29, 19), 1]]]
    Ds = [[3, 2, Q(-23, 6)], [-4, Q(1, 4), -23], [6, Q(-19, 6), Q(23, 19)]]
    Hs = [hermite(p, z) for z in [poly([1]), q, mul(q, q)]]
    for H, L, D in zip(Hs, Ls, Ds):
        verify(H == mmul(mmul(L, diagonal(D)), transpose(L)), 'certificate or invariant failed: H == mmul(mmul(L, diagonal(D)), transpose(L))')
    verify([sig(H) for H in Hs] == [1, -1, 1], 'certificate or invariant failed: [sig(H) for H in Hs] == [1, -1, 1]')
    verify([sig(hermite(P, z)) for z in [poly([1]), q, mul(q, q)]] == aggregate, 'certificate or invariant failed: [sig(hermite(P, z)) for z in [poly([1]), q, mul(q, q)]] == aggregate')
    Hx = hermite(p, poly([0, 1]))
    verify(Hx[0][0] == 0 and inertia(Hx) == (2, 1, 0), 'certificate or invariant failed: Hx[0][0] == 0 and inertia(Hx) == (2, 1, 0)')
    verify(inertia([[Q(0), Q(2)], [Q(2), Q(0)]]) == (1, 1, 0), 'certificate or invariant failed: inertia([[Q(0), Q(2)], [Q(2), Q(0)]]) == (1, 1, 0)')
    verify(equal_alg(q, I[3], mul(q, poly([-5, 1])), I[3]), 'certificate or invariant failed: equal_alg(q, I[3], mul(q, poly([-5, 1])), I[3])')
    verify(not equal_alg(p, I[2], q, I[3]), 'certificate or invariant failed: not equal_alg(p, I[2], q, I[3])')
    verify(not equal_alg(q, (Q(1), Q(2)), poly([-3, 0, 1]), (Q(1), Q(2))), 'certificate or invariant failed: not equal_alg(q, (Q(1), Q(2)), poly([-3, 0, 1]), (Q(1), Q(2)))')
    verify([query(P, z, -4, 4) for z in [poly([1]), linear, mul(linear, linear)]] == [4, 1, 3], 'certificate or invariant failed: [query(P, z, -4, 4) for z in [poly([1]), linear, mul(linear, linear)]] == [4, 1, 3]')
    auto_I, fallback = isolate(P, Q(1, 16))
    verify(len(auto_I) == 4 and sum(count(P, *J) for J in auto_I) == 4, 'certificate or invariant failed: len(auto_I) == 4 and sum(count(P, *J) for J in auto_I) == 4')
    # Force a rational root at the initial midpoint: the guarded split must run.
    zero_example = mul(poly([0, 1]), poly([-2, 0, 1]))
    zero_I, zero_fallback = isolate(zero_example, Q(1, 8))
    verify(len(zero_I) == 3 and zero_fallback > 0, 'certificate or invariant failed: len(zero_I) == 3 and zero_fallback > 0')
    verify(query(p, (), -2, 2) == 0, 'certificate or invariant failed: query(p, (), -2, 2) == 0')
    common = mul(linear, q)
    verify([query(common, z, -2, 2) for z in [poly([1]), linear, mul(linear, linear)]] == [3, 0, 2], 'certificate or invariant failed: [query(common, z, -2, 2) for z in [poly([1]), linear, mul(linear, linear)]] == [3, 0, 2]')
    verify(count(poly([1, 0, 1]), -2, 2) == 0, 'certificate or invariant failed: count(poly([1, 0, 1]), -2, 2) == 0')
    verify(sig(hermite(poly([1, 0, 1]), poly([1]))) == 0, 'certificate or invariant failed: sig(hermite(poly([1, 0, 1]), poly([1]))) == 0')
    # Deliberately corrupted certificates must not match the asserted result.
    wrong_chain = sturm[:-1]+[scale(sturm[-1], -1)]
    verify(variations(wrong_chain, -2)-variations(wrong_chain, 2) != 1, 'certificate or invariant failed: variations(wrong_chain, -2)-variations(wrong_chain, 2) != 1')
    verify(count(p, Q(4, 3), Q(3, 2)) == 0, 'certificate or invariant failed: count(p, Q(4, 3), Q(3, 2)) == 0')
    try:
        count(poly([0, 1]), -1, 0)
    except ValueError:
        pass
    else:
        raise AssertionError("root endpoint was incorrectly accepted")
    try:
        isolate(())
    except ValueError:
        pass
    else:
        raise AssertionError("zero polynomial was incorrectly accepted")
    # Independently known real roots, plus an irreducible positive quadratic.
    checks = 0
    known = poly([1])
    real_roots = [-2, 1, 3]
    for t in real_roots:
        known = mul(known, poly([-t, 1]))
    known = mul(known, poly([1, 0, 1]))
    for coeff in product(range(-2, 3), repeat=3):
        test_q = poly(coeff)
        expected = sum(sign(ev(test_q, Q(t))) for t in real_roots)
        verify(query(known, test_q, -4, 4) == expected, 'certificate or invariant failed: query(known, test_q, -4, 4) == expected')
        verify(sig(hermite(known, test_q)) == expected, 'certificate or invariant failed: sig(hermite(known, test_q)) == expected')
        checks += 1
    return {
        "status": "PASS",
        "arithmetic": "Python fractions.Fraction; exact, standard library only",
        "coefficients_increasing_degree": {"F": list(F), "squarefree_part": list(P)},
        "sturm_cubic_chain": sturm,
        "weighted_cubic_chain": weighted_chain,
        "manual_isolating_intervals": I,
        "distinct_real_roots": 4,
        "real_roots_counting_multiplicity": sum(multiplicities),
        "multiplicities": multiplicities,
        "q_sign_by_interval": [0, -1, -1, 0],
        "aggregate_queries_N_U_W": aggregate,
        "positive_negative_zero_counts": [0, 2, 2],
        "cubic_hermite_matrices": Hs,
        "cubic_hermite_signatures": [sig(H) for H in Hs],
        "automatic_isolating_intervals": auto_I,
        "rational_midpoint_guard_test": {"intervals": zero_I, "fallbacks": zero_fallback},
        "known_root_query_crosschecks": checks,
        "mutation_checks": ["negative remainder sign", "empty proposed isolator", "root endpoint", "zero input"],
        "scope": "mathematical certificate check; not proof-assistant or site-build validation"
    }


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--output", type=Path)
    args = parser.parse_args()
    report = run()
    encoded = json.dumps(report, ensure_ascii=False, indent=2, default=str)+"\n"
    if args.output:
        args.output.write_text(encoded, encoding="utf-8")
    print(encoded, end="")
