#!/usr/bin/env python3
"""Reproduce the operator-splitting unit using standard-library arithmetic.
Rational examples use Fraction. Entropy examples use Decimal at 70 digits,
with explicit numerical tolerance; those tests are not exact symbolic proof.
Run: python foundations-operator-splitting-check.py [--output PATH]
Finite verification remains active under python -O.
"""
from fractions import Fraction as Q
from pathlib import Path
import json
from decimal import Decimal, getcontext
getcontext().prec=70
class DecimalMath:
    mpf=staticmethod(Decimal)
    log=staticmethod(lambda x: Decimal(x).ln())
    exp=staticmethod(lambda x: Decimal(x).exp())
    root=staticmethod(lambda x,n: (Decimal(x).ln()/Decimal(n)).exp())
    nstr=staticmethod(lambda x,n: format(x,f'.{n}g'))
mp=DecimalMath()
import argparse
def check(condition,message):
    if not condition:
        raise AssertionError(message)

parser=argparse.ArgumentParser(description='Exact Fraction and 70-digit Decimal checks for operator splitting; Python standard library only.')
parser.add_argument('--output',type=Path,default=Path(__file__).with_name('foundations-operator-splitting-results.json'))
args=parser.parse_args()
root=Path(__file__).resolve().parent
v=lambda *a:tuple(map(Q,a))
add=lambda a,b:tuple(x+y for x,y in zip(a,b))
sub=lambda a,b:tuple(x-y for x,y in zip(a,b))
scale=lambda t,a:tuple(t*x for x in a)
dot=lambda a,b:sum((x*y for x,y in zip(a,b)),Q())
n2=lambda a:dot(a,a)
J=lambda a:(-a[1],a[0])
clip=lambda a:tuple(min(Q(1),max(Q(-1),x)) for x in a)
fmt=lambda a:[str(x) for x in a]
# Rotation, with all steps in the box for x0=(1/2,0).
x=v(Q(1,2),0);ave=v(0,0);rows=[];gam=Q(1,2)
for k in range(1,65):
 w=clip(sub(x,scale(gam,J(x))));p=clip(sub(x,scale(gam,J(w))))
 check(p==sub(scale(1-gam**2,x),scale(gam,J(x))), 'finite certificate 01: p==sub(scale(1-gam**2,x),scale(gam,J(x)))')
 ave=add(ave,w);avg=scale(Q(1,k),ave);gap=sum(abs(a) for a in avg)
 check(gap<=Q(13,4*k), 'finite certificate 02: gap<=Q(13,4*k)')
 if k in [1,2,3,4,8,16,32,64]:rows.append({'k':k,'predictor':fmt(w),'point':fmt(p),'average_predictor':fmt(avg),'minty_gap':str(gap),'bound':str(Q(13,4*k))})
 x=p
res={'extragradient_box_rotation':rows,'extragradient_rotation_norm_square_factor':'13/16'}
# General one-step VI inequality for affine monotone fields on [-1,1]^2.
grid=[v(a,b) for a in [Q(-1),Q(-1,2),Q(0),Q(1,2),Q(1)] for b in [Q(-1),Q(-1,2),Q(0),Q(1,2),Q(1)]]
count=0
for mu,a in [(Q(0),Q(1)),(Q(1),Q(0)),(Q(1),Q(1)),(Q(2),Q(-1))]:
 # use a rational certified Lipschitz upper bound, not a sampled estimate.
 L=abs(mu)+abs(a);g=1/(2*L)
 for b in [v(0,0),v(1,-1)]:
  F=lambda x:add(add(scale(mu,x),scale(a,J(x))),b)
  for x in grid:
   w=clip(sub(x,scale(g,F(x))));p=clip(sub(x,scale(g,F(w))))
   for u in grid:
    rhs=n2(sub(x,u))-n2(sub(p,u))-(1-g*L)*(n2(sub(x,w))+n2(sub(w,p)))
    check(2*g*dot(F(w),sub(w,u))<=rhs, 'finite certificate 03: 2*g*dot(F(w),sub(w,u))<=rhs');count+=1
