#!/usr/bin/env python3
"""Exact Bartlett HAC reader; Python standard library only.

Run without arguments for the page/capstone endpoints and negative tests.
Or: python hac-inference-reader.py --input request.json
Fields: scores (n by p integers or exact numeric strings), L (integer),
centering ('none' or 'sample'), A (optional p by p mean-equation sensitivity),
contrast (optional p-vector; requires A). Fractions are serialized as strings.
Arithmetic success never certifies a statistical model, CLT or coverage.
"""
from fractions import Fraction as F
import argparse
import json
import math
from pathlib import Path


class ContractError(ValueError):
    pass


def number(x):
    if isinstance(x, bool) or not isinstance(x, (int, str, F)):
        raise ContractError('Use integers or exact decimal/fraction strings; booleans and floats are rejected.')
    try:
        return F(x)
    except (ValueError, ZeroDivisionError):
        raise ContractError('Invalid finite exact number.') from None


def matrix(x, name, rows=None, cols=None):
    if not isinstance(x, (list, tuple)) or not x:
        raise ContractError(name + ' must be a nonempty rectangular matrix.')
    if any(not isinstance(row, (list, tuple)) or not row for row in x):
        raise ContractError(name + ' has an empty or invalid row.')
    p = len(x[0])
    if any(len(row) != p for row in x):
        raise ContractError(name + ' is ragged.')
    if rows is not None and len(x) != rows or cols is not None and p != cols:
        raise ContractError(name + ' has wrong dimensions.')
    return [[number(v) for v in row] for row in x]


def zero(p, q=None):
    return [[F(0) for _ in range(p if q is None else q)] for _ in range(p)]


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


def multiply(a, b):
    return [[sum((u*v for u, v in zip(row, col)), F(0)) for col in zip(*b)] for row in a]


def inverse(a):
    n = len(a)
    aug = [a[i][:] + [F(i == j) for j in range(n)] for i in range(n)]
    for j in range(n):
        pivot = next((i for i in range(j, n) if aug[i][j]), None)
        if pivot is None:
            raise ContractError('A is singular: parameter covariance is undefined under this interface.')
        aug[j], aug[pivot] = aug[pivot], aug[j]
        v = aug[j][j]
        aug[j] = [x/v for x in aug[j]]
        for i in range(n):
            if i != j:
                v = aug[i][j]
                aug[i] = [x-v*y for x, y in zip(aug[i], aug[j])]
    return [row[n:] for row in aug]


def psd_certificate(a):
    """Exact unpivoted LDL, including zero-pivot row checks."""
    n = len(a)
    if a != transpose(a):
        return {'status': 'failed_nonsymmetric'}
    d = []
    ell = [[F(i == j) for j in range(n)] for i in range(n)]
    for j in range(n):
        pivot = a[j][j] - sum((ell[j][k]**2*d[k] for k in range(j)), F(0))
        if pivot < 0:
            return {'status': 'failed_negative_pivot', 'index': j, 'pivot': pivot}
        d.append(pivot)
        for i in range(j+1, n):
            residual = a[i][j] - sum((ell[i][k]*ell[j][k]*d[k] for k in range(j)), F(0))
            if pivot == 0:
                if residual != 0:
                    return {'status': 'failed_zero_pivot_nonzero_row', 'index': j}
                ell[i][j] = F(0)
            else:
                ell[i][j] = residual/pivot
    reconstructed = multiply(multiply(ell, [[d[i] if i == j else F(0) for j in range(n)] for i in range(n)]), transpose(ell))
    if reconstructed != a:
        raise RuntimeError('Internal LDL reconstruction failure.')
    return {'status': 'exact_psd', 'L': ell, 'D': d, 'rank': sum(x > 0 for x in d)}


