import time
from typing import Union, Any

from pydantic import BaseModel
from torchaudio.pipelines import HDEMUCS_HIGH_MUSDB_PLUS

from suno_utils.audio import Audio

from suno_utils.audio import Audio
from suno_utils.tasks.gpt.generation import *
from suno_utils.tasks.gpt_v2 import chirp_v1
from suno_utils.tasks import mert_v2
from transformers import AutoModel, Wav2Vec2FeatureExtractor
from suno_utils.worker.loader import S3Loader
from suno_utils.worker.schema import QueueItem
from suno_utils.worker.utils import download_models_to_dir


class ChirpQueueItem(BaseModel):
    id: str
    prompt_audio: Union[str, None] = None
    prompt_npz: Union[str, None] = None
    prompt_text: Union[str, None] = None
    metadata: dict
    gen_duration: Union[int, None] = 12


CHIRP_V1_MODELS = dict(
    combo_path="georg/trained_models/chirp_v1/xl_2.pt",
    mert_path="georg/trained_models/chirp_v1/semantic_centroids.npy",
    codec_path="georg/trained_models/chirp_v1/codec.pt",
    genre_tags_paths="georg/trained_models/chirp_v1/common_genre_tags.json",
    fasttext_path="georg/trained_models/chirp_v1/lid.176.bin",
)

MOUNT_PATH = "/suno/models"


class ChirpV1Worker(S3Loader):
    gpu_id: int

    def __init__(self, gpu_id):
        super().__init__()

        self.gpu_id = gpu_id

    def preload(self):
        start_time = time.time()

        chirp_v1.preload_models(
            centroids_filepath=f"{MOUNT_PATH}/{CHIRP_V1_MODELS['mert_path']}",
            gpt_ckpt_path=f"{MOUNT_PATH}/{CHIRP_V1_MODELS['combo_path']}",
            codec_ckpt_path=f"{MOUNT_PATH}/{CHIRP_V1_MODELS['codec_path']}",
            common_genre_tags_path=f"{MOUNT_PATH}/{CHIRP_V1_MODELS['genre_tags_paths']}",
            fasttext_ckpt_path=f"{MOUNT_PATH}/{CHIRP_V1_MODELS['fasttext_path']}",
            local_whisper="/suno/models/whisper",
        )

        finish_time = time.time()
        print(f"Preloading took {finish_time - start_time}s")

    @staticmethod
    def download_models(dir_path=MOUNT_PATH):
        """Use AWS CLI to download models if they don't exist."""
        files = list(CHIRP_V1_MODELS.values())
        download_models_to_dir(files, dir_path)

        HDEMUCS_HIGH_MUSDB_PLUS.get_model()
        AutoModel.from_pretrained(
            mert_v2.DEFAULT_MODEL_NAME,
            trust_remote_code=True,
            revision=mert_v2.DEFAULT_REVISION,
        )
        Wav2Vec2FeatureExtractor.from_pretrained(
            mert_v2.DEFAULT_MODEL_NAME,
            trust_remote_code=True,
            revision=mert_v2.DEFAULT_REVISION,
        )
        import whisper

        whisper.load_model("small.en")
        whisper.load_model("small")

    def process_item(self, item: QueueItem) -> tuple[list[Audio], list[Any], list[Any]]:
        history_audio = self._load_audio_prompt(item, "1.0.0.0") if item.prompt_audio else None

        options = item.metadata.get("options", {}) or {}

        n_batch = 2
        chaos = options.get("chaos", 0)
        chaos = float(chaos)
        # cfg = float(options.get("lyrics_strength", 1.25))
        # cfg_coef_tags = float(options.get("tags_strength", 1.75))
        text = item.prompt_text
        text_tags = item.metadata.get("tags", None)

        if text_tags == "random":
            text_tags = None

        _, raw_arrays, audios, text_tags = chirp_v1.generate_audio(
            text=text,
            history_text_guess=item.metadata.get("continued_from_prompt", text),
            history_audio=history_audio,
            text_tags=text_tags or None,
            n_batch=n_batch,
            chaos_level=chaos,
            max_gen_duration_s=40 if history_audio is None else 30,
            max_history_duration_s=10.01,
            min_eos_p=0.25,
            return_raw_arrays=True,
            allow_genre_randomization=text_tags != "-",
        )

        return audios, raw_arrays
