import json
import time
from typing import Any
import pathlib
from uuid import uuid4

import modal
from torchaudio.pipelines import HDEMUCS_HIGH_MUSDB_PLUS
from openai import OpenAI


from suno_utils.audio import Audio
from suno_utils.gpt import chirp_v2_5 as chirp_v2
from suno_utils.worker.generate_song_lyrics import (
    generate_song_lyrics_with_genre_tags,
    moderate_user_inputs,
    ModerationSuccess,
    ModerationError,
)

from suno_utils.worker.loader import S3Loader
from suno_utils.worker.modal_base import MODAL_MOUNTS
from suno_utils.worker.schema import QueueItem
from suno_utils.worker.utils import recursive_ls_dir


base_image = (
    modal.Image.debian_slim()
    .apt_install("curl", "ffmpeg", "sox", "unzip", "libsox-fmt-mp3")
    .run_commands(
        [
            'curl "https://awscli.amazonaws.com/awscli-exe-linux-x86_64.zip" -o "awscliv2.zip"',
            "unzip -q awscliv2.zip",
            "./aws/install",
        ]
    )
    .pip_install(
        "torch==2.1.0.dev20230831+cu118",
        "torchaudio==2.2.0.dev20230831+cu118",
        index_url="https://download.pytorch.org/whl/nightly/cu118",
    )
    .pip_install_private_repos(
        "github.com/suno-ai/glockenspiel.git@f05e2f251#subdirectory=descript-audio-codec&egg=descript-audio-codec",
        git_user="mcamac",
        secrets=[modal.Secret.from_name("victor-modal-github-token")],
    )
    .pip_install(
        "boto3",
        "transformers",
        "tokenizers",
        "encodec",
        "ctc_segmentation",
        "psutil",
        "redis",
        "gradio",
        "pydantic",
        "nnAudio",
    )
    .pip_install_from_pyproject(
        str(pathlib.Path(__file__).parent.parent.parent / "pyproject.toml"),
    )
    .run_commands(
        "FLASH_ATTENTION_SKIP_CUDA_BUILD=TRUE pip install flash-attn --no-build-isolation",
    )
)


MOUNT_PATH = "/suno/models"