def eigen_diagnostic(a):
    """Jacobi approximation; residual radius is numerical, not an interval-arithmetic proof."""
    try:
        m = [[float(x) for x in row] for row in a]
    except OverflowError:
        return {'status': 'float_range_exceeded', 'exact_psd_certificate_remains_authoritative': True}
    if any(x != 0 and y == 0 for row, floats in zip(a, m) for x, y in zip(row, floats)):
        return {'status': 'float_range_exceeded_or_underflow', 'exact_psd_certificate_remains_authoritative': True}
    n = len(m)
    scale = max(1.0, max(sum(abs(v) for v in row) for row in m))
    if not math.isfinite(scale):
        return {'status': 'float_range_exceeded', 'exact_psd_certificate_remains_authoritative': True}
    tol = 1e-10*scale
    converged, rotations = n == 1, 0
    for _ in range(max(1, 100*n*n)):
        off = [(abs(m[i][j]), i, j) for i in range(n) for j in range(i+1, n)]
        if not off:
            break
        largest, i, j = max(off)
        if largest <= tol/(max(1, n)*10):
            converged = True
            break
        theta = 0.5*math.atan2(2*m[i][j], m[j][j]-m[i][i])
        c, s = math.cos(theta), math.sin(theta)
        aii, ajj, aij = m[i][i], m[j][j], m[i][j]
        for k in range(n):
            if k != i and k != j:
                u, v = m[k][i], m[k][j]
                m[k][i] = m[i][k] = c*u-s*v
                m[k][j] = m[j][k] = s*u+c*v
        m[i][i] = c*c*aii-2*c*s*aij+s*s*ajj
        m[j][j] = s*s*aii+2*c*s*aij+c*c*ajj
        m[i][j] = m[j][i] = 0.0
        rotations += 1
        if any(not math.isfinite(v) for row in m for v in row):
            return {'status': 'float_range_exceeded', 'exact_psd_certificate_remains_authoritative': True}
    radius = max(sum(abs(m[i][j]) for j in range(n) if j != i) for i in range(n))
    estimate = min(m[i][i] for i in range(n))
    return {'status': 'diagnostic_only', 'converged_to_off_diagonal_threshold': converged,
            'jacobi_rotations': rotations, 'smallest_eigenvalue_estimate': estimate,
            'off_diagonal_residual_radius': radius, 'negative_tolerance': tol,
            'within_psd_tolerance': estimate >= -tol,
            'roundoff_is_not_rigorously_bounded': True,
            'exact_psd_certificate_remains_authoritative': True}


def hac(scores, L, centering, A=None, contrast=None):
    scores = matrix(scores, 'scores')
    n, p = len(scores), len(scores[0])
    if isinstance(L, bool) or not isinstance(L, int) or not 0 <= L < n:
        raise ContractError('L must be an integer in [0,n-1].')
    if centering not in ('none', 'sample'):
        raise ContractError("centering must be explicitly 'none' or 'sample'.")
    means = [sum((row[j] for row in scores), F(0))/n for j in range(p)]
    x = [[v-means[j] if centering == 'sample' else v for j, v in enumerate(row)] for row in scores]
    gamma = []
    for k in range(L+1):
        gamma.append([[sum((x[t][i]*x[t-k][j] for t in range(k, n)), F(0))/n for j in range(p)] for i in range(p)])
    B = [row[:] for row in gamma[0]]
    for k in range(1, L+1):
        w = 1-F(k, L+1)
        for i in range(p):
            for j in range(p):
                B[i][j] += w*(gamma[k][i][j]+gamma[k][j][i])
    windows, running, W = [], [F(0)]*p, zero(p)
    for t in range(n+L):
        for j in range(p):
            if t < n:
                running[j] += x[t][j]
            if 0 <= t-L-1 < n:
                running[j] -= x[t-L-1][j]
        windows.append(running[:])
        for i in range(p):
            for j in range(p):
                W[i][j] += running[i]*running[j]/(n*(L+1))
    if B != W:
        raise RuntimeError('Lag/window identity failed.')
    cert = psd_certificate(B)
    if cert['status'] != 'exact_psd':
        raise RuntimeError('Bartlett PSD certificate failed.')
    out = {'arithmetic_status': 'passed_exact', 'model_assumptions_status': 'not_certified_from_input',
           'required_before_inference': ['correct zero-mean score target', 'applicable dependence model and CLT',
                                         'bandwidth and fitted-score consistency', 'nonsingular limiting sensitivity',
                                         'positive limiting contrast variance'],
           'n': n, 'p': p, 'L': L, 'centering': centering, 'input_score_mean': means,
           'lag_denominator': n, 'lag_covariances': gamma, 'B': B, 'window_sums': windows,
           'window_B': W, 'window_equals_lag_exactly': True, 'PSD': cert,
           'eigen_diagnostic': eigen_diagnostic(B)}
    if A is not None:
        a = matrix(A, 'A', p, p)
        ai = inverse(a)
        C = [[v/n for v in row] for row in multiply(multiply(ai, B), transpose(ai))]
        out.update({'A_mean_equation': a, 'C_parameter': C, 'scaling': 'A_inverse B A_inverse_transpose / n'})
        if contrast is not None:
            if not isinstance(contrast, (list, tuple)) or len(contrast) != p:
                raise ContractError('contrast must have length p.')
            c = [number(v) for v in contrast]
            var = sum((c[i]*C[i][j]*c[j] for i in range(p) for j in range(p)), F(0))
            if var < 0:
                raise RuntimeError('Negative exact contrast variance.')
            try:
                as_float = float(var)
                se = math.sqrt(as_float) if as_float > 0 or var == 0 else None
                se = se if se is not None and math.isfinite(se) else None
            except OverflowError:
                se = None
            out.update({'contrast': c, 'contrast_variance': var, 'contrast_se_approx': se,
                        'contrast_se_display_status': 'available' if se is not None else 'float_range_exceeded_or_underflow',
                        'studentization_status': 'nonzero_scale_only_model_unverified' if var > 0 else 'failed_zero_scale'})
    elif contrast is not None:
        raise ContractError('A is required to propagate a contrast.')
    return out


