import argparse, json, sys
from faster_whisper import WhisperModel


def run_model(model_name, device, compute_type, audio, language, beam_size, patience, prompt, vad_filter, cpu_threads):
    kwargs = {"device": device, "compute_type": compute_type}
    if device == "cpu" and cpu_threads:
        kwargs["cpu_threads"] = cpu_threads
    model = WhisperModel(model_name, **kwargs)
    segments, info = model.transcribe(
        audio,
        language=language,
        beam_size=beam_size,
        patience=patience,
        vad_filter=vad_filter,
        initial_prompt=prompt or None,
        condition_on_previous_text=False,
        word_timestamps=False,
        temperature=0.0,
    )
    segments = list(segments)
    text = " ".join((seg.text or "").strip() for seg in segments).strip()
    # Surface decoder quality signals to the server so weak/static audio can be
    # quarantined instead of becoming incident evidence.
    weighted = [(float(getattr(seg, "avg_logprob", 0.0)), max(0.01, float(getattr(seg, "end", 0.0)) - float(getattr(seg, "start", 0.0)))) for seg in segments]
    avg_logprob = (sum(v*w for v,w in weighted) / sum(w for _,w in weighted)) if weighted else None
    no_speech = [float(getattr(seg, "no_speech_prob", 0.0)) for seg in segments]
    no_speech_prob = (sum(no_speech) / len(no_speech)) if no_speech else None
    return {
        "text": text,
        "avg_logprob": avg_logprob,
        "no_speech_prob": no_speech_prob,
        "language": getattr(info, "language", language),
        "language_probability": getattr(info, "language_probability", None),
        "actual_model": model_name,
        "actual_device": device,
        "actual_compute_type": compute_type,
    }


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--audio", required=True)
    ap.add_argument("--model", default="distil-large-v3")
    ap.add_argument("--language", default="en")
    ap.add_argument("--compute-type", default="int8_float16")
    ap.add_argument("--device", default="cuda")
    ap.add_argument("--beam-size", type=int, default=5)
    ap.add_argument("--patience", type=float, default=1.5)
    ap.add_argument("--prompt", default="")
    ap.add_argument("--vad-filter", action="store_true")
    ap.add_argument("--cpu-threads", type=int, default=4)
    ap.add_argument("--fallback-model", default="small.en")
    ap.add_argument("--fallback-device", default="cpu")
    ap.add_argument("--fallback-compute-type", default="int8")
    a = ap.parse_args()

    try:
        result = run_model(a.model, a.device, a.compute_type, a.audio, a.language,
                           a.beam_size, a.patience, a.prompt, a.vad_filter, a.cpu_threads)
        result["fallback_used"] = False
        print(json.dumps(result, ensure_ascii=False))
        return
    except Exception as primary_exc:
        # A 4 GB GTX 1650 is useful, but Windows CUDA/cuDNN may not be installed yet.
        # Never punish the production scanner machine by falling back to large-v3 on CPU.
        if not a.fallback_model:
            raise
        try:
            result = run_model(a.fallback_model, a.fallback_device, a.fallback_compute_type,
                               a.audio, a.language, min(a.beam_size, 5), min(a.patience, 1.5),
                               a.prompt, a.vad_filter, a.cpu_threads)
            result["fallback_used"] = True
            result["primary_error"] = str(primary_exc)
            print(json.dumps(result, ensure_ascii=False))
            return
        except Exception:
            print(f"Primary GPU transcription failed: {primary_exc}", file=sys.stderr)
            raise


if __name__ == "__main__":
    try:
        main()
    except Exception as exc:
        print(str(exc), file=sys.stderr)
        raise
