#!/usr/bin/env python3
"""Exact finite shape-constraint certificates. Standard library only.

The small convex solver exhausts active sets; it is a verification reference,
not a large-data optimizer. A polyhedral normal-cone certificate has a
nonnegative representation by linearly independent active normals, so trying
independent subsets of at most n constraints finds a certificate. A Jensen
polytope vertex has affinely independent support of size at most d+1; trying
all smaller sizes also handles lower-dimensional designs.
All checks survive -O. Use separate --output paths.
"""
from fractions import Fraction as F
from itertools import combinations,product
from pathlib import Path
import argparse,json,random
COUNTS={}
def need(v,msg):
    if not v:raise ValueError(msg)
def check(v,msg):
    COUNTS[msg]=COUNTS.get(msg,0)+1
    if not v:raise RuntimeError(msg)
def exact(x):
    need(isinstance(x,(int,F))and not isinstance(x,bool),'integer or Fraction required')
    return F(x)
def dot(a,b):return sum((x*y for x,y in zip(a,b)),F(0))
def solve(A,b):
    """Unique solution of a consistent rectangular system; otherwise None."""
    n=len(A[0])if A else 0;need(len(A)==len(b)and all(len(r)==n for r in A),'system dimensions')
    B=[list(map(F,r))+[F(v)]for r,v in zip(A,b)];p=[];row=0
    for j in range(n):
        k=next((i for i in range(row,len(B))if B[i][j]),None)
        if k is None:continue
        B[row],B[k]=B[k],B[row];v=B[row][j];B[row]=[x/v for x in B[row]]
        for i in range(len(B)):
            if i!=row:
                v=B[i][j];B[i]=[x-v*y for x,y in zip(B[i],B[row])]
        p.append(j);row+=1
        if row==len(B):break
    if any(not any(r[:n])and r[n]for r in B)or len(p)!=n:return None
    out=[F(0)]*n
    for i,j in enumerate(p):out[j]=B[i][n]
    return out

def grenander(x,counts):
    x=list(map(exact,x));counts=list(counts);m=len(x)
    need(m>0 and len(counts)==m and x[0]>0 and all(a<b for a,b in zip(x,x[1:])),'distinct increasing positive locations')
    need(all(isinstance(c,int)and not isinstance(c,bool)and c>0 for c in counts),'positive integer counts')
    gap=[x[i]-(x[i-1]if i else 0)for i in range(m)];p=[F(c,sum(counts))for c in counts];stack=[]
    for i in range(m):
        stack.append([i,i,gap[i],p[i]])
        while len(stack)>1 and stack[-2][3]/stack[-2][2]<stack[-1][3]/stack[-1][2]:
            r=stack.pop();l=stack.pop();stack.append([l[0],r[1],l[2]+r[2],l[3]+r[3]])
    q=[F(0)]*m
    for a,b,length,mass in stack:
        for j in range(a,b+1):q[j]=mass/length
    eta=[];v=F(0)
    for length,mass,height in zip(gap,p,q):v+=length-mass/height;eta.append(v)
    return dict(x=x,counts=counts,gaps=gap,p=p,q=q,blocks=stack,eta=eta)
def density_at(f,t):
    t=exact(t);need(t>=0,'nonnegative query')
    if t==0:return f['q'][0]
    return next((q for x,q in zip(f['x'],f['q'])if t<=x),F(0))
def density_cdf(f,t):
    t=exact(t)
    return sum(q*max(F(0),min(t,x)-(x-d))for x,d,q in zip(f['x'],f['gaps'],f['q']))
def likelihood(q,counts):
    v=F(1)
    for a,c in zip(q,counts):v*=a**c
    return v
def partition_oracle(f):
    m=len(f['x']);best=None;winners=[]
    for mask in range(1<<max(0,m-1)):
        ends=[j+1 for j in range(m-1)if mask>>j&1]+[m];q=[];a=0
        for b in ends:
            v=sum(f['p'][a:b])/sum(f['gaps'][a:b]);q.extend([v]*(b-a));a=b
        if any(q[j]<q[j+1]for j in range(m-1)):continue
        val=likelihood(q,f['counts'])
        if best is None or val>best:best,winners=val,[q]
        elif val==best:winners.append(q)
    check(all(q==winners[0]for q in winners),'partition maximizer heights unique')
    return winners[0]

