"""Exact verifier for an explicit cycle with Gaussian moments.

Run with Python 3.9 or later:
    python verify_gaussian_tournaments.py

Only the Python standard library is used. All assertions and global bounds
use exact rational arithmetic. Floating point is used only to display the
final probability. This verifies the explicit example, not historical
novelty or the full quantified theorem in the accompanying research note.
"""
from fractions import Fraction as F
from math import factorial, gcd, lcm, sqrt, pi, comb, ceil


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


def add(a, b):
    out = [F(0)] * max(len(a), len(b))
    for p in (a, b):
        for j, c in enumerate(p):
            out[j] += c
    return trim(out)


def scale(a, c):
    return trim([x * c for x in a])


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


def derivative(a):
    return trim([j * a[j] for j in range(1, len(a))] or [F(0)])


def normal_moment(k, variance):
    if k % 2:
        return F(0)
    n = k // 2
    return F(factorial(2 * n), 2**n * factorial(n)) * variance**n


def normal_expectation(a, variance):
    return sum(c * normal_moment(j, variance) for j, c in enumerate(a))


def null_vector(a):
    a = [list(map(F, row)) for row in a]
    m, n = len(a), len(a[0])
    pivots = []
    r = 0
    for col in range(n):
        p = next((i for i in range(r, m) if a[i][col]), None)
        if p is None:
            continue
        a[r], a[p] = a[p], a[r]
        a[r] = [z / a[r][col] for z in a[r]]
        for i in range(m):
            if i != r:
                c = a[i][col]
                a[i] = [z - c * y for z, y in zip(a[i], a[r])]
        pivots.append(col)
        r += 1
        if r == m:
            break
    free = next(j for j in range(n) if j not in pivots)
    v = [F(0)] * n
    v[free] = F(1)
    for i, p in enumerate(pivots):
        v[p] = -a[i][free]
    return v


def physicists_hermite(n):
    a, b = [F(1)], [F(0), F(2)]
    if n == 0:
        return a
    for k in range(1, n):
        a, b = b, add(mul([F(0), F(2)], b), scale(a, -2 * k))
    return b


def gaussian_derivative_poly(a):
    return add(derivative(a), mul([F(0), F(-2)], a))


def gaussian_primitive_poly(a):
    remainder = list(a)
    out = [F(0)] * max(1, len(a)-1)
    for k in range(len(a)-1, 0, -1):
        c = -remainder[k]/2
        out[k-1] = c
        remainder[k] += 2*c
        if k >= 2:
            remainder[k-2] -= (k-1)*c
    assert all(c == 0 for c in remainder)
    assert gaussian_derivative_poly(out) == trim(a)
    return trim(out)


def negative_exp_upper(y):
    # Exact: exp(y) >= sum(y**k/k!, k=0..80), for y >= 0.
    assert y >= 0
    term = total = F(1)
    for k in range(1, 81):
        term *= y/k
        total += term
    return 1/total


def bernstein_absolute_bound(p, left, right):
    degree = len(p)-1
    power = [
        (right-left)**j * sum(p[k]*comb(k,j)*left**(k-j) for k in range(j, degree+1))
        for j in range(degree+1)
    ]
    bernstein = [sum(power[j]*F(comb(k,j),comb(degree,j)) for j in range(k+1)) for k in range(degree+1)]
    return max(abs(x) for x in bernstein)


def certified_gaussian_polynomial_bound(p):
    # sqrt(2*pi) < 3. Bounds all real x, using exact rational arithmetic.
    positive = list(p)
    negative = [c*(-1)**j for j,c in enumerate(p)]
    maximum = F(0)
    for index in range(160):
        left, right = F(index,16), F(index+1,16)
        interval_bound = max(bernstein_absolute_bound(q,left,right) for q in (positive,negative))
        maximum = max(maximum, 3*interval_bound*negative_exp_upper(left*left/2))
    assert len(p)-1 <= 100
    # For x >= 10, each x**k exp(-x*x/2) is nonincreasing if k <= 100.
    tail = 3 * sum(abs(c)*10**j for j,c in enumerate(p))*negative_exp_upper(F(50))
    return ceil(max(maximum,tail))


