#!/usr/bin/env python3
"""
G3: Gemini Direct Transcribe — Audio → TXT + SRT (1 API call)
Thay thế G3 (Chirp) + G4 (post-process) cũ.

Model: Gemini 3 Flash Preview qua OpenRouter
Output: transcript có timestamp inline → tự động extract .txt + .srt

Requirements:
    pip install openai

Environment:
    export OPENROUTER_API_KEY="sk-or-..."

Usage:
    python3 gemini_transcribe.py chunks/ -o transcripts/
    python3 gemini_transcribe.py chunks/012_chunk_000.wav -o transcripts/
"""

import argparse
import base64
import os
import re
import sys
import time
from pathlib import Path

try:
    from openai import OpenAI
except ImportError:
    print("❌ openai SDK not installed. Run: pip install openai")
    sys.exit(1)

OPENROUTER_BASE_URL = "https://openrouter.ai/api/v1"
DEFAULT_MODEL = "google/gemini-3-flash-preview"

# ─── System Prompt ───
SYSTEM_PROMPT = """Bạn là chuyên gia transcribe tiếng Myanmar cho bài giảng Phật giáo.

**Nhiệm vụ:** Nghe audio và tạo transcript chính xác với timestamp.

**Định dạng output BẮT BUỘC:**
Mỗi dòng theo format: [MM:SS.mmm] nội dung

Ví dụ:
[00:00.000] ကဲ ဦးဇင်းတို့ စာမျက်နှာ ၇၂ နော်။
[00:05.500] စာမျက်နှာ ၇၂ ထီးတောင်ဝှေး ဓားစသည်...
[00:10.200] တရားတော် နာယူသည့် ယောဂဝေဒနာ...

**Quy tắc:**
1. Timestamp là thời điểm BẮT ĐẦU của mỗi câu/phrase, format [MM:SS.mmm]
2. Mỗi dòng là 1 câu hoặc 1 ý hoàn chỉnh (5-15 giây nói)
3. Dòng đầu tiên LUÔN là [00:00.000], dòng cuối phải ≤ thời lượng audio
4. Phân bố đều theo thời gian thực — KHÔNG dồn timestamp về cuối
5. Giữ nguyên ngữ pháp và trợ từ Myanmar (၏, သာလျှင်, ပေါ့နော်, နော်, ပေါ့)
6. Từ Pāli giữ nguyên chính tả
7. Phân biệt đúng các từ đồng âm dựa trên ngữ cảnh bài giảng
8. KHÔNG thêm bất kỳ giải thích, tiêu đề, hay meta-data nào
9. Chỉ trả về các dòng transcript có timestamp, không thêm gì khác"""


def transcribe_audio(
    audio_path: str,
    api_key: str,
    model: str = DEFAULT_MODEL,
    max_retries: int = 3
) -> tuple[str, dict | None]:
    """
    Gửi audio tới Gemini, nhận transcript có timestamp.
    
    Returns:
        (raw_text, usage_dict) — usage_dict = None nếu không có
    """
    with open(audio_path, "rb") as f:
        audio_bytes = f.read()

    audio_b64 = base64.b64encode(audio_bytes).decode()
    duration = len(audio_bytes) / 16000 / 2
    size_mb = len(audio_bytes) / 1024 / 1024

    client = OpenAI(api_key=api_key, base_url=OPENROUTER_BASE_URL)

    user_prompt = (
        f"Audio bài giảng Phật giáo tiếng Myanmar dài {duration:.0f} giây.\n"
        "Hãy transcribe với timestamp theo đúng định dạng [MM:SS.mmm]."
    )

    for attempt in range(max_retries):
        if attempt > 0:
            wait = 2 ** attempt
            print(f"   🔄 Retry {attempt}/{max_retries} after {wait}s...")
            time.sleep(wait)

        try:
            t0 = time.time()
            response = client.chat.completions.create(
                model=model,
                messages=[
                    {"role": "system", "content": SYSTEM_PROMPT},
                    {"role": "user", "content": [
                        {"type": "text", "text": user_prompt},
                        {"type": "input_audio", "input_audio": {
                            "data": audio_b64, "format": "wav"
                        }},
                    ]},
                ],
                temperature=0.1,
                max_tokens=8192,
                extra_headers={
                    "HTTP-Referer": "https://github.com/openclaw/pipeline",
                    "X-Title": "Myanmar Audio Transcription",
                },
            )
            elapsed = time.time() - t0
            result = response.choices[0].message.content.strip()
            usage = None
            if hasattr(response, 'usage') and response.usage:
                usage = {
                    "prompt_tokens": response.usage.prompt_tokens,
                    "completion_tokens": response.usage.completion_tokens,
                    "total_tokens": response.usage.total_tokens,
                }
            chars = len(result)
            tok_info = f" | {usage['total_tokens']} tok" if usage else ""
            print(f"   ✅ {elapsed:.1f}s | {duration:.0f}s audio → {chars} chars{tok_info}")
            return result, usage

        except Exception as e:
            print(f"   ❌ API error: {e}")

    return "", None


