#!/usr/bin/env python3
"""Exact finite-moment certificates; standard library only.

All arithmetic is Fraction arithmetic. Inputs are exact rational data. Root
candidates are supplied and certified, not found by a floating root solver.
The symmetric irrational example is checked through rho**2 = 3/5, never by
rounding sqrt(3/5). Checks use explicit exceptions and remain active under -O.
"""
from fractions import Fraction as F
from itertools import combinations, product
import argparse
import json
from pathlib import Path

CHECKS = 0

def check(ok, message):
    global CHECKS
    CHECKS += 1
    if not ok:
        raise RuntimeError(message)

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

def rational(x):
    need(isinstance(x, (int, F)) and not isinstance(x, bool), 'exact rational input required')
    return F(x)

def dot(a, b):
    need(len(a) == len(b), 'dot dimensions')
    return sum((x*y for x, y in zip(a, b)), F(0))

def mv(A, b):
    return [dot(row, b) for row in A]

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

def mm(A, B):
    return [[dot(row, col) for col in transpose(B)] for row in A]

def poly_value(p, x):
    s = F(0)
    for a in reversed(p):
        s = s*x+a
    return s

def mul(p, q):
    out = [F(0)]*(len(p)+len(q)-1)
    for i, a in enumerate(p):
        for j, b in enumerate(q):
            out[i+j] += a*b
    return out

def node_polynomial(nodes):
    p = [F(1)]
    for z in nodes:
        p = mul(p, [-z, F(1)])
    return p

def moments(nodes, weights, top):
    need(len(nodes) == len(weights), 'node/weight dimensions')
    return [sum((w*z**k for z, w in zip(nodes, weights)), F(0)) for k in range(top+1)]

def hankel(m, d):
    need(d >= 0 and len(m) >= 2*d+1, 'Hankel dimensions')
    return [[m[i+j] for j in range(d+1)] for i in range(d+1)]

def solve(A, b):
    n = len(A)
    need(n > 0 and len(b) == n and all(len(row) == n for row in A), 'square solve dimensions')
    B = [[rational(x) for x in row]+[rational(v)] for row, v in zip(A, b)]
    for k in range(n):
        pivot = next((i for i in range(k, n) if B[i][k]), None)
        need(pivot is not None, 'singular system')
        B[k], B[pivot] = B[pivot], B[k]
        t = B[k][k]
        B[k] = [x/t for x in B[k]]
        for i in range(n):
            if i != k:
                t = B[i][k]
                B[i] = [x-t*y for x, y in zip(B[i], B[k])]
    return [row[-1] for row in B]

def determinant(A):
    n = len(A)
    need(all(len(row) == n for row in A), 'determinant dimensions')
    B = [[rational(x) for x in row] for row in A]
    d = F(1)
    for k in range(n):
        pivot = next((i for i in range(k, n) if B[i][k]), None)
        if pivot is None:
            return F(0)
        if pivot != k:
            B[k], B[pivot] = B[pivot], B[k]
            d = -d
        t = B[k][k]
        d *= t
        for i in range(k+1, n):
            ratio = B[i][k]/t
            for j in range(k+1, n):
                B[i][j] -= ratio*B[k][j]
    return d

def positive_ldl(A):
    n = len(A)
    need(n > 0 and all(len(row) == n for row in A), 'LDL dimensions')
    need(A == transpose(A), 'symmetric matrix required')
    L = [[F(i == j) for j in range(n)] for i in range(n)]
    D = []
    for j in range(n):
        d = A[j][j]-sum((L[j][k]**2*D[k] for k in range(j)), F(0))
        need(d > 0, 'matrix not positive definite')
        D.append(d)
        for i in range(j+1, n):
            L[i][j] = (A[i][j]-sum((L[i][k]*L[j][k]*D[k] for k in range(j)), F(0)))/d
    return L, D

