#!/usr/bin/env python3
"""Exact finite experiment and convex-order certificates; standard library only.
All checks are explicit exceptions and remain active with Python -O.
The binary optimizer enumerates primal vertices and a finite family of valid
risk witnesses; equality is checked, never inferred from optimizer termination.
"""
from fractions import Fraction as F
from itertools import product, combinations
from collections import Counter
from pathlib import Path
import argparse,json
C=Counter()
def require(v,label):
 if not v:raise ValueError(label)
def check(v,label):require(v,label);C[label]+=1
def kernel(P):
 require(P and P[0],'Nonempty kernel')
 require(all(len(r)==len(P[0])and all(v>=0 for v in r)and sum(r)==1 for r in P),'Invalid stochastic row')
 return P
def matmul(P,K):
 kernel(P);kernel(K);require(len(P[0])==len(K),'Matrix interface mismatch')
 return [[sum(p*k for p,k in zip(row,col))for col in zip(*K)]for row in P]
def tv(p,q):
 require(len(p)==len(q)and all(x>=0 for x in p+q)and sum(p)==sum(q)==1,'TV requires probability rows')
 return sum(abs(x-y)for x,y in zip(p,q))/2
def bayes(P,pi,L):
 kernel(P);require(len(pi)==len(P)==len(L)and min(pi)>=0 and sum(pi)==1,'Bayes prior/model')
 require(L and L[0]and all(len(r)==len(L[0])and all(0<=v<=1 for v in r)for r in L),'Bounded loss required')
 return sum(min(sum(pi[t]*P[t][x]*L[t][a]for t in range(len(P)))for a in range(len(L[0])))for x in range(len(P[0])))
def risk(P,D,L):
 A=matmul(P,D);return [sum(a*l for a,l in zip(r,loss))for r,loss in zip(A,L)]
def dual_value(P,Q,pi,u):
 require(len(P)==len(Q)==len(pi)==len(u)and min(pi)>=0 and sum(pi)==1,'Dual dimensions/prior')
 require(all(len(row)==len(Q[0])and all(0<=v<=1 for v in row)for row in u),'Dual rewards outside [0,1]')
 target=sum(pi[t]*sum(q*v for q,v in zip(Q[t],u[t]))for t in range(len(P)))
 best=sum(max(sum(pi[t]*P[t][x]*u[t][y]for t in range(len(P)))for y in range(len(Q[0])))for x in range(len(P[0])))
 return target-best

def det3(a):
 return a[0][0]*(a[1][1]*a[2][2]-a[1][2]*a[2][1])-a[0][1]*(a[1][0]*a[2][2]-a[1][2]*a[2][0])+a[0][2]*(a[1][0]*a[2][1]-a[1][1]*a[2][0])
def solve3(A,b):
 d=det3(A)
 if d==0:return None
 return tuple(det3([[b[i]if j==k else A[i][j]for j in range(3)]for i in range(3)])/d for k in range(3))
def binary_primal(P,Q):
 kernel(P);kernel(Q);require(len(P)==len(Q)==len(P[0])==len(Q[0])==2,'Binary experiment required')
 rows=[([-F(1),F(0),F(0)],F(0)),([F(1),F(0),F(0)],F(1)),([F(0),-F(1),F(0)],F(0)),([F(0),F(1),F(0)],F(1)),([F(0),F(0),-F(1)],F(0)),([F(0),F(0),F(1)],F(1))]
 for p,q in zip(P,Q):
  rows.extend([([p[0],p[1],-F(1)],q[1]),([-p[0],-p[1],-F(1)],-q[1])])
 best=None
 for active in combinations(rows,3):
  v=solve3([a for a,b in active],[b for a,b in active])
  if v is not None and all(sum(a*x for a,x in zip(row,v))<=b for row,b in rows):
   key=(v[2],v[0],v[1])
   if best is None or key<best:best=key
 require(best is not None,'No primal vertex found')
 value,r0,r1=best;K=[[1-r0,r0],[1-r1,r1]]
 check(max(tv(a,b)for a,b in zip(matmul(P,K),Q))==value,'Exact binary primal residual')
 return value,K

