import time
from typing import Union

from pydantic import BaseModel
from torchaudio.pipelines import HDEMUCS_HIGH_MUSDB_PLUS

from suno_utils.audio import Audio
from suno_utils.tasks.gpt import chirp_v0
from suno_utils.tasks import mert
from encodec import EncodecModel
from transformers import AutoModel

from .loader import S3Loader
from .schema import QueueItem
from .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_MODELS = dict(
    semantic_ckpt_path="georg/trained_models/chirp_v0/xl_semantic.pt",
    coarse_ckpt_path="georg/trained_models/chirp_v0/lg_coarse.pt",
    text_tokenizer_path="georg/trained_models/chirp_v0/tokenizer.json",
    centroids_filepath="georg/trained_models/chirp_v0/4x1k_centroids_mert.npy",
)

MOUNT_PATH = "/suno/models"


class ChirpV0Worker(S3Loader):
    gpu_id: int

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

        self.gpu_id = gpu_id

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

        chirp_v0.preload_models(**{k: f"{MOUNT_PATH}/{v}" for k, v in CHIRP_MODELS.items()})

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

        HDEMUCS_HIGH_MUSDB_PLUS.get_model()
        EncodecModel.encodec_model_24khz()
        AutoModel.from_pretrained(
            mert.DEFAULT_MODEL_NAME,
            trust_remote_code=True,
            revision=mert.DEFAULT_REVISION,
        )

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

        if item.prompt_text:
            audios = chirp_v0.text_to_audio(
                item.prompt_text + "\n\n-",
                history_audio=history_audio,
                n_batch=N_BATCH,
                max_gen_duration_s=30 if not history_audio else 20,
                cfg_gamma=1.8,
                # semantic_top_k=200,
                max_history_duration_s=10,
            )
        else:
            audios = chirp_v0.generate_audio(
                history_audio=history_audio,
                n_batch=N_BATCH,
                gen_duration_s=20,
                max_history_duration_s=15,
            )

        return audios
