#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""Standalone LIVE transcript-correction test (separate from the v2 program).

Proves/disproves: can we get CORRECT caller chat messages DURING a call,
instead of only after it ends?

Pipeline:
  1. Streams your WAV to DashScope Omni Realtime exactly like the AI server
     does for a live call (16 kHz PCM16, 100 ms chunks, same VAD defaults,
     same independent input-audio transcription pass).
  2. RAW line arrives per sentence (what the v2 dashboard shows today).
  3. Each RAW line is immediately corrected by an LLM using the lines so far
     as context (the "live incremental" version of the server's
     _correct_transcript) — timestamped so you see the real added delay.
  4. At the end, the server's original POST-CALL whole-transcript correction
     also runs, as the quality reference.

Output: console + <audio>.result.txt with, per sentence:
     RAW / LIVE-FIX (+delay) / POST-CALL-FIX side by side.
If LIVE-FIX ~= POST-CALL-FIX and both read correctly, the feature works.

Usage:
  python live_correction_test.py call.wav [--speed 2] [--llm-model qwen3-max]
Requires: DASHSCOPE_API_KEY env var.
"""

import argparse
import base64
import os
import queue
import sys
import threading
import time
import wave
from pathlib import Path

try:
    import numpy as np
    from openai import OpenAI
    import dashscope
    from dashscope.audio.qwen_omni import (
        MultiModality, OmniRealtimeConversation, OmniRealtimeCallback
    )
except ImportError as e:
    print(f"Missing dependency: {e}"); sys.exit(1)

LLM_BASE = 'https://dashscope-intl.aliyuncs.com/compatible-mode/v1'
OMNI_MODEL = 'qwen3.5-omni-plus-realtime'
OMNI_BASE = 'wss://dashscope-intl.aliyuncs.com/api-ws/v1/realtime'

API_KEY = os.getenv('DASHSCOPE_API_KEY', '')
if not API_KEY:
    print("ERROR: DASHSCOPE_API_KEY not set."); sys.exit(1)
dashscope.api_key = API_KEY

SAMPLE_RATE = 16000
CHUNK_MS = 100
CHUNK_BYTES = SAMPLE_RATE * 2 * CHUNK_MS // 1000


def load_wav_16k_mono(path):
    with wave.open(path, 'rb') as wf:
        n_ch, width, rate = wf.getnchannels(), wf.getsampwidth(), wf.getframerate()
        raw = wf.readframes(wf.getnframes())
    if width != 2:
        print(f"ERROR: only 16-bit PCM WAV supported (got {width * 8}-bit)."); sys.exit(1)
    samples = np.frombuffer(raw, dtype=np.int16)
    if n_ch > 1:
        samples = samples.reshape(-1, n_ch).mean(axis=1).astype(np.int16)
    if rate != SAMPLE_RATE:
        n_out = int(len(samples) * SAMPLE_RATE / rate)
        x_out = np.linspace(0, len(samples) - 1, n_out)
        samples = np.interp(x_out, np.arange(len(samples)),
                            samples.astype(np.float64)).astype(np.int16)
        print(f"[Audio] Resampled {rate} Hz -> {SAMPLE_RATE} Hz")
    dur = len(samples) / SAMPLE_RATE
    print(f"[Audio] {path}: {dur:.1f}s, {n_ch} ch -> mono 16 kHz")
    return samples.tobytes(), dur


class Seg:
    __slots__ = ('idx', 'raw', 't_raw', 'live', 't_live', 'post')
    def __init__(self, idx, raw, t_raw):
        self.idx, self.raw, self.t_raw = idx, raw, t_raw
        self.live = None; self.t_live = None; self.post = None


class TranscribeCallback(OmniRealtimeCallback):
    def __init__(self, on_raw):
        self.closed = False
        self._on_raw = on_raw
        self._last_event_t = time.time()
        self.count = 0

    def on_open(self): print('[Omni] WebSocket connected')
    def on_close(self, code=None, msg=None):
        print(f'[Omni] WebSocket closed: {code} {msg}'); self.closed = True

    def on_event(self, event):
        self._last_event_t = time.time()
        if event.get('type', '') == 'conversation.item.input_audio_transcription.completed':
            t = event.get('transcript', '').strip()
            if t:
                self.count += 1
                self._on_raw(t)


# ── LIVE incremental correction (the proposed feature) ─────────────
LIVE_PROMPT = """You are a transcript correction assistant for a live phone call. Below are the transcript lines so far, in order. They were produced by an automatic speech recognition system that sometimes makes errors, especially with:
- Proper nouns (product names, company names, people's names)
- Code-switching (mixing languages mid-sentence)
- Phonetically similar words

Your task: return the CORRECTED version of ONLY the LAST line. Use the earlier lines as context clues. If the last line is already correct, return it unchanged. Do NOT add explanations, quotes or prefixes — return ONLY the corrected text of the last line.

Transcript so far:
{transcript}"""


def live_worker(q, segs, llm, model, t0, lock):
    while True:
        item = q.get()
        if item is None:
            return
        idx = item
        with lock:
            lines = []
            for s in segs[:idx + 1]:
                lines.append(s.live if (s.live and s.idx < idx) else s.raw)
        prompt = LIVE_PROMPT.format(transcript="\n".join(lines))
        try:
            r = llm.chat.completions.create(
                model=model, temperature=0.1,
                messages=[{"role": "user", "content": prompt}])
            fixed = r.choices[0].message.content.strip().strip('"')
        except Exception as e:
            print(f'  [live-fix] LLM error on #{idx + 1}: {e}')
            fixed = None
        with lock:
            seg = segs[idx]
            seg.live = fixed if fixed else seg.raw
            seg.t_live = time.time() - t0
            delay = seg.t_live - seg.t_raw
            mark = 'CHANGED' if seg.live != seg.raw else 'same'
            print(f'  [live-fix +{delay:4.1f}s] #{idx + 1} ({mark}): {seg.live}')


# ── POST-CALL correction, verbatim server method (reference) ───────
POST_PROMPT = """You are a transcript correction assistant. Below is the transcript of a phone call. The speech was transcribed by an automatic speech recognition system that sometimes makes errors, especially with:
- Proper nouns (product names, company names, people's names)
- Code-switching (mixing languages mid-sentence)
- Phonetically similar words

Your task: Return the CORRECTED transcript. Fix ONLY lines where transcription errors are obvious from context. Keep the exact same format: one line per message, same number of lines, same order.

Return ONLY the corrected transcript lines, nothing else.

Transcript:
{transcript}"""


def post_correct(segs, llm, model):
    raw = "\n".join(s.raw for s in segs)
    try:
        r = llm.chat.completions.create(
            model=model, temperature=0.1,
            messages=[{"role": "user", "content": POST_PROMPT.format(transcript=raw)}])
        lines = [l.strip() for l in r.choices[0].message.content.strip().split('\n') if l.strip()]
        for i, s in enumerate(segs):
            s.post = lines[i] if i < len(lines) else s.raw
    except Exception as e:
        print(f'[post-fix] LLM error: {e}')
        for s in segs:
            s.post = s.raw


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument('audio')
    ap.add_argument('--model', default=OMNI_MODEL)
    ap.add_argument('--llm-model', default='qwen3-max')
    ap.add_argument('--vad-threshold', type=float, default=0.2)
    ap.add_argument('--silence-ms', type=int, default=800)
    ap.add_argument('--speed', type=float, default=2.0)
    args = ap.parse_args()

    pcm, duration = load_wav_16k_mono(args.audio)
    llm = OpenAI(api_key=API_KEY, base_url=LLM_BASE)

    segs = []
    lock = threading.Lock()
    fixq = queue.Queue()
    t0 = time.time()

    def on_raw(text):
        with lock:
            seg = Seg(len(segs), text, time.time() - t0)
            segs.append(seg)
            print(f'[raw     +{seg.t_raw:4.1f}s] #{seg.idx + 1}: {text}')
            fixq.put(seg.idx)

    worker = threading.Thread(target=live_worker,
                              args=(fixq, segs, llm, args.llm_model, t0, lock), daemon=True)
    worker.start()

    cb = TranscribeCallback(on_raw)
    conv = OmniRealtimeConversation(model=args.model, callback=cb, url=OMNI_BASE)
    conv.connect()
    conv.update_session(
        output_modalities=[MultiModality.AUDIO, MultiModality.TEXT],
        voice='Tina',
        turn_detection_threshold=args.vad_threshold,
        turn_detection_silence_duration_ms=args.silence_ms,
        input_audio_transcription_model=args.model,
        instructions='You are a silent transcription session. Reply with a single word.',
        temperature=0.8,
        max_response_output_tokens=1,
    )

    print(f'[Stream] Sending {duration:.1f}s of audio at {args.speed:g}x realtime...')
    delay = (CHUNK_MS / 1000.0) / max(args.speed, 0.1)
    for i in range(0, len(pcm), CHUNK_BYTES):
        conv.append_audio(base64.b64encode(pcm[i:i + CHUNK_BYTES]).decode())
        time.sleep(delay)
    silence = b'\x00' * CHUNK_BYTES
    for _ in range(int((args.silence_ms + 700) / CHUNK_MS)):
        conv.append_audio(base64.b64encode(silence).decode())
        time.sleep(delay)

    print('[Stream] Done, waiting for final transcripts...')
    deadline = time.time() + 15
    while time.time() < deadline and not cb.closed:
        if time.time() - cb._last_event_t > 4 and segs:
            break
        time.sleep(0.25)
    try:
        conv.close()
    except Exception:
        pass

    # Wait for pending live fixes, then stop the worker
    fixq.put(None)
    worker.join(timeout=120)

    if not segs:
        print('\nNo speech transcribed (no VAD-committed turns).')
        return

    print('\n[post-fix] Running the server\'s original whole-call correction (reference)...')
    post_correct(segs, llm, args.llm_model)

    # ── Verdict view ──
    out = []
    changed_live = sum(1 for s in segs if s.live != s.raw)
    match_post = sum(1 for s in segs if (s.live or s.raw) == s.post)
    out.append('=' * 70)
    out.append(f'RESULT — {len(segs)} sentences | live-fix changed {changed_live} | '
               f'live matches post-call on {match_post}/{len(segs)}')
    out.append('=' * 70)
    for s in segs:
        d = (s.t_live - s.t_raw) if s.t_live else 0
        out.append(f'\n#{s.idx + 1}  [heard at {s.t_raw:.1f}s]')
        out.append(f'  RAW (dashboard today) : {s.raw}')
        out.append(f'  LIVE-FIX (+{d:.1f}s)     : {s.live}')
        out.append(f'  POST-CALL (reference) : {s.post}')
    report = "\n".join(out)
    print('\n' + report)
    out_path = str(Path(args.audio).with_suffix('.result.txt'))
    with open(out_path, 'w', encoding='utf-8') as f:
        f.write(report + '\n')
    print(f'\nSaved: {out_path}')


if __name__ == '__main__':
    main()