def binary_dual(P,Q):
 best=None
 # Rewards with entries 0/1 suffice for these tested cases, verified by equality.
 # This search is a certificate finder, not a theorem about arbitrary dimensions.
 for z in product([F(0),F(1)],repeat=4):
  u=[list(z[:2]),list(z[2:])];points={F(0),F(1)}
  for x in range(2):
   a=P[0][x]*(u[0][0]-u[0][1]);b=P[1][x]*(u[1][0]-u[1][1])
   if a!=b:
    pi0=-b/(a-b)
    if 0<=pi0<=1:points.add(pi0)
  for pi0 in sorted(points):
   pi=[pi0,1-pi0];v=dual_value(P,Q,pi,u)
   if best is None or v>best[0]:best=(v,pi,u)
 return best

def law(xs,a):
 require(xs and len(xs)==len(a)and len(set(xs))==len(xs)and min(a)>=0 and sum(a)==1,'Invalid finite law')
 return dict(zip(xs,a))
def call(mu,t):return sum(a*max(x-t,F(0))for x,a in mu.items())
def mean(mu):return sum(x*a for x,a in mu.items())
def convex_certificate(mu,nu):
 if mean(mu)!=mean(nu):return False,{'mean_difference':mean(nu)-mean(mu)}
 for t in sorted(set(mu)|set(nu)):
  if call(mu,t)>call(nu,t):return False,{'threshold':t,'source':call(mu,t),'target':call(nu,t)}
 return True,None

def martingale_table(xs,a,ys,b,G):
 law(xs,a);law(ys,b)
 require(len(G)==len(xs)and all(len(r)==len(ys)and min(r)>=0 for r in G),'Invalid joint table')
 require([sum(r)for r in G]==a,'Wrong source marginal')
 require([sum(r)for r in zip(*G)]==b,'Wrong target marginal')
 require(all(sum((y-x)*w for y,w in zip(ys,row))==0 for x,row in zip(xs,G)),'Wrong row barycenter')
 check(convex_certificate(dict(zip(xs,a)),dict(zip(ys,b)))[0],'Every submitted martingale has convex order')
 cost=sum(w*(y-x)**2 for x,row in zip(xs,G)for y,w in zip(ys,row))
 variance_diff=sum(y*y*v for y,v in zip(ys,b))-sum(x*x*v for x,v in zip(xs,a))
 check(cost==variance_diff,'Martingale square cost identity')
 return cost

def compositions(n,k):
 if k==1:yield(n,);return
 for j in range(n+1):
  for rest in compositions(n-j,k-1):yield(j,)+rest

def reject(fn,label):
 try:fn()
 except ValueError:C[label]+=1;return
 raise ValueError('Bad input accepted: '+label)
def encode(x):
 if isinstance(x,F):return str(x)
 if isinstance(x,dict):return {str(k):encode(v)for k,v in x.items()}
 if isinstance(x,(list,tuple)):return [encode(v)for v in x]
 return x

