from transformers import WhisperProcessor 
from transformers import WhisperForConditionalGeneration
import string
import torch
import numpy as np
import traceback
import logging
import torch
import torchaudio
import cgi
import tempfile
import json
import os
import torchaudio.transforms as T
from contextlib import contextmanager
import time

# Configuration du logging
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logger = logging.getLogger(__name__)

# Variables globales pour le modèle
model_id = "openai/whisper-small"
processor = None
model = None
model_loaded = False
initialization_error = None
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# Configuration
MAX_FILE_SIZE = 50 * 1024 * 1024  # 50MB
SUPPORTED_FORMATS = ['.wav', '.mp3', '.m4a', '.flac', '.ogg']
TARGET_SAMPLE_RATE = 16000

def initialize_model():
    """Initialise le modèle Whisper une seule fois"""
    global processor, model, model_loaded, initialization_error

    if not model_loaded and initialization_error is None:
        logger.info("Chargement du modèle Whisper...")
        try:
            processor = WhisperProcessor.from_pretrained(model_id)
            model = WhisperForConditionalGeneration.from_pretrained(model_id).to(device)
            model.config.forced_decoder_ids = None
            model.config.suppress_tokens = []
            model.config.use_cache = False
            model_loaded = True
            logger.info(f"Modèle chargé avec succès sur {device}.")
        except Exception as e:
            initialization_error = str(e)
            logger.error("Échec de l'initialisation du modèle: %s", e)
            traceback.print_exc()

@contextmanager
def temporary_file(suffix=".wav"):
    """Context manager pour la gestion des fichiers temporaires"""
    temp_file = None
    try:
        temp_file = tempfile.NamedTemporaryFile(delete=False, suffix=suffix)
        yield temp_file
    finally:
        if temp_file:
            try:
                temp_file.close()
                if os.path.exists(temp_file.name):
                    os.unlink(temp_file.name)
            except Exception as e:
                logger.warning(f"Erreur lors de la suppression du fichier temporaire: {e}")

def validate_audio_file(file_path, max_size=MAX_FILE_SIZE):
    """Valide un fichier audio"""
    if not os.path.exists(file_path):
        raise ValueError("Le fichier audio n'existe pas")
    
    file_size = os.path.getsize(file_path)
    if file_size == 0:
        raise ValueError("Le fichier audio est vide")
    
    if file_size > max_size:
        raise ValueError(f"Le fichier est trop volumineux ({file_size} bytes). Limite: {max_size} bytes")
    
    # Vérifier l'extension (optionnel)
    file_ext = os.path.splitext(file_path)[1].lower()
    if file_ext not in SUPPORTED_FORMATS:
        logger.warning(f"Format de fichier non testé: {file_ext}")
    
    return file_size

def analyze_whisper_output(processor, model, audio_tensor, sampling_rate):
    """Analyse la sortie du modèle Whisper avec gestion d'erreurs améliorée"""
    try:
        forced_decoder_ids = processor.get_decoder_prompt_ids(language="french", task="transcribe")
        processed_input = processor(audio_tensor.squeeze().numpy(), sampling_rate=sampling_rate, return_tensors="pt")
        input_features = processed_input.input_features.to(device)

        with torch.no_grad():  # Économiser la mémoire
            output = model.generate(
                input_features=input_features,
                output_scores=True,
                return_dict_in_generate=True,
                forced_decoder_ids=forced_decoder_ids,
                max_length=448,  # Limite raisonnable
                num_beams=1,     # Greedy decoding pour plus de rapidité
            )

        sequence = output.sequences[0]
        scores = output.scores

        # Vérifier que nous avons des scores
        if not scores:
            logger.warning("Aucun score généré par le modèle")
            transcription = processor.batch_decode(output.sequences, skip_special_tokens=True)[0]
            return transcription, 0.0, []

        tokens = sequence[len(forced_decoder_ids) + 1:-1]
        scores = scores[:-1] if len(scores) > len(tokens) else scores

        # Assurer la cohérence entre tokens et scores
        min_length = min(len(tokens), len(scores))
        tokens = tokens[:min_length]
        scores = scores[:min_length]

        subtoken_ids = tokens.tolist()
        subtoken_strings = [processor.tokenizer.decode([tid], skip_special_tokens=True) for tid in subtoken_ids]

        token_probs = [
            torch.nn.functional.softmax(score, dim=-1)[0, tid].item()
            for score, tid in zip(scores, subtoken_ids)
        ]

        words, current_word, current_tokens, current_probs = [], '', [], []
        filtered_token_probs = []

        for subtoken, prob in zip(subtoken_strings, token_probs):
            if subtoken.strip() in string.punctuation:
                continue

            if subtoken.startswith(" "):
                if current_word:
                    words.append({
                        "word": current_word,
                        "subtokens": current_tokens,
                        "subtoken_probs": current_probs,
                        "avg_prob": float(np.mean(current_probs))
                    })
                    filtered_token_probs.extend(current_probs)
                current_word, current_tokens, current_probs = subtoken.strip(), [subtoken], [prob]
            else:
                current_word += subtoken
                current_tokens.append(subtoken)
                current_probs.append(prob)

        if current_word:
            words.append({
                "word": current_word,
                "subtokens": current_tokens,
                "subtoken_probs": current_probs,
                "avg_prob": float(np.mean(current_probs))
            })
            filtered_token_probs.extend(current_probs)

        transcription = processor.batch_decode(output.sequences, skip_special_tokens=True)[0]
        avg_transcription_prob = float(np.mean(filtered_token_probs)) if filtered_token_probs else 0.0

        return transcription, avg_transcription_prob, words

    except Exception as e:
        logger.error("Erreur dans analyze_whisper_output: %s", e)
        raise

