#!/usr/bin/env python3
"""Exact finite checks for moment fitting and inversion. Standard library only.

The finite identities and enumerations below do not prove asymptotic coverage.
Run with --output PATH; importing this file does not write files.
"""
from fractions import Fraction as F
from itertools import product
from collections import defaultdict
import argparse
import json
from pathlib import Path

CHECKS = 0

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


def mat(rows):
    return [[F(v) for v in row] for row in rows]


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


def tr(a):
    return [list(x) for x in zip(*a)]


def mul(a, b):
    if len(a[0]) != len(b):
        raise ValueError('incompatible matrix shapes')
    return [[sum((x*y for x, y in zip(row, col)), F(0)) for col in tr(b)] for row in a]


def add(a, b, sign=1):
    return [[x+sign*y for x, y in zip(r, s)] for r, s in zip(a, b)]


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


def inv(a):
    n = len(a)
    if any(len(row) != n for row in a):
        raise ValueError('inverse requires a square matrix')
    b = [list(row)+unit for row, unit in zip(a, eye(n))]
    for k in range(n):
        pivot = next((i for i in range(k, n) if b[i][k]), None)
        if pivot is None:
            raise ValueError('singular matrix')
        b[k], b[pivot] = b[pivot], b[k]
        v = b[k][k]
        b[k] = [x/v for x in b[k]]
        for i in range(n):
            if i != k:
                v = b[i][k]
                b[i] = [x-v*y for x, y in zip(b[i], b[k])]
    return [r[n:] for r in b]


def col(values):
    return [[F(x)] for x in values]


def join(*cols):
    return [sum((c[i] for c in cols), []) for i in range(len(cols[0]))]


def diag(values):
    return [[F(x) if i == j else F(0) for j in range(len(values))] for i, x in enumerate(values)]


def dot(x, y):
    return mul(tr(x), y)[0][0]


def quadratic(x, a):
    return dot(x, mul(a, x))


def projection(z):
    return mul(mul(z, inv(mul(tr(z), z))), tr(z))


def gmm(d, w, y, sigma):
    a = mul(mul(inv(mul(mul(tr(d), w), d)), tr(d)), w)
    beta = mul(a, y)
    residual = add(y, mul(d, beta), -1)
    cov = mul(mul(a, sigma), tr(a))
    return beta, residual, cov, a


def psd2(a):
    return a == tr(a) and a[0][0] >= 0 and a[1][1] >= 0 and a[0][0]*a[1][1] >= a[0][1]**2


