"""Lyrics alignment with hoot + HMM on modal.

Clip -> aligned lyrics
"""

import json
import logging
import pathlib
import modal
import time
from typing import List, Dict

from suno_utils.worker.schema import QueueItem
from suno_utils.utils.clip import SunoClip
from suno_utils.tasks.lyrics_alignment import AlignmentConfig


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

DEPLOYMENT_TYPE = "dev"  # dev, prod

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


aws_secret = modal.Secret.from_name("studio-aws")
SECRETS = [
    aws_secret,
    modal.Secret.from_name("openai-secret"),
    modal.Secret.from_name("api-callback-token"),
]

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(
        "boto3",
        "tokenizers",
        "encodec",
        "ctc_segmentation",
        "psutil",
        "pydantic",
        "nnAudio",
        "funcy",
    )
    .pip_install_from_pyproject(
        str(pathlib.Path(__file__).parent.parent.parent / "pyproject.toml"),
    )
    .pip_install("torch==2.4.0")  # this is cause flash-attn can't work with 2.5 yet
    .run_commands(
        "FLASH_ATTENTION_SKIP_CUDA_BUILD=TRUE pip install flash-attn==2.6.3 --no-build-isolation",
    )
    .pip_install(
        "torch==2.5.1",
        "torchaudio==2.5.1",
    )
    .pip_install(
        "pyphonetics==0.5.3",
        "rapidfuzz==3.11.0",
        "dtw-python==1.5.3",
    )
    .pip_install_private_repos(
        "github.com/suno-ai/hoot.git@96a3e65",
        git_user="mcamac",
        secrets=[modal.Secret.from_name("victor-modal-github-token")],
    )
)

logger = logging.getLogger(__name__)

APP_NAME = f"lyrics-alignment-{DEPLOYMENT_TYPE}"
app = modal.App(APP_NAME, image=image)


@app.cls(
    cpu=2,
    gpu=modal.gpu.H100(count=1),
    memory=15000,
    secrets=SECRETS,
    timeout=400,
    container_idle_timeout=480,
    retries=modal.Retries(
        max_retries=1,
        backoff_coefficient=2.0,
        initial_delay=5.0,
    ),
    concurrency_limit=1000,
    region="us-east",
    keep_warm=1 if DEPLOYMENT_TYPE == "dev" else 75,
    _experimental_buffer_containers=0 if DEPLOYMENT_TYPE == "dev" else 1,
)
class LyricsAlignmentStub:
    def __init__(self):
        from suno_utils.tasks.lyrics_alignment import init_lyrics_alignment

        init_lyrics_alignment()

    @modal.method()
    def align_clip(
        self,
        clip_id: str,
        debug: bool = False,
        use_phonemes: bool = True,
    ) -> List[Dict[str, float]]:
        from suno_utils.tasks.lyrics_alignment import word_timing

        clip = SunoClip(clip_id)
        return word_timing(clip, config=AlignmentConfig(debug=debug, use_phonemes=use_phonemes))

    @modal.method()
    def align_clip_callback(
        self,
        queue_item: str,
        debug: bool = False,
    ) -> None:
        from suno_utils.tasks.lyrics_alignment import word_timing

        queue_item = QueueItem(**json.loads(queue_item))
        try:
            clip = SunoClip(queue_item.id)
            alignment = word_timing(clip, config=AlignmentConfig(debug=debug))
        except Exception as e:
            queue_item.notify_progress(
                {"id": queue_item.id, "error": str(e)},
            )
            raise e
        queue_item.notify_progress(
            {"id": queue_item.id, "alignment": alignment},
        )


@app.local_entrypoint()
def main():
    from pprint import pprint

    # tests
    model = LyricsAlignmentStub()
    for _ in range(2):
        start = time.time()
        x = model.align_clip.remote("43055f66-c863-461b-8c4a-b7824090d033")
        end = time.time()
        print(f"Time taken: {end - start}")
        pprint(x)

    print("Done")