res['extragradient_one_step_inequalities']=count
# Forward-backward with a genuinely nongradient cocoercive field B=I+J.
x=v(1,0);fb=[]
for k in range(1,5):
 x=sub(x,scale(Q(1,2),add(x,J(x))));fb.append({'k':k,'point':fmt(x),'norm_squared':str(n2(x))})
res['forward_backward_damped_rotation']=fb
# FBF predictor is feasible while the correction can leave the quadrant.
x=v(1,0);B=lambda x:add(scale(-1,J(x)),v(1,2));g=Q(1,2)
p=tuple(max(Q(0),a) for a in sub(x,scale(g,B(x))));xp=add(p,scale(g,sub(B(x),B(p))));eg=tuple(max(Q(0),a) for a in sub(x,scale(g,B(p))))
check(p==v(Q(1,2),0) and xp==v(Q(1,2),Q(-1,4)) and eg==p, 'finite certificate 04: p==v(Q(1,2),0) and xp==v(Q(1,2),Q(-1,4)) and eg==p')
res['fbf_feasibility_boundary']={'input':fmt(x),'predictor':fmt(p),'corrected_point':fmt(xp),'two_projection_extragradient_point':fmt(eg)}
# FBF exact Fejer inequality, A=lambda I, B=mu I + a J and zero solution.
count=0
for lam in [Q(0),Q(1),Q(2)]:
 for mu,a in [(Q(0),Q(1)),(Q(1),Q(0)),(Q(1),Q(1)),(Q(2),Q(-1))]:
  L2=mu**2+a**2;g=Q(1,4)
  B=lambda x:add(scale(mu,x),scale(a,J(x)))
  for x in grid:
   p=scale(1/(1+g*lam),sub(x,scale(g,B(x))));xp=add(p,scale(g,sub(B(x),B(p))))
   r=add(scale(1/g,sub(x,p)),sub(B(p),B(x)))
   check(xp==sub(x,scale(g,r)), 'finite certificate 05: xp==sub(x,scale(g,r))')
   check(n2(xp)<=n2(x)-(1-g*g*L2)*n2(sub(x,p)), 'finite certificate 06: n2(xp)<=n2(x)-(1-g*g*L2)*n2(sub(x,p))');count+=1
res['fbf_fejer_inequalities']=count
# DR on lines y=0 and y=x.
PA=lambda z:(z[0],Q(0));PB=lambda z:scale(Q(1,2),v(sum(z),sum(z)))
z=v(1,0);dr=[]
for k in range(1,5):
 p=PA(z);q=PB(sub(scale(2,p),z));zp=add(z,sub(q,p));check(n2(zp)+n2(sub(zp,z))<=n2(z), 'finite certificate 07: n2(zp)+n2(sub(zp,z))<=n2(z)')
 dr.append({'k':k,'state':fmt(zp),'shadow_A':fmt(PA(zp))});z=zp
res['douglas_rachford_lines']=dr
z=Q(0);state=[]
for k in range(1,5):
 p=(z+2)/2;z=z-p;state.append({'k':k,'state':str(z),'new_shadow':str((z+2)/2)})
