#!/usr/bin/env python3
"""Exact finite resampling and regression certificates (Python standard library).
All checks are explicit and also execute under python -O. Finite enumeration
checks the worked records, not the asymptotic theorems proved in the pages.
"""
import argparse,json
from pathlib import Path
from fractions import Fraction as Q
from itertools import product,combinations
from collections import Counter
from math import comb

COUNTS=Counter()
def check(ok,kind):
    COUNTS[kind]+=1
    if not ok:raise ArithmeticError(kind)
def tr(A):return [list(row) for row in zip(*A)]
def mm(A,B):return [[sum((x*y for x,y in zip(row,col)),Q(0)) for col in zip(*B)] for row in A]
def mv(A,v):return [sum((x*y for x,y in zip(row,v)),Q(0)) for row in A]
def outer(v,w):return [[x*y for y in w] for x in v]
def add(A,B):return [[x+y for x,y in zip(a,b)] for a,b in zip(A,B)]
def scale(s,A):return [[s*x for x in row] for row in A]
def dot(x,y):return sum((a*b for a,b in zip(x,y)),Q(0))
def inv(A):
    n=len(A)
    if n==0 or any(len(row)!=n for row in A):raise ValueError('nonempty square matrix required')
    W=[[Q(x) for x in row]+[Q(i==j) for j in range(n)] for i,row in enumerate(A)]
    for j in range(n):
        k=next((i for i in range(j,n) if W[i][j]),None)
        if k is None:raise ValueError('singular matrix')
        W[j],W[k]=W[k],W[j];v=W[j][j];W[j]=[x/v for x in W[j]]
        for i in range(n):
            if i!=j:
                v=W[i][j];W[i]=[x-v*y for x,y in zip(W[i],W[j])]
    return [row[n:] for row in W]
def design(rows):
    X=[[Q(x) for x in row] for row in rows];A=inv(mm(tr(X),X));D=mm(A,tr(X));H=mm(X,D)
    return X,A,D,H

def fit(X,D,y):
    beta=mv(D,y);e=[Q(v)-u for v,u in zip(y,mv(X,beta))]
    return beta,e,dot(e,e)
def hc(X,A,e,h,power):
    if power and any(x>=1 for x in h):raise ValueError('leverage correction requires h<1')
    p=len(A);mid=[[Q(0)]*p for _ in range(p)]
    for x,r,g in zip(X,e,h):mid=add(mid,scale(r*r/(1-g)**power,outer(x,x)))
    return mm(mm(A,mid),A)
def record(v):
    if isinstance(v,Q):return str(v)
    if isinstance(v,list):return [record(x) for x in v]
    if isinstance(v,dict):return {str(k):record(x) for k,x in v.items()}
    return v

