#!/usr/bin/env python3
"""Realtime and file-based ASR processing pipeline using Silero VAD segmentation
and local Orukeet background workers.
"""

import argparse
import copy
import json
import os
import queue
import re
import subprocess
import sys
import tempfile
import threading
import wave
from pathlib import Path

import numpy as np
import torch
from orukeet import Orukeet

# Thread lock for synchronized access to the shared Orukeet engine instance
asr_lock = threading.Lock()


def float32_to_wav_file(float32_array: np.ndarray, sample_rate: int = 16000) -> str:
    """Converts a float32 PCM numpy array into a temporary 16-bit PCM WAV file."""
    int16_pcm = (np.clip(float32_array, -1.0, 1.0) * 32767).astype(np.int16).tobytes()

    with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp_file:
        tmp_filename = tmp_file.name

    with wave.open(tmp_filename, "wb") as wav_out:
        wav_out.setnchannels(1)
        wav_out.setsampwidth(2)
        wav_out.setframerate(sample_rate)
        wav_out.writeframes(int16_pcm)

    return tmp_filename


class SpacyTagger:
    """Optional spaCy integration for NLP tokenization and linguistic tagging."""
    def __init__(self, model_name: str = "fr_core_news_sm"):
        import spacy
        try:
            self.nlp = spacy.load(model_name)
        except Exception as e:
            sys.stderr.write(f"Failed to load spaCy model '{model_name}' ({e}). Attempting fallback download...\n")
            try:
                subprocess.run([sys.executable, "-m", "spacy", "download", model_name], check=True)
                self.nlp = spacy.load(model_name)
            except Exception as err:
                sys.stderr.write(f"Could not load spaCy model '{model_name}': {err}. Disabling spaCy tagging.\n")
                self.nlp = None

    def process_text(self, text: str) -> dict:
        if not self.nlp:
            return {}
        doc = self.nlp(text)
        tokens = [
            {
                "text": token.text,
                "lemma": token.lemma_,
                "pos": token.pos_,
                "tag": token.tag_
            }
            for token in doc
        ]
        return {"tokens": tokens}