check(z==Q(-15,8), 'finite certificate 08: z==Q(-15,8)')
res['douglas_rachford_state_not_solution']={'iterations':state,'state_limit':'-2','solution_limit':'0'}
# Primal-first PDHG fused pair, tau=sigma=1/2, K=[1,-1].
K=lambda x:x[0]-x[1];KT=lambda y:(y,-y)
f=lambda x:n2(sub(x,v(2,0)))/2
P=lambda x:f(x)+abs(K(x))
D=lambda y:2*y-y*y
energy=lambda x,y:2*n2(sub(x,v(1,1)))+2*(y-1)**2-2*K(sub(x,v(1,1)))*(y-1)
x=v(0,0);y=Q(0);rows=[];total_x=v(0,0);total_y=Q(0);count=0
for k in range(1,65):
 xp=scale(Q(1,3),add(sub(scale(2,x),KT(y)),v(2,0)))
 yp=min(Q(1),max(Q(-1),y+K(sub(scale(2,xp),x))/2))
 dx=sub(x,xp);dy=y-yp;stepM=2*n2(dx)+2*dy**2-2*K(dx)*dy
 check(energy(xp,yp)+stepM<=energy(x,y), 'finite certificate 09: energy(xp,yp)+stepM<=energy(x,y)')
 total_x=add(total_x,xp);total_y+=yp
 if k>=2:
  t=Q(2,3)**(k-2)
  check(xp==v(1-t/9,1-7*t/9) and yp==1, 'finite certificate 10: xp==v(1-t/9,1-7*t/9) and yp==1')
  check(P(xp)-D(yp)==Q(25,81)*Q(4,9)**(k-2), 'finite certificate 11: P(xp)-D(yp)==Q(25,81)*Q(4,9)**(k-2)')
 for ux in [v(0,0),v(1,1),v(2,0),v(-1,1)]:
  for uy in map(Q,[-1,0,1]):
   def metricdist(x,y):
    a=sub(x,ux);b=y-uy;return 2*n2(a)+2*b*b-2*K(a)*b
   gap=f(xp)+K(xp)*uy-f(ux)-K(ux)*yp
   check(2*gap<=metricdist(x,y)-metricdist(xp,yp)-stepM, 'finite certificate 12: 2*gap<=metricdist(x,y)-metricdist(xp,yp)-stepM');count+=1
 if k in [1,2,3,4,8,16,32,64]:rows.append({'k':k,'x':fmt(xp),'y':str(yp),'primal':str(P(xp)),'dual':str(D(yp)),'full_gap':str(P(xp)-D(yp))})
 x,y=xp,yp
res['pdhg_fused_pair']=rows;res['pdhg_one_step_gap_inequalities']=count
# Mirror-Prox matching pennies, independent high-precision transcendental checks.
mpv=lambda a:tuple(map(mp.mpf,a))
p=mpv(['0.75','0.25']);q=mpv(['0.5','0.5']);gamma=mp.log(2)
M=lambda p:(p[0]-p[1],p[1]-p[0])
neg=lambda p:tuple(-x for x in p)
prox=lambda p,g:tuple(w/sum(p[i]*mp.exp(-gamma*g[i]) for i in range(2)) for w in (p[i]*mp.exp(-gamma*g[i]) for i in range(2)))
sump=[mp.mpf(0),mp.mpf(0)];sumq=sump.copy();records=[];mp_inequalities=0
kl=lambda u,z:sum((a*(a/z[i]).ln() for i,a in enumerate(u) if a),mp.mpf(0))
comparisons=[((mp.mpf(a),1-mp.mpf(a)),(mp.mpf(b),1-mp.mpf(b))) for a in [0,1] for b in [0,1]]
for k in range(1,257):
 wp=prox(p,M(q));wq=prox(q,neg(M(p)))
 pp=prox(p,M(wq));qq=prox(q,neg(M(wp)))
 if k==1:
  check(abs(wq[0]-mp.mpf(2)/3)<mp.mpf('1e-65'), "finite certificate 13: abs(wq[0]-mp.mpf(2)/3)<mp.mpf('1e-65')")
  check(abs(pp[0]-3/(3+mp.root(4,3)))<mp.mpf('1e-65'), "finite certificate 14: abs(pp[0]-3/(3+mp.root(4,3)))<mp.mpf('1e-65')")
 for i in range(2):sump[i]+=wp[i];sumq[i]+=wq[i]
 gap=abs(2*sump[0]/k-1)+abs(2*sumq[0]/k-1)
 check(gap<=mp.mpf(3)/k+mp.mpf('1e-65'), "finite certificate 15: gap<=mp.mpf(3)/k+mp.mpf('1e-65')")
 if k in [1,2,3,4,8,16,32,64,128,256]:records.append({'k':k,'predictor_p':[mp.nstr(t,30) for t in wp],'predictor_q':[mp.nstr(t,30) for t in wq],'anchor_p':[mp.nstr(t,30) for t in pp],'anchor_q':[mp.nstr(t,30) for t in qq],'average_saddle_gap':mp.nstr(gap,30),'bound':mp.nstr(mp.mpf(3)/k,30)})
 for up,uq in comparisons:
  Fwp=M(wq);Fwq=neg(M(wp))
  lhs=gamma*sum((Fwp[i]*(wp[i]-up[i])+Fwq[i]*(wq[i]-uq[i]) for i in range(2)),mp.mpf(0))
  rhs=kl(up,p)+kl(uq,q)-kl(up,pp)-kl(uq,qq)
  check(lhs<=rhs+mp.mpf('1e-65'), "finite certificate 16: lhs<=rhs+mp.mpf('1e-65')");mp_inequalities+=1
 p,q=pp,qq