def ratio_second_poly(p):
    # (exp(-x*x/2) * p)'' = exp(-x*x/2) * result.
    return add(add(derivative(derivative(p)),
                   mul([F(0), F(-2)], derivative(p))),
               mul([F(-1), F(0), F(1)], p))


def main():
    # Derive the polynomial by homogeneous linear constraints on Hermites.
    basis = [physicists_hermite(n) for n in (6, 8, 10, 12)]
    constraints = [
        [p[0] for p in basis],
        [derivative(derivative(p))[0] for p in basis],
        [normal_expectation(p, F(1, 3)) for p in basis],
    ]
    p = [F(0)]
    for c, h in zip(null_vector(constraints), basis):
        p = add(p, scale(h, c))
    denominator = lcm(*(c.denominator for c in p))
    integers = [int(c*denominator) for c in p]
    divisor = gcd(*integers)
    p = [F(c//divisor) for c in integers]
    if p[-1] < 0:
        p = scale(p, -1)
    expected = list(map(F, [0, 0, 0, 0, 3465, 0, -6237, 0, 2970,
                            0, -484, 0, 24]))
    assert p == expected
    print('Derived R(x) coefficients (ascending):', [int(c) for c in p])

    # A = phi + eps*g; B = phi + eps*gprime; C = phi - eps*(g+gprime).
    # g = exp(-x*x) * R(x).
    pprime = gaussian_derivative_poly(p)
    perturbations = [p, pprime, scale(add(p, pprime), -1)]
    primitives = [gaussian_primitive_poly(q) for q in perturbations]
    for label, q, primitive in zip('ABC', perturbations, primitives):
        for k in range(6):
            # Integral x^k*q(x)*exp(-x*x) = sqrt(pi)*E[q(Z)*Z^k]
            # with Z normal of variance 1/2.
            assert normal_expectation([F(0)]*k+q, F(1, 2)) == 0
        # Primitive at zero gives the CDF perturbation at zero exactly.
        assert primitive[0] == 0
        # The density derivative at zero remains zero.
        assert gaussian_derivative_poly(q)[0] == 0
        # Integration by parts gives integral(Phi*q*exp(-x*x)) = 0.
        assert normal_expectation(primitive, F(1, 3)) == 0
        print(label, ': normalization, moments 1..5, median, stationary mode, '
                     'and zero linear comparison term verified')

    energy = normal_expectation(mul(p, p), F(1, 4))
    assert energy == F(47199651345, 262144)
    for i in range(3):
        for j in range(3):
            actual = normal_expectation(mul(perturbations[i], primitives[j]),
                                        F(1, 4))
            expected_pair = F(0) if i == j else (
                energy if (i, j) in ((0, 1), (1, 2), (2, 0)) else -energy)
            assert actual == expected_pair
    print('All nine comparison coefficients verified by exact integration')

    # Global bounds: Bernstein polynomial enclosures on 160 intervals on
    # each side of zero, rational exponential bounds, and analytic tails.
    bounds0 = [certified_gaussian_polynomial_bound(q) for q in perturbations]
    bounds2 = [certified_gaussian_polynomial_bound(ratio_second_poly(q))
               for q in perturbations]
    epsilon = F(1, 10_000_000)
    maximum0, maximum2 = max(bounds0), max(bounds2)
    assert epsilon*maximum0 <= F(1, 2)
    assert epsilon*maximum2 <= F(1, 4)
    curvature_bound = -1 + epsilon*maximum2/(1-epsilon*maximum0)
    assert curvature_bound <= -F(1, 2)
    # Positivity plus strict log-concavity makes the stationary mode unique.
    print('Certified |q/phi| bounds:', bounds0)
    print('Certified |(q/phi)second| bounds:', bounds2)
    print('Certified density ratio lower bound:', 1-epsilon*maximum0)
    print('Certified log-density curvature upper bound:', curvature_bound)
    margin_coefficient = epsilon**2*energy
    print('Exact cyclic win probability: 1/2 +', margin_coefficient,
          '* sqrt(pi/2)')
    print('Displayed cyclic win probability:',
          format(0.5+float(margin_coefficient)*sqrt(pi/2), '.16g'))
    print('PASS: every mathematical assertion above used exact arithmetic.')


if __name__ == '__main__':
    main()