def prony(m, r):
    need(isinstance(r, int) and not isinstance(r, bool) and r >= 1, 'positive integer order')
    need(len(m) >= 2*r, 'need 2r consecutive samples')
    m = [rational(x) for x in m]
    H, h = hankel(m, r-1), m[r:2*r]
    c = solve(H, h)
    return H, c, [-x for x in c]+[F(1)]

def verify_prony(m, nodes, weights):
    r = len(nodes)
    need(r > 0 and len(weights) == r, 'nonempty node/weight dimensions')
    nodes, weights = [rational(x) for x in nodes], [rational(x) for x in weights]
    need(len(set(nodes)) == r and all(weights), 'distinct nodes and nonzero amplitudes')
    H, c, q = prony(m, r)
    need(node_polynomial(nodes) == q, 'complete node polynomial mismatch')
    need(moments(nodes, weights, len(m)-1) == list(m), 'supplied moments mismatch')
    need(all(m[k+r] == dot(c, m[k:k+r]) for k in range(len(m)-r)), 'recurrence mismatch')
    return dict(order=r, determinant=determinant(H), recurrence=c, polynomial=q)

def flat_status(m):
    need(len(m) >= 3 and len(m) % 2 == 1, 'moments through degree 2r, r >= 1')
    m = [rational(x) for x in m]
    r = (len(m)-1)//2
    H = hankel(m, r-1)
    L, D = positive_ldl(H)
    h = m[r:2*r]
    c = solve(H, h)
    q = [-x for x in c]+[F(1)]
    s = m[-1]-dot(h, c)
    need(dot(mul(q, q), m) == s, 'internal Schur/square disagreement')
    return dict(status='FLAT' if s == 0 else 'INFEASIBLE_POSITIVE' if s < 0 else 'NOT_FLAT',
                order=r, schur=s, recurrence=c, polynomial=q, ldl_diagonal=D)

def verify_flat(m, nodes, weights):
    info = flat_status(m)
    need(info['status'] == 'FLAT', 'not a flat certificate')
    need(len(nodes) == info['order'] and all(w > 0 for w in weights), 'r positive atoms required')
    verify_prony(m, nodes, weights)
    return info

def christoffel(m, a):
    need(len(m) > 0 and len(m) % 2 == 1, 'moments through degree 2d')
    m, a = [rational(x) for x in m], rational(a)
    d = (len(m)-1)//2
    H = hankel(m, d)
    _, D = positive_ldl(H)
    v = [a**j for j in range(d+1)]
    u = solve(H, v)
    denominator = dot(v, u)
    p = [x/denominator for x in u]
    value = 1/denominator
    need(poly_value(p, a) == 1 and dot(mul(p, p), m) == value, 'internal variational mismatch')
    return dict(degree=d, point=a, bound=value, polynomial=p, square=mul(p, p), ldl_diagonal=D)

def robust_atom_bound(m, eps, a, p):
    m, eps, p = ([rational(x) for x in seq] for seq in (m, eps, p))
    a = rational(a)
    need(len(m) == len(eps) and len(m) >= 2*len(p)-1 and len(p) > 0, 'error/moment dimensions')
    need(all(e >= 0 for e in eps), 'negative moment error')
    need(poly_value(p, a) == 1, 'target value must equal one')
    b = mul(p, p)+[F(0)]*(len(m)-(2*len(p)-1))
    return dot(b, m)+dot([abs(x) for x in b], eps)

def rejected(call):
    try:
        call()
    except ValueError:
        return True
    return False

