"""
============================================================
HalaVoice — SILMA TTS RunPod Serverless Worker
============================================================
Receives synthesis jobs from the HalaVoice platform
(server/services/runpod-tts.ts, serverless mode) and runs
SILMA TTS voice cloning on the GPU.

Input (job["input"]):
  {
    "text": "...",                  # text to synthesize (required)
    "reference_audio_b64": "...",   # reference voice WAV, base64 (required)
    "reference_text": "...",        # transcript of the reference audio
    "speed": 1.0,                   # speech speed
    "seed": null                    # optional random seed
  }

Output:
  {
    "audio_base64": "...",          # synthesized WAV, base64
    "sample_rate": 24000,
    "duration": 3.2                 # seconds
  }
============================================================
"""

import base64
import os
import tempfile
import time

import runpod

# Loaded once per worker (cold start), reused across jobs
_MODEL = None


def get_model():
    global _MODEL
    if _MODEL is None:
        from silma_tts.api import SilmaTTS
        print("[Worker] Loading SILMA TTS model onto GPU...")
        t0 = time.time()
        _MODEL = SilmaTTS()
        print(f"[Worker] Model loaded in {time.time() - t0:.1f}s")
    return _MODEL


def handler(job):
    job_input = job.get("input") or {}

    text = (job_input.get("text") or "").strip()
    ref_b64 = job_input.get("reference_audio_b64") or ""
    if not text:
        return {"error": "text is required"}
    if not ref_b64:
        return {"error": "reference_audio_b64 is required"}

    reference_text = job_input.get("reference_text") or ""
    speed = float(job_input.get("speed") or 1.0)
    seed = job_input.get("seed")

    ref_path = out_path = None
    try:
        with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as f:
            f.write(base64.b64decode(ref_b64))
            ref_path = f.name
        out_path = ref_path.replace(".wav", "_out.wav")

        model = get_model()
        t0 = time.time()
        wav, sr, _spec = model.infer(
            ref_file=ref_path,
            ref_text=reference_text,
            gen_text=text,
            file_wave=out_path,
            seed=seed,
            speed=speed,
        )
        elapsed = time.time() - t0

        with open(out_path, "rb") as f:
            audio_b64 = base64.b64encode(f.read()).decode("ascii")

        duration = len(wav) / sr if sr else 0
        print(f"[Worker] Synthesized {len(text)} chars in {elapsed:.2f}s (audio {duration:.2f}s)")

        return {
            "audio_base64": audio_b64,
            "sample_rate": sr or 24000,
            "duration": duration,
        }
    except Exception as e:
        return {"error": f"synthesis failed: {e}"}
    finally:
        for p in (ref_path, out_path):
            if p and os.path.exists(p):
                try:
                    os.unlink(p)
                except OSError:
                    pass


runpod.serverless.start({"handler": handler})