res['mirror_prox_matching_pennies']=records
res['mirror_prox_one_step_inequalities']=mp_inequalities
# Forward-backward full squared budget, including the completed-square term.
count=0
for lam in [Q(0),Q(1),Q(2)]:
 for mu,a in [(Q(1),Q(0)),(Q(1),Q(1)),(Q(2),Q(-1)),(Q(3),Q(4))]:
  beta=mu/(mu*mu+a*a);g=beta
  B=lambda x:add(scale(mu,x),scale(a,J(x)))
  for x in grid:
   p=scale(1/(1+g*lam),sub(x,scale(g,B(x))));d=sub(x,p);b=B(x)
   rhs=n2(x)-(1-g/(2*beta))*n2(d)-2*g*beta*n2(sub(b,scale(1/(2*beta),d)))
   check(n2(p)<=rhs, 'finite certificate 17: n2(p)<=rhs');count+=1
res['forward_backward_full_square_inequalities']=count
# Transfer tasks: scaled rotation, damped rotation, shifted DR and changed K.
check(Q(1,4)*2==Q(1,2), 'finite certificate 18: Q(1,4)*2==Q(1,2)') # same points, doubled VI gap, doubled bound.
check(Q(13,2)/325==Q(1,50), 'finite certificate 19: Q(13,2)/325==Q(1,50)')
mu,a=Q(3),Q(4);beta=mu/(mu*mu+a*a)
check(beta==Q(3,25), 'finite certificate 20: beta==Q(3,25)')
check((1-Q(1,5)*mu)**2+(Q(1,5)*a)**2==Q(4,5), 'finite certificate 21: (1-Q(1,5)*mu)**2+(Q(1,5)*a)**2==Q(4,5)')
check((1-Q(1,4)*mu)**2+(Q(1,4)*a)**2==Q(17,16), 'finite certificate 22: (1-Q(1,4)*mu)**2+(Q(1,4)*a)**2==Q(17,16)')
check(Q(63,100)**2+Q(4,25)**2==Q(169,400), 'finite certificate 23: Q(63,100)**2+Q(4,25)**2==Q(169,400)')
check(Q(6,600)==Q(1,100), 'finite certificate 24: Q(6,600)==Q(1,100)')
z=Q(0);dr_transfer=[]
for k in range(1,5):
 p=(z+8)/3;z=z+1-p
 check(z==-5+5*Q(2,3)**k, 'finite certificate 25: z==-5+5*Q(2,3)**k')
 dr_transfer.append({'k':k,'state':str(z),'shadow':str((z+8)/3)})
