#!/usr/bin/env python3
"""Reproduce C6 examples and a shared constrained-QP comparison. No network."""
import math,json
from pathlib import Path
import numpy as np
from fractions import Fraction as F
import argparse
parser=argparse.ArgumentParser(description=__doc__)
parser.add_argument('--output-dir',default='C6-verification-output')
ROOT=Path(parser.parse_args().output_dir).resolve()
checks=[]
def check(name,ok,**details):
    assert bool(ok),name
    checks.append(dict(name=name,passed=True,**details))
def close(a,b,tol=1e-10): return np.max(np.abs(np.asarray(a)-np.asarray(b)))<tol
def fqp(x):return .5*float((x-np.array([2.,1.]))@(x-np.array([2.,1.])))
# Trust-region certificates, including the inaccessible negative eigendirection.
H=np.diag([-2.,1.]);g=np.array([0.,1.]);p=np.array([math.sqrt(8)/3,-1/3])
check('hard_case_certificate',close((H+2*np.eye(2))@p,-g) and close(p@p,1) and close(g@p+.5*p@H@p,-7/6))
check('hard_case_CG_is_not_global',close(g@np.array([0.,-1.])+.5*np.array([0.,-1.])@H@np.array([0.,-1.]),-.5))
f=lambda x:x**4-2*x*x+x
tr=[]
for x,p in [(0.,-2.),(0.,-1.),(-1.,-1/8)]:
 grad=4*x**3-4*x+1;hess=12*x*x-4;pred=-grad*p-.5*hess*p*p;ared=f(x)-f(x+p)
 tr.append(dict(x=x,p=p,pred=pred,ared=ared,ratio=ared/pred))
check('trust_region_acceptance',close([v['ared'] for v in tr],[-6,2,223/4096]),trace=tr)
tau=5*(math.sqrt(43)-3)/51;z=np.array([.4,.4]);d=np.array([24/25,-6/25]);p=z+tau*d
check('truncated_CG_second_boundary',close(p@p,.64) and (np.array([-1.,-1.])@p+.5*p@np.diag([1.,4.])@p)<-.4,tau=tau,p=p.tolist())
# The active-set method actually executes the release/block branches.
A=np.array([[-1.,0.],[0.,-1.],[1.,1.]]);b=np.array([0.,0.,2.]);a=np.array([2.,1.]);x=np.zeros(2);W=[0,1];active=[]
for k in range(20):
 Aw=A[W];K=np.block([[np.eye(2),Aw.T],[Aw,np.zeros((len(W),len(W)))]])
 sol=np.linalg.solve(K,np.r_[-(x-a),np.zeros(len(W))]);p=sol[:2];nu=sol[2:]
 row=dict(k=k,x=x.tolist(),working=[i+1 for i in W],p=p.tolist(),multipliers=nu.tolist())
 if np.linalg.norm(p)<1e-12:
  if len(nu)==0 or min(nu)>=-1e-12:row['event']='optimal';active.append(row);break
  j=int(np.argmin(nu));row['event']='release '+str(W[j]+1);W.pop(j)
 else:
  ratios=[((b[i]-A[i]@x)/(A[i]@p),i) for i in range(3) if i not in W and A[i]@p>1e-12]
  alpha=min(1.,min((r for r,i in ratios),default=math.inf));x=x+alpha*p
  blocking=[i for r,i in ratios if abs(r-alpha)<1e-10 and r<=1+1e-12]
  if blocking:W.append(min(blocking));row['event']='block '+str(min(blocking)+1)
  else:row['event']='full step'
  row['alpha']=alpha
 active.append(row)
check('active_set_trajectory',close(x,[1.5,.5]) and len(active)==5,trace=active)
# Equality article and inequality shared-QP penalty solutions.
penalty=[]
for mu in [1.,10.,100.,10000.]:
 r=1/(1+2*mu);lam=mu*r;xe=np.array([2-lam,-lam]);xq=np.array([2-lam,1-lam]);K=np.eye(2)+mu*np.ones((2,2))
 check('penalty_'+str(mu),close(sum(xe)-1,r) and close(K@xe,[2+mu,mu]) and close(np.linalg.cond(K),1+2*mu,tol=1e-6))
 penalty.append(dict(mu=mu,x=xq.tolist(),violation=r,lambda_estimate=lam,condition_number=1+2*mu))
check('exact_l1_threshold',all(max(2-rho,0)==expected for rho,expected in [(1,1),(2,0),(3,0)]))
alm=[];lam=F(0)
for k in range(1,6):
 lam=(lam+1)/3;r=1-2*lam
 check('ALM_exact_'+str(k),r==F(1,3**k))
 alm.append(dict(k=k,x=[str(2-lam),str(1-lam)],lambda_=str(lam),violation=str(r)))
# Curved SQP using exact rationals.
p=[F(16,15),F(-4,5)];z=[F(3,5),F(4,5)];c=lambda z:z[0]**2+z[1]**2-1;fs=lambda z:(z[0]-2)**2+z[1]**2
full=[z[i]+p[i] for i in range(2)];half=[z[i]+p[i]/2 for i in range(2)]
check('SQP_rational_step',c(full)==F(16,9) and c(half)==F(4,9) and fs(full)+2*abs(c(full))==F(11,3) and fs(half)+2*abs(c(half))==F(9,5))
# Barrier LP central path and exact dual gap.
barrier=[]
for t in [1,4,16,64]:
 y=2/(t+2+math.sqrt(t*t+4));x=1-y;lx=1/(t*x);ly=1/(t*y);nu=lx-1
 check('barrier_'+str(t),close(ly-lx,1) and close((x+2*y)+nu,2/t))
 barrier.append(dict(t=t,x=x,y=y,lambda_=[lx,ly],nu=nu,gap=2/t,error=y))
