import numpy as np, scipy.io.wavfile as wf, sys
from scipy.linalg import solve_toeplitz
from scipy.signal import butter, lfilter, filtfilt

SR_T = -50.0     # limiar de dropout (dBFS, env RMS 5ms)
P    = 64        # ordem do modelo AR
MAX_INTERP = 0.040   # ate 40ms reconstroi por AR; acima disso e' conteudo perdido

def regions(mask):
    d=np.diff(mask.astype(np.int8)); st=np.flatnonzero(d==1)+1; en=np.flatnonzero(d==-1)+1
    if mask[0]: st=np.r_[0,st]
    if mask[-1]: en=np.r_[en,len(mask)]
    return st,en

def ar_interp(x, s, e, p=P):
    """Interpolacao autorregressiva (Janssen): estima AR do contexto e resolve
    as amostras faltantes que minimizam o erro de predicao."""
    L=e-s
    ctx=max(4*p, 3*L)
    a0=max(0,s-ctx); b1=min(len(x), e+ctx)
    if s-a0 < p+1 or b1-e < p+1: return False
    seg=x[a0:b1].astype(np.float64); ms=s-a0; me=e-a0; N=len(seg)
    known=np.r_[seg[:ms], seg[me:]]
    if len(known) < 3*p or not np.any(known): return False
    r=np.correlate(known,known,'full')[len(known)-1:len(known)+p].astype(np.float64)
    if r[0]<=0: return False
    r[0]*=1.0001
    try: a=solve_toeplitz((r[:p],r[:p]), r[1:p+1])
    except Exception: return False
    coef=np.r_[1.0,-a]
    segz=seg.copy(); segz[ms:me]=0.0
    c=np.convolve(segz,coef,'valid')             # erro de predicao com buraco zerado
    lo=ms-p
    if lo<0 or ms+L>len(c): return False
    rhs=-np.correlate(c[lo:ms+L], coef,'valid')[:L]
    rc=np.correlate(coef,coef,'full')
    col=np.zeros(L); m=min(L,p+1); col[:m]=rc[p:p+m]
    col[0]*=1.0001
    try: sol=solve_toeplitz((col,col), rhs)
    except Exception: return False
    if not np.all(np.isfinite(sol)): return False
    lim=4*np.max(np.abs(known))
    if np.max(np.abs(sol))>lim: return False
    x[s:e]=sol
    return True

def main(inp, outp):
    sr,x0=wf.read(inp); x=x0.astype(np.float64)
    if x.ndim>1: x=x.mean(1)
    n=len(x); orig=x.copy()

    w=int(0.005*sr); k=np.ones(w)/w
    env=np.sqrt(np.convolve(x*x,k,'same'))
    envdb=20*np.log10(np.maximum(env,1e-12))
    st,en=regions(envdb<SR_T)
    keep=(en-st)>=int(0.003*sr); st,en=st[keep],en[keep]

    # ruido de sala: menor RMS entre trechos de 200ms que NAO sao dropout
    ok=np.ones(n,bool)
    for s,e in zip(st,en): ok[max(0,s-int(.01*sr)):min(n,e+int(.01*sr))]=False
    blk=int(0.2*sr); best=None; bl=1e9
    for i in range(0,n-blk,blk//2):
        if ok[i:i+blk].all():
            v=np.sqrt(np.mean(x[i:i+blk]**2))
            if 1e-6<v<bl: bl,best=v,x[i:i+blk].copy()
    tone = best if best is not None else np.random.randn(blk)*1e-4

    rec=lost=0; falhas=0
    for s,e in zip(st,en):
        L=(e-s)/sr
        if L<=MAX_INTERP:
            if ar_interp(x,s,e): rec+=L
            else: falhas+=1
        else:
            lost+=L
            need=e-s
            reps=int(np.ceil(need/len(tone)))
            fill=np.tile(tone,reps)[:need]*0.9
            f=int(min(0.008*sr, need//4))
            if f>2:
                fill[:f]*=np.linspace(0,1,f); fill[-f:]*=np.linspace(1,0,f)
            x[s:e]=fill

    # de-click residual: transientes HF isolados -> micro-interpolacao de 2ms
    b,a=butter(4, 6000/(sr/2),'high'); hp=filtfilt(b,a,x)
    loc=np.sqrt(np.convolve(hp*hp,np.ones(int(0.05*sr))/int(0.05*sr),'same'))
    ratio=np.abs(hp)/np.maximum(loc,1e-9)
    idx=np.flatnonzero(ratio>14)
    dc=0
    if len(idx):
        for g in np.split(idx, np.flatnonzero(np.diff(idx)>int(0.005*sr))+1):
            s=max(0,g[0]-int(0.0015*sr)); e=min(n,g[-1]+int(0.0015*sr))
            if e-s < int(0.012*sr) and ar_interp(x,s,e,p=48): dc+=1

    x=np.clip(x,-1,1)
    wf.write(outp, sr, x.astype(np.float32))
    print(f"dropouts encontrados : {len(st)}")
    print(f"reconstruidos por AR : {(en-st)[ (en-st)<=MAX_INTERP*sr ].shape[0]} gaps  ({rec:.2f}s de audio recomposto)")
    print(f"nao recuperaveis     : {(en-st)[ (en-st)>MAX_INTERP*sr ].shape[0]} gaps  ({lost:.2f}s -> preenchidos com ruido de sala)")
    print(f"estalos removidos    : {dc}")
    if falhas: print(f"(gaps sem contexto suficiente: {falhas})")

main(sys.argv[1], sys.argv[2])