def ma_mean_variance(n, coefficients, sigma2=1):
    if isinstance(n, bool) or not isinstance(n, int) or n < 1:
        raise ContractError('n must be a positive integer.')
    if not isinstance(coefficients, (list, tuple)) or not coefficients:
        raise ContractError('MA coefficients must be nonempty.')
    a, v = [number(z) for z in coefficients], number(sigma2)
    if v <= 0:
        raise ContractError('innovation variance must be positive.')
    q = len(a)-1
    gamma = [v*sum((a[j]*a[j+k] for j in range(q+1-k)), F(0)) for k in range(q+1)]
    covariance_sum = (n*gamma[0]+2*sum(((n-k)*gamma[k] for k in range(1, min(q, n-1)+1)), F(0)))/(n*n)
    innovation_weights = {}
    for t in range(1, n+1):
        for j, coefficient in enumerate(a):
            innovation_weights[t-j] = innovation_weights.get(t-j, F(0))+coefficient/n
    independent_sum = v*sum((w*w for w in innovation_weights.values()), F(0))
    if covariance_sum != independent_sum:
        raise RuntimeError('Independent MA mean variance reconstructions differ.')
    omega = v*sum(a, F(0))**2
    return {'n': n, 'gamma': gamma, 'exact_by_covariance': covariance_sum,
            'exact_by_innovation_coefficients': independent_sum, 'iid_comparison': gamma[0]/n,
            'long_run_variance': omega, 'long_run_over_n': omega/n,
            'positive_lrv': omega > 0,
            'model': 'fixed coefficients and independent Gaussian innovations are assumptions, not inferred'}


def check(ok, message):
    if not ok:
        raise RuntimeError(message)


