#!/usr/bin/env python3
"""Recompute the fixed group/matrix examples using Python 3's standard library.

Exact Fraction checks establish the displayed rational certificates. Floating
iterations are diagnostics for these small symmetric examples, not interval-
arithmetic certificates or a general-purpose matrix completion package.
No files are written and no network or third-party packages are used.
"""
import argparse
import json
import math
from fractions import Fraction as Q


def require(condition, message):
    if not condition:
        raise ValueError(message)


def dot(a, b):
    return sum(x*y for x,y in zip(a,b))


def norm(a):
    return math.sqrt(dot(a,a))


def mv(a, x):
    return [dot(row,x) for row in a]


def tr(a):
    return list(map(list,zip(*a)))


def mm(a,b):
    return [[dot(row,col) for col in tr(b)] for row in a]


def flat(a):
    return [x for row in a for x in row]


def add(a,b,scale=1):
    return [[x+scale*y for x,y in zip(ar,br)] for ar,br in zip(a,b)]


def scale(a,c):
    return [[c*x for x in row] for row in a]


def inner(a,b):
    return dot(flat(a),flat(b))


def spectral_symmetric(a):
    require(len(a)==2 and all(len(r)==2 for r in a), 'Expected 2 by 2')
    require(abs(float(a[0][1]-a[1][0]))<1e-12, 'Expected symmetric input')
    mid=(a[0][0]+a[1][1])/2
    d=math.hypot(float((a[0][0]-a[1][1])/2),float(a[0][1]))
    return float(mid+d),float(mid-d)


def opnorm(a):
    return max(abs(v) for v in spectral_symmetric(a))


def nuclear(a):
    return sum(abs(v) for v in spectral_symmetric(a))


def shrink(t,tau):
    return math.copysign(max(abs(t)-tau,0),t)


def symmetric_svt(a,tau):
    """SVD shrinkage for real symmetric 2 by 2 input, including negative eigenvalues."""
    hi,lo=spectral_symmetric(a)
    if abs(hi-lo)<1e-15:
        return [[shrink(hi,tau),0.0],[0.0,shrink(hi,tau)]]
    slope=(shrink(hi,tau)-shrink(lo,tau))/(hi-lo)
    intercept=shrink(lo,tau)-slope*lo
    return [[slope*a[i][j]+(intercept if i==j else 0) for j in range(2)] for i in range(2)]


def mask(a):
    return [[a[0][0],a[0][1]],[a[1][0],0*a[1][1]]]


def exact_checks():
    A=[[Q(1),Q(0),Q(3,5)],[Q(0),Q(1),Q(0)],[Q(0),Q(0),Q(4,5)]]
    b=[Q(18,5),Q(24,5),Q(0)]; xs=[Q(3),Q(4),Q(0)]
    r=[u-v for u,v in zip(b,mv(A,xs))]; c=mv(tr(A),r)
    require(r==[Q(3,5),Q(4,5),0], 'group residual')
    require(c==[Q(3,5),Q(4,5),Q(9,25)], 'group KKT')
    require(dot(b,r)-dot(r,r)/2==Q(11,2), 'group exact dual value')
    x1=[Q(15,8),Q(5,2),Q(29,40)]
    r1=[u-v for u,v in zip(b,mv(A,x1))]
    require(r1==[Q(129,100),Q(23,10),-Q(29,50)], 'group first residual')
    require(dot(r1,r1)==Q(14581,2000), 'group first residual square')
    require(dot(r1,r1)/2+Q(77,20)==Q(29981,4000), 'group first objective')
    require(dot(b,r1)==Q(3921,250), 'group first dual numerator')
    M=[[Q(4),Q(2)],[Q(2),Q(1)]]
    U=[[Q(4,5),Q(2,5)],[Q(2,5),Q(1,5)]]
    V=[[Q(1,5),-Q(2,5)],[-Q(2,5),Q(4,5)]]
    Y=[[Q(3,4),Q(1,2)],[Q(1,2),Q(0)]]
    require(Y==add(U,V,-Q(1,4)), 'completion dual decomposition')
    require(mm(U,V)==[[0,0],[0,0]], 'orthogonal supports')
    require(inner(mask(M),Y)==Q(5), 'exact completion equality')
    E=[[Q(0),Q(0)],[Q(0),Q(1)]]
    require(mm(mm(V,E),V)==scale(V,Q(4,5)), 'tangent complement')
    BP=mask(add(M,Y)); residual=add(BP,mask(M),-1)
    require(residual==Y, 'penalized residual')
    require(inner(Y,Y)==Q(17,16), 'penalized residual square')
    require(inner(BP,Y)==Q(97,16), 'penalized dual pairing')
    require(inner(BP,Y)-inner(Y,Y)/2==Q(177,32), 'penalized exact value')
    # Nuclear denoising example and its two wrong-objective alternatives.
    Z=[[Q(3),Q(2)],[Q(2),Q(3)]]; X=[[Q(3,2)]*2 for _ in range(2)]
    R=add(Z,X,-1)
    require(inner(R,R)==5 and inner(Z,R)==11, 'SVT exact pair')
    require(inner(R,R)/2+2*3==Q(17,2), 'SVT objective')
    # The scalar reductions for overlap and rank-one identifiable failure.
    require(Q(1,2)*(1-3)**2+2*abs(1)<Q(1,2)*(2-3)**2+2*abs(2), 'overlap failure')
    require(9+1+2*abs(1-4)==16 and 9+16+2*abs(4-4)==25, 'rank relaxation failure')
    return {'rational_certificate_checks':'passed','group_optimum':'11/2','svt_optimum':'17/2',
            'exact_completion_optimum':'5','penalized_completion_optimum':'177/32',
            'completion_uniqueness_lower_bound':'3|h|/5'}