def gmm_checks():
    d = mat([[1, 0], [0, 1], [1, 1]])
    sigma = diag([1, 4, 1])
    y = col([1, 2, 4])
    c = mat([[1, 1, 0], [0, 1, 1], [0, 0, 1]])
    b, r, v, a = gmm(d, inv(sigma), y, sigma)
    bi, ri, vi, ai = gmm(d, eye(3), y, sigma)
    check(b == col([F(7, 6), F(8, 3)]), 'efficient coefficients')
    check(bi == col([F(4, 3), F(7, 3)]), 'identity coefficients')
    check(quadratic(r, inv(sigma)) == F(1, 6), 'efficient objective')
    check(vi == mat([[1, -1], [-1, 2]]), 'identity covariance')
    check(v == mat([[F(5, 6), -F(2, 3)], [-F(2, 3), F(4, 3)]]), 'efficient covariance')
    check(add(vi, v, -1) == mat([[F(1, 6), -F(1, 3)], [-F(1, 3), F(2, 3)]]), 'PSD difference')
    dp, sp, yp = mul(c, d), mul(mul(c, sigma), tr(c)), mul(c, y)
    check(gmm(dp, inv(sp), yp, sp)[0] == b, 'transformed efficient fit')
    check(gmm(dp, eye(3), yp, sp)[0] == col([1, F(5, 2)]), 'untransformed identity is another problem')
    blind = mul(d, col([1, -1]))
    check(gmm(d, inv(sigma), add(y, blind), sigma)[1] == r, 'column-space violation invisible to residual')
    orthogonal = col([-1, -4, 1])
    check(mul(mul(tr(d), inv(sigma)), orthogonal) == col([0, 0]), 'orthogonal violation')
    check(quadratic(orthogonal, inv(sigma)) == 6, 'violation cost')
    # Independent finite matrix families, including correlated covariance and weights.
    cases = 0
    for h, j, k in product(range(-2, 3), range(-1, 2), range(1, 4)):
        c = mat([[1, h, j], [0, 1, -h], [0, 0, 1]])
        s = mul(mul(c, diag([k, k+1, k+3])), tr(c))
        l = mat([[1, 0, 0], [j, 1, 0], [h, -j, 1]])
        w = mul(l, tr(l))
        dw = mat([[1, 0], [0, 1], [h, k]])
        yy = col([h-1, j+2, k+3])
        bb, rr, vv, aa = gmm(dw, w, yy, s)
        bs, rs, vs, astar = gmm(dw, inv(s), yy, s)
        check(mul(aa, dw) == eye(2), 'GMM left inverse')
        check(mul(mul(tr(dw), w), rr) == col([0, 0]), 'weighted normal equations')
        delta = add(aa, astar, -1)
        check(add(vv, vs, -1) == mul(mul(delta, s), tr(delta)), 'efficiency identity')
        check(psd2(add(vv, vs, -1)), 'efficiency all contrasts')
        residual_map = add(eye(3), mul(dw, astar), -1)
        check(mul(residual_map, residual_map) == residual_map, 'weighted residual idempotence')
        check(mul(mul( astar, s), tr(residual_map)) == mat([[0]*3]*2), 'estimate and efficient residual uncorrelated')
        check(sum(residual_map[i][i] for i in range(3)) == 1, 'one residual dimension')
        cc = mat([[1, j, h], [0, 1, k], [0, 0, 1]])
        wt = mul(mul(tr(inv(cc)), w), inv(cc))
        transformed = gmm(mul(cc, dw), wt, mul(cc, yy), mul(mul(cc, s), tr(cc)))
        check(transformed[0] == bb, 'arbitrary weight coordinate invariance')
        check(transformed[2] == vv, 'covariance coordinate invariance')
        check(quadratic(transformed[1], wt) == quadratic(rr, w), 'objective coordinate invariance')
        cases += 1
    return {'efficient_beta': b, 'identity_beta': bi, 'efficient_covariance_n': v,
            'identity_covariance_n': vi, 'efficient_objective': F(1, 6),
            'transformed_covariance': sp, 'transformed_weight': inv(sp), 'matrix_families': cases}


def iv_fit(x, z, y):
    p = projection(z)
    xh = mul(p, x)
    bread_inv = inv(mul(tr(xh), xh))
    beta = mul(mul(bread_inv, tr(xh)), y)
    u = add(y, mul(x, beta), -1)
    stage = add(y, mul(xh, beta), -1)
    def cov(residual):
        meat = mul(mul(tr(xh), diag([v[0]**2 for v in residual])), xh)
        return mul(mul(bread_inv, meat), bread_inv)
    return beta, u, stage, cov(u), cov(stage), p


def iv_checks():
    x = mat([[1, 0], [0, -1], [0, 0], [1, 1]])
    z = mat([[F(1, 2), 0, -1], [F(1, 2), 0, 0], [F(1, 2), 1, 0], [F(1, 2), 1, 1]])
    y = col([F(3, 2), -F(5, 2), -F(3, 2), F(9, 2)])
    beta, u, stage, cv, bad, p = iv_fit(x, z, y)
    check(beta == col([1, 2]), '2SLS coefficient')
    check(u == col([F(1, 2), -F(1, 2), -F(3, 2), F(3, 2)]), 'structural residual')
    check(stage == col([2, -2, -3, 3]), 'second-stage residual')
    check(cv == mat([[F(5, 4), 1], [1, F(5, 4)]]), 'structural HC0')
    check(bad == mat([[F(13, 2), F(5, 2)], [F(5, 2), F(13, 2)]]), 'wrong-stage HC0')
    check(mul(mul(inv(mul(tr(x), x)), tr(x)), y) == col([F(5, 3), F(8, 3)]), 'OLS differs')
    z12 = [[r[0], r[1]] for r in z]
    z13 = [[r[0], r[2]] for r in z]
    check(iv_fit(x, z12, y)[:4] == (beta, u, stage, cv), 'delete unused fitted direction')
    changed = iv_fit(x, z13, y)
    check(changed[0] == col([1, 3]), 'delete a different instrument changes fit')
    check(changed[1] == col([F(1, 2), F(1, 2), -F(3, 2), F(1, 2)]), 'changed residual')
    for a, b, c in product(range(-2, 3), repeat=3):
        transform = mat([[1, a, b], [0, 1, c], [0, 0, 1]])
        fit = iv_fit(x, mul(z, transform), y)
        check(fit == (beta, u, stage, cv, bad, p), 'all outputs instrument basis invariance')
        shift = col([a, b])
        shifted = iv_fit(x, z, add(y, mul(x, shift)))
        check(shifted[0] == add(beta, shift), 'structural coefficient shift')
        check(shifted[1] == u and shifted[3] == cv, 'structural residual shift invariance')
        # Exact GMM equivalence using averages, to catch an extra factor of n.
        qzx, qzz, qzy = scale(mul(tr(z), x), F(1, 4)), scale(mul(tr(z), z), F(1, 4)), scale(mul(tr(z), y), F(1, 4))
        om = scale(mul(mul(tr(z), diag([v[0]**2 for v in u])), z), F(1, 4))
        gb, _, gv, _ = gmm(qzx, inv(qzz), qzy, om)
        check(gb == beta and scale(gv, F(1, 4)) == cv, 'GMM average and covariance scale')
    return {'beta': beta, 'structural_residual': u, 'second_stage_residual': stage,
            'HC0': cv, 'wrong_stage_HC0': bad, 'delete_columns_1_3_beta': changed[0]}