def load_audio(file_path):
    """Charge un fichier audio, convertit en mono et resample en 16 kHz"""
    try:
        # Valider le fichier d'abord
        validate_audio_file(file_path)
        
        waveform, sample_rate = torchaudio.load(file_path)
        
        # Conversion en mono si nécessaire
        if waveform.shape[0] > 1:
            waveform = torch.mean(waveform, dim=0, keepdim=True)
        
        # Resample si nécessaire
        if sample_rate != TARGET_SAMPLE_RATE:
            resampler = T.Resample(orig_freq=sample_rate, new_freq=TARGET_SAMPLE_RATE)
            waveform = resampler(waveform)
            sample_rate = TARGET_SAMPLE_RATE
        
        logger.info(f"Audio chargé: {waveform.shape}, sample_rate: {sample_rate}")
        return waveform, sample_rate
        
    except Exception as e:
        logger.error("Erreur lors du chargement audio: %s", e)
        raise

def transcribe_audio_file(file_path):
    """Transcrit un fichier audio complet"""
    if not model_loaded:
        if initialization_error:
            raise RuntimeError(f"Modèle non initialisé: {initialization_error}")
        else:
            raise RuntimeError("Modèle en cours de chargement, veuillez réessayer")
    
    start_time = time.time()
    
    try:
        # Charger l'audio
        waveform, sample_rate = load_audio(file_path)
        
        # Transcription avec analyse détaillée
        transcription, avg_prob, word_details = analyze_whisper_output(
            processor, model, waveform, sample_rate
        )
        
        processing_time = time.time() - start_time
        
        result = {
            "transcription": transcription.strip(),
            "average_probability": avg_prob,
            "word_details": word_details,
            "processing_time_seconds": round(processing_time, 2),
            "audio_duration_seconds": round(waveform.shape[-1] / sample_rate, 2),
            "device_used": str(device)
        }
        
        logger.info(f"Transcription terminée en {processing_time:.2f}s")
        return result
        
    except Exception as e:
        logger.error(f"Erreur lors de la transcription: {e}")
        raise

def create_error_response(error_message, status_code=400, additional_info=None):
    """Crée une réponse d'erreur standardisée"""
    error_data = {
        "error": error_message,
        "status": "error",
        "timestamp": time.time()
    }
    
    if additional_info:
        error_data.update(additional_info)
    
    return json.dumps(error_data), status_code

def create_success_response(data):
    """Crée une réponse de succès standardisée"""
    response_data = {
        "status": "success",
        "timestamp": time.time(),
        "data": data
    }
    return json.dumps(response_data, indent=2)