def group_run(epsilon,budget):
    A=[[1,0,.6],[0,1,0],[0,0,.8]]; AT=tr(A);b=[3.6,4.8,0];x=[0.,0.,0.];alpha=5/8
    rows=[];previous=float('inf')
    for k in range(budget+1):
        r=[u-v for u,v in zip(b,mv(A,x))]; c=mv(AT,r)
        rho=max(1.,norm(c[:2]),abs(c[2]));theta=[v/rho for v in r]
        P=dot(r,r)/2+norm(x[:2])+abs(x[2]);D=dot(b,theta)-dot(theta,theta)/2;gap=P-D
        require(P<=previous+1e-11,'group descent');previous=P
        feasible=mv(AT,theta)
        require(norm(feasible[:2])<=1+1e-12 and abs(feasible[2])<=1+1e-12,'group dual feasibility')
        require(gap>=-1e-11 and gap>=P-5.5-1e-11,'group gap bound')
        if k in {0,1,2,5,10,20,budget} or gap<=epsilon:
            rows.append({'k':k,'x':x[:],'P':P,'D':D,'gap':gap})
        if gap<0:
            return {'status':'numerical_resolution_limited','epsilon':epsilon,'iterations':k,'final_gap':gap,'final_state':rows[-1],'rows':rows}
        if gap<=epsilon:
            return {'status':'gap_met','epsilon':epsilon,'iterations':k,'rows':rows,
                    'distance_to_exact':norm([x[0]-3,x[1]-4,x[2]])}
        if k==budget:
            return {'status':'budget_exhausted','epsilon':epsilon,'iterations':k,'final_gap':gap,'final_state':rows[-1],'rows':rows}
        z=[a+alpha*d for a,d in zip(x,c)];length=norm(z[:2]);factor=max(0.,1-alpha/length) if length else 0.
        x=[factor*z[0],factor*z[1],shrink(z[2],alpha)]
    raise RuntimeError('Unreachable loop exit')


def completion_run(epsilon,budget):
    B=[[4.75,2.5],[2.5,0.]];X=[[0.,0.],[0.,0.]];M=[[4.,2.],[2.,1.]]
    previous=float('inf');rows=[]
    for k in range(budget+1):
        R=add(B,mask(X),-1);rho=max(1.,opnorm(R));Y=scale(R,1/rho)
        P=inner(R,R)/2+nuclear(X);D=inner(B,Y)-inner(Y,Y)/2;gap=P-D
        require(Y[1][1]==0,'completion dual support')
        require(opnorm(Y)<=1+1e-12,'completion spectral feasibility')
        require(P<=previous+1e-10,'completion descent');previous=P
        require(gap>=-1e-10 and gap>=P-177/32-1e-10,'completion gap bound')
        if k in {0,1,2,5,10,20,budget} or gap<=epsilon:
            rows.append({'k':k,'X':X,'P':P,'D':D,'gap':gap,'rank_at_1e-10':sum(abs(z)>1e-10 for z in spectral_symmetric(X))})
        if gap<0:
            return {'status':'numerical_resolution_limited','epsilon':epsilon,'iterations':k,'final_gap':gap,'final_state':rows[-1],'rows':rows}
        if gap<=epsilon:
            return {'status':'gap_met','epsilon':epsilon,'iterations':k,'rows':rows,
                    'distance_to_exact':math.sqrt(inner(add(X,M,-1),add(X,M,-1)))}
        if k==budget:
            return {'status':'budget_exhausted','epsilon':epsilon,'iterations':k,'final_gap':gap,'final_state':rows[-1],'rows':rows}
        Z=add(X,R)
        require(Z[1][1]==X[1][1] and all(abs(Z[i][j]-B[i][j])<1e-12 for i,j in [(0,0),(0,1),(1,0)]),'mask input invariant')
        X=symmetric_svt(Z,1.)
    raise RuntimeError('Unreachable loop exit')


def main():
    parser=argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--epsilon',type=float,default=1e-8)
    parser.add_argument('--max-iterations',type=int,default=10000)
    args=parser.parse_args()
    require(math.isfinite(args.epsilon) and args.epsilon>0,'epsilon must be positive and finite')
    if args.epsilon<1e-10:
        parser.error('epsilon below 1e-10 is outside this floating diagnostic; exact rational endpoints remain separately certified')
    require(args.max_iterations>=0,'max-iterations must be nonnegative')
    result={'exact':exact_checks(),'group':group_run(args.epsilon,args.max_iterations),'completion':completion_run(args.epsilon,args.max_iterations),
            'numerical_scope':'double-precision diagnostics; minimum epsilon 1e-10; negative computed gap reports numerical_resolution_limited; rational endpoints checked exactly'}
    print(json.dumps(result,ensure_ascii=False,indent=2))
    if result['group']['status']!='gap_met' or result['completion']['status']!='gap_met':
        raise SystemExit(1)

if __name__=='__main__':
    main()
