#!/usr/bin/env python3
"""Recompute the wavelet capstone with Python 3 standard library only.

Run: python wavelet-inference-reader.py --self-test
No random simulation is used. Outputs are deterministic floating-point values;
checks use explicit exceptions, so python -O retains all validation.
Order: scale, coarsest detail, ..., finest detail. Left minus right convention.
"""
import argparse
import json
import math
from statistics import NormalDist
from functools import wraps


def checked(fn):
    """Reject non-finite public results, including intermediate-overflow fallout."""
    def verify(obj):
        if isinstance(obj,float) and not math.isfinite(obj):
            raise ValueError('floating-point range exceeded; rescale the input')
        if isinstance(obj,dict):
            for value in obj.values(): verify(value)
        if isinstance(obj,(list,tuple)):
            for value in obj: verify(value)
    @wraps(fn)
    def wrapped(*args,**kwargs):
        try: out=fn(*args,**kwargs)
        except OverflowError as exc:
            raise ValueError('floating-point range exceeded; rescale the input') from exc
        verify(out)
        return out
    return wrapped


def finite(x, name):
    if isinstance(x, bool) or not isinstance(x, (int, float)) or not math.isfinite(x):
        raise ValueError(name + ' must be a finite real number')
    return float(x)


def vector(xs, power_two=False):
    out = [finite(x, 'coordinate') for x in xs]
    n = len(out)
    if not n or (power_two and n & (n - 1)):
        raise ValueError('length must be positive' + (' and a power of two' if power_two else ''))
    return out


def threshold(t):
    t = finite(t, 'threshold')
    if t < 0:
        raise ValueError('threshold must be nonnegative')
    return t


def noise(tau):
    tau = finite(tau, 'noise standard deviation')
    if tau <= 0:
        raise ValueError('noise standard deviation must be positive')
    return tau


@checked
def haar(xs):
    a = vector(xs, True)
    details = []
    while len(a) > 1:
        details.append([(a[i]-a[i+1])/math.sqrt(2) for i in range(0,len(a),2)])
        a = [(a[i]+a[i+1])/math.sqrt(2) for i in range(0,len(a),2)]
    return a + [x for block in reversed(details) for x in block]


@checked
def inverse(cs):
    cs = vector(cs, True)
    a, start = cs[:1], 1
    while start < len(cs):
        d = cs[start:2*start]
        a = [x for av,dv in zip(a,d) for x in ((av+dv)/math.sqrt(2),(av-dv)/math.sqrt(2))]
        start *= 2
    return a


@checked
def soft(x,t):
    x,t = finite(x,'coordinate'),threshold(t)
    return math.copysign(max(abs(x)-t,0),x)


@checked
def hard(x,t):
    x,t = finite(x,'coordinate'),threshold(t)
    return x if abs(x)>t else 0.


@checked
def denoise(xs,t):
    t=threshold(t)
    cs=haar(xs)
    return inverse(cs[:1]+[soft(x,t) for x in cs[1:]])


def phi(x):
    return math.exp(-x*x/2)/math.sqrt(2*math.pi)


def Q(x):
    return .5*math.erfc(x/math.sqrt(2))


@checked
def soft_risk(lam,u):
    """Risk at standardized mean u and threshold lam, by truncated moments."""
    lam,u=threshold(lam),finite(u,'standardized mean')
    a,b=-lam-u,lam-u
    prob=Q(a)-Q(b)
    inside=(u*u+1)*prob+2*u*(phi(a)-phi(b))+a*phi(a)-b*phi(b)
    return 1+inside+lam*lam*(1-prob)-2*prob


@checked
def sure(xs,t,tau):
    xs,t,tau=vector(xs),threshold(t),noise(tau)
    variance=finite(tau*tau,'noise variance')
    return len(xs)*variance+sum(finite(min(abs(x),t)**2,'clipped square')-2*variance*(abs(x)<=t) for x in xs)


@checked
def choose_sure(xs,tau,candidates=None):
    xs,tau=vector(xs),noise(tau)
    if candidates is None:
        ordered=sorted(abs(x) for x in xs)
        ts=sorted(set([0.]+ordered))
        k,prefix=0,0.
        scores=[]
        for t in ts:
            while k<len(xs) and ordered[k]<=t:
                prefix=finite(prefix+finite(ordered[k]*ordered[k],'coefficient square'),'prefix energy')
                k+=1
            scores.append((t,len(xs)*tau*tau+prefix+(len(xs)-k)*t*t-2*tau*tau*k))
    else:
        ts=sorted(set(threshold(t) for t in candidates))
        if not ts:
            raise ValueError('candidate set must not be empty')
        scores=[(t,sure(xs,t,tau)) for t in ts]
    best=min(scores,key=lambda pair:(pair[1],pair[0]))
    return {'threshold':best[0],'score':best[1],'scores':scores,'estimate':[soft(x,best[0]) for x in xs]}