def ranges():
    for n in range(2,8):
        for m in range(2,5):
            pairs=Counter((min(v),max(v)) for v in product(range(n),repeat=m))
            for i in range(n):
                for j in range(i,n):
                    num=1 if i==j else (j-i+1)**m-2*(j-i)**m+(j-i-1)**m
                    check(pairs[(i,j)]==num,'with-replacement endpoint inclusion-exclusion')
            for k in range(n):
                check(sum(v for (i,j),v in pairs.items() if j<=k)==(k+1)**m,'maximum exact CDF')
        for b in range(2,n+1):
            pairs=Counter((min(v),max(v)) for v in combinations(range(n),b))
            for i in range(n):
                for j in range(i+1,n):
                    expected=comb(j-i-1,b-2) if j-i-1>=b-2 else 0
                    check(pairs[(i,j)]==expected,'without-replacement endpoint count')
    x=[1,2,4,5,8,10];R=Q(9)
    laws=[]
    for iterator,den in [(product(x,repeat=3),216),(combinations(x,3),20)]:
        C=Counter(Q(3)*(R-(max(v)-min(v)))/R for v in iterator)
        law={k:Q(v,den) for k,v in sorted(C.items())};check(sum(law.values())==1,'range law normalization');laws.append(law)
    check(laws[0][0]==Q(5,36) and laws[1][0]==Q(1,5),'both original endpoints retained')
    def quantile(law,p):
        acc=Q(0)
        for x,m in sorted(law.items()):
            acc+=m
            if acc>=p:return x
        raise ArithmeticError('missing quantile')
    qs=[quantile(law,Q(3,4)) for law in laws];check(qs==[2,Q(4,3)],'range root quantiles')
    uppers=[R/(1-c/6) for c in qs];check(uppers==[Q(27,2),Q(81,7)],'original n interval scale')
    lo,hi=Q(3,5),Q(31,50);F=lambda z:6*z**5-5*z**6
    check(F(lo)==Q(729,3125)<Q(1,4)<F(hi)==Q(830245379,3125000000),'exact Uniform range quantile bracket')
    # Integral of n(n-1)v^(n-2)(1-v), including n=2 boundary.
    for n in range(2,20):
        check(Q(n*(n-1),n-1)-Q(n*(n-1),n)==1,'range density normalization')
        for j in range(21):
            r=Q(j,20);cdf=n*r**(n-1)-(n-1)*r**n
            integral=n*(n-1)*(r**(n-1)/(n-1)-r**n/n)
            check(cdf==integral and 0<=cdf<=1,'range density integral')
    # Exact binomial moment calculation for the shrinking endpoint-count argument.
    for n in [5,9,17]:
        for m in range(2,n+1):
            for t in [Q(1,3),Q(2,3),Q(3,2)]:
                p=t/m
                weights=[Q(comb(n-1,k))*p**k*(1-p)**(n-1-k) for k in range(n)]
                z=[Q(m*(k+1),n) for k in range(n)]
                mean=sum(w*v for w,v in zip(weights,z));var=sum(w*(v-mean)**2 for w,v in zip(weights,z))
                check(mean==Q(m,n)+Q(n-1,n)*t,'endpoint scaled mean')
                check(var==Q(m*m*(n-1),n*n)*p*(1-p)<=t*m/n,'endpoint scaled variance')
    return {'with_replacement':laws[0],'without_replacement':laws[1],'root_quantiles':qs,'upper_endpoints':uppers,'exact_quantile_bracket':[lo,hi],'exact_upper_endpoint_bracket':[R/hi,R/lo]}

def subsampling():
    # Full original Bernoulli data law reduced exactly by its sufficient count.
    for n in range(2,25):
        for b in range(1,n):
            for cutoff in range(b+1):
                oracle=[]
                for K in range(n+1):
                    count=sum(comb(K,k)*comb(n-K,b-k) for k in range(max(0,b-(n-K)),min(b,K)+1) if abs(2*k-b)<=cutoff)
                    oracle.append(Q(count,comb(n,b)))
                w=[Q(comb(n,K),2**n) for K in range(n+1)]
                mean=sum(a*v for a,v in zip(w,oracle));true=Q(sum(comb(b,k) for k in range(b+1) if abs(2*k-b)<=cutoff),2**b)
                var=sum(a*(v-mean)**2 for a,v in zip(w,oracle));q=n//b
                check(mean==true,'oracle subset CDF expectation')
                check(var<=true*(1-true)/q<=Q(1,4*q),'random-block variance budget')
    # Verify the actual sequential uniform subset sampler path probability.
    for n in range(1,9):
        for b in range(1,n+1):
            for I in combinations(range(n),b):
                left=b;prob=Q(1)
                for j in range(n):
                    remaining=n-j
                    if j in I:prob*=Q(left,remaining);left-=1
                    else:prob*=1-Q(left,remaining)
                check(left==0 and prob==Q(1,comb(n,b)),'sequential uniform subset probability')
    data=[1]*6+[-1]*2
    # Store coefficient of sqrt(3), so every support comparison is rational.
    C=Counter(abs(Q(sum(data[i] for i in I),3))-Q(1,2) for I in combinations(range(8),3))
    law={v:Q(c,56) for v,c in sorted(C.items())}
    D=Counter(abs(Q(sum(z),3))-Q(1,2) for z in product(data,repeat=3));rep={v:Q(c,512) for v,c in sorted(D.items())}
    check(law=={Q(-1,6):Q(9,14),Q(1,2):Q(5,14)},'absolute mean subset law')
    check(rep=={Q(-1,6):Q(9,16),Q(1,2):Q(7,16)},'absolute mean replacement law')
    check(rep[Q(-1,6)]<Q(3,5)<law[Q(-1,6)],'absolute mean quantile separation')
    return {'support_is_coefficient_of_sqrt3':True,'without_replacement':law,'with_replacement':rep}