def demo():
    scalar = hac([[1], [-2], [1]], 1, 'none', [[-1]], [1])
    check(scalar['B'] == [[F(2, 3)]], 'Scalar endpoint mismatch.')
    raw = scalar['lag_covariances'][0][0][0]+2*scalar['lag_covariances'][1][0][0]
    check(raw == F(-2, 3) and scalar['contrast_variance'] == F(2, 9), 'Raw/scaling mismatch.')
    matrix_case = hac([[1, 2], [-2, 0], [1, -1], [0, -1]], 2, 'sample', [[-2, 0], [0, -1]], [1, 1])
    shifted = hac([[8, 5], [9, 3], [11, 1]], 1, 'sample', [[-2, 0], [0, -1]], [1, 1])
    centered = hac([[1, 3], [2, 1], [4, -1]], 1, 'sample', [[-2, 0], [0, -1]], [1, 1])
    check(centered['contrast_variance'] == F(113, 324), 'Capstone contrast mismatch.')
    check(shifted['B'] == centered['B'] and shifted['C_parameter'] == centered['C_parameter'] and shifted['contrast_variance'] == centered['contrast_variance'], 'Centered shift invariance failed.')
    ma = [ma_mean_variance(n, [1, '1/2']) for n in [1, 4, 10, 100]]
    for result, expected in zip(ma, [F(5, 4), F(1, 2), F(43, 200), F(14, 625)]):
        check(result['exact_by_covariance'] == expected, 'MA endpoint mismatch.')
    tiny = hac([['1/'+'1'+'0'*200], ['-1/'+'1'+'0'*200]], 0, 'none', [[-1]], [1])
    check(tiny['contrast_variance'] > 0 and tiny['contrast_se_approx'] is None, 'Positive variance underflow missed.')
    negative_tests = []
    invalid = [('empty', [], 0, 'none', None, None), ('ragged', [[1], [2, 3]], 0, 'none', None, None),
               ('float', [[1.0]], 0, 'none', None, None), ('nan', [['NaN']], 0, 'none', None, None),
               ('bool', [[True]], 0, 'none', None, None), ('large_L', [[1]], 1, 'none', None, None),
               ('bool_L', [[1]], True, 'none', None, None), ('negative_L', [[1]], -1, 'none', None, None),
               ('centering', [[1]], 0, 'guess', None, None), ('singular_A', [[1]], 0, 'none', [[0]], [1]),
               ('bad_A_shape', [[1]], 0, 'none', [[1, 2]], None), ('contrast_without_A', [[1]], 0, 'none', None, [1])]
    for label, x, lag, center, a, c in invalid:
        try:
            hac(x, lag, center, a, c)
        except ContractError as exc:
            negative_tests.append({'case': label, 'status': 'rejected', 'message': str(exc)})
        else:
            raise RuntimeError('Failed to reject '+label)
    check(psd_certificate([[F(1), F(2)], [F(2), F(1)]])['status'] == 'failed_negative_pivot', 'Indefinite matrix missed.')
    check(psd_certificate([[F(0), F(1)], [F(1), F(0)]])['status'] == 'failed_zero_pivot_nonzero_row', 'Zero-pivot failure missed.')
    degenerate = hac([[2], [2]], 1, 'sample', [[-1]], [1])
    check(degenerate['studentization_status'] == 'failed_zero_scale', 'Zero scale not diagnosed.')
    denom = hac([[1], [-2], [2], [-1]], 1, 'none')
    wrong_denominator = denom['lag_covariances'][0][0][0]+F(4, 3)*denom['lag_covariances'][1][0][0]
    check(wrong_denominator == -F(1, 6), 'Denominator failure mismatch.')
    return {'scalar': scalar, 'raw_rectangular': raw, 'matrix_transfer': matrix_case, 'centering_transfer': centered,
            'finite_MA': ma, 'fixed_L1_population_limit': F(7, 4),
            'negative_correlation_MA': ma_mean_variance(4, [1, '-1/2']),
            'denominator_failure': {'common_n_B': denom['B'], 'mixed_n_minus_k_B': wrong_denominator},
            'degenerate_sample': degenerate, 'underflow_test': {'status': tiny['contrast_se_display_status'], 'exact_variance_positive': tiny['contrast_variance'] > 0}, 'rejected_inputs': negative_tests,
            'self_checks': 'passed; checks remain active under python -O'}


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


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--input', type=Path)
    args = parser.parse_args()
    try:
        if args.input:
            request = json.loads(args.input.read_text())
            if not isinstance(request, dict):
                raise ContractError('Input must be a JSON object.')
            allowed = {'scores', 'L', 'centering', 'A', 'contrast'}
            if set(request)-allowed or not {'scores', 'L', 'centering'} <= set(request):
                raise ContractError('Required fields: scores, L, centering; optional fields: A, contrast.')
            result = hac(**request)
        else:
            result = demo()
    except (ContractError, OSError, json.JSONDecodeError) as exc:
        print(json.dumps({'arithmetic_status': 'failed_input_contract', 'message': str(exc),
                          'model_assumptions_status': 'not_certified'}, ensure_ascii=False, indent=2))
        return 2
    print(json.dumps(serialize(result), ensure_ascii=False, indent=2))
    return 0


if __name__ == '__main__':
    raise SystemExit(main())