def projection(y,w,C):
    """Minimize half weighted squared error subject to C theta <= 0."""
    y=list(map(exact,y));w=list(map(exact,w));n=len(y)
    need(n>0 and len(w)==n and all(v>0 for v in w),'positive weights and dimensions')
    C=[list(map(exact,r))for r in C];need(all(len(r)==n for r in C),'constraint dimensions')
    for size in range(min(n,len(C))+1):
        for active in combinations(range(len(C)),size):
            A=[C[j]for j in active]
            if size:
                G=[[sum(a[k]*b[k]/w[k]for k in range(n))for b in A]for a in A]
                lam=solve(G,[dot(a,y)for a in A])
                if lam is None or any(v<0 for v in lam):continue
            else:lam=[]
            theta=[y[k]-sum(l*a[k]for l,a in zip(lam,A))/w[k]for k in range(n)]
            if any(dot(c,theta)>0 for c in C):continue
            multipliers=[F(0)]*len(C)
            for j,v in zip(active,lam):multipliers[j]=v
            return dict(theta=theta,multipliers=multipliers,constraints=C,loss=sum(wi*(a-b)**2 for wi,a,b in zip(w,theta,y))/2)
    raise RuntimeError('no active-set certificate found')
def convex_1d(x,y,w):
    x=list(map(exact,x));need(len(x)==len(y)==len(w)>0 and all(a<b for a,b in zip(x,x[1:])),'increasing one-dimensional design')
    n=len(x);C=[]
    for i in range(n-2):
        h=x[i+1]-x[i];k=x[i+2]-x[i+1];r=[F(0)]*n;r[i:i+3]=[-k,h+k,-h];C.append(r)
    return projection(y,w,C)
def jensen_vertices(x,target):
    m=len(x);d=len(target);out=set()
    for k in range(1,min(m,d+1)+1):
        for ids in combinations(range(m),k):
            A=[[F(1)]*k]+[[x[j][r]for j in ids]for r in range(d)];v=solve(A,[F(1)]+list(target))
            if v is None or any(a<0 for a in v):continue
            alpha=[F(0)]*m
            for j,a in zip(ids,v):alpha[j]=a
            out.add(tuple(alpha))
    return sorted(out)
def convex_fit(x,y,w):
    x=[tuple(map(exact,v))for v in x];need(len(x)==len(y)==len(w)>0,'design sizes');d=len(x[0])
    need(d>0 and all(len(v)==d for v in x)and len(set(x))==len(x),'distinct rectangular design')
    C=set()
    for i,t in enumerate(x):
        for a in jensen_vertices(x,t):
            r=[-v for v in a];r[i]+=1
            if any(r):C.add(tuple(r))
    return projection(y,w,sorted(C))
def aggregate(x,y,w):
    """Collapse repeated inputs; return the exact half-squared-loss constant."""
    x=[tuple(map(exact,v))for v in x];y=list(map(exact,y));w=list(map(exact,w))
    need(len(x)==len(y)==len(w)>0 and all(v>0 for v in w),'aggregation sizes and weights')
    d=len(x[0]);need(d>0 and all(len(v)==d for v in x),'aggregation design dimensions')
    buckets={}
    for p,a,b in zip(x,y,w):buckets.setdefault(p,[]).append((a,b))
    points=sorted(buckets);means=[];weights=[];constant=F(0)
    for p in points:
        rows=buckets[p];weight=sum(b for a,b in rows);mean=sum(a*b for a,b in rows)/weight
        means.append(mean);weights.append(weight);constant+=sum(b*(a-mean)**2 for a,b in rows)/2
    return points,means,weights,constant
def upper_envelope(x,theta,t):
    combos=jensen_vertices(x,t)
    return min((dot(a,theta)for a in combos),default=None)
def extension(x,theta,zeta,t):return max(a+dot(z,[u-v for u,v in zip(t,p)])for p,a,z in zip(x,theta,zeta))
def check_projection(y,w,c):
    t=c['theta'];C=c['constraints'];lam=c['multipliers'];n=len(t)
    check(all(dot(a,t)<=0 for a in C),'projection feasible')
    check(all(l>=0 and l*dot(a,t)==0 for l,a in zip(lam,C)),'projection complementarity')
    check(all(w[i]*(t[i]-y[i])+sum(l*a[i]for l,a in zip(lam,C))==0 for i in range(n)),'projection stationarity')
def check_pair_certificate(x,y,w,theta,zeta,lam):
    n=len(x);d=len(x[0]);g=[[theta[i]+dot(zeta[i],[x[j][r]-x[i][r]for r in range(d)])-theta[j]for j in range(n)]for i in range(n)]
    check(all(g[i][j]<=0 and lam[i][j]>=0 and g[i][j]*lam[i][j]==0 for i in range(n)for j in range(n)if i!=j),'pair feasibility and complementarity')
    check(all(w[i]*(theta[i]-y[i])+sum(lam[i][j]-lam[j][i]for j in range(n))==0 for i in range(n)),'pair height stationarity')
    check(all(sum(lam[i][j]*(x[j][r]-x[i][r])for j in range(n))==0 for i in range(n)for r in range(d)),'pair slope stationarity')