def regression():
    models=[[[1],[1],[1]],[[1,x] for x in [-1,0,1,2]],[[1,x] for x in [-1,0,1,2,8]],[[1,x,x*x] for x in [-2,-1,0,1,2]],[[1,0],[1,0],[0,1]],[[1,0],[1,0],[1,0],[0,1],[0,1]],[[1,0,1],[1,1,0],[1,2,1],[1,3,0],[1,4,1],[1,5,0]]]
    for rows in models:
        X,A,D,H=design(rows);n,p=len(X),len(A);h=[H[i][i] for i in range(n)];nu=n-p
        check(mm(H,H)==H and H==tr(H),'hat matrix orthogonal projection')
        check(sum(h)==p and all(0<=v<=1 for v in h),'leverage bounds and trace')
        deletions=[]
        for i in range(n):
            xx=X[:i]+X[i+1:]
            try:Ax=inv(mm(tr(xx),xx));Dx=mm(Ax,tr(xx))
            except ValueError:check(h[i]==1,'rank failure exactly unit leverage');deletions.append(None);continue
            check(h[i]<1,'rank preserved exactly subunit leverage');deletions.append((xx,Dx))
        ys=product([-1,0,2],repeat=n) if n<=5 else product([-1,2],repeat=n)
        for y in ys:
            beta,e,rss=fit(X,D,y);check(mv(tr(X),e)==[0]*p,'OLS residual orthogonality')
            shifts=[]
            for i,item in enumerate(deletions):
                if item is None:check(e[i]==0,'unit leverage forces zero residual');continue
                xx,Dx=item;bm,em,rm=fit(xx,Dx,y[:i]+y[i+1:]);diff=[u-v for u,v in zip(beta,bm)];formula=[v*e[i]/(1-h[i]) for v in mv(A,X[i])]
                check(diff==formula,'deleted coefficient direct refit')
                check(Q(y[i])-dot(X[i],bm)==e[i]/(1-h[i]),'deleted prediction residual')
                check(rm==rss-e[i]**2/(1-h[i]),'deleted RSS direct refit')
                shifts.append(diff)
                if rss>0:
                    r2=e[i]**2/(rss/nu*(1-h[i]));check(r2<=nu,'internal residual support bound')
                    cook=dot(mv(X,diff),mv(X,diff))/(p*rss/nu)
                    check(cook==r2*h[i]/(p*(1-h[i])),'Cook full fitted-vector norm')
                    if nu>1 and rm>0:
                        t2=e[i]**2/(rm/(nu-1)*(1-h[i]));check(t2==r2*(nu-1)/(nu-r2),'external studentization algebra')
            if all(g<1 for g in h):
                hc3=hc(X,A,e,h,2);sumsh=[[Q(0)]*p for _ in range(p)]
                for d in shifts:sumsh=add(sumsh,outer(d,d))
                check(hc3==sumsh,'HC3 uncentered deletion outer products')
        # Homoskedastic HC2 expectation, from residual row variance not simulations.
        if all(g<1 for g in h):
            mid=[[Q(0)]*p for _ in range(p)]
            for x,g in zip(X,h):mid=add(mid,scale((1-g)/(1-g),outer(x,x)))
            check(mm(mm(A,mid),A)==A,'homoskedastic HC2 expectation')
    X,A,D,H=design([[1,x] for x in [-1,0,1,2,8]]);y=[1,0,1,0,7];beta,e,rss=fit(X,D,y);h=[H[i][i] for i in range(5)]
    Xm=X[:-1];bm,em,rm=fit(Xm,mm(inv(mm(tr(Xm),Xm)),tr(Xm)),y[:-1]);r2=e[-1]**2/(rss/3*(1-h[-1]));t2=e[-1]**2/(rm/2*(1-h[-1]));cook=r2*h[-1]/(2*(1-h[-1]))
    check(beta==[Q(7,25),Q(19,25)] and bm==[Q(3,5),Q(-1,5)],'high leverage terminal coefficients')
    check((rss,rm,r2,t2,cook)==(Q(148,25),Q(4,5),Q(96,37),Q(64,5),Q(552,37)),'high leverage terminal diagnostics')
    R=[[Q(1),Q(1)],[Q(0),Q(2)]];Xp,Ap,Dp,Hp=design(mm(X,R));bp,ep,rp=fit(Xp,Dp,y);Xpm=Xp[:-1];bpm,_,_=fit(Xpm,mm(inv(mm(tr(Xpm),Xpm)),tr(Xpm)),y[:-1])
    check(Hp==H and ep==e and rp==rss,'invertible feature change keeps diagnostics')
    check(bp==[Q(-1,10),Q(19,50)] and bpm==[Q(7,10),Q(-1,10)],'changed coordinate coefficients')
    d1=[u-v for u,v in zip(beta,bm)];d2=[u-v for u,v in zip(bp,bpm)];check(dot(d1,d1)==Q(640,625) and dot(d2,d2)==Q(544,625),'raw coefficient metric changes')
    # Actual heteroskedastic expectation of HC2 for X=(1,2), sigma²=(1,4).
    xx,aa,dd,hh=design([[1],[2]]);M=[[Q(i==j)-hh[i][j] for j in range(2)] for i in range(2)];vars=[sum(M[i][j]**2*[1,4][j] for j in range(2)) for i in range(2)];expected=aa[0][0]**2*sum(vars[i]*xx[i][0]**2/(1-hh[i][i]) for i in range(2));truth=aa[0][0]**2*17
    check((expected,truth)==(Q(8,25),Q(17,25)),'heteroskedastic HC2 bias counterexample')
    return {'A':A,'beta':beta,'residuals':e,'leverages':h,'RSS':rss,'deleted_beta':bm,'deleted_RSS':rm,'internal_square':r2,'external_square':t2,'Cook':cook,'changed_beta':bp,'changed_deleted_beta':bpm,'heteroskedastic_HC2_expectation':expected,'heteroskedastic_true_variance':truth}

