#!/usr/bin/env python3
"""Create an SRT file from audio or video with Groq-hosted Whisper."""

import argparse
import json
import math
import mimetypes
import os
import subprocess
import sys
import tempfile
import time
import uuid
from pathlib import Path
from urllib import request as url_request
from urllib.error import HTTPError, URLError


GROQ_URL = "https://api.groq.com/openai/v1/audio/transcriptions"
DEFAULT_MODEL = os.environ.get("GROQ_WHISPER_MODEL", "whisper-large-v3")
MAX_BYTES = 24 * 1024 * 1024
OPUS_BITRATE = "32k"
CHUNK_SECONDS = 1500
PROMPT_MAX_CHARS = 800
TAIL_MAX_CHARS = 200
USER_AGENT = "transcribe-groq/1.0"

def run(command: list[str]) -> subprocess.CompletedProcess:
    result = subprocess.run(command, capture_output=True, text=True)
    if result.returncode != 0:
        sys.exit(f"Command failed: {' '.join(command)}\n{result.stderr}")
    return result


def duration(path: Path) -> float:
    result = run([
        "ffprobe",
        "-v",
        "error",
        "-show_entries",
        "format=duration",
        "-of",
        "csv=p=0",
        str(path),
    ])
    return float(result.stdout.strip())


def encode_ogg(
    source: Path,
    destination: Path,
    start: float | None = None,
    length: float | None = None,
) -> None:
    command = ["ffmpeg", "-y", "-loglevel", "error"]
    if start is not None:
        command += ["-ss", str(start)]
    if length is not None:
        command += ["-t", str(length)]
    command += [
        "-i",
        str(source),
        "-c:a",
        "libopus",
        "-b:a",
        OPUS_BITRATE,
        "-ar",
        "16000",
        "-ac",
        "1",
        "-vn",
        "-application",
        "voip",
        str(destination),
    ]
    run(command)


def post_multipart(
    url: str,
    headers: dict[str, str],
    fields: dict[str, str],
    file_path: Path,
) -> dict:
    boundary = uuid.uuid4().hex
    crlf = b"\r\n"
    body = bytearray()

    for key, value in fields.items():
        body += f"--{boundary}\r\n".encode()
        body += f'Content-Disposition: form-data; name="{key}"\r\n\r\n'.encode()
        body += str(value).encode("utf-8") + crlf

    filename = file_path.name
    content_type = mimetypes.guess_type(filename)[0] or "application/octet-stream"
    body += f"--{boundary}\r\n".encode()
    body += (
        f'Content-Disposition: form-data; name="file"; filename="{filename}"\r\n'
    ).encode()
    body += f"Content-Type: {content_type}\r\n\r\n".encode()
    body += file_path.read_bytes()
    body += crlf
    body += f"--{boundary}--\r\n".encode()

    request_headers = dict(headers)
    request_headers["Content-Type"] = f"multipart/form-data; boundary={boundary}"
    request_headers["User-Agent"] = USER_AGENT
    request_headers["Accept"] = "application/json"

    for attempt in range(1, 4):
        request = url_request.Request(
            url,
            data=bytes(body),
            headers=request_headers,
            method="POST",
        )
        try:
            with url_request.urlopen(request, timeout=600) as response:
                return json.loads(response.read().decode("utf-8"))
        except HTTPError as error:
            retryable = error.code == 429 or 500 <= error.code < 600
            if not retryable or attempt == 3:
                detail = error.read().decode("utf-8", "replace")
                sys.exit(f"Groq returned HTTP {error.code}: {detail}")
        except (URLError, OSError) as error:
            if attempt == 3:
                sys.exit(f"Groq connection failed after 3 attempts: {error}")

        print(f"  Retrying Groq request ({attempt}/3)...")
        time.sleep(attempt * 2)

    raise RuntimeError("unreachable")


def transcribe_chunk(
    path: Path,
    api_key: str,
    model: str,
    language: str,
    prompt: str,
) -> list[dict]:
    fields = {
        "model": model,
        "response_format": "verbose_json",
        "temperature": "0",
    }
    if language:
        fields["language"] = language
    if prompt:
        fields["prompt"] = prompt[:PROMPT_MAX_CHARS]

    data = post_multipart(
        GROQ_URL,
        {"Authorization": f"Bearer {api_key}"},
        fields,
        path,
    )
    if not isinstance(data, dict) or not isinstance(data.get("segments"), list):
        raise ValueError("Groq response must contain a segments list.")
    return clean_segments(data["segments"])


def clean_segments(segments: list[dict]) -> list[dict]:
    """Validate timestamps and strip whitespace without guessing at hallucinations."""
    cleaned = []
    for segment in segments:
        if not isinstance(segment, dict) or not isinstance(segment.get("text"), str):
            raise ValueError("Invalid transcript segment text.")
        start, end = segment.get("start"), segment.get("end")
        if any(isinstance(value, bool) or not isinstance(value, (int, float))
               or not math.isfinite(value) for value in (start, end)):
            raise ValueError("Invalid transcript segment timestamps.")
        if start < 0 or end <= start or round(end * 1000) <= round(start * 1000):
            raise ValueError("Transcript segment must have positive duration.")
        text = segment["text"].strip()
        if text:
            cleaned.append({"start": float(start), "end": float(end), "text": text})
    return cleaned