def parse_timestamped_transcript(raw: str, duration: float) -> list[dict]:
    """
    Parse [MM:SS.mmm] text format thành list of segments.
    Normalize timestamps để khớp với audio duration thực tế.
    
    Returns:
        [{"start": 0.0, "text": "..."}, ...]
    """
    segments = []
    pattern = re.compile(r'\[(\d{1,2}):(\d{2})\.(\d{3})\]\s*(.*)')

    for line in raw.split("\n"):
        line = line.strip()
        if not line:
            continue
        m = pattern.match(line)
        if m:
            minutes = int(m.group(1))
            seconds = int(m.group(2))
            millis = int(m.group(3))
            start_time = minutes * 60 + seconds + millis / 1000.0
            text = m.group(4).strip()
            if text:
                segments.append({"start": start_time, "text": text})

    if not segments:
        return segments

    # ── Normalize timestamps ──
    # Gemini sometimes overshoots or compresses. Scale linearly to fit [0, duration].
    max_raw = segments[-1]["start"]
    if max_raw > 0 and (max_raw > duration * 1.05 or max_raw < duration * 0.85):
        scale = duration / max_raw
        print(f"      ⚠️  Normalizing timestamps: max_raw={max_raw:.1f}s → {duration:.1f}s (scale={scale:.3f})")
        for seg in segments:
            seg["start"] *= scale
    
    # Ensure last segment has room: if it ends at/after duration, pull it back
    if len(segments) >= 2:
        second_last = segments[-2]["start"]
        if segments[-1]["start"] >= duration - 0.5:
            segments[-1]["start"] = max(second_last + 1.5, duration - 3.0)
    elif segments and segments[-1]["start"] >= duration - 0.5:
        segments[-1]["start"] = max(0, duration - 3.0)

    return segments


def segments_to_txt(segments: list[dict]) -> str:
    """Convert segments to clean TXT (merged by pauses)."""
    if not segments:
        return ""

    lines = []
    current_paragraph = []
    last_end = segments[0]["start"]

    for i, seg in enumerate(segments):
        gap = seg["start"] - last_end if i > 0 else 0
        current_paragraph.append(seg["text"])

        # Merge into paragraphs: break on gaps > 1.5s
        next_gap = (
            segments[i + 1]["start"] - seg["start"]
            if i + 1 < len(segments) else 999
        )

        if gap > 1.5 and i > 0 and current_paragraph:
            # Flush previous paragraph
            lines.append(" ".join(current_paragraph[:-1]))
            current_paragraph = [seg["text"]]
        elif next_gap > 1.5:
            lines.append(" ".join(current_paragraph))
            current_paragraph = []
            lines.append("")  # paragraph break

        last_end = seg["start"]

    if current_paragraph:
        lines.append(" ".join(current_paragraph))

    return "\n".join(lines).strip()


