import numpy as np, scipy.io.wavfile as wf, sys
from scipy.linalg import solve_toeplitz
S=sys.argv[1] if len(sys.argv)>1 else "."

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(float); 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

sr,a=wf.read(S+"/adobe.wav"); a=a.astype(float)
if a.ndim>1: a=a.mean(1)
_,o=wf.read(S+"/orig.wav"); o=o.astype(float)
if o.ndim>1: o=o.mean(1)
n=min(len(a),len(o)); a=a[:n].copy(); o=o[:n]

def env(x,ms):
    w=max(2,int(ms/1000*sr)); return np.sqrt(np.convolve(x*x,np.ones(w)/w,'same'))

# --- falhas reais no ORIGINAL (referencia de onde houve defeito) ---
eo=env(o,4.0); eodb=20*np.log10(np.maximum(eo,1e-12))
W=int(.25*sr); lo_=20*np.log10(np.maximum(np.convolve(eo,np.ones(W)/W,'same'),1e-12))
so,en_=regions((eodb<lo_-18)|(eodb<-52))
k=(en_-so)>=int(.004*sr); so,en_=so[k],en_[k]
bad=np.zeros(n,bool)
for s,e in zip(so,en_): bad[s:min(e,n)]=True

# --- 1. LEITO DE RUIDO DE SALA, tirado do proprio evento ---
okm=~bad
blk=int(.35*sr); cands=[]
for i in range(0,n-blk,blk//3):
    if okm[i:i+blk].all():
        v=np.sqrt(np.mean(o[i:i+blk]**2))
        if 1e-6<v: cands.append((v,i))
cands.sort()
pieces=[o[i:i+blk].copy() for _,i in cands[:6]] or [np.random.randn(blk)*1e-4]
xf=int(.05*sr); bed=np.zeros(0)
while len(bed)<n+blk:
    p=pieces[np.random.randint(len(pieces))].copy()
    if len(bed)==0: bed=p
    else:
        h=np.minimum(xf,min(len(bed),len(p))//2)
        bed[-h:]=bed[-h:]*np.linspace(1,0,h)+p[:h]*np.linspace(0,1,h)
        bed=np.r_[bed,p[h:]]
bed=bed[:n]
import os
TONE_DB=float(os.environ.get("TONE_DB","-46"))
bed*= 10**(TONE_DB/20)/max(np.sqrt(np.mean(bed**2)),1e-12)
y=a+bed
print(f"1. leito de sala   : ruido do proprio evento a {TONE_DB:.0f} dBFS sob todo o audio")
print(f"                     (mata o silencio digital nas 59 pausas entre palavras)")

# --- 2. completa os defeitos CURTOS que a Adobe deixou ---
ey=env(y,4.0); eydb=20*np.log10(np.maximum(ey,1e-12))
ly=20*np.log10(np.maximum(np.convolve(ey,np.ones(W)/W,'same'),1e-12))
sy,ey2=regions((eydb<ly-18)|(eydb<-52))
k=(ey2-sy)>=int(.004*sr); sy,ey2=sy[k],ey2[k]
ni=0; rec=0.0; grandes=[]
for s,e in zip(sy,ey2):
    j0=max(0,s-int(.03*sr)); j1=min(n,e+int(.03*sr))
    if bad[j0:j1].mean()<=0.25: continue          # pausa natural: nao mexe
    L=(e-s)/sr
    if L<=0.040:
        s2=max(0,s-int(.001*sr)); e2=min(n,e+int(.001*sr))
        if ar_interp(y,s2,e2): ni+=1; rec+=L
    else: grandes.append((s/sr,L))
print(f"2. defeitos curtos : {ni} completados por interpolacao ({rec*1000:.0f} ms)")
print(f"3. buracos grandes : {len(grandes)} sem conteudo ({sum(l for _,l in grandes):.2f}s) - so regravando")

y=np.clip(y,-1,1)
wf.write(S+"/final.wav",sr,y.astype(np.float32))
np.save(S+"/grandes.npy", np.array(grandes))