def log_bounds(x, terms=24):
    """Rigorous rational bounds from log(x)=2 sum t^(2j+1)/(2j+1)."""
    x = F(x)
    if x <= 0:
        raise ValueError('log requires x>0')
    if x < 1:
        lo, hi = log_bounds(1/x, terms)
        return -hi, -lo
    t = (x-1)/(x+1)
    val = 2*sum((t**(2*j+1)/F(2*j+1) for j in range(terms)), F(0))
    rem = 2*t**(2*terms+1)/((2*terms+1)*(1-t*t))
    return val, val+rem


def el_special(values):
    """Exact empirical likelihood at zero for observations in {-1,0,2}."""
    n = len(values)
    if not n:
        raise ValueError('empty sample')
    if all(v == 0 for v in values):
        return F(1), [F(1, n)]*n, F(0)
    a, c = values.count(-1), values.count(2)
    if not a or not c:
        return F(0), None, None
    lam = F(2*c-a, 2*(a+c))
    weights = [1/(n*(1+lam*v)) for v in values]
    ratio = F(1)
    for w in weights:
        ratio *= n*w
    return ratio, weights, lam


def h(values, theta, lam):
    return sum(((v-theta)/(1+lam*(v-theta)) for v in values), F(0))


def el_bracket(values, theta, steps=60):
    """Exact rational bisection; finite pole-approach initial bracketing."""
    values = list(map(F, values)); theta = F(theta)
    z = [v-theta for v in values]
    if not min(z) < 0 < max(z):
        raise ValueError('strict interior target required')
    hz = sum(z)
    if hz == 0:
        return F(0), F(0)
    if hz > 0:
        lo, pole, hi = F(0), -1/min(z), -1/(2*min(z))
        while h(values, theta, hi) > 0:
            hi = (hi+pole)/2
    else:
        hi, pole, lo = F(0), -1/max(z), -1/(2*max(z))
        while h(values, theta, lo) < 0:
            lo = (lo+pole)/2
    for _ in range(steps):
        mid = (lo+hi)/2
        sign = h(values, theta, mid)
        if sign == 0:
            return mid, mid
        if sign > 0:
            lo = mid
        else:
            hi = mid
    return lo, hi


