#!/usr/bin/env python3
"""Finite optimal-design certificates using exact rational arithmetic only.
No optimizer dependencies; exhaustive small basic sets and integer allocations.
Infinity means an inestimable target, never a pseudoinverse zero.
"""
from fractions import Fraction as Q
from itertools import combinations,product
from collections import Counter
from pathlib import Path
import argparse,json,random
COUNTS=Counter()
def check(ok,kind):
    COUNTS[kind]+=1
    if not ok: raise ValueError(kind)
def rational(x):
    if type(x) not in (int,Q): raise ValueError('only exact int/Fraction inputs')
    return Q(x)
def vectors(A):
    A=[list(map(rational,a)) for a in A]
    if not A or not A[0] or any(len(a)!=len(A[0]) for a in A): raise ValueError('empty/ragged candidates')
    return A

def solve(A,b):
    """Rectangular exact elimination; one solution, or None if inconsistent."""
    A=[list(map(rational,r)) for r in A];b=list(map(rational,b))
    if not A or len(A)!=len(b) or any(len(r)!=len(A[0]) for r in A): raise ValueError('linear system shape')
    M=[r+[v] for r,v in zip(A,b)];piv=[];h=0;n=len(A[0])
    for j in range(n):
        k=next((i for i in range(h,len(M)) if M[i][j]),None)
        if k is None: continue
        M[h],M[k]=M[k],M[h];c=M[h][j];M[h]=[v/c for v in M[h]]
        for i in range(len(M)):
            if i!=h:
                c=M[i][j];M[i]=[x-c*y for x,y in zip(M[i],M[h])]
        piv.append(j);h+=1
        if h==len(M): break
    if any(not any(r[:n]) and r[-1] for r in M): return None
    x=[Q(0)]*n
    for i,j in enumerate(piv): x[j]=M[i][-1]
    return x

def det(A):
    A=[list(map(rational,r)) for r in A];n=len(A)
    if not n or any(len(r)!=n for r in A): raise ValueError('det square matrix')
    z=Q(1)
    for j in range(n):
        k=next((i for i in range(j,n) if A[i][j]),None)
        if k is None:return Q(0)
        if k!=j:A[j],A[k]=A[k],A[j];z=-z
        c=A[j][j];z*=c
        for i in range(j+1,n):
            d=A[i][j]/c
            for h in range(j+1,n):A[i][h]-=d*A[j][h]
    return z

def dot(x,y):
    if len(x)!=len(y):raise ValueError('dot shape')
    return sum((a*b for a,b in zip(x,y)),Q(0))
def weights(w,m):
    w=list(map(rational,w))
    if len(w)!=m or any(v<0 for v in w) or sum(w)!=1:raise ValueError('probability weights')
    return w

def info(A,w):
    A=vectors(A);w=weights(w,len(A));p=len(A[0])
    return [[sum((v*a[i]*a[j] for a,v in zip(A,w)),Q(0)) for j in range(p)] for i in range(p)]
def sensitivity(A,w):
    A=vectors(A);M=info(A,w)
    if det(M)==0:raise ValueError('singular design has infinite G criterion')
    return [dot(a,solve(M,a)) for a in A]
def variance(A,w,c):
    A=vectors(A);w=weights(w,len(A));c=list(map(rational,c));p=len(A[0])
    if len(c)!=p:raise ValueError('target dimension')
    M=info(A,w);u=solve(M,c)
    if u is None:return None,None
    b=[v*dot(a,u) for a,v in zip(A,w)]
    return dot(c,u),b

def lp_certificate(A,c):
    """Enumerate every independent p-column primal basis and signed dual vertex."""
    A=vectors(A);p=len(A[0]);c=list(map(rational,c));m=len(A)
    if len(c)!=p or m<p:raise ValueError('target/span shape')
    prim=[];dual=[]
    for J in combinations(range(m),p):
        B=[A[j] for j in J]
        if det(B)==0:continue
        beta=solve(list(map(list,zip(*B))),c);b=[Q(0)]*m
        for j,v in zip(J,beta):b[j]=v
        prim.append((sum(map(abs,b)),b))
        for signs in product((-1,1),repeat=p):
            z=solve(B,signs)
            if all(abs(dot(a,z))<=1 for a in A):dual.append((dot(c,z),z))
    if not prim:raise ValueError('candidates must span parameter space')
    s,b=min(prim,key=lambda v:v[0]);sdual,z=max(dual,key=lambda v:v[0])
    check(s==sdual,'separate primal basis and dual vertex enumeration agree')
    check(all(sum(b[i]*A[i][j] for i in range(m))==c[j] for j in range(p)),'LP unbiasedness')
    check(all(abs(dot(a,z))<=1 for a in A),'LP all-candidate dual feasibility')
    if s:
        w=[abs(v)/s for v in b];v,best=variance(A,w,c)
        check(v==s*s,'LP weights attain squared norm even if singular')
        check(all(not b[i] or dot(A[i],z)==(1 if b[i]>0 else -1) for i in range(m)),'LP complementarity signs')
    else:w=[Q(1,m)]*m
    return {'s':s,'b':b,'z':z,'w':w,'variance':s*s}