x=v(Q(2,3),0);y=Q(1,9)
P3=lambda x:f(x)+3*abs(K(x));D3=lambda y:6*y-9*y*y
check(P3(x)-D3(y)==Q(7,3), 'finite certificate 26: P3(x)-D3(y)==Q(7,3)')
res['scale_and_interface_checks']={'scaled_rotation_L':'2','scaled_rotation_gamma':'1/4','rounds_for_gap_1_over_50':325,'damped_rotation_beta':'3/25','safe_gamma_1_over_5_squared_factor':'4/5','unsafe_gamma_1_over_4_squared_factor':'17/16','shifted_DR':dr_transfer,'pdhg_scaled_K_first_x':fmt(x),'pdhg_scaled_K_first_y':str(y),'pdhg_scaled_K_first_gap':'7/3'}
# Genuine structural transfer: three signal nodes, two edges, distinct active roles.
b3=v(0,0,3)
K3=lambda x:(x[1]-x[0],x[2]-x[1])
KT3=lambda y:(-y[0],y[0]-y[1],y[1])
f3=lambda x:n2(sub(x,b3))/2
P3=lambda x:f3(x)+sum(abs(t) for t in K3(x))
D3=lambda y:dot(b3,KT3(y))-n2(KT3(y))/2
xs=v(Q(1,2),Q(1,2),2);ys=v(Q(1,2),1)
check(add(sub(xs,b3),KT3(ys))==v(0,0,0), 'finite certificate 27: add(sub(xs,b3),KT3(ys))==v(0,0,0)')
check(K3(xs)==v(0,Q(3,2)) and P3(xs)==D3(ys)==Q(9,4), 'finite certificate 28: K3(xs)==v(0,Q(3,2)) and P3(xs)==D3(ys)==Q(9,4)')
wrong_y=v(1,1);wrong_x=sub(b3,KT3(wrong_y))
check(wrong_x==v(1,0,2) and P3(wrong_x)==4 and D3(wrong_y)==2, 'finite certificate 29: wrong_x==v(1,0,2) and P3(wrong_x)==4 and D3(wrong_y)==2')
x=v(0,0,0);y=v(0,0);records=[];three_checks=0
xcmp=[v(0,0,0),b3,xs,v(1,0,2),v(-1,2,0)]
ycmp=[v(-1,-1),v(-1,1),v(1,-1),v(1,1),ys]
metric=lambda dx,dy:2*n2(dx)+2*n2(dy)-2*dot(K3(dx),dy)
for k in range(1,129):
 xp=scale(Q(1,3),add(sub(scale(2,x),KT3(y)),b3))
 yp=clip(add(y,scale(Q(1,2),K3(sub(scale(2,xp),x)))))
 gap=P3(xp)-D3(yp)
 check(gap>=0 and n2(sub(xp,xs))<=2*gap, 'finite certificate 30: gap>=0 and n2(sub(xp,xs))<=2*gap')
 if k==1:check(xp==v(0,0,1) and yp==v(0,1) and gap==1, 'finite certificate 31: xp==v(0,0,1) and yp==v(0,1) and gap==1')
 if k==2:check(xp==v(0,Q(1,3),Q(4,3)) and yp==v(Q(1,3),1) and gap==Q(5,9), 'finite certificate 32: xp==v(0,Q(1,3),Q(4,3)) and yp==v(Q(1,3),1) and gap==Q(5,9)')
 step=metric(sub(x,xp),sub(y,yp))
 for u in xcmp:
  for z in ycmp:
   saddle_gap=f3(xp)+dot(K3(xp),z)-f3(u)-dot(K3(u),yp)
   before=metric(sub(x,u),sub(y,z));after=metric(sub(xp,u),sub(yp,z))
   check(2*saddle_gap<=before-after-step, 'finite certificate 33: 2*saddle_gap<=before-after-step');three_checks+=1
 if k in [1,2,3,4,8,16,32,64,128]:records.append({'k':k,'x':fmt(xp),'y':fmt(yp),'edges':fmt(K3(xp)),'primal':str(P3(xp)),'dual':str(D3(yp)),'full_gap':str(gap)})
 x,y=xp,yp
res['three_node_fusion_one_step_gap_checks']=three_checks
res['capstone_three_node_fusion']={'optimal_x':fmt(xs),'optimal_y':fmt(ys),'optimal_value':'9/4','K_norm_squared':'3','step_product':'3/4','iterations':records}
res['status']='PASS'
args.output.parent.mkdir(parents=True,exist_ok=True)
args.output.write_text(json.dumps(res,ensure_ascii=False,indent=2)+'\n')
print(json.dumps({k:v for k,v in res.items() if isinstance(v,(str,int))},ensure_ascii=False,indent=2))
