import time
from typing import Any


from suno_utils.audio import Audio
from suno_utils.tasks.gpt.generation import *
from suno_utils.tasks.gpt_v2 import bark_v2
from suno_utils.tasks import wavlm_v2
from transformers import WavLMModel
from suno_utils.worker.loader import S3Loader
from suno_utils.worker.schema import QueueItem
from suno_utils.worker.utils import download_models_to_dir


BARK_V2_MODELS = dict(
    combo_path="georg/trained_models/bark_v2/xl.pt",
    centroids_filepath="georg/trained_models/bark_v2/semantic_centroids.npy",
    codec_path="georg/trained_models/bark_v2/codec.pt",
)

MOUNT_PATH = "/suno/models"


class BarkV2Worker(S3Loader):
    gpu_id: int

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

        self.gpu_id = gpu_id

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

        bark_v2.preload_models(
            centroids_filepath=f"{MOUNT_PATH}/{BARK_V2_MODELS['centroids_filepath']}",
            gpt_ckpt_path=f"{MOUNT_PATH}/{BARK_V2_MODELS['combo_path']}",
            codec_ckpt_path=f"{MOUNT_PATH}/{BARK_V2_MODELS['codec_path']}",
        )

        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(BARK_V2_MODELS.values())
        download_models_to_dir(files, dir_path)

        WavLMModel.from_pretrained(wavlm_v2.DEFAULT_MODEL_NAME)

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

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

        is_variations = item.metadata.get("variations", False)

        oracle_arr = None
        if is_variations and item.prompt_audio:
            item.prompt_npz = item.prompt_audio.split(".")[0]
            history_npz = self._load_history_prompt(item)
            oracle_arr = history_npz["semantic_prompt"]

        n_batch = 2
        chaos = float(options.get("chaos", 1))
        cfg = float(options.get("text_strength", 1.15))
        text = item.prompt_text

        audios, semantic_arrays, coarse_arrays = bark_v2.generate_audio(
            text=text,
            history_audio=history_audio,
            n_batch=n_batch,
            temp_semantic=0.9 * chaos,
            cfg_coef=cfg,
            max_gen_duration_s=int(options.get("seconds", 45)),
            return_raw_arrays=True,
            oracle_arr=oracle_arr,
        )

        return audios, (semantic_arrays[1], coarse_arrays[1])