@checked
def bump(n,r,c,sigma=1.,alpha=.05,L=1.):
    if isinstance(n,bool) or not isinstance(n,int) or n<1:
        raise ValueError('n must be a positive integer')
    r,c,sigma,alpha,L=finite(r,'r'),finite(c,'c'),noise(sigma),finite(alpha,'alpha'),finite(L,'L')
    if not 0<r<=1 or c<=0 or L<=0 or not 0<alpha<.5 or c*sigma>L:
        raise ValueError('need 0<r<=1, c>0, L>0, 0<alpha<1/2 and c*sigma<=L')
    h=n**(-1/(2*r+1))
    kappa=1-2*alpha-math.sqrt(3)/2*c
    if h>.5 or kappa<=0:
        raise ValueError('the stated bound needs h<=1/2 and positive kappa')
    vals=[c*sigma*h**r*max(1-abs((i/n-.5)/h),0) for i in range(1,n+1)]
    kl=sum(x*x for x in vals)/(2*sigma*sigma)
    peak=c*sigma*h**r
    prob=1-2*alpha-math.sqrt(kl/2)
    return {'n':n,'r':r,'h':h,'continuum_peak':peak,'sample_peak':max(vals),'nonzero_samples':sum(x>0 for x in vals),'KL_nats':kl,'KL_upper':1.5*c*c,'probability_lower':prob,'mean_width_lower':peak*prob,'uniform_probability_lower':kappa}


def close(x,y,tol=1e-10):
    if not math.isclose(x,y,rel_tol=tol,abs_tol=tol):
        raise RuntimeError('check failed: '+repr(x)+' != '+repr(y))


def self_test():
    nchecks=0
    for xs in [[3,1,0,0],[4,2,0,2],[1],[0,0],list(range(16))]:
        cs=haar(xs)
        for a,b in zip(inverse(cs),xs): close(a,b);nchecks+=1
        close(sum(x*x for x in xs),sum(x*x for x in cs));nchecks+=1
    close(soft_risk(1,0),2*(2*Q(1)-phi(1)));nchecks+=1
    close(sure([.2,-.6,1.5,3],1,1),2.4);nchecks+=1
    close(choose_sure([.2,-.6,1.5,3],1)['threshold'],.6);nchecks+=1
    close(choose_sure([.1,.1],1)['threshold'],.1);nchecks+=1
    close(choose_sure([0,0],1)['threshold'],0);nchecks+=1
    close(choose_sure([1],1,[1,2])['threshold'],1);nchecks+=1
    for r in [.1,.5,1.]:
        for n in [16,63,64,255,256,1001]:
            out=bump(n,r,.2)
            if out['KL_nats']>out['KL_upper']+1e-12:raise RuntimeError('KL bound failed')
            nchecks+=1
    failures=[lambda:haar([]),lambda:haar([1,2,3]),lambda:haar([math.nan,0]),lambda:inverse([1,2,3]),lambda:soft(1,-1),lambda:soft(math.inf,1),lambda:sure([0],1,0),lambda:choose_sure([0],1,[]),lambda:choose_sure([0],1,[math.nan]),lambda:bump(1,.5,.2),lambda:bump(256,2,.2),lambda:bump(256,.5,2),lambda:bump(256,.5,.2,alpha=.5),lambda:haar([1e308,1e308]),lambda:sure([1e308],1e308,1),lambda:choose_sure([1e308],1)]
    for call in failures:
        try:call()
        except ValueError:pass
        else:raise RuntimeError('invalid input was accepted')
    return {'status':'passed','numeric_checks':nchecks,'rejected_invalid_inputs':len(failures),'scope':'finite floating-point classroom inputs; no arbitrary-precision or extreme-tail guarantee'}


def demo():
    y=[4,2,0,2];mu=[3,3,1,1];out=denoise(y,1)
    a=math.sqrt(2)
    return {'haar':{'input':y,'coefficients':haar(y),'inverse':inverse(haar(y)),'soft_reconstruction':out,'realized_loss':sum((x-z)**2 for x,z in zip(out,mu))/4,'expected_risk':(1+soft_risk(1,2)+2*soft_risk(1,0))/4},'SURE_finite':choose_sure([.2,-.6,1.5,3],1,[0,.5,1,2]),'SURE_breakpoints':choose_sure([.2,-.6,1.5,3],1),'selection_optimism':{'mean_min_score':2*(Q(a)-a*phi(a)),'selected_risk':2*(Q(a)+a*phi(a))},'hard_discontinuity':{'naive_score_mean':2*(phi(1)+Q(1))-4*phi(1),'risk':2*(phi(1)+Q(1)),'gap':4*phi(1)},'jump_best_m_error':[{'m':m,'squared_error':2/9*2**(-m)} for m in [1,3,10]],'bump':bump(256,.5,.2),'honest_bin_halfwidth':.25+NormalDist().inv_cdf(1-.05/8)/4}


if __name__=='__main__':
    ap=argparse.ArgumentParser(description=__doc__)
    ap.add_argument('--self-test',action='store_true')
    args=ap.parse_args()
    result=demo()
    if args.self_test:result['self_test']=self_test()
    print(json.dumps(result,ensure_ascii=False,indent=2,allow_nan=False))
