"""Studio session bounce on modal.

Studio session JSON -> WAV audio
"""

import asyncio
import json
import logging
import modal
import time
import traceback
import tempfile
import subprocess
import os
from concurrent.futures import ThreadPoolExecutor
from suno_utils.worker.schema import QueueItem
from suno_utils.worker.modal_base import get_modal_base_image
from suno_utils.audio import Audio
from suno_utils.worker.loader import S3Loader, retry_s3_download
from suno_utils.worker.settings import s3_client
from suno_utils.worker.modal_model_configs import VAEVersion

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

DEPLOYMENT_TYPE = "dev"  # dev, prod

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

UPLOADS_S3_BUCKET = "suno-data-uploads"
UPLOADS_S3_PREFIX = "studio/uploads"

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()
    .env(
        {
            "NVM_DIR": "/root/.nvm",
            "PATH": "$PATH:/root/.nvm/versions/node/v23.11.0/bin",
        }
    )
    .run_commands(
        [
            "curl -o- https://raw.githubusercontent.com/nvm-sh/nvm/v0.40.2/install.sh | bash",
            "bash -c 'source $NVM_DIR/nvm.sh && nvm install 23.11.0 && nvm use 23.11.0 && npm install -g yarn'",
        ]
    )
    .add_local_dir(
        "../dsp-engine", "/dsp-engine", copy=True, ignore=["**/node_modules", "**/build", "**/dist"]
    )
    .run_commands(
        ["cd /dsp-engine/buildtool && yarn install --frozen-lockfile && node ./index.js pull"],
        secrets=[aws_secret],
    )
    .add_local_python_source("suno_utils", copy=False)
)

logger = logging.getLogger(__name__)

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


@app.cls(
    secrets=SECRETS,
    timeout=400,
    scaledown_window=480,
    memory=15000,
    cpu=4,
    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=4)
class StudioBounceStub(S3Loader):
    def __init__(self):
        super().__init__()
        self.modal_f_encode_audio = modal.Cls.from_name(
            f"cycle-{DEPLOYMENT_TYPE}",
            "CycleStub",
        )().encode_audio
        self.modal_f_convert_to_wav = modal.Cls.from_name(
            f"cycle-{DEPLOYMENT_TYPE}",
            "WavCycleStub",
        )().convert_to_wav

    def _get_s3_ids_from_studio_state(self, studio_state_json):
        return set(
            filter(
                None,
                (
                    clip.get("clipId")
                    for track in studio_state_json.get("tracks", [])
                    for clip in track.get("clips", [])
                ),
            )
        )

    def _ensure_opus(self, s3_id):
        try:
            s3_client.head_object(Bucket=UPLOADS_S3_BUCKET, Key=f"{UPLOADS_S3_PREFIX}/{s3_id}.opus")
            return
        except (s3_client.exceptions.NoSuchKey, s3_client.exceptions.ClientError):
            print(f"Opus file for {s3_id} not found in S3, will try to convert")

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

    @modal.method()
    def bounce_callback(
        self,
        queue_item: str,
    ) -> None:
        local_queue_item = QueueItem(**json.loads(queue_item))

        with (
            tempfile.NamedTemporaryFile(delete=False, suffix=".json", mode="w") as studio_state_file,
            tempfile.NamedTemporaryFile(delete=False, suffix=".wav", mode="w") as output_file,
        ):
            output_file.close()
            try:
                if studio_state_json := local_queue_item.metadata.get("studio_state"):
                    json.dump(studio_state_json, studio_state_file)
                    studio_state_file.close()
                else:
                    studio_state_file.close()
                    with open(studio_state_file.name, "wb") as f:
                        retry_s3_download(
                            UPLOADS_S3_BUCKET, f"{UPLOADS_S3_PREFIX}/{local_queue_item.id}.json", f
                        )
                    with open(studio_state_file.name, "r") as f:
                        studio_state_json = json.load(f)
                with ThreadPoolExecutor(max_workers=4) as executor:
                    futures = [
                        executor.submit(
                            self._ensure_opus,
                            s3_id,
                        )
                        for s3_id in self._get_s3_ids_from_studio_state(studio_state_json)
                    ]
                    for future in futures:
                        future.result()
                subprocess.run(
                    [
                        "node",
                        "/dsp-engine/bounce/dist/index.js",
                        studio_state_file.name,
                        "--start",
                        str(float(local_queue_item.metadata.get("start_beats", "0"))),
                        "--end",
                        str(float(local_queue_item.metadata.get("end_beats", "100"))),
                        "--output",
                        output_file.name,
                    ],
                    check=True,
                )
                audio = Audio.from_file(output_file.name, n_channels=2)

                asyncio.run(
                    self._write_audio_only_async(
                        item=local_queue_item,
                        audio=audio,
                        s3_bucket=UPLOADS_S3_BUCKET,
                        s3_folder=UPLOADS_S3_PREFIX,
                    )
                )

                # TODO: this operation should be blocking cause untils this is finished other operations can't be done
                self.modal_f_encode_audio.spawn(
                    audio=f"s3://{UPLOADS_S3_BUCKET}/{UPLOADS_S3_PREFIX}/{local_queue_item.id}.opus",
                    s3_npz_id=local_queue_item.id,
                    encode_vae_version=VAEVersion.V_VAE_25_TUNED_2.value,
                )

                local_queue_item.notify_progress(
                    {
                        "id": local_queue_item.id,
                        "success": True,
                        "metadata": {
                            "duration": audio.duration_s,
                        },
                        "is_remix": local_queue_item.metadata.get("is_remix", True),
                        "parent_relationship_type": local_queue_item.metadata.get(
                            "parent_relationship_type", ""
                        ),
                        "is_user_direct_parent_owner": local_queue_item.metadata.get(
                            "is_user_direct_parent_owner", None
                        ),
                    },
                )
                print(f"{local_queue_item.id}: Notified progress for bounce.")
            except Exception as e:
                print(f"Error bouncing for {local_queue_item.id}: {e}")
                traceback.print_exc()
                local_queue_item.notify_progress(
                    {"id": local_queue_item.id, "error": str(e)},
                )
            finally:
                if os.path.exists(studio_state_file.name):
                    os.remove(studio_state_file.name)
                if os.path.exists(output_file.name):
                    os.remove(output_file.name)


@app.local_entrypoint()
def main():
    # tests
    model = StudioBounceStub()
    for gen_id in ["71e291bf-9d2f-43c3-a787-172a7ac8f897"]:
        for i in range(2):
            start = time.time()
            model.bounce_callback.remote(
                f'{{"id": "{gen_id}", "metadata": {{"start_beats": 0, "end_beats": 100}}}}'
            )
            print(f"Time taken for bounce: {time.time() - start}")

    print("Done")
