import os
import glob
import subprocess
import gc
from pathlib import Path
from datetime import datetime

# Externe Bibliotheken (müssen via pip installiert sein)
from docx import Document
from dotenv import load_dotenv
import torch
import whisperx

# --- KONFIGURATION & SICHERHEIT ---
# Projektpfade relativ zu dieser Datei bestimmen, damit die App nach Clone/Move
# nicht versehentlich in ~/Projekte/Autoprotocol liest oder schreibt.
BASE_DIR = Path(__file__).resolve().parent
INPUT_DIR = BASE_DIR / "input"
OUTPUT_DIR = BASE_DIR / "output"

# Lädt die Umgebungsvariablen aus der .env Datei
load_dotenv() 
HF_TOKEN = os.getenv("HF_TOKEN")

# AutoProtocol nutzt bewusst eine eigene STT-Konfiguration, getrennt von Hermes
# Telegram-STT. Kurze Telegram-Sprachclips bleiben dadurch schnell (`medium`),
# während Sitzungsaufzeichnungen mit dem besten lokalen Modell auf der A5000
# transkribiert werden können. Defaults sind auf "parallel zu Chatterbox testen"
# ausgelegt: large-v3, CUDA, float16, kleiner Batch.
WHISPER_MODEL = os.getenv("AUTOPROTOCOL_WHISPER_MODEL", "large-v3")
WHISPER_DEVICE = os.getenv("AUTOPROTOCOL_WHISPER_DEVICE", "cuda")
WHISPER_COMPUTE_TYPE = os.getenv("AUTOPROTOCOL_WHISPER_COMPUTE_TYPE", "float16")
WHISPER_FALLBACK_COMPUTE_TYPE = os.getenv("AUTOPROTOCOL_WHISPER_FALLBACK_COMPUTE_TYPE", "int8_float16")
WHISPER_BATCH_SIZE = int(os.getenv("AUTOPROTOCOL_WHISPER_BATCH_SIZE", "2"))
WHISPER_LANGUAGE = os.getenv("AUTOPROTOCOL_WHISPER_LANGUAGE", "de")


def require_hf_token() -> str:
    """Return the HuggingFace token or raise a web-app friendly error."""
    token = os.getenv("HF_TOKEN") or HF_TOKEN
    if not token:
        raise RuntimeError(
            "HF_TOKEN fehlt. Bitte in der lokalen .env setzen; ohne Token kann WhisperX keine Sprecher-Diarisierung ausführen."
        )
    return token

def setup_directories():
    """Stellt sicher, dass die Input/Output-Ordner existieren."""
    os.makedirs(INPUT_DIR, exist_ok=True)
    os.makedirs(OUTPUT_DIR, exist_ok=True)


def log_cuda_memory(label: str) -> None:
    """Print a compact VRAM snapshot without failing on CPU-only systems."""
    if not torch.cuda.is_available():
        print(f"[*] VRAM {label}: CUDA nicht verfügbar.")
        return
    free, total = torch.cuda.mem_get_info()
    used = total - free
    mib = 1024 * 1024
    print(
        f"[*] VRAM {label}: used={used / mib:.0f} MiB, "
        f"free={free / mib:.0f} MiB, total={total / mib:.0f} MiB"
    )


def load_whisperx_model_with_fallback(device: str):
    """Load the configured WhisperX model; on CUDA OOM retry with lower VRAM compute type."""
    print(
        f"[*] Lade WhisperX-Modell '{WHISPER_MODEL}' "
        f"auf {device} mit compute_type={WHISPER_COMPUTE_TYPE}..."
    )
    try:
        return whisperx.load_model(WHISPER_MODEL, device, compute_type=WHISPER_COMPUTE_TYPE), WHISPER_COMPUTE_TYPE
    except torch.cuda.OutOfMemoryError:
        if device != "cuda" or WHISPER_FALLBACK_COMPUTE_TYPE == WHISPER_COMPUTE_TYPE:
            raise
        print(
            "[!] CUDA VRAM war für den primären Lauf zu knapp. "
            f"Versuche erneut mit compute_type={WHISPER_FALLBACK_COMPUTE_TYPE}, "
            "ohne Chatterbox zu stoppen."
        )
        gc.collect()
        torch.cuda.empty_cache()
        if torch.cuda.is_available():
            torch.cuda.ipc_collect()
        log_cuda_memory("nach OOM cleanup")
        return whisperx.load_model(WHISPER_MODEL, device, compute_type=WHISPER_FALLBACK_COMPUTE_TYPE), WHISPER_FALLBACK_COMPUTE_TYPE