def empirical_checks():
    # All ordered tables retain their actual probability, including the atom at the target.
    law = {-1: F(1, 2), 0: F(1, 4), 2: F(1, 4)}
    coverage, distribution = F(0), defaultdict(F)
    tables = 0
    for n in range(1, 8):
        for values0 in product((-1, 0, 2), repeat=n):
            values = list(values0)
            ratio, weights, lam = el_special(values)
            check(0 <= ratio <= 1, 'empirical likelihood ratio range')
            if weights is not None:
                check(sum(weights) == 1, 'EL normalization')
                check(sum(v*w for v, w in zip(values, weights)) == 0, 'EL mean constraint')
                check(all(w > 0 for w in weights), 'EL strictly positive weights')
                if min(values) < 0 < max(values):
                    check(h(values, F(0), lam) == 0, 'exact multiplier root')
            else:
                check(not min(values) < 0 < max(values), 'support failure only without two signs')
            if n == 3:
                prob = F(1)
                for v in values:
                    prob *= law[v]
                distribution[str(ratio)] += prob
                if ratio:
                    coverage += prob
            tables += 1
    check(coverage == F(31, 64), 'three-point exact coverage')
    check(dict(distribution) == {'0': F(33, 64), '1': F(13, 64), '8/9': F(12, 64), '1/2': F(6, 64)}, 'three-point likelihood law')
    check(2*log_bounds(F(27, 16))[1] < F(9, 8), 'binary finite log value bound')
    check(2*log_bounds(2)[1] < F(7, 5), 'three-point finite log value bound')
    # A rigorous sufficient certificate that chi-square(1)'s 95th quantile exceeds 7/5:
    # Integrate 1 - x^2/2 + x^4/8 as an upper bound for exp(-x^2/2).
    # Over [0,sqrt(7/5)] the bound gives sqrt(14/(5*pi))*(1-7/30+49/1000).
    # pi>3 then certifies the squared upper probability is below (19/20)^2.
    bound_square = F(14, 15)*(1-F(7, 30)+F(49, 1000))**2
    check(bound_square < F(19, 20)**2, 'finite log values below chi-square 95% quantile')
    binary_coverage = F(0)
    for sample in product((0, 1), repeat=4):
        k = sum(sample)
        if 0 < k < 4:
            ratio = (F(2, k))**k*(F(2, 4-k))**(4-k)
            check(ratio in (1, F(16, 27)), 'binary profile ratio')
            binary_coverage += F(1, 16)
    check(binary_coverage == F(7, 8), 'binary exact coverage')
    # General targets, not just the closed-form zero-mean examples.
    brackets = []
    for values in ([-3, -1, 2, 5], [-1, 0, 2], [-2, -2, 1, 4, 4]):
        for j in range(1, 10):
            theta = F(min(values)) + F(j, 10)*(max(values)-min(values))
            lo, hi = el_bracket(values, theta)
            check(lo <= hi and hi-lo < F(1, 10**15), 'EL bracket width')
            check(h(values, theta, lo) >= 0 and h(values, theta, hi) <= 0, 'EL bracket signs')
            check(all(1+t*(v-theta) > 0 for v in values for t in (lo, hi)), 'EL bracket inside poles')
            brackets.append({'x': values, 'theta': theta, 'lambda': [lo, hi]})
    check(el_bracket([-1, 0, 2], 0) == (F(1, 4), F(1, 4)), 'hand example exact root')
    return {'three_point_n3_coverage': coverage, 'three_point_R_distribution': dict(distribution),
            'binary_n4_coverage': binary_coverage, 'ordered_tables': tables,
            'log_2_bounds': log_bounds(2), 'general_target_brackets': brackets}


def ar_coefficients(x, y, z, critical):
    n, q = len(x), len(z[0])
    p = projection(z); m = add(eye(n), p, -1)
    a = add(p, scale(m, F(critical)*q/(n-q)), -1)
    return quadratic(x, a), -2*dot(x, mul(a, y)), quadratic(y, a), p


def poly_at(coeff, x):
    a, b, c = coeff[:3]
    return a*x*x+b*x+c


