from transformers import AutoProcessor, AutoModelForCTC, AutoConfig, AutoTokenizer, AutoFeatureExtractor,Wav2Vec2Processor
import torchaudio
import torch
import numpy as np
from typing import Optional, Union
import logging
import os
import sys
import json
import tempfile

sys.path.insert(0, os.path.dirname(__file__))


class TokenSet:
    """
    TokenSet

    Parameters
    ----------
    tokens : list[str]
        List of tokens.

    blank_token : Optional[str] = "<pad>"
        Blank token

    silence_token : Optional[str] = "|"
        Silence token

    unk_token : Optional[str] = "<unk>"
        Unk token

    bos_token : Optional[str] = "<s>"
        BOS token

    eos_token : Optional[str] = "</s>"
        EOS token

    letter_case: str
        Case mode to be applied to the transcription, can be 'lowercase', 'uppercase'
        or None (None == keep the original letter case). Default is lowercase.

    """

    def __init__(self, tokens: list[str], blank_token: Optional[str] = "<pad>",
                 silence_token: Optional[str] = "|",
                 unk_token: Optional[str] = "<unk>",
                 bos_token: Optional[str] = "<s>",
                 eos_token: Optional[str] = "</s>",
                 letter_case: str = "lowercase"):

        self.tokens = tokens
        self.blank_token = blank_token
        self.silence_token = silence_token
        self.unk_token = unk_token
        self.bos_token = bos_token
        self.eos_token = eos_token
        self.letter_case = letter_case

        if self.letter_case == "lowercase":
            self.tokens = [
                token.lower() if token not in self.special_tokens else token
                for token in self.tokens]
        elif self.letter_case == "uppercase":
            self.tokens = [
                token.upper() if token not in self.special_tokens else token
                for token in self.tokens]

        if blank_token not in tokens:
            logging.warning(
                f"blank_token {blank_token} not in provided tokens. It will be added to the list of tokens")
            self.tokens.append(blank_token)

        if silence_token not in tokens:
            logging.warning(
                f"silence_token {silence_token} not in provided tokens. It will be added to the list of tokens")
            self.tokens.append(silence_token)

        if unk_token not in tokens:
            logging.warning(
                f"unk_token {unk_token} not in provided tokens. It will be added to the list of tokens")
            self.tokens.append(unk_token)

        if bos_token not in tokens:
            logging.warning(
                f"bos_token {bos_token} not in provided tokens. It will be added to the list of tokens")
            self.tokens.append(bos_token)

        if eos_token not in tokens:
            logging.warning(
                f"eos_token {eos_token} not in provided tokens. It will be added to the list of tokens")
            self.tokens.append(eos_token)

        self.id_by_token = {token: i for i, token in enumerate(self.tokens)}
        self.token_by_id = {i: token for i, token in enumerate(self.tokens)}

    @property
    def blank_token_id(self):
        return self.id_by_token[self.blank_token]

    @property
    def silence_token_id(self):
        return self.id_by_token[self.silence_token]

    @property
    def unk_token_id(self):
        return self.id_by_token[self.unk_token]

    @property
    def bos_token_id(self):
        return self.id_by_token[self.bos_token]

    @property
    def eos_token_id(self):
        return self.id_by_token[self.eos_token]

    @property
    def non_special_tokens(self):
        return [token for token in self.tokens if
                token not in self.special_tokens]

    @property
    def special_tokens(self):
        return [self.blank_token, self.silence_token, self.unk_token,
                self.bos_token, self.eos_token]

    @property
    def size(self):
        return len(self.tokens)

    def to_processor(self,
                     model_name_or_path: str = "facebook/wav2vec2-large-xlsr-53"):

        tokens_dict = {v: i for i, v in enumerate(self.tokens)}

        with tempfile.TemporaryDirectory() as tmpdirname:
            vocab_path = os.path.join(tmpdirname, "vocab.json")

            with open(vocab_path, "w") as vocab_file:
                json.dump(tokens_dict, vocab_file)

            config = AutoConfig.from_pretrained(model_name_or_path)
            config_for_tokenizer = config if config.tokenizer_class is not None else None
            tokenizer_type = config.model_type if config.tokenizer_class is None else None

            tokenizer = AutoTokenizer.from_pretrained(
                tmpdirname,
                config=config_for_tokenizer,
                tokenizer_type=tokenizer_type,
                bos_token=self.bos_token,
                eos_token=self.eos_token,
                unk_token=self.unk_token,
                pad_token=self.blank_token,
                word_delimiter_token=self.silence_token,
                do_lower_case=False,
                # do_lower_case=self.letter_case == "lowercase",
                # TODO: fix transformers/models/wav2vec2/tokenization_wav2vec2.py:199
            )

            feature_extractor = AutoFeatureExtractor.from_pretrained(
                model_name_or_path)

            return Wav2Vec2Processor(feature_extractor=feature_extractor,
                                     tokenizer=tokenizer)

    @classmethod
    def from_processor(cls, processor: Wav2Vec2Processor,
                       letter_case: str = "lowercase"):

        blank_token = processor.tokenizer.pad_token
        silence_token = processor.tokenizer.word_delimiter_token
        unk_token = processor.tokenizer.unk_token
        bos_token = processor.tokenizer.bos_token
        eos_token = processor.tokenizer.eos_token
        tokens = [x for x in processor.tokenizer.convert_ids_to_tokens(
            range(0, processor.tokenizer.vocab_size))]

        return cls(tokens, blank_token, silence_token, unk_token, bos_token,
                   eos_token, letter_case)

    def save(self, path: str):

        with open(path, "w", encoding="utf-8") as f:
            json.dump(self.__dict__, f, indent=2, ensure_ascii=False)

    @classmethod
    def load(cls, path: str):

        with open(path, encoding="utf-8") as f:
            o = json.load(f)
            return cls(o["tokens"], o["blank_token"], o["silence_token"],
                       o["unk_token"], o["bos_token"], o["eos_token"],
                       o["letter_case"])