class RealtimeVADandOrukeetStreamer:
    """Streams an audio track via ffmpeg/parec, detects boundaries with Silero VAD,
    and enqueues audio turns for Orukeet transcription.
    """
    def __init__(
        self,
        source: str,
        label: str,
        vad_model,
        orukeet_engine: Orukeet,
        min_silence_ms: int,
        track_id: int,
        out_file,
        is_file_input: bool = False,
        spacy_tagger: SpacyTagger = None
    ):
        self.source = source
        self.label = label
        self.vad_model = copy.deepcopy(vad_model)
        self.orukeet_engine = orukeet_engine
        self.out_file = out_file
        self.track_id = track_id
        self.is_file_input = is_file_input
        self.spacy_tagger = spacy_tagger

        self.sample_rate = 16000
        self.frame_samples = 512  # 32ms frames @ 16kHz
        self.frame_bytes = self.frame_samples * 2

        self.min_silence_samples = int(self.sample_rate * (min_silence_ms / 1000.0))
        self.is_speaking = False
        self.speech_start_sample = 0
        self.silence_counter_samples = 0

        self.audio_turn_buffer = []

        # Worker queue offloading Orukeet transcription
        self.asr_queue = queue.Queue()
        self.running = False
        self.read_thread = None
        self.worker_thread = None
        self.process = None

    def emit_event(
        self,
        event_type: str,
        timestamp: float,
        start: float = None,
        text: str = None,
        phrases: list = None,
        duration: float = None,
        extra: dict = None
    ):
        """Emits structured JSONL event payloads."""
        payload = {
            "label": self.label,
            "event": event_type,
            "timestamp": round(timestamp, 3)
        }
        if start is not None:
            payload["start"] = round(start, 3)
        if text is not None:
            payload["text"] = text
        if phrases:
            payload["phrases"] = phrases
        if duration is not None:
            payload["duration"] = round(duration, 3)
        if extra:
            payload.update(extra)

        self.out_file.write(json.dumps(payload, ensure_ascii=False) + "\n")
        self.out_file.flush()

    def _asr_worker_loop(self):
        """Worker thread processing completed audio turns through Orukeet."""
        while self.running or not self.asr_queue.empty():
            try:
                task = self.asr_queue.get(timeout=0.1)
            except queue.Empty:
                continue

            pcm_data, start_sec, end_sec, duration = task
            wav_path = float32_to_wav_file(pcm_data, self.sample_rate)

            try:
                with asr_lock:
                    result = self.orukeet_engine.transcribe(Path(wav_path))

                text = ""
                phrases_info = []

                if isinstance(result, dict):
                    text = result.get("text", "").strip()

                    raw_phrases = result.get("phrases") or result.get("segments") or []
                    for phrase in raw_phrases:
                        p_text = phrase.get("text", "").strip()
                        p_start = round(start_sec + phrase.get("start", 0.0), 3)
                        p_end = round(start_sec + phrase.get("end", 0.0), 3)

                        phrases_info.append({
                            "text": p_text,
                            "start": p_start,
                            "end": p_end
                        })

                elif isinstance(result, str):
                    text = result.strip()

                if text:
                    extra_data = {}
                    if self.spacy_tagger:
                        extra_data["spacy"] = self.spacy_tagger.process_text(text)

                    self.emit_event(
                        "final",
                        timestamp=end_sec,
                        start=start_sec,
                        text=text,
                        phrases=phrases_info if phrases_info else None,
                        duration=duration,
                        extra=extra_data
                    )

            except Exception as exc:
                sys.stderr.write(f"[{self.label}] Orukeet transcription error: {exc}\n")

            finally:
                if os.path.exists(wav_path):
                    os.remove(wav_path)
                self.asr_queue.task_done()

    def _get_subprocess_command(self):
        """Builds reading command depending on input type (file, PulseAudio device, standard ffmpeg)."""
        if self.source.startswith("parec:"):
            device = self.source.replace("parec:", "")
            return ["parec", "--device=" + device, "--format=s16le", "--channels=1", f"--rate={self.sample_rate}"]
        elif self.is_file_input:
            return [
                "ffmpeg",
                "-loglevel", "quiet",
                "-i", self.source,
                "-map", f"0:a:{self.track_id}",
                "-f", "s16le",
                "-ac", "1",
                "-ar", str(self.sample_rate),
                "pipe:1"
            ]
        else:
            return [
                "ffmpeg",
                "-loglevel", "quiet",
                "-i", self.source,
                "-f", "s16le",
                "-ac", "1",
                "-ar", str(self.sample_rate),
                "pipe:1"
            ]

    def _read_loop(self):
        """Streams raw audio samples, evaluates VAD activity, and pushes turns to worker."""
        cmd = self._get_subprocess_command()

        try:
            self.process = subprocess.Popen(cmd, stdout=subprocess.PIPE, stderr=subprocess.DEVNULL)
        except Exception as e:
            sys.stderr.write(f"[{self.label}] Failed to start subprocess for '{self.source}': {e}\n")
            return

        if hasattr(self.vad_model, "reset_states"):
            self.vad_model.reset_states()

        total_samples_processed = 0

        while self.running:
            raw_bytes = self.process.stdout.read(self.frame_bytes)
            if len(raw_bytes) < self.frame_bytes:
                break

            audio_int16 = np.frombuffer(raw_bytes, dtype=np.int16)
            float32_frame = audio_int16.astype(np.float32) / 32768.0
            tensor_frame = torch.from_numpy(float32_frame)

            with torch.no_grad():
                speech_prob = self.vad_model(tensor_frame, self.sample_rate).item()

            threshold = 0.5
            current_sec = total_samples_processed / self.sample_rate

            if speech_prob >= threshold:
                self.silence_counter_samples = 0
                if not self.is_speaking:
                    self.is_speaking = True
                    self.speech_start_sample = total_samples_processed
                    self.audio_turn_buffer = []

                    # Emit immediate VAD start event
                    self.emit_event("start", timestamp=current_sec)

                self.audio_turn_buffer.extend(float32_frame)

            else:
                if self.is_speaking:
                    self.silence_counter_samples += self.frame_samples
                    self.audio_turn_buffer.extend(float32_frame)

                    if self.silence_counter_samples >= self.min_silence_samples:
                        self.is_speaking = False

                        speech_end_sample = (total_samples_processed + self.frame_samples) - self.silence_counter_samples
                        speech_audio = np.array(self.audio_turn_buffer, dtype=np.float32)

                        start_sec = self.speech_start_sample / self.sample_rate
                        end_sec = max(speech_end_sample / self.sample_rate, start_sec)
                        duration = end_sec - start_sec

                        # Emit immediate VAD end event
                        self.emit_event("end", timestamp=end_sec, start=start_sec, duration=duration)

                        # Enqueue audio segment to worker thread
                        self.asr_queue.put((speech_audio, start_sec, end_sec, duration))

                        self.audio_turn_buffer = []
                        self.silence_counter_samples = 0

            total_samples_processed += self.frame_samples

    def start(self):
        self.running = True
        self.read_thread = threading.Thread(target=self._read_loop, daemon=True)
        self.worker_thread = threading.Thread(target=self._asr_worker_loop, daemon=True)
        self.read_thread.start()
        self.worker_thread.start()

    def stop(self):
        self.running = False
        if self.process:
            self.process.terminate()
            self.process.wait()

    def join(self):
        if self.read_thread:
            self.read_thread.join()
        self.asr_queue.join()


