#!/usr/bin/env python3
"""Exact educational simplex basis-pursuit decoder. Python 3.10+, standard library.

No LP package or floating-point eigensolver is used. For p columns, A = D B,
B[k-1] = [1]*k + [-k] + [0]*(p-k-1), D[k-1]^2=p/((p-1)*k*(k+1)).
Input w = D^{-1} y contains p-1 rational row-scaled measurements. The decoder
receives w only, solves B z0=w with sum(z0)=0, and minimizes along ker B.
This is not a general LP solver and does not certify noisy exact recovery.

Run:
  python3 basis-pursuit-recovery-reader.py
  python3 -O basis-pursuit-recovery-reader.py
  python3 basis-pursuit-recovery-reader.py --p 8 --measurement 1,1,1,1,1,1,1
"""
from fractions import Fraction as F
from itertools import combinations, product
import argparse
import json
import re
import sys


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


def dot(a, b):
    require(len(a) == len(b), 'dot dimensions differ')
    return sum((x*y for x, y in zip(a, b)), F(0))


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


def inverse(a):
    n = len(a)
    require(n > 0 and all(len(row) == n for row in a), 'inverse needs a square matrix')
    r = [[F(x) for x in row] + [F(i == j) for j in range(n)] for i, row in enumerate(a)]
    for j in range(n):
        pivot = next((i for i in range(j, n) if r[i][j]), None)
        require(pivot is not None, 'singular matrix')
        r[j], r[pivot] = r[pivot], r[j]
        q = r[j][j]
        r[j] = [v/q for v in r[j]]
        for i in range(n):
            if i != j and r[i][j]:
                q = r[i][j]
                r[i] = [a-q*b for a, b in zip(r[i], r[j])]
    return [row[n:] for row in r]