def rejected(fn):
    try:fn()
    except ValueError:return True
    return False

def run():
    COUNTS.clear();rng=random.Random(20261020)
    main=grenander([1,2,4,5],[1,3,1,3]);check(main['q']==[F(1,4)]*2+[F(1,6)]*2 and main['eta']==[F(1,2),0,F(5,4),0],'terminal density')
    changed=grenander([1,2,4,5],[1,1,1,5]);check(changed['q']==[F(1,5)]*4 and changed['eta']==[F(3,8),F(3,4),F(17,8),0],'changed counts merge all')
    density_cases=0
    data=[]
    for m in range(1,5):
        for z in product(list(product([1,2],[1,2,3])),repeat=m):data.append(([a for a,b in z],[b for a,b in z]))
    for _ in range(180):
        m=rng.randrange(5,8);data.append(([rng.randrange(1,6)for _ in range(m)],[rng.randrange(1,5)for _ in range(m)]))
    for gaps,counts in data:
        x=[];v=0
        for h in gaps:v+=h;x.append(v)
        f=grenander(x,counts);q=f['q'];p=f['p'];eta=f['eta'];m=len(x)
        check(q==partition_oracle(f),'PAVA versus all partitions');check(dot(q,gaps)==1 and all(v>0 for v in q)and all(a>=b for a,b in zip(q,q[1:])),'density area and order')
        check(eta[-1]==0 and all(v>=0 for v in eta)and all(eta[j]*(q[j]-q[j+1])==0 for j in range(m-1)),'likelihood multipliers')
        for j in range(m):check(density_at(f,x[j])==q[j]and density_cdf(f,x[j])>=sum(p[:j+1]),'left version and CDF majorant')
        check(density_at(f,x[-1]+1)==0 and density_cdf(f,x[-1]+1)==1,'zero tail')
        scaled=grenander([3*t for t in x],counts);check(scaled['q']==[v/3 for v in q]and scaled['blocks']==[[a,b,3*length,mass]for a,b,length,mass in f['blocks']],'positive scaling')
        uniform=[F(1,x[-1])]*m;check(likelihood(uniform,counts)<=likelihood(q,counts),'uniform likelihood bound');density_cases+=1
    first=convex_1d([0,1,3],[0,2,0],[1,2,1]);check(first['theta']==[F(24,19),F(20,19),F(12,19)]and first['loss']==F(36,19)and first['multipliers']==[F(12,19)],'terminal nonuniform regression')
    check(F(2)>first['loss'],'naive slope pooling loses')
    check(dot(first['constraints'][0],[0,2,0])>0,'reversed-constraint candidate refused')
    check([w*(t-y) for w,t,y in zip([1,2,1],first['theta'],[0,2,0])]!=[0,0,0],'zeroed multiplier fails stationarity')
    pairlam=[[F(0)]*3 for _ in range(3)];pairlam[1][0]=F(24,19);pairlam[1][2]=F(12,19)
    check_pair_certificate([(0,),(1,),(3,)],[0,2,0],[1,2,1],first['theta'],[(F(-4,19),)]*3,pairlam)
    ax,ay,aw,constant=aggregate([(0,),(0,),(1,),(1,),(3,)],[0,2,1,3,0],[1,3,2,2,1])
    aggregated=convex_fit(ax,ay,aw)
    fitted={p:t for p,t in zip(ax,aggregated['theta'])}
    direct=sum(F(w)*(fitted[(x,)]-y)**2 for x,y,w in zip([0,0,1,1,3],[0,2,1,3,0],[1,3,2,2,1]))/2
    check(ay==[F(3,2),F(2),F(0)]and aw==[4,4,1]and constant==F(7,2),'repeated input aggregation')
    check(direct==aggregated['loss']+constant,'aggregation preserves objective up to constant')
    single=convex_fit([(2,3)],[F(7,3)],[4]);check(single['theta']==[F(7,3)]and single['loss']==0,'singleton exact interpolation')
    X=[(F(0),F(0)),(F(1),F(0)),(F(0),F(1)),(F(1),F(1)),(F(1,2),F(1,2))];Y=[0,0,0,0,2];W=[1]*5
    square=convex_fit(X,Y,W);check(square['theta']==[F(2,5)]*5 and square['loss']==F(8,5),'terminal two-dimensional fit')
    lam=[[F(0)]*5 for _ in range(5)]
    for j in range(4):lam[4][j]=F(2,5)
    check_pair_certificate(X,Y,W,square['theta'],[(0,0)]*5,lam)
    zeta=[(-1,0),(0,-1),(0,1),(1,0)];corners=X[:4]
    check_pair_certificate(corners,[0]*4,[1]*4,[0]*4,zeta,[[0]*4 for _ in range(4)])
    check(extension(corners,[0]*4,zeta,(F(1,2),F(1,2)))==F(-1,2)and upper_envelope(corners,[0]*4,(F(1,2),F(1,2)))==0,'distinct in-hull extensions')
    check(upper_envelope(corners,[0]*4,(2,2))is None,'outside-hull refusal')
    one_cases=0;general_cases=0
    for m in range(1,7):
        for _ in range(65):
            x=[];a=F(-2)
            for i in range(m):a+=F(rng.randrange(1,5),rng.randrange(1,4));x.append(a)
            y=[F(rng.randrange(-4,5))for _ in range(m)];w=[F(rng.randrange(1,4))for _ in range(m)];c=convex_1d(x,y,w);check_projection(y,w,c)
            t=c['theta'];check(sum(wi*(a-b)for wi,a,b in zip(w,t,y))==0 and sum(wi*(a-b)*xi for wi,a,b,xi in zip(w,t,y,x))==0,'affine residual moments')
            yp=[v+F(rng.randrange(-2,3))for v in y];cp=convex_1d(x,yp,w);dt=[a-b for a,b in zip(t,cp['theta'])];dy=[a-b for a,b in zip(y,yp)]
            check(sum(wi*a*a for wi,a in zip(w,dt))<=sum(wi*a*b for wi,a,b in zip(w,dt,dy)),'firm nonexpansion')
            shifted=convex_1d(x,[v+2+3*xi for v,xi in zip(y,x)],w);check(shifted['theta']==[v+2+3*xi for v,xi in zip(t,x)],'affine response equivariance');one_cases+=1
            if m<=4:
                general=convex_fit([(v,)for v in x],y,w);check(general['theta']==t and general['loss']==c['loss'],'all Jensen versus adjacent slopes');general_cases+=1
    designs=[X,[(0,0),(1,0),(2,0),(F(1,2),0)],[(0,0),(2,0),(0,2),(F(1,2),F(1,2))],[(0,0),(1,0),(0,1)]]
    for design in designs:
        for _ in range(35):
            y=[F(rng.randrange(-3,4))for _ in design];w=[F(rng.randrange(1,4))for _ in design];c=convex_fit(design,y,w);check_projection(y,w,c)
            check(all(upper_envelope(design,c['theta'],x)==v for x,v in zip(design,c['theta'])),'Jensen interpolation values');general_cases+=1
    for fn in [lambda:grenander([0,1],[1,1]),lambda:grenander([1,1],[1,1]),lambda:grenander([1],[0]),lambda:grenander([1.0],[1]),lambda:grenander([1],[True]),lambda:convex_1d([0,0],[1,2],[1,1]),lambda:convex_1d([0],[1],[0]),lambda:convex_1d([0],[1],[-1]),lambda:convex_fit([(0,),(0,)],[1,2],[1,1]),lambda:convex_fit([(0,),(1,2)],[1,2],[1,1]),lambda:convex_fit([(0,)],[1.0],[1])]:check(rejected(fn),'invalid input refused')
    return dict(status='PASS',checks=sum(COUNTS.values()),groups=COUNTS,density_cases=density_cases,one_dimensional_cases=one_cases,general_Jensen_cases=general_cases,terminal_density=main,changed_density=changed,terminal_nonuniform=first,terminal_square=square,terminal_pair_multipliers=lam,scope='Exact finite data and certificates; exponential small-instance reference, not a generalization guarantee or fast production optimizer.')
def serial(x):
    if isinstance(x,F):return str(x)
    if isinstance(x,(tuple,list)):return [serial(v)for v in x]
    if isinstance(x,dict):return {k:serial(v)for k,v in x.items()}
    return x
if __name__=='__main__':
    ap=argparse.ArgumentParser(description=__doc__);ap.add_argument('--output',type=Path,required=True);a=ap.parse_args();out=run();a.output.parent.mkdir(parents=True,exist_ok=True);a.output.write_text(json.dumps(serial(out),ensure_ascii=False,indent=2)+'\n');print('PASS',out['checks'],'checks;',a.output)
