import time
import pathlib
import modal
import torch

from suno_utils.audio import Audio
from suno_utils.worker.loader import S3Loader
from suno_utils.worker.modal_base import MODAL_MOUNTS
from suno_utils.worker.utils import print_gpu_memory_usage
from suno_utils.gpt import chirp_v2_5 as chirp_v3
from suno_utils.models.ditto.ditto import Ditto

############## CHANGE THESE ##############

DEPLOYMENT_TYPE = "dev"  # dev, prod

##########################################


ENCODER_CONCURRENCY_LIMITS = {
    "dev": 50,
    "prod": 250,
}
KEEP_WARM = {
    "dev": 1,
    "prod": 1,
}


# set number of cpus.
ENCODER_MAX_INPUT = 10  # encoder can handle much more traffic, but let's be conservative...
VERBOSE_MESSAGE = DEPLOYMENT_TYPE == "dev"
MOUNT_PATH = "/suno/models"


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"),
]

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.2.0.+cu118",
        "torchaudio==2.2.0+cu118",
        index_url="https://download.pytorch.org/whl/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_private_repos(
        "github.com/suno-ai/hoot.git@1ad12a3",
        git_user="mcamac",
        secrets=[modal.Secret.from_name("victor-modal-github-token")],
    )
    .pip_install(
        "boto3",
        "tokenizers",
        "encodec",
        "ctc_segmentation",
        "psutil",
        "pydantic",
        "nnAudio",
        "biopython>=1.81",  # TODO: don't love this depdendency, for hoot
        "turbopuffer",
    )
    .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==2.5.2 --no-build-isolation",
    )
)

DITTO_S3_PATH = "s3://suno-data/victor/checkpoints/ditto/step_370k.pt"


class DittoWorker(S3Loader):
    def __init__(self):
        S3Loader.__init__(self)
        print("Start loading models")
        self.ditto_path = chirp_v3._get_model_if_needed(DITTO_S3_PATH, cache_dir=MOUNT_PATH)
        self.music_encoder_path = chirp_v3._get_model_if_needed(
            "s3://suno-data/victor/checkpoints/ditto/musicfm_concat_epoch=51.pt",
            cache_dir=MOUNT_PATH,
        )
        print(f"Downloaded model to {self.ditto_path}")
        self.ditto = Ditto(
            music_encoder_name="musicfm_concat",
            latent_dim=128,
            model_path=self.ditto_path,
            music_encoder_path=self.music_encoder_path,
            is_flash=False,
        )

        self.ditto = self.ditto.eval().cuda()
        print("Finish loading models")

    @staticmethod
    def download_models(dir_path=MOUNT_PATH):
        print("Start downloading models")
        dl_path = chirp_v3._get_model_if_needed(DITTO_S3_PATH, cache_dir=dir_path)
        music_encoder_path = chirp_v3._get_model_if_needed(
            "s3://suno-data/victor/checkpoints/ditto/musicfm_concat_epoch=51.pt",
            cache_dir=dir_path,
        )
        print(f"Downloaded model to {dl_path}")
        print(f"Downloaded model to {music_encoder_path}")
        print("Finish downloading models")


def download_model_wrapper_e():
    # this print is necessary to have modal rerun this when MODEL changes
    # Modal tracks referenced global variables
    # Change the name of the function to force a rerun
    print("Downloading model", DITTO_S3_PATH)
    DittoWorker.download_models()


image = base_image.run_function(download_model_wrapper_e, secrets=SECRETS)
STUB_NAME = f"ditto-{DEPLOYMENT_TYPE}"
stub = modal.Stub(STUB_NAME, image=image)