def segments_to_srt(segments: list[dict], total_duration: float) -> str:
    """Convert segments to SRT format."""
    if not segments:
        return ""

    def fmt_time(seconds: float) -> str:
        h = int(seconds // 3600)
        m = int((seconds % 3600) // 60)
        s = int(seconds % 60)
        ms = int((seconds % 1) * 1000)
        return f"{h:02d}:{m:02d}:{s:02d},{ms:03d}"

    blocks = []
    for i, seg in enumerate(segments):
        start = seg["start"]
        # End = start of next segment, or estimated from text length
        if i + 1 < len(segments):
            end = segments[i + 1]["start"]
        else:
            # Last segment: estimate duration from text length (~10 chars/s for Burmese speech)
            end = start + max(2.0, len(seg["text"]) / 10.0)

        # Clamp to audio duration, but ensure minimum 1.5s
        end = min(end, total_duration)
        if end <= start:
            end = min(start + 2.0, total_duration)
            if end <= start:
                start = max(0, total_duration - 2.0)
                end = total_duration

        blocks.append(
            f"{i + 1}\n"
            f"{fmt_time(start)} --> {fmt_time(end)}\n"
            f"{seg['text']}"
        )

    return "\n\n".join(blocks) + "\n"


def process_file(
    api_key: str,
    audio_path: str,
    output_dir: str,
    model: str = DEFAULT_MODEL,
    max_retries: int = 3
) -> bool:
    """
    Process 1 audio file → .txt + .srt
    """
    name = Path(audio_path).stem
    txt_path = os.path.join(output_dir, f"{name}_gemini.txt")
    srt_path = os.path.join(output_dir, f"{name}_gemini.srt")
    raw_path = os.path.join(output_dir, f"{name}_gemini_raw.txt")

    # Skip if already done
    if os.path.exists(txt_path) and os.path.exists(srt_path):
        print(f"   ⏭️  {name} (already done)")
        return True

    print(f"   🎤 {name}")

    # Duration
    with open(audio_path, "rb") as f:
        audio_bytes = f.read()
    duration = len(audio_bytes) / 16000 / 2
    print(f"      {duration:.1f}s, {len(audio_bytes)/1024/1024:.1f}MB")

    raw, usage = transcribe_audio(audio_path, api_key, model, max_retries)
    if not raw:
        return False

    # Save raw output for debugging
    os.makedirs(output_dir, exist_ok=True)
    with open(raw_path, "w", encoding="utf-8") as f:
        f.write(raw)

    # Save usage
    if usage:
        import json
        usage_path = os.path.join(output_dir, f"{name}_gemini_usage.json")
        usage["audio_duration_s"] = duration
        usage["model"] = model
        with open(usage_path, "w") as f:
            json.dump(usage, f)
        print(f"      📊 {usage['prompt_tokens']} in + {usage['completion_tokens']} out = {usage['total_tokens']} tok")

    segments = parse_timestamped_transcript(raw, duration)
    if not segments:
        print(f"   ⚠️  No timestamped segments found in output")
        # Fallback: save raw as TXT
        with open(txt_path, "w", encoding="utf-8") as f:
            f.write(raw)
        return True

    print(f"      {len(segments)} segments extracted")

    # Write TXT
    txt = segments_to_txt(segments)
    with open(txt_path, "w", encoding="utf-8") as f:
        f.write(txt)

    # Write SRT
    srt = segments_to_srt(segments, duration)
    with open(srt_path, "w", encoding="utf-8") as f:
        f.write(srt)

    print(f"      → {txt_path} ({len(txt)} chars)")
    print(f"      → {srt_path} ({segments[-1]['start']:.1f}s end)")

    return True


def process_directory(
    api_key: str,
    input_dir: str,
    output_dir: str,
    model: str = DEFAULT_MODEL,
    max_retries: int = 3
):
    """Process all .wav files in directory."""
    wav_files = sorted(Path(input_dir).glob("*.wav"))
    if not wav_files:
        print(f"❌ No .wav files in {input_dir}")
        return

    os.makedirs(output_dir, exist_ok=True)
    print(f"\n🎯 Gemini Transcribe: {len(wav_files)} file(s)")
    print(f"   Model: {model} | Provider: OpenRouter")
    print(f"   Output: TXT + SRT (timestamped)")
    print()

    success = 0
    failed = []

    for wav_path in wav_files:
        ok = process_file(api_key, str(wav_path), output_dir, model, max_retries)
        if ok:
            success += 1
        else:
            failed.append(wav_path.name)

    print(f"\n{'=' * 50}")
    print(f"📊 {success}/{len(wav_files)} succeeded")
    if failed:
        print(f"❌ Failed: {', '.join(failed)}")
    return failed


def main():
    parser = argparse.ArgumentParser(
        description="G3: Gemini Direct Transcribe — Audio → TXT + SRT"
    )
    parser.add_argument("input", help=".wav file or directory of .wav files")
    parser.add_argument("-o", "--output-dir", default="transcripts/",
                        help="Output directory")
    parser.add_argument("-k", "--api-key", default=None,
                        help="OpenRouter API key (default: OPENROUTER_API_KEY env)")
    parser.add_argument("-m", "--model", default=DEFAULT_MODEL,
                        help=f"Model (default: {DEFAULT_MODEL})")
    parser.add_argument("--max-retries", type=int, default=3,
                        help="Max retries per chunk")
    args = parser.parse_args()

    api_key = args.api_key or os.environ.get("OPENROUTER_API_KEY")
    if not api_key:
        print("❌ OPENROUTER_API_KEY required")
        sys.exit(1)

    if os.path.isdir(args.input):
        process_directory(api_key, args.input, args.output_dir, args.model, args.max_retries)
    else:
        os.makedirs(args.output_dir, exist_ok=True)
        process_file(api_key, args.input, args.output_dir, args.model, args.max_retries)


if __name__ == "__main__":
    main()