def format_timestamp(value: float) -> str:
    total_milliseconds = max(0, round(value * 1000))
    hours, remainder = divmod(total_milliseconds, 3_600_000)
    minutes, remainder = divmod(remainder, 60_000)
    seconds, milliseconds = divmod(remainder, 1_000)
    return f"{hours:02d}:{minutes:02d}:{seconds:02d},{milliseconds:03d}"


def write_srt(segments: list[dict], path: Path) -> None:
    with path.open("w", encoding="utf-8") as output:
        for index, segment in enumerate(segments, 1):
            output.write(
                f"{index}\n"
                f"{format_timestamp(segment['start'])} --> "
                f"{format_timestamp(segment['end'])}\n"
                f"{segment['text']}\n\n"
            )


def build_prompt(base: str, tail: str) -> str:
    if not tail:
        return base[:PROMPT_MAX_CHARS]
    budget = PROMPT_MAX_CHARS - len(base) - 1
    if budget <= 0:
        return base[:PROMPT_MAX_CHARS]
    return f"{base} {tail[-budget:]}"[:PROMPT_MAX_CHARS]


def parse_arguments() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="Create an SRT file from audio or video with Groq Whisper."
    )
    parser.add_argument("input", type=Path)
    parser.add_argument("-o", "--output", type=Path)
    parser.add_argument(
        "--language",
        default=os.environ.get("TRANSCRIPTION_LANGUAGE", "en"),
        help="ISO-639-1 language code for the source audio; use an empty value to detect it",
    )
    parser.add_argument(
        "--prompt",
        default="",
        help="Optional names or terms that help Whisper spell domain-specific words",
    )
    parser.add_argument("--model", default=DEFAULT_MODEL)
    parser.add_argument("--keep-tmp", action="store_true")
    parser.add_argument("--force", action="store_true", help="Replace an existing SRT only after a valid nonempty result")
    return parser.parse_args()


def main() -> int:
    arguments = parse_arguments()
    api_key = os.environ.get("GROQ_API_KEY")
    if not api_key:
        sys.exit("Set GROQ_API_KEY before running this script.")

    source = arguments.input.expanduser().resolve()
    if not source.is_file():
        sys.exit(f"Input file not found: {source}")

    destination = (
        arguments.output.expanduser().resolve()
        if arguments.output
        else source.with_suffix(".srt")
    )
    if destination == source:
        sys.exit("Output must differ from the input file.")
    if destination.exists() and not arguments.force:
        sys.exit("Output already exists. Use --force to replace it after validation.")
    destination.parent.mkdir(parents=True, exist_ok=True)

    total_seconds = duration(source)
    if not math.isfinite(total_seconds) or total_seconds <= 0:
        sys.exit("Input duration must be a positive finite number.")
    print(f"Duration: {total_seconds / 60:.1f} minutes")
    temporary_directory = Path(tempfile.mkdtemp(prefix="transcribe_groq_"))
    print(f"Temporary directory: {temporary_directory}")

    try:
        all_segments: list[dict] = []
        previous_tail = ""
        chunk_count = math.ceil(total_seconds / CHUNK_SECONDS)

        for index in range(chunk_count):
            start = index * CHUNK_SECONDS
            length = min(CHUNK_SECONDS, total_seconds - start)
            chunk = temporary_directory / f"chunk-{index + 1:03d}.ogg"
            encode_ogg(source, chunk, start=start if index else None, length=length)

            size = chunk.stat().st_size
            print(
                f"[{index + 1}/{chunk_count}] {start / 60:.1f}-"
                f"{(start + length) / 60:.1f} minutes, {size / 1024 / 1024:.1f} MB"
            )
            if size > MAX_BYTES:
                sys.exit("Compressed chunk exceeds 24 MB. Lower CHUNK_SECONDS and rerun.")

            prompt = build_prompt(arguments.prompt, previous_tail)
            segments = transcribe_chunk(
                chunk,
                api_key,
                arguments.model,
                arguments.language,
                prompt,
            )
            for segment in segments:
                segment["start"] += start
                segment["end"] += start
            all_segments.extend(segments)
            previous_tail = " ".join(
                str(segment.get("text", "")) for segment in segments[-3:]
            ).strip()[-TAIL_MAX_CHARS:]

        cleaned = clean_segments(all_segments)
        if not cleaned:
            sys.exit("No usable transcript segments; existing output was preserved.")
        with tempfile.NamedTemporaryFile(dir=destination.parent, prefix=destination.name + '.', suffix='.tmp', delete=False) as handle:
            temporary_output = Path(handle.name)
        try:
            write_srt(cleaned, temporary_output)
            if arguments.force:
                os.replace(temporary_output, destination)
            else:
                # Atomic no-clobber publication, even if another writer won the race.
                os.link(temporary_output, destination)
        finally:
            temporary_output.unlink(missing_ok=True)

        print(f"{len(all_segments)} raw -> {len(cleaned)} clean segments")
        print(f"Saved: {destination}")
        return 0
    finally:
        if arguments.keep_tmp:
            print(f"Kept temporary files: {temporary_directory}")
        else:
            for path in temporary_directory.iterdir():
                path.unlink()
            temporary_directory.rmdir()


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