import os
import subprocess
import time
import traceback

import torch

from suno_utils.tasks.encodec import decode as codec_decode
from suno_utils.tasks.encodec import encode as codec_encode
from suno_utils.tasks.nano.gpt import (
    generate_coarse,
    generate_fine,
    generate_semantic,
    generate_text_semantic,
)
from suno_utils.tasks.nano.gpt import preload_models as preload_gpt_models
from suno_utils.tasks.wavlm import (
    encode as wavlm_encode,
)
from suno_utils.tasks.wavlm import (
    preload_models as preload_wavlm_models,
)
from suno_utils.worker.schema import QueueItem
from suno_utils.tasks import real_or_fake
from suno_utils.worker.loader import S3Loader


def _clear_cuda_cache():
    if torch.cuda.is_available():
        torch.cuda.empty_cache()
        torch.cuda.synchronize()


class ModelV1Worker(S3Loader):
    gpu_id: int

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

        self.gpu_id = gpu_id

    def preload_wavlm(self):
        centroids_filepath = "/suno/models/wavlm_v0/10k_centroids_wavlm.npy"
        _ = preload_wavlm_models(device=f"cuda:{self.gpu_id}", centroids_filepath=centroids_filepath)
        print("Preloaded WavLM models.")

    def preload(self):
        """Preload models onto GPU."""
        # _clear_cuda_cache()
        start_time = time.time()
        semantic_ckpt_path = "/suno/models/wavlm_v0/lg_semantic.pt"
        coarse_ckpt_path = "/suno/models/wavlm_v0/lg_coarse.pt"
        fine_ckpt_path = "/suno/models/md_fine.pt"
        text_ckpt_path = "/suno/models/wavlm_v0/md_text_semantic.pt"

        centroids_filepath = "/suno/models/wavlm_v0/10k_centroids_wavlm.npy"
        _ = preload_wavlm_models(device=f"cuda:{self.gpu_id}", centroids_filepath=centroids_filepath)
        print("Preloaded WavLM models.")

        preload_gpt_models(
            semantic_ckpt_path=semantic_ckpt_path,
            coarse_ckpt_path=coarse_ckpt_path,
            fine_ckpt_path=fine_ckpt_path,
            text_ckpt_path=text_ckpt_path,
        )

        real_or_fake.preload_model(
            "/suno/models/wavlm_v0/real_or_fake/model.pt", device=f"cuda:{self.gpu_id}"
        )

        preload_time = time.time() - start_time
        print(f"Finished preloading models in {preload_time}s.")

    def generate_saved(self, item: QueueItem):
        """Use a saved history prompt (NPZ) and text to generate."""
        history_prompt = self._load_history_prompt(item)

        prompt_semantic_arr = history_prompt["semantic_prompt"]
        prompt_coarse_arr = history_prompt["coarse_prompt"]
        prompt_fine_arr = history_prompt["fine_prompt"]

        text_prompt = item.prompt_text

        gen_semantic_arr = generate_text_semantic(
            text_prompt,
            temp=item.metadata.get("gen_semantic_temp", 0.8),
            semantic_history=prompt_semantic_arr,
        )
        gen_coarse_arr = generate_coarse(
            gen_semantic_arr,
            x_history=(prompt_semantic_arr, prompt_coarse_arr),
            temp=item.metadata.get("gen_coarse_temp", 0.9),
        )
        gen_fine_arr = generate_fine(
            gen_coarse_arr,
            x_fine_history=prompt_fine_arr,
            temp=item.metadata.get("gen_fine_temp", 0.3),
        )

        audio = codec_decode(gen_fine_arr, n_codebooks=8)
        self._write_audio(item, audio, [gen_semantic_arr, gen_coarse_arr, gen_fine_arr])

    def render_history_prompt(self, id: str):
        history_prompt = self._decode_history_prompt(id)
        prompt_fine_arr = history_prompt["fine_prompt"]
        audio = codec_decode(prompt_fine_arr, n_codebooks=8)

        self._write_audio(QueueItem(id=id, metadata={}), audio, [])

    def generate_prompt(self, item: QueueItem):
        prompt_audio = self._load_audio_prompt(item)

        # make embeddings
        prompt_semantic_arr = wavlm_encode(prompt_audio, device=f"cuda:{self.gpu_id}")
        prompt_coarse_arr = codec_encode(prompt_audio, n_codebooks=2)
        prompt_fine_arr = codec_encode(prompt_audio, n_codebooks=8)

        text_prompt = item.prompt_text

        gen_semantic_arr = generate_text_semantic(
            text_prompt,
            temp=item.metadata.get("gen_semantic_temp", 0.8),
            semantic_history=prompt_semantic_arr,
        )
        gen_coarse_arr = generate_coarse(
            gen_semantic_arr,
            x_history=(prompt_semantic_arr, prompt_coarse_arr),
            temp=item.metadata.get("gen_coarse_temp", 0.9),
        )
        gen_fine_arr = generate_fine(
            gen_coarse_arr,
            x_fine_history=prompt_fine_arr,
            temp=item.metadata.get("gen_fine_temp", 0.3),
        )

        audio = codec_decode(gen_fine_arr, n_codebooks=8)
        self._write_audio(item, audio, [gen_semantic_arr, gen_coarse_arr, gen_fine_arr])

    def generate_continuation(self, item: QueueItem):
        prompt_audio = self._load_audio_prompt(item)

        # make embeddings
        prompt_semantic_arr = wavlm_encode(prompt_audio, device=f"cuda:{self.gpu_id}")
        prompt_coarse_arr = codec_encode(prompt_audio, n_codebooks=2)
        prompt_fine_arr = codec_encode(prompt_audio, n_codebooks=8)

        gen_semantic_arr = generate_semantic(
            x=prompt_semantic_arr,
            gen_duration_s=item.gen_duration,
            temp=item.metadata.get("gen_semantic_temp", 0.8),
        )
        gen_coarse_arr = generate_coarse(
            gen_semantic_arr,
            x_history=(prompt_semantic_arr, prompt_coarse_arr),
            temp=item.metadata.get("gen_coarse_temp", 0.9),
        )
        gen_fine_arr = generate_fine(
            gen_coarse_arr,
            x_fine_history=prompt_fine_arr,
            temp=item.metadata.get("gen_fine_temp", 0.3),
        )

        audio = codec_decode(gen_fine_arr, n_codebooks=8)
        self._write_audio(item, audio, [gen_semantic_arr, gen_coarse_arr, gen_fine_arr])

    def generate_text(self, item: QueueItem):
        gen_semantic_arrs = [
            generate_text_semantic(
                item.prompt_text,
                temp=item.metadata.get("gen_semantic_temp", 0.8),
            )
            for i in range(3)
        ]

        scored = list(
            zip(
                real_or_fake.score_audio(gen_semantic_arrs, device=f"cuda:{self.gpu_id}"),
                gen_semantic_arrs,
            )
        )
        scored = sorted(scored, key=lambda x: x[0])

        gen_semantic_arr = scored[-1][1]

        gen_coarse_arr = generate_coarse(
            gen_semantic_arr,
            temp=item.metadata.get("gen_coarse_temp", 0.9),
        )
        gen_fine_arr = generate_fine(
            gen_coarse_arr,
            temp=item.metadata.get("gen_fine_temp", 0.3),
        )

        audio = codec_decode(gen_fine_arr, n_codebooks=8)
        self._write_audio(item, audio, [gen_semantic_arr, gen_coarse_arr, gen_fine_arr])

    def generate_unconditional(self, item: QueueItem):
        gen_semantic_arr = generate_semantic(
            gen_duration_s=item.gen_duration,
            temp=item.metadata.get("gen_semantic_temp", 0.8),
        )
        gen_coarse_arr = generate_coarse(
            gen_semantic_arr,
            temp=item.metadata.get("gen_coarse_temp", 0.9),
        )
        gen_fine_arr = generate_fine(
            gen_coarse_arr,
            temp=item.metadata.get("gen_fine_temp", 0.3),
        )

        audio = codec_decode(gen_fine_arr, n_codebooks=8)
        self._write_audio(item, audio, [gen_semantic_arr, gen_coarse_arr, gen_fine_arr])

    def process_item(self, item):
        try:
            if item.prompt_npz and item.prompt_text:
                self.generate_saved(item)
            elif not item.prompt_text and not item.prompt_audio:
                self.generate_unconditional(item)
            elif not item.prompt_text:
                self.generate_continuation(item)
            elif not item.prompt_audio:
                self.generate_text(item)
            else:
                self.generate_prompt(item)
            ok = True
        except:
            traceback.print_exc()
            print("Errored", item)
            ok = False

        return ok

    @staticmethod
    def download_models():
        """Use AWS CLI to download models if they don't exist."""
        dir_path = "/suno/models/"

        files = [
            "wavlm_v0/md_text_semantic.pt",
            "wavlm_v0/lg_coarse.pt",
            "wavlm_v0/lg_semantic.pt",
            "wavlm_v0/10k_centroids_wavlm.npy",
            "md_fine.pt",
            "wavlm_v0/real_or_fake/model.pt",
        ]

        for f in files:
            full_path = os.path.abspath(os.path.join(dir_path, f))
            if not os.path.exists(full_path):
                parent_dir = os.path.dirname(full_path)
                os.makedirs(parent_dir, exist_ok=True)

                print("Downloading", full_path)
                cmd = [
                    "aws",
                    "s3",
                    "cp",
                    f"s3://suno-data/georg/checkpoints/{f}",
                    full_path,
                ]
                subprocess.run(cmd)
            else:
                print("Skipping", f)