def application(environ, start_response):
    """Application WSGI principale"""
    
    # Vérifier le modèle
    if not model_loaded and initialization_error:
        error_response, status_code = create_error_response(
            f"Service indisponible: {initialization_error}", 503
        )
        start_response(f'{status_code} Service Unavailable', [
            ('Content-Type', 'application/json'),
            ('Content-Length', str(len(error_response)))
        ])
        return [error_response.encode('utf-8')]
    
    # Endpoint de santé
    if environ['REQUEST_METHOD'] == 'GET' and environ.get('PATH_INFO', '/') == '/health':
        health_data = {
            "model_loaded": model_loaded,
            "device": str(device),
            "model_id": model_id
        }
        response = create_success_response(health_data)
        start_response('200 OK', [
            ('Content-Type', 'application/json'),
            ('Content-Length', str(len(response)))
        ])
        return [response.encode('utf-8')]
    
    # Traitement des requêtes POST
    if environ['REQUEST_METHOD'] == 'POST' and environ.get('PATH_INFO', '/') == '/transcribe':
        try:
            # Obtenir la taille du contenu
            try:
                content_length = int(environ.get('CONTENT_LENGTH', 0))
            except ValueError:
                content_length = 0
            
            if content_length == 0:
                error_response, status_code = create_error_response("Aucun contenu dans la requête")
                start_response(f'{status_code} Bad Request', [
                    ('Content-Type', 'application/json'),
                    ('Content-Length', str(len(error_response)))
                ])
                return [error_response.encode('utf-8')]
            
            if content_length > MAX_FILE_SIZE:
                error_response, status_code = create_error_response(
                    f"Fichier trop volumineux. Limite: {MAX_FILE_SIZE} bytes"
                )
                start_response(f'{status_code} Bad Request', [
                    ('Content-Type', 'application/json'),
                    ('Content-Length', str(len(error_response)))
                ])
                return [error_response.encode('utf-8')]
            
            # Parser le contenu multipart/form-data
            form = cgi.FieldStorage(fp=environ['wsgi.input'], environ=environ, keep_blank_values=True)
            
            if 'file' not in form:
                error_response, status_code = create_error_response(
                    "Aucun fichier trouvé. Utilisez le champ 'file'"
                )
                start_response(f'{status_code} Bad Request', [
                    ('Content-Type', 'application/json'),
                    ('Content-Length', str(len(error_response)))
                ])
                return [error_response.encode('utf-8')]
            
            file_item = form['file']
            
            if not file_item.file:
                error_response, status_code = create_error_response("Fichier vide ou invalide")
                start_response(f'{status_code} Bad Request', [
                    ('Content-Type', 'application/json'),
                    ('Content-Length', str(len(error_response)))
                ])
                return [error_response.encode('utf-8')]
            
            # Sauvegarder et traiter le fichier
            with temporary_file(suffix=".wav") as temp_file:
                # Écrire le contenu du fichier
                file_content = file_item.file.read()
                temp_file.write(file_content)
                temp_file.flush()
                
                # Transcription
                result = transcribe_audio_file(temp_file.name)
                
                # Ajouter des métadonnées
                result["file_size_bytes"] = len(file_content)
                result["filename"] = getattr(file_item, 'filename', 'unknown')
                
                response_body = create_success_response(result)
                
                start_response('200 OK', [
                    ('Content-Type', 'application/json'),
                    ('Content-Length', str(len(response_body)))
                ])
                
                return [response_body.encode('utf-8')]
                
        except Exception as e:
            logger.error(f"Erreur lors du traitement: {e}")
            traceback.print_exc()
            
            error_response, status_code = create_error_response(
                f"Erreur de traitement: {str(e)}", 500
            )
            start_response(f'{status_code} Internal Server Error', [
                ('Content-Type', 'application/json'),
                ('Content-Length', str(len(error_response)))
            ])
            return [error_response.encode('utf-8')]
    
    else:
        # Méthode non supportée
        error_response, status_code = create_error_response(
            "Seules les requêtes POST sont supportées", 405
        )
        start_response(f'{status_code} Method Not Allowed', [
            ('Content-Type', 'application/json'),
            ('Allow', 'POST'),
            ('Content-Length', str(len(error_response)))
        ])
        return [error_response.encode('utf-8')]

# Initialiser le modèle au démarrage
logger.info("Démarrage du service de transcription Whisper...")
initialize_model()