#!/usr/bin/env python3
"""Pinned Microsoft VibeVoice-ASR JSONL worker managed by Omnius.

This worker never installs dependencies or downloads weights intentionally.
The Node runtime owns setup, revision pinning, disk checks, and CUDA selection.
"""

import argparse
import json
import os
import sys
import time
from pathlib import Path


def emit(payload):
    print(json.dumps(payload, ensure_ascii=False), flush=True)


def check_only():
    required = [
        "OMNIUS_VIBEVOICE_MODEL_ID",
        "OMNIUS_VIBEVOICE_MODEL_REVISION",
        "OMNIUS_VIBEVOICE_TOKENIZER_ID",
        "OMNIUS_VIBEVOICE_TOKENIZER_REVISION",
    ]
    missing = [key for key in required if not os.environ.get(key)]
    emit({"type": "check", "ok": not missing, "missing": missing})
    return 0 if not missing else 2


def normalize_seconds(value):
    if value is None:
        return 0.0
    if isinstance(value, (int, float)):
        return float(value)
    text = str(value).strip()
    if not text:
        return 0.0
    parts = text.split(":")
    try:
        if len(parts) == 3:
            return float(parts[0]) * 3600 + float(parts[1]) * 60 + float(parts[2])
        if len(parts) == 2:
            return float(parts[0]) * 60 + float(parts[1])
        return float(text)
    except ValueError:
        return 0.0


class VibeVoiceWorker:
    def __init__(self):
        import torch
        from huggingface_hub import snapshot_download
        from vibevoice.modular.modeling_vibevoice_asr import VibeVoiceASRForConditionalGeneration
        from vibevoice.processor.vibevoice_asr_processor import VibeVoiceASRProcessor

        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is unavailable; CPU fallback is forbidden for VibeVoice ASR")
        visible = os.environ.get("CUDA_VISIBLE_DEVICES", "").strip()
        if not visible or "," in visible:
            raise RuntimeError("exactly one CUDA_VISIBLE_DEVICES entry is required")
        self.torch = torch
        self.model_id = os.environ["OMNIUS_VIBEVOICE_MODEL_ID"]
        self.model_revision = os.environ["OMNIUS_VIBEVOICE_MODEL_REVISION"]
        tokenizer_id = os.environ["OMNIUS_VIBEVOICE_TOKENIZER_ID"]
        tokenizer_revision = os.environ["OMNIUS_VIBEVOICE_TOKENIZER_REVISION"]
        model_path = snapshot_download(
            repo_id=self.model_id,
            revision=self.model_revision,
            local_files_only=True,
        )
        tokenizer_path = snapshot_download(
            repo_id=tokenizer_id,
            revision=tokenizer_revision,
            local_files_only=True,
            allow_patterns=["*.json", "*.model", "*.txt", "merges.txt", "vocab.json", "tokenizer*"],
        )
        self.processor = VibeVoiceASRProcessor.from_pretrained(
            model_path,
            language_model_pretrained_name=tokenizer_path,
            local_files_only=True,
        )
        self.model = VibeVoiceASRForConditionalGeneration.from_pretrained(
            model_path,
            dtype=torch.bfloat16,
            attn_implementation=os.environ.get("OMNIUS_VIBEVOICE_ATTN", "sdpa"),
            local_files_only=True,
        ).to("cuda:0")
        self.model.eval()
        emit({
            "type": "ready",
            "pid": os.getpid(),
            "device": torch.cuda.get_device_name(0),
            "computeCapability": ".".join(str(part) for part in torch.cuda.get_device_capability(0)),
            "totalMemoryBytes": torch.cuda.get_device_properties(0).total_memory,
            "cudaVisibleDevices": visible,
            "modelId": self.model_id,
            "modelRevision": self.model_revision,
        })

    def transcribe(self, request):
        file_path = str(request.get("file") or "")
        if not file_path or not Path(file_path).is_file():
            raise ValueError("a readable audio file is required")
        context = str(request.get("context") or "").strip() or None
        started = time.monotonic()
        inputs = self.processor(
            audio=[file_path],
            return_tensors="pt",
            padding=True,
            context_info=context,
        )
        inputs = {
            key: value.to("cuda:0") if isinstance(value, self.torch.Tensor) else value
            for key, value in inputs.items()
        }
        generation = {
            "max_new_tokens": int(os.environ.get("OMNIUS_VIBEVOICE_MAX_NEW_TOKENS", "32768")),
            "do_sample": False,
            "pad_token_id": self.processor.pad_id,
            "eos_token_id": self.processor.tokenizer.eos_token_id,
        }
        with self.torch.inference_mode():
            output_ids = self.model.generate(**inputs, **generation)
        input_length = inputs["input_ids"].shape[1]
        generated_ids = output_ids[0, input_length:]
        raw_text = self.processor.decode(generated_ids, skip_special_tokens=True)
        parsed = self.processor.post_process_transcription(raw_text)
        warnings = []
        segments = []
        speakers = []
        for item in parsed:
            text = str(item.get("text") or "").strip()
            speaker = str(item.get("speaker_id") or "").strip()
            if speaker and speaker not in speakers:
                speakers.append(speaker)
            segments.append({
                "startMs": round(normalize_seconds(item.get("start_time")) * 1000),
                "endMs": round(normalize_seconds(item.get("end_time")) * 1000),
                "text": text,
                **({"speakerId": speaker} if speaker else {}),
            })
        if not segments and raw_text.strip():
            warnings.append("VibeVoice structured output could not be parsed; raw transcription was preserved")
        text = " ".join(segment["text"] for segment in segments if segment["text"]).strip()
        if not text:
            text = raw_text.strip()
        duration_ms = max((segment["endMs"] for segment in segments), default=0)
        return {
            "text": text,
            "rawText": raw_text,
            "segments": segments,
            "speakers": speakers,
            "durationMs": duration_ms or None,
            "engineId": "vibevoice-transformers",
            "modelId": "vibevoice-asr-7b",
            "device": self.torch.cuda.get_device_name(0),
            "warnings": warnings,
            "latencyMs": round((time.monotonic() - started) * 1000),
        }


def serve():
    try:
        worker = VibeVoiceWorker()
    except Exception as exc:
        emit({"type": "error", "message": f"model activation failed: {exc}"})
        return 1
    for line in sys.stdin:
        try:
            request = json.loads(line)
            request_id = str(request.get("id") or "")
            if request.get("action") != "transcribe":
                raise ValueError(f"unknown action: {request.get('action')}")
            emit({"type": "result", "id": request_id, "result": worker.transcribe(request)})
        except Exception as exc:
            emit({"type": "error", "id": locals().get("request_id", ""), "message": str(exc)})
    return 0


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--check", action="store_true")
    parser.add_argument("--serve", action="store_true")
    args = parser.parse_args()
    if args.check:
        return check_only()
    if args.serve:
        return serve()
    parser.error("use --check or --serve")


if __name__ == "__main__":
    raise SystemExit(main())