class Decoder:
    """
    Decoder

    Parameters
    ----------
    token_set : TokenSet
        The TokenSet object to use for decoding.

    skip_special_tokens: Optional[bool] = True
        If True, skip the special tokens in the TokenSet during decoding.

    ms_per_timestep : Optional[int] = 20
        The number of milliseconds per timestep.

        The magic number 20 comes from the Wav2Vec2 convolutional layers. I'll try to explain this a little:
            The convolutional layers of the feature extractor have padding=0, kernel=(10,3,3,3,3,2,2) and strides=(5,2,2,2,2,2,2).
            So to map the model output (in timesteps) to the original audio (in milliseconds) we need to apply
            the convolution output size formula ( ((input-kernel+2*padding)/stride)+1) ) on each layer, given some waveform input.
            However, I don't recommend that approach because it needs some extra computation and it's not really precise,
            because there're a lot of overlappings due to the kernel and stride sizes desynchronization.
            So a good and simple approximation can be extracted using just the convolutional layers strides, more specifically,
            using the product of the strides:

            MS_PER_TIMESTEP = STRIDES_PRODUCT / WAVEFORM_POINTS_PER_MS
            MS_PER_TIMESTEP = 5*2^6 / 16 <- remember that the wav2vec input is in 16000Hz
            MS_PER_TIMESTEP = 20 <- the magic number :)

    probability_offset : Optional[float] = 1
        The probability offset to use when calculating the probability of a token when the end_timestep of a token isn't provided,
        i. e., how many timesteps to look around to calc the probability.

    """

    def __init__(self, token_set: TokenSet,
                 skip_special_tokens: Optional[bool] = True,
                 ms_per_timestep: Optional[int] = 20,
                 probability_offset: Optional[float] = 1):

        self.token_set = token_set
        self.skip_special_tokens = skip_special_tokens
        self.ms_per_timestep = ms_per_timestep
        self.probability_offset = probability_offset

    def _get_predictions(self, logits: torch.Tensor) -> list[dict]:
        """
        Get the predicted ids given the model's output logits.

        Parameters:
        ----------
            logits: torch.Tensor
                Model's output tensors of shape (BATCH_SIZE, TIMESTEPS, TOKEN_SET_SIZE)

        Returns:
        ----------
            list[dict]: prediction list of size BATCH_SIZE with format:
                [{
                    "ids": list,
                    "start_timesteps": list,
                    "end_timesteps": list,
                }, ...]
        """

        raise NotImplementedError()

    def _ctc_decode(self,
                    predicted_ids: Union[torch.Tensor, np.ndarray, list[int]],
                    return_timesteps: Optional[bool] = True) -> list[dict]:
        """
        Decode the predicted using the CTC decoding algorithm.

        Parameters:
        ----------
            predicted_ids: Union[torch.Tensor, np.ndarray, list[int]]
                Model's output tensors of shape (BATCH_SIZE, TIMESTEPS, TOKEN_SET_SIZE)

            return_timesteps: Optional[bool] = True
                If True, return the timesteps of the decoded sequence.

        Returns:
        ----------
            list[dict]: decoded prediction list of size BATCH_SIZE with format:
                [{
                    "ids": list,
                    "start_timesteps": list,
                    "end_timesteps": list,
                }, ...]
        """

        predictions = []

        for i in range(len(predicted_ids)):  # for each item in the batch

            i_predicted_ids = []
            i_start_timesteps = []
            i_end_timesteps = []
            previous_predicted_id = None

            for t in range(
                    len(predicted_ids[i])):  # for each timestep in the item

                predicted_id = int(predicted_ids[i][t])

                if predicted_id != self.token_set.blank_token_id:

                    if len(i_predicted_ids) == 0 or previous_predicted_id == self.token_set.blank_token_id or predicted_id != \
                            i_predicted_ids[-1]:
                        i_predicted_ids.append(predicted_id)
                        if return_timesteps:
                            i_start_timesteps.append(t)
                            i_end_timesteps.append(t + 1)

                    elif predicted_id == i_predicted_ids[
                        -1] and return_timesteps:
                        i_end_timesteps[-1] = t

                previous_predicted_id = predicted_id

            predictions.append({
                "ids": i_predicted_ids,
                "start_timesteps": i_start_timesteps if return_timesteps else None,
                "end_timesteps": i_end_timesteps if return_timesteps else None,
            })

        return predictions

    def __call__(self, logits: torch.Tensor) -> list[dict]:
        """
        Getting the predictions given the model's output logits.

        Parameters:
        ----------
            logits: torch.Tensor
                Model's output tensors of shape (BATCH_SIZE, TIMESTEPS, TOKEN_SET_SIZE)

        Returns:
        ----------
            list[dict]: Decoded prediction list of size BATCH_SIZE with format:
                [{
                    "transcription": list,
                    "start_timesteps": list,
                    "end_timesteps": list,
                }, ...]
        """

        result = []

        predictions = self._get_predictions(logits)
        logits_probs = torch.nn.functional.softmax(logits.float(), dim=-1).to(
            "cpu").detach()

        for i, prediction in enumerate(predictions):

            transcription = ""
            transcription_start_timestamps = []
            transcription_end_timestamps = []
            transcription_probabilities = []

            if "transcription" in prediction:
                transcription = prediction["transcription"]
                J = range(len(transcription))
            else:
                J = range(len(prediction["ids"]))

            for j in J:

                if "transcription" in prediction:
                    token = transcription[j] if transcription[
                                                    j] != " " else self.token_set.silence_token
                    if token not in self.token_set.tokens:
                        token = self.token_set.unk_token
                    predicted_id = self.token_set.id_by_token[token]
                else:
                    predicted_id = prediction["ids"][j]

                    if predicted_id == self.token_set.silence_token_id:
                        transcription += " "
                    elif self.skip_special_tokens and self.token_set.tokens[
                        predicted_id] in self.token_set.special_tokens:
                        continue
                    else:
                        transcription += self.token_set.tokens[predicted_id]

                if prediction["start_timesteps"] is not None:
                    start_timestep = prediction["start_timesteps"][j]
                    transcription_start_timestamps.append(
                        int(start_timestep * self.ms_per_timestep))

                    # as we report the character based probability and more than one timestep can be responsable for the character prediction,
                    # when a start_timestep and end_timestep are provided we'll report the mean value of the this range of timesteps,
                    # otherwise we'll report the mean probability of a window defined be the start_timestep_t and start_timestep_t+1

                    if prediction["end_timesteps"] is not None:
                        window_end_timestep = prediction["end_timesteps"][j]
                    else:
                        window_end_timestep = start_timestep + 1

                    if start_timestep == window_end_timestep:  # it needs to have at least one timestep of difference
                        window_end_timestep += 1

                    window_probabilities = [x[predicted_id] for x in
                                            logits_probs[i][
                                            start_timestep:window_end_timestep]]
                    probability = float(np.mean(window_probabilities))
                    transcription_probabilities.append(probability)

                if prediction["end_timesteps"] is not None:
                    end_timestep = prediction["end_timesteps"][j]
                    transcription_end_timestamps.append(
                        int(end_timestep * self.ms_per_timestep))

                    # probability = float(logits_probs[i][start_timestep][predicted_id])
                    # transcription_probabilities[-1] = probability

            # transcription trimming
            if len(transcription) > 0:

                left_offset = len(transcription) - len(transcription.lstrip())
                right_offset = len(transcription) - len(transcription.rstrip())

                transcription = transcription[
                                left_offset:len(transcription) - right_offset]

                if len(transcription_start_timestamps) > 0:
                    transcription_start_timestamps = transcription_start_timestamps[
                                                     left_offset:len(
                                                         transcription_start_timestamps) - right_offset]
                if len(transcription_end_timestamps) > 0:
                    transcription_end_timestamps = transcription_end_timestamps[
                                                   left_offset:len(
                                                       transcription_end_timestamps) - right_offset]
                if len(transcription_probabilities) > 0:
                    transcription_probabilities = transcription_probabilities[
                                                  left_offset:len(
                                                      transcription_probabilities) - right_offset]

            result.append({
                "transcription": transcription,
                "start_timestamps": transcription_start_timestamps if len(
                    transcription_start_timestamps) > 0 else None,
                "end_timestamps": transcription_end_timestamps if len(
                    transcription_end_timestamps) > 0 else None,
                "probabilities": transcription_probabilities if len(
                    transcription_probabilities) > 0 else None,
            })

        return result