class Simplex:
    def __init__(self, p):
        require(type(p) is int and 2 <= p <= 32, 'demo p must be an integer in [2,32]')
        self.p = p
        self.b = [[F(1)]*k + [F(-k)] + [F(0)]*(p-k-1) for k in range(1, p)]
        self.d2 = [F(p, (p-1)*k*(k+1)) for k in range(1, p)]
        self.g = [[sum((self.b[k][i]*self.d2[k]*self.b[k][j] for k in range(p-1)), F(0))
                   for j in range(p)] for i in range(p)]
        self.inv = inverse(self.b + [[F(1)]*p])
        for i in range(p):
            for j in range(p):
                require(self.g[i][j] == (F(1) if i == j else F(-1,p-1)), 'Gram construction mismatch')
        require(mv(self.b, [F(1)]*p) == [F(0)]*(p-1), 'ones must lie in kernel')

    def measure(self, x):
        require(len(x) == self.p, 'signal has wrong dimension')
        return mv(self.b, x)

    def decode(self, w):
        """No true signal/support input; return the entire minimizing t interval."""
        require(len(w) == self.p-1, 'measurement has wrong dimension')
        require(all(isinstance(q, (int, F)) and not isinstance(q, bool) for q in w),
                'measurement must use exact rational numbers')
        z0 = mv(self.inv, list(map(F, w)) + [F(0)])
        ordered = sorted(-a for a in z0)
        lo, hi = ordered[(self.p-1)//2], ordered[self.p//2]
        z = [a+lo for a in z0]
        residual = [a-b for a,b in zip(self.measure(z), w)]
        residual2 = dot(self.d2, [r*r for r in residual])
        require(residual2 == 0, 'exact primal feasibility failed')
        return {'z':z, 'z0':z0, 't_interval':[lo,hi], 'unique':lo == hi,
                'primal_l1':sum(map(abs,z), F(0)), 'residual_squared':residual2}

    def dual(self, w, z):
        """Try the support-interpolating u=A v; if infeasible, return that fact."""
        require(len(w) == self.p-1 and len(z) == self.p, 'dual dimensions differ')
        require(all(isinstance(q, (int,F)) and not isinstance(q,bool) for q in list(w)+list(z)),
                'dual inputs must be exact rational numbers')
        residual = [a-b for a,b in zip(self.measure(z),w)]
        residual2 = dot(self.d2,[r*r for r in residual])
        s = [i for i,a in enumerate(z) if a]
        v = [F(0)]*self.p
        if s:
            gs = [[self.g[i][j] for j in s] for i in s]
            coeff = mv(inverse(gs), [F(1 if z[i]>0 else -1) for i in s])
            for i,c in zip(s,coeff):
                v[i] = c
        corr = mv(self.g,v)
        dual_value = dot(w, [d*b for d,b in zip(self.d2,mv(self.b,v))])
        feasible = max(map(abs,corr), default=F(0)) <= 1
        outside = max((abs(corr[i]) for i in range(self.p) if i not in s), default=F(0))
        gap = sum(map(abs,z),F(0))-dual_value if feasible and residual2 == 0 else None
        return {'u_representation':'u=A v', 'v':v, 'correlations':corr,
                'dual_feasible':feasible, 'outside_max':outside,
                'dual_value':dual_value,
                'primal_residual_squared':residual2, 'gap':gap,
                'strict_support_certificate':feasible and outside < 1 and residual2 == 0 and gap == 0}

    def restricted_gram_check(self, k):
        require(type(k) is int and 1 <= k <= self.p, 'invalid support size')
        lo, hi = F(self.p-k,self.p-1), F(self.p,self.p-1)
        count = 0
        for support in combinations(range(self.p),k):
            g = [[self.g[i][j] for j in support] for i in support]
            require(mv(g,[F(1)]*k) == [lo]*k, 'constant eigenspace failed')
            # e_i-e_last span the k-1 dimensional zero-sum eigenspace.
            for i in range(k-1):
                require(all(g[j][i]-g[j][-1] == hi*(int(j==i)-int(j==k-1)) for j in range(k)),
                        'zero-sum eigenspace failed')
            count += 1
        return {'support_size':k, 'supports_checked':count, 'alpha_squared':lo,
                'beta_squared':hi, 'spectrum_multiplicity':[1,k-1]}


def failure_case():
    c = [[F(1),F(0),F(1,3)],[F(0),F(1),F(1,3)]]
    x, z, u = [F(0),F(0),F(1)], [F(1,3),F(1,3),F(0)], [F(1),F(1)]
    y = mv(c,x)
    require(mv(c,z) == y, 'counterexample feasibility failed')
    kernel = [F(-1),F(-1),F(3)]
    require(mv(c,kernel) == [0,0], 'counterexample kernel failed')
    dets = []
    for i,j in combinations(range(3),2):
        det = c[0][i]*c[1][j]-c[0][j]*c[1][i]
        require(det != 0, 'two columns must be independent')
        dets.append(det)
    corr = [sum(c[i][j]*u[i] for i in range(2)) for j in range(3)]
    primal, dual = sum(map(abs,z)), dot(y,u)
    require(max(map(abs,corr)) <= 1 and primal == dual == F(2,3), 'counterexample dual failed')
    # For pairs (1,3),(2,3), characteristic polynomial 9 t^2 - 11 t + 1.
    # 1 < sqrt(85) < 11 and sqrt(85)<11 give min eigenvalue>0 and delta2<1.
    require(85 < 121 and 85 > 1, 'radical enclosure failed')
    slopes = [F(-5,3),F(1,3),F(5,3)]
    require(slopes[0] < 0 < slopes[1] < slopes[2], 'piecewise minimum proof failed')
    return {'A':c,'truth':x,'decoded':z,'measurement':y,'kernel':kernel,
            'two_column_determinants':dets,'pair_characteristic_coefficients':[9,-11,1],
            'delta2_exact':'(7+sqrt(85))/18 < 1', 'u':u,'correlations':corr,
            'truth_l1':1,'decoded_l1':primal,'gap':primal-dual,
            'residual_squared':0,'objective_slopes':slopes,
            'recovery':False,'decoder_optimal_and_unique':True}


def self_test():
    rows, total = [], 0
    for p,s in [(8,1),(16,2)]:
        a = Simplex(p)
        cert = a.restricted_gram_check(3*s)
        require(F(2) > cert['beta_squared']/cert['alpha_squared'], 'RIP ratio failed')
        cases = 0
        for r in range(s+1):
            for support in combinations(range(p),r):
                for signs in product([-1,1],repeat=r):
                    x = [F(0)]*p
                    for i,(j,sign) in enumerate(zip(support,signs)):
                        x[j] = F(sign*(i+1))
                    w = a.measure(x)
                    dec = a.decode(w)
                    require(dec['z'] == x and dec['unique'], 'sparse decode failed')
                    d = a.dual(w,dec['z'])
                    require(d['dual_feasible'] and d['gap']==0 and d['strict_support_certificate'],
                            'sparse dual failed')
                    cases += 1
        x = [F(0)]*p
        if p == 8:
            x[0] = 1
        else:
            x[2],x[11] = 1,-2
        w = a.measure(x); dec = a.decode(w)
        rows.append({'m':p-1,'p':p,'s':s,'lambda':2,'matrix_certificate':cert,
                     'strict_ratio_margin':F(2)-cert['beta_squared']/cert['alpha_squared'],
                     'signed_support_decodes':cases,'example_measurement_w':w,
                     'example':dec,'example_dual':a.dual(w,dec['z']),
                     'nsp_support_mass_bound':s,'nsp_complement_mass_bound':p-s})
        total += cases
    a = Simplex(8)
    tie = a.decode(a.measure([F(1)]*4+[F(0)]*4))
    require(tie['t_interval']==[F(-1,2),F(1,2)] and not tie['unique'], 'median tie handling failed')
    bad = a.dual([F(0)]*7,[F(1)]+[F(0)]*7)
    require(not bad['strict_support_certificate'] and bad['gap'] is None and
            bad['primal_residual_squared'] == 1, 'infeasible candidate got a certificate')
    rejected = []
    for label,fn in [('p=1',lambda:Simplex(1)),('boolean p',lambda:Simplex(True)),
                     ('short measurement',lambda:a.decode([F(0)])),
                     ('floating measurement',lambda:a.decode([0.0]*7)),
                     ('singular inverse',lambda:inverse([[1,1],[1,1]])),
                     ('empty inverse',lambda:inverse([])),
                     ('dual short measurement',lambda:a.dual([F(0)],[F(0)]*8)),
                     ('dual short candidate',lambda:a.dual([F(0)]*7,[F(0)])),
                     ('dual floating candidate',lambda:a.dual([F(0)]*7,[0.0]*8)),
                     ('huge exponent',lambda:rational('1e999999999')),
                     ('NaN token',lambda:rational('nan')),
                     ('zero denominator',lambda:rational('1/0'))]:
        try:
            fn()
        except ValueError:
            rejected.append(label)
        else:
            raise ValueError('invalid input accepted: '+label)
    require(total == 530, 'signed support count mismatch')
    return {'arithmetic':'fractions.Fraction; exact row-scaled measurements, Gram and dual checks',
            'uniform_guarantee':'proved analytically by the page; enumeration is independent verification',
            'examples':rows, 'total_signed_support_decodes_including_zero':total,
            'counterexample':failure_case(),'half_support_nonunique':tie,
            'infeasible_zero_gap':{'A':[[1]],'y':[1],'z':[0],'u':[0],'reported_difference':0,
                                   'residual':1,'valid_primal_gap':False},
            'invalid_inputs_rejected':rejected,'status':'PASS'}


def rational(text):
    require(0 < len(text) <= 120, 'each rational token must have 1..120 characters')
    require(re.fullmatch(r'[+-]?(?:[0-9]+(?:/[0-9]+)?|[0-9]+\.[0-9]*|\.[0-9]+)', text) is not None,
            'use an integer, decimal, or integer/positive-integer fraction; no exponent notation')
    try:
        q = F(text)
    except (ValueError,ZeroDivisionError,OverflowError) as exc:
        raise ValueError('measurement entries must be finite rationals') from exc
    require(q.numerator.bit_length() <= 1024 and q.denominator.bit_length() <= 1024,
            'demo rational precision limit exceeded')
    return q


def json_value(x):
    if isinstance(x,F):
        return str(x)
    raise TypeError(type(x).__name__)


def main():
    parser = argparse.ArgumentParser(description=__doc__,formatter_class=argparse.RawDescriptionHelpFormatter)
    parser.add_argument('--p',type=int,default=8)
    parser.add_argument('--measurement',help='p-1 comma-separated exact rational entries of w=D^-1 y')
    args = parser.parse_args()
    try:
        if args.measurement is None:
            require(args.p==8, '--p without --measurement is not a separate self-test configuration')
            out = self_test()
        else:
            a = Simplex(args.p)
            w = [rational(t.strip()) for t in args.measurement.split(',')]
            out = {'p':args.p,'measurement_w':w,'decoder':a.decode(w),
                   'scope':'Exact simplex-family optimization; no supplied sparsity assumption or uniform RIP certificate'}
            # Arbitrary dense optima need not admit this particular interpolating dual.
        print(json.dumps(out,ensure_ascii=False,indent=2,default=json_value))
    except ValueError as exc:
        parser.error(str(exc))


if __name__ == '__main__':
    main()
