"""
faster_whisper_server.py
========================
FastAPI transcription service for Windows PC with RTX 5080.
Receives audio files via multipart POST, returns transcript JSON.

Dependencies (install once on Windows, Python 3.11+):
    pip install faster-whisper fastapi uvicorn python-multipart

Start:
    uvicorn faster_whisper_server:app --host 0.0.0.0 --port 8765

Windows Firewall: allow inbound TCP on port 8765 from Ubuntu server IP.
"""

import gc
import os
import sys
import tempfile
import threading
from pathlib import Path

import numpy as np
from faster_whisper import WhisperModel, BatchedInferencePipeline
from fastapi import FastAPI, File, Form, UploadFile
from fastapi.responses import JSONResponse

MODEL_NAME   = os.environ.get("WHISPER_MODEL", "large-v3")
DEVICE       = os.environ.get("WHISPER_DEVICE", "cuda")
COMPUTE_TYPE = os.environ.get("WHISPER_COMPUTE_TYPE", "float16")
# >1 enables BatchedInferencePipeline (GPU chunk parallelism). 0/1 = plain model.
BATCH_SIZE   = int(os.environ.get("WHISPER_BATCH_SIZE", "0"))

# Named whisper_model, NOT model: the /transcribe endpoint has a `model` form
# param that would shadow this global and crash every request.
whisper_model = None
batched_model = None
_paused = False


def _load_model() -> WhisperModel:
    global whisper_model, batched_model
    if whisper_model is None:
        print(f"Loading Whisper model: {MODEL_NAME} on {DEVICE} ({COMPUTE_TYPE})...")
        whisper_model = WhisperModel(MODEL_NAME, device=DEVICE, compute_type=COMPUTE_TYPE)
        if BATCH_SIZE > 1:
            batched_model = BatchedInferencePipeline(model=whisper_model)
            print(f"Batched pipeline enabled (batch_size={BATCH_SIZE}).")
        print("Model loaded.")
    return whisper_model


def _unload_model_locked() -> None:
    """Free GPU VRAM. Caller must hold _transcribe_lock."""
    global whisper_model, batched_model
    if whisper_model is not None:
        whisper_model = None
        batched_model = None
        gc.collect()
        print("Model unloaded — VRAM freed for gaming.")


_load_model()

# Self-test: run real inference at startup so CUDA/cuDNN failures surface at
# boot instead of as a 500 on the first request. vad_filter must stay off here
# or the silent clip is skipped before any kernel runs.
print("Running startup inference self-test...")
try:
    _segments, _info = whisper_model.transcribe(
        np.zeros(16000, dtype=np.float32),  # 1s of silence @ 16kHz
        language="en",
        beam_size=1,
        vad_filter=False,
    )
    list(_segments)
    print("Self-test OK — inference path working.")
except Exception:
    import traceback
    print("FATAL: startup inference self-test failed:", file=sys.stderr)
    traceback.print_exc()
    sys.exit(1)

# One request at a time on the GPU; concurrent transcriptions risk OOM.
_transcribe_lock = threading.Lock()

app = FastAPI(title="faster-whisper-server")


@app.get("/health")
def health():
    return {
        "status": "paused" if _paused else "ok",
        "model": MODEL_NAME,
        "device": DEVICE,
        "model_loaded": whisper_model is not None,
    }


@app.post("/pause")
def pause():
    """Stop accepting new jobs and free VRAM (immediately if idle, else after
    the in-flight transcription finishes). For gaming on this PC."""
    global _paused
    _paused = True
    if _transcribe_lock.acquire(blocking=False):
        try:
            _unload_model_locked()
        finally:
            _transcribe_lock.release()
        return {"status": "paused", "model_unloaded": True}
    return {
        "status": "paused",
        "model_unloaded": False,
        "note": "current transcription will finish first, then VRAM is freed",
    }


@app.post("/resume")
def resume():
    global _paused
    _paused = False
    with _transcribe_lock:
        _load_model()
    return {"status": "ok", "model_loaded": True}


@app.post("/transcribe")
def transcribe(
    file: UploadFile = File(...),
    model: str = Form(default=MODEL_NAME),
    language: str = Form(default="en"),
    beam_size: int = Form(default=1),
):
    """Transcribe an audio file. Returns {text, language, segments}.

    Deliberately a sync handler: FastAPI runs it in the threadpool, keeping
    the event loop free so /pause and /health respond instantly even while
    a long transcription holds the GPU. An async handler here would block
    the whole server for the duration of the inference.
    """
    if _paused:
        return JSONResponse({"error": "paused"}, status_code=503)

    audio_bytes = file.file.read()

    with tempfile.NamedTemporaryFile(
        suffix=Path(file.filename or "audio.wav").suffix or ".wav",
        delete=False,
    ) as tmp:
        tmp.write(audio_bytes)
        tmp_path = tmp.name

    try:
        with _transcribe_lock:
            _load_model()
            engine = batched_model if BATCH_SIZE > 1 else whisper_model
            kwargs = dict(
                language=language,
                beam_size=beam_size,
                vad_filter=True,          # skip silence
                vad_parameters={"min_silence_duration_ms": 500},
            )
            if BATCH_SIZE > 1:
                kwargs["batch_size"] = BATCH_SIZE
            segments_iter, info = engine.transcribe(tmp_path, **kwargs)
            segments = list(segments_iter)
            if _paused:  # pause arrived mid-job: free VRAM now that we're done
                _unload_model_locked()
        text = " ".join(s.text.strip() for s in segments)
        segment_data = [
            {"start": s.start, "end": s.end, "text": s.text}
            for s in segments
        ]
    except Exception as exc:
        import traceback
        os.unlink(tmp_path)
        return JSONResponse(
            {"error": str(exc), "traceback": traceback.format_exc()},
            status_code=500,
        )
    finally:
        try:
            os.unlink(tmp_path)
        except FileNotFoundError:
            pass

    return JSONResponse({
        "text":     text,
        "language": info.language,
        "segments": segment_data,
    })


if __name__ == "__main__":
    import uvicorn
    uvicorn.run(app, host="0.0.0.0", port=8765)