def run():
 C.clear();grid=[F(k,4)for k in range(5)];experiments=[[[1-p,p],[1-q,q]]for p,q in product(grid,repeat=2)];hist=Counter()
 for P,Q in product(experiments,repeat=2):
  v,K=binary_primal(P,Q);lo,pi,u=binary_dual(P,Q)
  check(lo==v,'Binary primal equals valid risk witness');hist[str(v)]+=1
  L=[[1-z for z in row]for row in u]
  check(bayes(P,pi,L)-bayes(Q,pi,L)==v,'Optimal decision risk attains binary deficiency')
  # Every one of the 16 deterministic binary loss tables and five priors.
  for l in product([F(0),F(1)],repeat=4):
   loss=[list(l[:2]),list(l[2:])]
   for pi0 in grid:
    check(bayes(P,[pi0,1-pi0],loss)-bayes(Q,[pi0,1-pi0],loss)<=v,'All tested bounded Bayes transfers')
 # Three-state irreversible compression whose one classification risk is unchanged.
 P=[[F(1,2),F(1,2),0],[0,F(1,2),F(1,2)],[F(1,2),0,F(1,2)]]
 K=[[F(1),F(0)],[F(1,3),F(2,3)],[F(0),F(1)]]
 Q=[[F(2,3),F(1,3)],[F(1,6),F(5,6)],[F(1,2),F(1,2)]]
 check(matmul(P,K)==Q and det3(P)==F(1,4),'Three-state unique forward kernel and rank witness')
 L=[[F(i!=j)for j in range(3)]for i in range(3)]
 check(bayes(P,[F(1,3)]*3,L)==bayes(Q,[F(1,3)]*3,L)==F(1,2),'Same classification risk is insufficient')
 losses=[[F(0),F(0)],[F(0),F(1)],[F(1),F(0)]]
 check(bayes(Q,[0,F(1,2),F(1,2)],losses)-bayes(P,[0,F(1,2),F(1,2)],losses)==F(1,12),'Restricted binary loss detects lost information')
 check(bayes(Q,[F(1,3)]*3,losses)-bayes(P,[F(1,3)]*3,losses)==F(1,18),'Full-support loss detects lost information')
 # Asymmetric minimax certificate.
 A=[[F(4,5),F(1,5)],[F(2,5),F(3,5)]];I=[[F(1),F(0)],[F(0),F(1)]]
 asym,Kstar=binary_primal(A,I);lower,pistar,ustar=binary_dual(A,I)
 check(asym==lower==F(1,3),'Asymmetric exact deficiency')
 explicitK=[[F(5,6),F(1,6)],[F(0),F(1)]];explicitpi=[F(1,3),F(2,3)]
 check(max(tv(a,b)for a,b in zip(matmul(A,explicitK),I))==dual_value(A,I,explicitpi,I)==F(1,3),'Asymmetric explicit primal and dual')
 check(bayes(A,[F(1,2)]*2,[[0,1],[1,0]])==F(3,10),'Fair prior misses minimax witness')
 # Binary symmetric examples including the zero denominator boundary.
 for p,q in product([F(i,20)for i in range(11)],repeat=2):
  if p>q:continue
  PP=[[1-p,p],[p,1-p]];QQ=[[1-q,q],[q,1-q]]
  k=(q-p)/(1-2*p)if p<F(1,2)else F(0)
  check(matmul(PP,[[1-k,k],[k,1-k]])==QQ,'BSC exact garbling including fair boundary')
  check(max(tv(a,b)for a,b in zip(PP,QQ))==q-p,'BSC reverse identity upper')
  check(bayes(QQ,[F(1,2)]*2,[[0,1],[1,0]])-bayes(PP,[F(1,2)]*2,[[0,1],[1,0]])==q-p,'BSC matching risk lower')
 # Coarse-to-fine posterior martingale and fine-to-coarse observation kernel.
 V=[F(1,3),F(2,3)];av=[F(1,2)]*2;U=[F(0),F(1,2),F(1)];bu=[F(1,6),F(2,3),F(1,6)]
 G=[[F(1,6),F(1,3),F(0)],[F(0),F(1,3),F(1,6)]]
 check(martingale_table(V,av,U,bu,G)==F(1,18),'Posterior martingale square cost')
 fine=[[2*m*(1-u)for m,u in zip(bu,U)],[2*m*u for m,u in zip(bu,U)]]
 coarse=[[2*m*(1-v)for m,v in zip(av,V)],[2*m*v for m,v in zip(av,V)]]
 postK=[[G[j][i]/bu[i]for j in range(2)]for i in range(3)]
 check(matmul(fine,postK)==coarse,'Posterior coupling constructs actual garbling')
 loss=[[F(3,4),F(1)],[F(3,4),F(0)]]
 check(bayes(fine,[F(1,2)]*2,loss)==F(11,24)and bayes(coarse,[F(1,2)]*2,loss)==F(1,2),'Hinge decision strict advantage')
 check(tv(fine[0],fine[1])==tv(coarse[0],coarse[1])==F(1,3),'Same binary TV need not mean experiment equivalence')
 # Convex-order checks over all equal-mean five-site denominator-four laws.
 support=list(map(F,[-2,-1,0,1,2]));laws=[law(support,[F(v,4)for v in c])for c in compositions(4,5)];equalmean=0;ordered=0
 for mu,nu in product(laws,repeat=2):
  if mean(mu)!=mean(nu):continue
  equalmean+=1;ok,witness=convex_certificate(mu,nu)
  if ok:
   ordered+=1
   check(sum(x*x*a for x,a in mu.items())<=sum(x*x*a for x,a in nu.items()),'Convex-order variance necessity')
   for t in [F(j,4)for j in range(-12,13)]:check(call(mu,t)<=call(nu,t),'Knot certificate agrees with intervening thresholds')
   # Max of affine functions is convex; direct expectation comparison is separate from call checks.
   for m in range(1,5):
    phi=lambda x:max(-m*x-1,x,m*x-m)
    check(sum(a*phi(x)for x,a in mu.items())<=sum(a*phi(x)for x,a in nu.items()),'Independent max-affine convex probes')
  else:
   t=witness['threshold'];check(call(mu,t)>call(nu,t),'Failed order supplies concrete hinge witness')
 # All comparable three-site denominator-six laws have an explicit middle-row split.
 triples=[list(map(lambda z:F(z,6),c))for c in compositions(6,3)];threeplans=0
 for a,b in product(triples,repeat=2):
  if a[2]-a[0]!=b[2]-b[0]:continue
  ok,_=convex_certificate(law([-F(1),F(0),F(1)],a),law([-F(1),F(0),F(1)],b))
  if ok:
   G3=[[a[0],0,0],[b[0]-a[0],b[1],b[2]-a[2]],[0,0,a[2]]]
   martingale_table([-F(1),F(0),F(1)],a,[-F(1),F(0),F(1)],b,G3);threeplans+=1
 # Ordinary transport versus two different martingale plans.
 xs=[F(0),F(1)];a=[F(1,2)]*2;ys=list(map(F,[-1,0,1,2]));b=[F(1,4)]*4
 ordinary=[[F(1,4),F(1,4),0,0],[0,0,F(1,4),F(1,4)]]
 g1=[[F(1,4),0,F(1,4),0],[0,F(1,4),0,F(1,4)]]
 g2=[[F(1,6),F(1,4),0,F(1,12)],[F(1,12),0,F(1,4),F(1,6)]]
 check(martingale_table(xs,a,ys,b,g1)==martingale_table(xs,a,ys,b,g2)==1 and g1!=g2,'Distinct martingale plans same squared cost')
 check(sum(v*(y-x)**2 for x,row in zip(xs,ordinary)for y,v in zip(ys,row))==F(1,2),'Ordinary plan cheaper but not martingale')
 reject(lambda:martingale_table(xs,a,ys,b,ordinary),'Ordinary plan row moments rejected')
 mu=law([-F(1),F(1)],[F(1,2)]*2);nu=law([-F(4),F(0),F(4)],[F(1,20),F(9,10),F(1,20)])
 check(convex_certificate(mu,nu)[0]==convex_certificate(nu,mu)[0]==False,'Variance counterexample incomparable')
 reject(lambda:kernel([[F(3,2),-F(1,2)]]),'Negative kernel rejected')
 reject(lambda:kernel([[F(1,2),F(1,3)]]),'Wrong row sum rejected')
 reject(lambda:law([F(0),F(0)],[F(1,2)]*2),'Duplicate supports rejected')
 reject(lambda:dual_value(I,I,[F(1,2)]*2,[[F(2),0],[0,1]]),'Unbounded dual reward rejected')
 reject(lambda:bayes(I,[F(1,2),F(1,3)],I),'Invalid prior rejected')
 return encode({'status':'PASS','checks':sum(C.values()),'categories':dict(C),'binary_experiment_pairs':625,'binary_deficiency_histogram':dict(sorted(hist.items())),'five_site_laws':len(laws),'equal_mean_pairs':equalmean,'convex_ordered_pairs':ordered,'three_site_martingale_plans':threeplans,'asymmetric_certificate':{'value':asym,'kernel':explicitK,'prior':explicitpi,'reward':I},'posterior_certificate':{'fine_experiment':fine,'coarse_experiment':coarse,'garbling':postK,'coarse_fine_joint':G},'scope':'Finite rational certificates. General finite theorems are proved in prose; no claims for arbitrary infinite experiments, unrestricted losses, or generic martingale optimization.'})
if __name__=='__main__':
 p=argparse.ArgumentParser();p.add_argument('--output',type=Path,required=True);args=p.parse_args();res=run();args.output.write_text(json.dumps(res,ensure_ascii=False,indent=2)+'\n');print(json.dumps({'status':res['status'],'checks':res['checks']}))