def list_and_select_files():
    """Listet alle Medien-Dateien im Input-Ordner auf und lässt den Nutzer wählen."""
    files = glob.glob(os.path.join(INPUT_DIR, "*.*"))
    media_files = [
        f for f in files
        if f.lower().endswith(('.mp4', '.mov', '.m4v', '.mkv', '.avi', '.webm', '.mp3', '.m4a', '.aac', '.ogg', '.oga', '.flac', '.wav'))
    ]
    
    if not media_files:
        print("[!] Keine Audio- oder Videodateien im 'input'-Ordner gefunden.")
        return None
        
    print("\n--- Verfügbare Dateien im Input-Ordner ---")
    for i, file_path in enumerate(media_files):
        filename = os.path.basename(file_path)
        print(f"[{i + 1}] {filename}")
        
    try:
        selection = int(input("\nBitte wähle die Nummer der Datei für die Transkription: ")) - 1
        if 0 <= selection < len(media_files):
            return media_files[selection]
        else:
            print("[!] Ungültige Auswahl.")
            return None
    except ValueError:
        print("[!] Bitte eine Zahl eingeben.")
        return None

def extract_audio(media_path):
    """Extrahiert die Audiospur aus einem Video via ffmpeg für schnellere Verarbeitung."""
    filename = os.path.basename(media_path)
    name, ext = os.path.splitext(filename)
    
    # Wenn es bereits Audio ist, überspringen wir die Extraktion
    if ext.lower() in ['.wav', '.mp3', '.m4a', '.aac', '.ogg', '.oga', '.flac']:
        return media_path
        
    audio_path = os.path.join(INPUT_DIR, f"{name}_extracted.wav")
    print(f"\n[*] Extrahiere Audio aus {filename}...")
    
    # ffmpeg Kommando: 16kHz, mono, wav (Ideal für Whisper)
    command = ["ffmpeg", "-y", "-i", media_path, "-vn", "-ar", "16000", "-ac", "1", audio_path]
    subprocess.run(command, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
    
    return audio_path

def run_whisperx(audio_path):
    """
    Führt WhisperX lokal auf der GPU aus:
    1. Transkription (Schweizerdeutsch -> Hochdeutsch)
    2. Alignment (Zeitsynchronisation)
    3. Diarisierung (Sprechererkennung)
    """
    # Direkter Import der Diarisierungs-Funktionen, um den AttributeError zu vermeiden
    from whisperx.diarize import DiarizationPipeline, assign_word_speakers
    
    print(f"\n[*] Starte KI-Transkription (WhisperX) für: {os.path.basename(audio_path)}")
    print(
        "[*] AutoProtocol-STT läuft getrennt von Hermes Telegram-STT: "
        f"model={WHISPER_MODEL}, language={WHISPER_LANGUAGE}, "
        f"batch_size={WHISPER_BATCH_SIZE}."
    )

    device = WHISPER_DEVICE
    if device == "cuda" and not torch.cuda.is_available():
        print("[!] CUDA ist nicht verfügbar; falle auf CPU zurück.")
        device = "cpu"

    log_cuda_memory("vor WhisperX load")

    # 1. Transkription
    model, active_compute_type = load_whisperx_model_with_fallback(device)
    print(f"[*] WhisperX aktiv mit compute_type={active_compute_type}.")
    audio = whisperx.load_audio(audio_path)

    # Wir zwingen das Modell auf Deutsch ("de"), das deckt Schweizer Hochdeutsch
    # und viele Dialekt-Passagen besser ab als Autodetection.
    result = model.transcribe(audio, batch_size=WHISPER_BATCH_SIZE, language=WHISPER_LANGUAGE)
    print(f"[*] Transkription abgeschlossen. Festgesetzte Sprache: {result['language']}")

    # VRAM vor Alignment freigeben; sonst bleibt large-v3 im Speicher und konkurriert
    # unnötig mit Alignment/Diarization/Chatterbox.
    del model
    gc.collect()
    if torch.cuda.is_available():
        torch.cuda.empty_cache()
        torch.cuda.ipc_collect()
    log_cuda_memory("nach Transkriptionsmodell unload")

    # 2. Alignment
    print("[*] Führe Alignment durch (Text-Audio-Synchronisation)...")
    model_a, metadata = whisperx.load_align_model(language_code=result["language"], device=device)
    result = whisperx.align(result["segments"], model_a, metadata, audio, device, return_char_alignments=False)

    # VRAM aufräumen
    del model_a
    gc.collect()
    if torch.cuda.is_available():
        torch.cuda.empty_cache()
        torch.cuda.ipc_collect()
    log_cuda_memory("nach Alignment unload")

    # 3. Diarisierung (Wer hat wann gesprochen?)
    print("[*] Starte Sprechererkennung (Diarisierung)...")
    # NEU: Direkter Aufruf der Pipeline über den neuen Import
    token = require_hf_token()
    diarize_model = DiarizationPipeline(token=token, device=device)
    diarize_segments = diarize_model(audio_path)
    del diarize_model
    gc.collect()
    if torch.cuda.is_available():
        torch.cuda.empty_cache()
        torch.cuda.ipc_collect()
    log_cuda_memory("nach Diarization unload")

    # 4. Zusammenführen von Text und Sprechern
    print("[*] Führe Text und Sprecher zusammen...")
    # NEU: Direkter Aufruf der Zuweisungs-Funktion
    result = assign_word_speakers(diarize_segments, result)
    
    # 5. Formatieren für unseren Word-Export
    formatted_result = []
    for segment in result["segments"]:
        speaker = segment.get("speaker", "UNBEKANNT")
        text = segment.get("text", "").strip()
        formatted_result.append({"speaker": speaker, "text": text})
        
    return formatted_result

def export_to_word(transcription_data, original_filename):
    """Speichert das fertige Transkript als formatiertes Word-Dokument."""
    print("\n[*] Erstelle Word-Dokument...")
    doc = Document()
    doc.add_heading('Sitzungstranskript', 0)
    
    now = datetime.now()
    date_str = now.strftime("%d.%m.%Y %H:%M")
    file_date_str = now.strftime("%Y%m%d_%H%M")
    
    doc.add_paragraph(f"Datum der Erstellung: {date_str}")
    doc.add_paragraph(f"Quelle: {os.path.basename(original_filename)}")
    doc.add_paragraph("-" * 50)
    
    for segment in transcription_data:
        p = doc.add_paragraph()
        runner = p.add_run(f"{segment['speaker']}: ")
        runner.bold = True
        p.add_run(segment['text'])
        
    output_filename = f"Transkript_{file_date_str}.docx"
    output_path = os.path.join(OUTPUT_DIR, output_filename)
    
    doc.save(output_path)
    print(f"[+] ERFOLG! Transkript gespeichert unter: {output_path}\n")
    return output_path

def process_direct(file_path):
    """Diese Funktion wird unsichtbar vom Streamlit-UI aufgerufen"""
    setup_directories()
    # 1. Audio extrahieren
    audio_file = extract_audio(file_path)
    
    # 2. WhisperX starten
    transcription_data = run_whisperx(audio_file)
    
    # --- NEU: DER HOLZHAMMER FÜR DEN VRAM ---
    print("[*] Zwinge PyTorch, den VRAM komplett freizugeben...")
    import gc
    import torch
    gc.collect()
    torch.cuda.empty_cache()
    # Diese Umgebungsvariable zwingt PyTorch, Speicherblöcke zu defragmentieren
    os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True" 
    
    # Optional, aber oft der Retter bei hartnäckigem PyTorch-Caching:
    if torch.cuda.is_available():
        torch.cuda.ipc_collect()
    # ----------------------------------------
    
    return transcription_data

def main():
    setup_directories()
    
    selected_file = list_and_select_files()
    if not selected_file:
        return
        
    # 1. Audio extrahieren (falls Video)
    audio_file = extract_audio(selected_file)
    
    # 2. Transkribieren (mit lokaler KI)
    transcription_data = run_whisperx(audio_file)
    
    # 3. Als Word exportieren
    export_to_word(transcription_data, selected_file)

if __name__ == "__main__":
    main()