"""Audio feature extraction on modal.

Clip -> features
"""

import json
import logging
import modal
import time
import traceback
from typing import Optional, Dict, Any
from suno_utils.worker.schema import QueueItem
from suno_utils.utils.clip import SunoClip
from suno_utils.worker.utils import retry_decorator
from suno_utils.worker.modal_base import get_modal_base_image
from suno_utils.audio import Audio
import tempfile
import boto3
import torch
from modal.experimental import stop_fetching_inputs

############## 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 = (
    get_modal_base_image()
    .pip_install(
        "git+https://github.com/marc-suno/madmom.git@bf7d502#egg=madmom",
        "git+https://github.com/CPJKU/beat_this@117ff34",
        "cvxpy==1.6.5",
    )
    .add_local_python_source("suno_utils", copy=False)
)

logger = logging.getLogger(__name__)

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

UPLOADS_BUCKET = "suno-data-uploads"
UPLOADS_KEY_PREFIX = "studio/uploads"


@app.cls(
    secrets=SECRETS,
    gpu="T4",
    timeout=60,
    scaledown_window=480,
    memory=3000,
    cpu=2,
    retries=modal.Retries(
        max_retries=1,
        backoff_coefficient=2.0,
        initial_delay=5.0,
    ),
    max_containers=1000,
    region="us-east",
    buffer_containers=0 if DEPLOYMENT_TYPE == "dev" else 1,
    min_containers=1,
)
@modal.concurrent(max_inputs=2)
class AudioFeaturesStub:
    def __init__(self):
        from suno_utils.tasks.audio_features.beat_this_downbeat import BeatThisDownbeatExtractor
        from suno_utils.tasks.audio_features.key import KeyExtractor
        from suno_utils.tasks.audio_features.instrument import InstrumentExtractor

        print("Initializing AudioFeaturesStub")

        self.modal_f_convert_to_wav = modal.Cls.from_name(
            f"cycle-{DEPLOYMENT_TYPE}",
            "WavCycleStub",
        )().convert_to_wav
        self.downbeat_extractor = BeatThisDownbeatExtractor(
            device="cuda", model_path="s3://suno-data/m4burns/beat_this_beta.pt"
        )
        self.key_extractor = KeyExtractor()
        self.instrument_extractor = InstrumentExtractor()
        self.s3_client = boto3.client("s3")
        self.retry_s3_download = retry_decorator(3, wait_seconds=5)(self.s3_client.download_fileobj)

        torch.backends.cuda.cufft_plan_cache[0].max_size = 0

        # warmup
        warmup_audio = self._get_aligned_audio("40ec1fc7-0b0a-4f01-89de-6a6a6db65475")
        self.downbeat_extractor.extract(warmup_audio)

    def _get_aligned_audio(self, gen_id: str) -> Audio:
        # try grabbing the opus from s3
        key = f"{UPLOADS_KEY_PREFIX}/{gen_id}.opus"
        with tempfile.NamedTemporaryFile(suffix=".opus") as f:
            try:
                self.s3_client.download_fileobj(Bucket=UPLOADS_BUCKET, Key=key, Fileobj=f)
                f.flush()
                return Audio.from_file(f.name, n_channels=1, sample_rate=48000)
            except (self.s3_client.exceptions.NoSuchKey, self.s3_client.exceptions.ClientError):
                print(f"Opus not found for {gen_id}, trying to generate...")
            except Exception as e:
                print(f"Error getting opus for {gen_id}: {e}")
                raise e

            self.modal_f_convert_to_wav.remote(
                f'{{"id": "{gen_id}", "metadata": {{"convert_to_opus": true, "user_id": 0, "clip_user_id": 0}}}}'
            )

            try:
                f.truncate(0)
                self.retry_s3_download(Bucket=UPLOADS_BUCKET, Key=key, Fileobj=f)
                f.flush()
                return Audio.from_file(f.name, n_channels=1, sample_rate=48000)
            except Exception as e:
                print(f"Error getting opus for {gen_id}: {e}")
                raise e

    @modal.method()
    def extract_downbeats_callback(
        self,
        queue_item_str: str,
        return_result: bool = False,
    ) -> Optional[Dict[str, Any]]:
        queue_item = QueueItem(**json.loads(queue_item_str))
        print("extract_downbeats_callback", queue_item.model_dump_json())

        try:
            audio = self._get_aligned_audio(queue_item.id)
            result_dict = self.downbeat_extractor.extract(audio)
            queue_item.notify_progress(
                {
                    "id": queue_item.id,
                    "metadata": queue_item.metadata,
                    **result_dict,
                },
            )
        except torch.cuda.OutOfMemoryError:
            # container is likely to be in a bad state
            # stop fetching inputs and raise to trigger a restart
            print(f"Out of memory for {queue_item.id} - exiting")
            stop_fetching_inputs()
            raise
        except Exception as e:
            print(f"Error extracting downbeats for {queue_item.id}: {e}")
            traceback.print_exc()
            queue_item.notify_progress(
                {"id": queue_item.id, "error": str(e), "metadata": queue_item.metadata},
            )
            return None

        if return_result:
            return result_dict

    @modal.method()
    def extract_key_callback(
        self,
        queue_item_str: str,
        return_result: bool = False,
    ) -> Optional[str]:
        queue_item = QueueItem(**json.loads(queue_item_str))
        try:
            key = self.key_extractor.extract(SunoClip(queue_item.id))
        except Exception as e:
            print(f"Error extracting key for {queue_item.id}: {e}")
            traceback.print_exc()
            queue_item.notify_progress(
                {"id": queue_item.id, "error": str(e), "metadata": queue_item.metadata},
            )
            return

        queue_item.notify_progress(
            {"id": queue_item.id, "key": key, "metadata": queue_item.metadata},
        )

        if return_result:
            return key

    @modal.method()
    def extract_instruments_callback(
        self,
        queue_item_str: str,
        return_result: bool = False,
    ) -> Optional[Dict[str, Any]]:
        queue_item = QueueItem(**json.loads(queue_item_str))
        try:
            _, _, group_class, inst_class = self.instrument_extractor.extract(SunoClip(queue_item.id))
            result = {"groups": group_class, "instruments": inst_class}
        except Exception as e:
            print(f"Error extracting instruments for {queue_item.id}: {e}")
            traceback.print_exc()
            queue_item.notify_progress(
                {"id": queue_item.id, "error": str(e), "metadata": queue_item.metadata},
            )
            return

        queue_item.notify_progress(
            {"id": queue_item.id, "result": result, "metadata": queue_item.metadata},
        )

        if return_result:
            return result


