import os, numpy as np, scipy.io.wavfile as wf
from scipy.signal import resample_poly
os.environ["COQUI_TOS_AGREED"]="1"
S=os.path.dirname(os.path.abspath(__file__))

# torchcodec nao carrega neste ambiente: IO de audio via scipy
import torch, torchaudio
def _load(uri, *a, **k):
    sr, x = wf.read(uri)
    x = x.astype(np.float32)
    if np.issubdtype(x.dtype, np.integer) or x.max() > 1.5: x = x/32768.0
    if x.ndim == 1: x = x[None, :]
    else: x = x.T
    return torch.from_numpy(np.ascontiguousarray(x)), sr
def _save(uri, t, sr, *a, **k):
    y = t.detach().cpu().numpy()
    if y.ndim > 1: y = y.T if y.shape[0] < y.shape[1] else y
    wf.write(uri, sr, np.clip(y.squeeze(), -1, 1).astype(np.float32))
torchaudio.load = _load
torchaudio.save = _save
from TTS.api import TTS

FRASE="que vai levar a vitória"      # contexto: prosodia natural
ALVO ="levar"                          # so isto entra no arquivo
T0,T1=7.04,7.44                        # buraco no original

tts=TTS("tts_models/multilingual/multi-dataset/xtts_v2", progress_bar=False)
tts.tts_to_file(text=FRASE, speaker_wav=S+"/ref_alto.wav", language="pt",
                file_path=S+"/gen_raw2.wav")
print("gerado:", FRASE)

# localiza a palavra alvo dentro do que foi gerado
from faster_whisper import WhisperModel
m=WhisperModel("small", device="cpu", compute_type="int8")
seg,_=m.transcribe(S+"/gen_raw2.wav", language="pt", word_timestamps=True)
ws=[w for s in seg for w in (s.words or [])]
print("  alinhamento:", [(w.word.strip(), round(w.start,2), round(w.end,2)) for w in ws])
hit=[w for w in ws if ALVO in w.word.lower()]
if not hit: raise SystemExit(f"nao localizei '{ALVO}' no audio gerado")
g0,g1=hit[0].start,hit[0].end
print(f"  '{ALVO}' gerado em {g0:.2f}-{g1:.2f}s ({(g1-g0)*1000:.0f} ms)")

sg,gen=wf.read(S+"/gen_raw2.wav"); gen=gen.astype(float)
if gen.ndim>1: gen=gen.mean(1)
if np.issubdtype(np.dtype(gen.dtype),np.integer): gen/=32768.
piece=gen[int(g0*sg):int(g1*sg)]
sr=48000
piece=resample_poly(piece, sr, sg)
tgt=len(piece)   # duracao natural: esticar por reamostragem baixaria o tom

sr2,y=wf.read(S+"/final2.wav"); y=y.astype(float)
if y.ndim>1: y=y.mean(1)
y=y.copy()
a0=int(T0*sr); a1=a0+tgt
ref=np.sqrt(np.mean(y[int(6.0*sr):int(6.9*sr)]**2))
piece*= ref/max(np.sqrt(np.mean(piece**2)),1e-9)
xf=int(.015*sr)
new=np.r_[y[a0-xf:a0], piece, y[a1:a1+xf]]
new[:xf]=y[a0-xf:a0]*np.linspace(1,0,xf)+new[:xf]*np.linspace(0,1,xf)
new[-xf:]=new[-xf:]*np.linspace(1,0,xf)+y[a1:a1+xf]*np.linspace(0,1,xf)
y[a0-xf:a1+xf]=new
wf.write(S+"/final3.wav",sr,np.clip(y,-1,1).astype(np.float32))
print(f"encaixado em {T0}s. -> final3.wav")