class ChirpV2Worker(S3Loader):
    gpu_id: int

    def __init__(self, gpu_id, bg_image=None):
        super().__init__(bg_image=bg_image)

        self.gpu_id = gpu_id
        self.bg_image = bg_image

    def preload(self):
        start_time = time.time()
        chirp_v2.preload_models(
            cache_dir=MOUNT_PATH,
            local_whisper="/suno/models/whisper",
            codec_ckpt_path="s3://suno-data/georg/models/codec/dac_2c_25x12.pt",
            gpt_ckpt_path="s3://suno-data/georg/checkpoints/chirp_v2_5/7b_ft_fix.pt",
        )

        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."""
        chirp_v2.preload_models(
            fetch_only=True,
            cache_dir=dir_path,
            codec_ckpt_path="s3://suno-data/georg/models/codec/dac_2c_25x12.pt",
            gpt_ckpt_path="s3://suno-data/georg/checkpoints/chirp_v2_5/7b_ft_fix.pt",
        )
        recursive_ls_dir(dir_path)

        HDEMUCS_HIGH_MUSDB_PLUS.get_model()

        import whisper

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

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

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

        n_batch = 2
        chaos = options.get("chaos", 1)
        chaos = float(chaos)

        text = item.prompt_text or ""
        text_tags = item.metadata.get("tags", None)
        if text_tags == "random":
            text_tags = None
        detected_lang = chirp_v2._get_text_lang(text)
        cfg_coef = 1.25
        # cfg_coef = float(options.get("lyrics_strength", 1.2))
        # condition on non-english, higher cfg text
        if (not detected_lang == "en") and text:
            cfg_coef = 1.7
            print("detected non-english lyrics, increase cfg")
        # condition on genre, rap
        elif text_tags and ("rap" in text_tags or "hip-hop" in text_tags):
            cfg_coef = 1.4
        text_neg_tags = None
        if "instrumental" not in (text_tags or "") and len(text.split()) > 5:
            text_neg_tags = "instrumental noise"
        # cfg_coef_tags = float(options.get("tags_strength", 1.75))
        cfg = chirp_v2.GenerationConfig(
            text=text,
            text_tags=text_tags or None,
            text_neg_tags=text_neg_tags,
            n_batch=n_batch,
            max_gen_duration_s=80,
            cfg_coef_neg_tags=0,
            # cfg_coef=cfg_coef,
            # cfg_coef_tags=1.9,
            # temp_semantic=0.9,
            # temp_coarse=0.85,
            # min_text_offset=128,
            # top_k_semantic=1000,
            # top_k_coarse=1000,
            # top_p_semantic=None,
            # top_p_coarse=None,
            # rep_penality=0,
        )

        MAX_HISTORY_DURATION_S = 50
        MAX_CONTINUATION_DURATION_S = 60
        if isinstance(history_audio, Audio):
            history_audio = history_audio.get_segment(
                from_s=max(0, history_audio.duration_s - MAX_HISTORY_DURATION_S)
            )
            cfg = cfg.modify(max_gen_duration_s=MAX_CONTINUATION_DURATION_S)
        elif history_audio is not None:  # NPZ
            # history_arrs are passed through GenerationConfig
            cfg = cfg.modify(
                max_gen_duration_s=MAX_CONTINUATION_DURATION_S,
                history_arr=history_audio,
                cfg_coef_tags=0,
            )
            history_audio = None

        _, raw_arrays, audios = chirp_v2.generate_audio(
            cfg,
            chaos_level=chaos,
            history_audio=history_audio,
            return_raw_arrays=True,
            max_history_duration_s=MAX_HISTORY_DURATION_S,
        )

        return audios, raw_arrays


aws_secret = modal.Secret.from_name("studio-aws")
SECRETS = [
    aws_secret,
    modal.Secret.from_dict(
        {
            "SUNO_ASSETS_PATH": "/suno/models/assets",
            "XDG_CACHE_HOME": "/suno/models/",
        }
    ),
    modal.Secret.from_name("openai-secret"),
]


def download_model_wrapper_a5():
    ChirpV2Worker.download_models()


image = base_image.run_function(download_model_wrapper_a5, secrets=SECRETS)
STUB_NAME = "chirp-v2-7b-staging"
stub = modal.Stub(STUB_NAME, image=image)


@stub.cls(
    cpu=2.0,
    # memory=16384,
    gpu=modal.gpu.A10G(count=1),
    secrets=SECRETS,
    timeout=300,
    container_idle_timeout=400,
    mounts=MODAL_MOUNTS,
    memory=60000,
    retries=modal.Retries(
        max_retries=3,
        backoff_coefficient=2.0,
        initial_delay=10.0,
    ),
    keep_warm=1,
    concurrency_limit=200,
)
class ChirpV2Stub:
    def __enter__(self):
        import torch

        num_gpus = torch.cuda.device_count()
        print(f"Found {num_gpus} GPUs.")
        recursive_ls_dir("/suno/models")

        self.worker = ChirpV2Worker(0, bg_image="/suno/models/assets/wave-bg-2.png")
        self.worker.preload()

        self.openai_client = OpenAI()

    @modal.method()
    def generate(self, queue_item: str):
        import json

        from suno_utils.worker.schema import QueueItem

        print(queue_item)
        start_time = time.time()
        item = QueueItem(**json.loads(queue_item))
        item_id = item.id
        print(item.metadata)
        is_square = True
        ids = []

        if user_prompt := item.metadata.get("gpt_description_prompt"):
            try:
                title, lyrics, genre_tags = generate_song_lyrics_with_genre_tags(
                    self.openai_client, chirp_v2.text_lang_model, user_prompt
                )
            except ModerationError as e:
                import traceback

                self.worker.notify_finish(
                    item,
                    {
                        "id": item.id,
                        "type": "error",
                        "error_message": str(e),
                    },
                )
                traceback.print_exc()
                return

            item.prompt_text = lyrics
            item.metadata["tags"] = " ".join(genre_tags)
            print("generated", item.metadata["tags"], lyrics)

            self.worker.notify_progress(
                item,
                {
                    "id": item.id,
                    "type": "lyrics",  # is 'lyrics' still accurate?
                    "title": title,
                    "text": item.prompt_text,
                    "tags": item.metadata["tags"],
                },
            )
        else:
            tags = item.metadata.get("tags") or ""
            moderation_result = moderate_user_inputs(self.openai_client, item.prompt_text or "", tags)
            print(moderation_result)
            if not isinstance(moderation_result, ModerationSuccess):
                self.worker.notify_finish(
                    item,
                    {
                        "id": item.id,
                        "type": "error",
                        "error_type": "moderation_failure",
                        "error_message": moderation_result.err_msg,
                    },
                )
                return

        f = modal.Function.lookup("stable-diffusion-xl-beta-3", "StableDiffusion.generate_image")
        fn_call = f.spawn(
            item.prompt_text or "",
            item_id,
            num_images=1,
            tags=item.metadata.get("tags", None),
        )

        try:
            audios, raw_arrays = self.worker.process_item(item)
            items = []
            for i, audio in enumerate(audios):
                new_id = f"{item_id}_{i}"
                ids.append(new_id)
                item.id = new_id
                # item.prompt_text = trimmed[i]
                items.append(item.json())
                self.worker._write_audio_only(
                    item,
                    audio,
                )

                self.worker._write_npz(
                    item,
                    raw_arrays[i],
                )

            image_url = fn_call.get(timeout=None)
            print("image located", image_url)

            self.worker.notify_finish(
                item,
                {
                    "id": item_id,
                    "model": "chirp-7B",
                    "n_audios": len(ids),
                    "ok": 1,
                    "gen_duration": time.time() - start_time,
                },
                queue_name="results:q",
            )

            video_time = time.time()
            f = modal.Function.lookup("dummy-v2-staging", "DummyV0Stub.write_video")
            # Start the video jobs.
            for item in items:
                f.spawn(item, image_url.split("/")[-1], is_square)

            print("Spawning videos", time.time() - video_time)
            return item_id

        except Exception as e:
            import traceback

            finish_time = time.time()
            self.worker.notify_finish(
                item,
                {
                    "id": item_id,
                    "model": "chirp-7B",
                    "n_audios": len(ids),
                    "ok": 0,
                    "gen_duration": finish_time - start_time,
                },
                queue_name="results:q",
            )

            traceback.print_exc()

            raise e

        return item.id


@stub.local_entrypoint()
def main():
    uid = str(uuid4())
    inputs = [
        json.dumps(
            dict(
                id=uid,
                prompt_text="""Hello world""",
                metadata={"tags": "pop"},
            )
        )
    ]
    model = ChirpV2Stub()
    print(uid)
    for input in inputs:
        model.generate.remote(
            input,
        )