class GreedyDecoder(Decoder):
    """
    Greedy decoder

    Parameters
    ----------
    token_set : TokenSet
        The TokenSet object to use for decoding.
    """

    def __init__(self, token_set: TokenSet):
        super().__init__(token_set)

    def _get_predictions(self, logits: torch.Tensor):
        predicted_ids = torch.argmax(logits, dim=-1)
        predictions = self._ctc_decode(predicted_ids)

        return predictions




class AudioTranscriber:
    def __init__(self, model: AutoModelForCTC, processor: AutoProcessor):
        self.model = model
        self.processor = processor
        self.token_set = TokenSet.from_processor(processor)
        self.greedy_decoder = GreedyDecoder(self.token_set)

    def transcribe_audio(self, audio_waveform, sample_rate):
        # If the audio sample rate is not 16kHz, resample it
        if sample_rate != 16000:
            audio_waveform = torchaudio.functional.resample(
                audio_waveform, orig_freq=sample_rate, new_freq=16000
            )
            sample_rate = 16000

        # If stereo, convert to mono by averaging the channels
        if audio_waveform.shape[0] > 1:
            audio_waveform = torch.mean(audio_waveform, dim=0)

        # Prepare the inputs for the model
        inputs = self.processor(
            audio_waveform.squeeze(), sampling_rate=sample_rate,
            return_tensors="pt"
        )

        # Perform inference
        with torch.no_grad():
            logits = self.model(inputs.input_values).logits
        result = self.greedy_decoder(logits)[0]
        return result