@app.local_entrypoint()
def main():
    # tests
    model = AudioFeaturesStub()
    for gen_id in [
        "3fabec37-fe9e-489d-b107-f0517feb3cf1",
        "40ec1fc7-0b0a-4f01-89de-6a6a6db65475",
        "879c3b07-87a1-4926-bdea-97fa77cc2d65",
        "ba6e3678-8fa6-4d9f-87f0-72530ccf3245",
        "a16c1c28-aaf8-4cb2-8ba9-4855f8b7054c",
        "a67cda87-cc19-49a2-90bc-bcfa211c29aa",
    ]:
        start = time.time()
        q = model.extract_downbeats_callback.remote(
            f'{{"id": "{gen_id}", "metadata": {{}}}}', return_result=True
        )
        assert q is not None
        json.dumps(q)
        print(f"Time taken for downbeats: {time.time() - start}")

        start = time.time()
        q = model.extract_key_callback.remote(
            f'{{"id": "{gen_id}", "metadata": {{}}}}', return_result=True
        )
        assert q is not None
        json.dumps(q)
        print(f"Time taken for modulation: {time.time() - start}")

        start = time.time()
        q = model.extract_instruments_callback.remote(
            f'{{"id": "{gen_id}", "metadata": {{}}}}', return_result=True
        )
        assert q is not None
        json.dumps(q)
        print(f"Time taken for instruments: {time.time() - start}")

    print("Done")