def wild():
    cases=[([[1],[2],[3]],[1,0,2]),([[1,0]]*3+[[0,1]]*2,[1,2,3,5,9]),([[1,x] for x in [-1,0,1,2,8]],[1,0,1,0,7]),([[1,x,x*x] for x in [-2,-1,0,1,2]],[2,0,1,3,0])]
    laws=[]
    for rows,y in cases:
        X,A,D,H=design(rows);n,p=len(X),len(A);beta,e,rss=fit(X,D,y);h=[H[i][i] for i in range(n)];total=2**n;mean=[Q(0)]*p;second=[[Q(0)]*p for _ in range(p)];law=Counter()
        for signs in product([-1,1],repeat=n):
            disturb=[z*v for z,v in zip(signs,e)];delta=mv(D,disturb);star=[v+w for v,w in zip(mv(X,beta),disturb)];bs,es,_=fit(X,D,star)
            check([x-y for x,y in zip(bs,beta)]==delta,'wild direct refit equals linear update')
            mean=[u+v/total for u,v in zip(mean,delta)];second=add(second,scale(Q(1,total),outer(delta,delta)))
            if len(rows)==5 and rows[0]==[1,0]:law[delta[1]-delta[0]]+=1
        check(mean==[0]*p and second==hc(X,A,e,h,0),'exact wild conditional mean and covariance')
        if law:laws.append({k:Q(v,total) for k,v in sorted(law.items())})
    X,A,D,H=design([[1,0]]*3+[[0,1]]*2);beta,e,rss=fit(X,D,[1,2,3,5,9]);h=[H[i][i] for i in range(5)];covs=[hc(X,A,e,h,p) for p in [0,1,2]]
    check(covs==[[[Q(2,9),0],[0,2]],[[Q(1,3),0],[0,4]],[[Q(1,2),0],[0,8]]],'group HC0 HC2 HC3')
    signs=[1,-1,-1,1,-1];star=[u+v*z for u,v,z in zip(mv(X,beta),e,signs)];bs,es,_=fit(X,D,star);Vs=hc(X,A,es,h,0);contrast=bs[1]-bs[0]-(beta[1]-beta[0]);Vsc=Vs[0][0]+Vs[1][1]-2*Vs[0][1]
    check(bs==[Q(4,3),5] and es==[Q(-1,3),Q(2,3),Q(-1,3),0,0],'group selected wild record')
    check(contrast==Q(-4,3) and Vsc==Q(2,27) and contrast**2/Vsc==24,'refitted bootstrap studentization')
    check(contrast**2/Q(20,9)==Q(4,5),'frozen studentization is distinct')
    raw_pool=Q(2)*(Q(1,3)+Q(1,2));truth=1/Q(3,5)+4/Q(2,5);pooled=(Q(3,5)+4*Q(2,5))*(1/Q(3,5)+1/Q(2,5));check((raw_pool,truth,pooled,pooled/truth)==(Q(5,3),Q(35,3),Q(55,6),Q(11,14)),'heteroskedastic pooling limit')
    # Mammen multipliers in Q(sqrt(5)): pairs represent a+b sqrt(5).
    def mul5(x,y):return (x[0]*y[0]+5*x[1]*y[1],x[0]*y[1]+x[1]*y[0])
    def pow5(x,k):
        r=(Q(1),Q(0))
        for _ in range(k):r=mul5(r,x)
        return r
    support=[(Q(1,2),Q(-1,2)),(Q(1,2),Q(1,2))];prob=[(Q(1,2),Q(1,10)),(Q(1,2),Q(-1,10))]
    for k,want in enumerate([1,0,1,1,2]):
        terms=[mul5(w,pow5(x,k)) for w,x in zip(prob,support)]
        check((sum(z[0] for z in terms),sum(z[1] for z in terms))==(want,0),'Mammen exact moments')
    return {'group_contrast_law':laws[0],'group_covariances':covs,'selected_star_beta':bs,'selected_star_residuals':es,'selected_star_variance':Vsc,'selected_studentized_square':Q(24),'pooled_finite_variance':raw_pool,'true_n_variance_limit':truth,'pooled_n_variance_limit':pooled}

def main():
    COUNTS.clear()
    result={'status':'PASS','range_task':ranges(),'subsampling_task':subsampling(),'regression_task':regression(),'wild_task':wild()}
    result['checks']=sum(COUNTS.values());result['categories']=dict(COUNTS);result['scope']='Exact finite laws and algebra. Asymptotic validity and Gaussian t laws are proved in the pages, not established by these finite checks.'
    return record(result)
if __name__=='__main__':
    p=argparse.ArgumentParser();p.add_argument('--output',required=True,type=Path);args=p.parse_args();result=main();args.output.parent.mkdir(parents=True,exist_ok=True);args.output.write_text(json.dumps(result,ensure_ascii=False,indent=2)+'\n');print(json.dumps({'status':result['status'],'checks':result['checks'],'output':str(args.output)}))