def compositions(N,m):
    if m==1:yield (N,);return
    for n in range(N+1):
        for tail in compositions(N-n,m-1):yield (n,)+tail

def integer_designs(A,N):
    if type(N) is not int or N<1:raise ValueError('positive integer budget')
    rows=[]
    for n in compositions(N,len(A)):
        w=[Q(v,N) for v in n];M=info(A,w);D=det(M)
        rows.append({'n':list(n),'det':D,'sensitivity':sensitivity(A,w) if D else None})
    return rows

def transform(T,x):return [dot(row,x) for row in T]
def reject(fn):
    try:fn()
    except ValueError:check(True,'invalid exact contract rejected')
    else:check(False,'invalid exact contract rejected')

def run():
    COUNTS.clear()
    A=[[Q(1),Q(0)],[Q(0),Q(1)],[Q(1),Q(1)],[Q(1,2),Q(-1,2)]];p=2
    w=[Q(1,3)]*3+[Q(0)];M=info(A,w);d=sensitivity(A,w)
    check(M==[[Q(2,3),Q(1,3)],[Q(1,3),Q(2,3)]] and d==[Q(2),Q(2),Q(2),Q(3,2)],'terminal fractional certificate')
    rows=integer_designs(A,5);full=[r for r in rows if r['det']];bestD=max(r['det'] for r in full);bestG=min(max(r['sensitivity']) for r in full)
    Drows=[r for r in full if r['det']==bestD];Grows=[r for r in full if max(r['sensitivity'])==bestG]
    check(len(rows)==56,'all N5 allocations enumerated')
    check(bestD==Q(8,25) and {tuple(r['n']) for r in Drows}=={(1,2,2,0),(2,1,2,0),(2,2,1,0)},'integer D optimum')
    check(bestG==Q(13,6) and len(Grows)==1 and Grows[0]['n']==[1,1,2,1] and Grows[0]['det']==Q(3,10),'integer G optimum differs')
    check(all(max(r['sensitivity'])==Q(5,2) for r in Drows),'integer D sensitivity exceeds G optimum')
    c=[Q(2),Q(-1)];z=[Q(1),Q(-1)];cdesigns=[]
    for b in [[Q(2),Q(-1),Q(0),Q(0)],[Q(1),Q(0),Q(0),Q(2)]]:
        w=[abs(v)/3 for v in b];v,best=variance(A,w,c)
        check(v==9 and b==best and dot(c,z)==3 and all(abs(dot(a,z))<=1 for a in A),'terminal two distinct c-optimal certificates')
        cdesigns.append({'b':b,'w':w,'M':info(A,w),'variance':v})
    single=[Q(0),Q(0),Q(0),Q(1)];target=[Q(1),Q(-1)]
    check(variance(A,single,target)==(Q(4),[Q(0),Q(0),Q(0),Q(2)]),'singular estimable certificate')
    check(variance(A,single,[1,1])==(None,None),'singular inestimable is not pseudoinverse zero')
    T=[[Q(1),Q(1)],[Q(0),Q(2)]];TA=[transform(T,a) for a in A];tc=transform(T,target);tz=solve(list(map(list,zip(*T))),z)
    check(tc==[0,-2] and tz==[1,-1],'terminal transformed target and dual')
    check(variance(TA,single,tc)[0]==4 and variance(TA,single,target)[0] is None,'coordinate change versus changed target')
    terminal={'A':A,'fractional':{'w':[Q(1,3)]*3+[Q(0)],'M':M,'sensitivity':d,'det':det(M)},'N5_all_allocations':rows,'integer_D_optima':Drows,'integer_G_optima':Grows,'c_target':c,'c_designs':cdesigns,'singular':{'target':target,'M':info(A,single),'variance':4,'T':T,'transformed_candidates':TA,'transformed_target':tc,'transformed_dual':tz}}
    # All rational grids test trace budgets, exact one-point updates and efficiency bounds.
    rational_cases=0
    designs=[A,[[Q(1),Q(x)] for x in [-2,-1,0,1,2]],[[Q(1),Q(x),Q(x*x)] for x in [-1,0,1,2]],[[Q(0)],[Q(1)],[Q(3)]]]
    for X in designs:
        p=len(X[0]);grid=integer_designs(X,8);positive=[r for r in grid if r['det']]
        for row in positive:
            w=[Q(v,8) for v in row['n']];M=info(X,w);d=row['sensitivity'];D=max(d)
            check(sum(wi*di for wi,di in zip(w,d))==p and D>=p,'trace budget and maximum lower bound')
            for other in grid[::max(1,len(grid)//13)]:
                check(other['det']*p**p<=row['det']*D**p,'AMGM efficiency bound for every sampled competitor')
            for i,di in enumerate(d):
                t=Q(1,5);ww=[(1-t)*v for v in w];ww[i]+=t
                check(det(info(X,ww))==row['det']*(1-t)**(p-1)*(1+t*(di-1)),'rank-one determinant path exact identity')
                if p>1 and di>p:
                    step=(di-p)/(p*(di-1));ww=[(1-step)*v for v in w];ww[i]+=step
                    check(0<step<1 and det(info(X,ww))>row['det'],'analytic improving step')
            rational_cases+=1
    # Different signed convex hulls and all small target directions, including zero.
    families=[A,[[1,0],[0,1],[2,1],[-1,2]],[[1,-1],[1,0],[1,1]],[[1,0,0],[0,1,0],[0,0,1],[1,1,1],[1,-1,0]]]
    lp_cases=0
    for X in families:
        X=vectors(X);p=len(X[0])
        for c in product(range(-2,3),repeat=p):
            cert=lp_certificate(X,c);s=cert['s'];lp_cases+=1
            # Uniform design gives a feasible estimator upper bound for every target.
            w=[Q(1,len(X))]*len(X);v,b=variance(X,w,c)
            check(v is not None and v>=s*s,'full rank candidate variance upper bound')
            check(sum((bi*bi/wi for bi,wi in zip(b,w)),Q(0))==v,'normal equation variance equals direct estimator sum')
            for scale in [-3,2]:
                vv,bb=variance(X,w,[scale*ci for ci in c]);check(vv==scale*scale*v,'target scaling square')
    # Explicit body counterexamples.
    X=[[1,-1],[1,0],[1,1]]
    check(sensitivity(X,[0,Q(1,2),Q(1,2)])==[Q(10),Q(2),Q(2)],'support-only sensitivity trap')
    check(variance(X,[Q(1,2),0,Q(1,2)],[1,2])[0]==5 and variance(X,[Q(1,4),0,Q(3,4)],[1,2])[0]==4,'D versus c external target')
    check(variance(X,[0,1,0],[1,0])[0]==1 and variance(X,[0,1,0],[0,1])[0] is None,'body intercept versus slope estimability')
    for fn in [lambda:info(A,[1,0,0]),lambda:info(A,[2,-1,0,0]),lambda:info(A,[0,0,0,0]),lambda:info([[1],[1,2]],[Q(1,2)]*2),lambda:info(A,[0.25]*4),lambda:variance(A,[1,0,0,0],[1]),lambda:sensitivity(A,[1,0,0,0]),lambda:lp_certificate([[1,0],[2,0]],[0,1]),lambda:integer_designs(A,0)]:reject(fn)
    return {'schema':1,'status':'PASS','checks':sum(COUNTS.values()),'by_kind':dict(COUNTS),'rational_designs':rational_cases,'LP_target_cases':lp_cases,'terminal':terminal,'scope':'Exact finite arithmetic certificates; rational weight grids and explicitly listed target families, not a numerical optimizer or a proof by sampling.'}

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

def main():
    parser=argparse.ArgumentParser();parser.add_argument('--output',type=Path,default=Path(__file__).with_name('foundations-optimal-design-results.json'));args=parser.parse_args()
    result=run();args.output.parent.mkdir(parents=True,exist_ok=True);args.output.write_text(json.dumps(serialize(result),ensure_ascii=False,indent=2)+'\n');print(json.dumps({'status':result['status'],'checks':result['checks'],'output':str(args.output)}))
if __name__=='__main__':main()
