#!/usr/bin/env python3
"""Exact finite checks for Gaussian shrinkage and mixture certificates.
Standard library only; no simulation, network access, or disabled assertions.
Finite checks supplement, rather than replace, the mathematical proofs.
"""
from fractions import Fraction as F
from itertools import product, combinations
from collections import Counter
from pathlib import Path
import argparse
import json

COUNTS = Counter()

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

def check(ok, category):
    require(ok, category)
    COUNTS[category] += 1

def rejected(thunk, category):
    try:
        thunk()
    except (ValueError, ZeroDivisionError):
        COUNTS[category] += 1
        return
    raise ValueError("Invalid input accepted: " + category)

def dot(x, y):
    return sum((a*b for a, b in zip(x, y)), F(0))

def block_shrink(x, groups, variance=F(1)):
    x = tuple(map(F, x))
    variance = F(variance)
    require(variance > 0, "Positive known variance required")
    require(groups and all(groups), "Nonempty partition blocks required")
    flat = [i for block in groups for i in block]
    require(sorted(flat) == list(range(len(x))), "Blocks must partition coordinates")
    m = len(x)-len(groups)
    require(m >= 3, "At least three residual dimensions required")
    mean = [F(0)]*len(x)
    for block in groups:
        v = sum((x[i] for i in block), F(0))/len(block)
        for i in block:
            mean[i] = v
    residual = tuple(x[i]-mean[i] for i in range(len(x)))
    s = dot(residual, residual)
    factor = 1-F(m-2)*variance/s if s else F(0)
    result = tuple(mean[i]+factor*residual[i] for i in range(len(x)))
    return result, tuple(mean), residual, factor

def exp_minus_bracket(h, terms=120):
    """Bound exp(-h) by reciprocals of an explicit positive Taylor bracket."""
    h = F(h)
    require(h >= 0 and terms >= 1 and h < terms+2, "Exponential bracket contract")
    term = F(1)
    total = term
    for k in range(1, terms+1):
        term *= h/k
        total += term
    next_term = term*h/(terms+1)
    tail = next_term/(1-h/(terms+2))
    return 1/(total+tail), 1/total

def inverse_chi_square_bracket(d, lam, terms=80):
    d = int(d)
    lam = F(lam)
    require(d >= 3 and lam >= 0 and terms >= 1, "Inverse chi-square contract")
    h = lam/2
    lo, hi = exp_minus_bracket(h)
    term = F(1)
    mass_sum = term
    weighted_sum = term/(d-2)
    for k in range(1, terms+1):
        term *= h/k
        mass_sum += term
        weighted_sum += term/(d+2*k-2)
    lower = lo*weighted_sum
    upper = hi*weighted_sum + max(F(0), 1-lo*mass_sum)/(d+2*terms)
    return lower, upper

def js_risk_bracket(d, lam, variance=F(1)):
    variance = F(variance)
    require(variance > 0, "Positive variance")
    lo, hi = inverse_chi_square_bracket(d, lam)
    return variance*(d-(d-2)**2*hi), variance*(d-(d-2)**2*lo)