def ar_checks():
    u = scale(mat([[1, -1, -1, 1], [1, -1, 1, -1], [1, 1, -1, -1], [1, 1, 1, 1]]), F(1, 2))
    e = [[[u[i][j]] for i in range(4)] for j in range(4)]
    check(mul(tr(u), u) == eye(4), 'orthonormal rational basis')
    z = join(e[0], e[1]); x = add(scale(e[0], F(1, 2)), e[2]); y = add(scale(e[0], 3), e[3])
    coeff = ar_coefficients(x, y, z, 4)
    check(coeff[:3] == (-F(15, 4), -3, 5), 'AR quadratic coefficients')
    brackets = [(F(-1623, 1000), F(-1622, 1000)), (F(822, 1000), F(823, 1000))]
    for lo, hi in brackets:
        check(poly_at(coeff, lo)*poly_at(coeff, hi) < 0, 'root bracket opposite signs')
    for b in [F(j, 11) for j in range(-100, 101)]:
        residual = add(y, scale(x, b), -1)
        num = quadratic(residual, coeff[3]); den = quadratic(residual, add(eye(4), coeff[3], -1))
        check(num == (3-b/2)**2 and den == b*b+1, 'direct AR projected sums')
        check((num <= 4*den) == (poly_at(coeff, b) <= 0), 'AR inversion membership')
    empty = ar_coefficients(e[0], add(scale(e[1], 3), e[2]), z, 4)
    whole = ar_coefficients(e[2], e[3], z, 4)
    check(empty[:3] == (1, 0, 5), 'empty-set polynomial')
    check(whole[:3] == (-4, 0, -4), 'whole-line polynomial')
    # Include zero denominators and zero residuals. Comparison of nonnegative sums
    # directly implements the stated infinity/zero conventions.
    for coords_x in product(range(-1, 2), repeat=4):
        xx = mul(u, col(coords_x))
        yy = mul(u, col([coords_x[3]+1, coords_x[2]-1, coords_x[1], coords_x[0]]))
        cf = ar_coefficients(xx, yy, z, 4)
        for b in [F(j, 2) for j in range(-8, 9)]:
            rr = add(yy, scale(xx, b), -1)
            num = quadratic(rr, cf[3]); den = dot(rr, rr)-num
            direct = num <= 4*den
            check(direct == (poly_at(cf, b) <= 0), 'AR all sign/zero denominator branches')
            if den == 0:
                check(direct == (num == 0), 'zero denominator convention')
        for h in range(-2, 3):
            zz = mul(z, mat([[1, h], [0, 1]]))
            check(ar_coefficients(xx, yy, zz, 4)[:3] == cf[:3], 'AR instrument basis invariance')
    # New control-variable experiment; use a five-dimensional orthonormal coordinate system.
    e5 = [col([int(i == j) for i in range(5)]) for j in range(5)]
    w = e5[4]; mw = add(eye(5), projection(w), -1)
    zraw = join(add(e5[0], w), add(e5[1], w, -1))
    xx = add(add(scale(e5[0], F(1, 2)), e5[2]), scale(w, 2))
    yy = add(add(scale(e5[0], 3), e5[3]), scale(w, 7))
    zres = mul(mw, zraw); p = projection(zres); remainder = add(mw, p, -1)
    check(sum(mw[i][i] for i in range(5)) == 4, 'control residual dimension')
    check(sum(remainder[i][i] for i in range(5)) == 2, 'control-adjusted denominator degrees')
    check(mul(p, remainder) == mat([[0]*5]*5), 'orthogonal numerator/denominator after control')
    for b in [F(j, 7) for j in range(-70, 71)]:
        rr = mul(mw, add(yy, scale(xx, b), -1))
        num, den = quadratic(rr, p), quadratic(rr, remainder)
        check(num == (3-b/2)**2 and den == b*b+1, 'control-removal identity')
        check((num <= 4*den) == (poly_at(coeff, b) <= 0), 'control-adjusted same acceptance set')
    wrong_f = (quadratic(yy, p)/2)/(quadratic(yy, add(eye(5), p, -1))/3)
    check(wrong_f == F(27, 100), 'ignored control contaminates denominator')
    check(1-F(50, 59)**2 < F(3, 10), 'wrong F critical comparison certificate')
    check(F(4, 5) == F(4, 1)/(1+4), 'exact F22 80 percent critical')
    return {'ray_quadratic': coeff[:3], 'root_brackets': brackets,
            'empty_polynomial': empty[:3], 'whole_polynomial': whole[:3],
            'control_residual_dimension': 4, 'control_denominator_df': 2, 'ignored_control_F': wrong_f}


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


def main():
    global CHECKS
    CHECKS = 0
    result = {'GMM': gmm_checks(), '2SLS': iv_checks(),
              'empirical_likelihood': empirical_checks(), 'Anderson_Rubin': ar_checks()}
    result.update(status='PASS', checks=CHECKS,
                  scope='Exact finite matrix identities, rational root brackets, and full finite sampling laws; asymptotic and Gaussian laws are proved in the pages, not by enumeration.')
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--output', type=Path)
    args = parser.parse_args()
    text = json.dumps(serial(result), ensure_ascii=False, indent=2)+'\n'
    if args.output:
        args.output.parent.mkdir(parents=True, exist_ok=True)
        args.output.write_text(text)
    print(text, end='')


if __name__ == '__main__':
    main()