@stub.cls(
    gpu=modal.gpu.T4(count=1),
    secrets=SECRETS,
    timeout=100,
    container_idle_timeout=400,
    mounts=MODAL_MOUNTS,
    retries=modal.Retries(
        max_retries=1,
        backoff_coefficient=2.0,
        initial_delay=5.0,
    ),
    keep_warm=KEEP_WARM[DEPLOYMENT_TYPE],
    concurrency_limit=ENCODER_CONCURRENCY_LIMITS[DEPLOYMENT_TYPE],
    allow_concurrent_inputs=ENCODER_MAX_INPUT,
)
class DittoWorkerStub:
    def __enter__(self):
        import torch

        num_gpus = torch.cuda.device_count()
        print(f"Found {num_gpus} GPUs.")

        self.worker = DittoWorker()

    @modal.method()
    def encode_audio(
        self, id: str = None, s3_url: str = None, start: float = 0, dur: float = 30, callback_url=None
    ) -> None:
        if VERBOSE_MESSAGE:
            print_gpu_memory_usage(self.__class__.__name__)
            torch.cuda.reset_max_memory_allocated()

        if dur > 30:
            raise ValueError("Duration must be less than 30 seconds")

        assert id or s3_url, "Either id or s3_url must be provided"

        if id is not None:
            s3_url = f"s3://suno-data-uploads/studio/uploads/{id}.mp3"

        audio = Audio.from_s3(s3_url, n_channels=1, sample_rate=24000)
        audio = audio.get_segment(from_s=start, to_s=start + dur)
        wav = torch.tensor(audio.array_float).unsqueeze(0).cuda()
        emb = self.worker.ditto.music_to_latent(wav)[0].detach().cpu().numpy()

        if id and callback_url:
            import requests

            requests.post(
                callback_url,
                json={"id": id, "vector": emb.tolist()},
                headers={
                    "Authentication": "Bearer 562a512f-0dce-4acd-bf23-ad9bb8d8a084",
                },
            )

        return emb

    @modal.method()
    def encode_audio_and_upload(
        self, id: str = None, s3_url: str = None, start: float = 0, dur: float = 30
    ) -> None:
        """Use this method to backfill existing clips with embeddings.
        Should be used for one-off tasks (e.g. from your notebook) only.
        Please note that the DEPLOYMENT_TYPE should be set properly.
        For exmaple, you can't upload staging clips to prod namespace"""
        if VERBOSE_MESSAGE:
            print_gpu_memory_usage(self.__class__.__name__)
            torch.cuda.reset_max_memory_allocated()

        if dur > 30:
            raise ValueError("Duration must be less than 30 seconds")

        assert id or s3_url, "Either id or s3_url must be provided"

        if id is not None:
            s3_url = f"s3://suno-data-uploads/studio/uploads/{id}.mp3"

        audio = Audio.from_s3(s3_url, n_channels=1, sample_rate=24000)
        audio = audio.get_segment(from_s=start, to_s=start + dur)
        wav = torch.tensor(audio.array_float).unsqueeze(0).cuda()
        emb = self.worker.ditto.music_to_latent(wav)[0].detach().cpu().numpy()

        import turbopuffer as tpuf

        tpuf.api_key = "C6FVrjLLHJ65WwP8WbDJDut8NSHXT1rj"
        ns = tpuf.Namespace(f"energy-{DEPLOYMENT_TYPE}")
        ns.upsert(
            ids=[id],
            vectors=[emb.tolist()],
            distance_metric="cosine_distance",
        )

        return emb

    @modal.method()
    def encode_text(self, text: str) -> None:
        if VERBOSE_MESSAGE:
            print_gpu_memory_usage(self.__class__.__name__)
            torch.cuda.reset_max_memory_allocated()

        te = self.worker.ditto.text_to_latent("[CLS]" + text)[0].detach().cpu().numpy()
        return te


@stub.local_entrypoint()
def main():
    ditto_worker = DittoWorkerStub()
    print(ditto_worker.encode_audio.remote(id="4a77dea7-19f3-46d2-8b0a-b2b7e9ea9a05"))
    print(
        ditto_worker.encode_audio.remote(
            s3_url="s3://suno-data-uploads/studio/uploads/4a77dea7-19f3-46d2-8b0a-b2b7e9ea9a05.mp3"
        )
    )
    print(ditto_worker.encode_text.remote("jazz"))

    for i in range(30):
        ditto_worker.encode_audio.spawn("4a77dea7-19f3-46d2-8b0a-b2b7e9ea9a05")

    time.sleep(100)
