import numpy as np, scipy.io.wavfile as wf, sys
from scipy.linalg import solve_toeplitz

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

def ar_interp(x,s,e,p=64):
    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
    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(float)
    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'); 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)) or np.max(np.abs(sol))>3*np.max(np.abs(known)): return False
    x[s:e]=sol; return True

inp,outp=sys.argv[1],sys.argv[2]
sr,x=wf.read(inp); x=x.astype(np.float64)
if x.ndim>1: x=x.mean(1)
n=len(x); pk0=np.max(np.abs(x))

w=max(2,int(.004*sr)); e=np.sqrt(np.convolve(x*x,np.ones(w)/w,'same'))
edb=20*np.log10(np.maximum(e,1e-12))
# nivel local (250ms) -> detecta queda RELATIVA, pega dropout no meio da fala alta
W=int(.25*sr); loc=np.convolve(e,np.ones(W)/W,'same')
locdb=20*np.log10(np.maximum(loc,1e-12))
gap=(edb < locdb-18) | (edb < -52)

st,en=regions(gap)
keep=(en-st)>=int(.0015*sr); st,en=st[keep],en[keep]
# junta gaps separados por <3ms
if len(st)>1:
    ns,ne=[st[0]],[en[0]]
    for s,t in zip(st[1:],en[1:]):
        if s-ne[-1] < int(.003*sr): ne[-1]=t
        else: ns.append(s); ne.append(t)
    st,en=np.array(ns),np.array(ne)

# ruido de sala p/ buracos longos
ok=np.ones(n,bool)
for s,t in zip(st,en): ok[max(0,s-480):min(n,t+480)]=False
blk=int(.2*sr); tone=None; bl=9e9
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,tone=v,x[i:i+blk].copy()
if tone is None: tone=np.random.randn(blk)*1e-4

MAX=0.045
ni=nf=nx=0; rec=lost=0.0
for s,t in zip(st,en):
    # alarga 1ms de cada lado: cobre a rampa do corte (e' ela que estala)
    s2=max(0,s-int(.001*sr)); t2=min(n,t+int(.001*sr))
    L=(t2-s2)/sr
    if L<=MAX:
        if ar_interp(x,s2,t2): ni+=1; rec+=L
        else: nx+=1
    else:
        nf+=1; lost+=L; need=t2-s2
        fill=np.tile(tone,int(np.ceil(need/len(tone))))[:need]*0.9
        f=int(min(.010*sr,need//4))
        if f>2: fill[:f]*=np.linspace(0,1,f); fill[-f:]*=np.linspace(1,0,f)
        x[s2:t2]=fill

print(f"falhas de transmissao detectadas : {len(st)}")
print(f"  COMPLETADAS por interpolacao   : {ni}   ({rec:.2f}s de fala reconstruida)")
print(f"  longas demais (sem conteudo)   : {nf}   ({lost:.2f}s -> ruido de sala)")
if nx: print(f"  sem contexto suficiente        : {nx}")
print("nenhum outro processamento aplicado: sem NR, sem EQ, sem compressor")
x=np.clip(x,-1,1); x*= pk0/max(np.max(np.abs(x)),1e-9)
wf.write(outp,sr,x.astype(np.float32))