def parse_track_args(track_args):
    """Parses track mappings in format `-t 0=speaker1 1=speaker2` or `-t device_name=label` or `-t device_name`."""
    tracks = {}
    if not track_args:
        return {0: "track_0"}

    for item in track_args:
        if "=" in item:
            key, label = item.split("=", 1)
            try:
                tracks[int(key)] = label
            except ValueError:
                tracks[key] = label
        else:
            try:
                tracks[int(item)] = f"track_{item}"
            except ValueError:
                tracks[item] = item
    return tracks


def main():
    parser = argparse.ArgumentParser(description="Realtime VAD and Orukeet Speech Recognition Pipeline")
    parser.add_argument("-i", "--input", nargs="?", default=None, help="Audio/video file path or stream URL.")
    parser.add_argument("-t", "--track", action="append", nargs="+", help="Audio track or device mapping (e.g. -t dev=label)")
    parser.add_argument("-s", "--silence", "--min-silence-ms", dest="min_silence_ms", type=int, default=500, help="VAD silence threshold in milliseconds")
    parser.add_argument("--spacy-model", "--spacy_model", dest="spacy_model", nargs="?", const="fr_core_news_sm", default=None, help="Optional spaCy model (e.g., fr_core_news_sm)")
    parser.add_argument("--installation", type=Path, default=Path("installation.json"), help="Path to Orukeet installation receipt")
    parser.add_argument("-o", "--output", default=None, help="Output JSONL log path (defaults to stdout)")

    # Legacy CLI compatibility flags (ignored gracefully by Orukeet)
    parser.add_argument("--transcribe", action="store_true", help="Legacy flag (ignored)")
    parser.add_argument("--model", type=str, default=None, help="Legacy Whisper model size flag (ignored)")
    parser.add_argument("--language", type=str, default=None, help="Legacy language flag (ignored)")

    args = parser.parse_args()

    if hasattr(sys.stdout, "reconfigure"):
        sys.stdout.reconfigure(encoding="utf-8")

    # Load Orukeet configuration from installation receipt
    config_path = args.installation
    if not config_path.is_file():
        sys.stderr.write(f"Orukeet installation config not found: {config_path}\n")
        sys.exit(1)

    try:
        config = json.loads(config_path.read_text(encoding="utf-8-sig"))
    except Exception as e:
        sys.stderr.write(f"Failed to parse Orukeet installation JSON: {e}\n")
        sys.exit(1)

    # Flatten nested -t lists if user passes multiple -t options
    raw_tracks = []
    if args.track:
        for item in args.track:
            if isinstance(item, list):
                raw_tracks.extend(item)
            else:
                raw_tracks.append(item)

    track_mappings = parse_track_args(raw_tracks)
    is_file = os.path.isfile(args.input) if args.input else False

    sys.stderr.write("Loading Silero VAD model...\n")
    vad_model, _ = torch.hub.load(
        repo_or_dir="snakers4/silero-vad",
        model="silero_vad",
        force_reload=False,
        onnx=False
    )

    spacy_tagger = SpacyTagger(args.spacy_model) if args.spacy_model else None
    out_file = open(args.output, "w", encoding="utf-8") if args.output else sys.stdout

    try:
        sys.stderr.write("Initializing Orukeet engine...\n")
        with Orukeet(config["model"], config["runtime"], device=config["device"]) as asr:
            streamers = []

            for track_key, label in track_mappings.items():
                if args.input:
                    source = args.input
                    track_id = track_key if isinstance(track_key, int) else 0
                else:
                    pa_device = str(track_key) if isinstance(track_key, str) else label
                    source = f"parec:{pa_device}"
                    track_id = 0

                sys.stderr.write(f"[*] Monitoring source: '{source}' (label: '{label}')\n")

                streamer = RealtimeVADandOrukeetStreamer(
                    source=source,
                    label=label,
                    vad_model=vad_model,
                    orukeet_engine=asr,
                    min_silence_ms=args.min_silence_ms,
                    track_id=track_id,
                    out_file=out_file,
                    is_file_input=is_file,
                    spacy_tagger=spacy_tagger
                )
                streamers.append(streamer)

            for streamer in streamers:
                streamer.start()

            for streamer in streamers:
                streamer.join()

    except Exception as exc:
        sys.stderr.write(f"Execution error: {exc}\n")
        sys.exit(1)
    finally:
        if args.output and out_file != sys.stdout:
            out_file.close()


if __name__ == "__main__":
    main()