# Initialize the transcriber
model_name = "Cnam-LMSSC/wav2vec2-french-phonemizer"
processor = AutoProcessor.from_pretrained(model_name)
model = AutoModelForCTC.from_pretrained(model_name)
transcriber = AudioTranscriber(model=model, processor=processor)


import cgi
# curl -X POST --data-binary "@charmant.wav" https://alex.fonetix.org/
def application(environ, start_response):
    # Vérifier le type de requête
    if environ['REQUEST_METHOD'] == 'POST':
        try:
            try:
                request_body_size = int(environ.get('CONTENT_LENGTH', 0))
            except (ValueError):
                request_body_size = 0


            # Parse le contenu de la requête multipart/form-data
            form = cgi.FieldStorage(fp=environ['wsgi.input'], environ=environ, keep_blank_values=True)
            file_item = form['file']  # Assurez-vous que le champ du fichier s'appelle bien 'file'

            if not file_item.file:
                raise Exception("Aucun fichier audio trouvé dans la requête.")

            # Enregistrer le fichier audio dans un fichier temporaire
            with tempfile.NamedTemporaryFile(delete=False, suffix=".wav") as audio_file:
                audio_file.write(file_item.file.read())
                audio_file_path = audio_file.name

            # Vérifier si le fichier existe sur le disque
            file_exists = os.path.exists(audio_file_path)
            if not file_exists:
                raise Exception("Le fichier temporaire n'existe pas sur le disque après la tentative de création.")

            # Vérifier la taille du fichier temporaire
            file_size = os.path.getsize(audio_file_path)
            if file_size == 0:
                raise Exception("Le fichier temporaire est vide après l'enregistrement.")

            # Charger les données audio en utilisant torchaudio et vérifier le format
            try:
                audio_waveform, sample_rate = torchaudio.load(audio_file_path)
            except Exception as e:
                raise Exception(f"Erreur lors du chargement du fichier audio : {str(e)}")

            # Transcrire l'audio
            transcription = transcriber.transcribe_audio(audio_waveform, sample_rate)
            # Créer la réponse JSON avec la transcription
            response_data = {
                "message": "Transcription effectuée avec succès",
                "transcription": transcription,
                "file_size": file_size
            }
            response_body = json.dumps(response_data)

            # Définir le statut et les en-têtes de réponse
            status = '200 OK'
            response_headers = [
                ('Content-Type', 'application/json'),
                ('Content-Length', str(len(response_body)))
            ]
            start_response(status, response_headers)

            return [response_body.encode('utf-8')]

        except Exception as e:
            # En cas d'erreur, retourner une réponse d'erreur JSON
            error_info = {
                "error": str(e),
                "audio_file_path": audio_file_path if 'audio_file_path' in locals() else None,
                "content_length": environ.get('CONTENT_LENGTH'),
                "file_size": file_size if 'file_size' in locals() else None,
                "file_exists": file_exists if 'file_exists' in locals() else None
            }
            error_response = json.dumps(error_info)
            status = '400 Bad Request'
            response_headers = [
                ('Content-Type', 'application/json'),
                ('Content-Length', str(len(error_response)))
            ]
            start_response(status, response_headers)
            return [error_response.encode('utf-8')]

    else:
        # Pour les autres méthodes, retourner un message simple
        start_response('200 OK', [('Content-Type', 'text/plain')])
        return [b"This endpoint only supports POST requests."]