# LP infeasible-start block Newton example.
Alp=np.ones((1,3));clp=np.array([1.,2.,3.]);xl=np.ones(3);sl=np.ones(3);yl=np.zeros(1)
K=np.block([[Alp,np.zeros((1,1)),np.zeros((1,3))],[np.zeros((3,3)),Alp.T,np.eye(3)],[np.diag(sl),np.zeros((3,1)),np.diag(xl)]])
rp=Alp@xl-1;rd=Alp.T@yl+sl-clp
sol=np.linalg.solve(K,-np.r_[rp,rd,xl*sl-.2]);dx,dy,ds=sol[:3],sol[3:4],sol[4:]
alpha=.54;xx=xl+alpha*dx;ss=sl+alpha*ds;yy=yl+alpha*dy
check('LP_infeasible_start',close(dx,[1/3,-2/3,-5/3]) and close(ds,[-17/15,-2/15,13/15]) and close(xx@ss,7491/6250) and close(clp@xx-yy[0],2.148),x=xx.tolist(),s=ss.tolist(),y=yy.tolist(),complementarity=float(xx@ss),objective_difference=float(clp@xx-yy[0]))
# Shared-QP primal-dual steps, starting exactly primal and dual stationary.
x=np.array([.5,.5]);s=b-A@x;lam=np.array([.5,1.5,2.]);pd=[]
for k in range(60):
 rp=A@x+s-b;rd=x-a+A.T@lam;gap=float(s@lam);dual=float(lam@(A@a-b)-.5*(A.T@lam)@(A.T@lam))
 check('QP_PDI_invariants_'+str(k),np.min(s)>0 and np.min(lam)>0 and close(rp,0) and close(rd,0) and close(fqp(x)-dual,gap,tol=1e-9))
 row=dict(k=k,x=x.tolist(),s=s.tolist(),lambda_=lam.tolist(),objective=fqp(x),dual=dual,gap=gap)
 pd.append(row)
 if gap<=1e-9:break
 mu=gap/3;K=np.block([[np.eye(2),np.zeros((2,3)),A.T],[A,np.eye(3),np.zeros((3,3))],[np.zeros((3,2)),np.diag(lam),np.diag(s)]])
 sol=np.linalg.solve(K,-np.r_[rd,rp,s*lam-.2*mu]);dx,ds,dl=sol[:2],sol[2:5],sol[5:]
 abd=min([1.]+[-s[i]/ds[i] for i in range(3) if ds[i]<0]+[-lam[i]/dl[i] for i in range(3) if dl[i]<0]);alpha=.9*abd
 # A centering/backtracking safeguard ensures actual gap decrease in this executable case.
 while (s+alpha*ds)@(lam+alpha*dl)>(1-.1*alpha)*gap:alpha*=.5
 row.update(alpha=alpha,dx=dx.tolist(),ds=ds.tolist(),dl=dl.tolist())
 x=x+alpha*dx;s=s+alpha*ds;lam=lam+alpha*dl
check('QP_PDI_target',pd[-1]['gap']<=1e-9 and close(x,[1.5,.5],1e-7),iterations=len(pd)-1)
# Shared QP logarithmic barrier via damped Newton, independent of PD solve.
centers=[]
for t in [1,4,16,64]:
 x=np.array([.5,.5])
 for it in range(100):
  s=b-A@x;G=t*(x-a)+A.T@(1/s);B=t*np.eye(2)+A.T@np.diag(1/s**2)@A;p=np.linalg.solve(B,-G)
  if -G@p/2<1e-25:break
  al=1.;val=t*fqp(x)-sum(np.log(s))
  while np.min(b-A@(x+al*p))<=0 or t*fqp(x+al*p)-sum(np.log(b-A@(x+al*p)))>val+.01*al*G@p:al*=.5
  x+=al*p
 s=b-A@x;lam=1/(t*s);dual=float(lam@(A@a-b)-.5*(A.T@lam)@(A.T@lam));rd=x-a+A.T@lam
 check('QP_barrier_'+str(t),np.linalg.norm(rd)<1e-7 and abs(fqp(x)-dual-3/t)<1e-7)
 centers.append(dict(t=t,x=x.tolist(),gap=3/t,objective=fqp(x),dual=dual,stationarity=float(np.linalg.norm(rd))))
result=dict(check_count=len(checks),all_passed=True,checks=checks,shared_qp=dict(problem='min .5*((x-2)^2+(y-1)^2), x,y>=0,x+y<=2',optimum=[1.5,.5],optimal_value=.25,active_set=active,quadratic_penalty=penalty,augmented_lagrangian=alm,barrier=centers,primal_dual=pd),article_barrier=barrier)
out=ROOT/'qa/C6/C6-numerical-results.json';out.parent.mkdir(parents=True,exist_ok=True);out.write_text(json.dumps(result,ensure_ascii=False,indent=2)+'\n')
print(json.dumps({'check_count':len(checks),'all_passed':True,'PD_steps':len(pd)-1,'first_QP_PD_step':pd[0],'last_QP_PD':pd[-1],'barrier_centers':centers},ensure_ascii=False,indent=2))