def fixed_decimal_bounds(interval, digits=16):
    scale = 10**digits
    lo, hi = interval
    low = (lo.numerator*scale)//lo.denominator
    high = -((-hi.numerator*scale)//hi.denominator)
    def encode(k):
        sign = "-" if k < 0 else ""
        k = abs(k)
        return sign+str(k//scale)+"."+str(k % scale).zfill(digits)
    return [encode(low), encode(high)]

def normalize(weights):
    weights = tuple(map(F, weights))
    require(weights and all(w >= 0 for w in weights) and sum(weights) > 0,
            "Nonnegative, nonzero weights")
    z = sum(weights)
    return tuple(w/z for w in weights)

def kernel_value(y, theta):
    """y and theta are integer multiples of a=sqrt(2 log 2), sigma^2=1."""
    require(isinstance(y, int) and isinstance(theta, int), "Integer a-coordinates")
    return F(1, 2**((y-theta)**2))

def posterior(prior, support, y):
    prior = normalize(prior)
    require(len(prior) == len(support), "Posterior dimensions")
    raw = tuple(w*kernel_value(y, theta) for w, theta in zip(prior, support))
    p = normalize(raw)
    mean = sum(w*theta for w, theta in zip(p, support))
    second = sum(w*theta*theta for w, theta in zip(p, support))
    variance = second-mean*mean
    return p, mean, variance

def kernel_rows(observations, support):
    require(observations and support, "Nonempty data and grid")
    return tuple(tuple(kernel_value(y, t) for t in support) for y in observations)

def fit_values(K, w):
    w = tuple(map(F, w))
    require(K and w and sum(w) == 1 and all(q >= 0 for q in w), "Probability weights")
    require(all(len(row) == len(w) and all(a > 0 for a in row) for row in K),
            "Positive rectangular kernel")
    return tuple(dot(row, w) for row in K)

def certificate(K, w):
    f = fit_values(K, w)
    D = tuple(sum(row[j]/v for row, v in zip(K, f)) for j in range(len(w)))
    return f, D

def em_step(K, w):
    _, D = certificate(K, w)
    return tuple(F(q)*d/len(K) for q, d in zip(w, D))

def likelihood(values):
    result = F(1)
    for x in values:
        result *= x
    return result

def weights_on_grid(total, size=3):
    return [tuple(F(k, total) for k in row)
            for row in product(range(total+1), repeat=size) if sum(row) == total]

def run():
    COUNTS.clear()
    # A projection computed from explicit block means; verify its geometry exactly.
    configurations = [(6, ((0,1,2),(3,4,5))),
                      (5, ((0,1),(2,3,4))),
                      (7, ((0,1,2,3,4,5,6),)),
                      (8, ((0,), (1,2,3,4), (5,6,7)))]
    for d, groups in configurations:
        for seed in range(45):
            x = tuple(F(((seed+3)*(i+2)+i*i)%13-6, 2) for i in range(d))
            out, mean, residual, factor = block_shrink(x, groups)
            check(dot(mean, residual) == 0, "projection_orthogonality")
            check(all(sum(residual[i] for i in block) == 0 for block in groups),
                  "residual_zero_group_sums")
            for scale in [F(1,3), F(2), F(3)]:
                scaled = block_shrink([scale*v for v in x], groups, scale*scale)[0]
                check(scaled == tuple(scale*v for v in out), "joint_scale_equivariance")
            shift = [F((b+1)*3) for b in range(len(groups))]
            translated = list(x)
            target = list(out)
            for b, block in enumerate(groups):
                for i in block:
                    translated[i] += shift[b]
                    target[i] += shift[b]
            check(block_shrink(translated, groups)[0] == tuple(target),
                  "preserved_subspace_translation")
    terminal1 = block_shrink([3,1,2,-1,1,0], ((0,1,2),(3,4,5)))
    check(terminal1[0] == (F(5,2),F(3,2),F(2),F(-1,2),F(1,2),F(0)),
          "terminal_two_group_output")
    check(terminal1[3] == F(1,2), "terminal_residual_dimension_coefficient")
    check(F(4)**2-2*4*(4-2) == 0, "wrong_full_dimension_has_zero_risk_gain")

    # Certified inverse-moment and risk brackets, including exact lambda=0.
    risk_rows = []
    lambdas = [F(0),F(1,2),F(1),F(2),F(4),F(9),F(16),F(25),F(36)]
    for d in range(3,9):
        for lam in lambdas:
            jl, ju = inverse_chi_square_bracket(d, lam)
            check(0 < jl <= ju <= F(1,d-2)+F(1,10**25),
                  "positive_inverse_moment_bracket")
            rl, ru = js_risk_bracket(d, lam)
            check(0 <= rl <= ru < d, "strict_joint_risk_dominance")
            check(ru-rl < F(1,10**20), "risk_bracket_width")
            if lam == 0:
                check(jl == ju == F(1,d-2) and rl == ru == 2,
                      "central_risk_exactly_two")
            for ratio in [F(1,4),F(1,2),F(1),F(3,2),F(2),F(3)]:
                c = ratio*(d-2)
                coefficient = c*c-2*c*(d-2)
                check((coefficient < 0) == (ratio < 2)
                      and (coefficient == 0) == (ratio == 2),
                      "family_risk_sign_boundary")
            if d == 5:
                risk_rows.append({"lambda":str(lam),"risk_sigma_squared_units":
                                  fixed_decimal_bounds((rl,ru))})
    # A separate closed-form recurrence for even d from integration by parts.
    for lam in lambdas[1:]:
        h = lam/2
        el, eu = exp_minus_bracket(h)
        lo, hi = (1-eu)/(2*h), (1-el)/(2*h)
        for d in [4,6,8]:
            sl, su = inverse_chi_square_bracket(d, lam)
            check(max(lo,sl) <= min(hi,su), "even_dimension_integral_recurrence")
            lo, hi = (1-(d-2)*hi)/(2*h), (1-(d-2)*lo)/(2*h)

    # Posterior moments use exact Gaussian kernel ratios at integer a-coordinates.
    grid_weights = weights_on_grid(6)
    for support in combinations(range(-2,3),3):
        for weights in grid_weights:
            for y in range(-3,4):
                p, mean, variance = posterior(weights,support,y)
                raw = tuple(w*kernel_value(y,t) for w,t in zip(weights,support))
                mass = sum(raw)
                first_derivative_coefficient = sum(w*(t-y) for w,t in zip(raw,support))/mass
                second_polynomial_coefficient = sum(w*(t-y)**2 for w,t in zip(raw,support))/mass
                pair_variance = sum(pi*pj*(ti-tj)**2 for pi,ti in zip(p,support)
                                    for pj,tj in zip(p,support))/2
                check(first_derivative_coefficient == mean-y, "tweedie_first_derivative_identity")
                check(second_polynomial_coefficient-first_derivative_coefficient**2 == variance,
                      "tweedie_second_derivative_identity")
                check(variance == pair_variance and variance >= 0, "posterior_pairwise_variance")
                check(min(support) <= mean <= max(support), "posterior_mean_support")
    p0,m0,v0 = posterior([1,2,1],[-1,0,1],0)
    pa,ma,va = posterior([1,2,1],[-1,0,1],1)
    check((p0,m0,v0)==((F(1,6),F(2,3),F(1,6)),0,F(1,3)),
          "terminal_three_point_zero")
    check((pa,ma,va)==((F(1,33),F(16,33),F(16,33)),F(5,11),F(112,363)),
          "terminal_three_point_positive")
    check(1-F(1)/F(1,4) == -3, "too_narrow_density_negative_variance")

    # Exact likelihood monotonicity and directional certificates on finite grids.
    cases = 0
    for observations in [(0,),(-1,1),(-1,-1,1),(-2,0,2),(0,0,1,2),(-2,-1,1,2)]:
        for support in [(-1,0,1),(-2,0,2),(-2,-1,1)]:
            K = kernel_rows(observations,support)
            for w in grid_weights:
                f,D = certificate(K,w);new = em_step(K,w)
                check(sum(new)==1 and all(q>=0 for q in new), "em_probability_simplex")
                check(all(q or new[j]==0 for j,q in enumerate(w)), "em_preserves_zero_weights")
                check(likelihood(fit_values(K,new)) >= likelihood(f), "em_exact_likelihood_monotonicity")
                check(sum(q*d for q,d in zip(w,D)) == len(K), "certificate_weighted_equality")
                for other in grid_weights[::7]:
                    otherf = fit_values(K,other)
                    ratios = [a/b for a,b in zip(otherf,f)]
                    average = sum(ratios)/len(K)
                    check(likelihood(ratios) <= average**len(K), "likelihood_tangent_amgm_check")
                    check(sum(ratios) <= max(D), "directional_upper_certificate")
                    if max(D) <= len(K):
                        check(likelihood(otherf) <= likelihood(f), "certified_grid_global_comparison")
                cases += 1
    K = kernel_rows((-1,1),(-1,0,1))
    f,D = certificate(K,(F(1,2),0,F(1,2)))
    check(f==(F(17,32),F(17,32)) and D==(2,F(32,17),2), "symmetric_grid_optimum")
    _,badD = certificate(K,(0,1,0))
    check(em_step(K,(0,1,0))==(0,1,0) and max(badD)==F(17,8),
          "em_fixed_point_not_global")
    check(F(20,17)**4<2, "one_extra_midpoint_still_passes")
    # The derivative at the right support point is negative in units of a.
    slope = sum(F(y-1)*kernel_value(y,1)/value for y,value in zip((-1,1),f))
    check(slope==F(-4,17), "continuous_domain_derivative_rejects_grid")
    K3 = kernel_rows((-1,-1,1),(-1,0,1));w3=(F(31,45),0,F(14,45))
    f3,D3 = certificate(K3,w3)
    check(f3==(F(17,24),F(17,24),F(17,48)) and D3==(3,F(48,17),3),
          "terminal_asymmetric_grid_certificate")
    posterior3 = normalize((w3[0]*F(1,16),w3[2]))
    check(posterior3==(F(31,255),F(224,255)), "terminal_grid_posterior")
    slope3 = sum(F(y-1)*kernel_value(y,1)/value for y,value in zip((-1,-1,1),f3))
    check(slope3==F(-6,17), "terminal_continuous_domain_violation")

    # The full refit rule is a linear map: independent exact risk and SURE algebra.
    A=((F(3,4),F(1,4)),(F(1,4),F(3,4)))
    covariance=tuple(tuple(dot(A[i],A[j]) for j in range(2)) for i in range(2))
    check(covariance==((F(5,8),F(3,8)),(F(3,8),F(5,8))), "refit_sampling_covariance")
    for theta in product(range(-5,6),repeat=2):
        bias=[dot(row,theta)-theta[i] for i,row in enumerate(A)]
        risk=sum(covariance[i][i] for i in range(2))+dot(bias,bias)
        target=F(5,4)+F((theta[0]-theta[1])**2,8)
        correct_sure=1+F((theta[0]-theta[1])**2+2,8)
        wrong_sure=F((theta[0]-theta[1])**2+2,8)
        check(risk==target==correct_sure,"refit_risk_equals_correct_sure")
        check(risk-wrong_sure==1,"frozen_derivative_underestimates_risk_by_one")
        check((risk<2)==((theta[0]-theta[1])**2<6),"refit_no_uniform_dominance")

    rejected(lambda:block_shrink([1,2,3],((0,1,2),)),"residual_dimension_rejected")
    rejected(lambda:block_shrink([1]*6,((0,1,2),(2,3,4))),"invalid_partition_rejected")
    rejected(lambda:block_shrink([1]*6,((0,1,2),(3,4,5)),0),"zero_variance_rejected")
    rejected(lambda:inverse_chi_square_bracket(2,0),"singular_inverse_moment_rejected")
    rejected(lambda:normalize([0,0]),"zero_prior_mass_rejected")
    rejected(lambda:fit_values(K,[F(1,2),0,0]),"unnormalized_mixing_weights_rejected")
    return {"status":"PASS","checks":sum(COUNTS.values()),
            "groups":dict(sorted(COUNTS.items())),
            "exact_em_candidates":cases,"certified_five_dimensional_risks":risk_rows,
            "terminal1_output":[str(x) for x in terminal1[0]],
            "terminal2_positive_posterior":[str(x) for x in pa],
            "terminal2_mean_in_a_units":str(ma),"terminal2_variance_in_a_squared_units":str(va),
            "terminal3_weights":[str(x) for x in w3],"terminal3_certificate":[str(x) for x in D3],
            "terminal4_covariance":[[str(x) for x in row] for row in covariance],
            "scope":"Exact finite rational kernel, posterior, projection and likelihood checks; certified inverse-moment series bounds; not an arbitrary-prior risk theorem or a continuous-domain optimizer."}

if __name__ == "__main__":
    parser=argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--output",type=Path,required=True)
    args=parser.parse_args()
    result=run()
    args.output.write_text(json.dumps(result,ensure_ascii=False,indent=2)+"\n")
    print(json.dumps({"status":result["status"],"checks":result["checks"]}))