def run():
    global CHECKS
    CHECKS = 0
    m = list(map(F, [1]))+[F(1,2), F(3,2), F(5,2), F(11,2), F(21,2), F(43,2)]
    nodes, weights = list(map(F, [-1, 0, 2])), [F(1,6), F(1,2), F(1,3)]
    main = verify_flat(m, nodes, weights)
    check(main['polynomial'] == [0, -2, -1, 1], 'main annihilator')
    check(main['ldl_diagonal'] == [1, F(5,4), F(4,5)], 'main LDL')
    check(determinant(hankel(m, 2)) == 1, 'main determinant')
    check(moments(nodes, weights, 6) == m, 'all seven main moments')
    bad_m = m[:-1]+[m[-1]-F(1,100)]
    check(flat_status(bad_m)['schur'] == -F(1,100), 'negative square rejection')
    check(rejected(lambda: verify_flat(bad_m, nodes, weights)), 'reject negative flat input')
    signed = verify_prony([F(x) for x in [0,-1,-3,-7]], [F(1),F(2)], [F(1),F(-1)])
    check(signed['polynomial'] == [2,-3,1] and signed['determinant'] == -1, 'signed cancellation')
    _, _, repeated = prony(list(map(F, [1,4,12,32])), 2)
    check(repeated == [4,-4,1], 'confluent candidate')
    check(rejected(lambda: verify_prony(list(map(F,[1,4,12,32])),[F(2),F(2)],[F(1),F(1)])), 'duplicate rejection')
    check(rejected(lambda: verify_prony(m, nodes, weights[::-1])), 'weight mutation')
    check(rejected(lambda: verify_prony(m[:-1]+[m[-1]+1], nodes, weights)), 'new sample rejection')
    check(rejected(lambda: flat_status([F(x) for x in [1,0,0,0,1]])), 'singular predecessor rejected')
    check(flat_status([F(1),F(0),F(0)])['status'] == 'FLAT', 'lower-order atom accepted')
    check(rejected(lambda: prony([F(1)]*4, 2)), 'lower-order singular Prony')
    check(rejected(lambda: christoffel([F(1),F(0),F(0)], F(1))), 'no pseudoinverse substitution')
    uniform = [F(0) if k%2 else F(1,k+1) for k in range(7)]
    cb = christoffel(uniform[:5], F(0))
    check(cb['bound'] == F(4,9) and cb['polynomial'] == [1,0,-F(5,3)], 'uniform Christoffel')
    # The pair at +/-rho uses exactly rho**2=3/5. Odd moments cancel.
    rho2, pair_weight, zero_weight = F(3,5), F(5,18), F(4,9)
    gauss = [F(1)]+[F(0) if k%2 else 2*pair_weight*rho2**(k//2) for k in range(1,7)]
    check(gauss[:6] == uniform[:6], 'six shared moments')
    check(gauss[6] == F(3,25), 'Gauss sixth moment')
    check(1-F(5,3)*rho2 == 0 and zero_weight == cb['bound'], 'attainment witness')
    check(flat_status(gauss)['status'] == 'FLAT', 'Gauss flat extension')
    check(flat_status(uniform)['schur'] == F(4,175), 'uniform nonflat extension')
    uniform_five_atoms = [F(1)]+[F(0) if k%2 else F(49,90)*F(3,7)**(k//2)+F(1,10) for k in range(1,7)]
    check(F(16,45)+F(49,90)+F(1,10)==1, 'five-atom positive mass')
    check(uniform_five_atoms==uniform, 'same-table atomic representative through degree six')
    eps = [F(0), F(1,1000), F(1,1000), F(1,1000), F(1,1000)]
    robust = robust_atom_bound(uniform[:5], eps, F(0), cb['polynomial'])
    check(robust == F(811,1800), 'robust fixed-polynomial bound')
    check(rejected(lambda: robust_atom_bound(uniform[:5], [F(-1)]+eps[1:], F(0), cb['polynomial'])), 'negative error rejection')
    check(rejected(lambda: robust_atom_bound(uniform[:5], eps, F(1), cb['polynomial'])), 'target normalization rejection')
    check(christoffel(uniform[:3], F(0))['bound'] == 1, 'unattained supremum bound')
    compact_nodes, compact_weights = list(map(F,[-1,0,1])), [F(1,6),F(2,3),F(1,6)]
    check(moments(compact_nodes, compact_weights, 2) == uniform[:3], 'compact sharp witness')
    check(moments([F(0),F(1),F(2)],compact_weights,2)==[1,1,F(4,3)], 'translated compact moments')
    check(dot([F(0),F(2),F(-1)],[F(1),F(1),F(4,3)])==F(2,3), 'translated support majorant')
    for j in range(-100,101):
        x = F(j,100)
        check(1-x*x >= (1 if x == 0 else 0), 'compact pointwise majorant sampled cross-check')
    escaping = []
    for a in range(1,65):
        e = F(1,3*a*a)
        w = [e/2,1-e,e/2]
        check(moments([F(-a),F(0),F(a)],w,2) == uniform[:3], 'escaping family moments')
        check(w[1] < 1 and w[1] >= F(2,3), 'unattained positive gap')
        if a in [1,2,4,8]:
            escaping.append(dict(outer_node=F(a), outer_weight=e/2, zero_weight=1-e))
    # Nearby four-atom measure defeats exact-rank inference from a small box.
    e = F(1,10**9)
    near_nodes, near_weights = nodes+[F(3)], [w*(1-e) for w in weights]+[e]
    near = moments(near_nodes, near_weights, 6)
    check(all(abs(a-b) <= F(1,10**6) for a,b in zip(near,m)), 'four-atom measure in error box')
    check(determinant(hankel(near,3)) == 144*(1-e)**3*e > 0, 'rank growth exact determinant')
    check(flat_status(near)['status'] == 'NOT_FLAT', 'nearby four-atom nonflat')
    q2 = mul(main['polynomial'], main['polynomial'])
    check(q2 == [0,0,4,4,-3,-2,1], 'main square coefficients')
    check(dot(q2, near) == 144*e, 'nearby square moment')
    square_budget = sum(abs(b) for b in q2)*F(1,10**6)
    check(square_budget == F(7,500000), 'square uncertainty budget')
    outside_mass = square_budget/F(1,100)
    check(outside_mass == F(7,5000), 'off-zero-set mass bound')
    # Broad exact families: infer recurrence, verify atoms, matrix factorization,
    # self-adjoint multiplication and robustness to deliberate certificate errors.
    cases = 0
    for r in range(1,5):
        for ns in combinations([-3,-1,0,1,2,4], r):
            ns = list(map(F, ns))
            ws = [F(i+1, sum(range(1,r+1))) for i in range(r)]
            ms = moments(ns,ws,2*r)
            cert = verify_flat(ms,ns,ws)
            check(cert['status'] == 'FLAT', 'finite positive family')
            V = [[z**i for z in ns] for i in range(r)]
            VW = [[V[i][j]*ws[j] for j in range(r)] for i in range(r)]
            H = hankel(ms,r-1)
            check(mm(VW,transpose(V)) == H, 'V diag(w) V^T')
            c = cert['recurrence']
            M = [[F(i == j+1) if j<r-1 else c[i] for j in range(r)] for i in range(r)]
            HM = mm(H,M)
            check(HM == transpose(HM), 'H-self-adjoint multiplication')
            check(HM == [[ms[i+j+1] for j in range(r)] for i in range(r)], 'shifted moment identity')
            for top in range(2*r,2*r+5):
                tail = moments(ns,ws,top)
                check(all(tail[k+r] == dot(c,tail[k:k+r]) for k in range(len(tail)-r)), 'unseen algebraic model tails')
            check(rejected(lambda: verify_flat(ms[:-1]+[ms[-1]+F(1,97)],ns,ws)), 'nonflat mutation')
            for a in [F(-2),F(0),F(3,2),F(5)]:
                b = christoffel(ms[:2*r-1],a)
                actual = sum((w for z,w in zip(ns,ws) if z==a),F(0))
                check(actual <= b['bound'], 'atom upper bound')
                check(dot(mul(b['polynomial'],b['polynomial']),ms[:2*r-1]) == b['bound'], 'optimal square value')
                for t in [F(-2),F(1,3),F(2)]:
                    # Any variation u with u(a)=0 preserves the constraint.
                    if r>=2:
                        p = b['polynomial'][:]
                        p[0] -= t*a
                        p[1] += t
                        check(poly_value(p,a)==1 and dot(mul(p,p),ms[:2*r-1])>=b['bound'], 'variational competitor')
            for j in range(r):
                b = christoffel(ms[:2*r-1],ns[j])
                check(b['bound']==ws[j], 'r-atom interpolation equality')
            cases += 1
    # Signed models also obey the determinant/annihilator theorem, without PSD.
    for ns in combinations([-2,0,1,3],3):
        for ws in product([-2,-1,1,2], repeat=3):
            nsf, wsf = list(map(F,ns)), list(map(F,ws))
            ms = moments(nsf,wsf,8)
            cert = verify_prony(ms,nsf,wsf)
            vandermonde = F(1)
            for i,j in combinations(range(3),2):
                vandermonde *= (nsf[j]-nsf[i])**2
            check(cert['determinant']==wsf[0]*wsf[1]*wsf[2]*vandermonde, 'signed determinant')
            check(cert['polynomial']==node_polynomial(nsf), 'signed annihilator')
    collisions = []
    for n in range(1,21):
        delta = F(1,2**n)
        ms = moments([1-delta,1+delta],[F(1,2),F(1,2)],3)
        check(determinant(hankel(ms,1))==delta**2, 'near-collision exact rank')
        check(ms==[1,1,1+delta**2,1+3*delta**2], 'near-collision exact samples')
        if n in [1,10,20]:
            collisions.append(dict(delta=delta,determinant=delta**2,samples=ms))
    # Basis invariance on the nondegenerate uniform degree-two matrix.
    H = hankel(uniform,2)
    for k in range(-7,8):
        T = [[F(1),F(k),F(k*k)],[F(0),F(1),F(k)],[F(0),F(0),F(1)]]
        G = mm(mm(T,H),transpose(T))
        for a in [F(-2),F(-1,3),F(0),F(1,2),F(3)]:
            v = [F(1),a,a*a]
            tv = mv(T,v)
            check(1/dot(tv,solve(G,tv)) == christoffel(uniform[:5],a)['bound'], 'basis invariance')
    check(rejected(lambda: solve([[F(1)]],[0.5])), 'floating input rejected')
    return dict(status='PASS', checks=CHECKS, positive_atomic_families=cases,
        main_flat=main, signed_prony=signed, confluent_polynomial=repeated,
        uniform_christoffel=cb, gaussian_moments=gauss,
        uniform_schur=flat_status(uniform)['schur'], robust_atom_bound=robust,
        compact_maximum=F(2,3), escaping_examples=escaping,
        noise_counterexample=dict(epsilon=e,moments=near,rank4_determinant=determinant(hankel(near,3)),
                                  square_value=dot(q2,near),square_budget=square_budget,outside_mass_bound=outside_mass),
        collision_examples=collisions,
        scope='Exact rational examples and finite cross-checks, not a floating root finder or proof of the general theorems')

def serial(value):
    if isinstance(value,F):
        return str(value)
    if isinstance(value,list):
        return [serial(x) for x in value]
    if isinstance(value,dict):
        return {k:serial(v) for k,v in value.items()}
    return value

if __name__ == '__main__':
    parser=argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--output',type=Path,help='write exact JSON to this path; otherwise print it')
    args=parser.parse_args()
    result=json.dumps(serial(run()),ensure_ascii=False,indent=2)+'\n'
    if args.output:
        args.output.parent.mkdir(parents=True,exist_ok=True)
        args.output.write_text(result)
        print('PASS:',CHECKS,'exact checks; wrote',args.output)
    else:
        print(result,end='')
