"""Audio to MIDI transcription application on modal.

Audio --> Model --> MIDI
"""

import modal
import tempfile
import time
import os
from uuid import uuid4
import pathlib
from typing import Optional

from suno_utils.audio import Audio
from suno_utils.worker.settings import s3_client
from suno_utils.tasks.midi_transcription import MidiTranscriber
from suno_utils.worker.modal_base import get_modal_base_image_with_flash_attention
from suno_utils.worker.modal_model_volume import (
    MODEL_STORE_VOLUME_DIR,
    model_store_volume,
)

base_image = (
    get_modal_base_image_with_flash_attention()
    .pip_install_from_pyproject(
        str(pathlib.Path(__file__).parent.parent.parent / "pyproject.toml"),
    )
    .pip_install(
        "tqdm",
        "transformers==4.46.1",
        "pretty_midi",
        "midi_player",
        "miditok==3.0.6",
        "tokenizers==0.20.2",
    )
    .add_local_python_source("suno_utils", copy=False)
)

APP_NAME = "midi-transcription-dev"
app = modal.App(APP_NAME, image=base_image)

aws_secret = modal.Secret.from_name("studio-aws")


@app.cls(
    gpu="A10G",
    secrets=[aws_secret],
    timeout=600,  # 10 minutes
    volumes={MODEL_STORE_VOLUME_DIR: model_store_volume},
    min_containers=0,
    max_containers=20,
    keep_warm=2,
    scaledown_window=60 * 3,
)
class MidiTranscriptionStub:
    def __init__(self):
        print("Loading MidiTranscriber model...")
        self.transcriber = MidiTranscriber("s3://suno-data/victor/checkpoints/midi_transcription_v0.pt")
        print("Model loaded!")

        # warmup with a sample audio
        print("Warming up model...")
        self.transcriber.transcribe(
            # chosen because it's a short audio
            Audio.from_s3(
                "s3://suno-data-uploads/studio/uploads/c64e0f95-b907-4bf1-aa74-e07c25f23ded.mp3"
            )
        )
        print("Warmup complete!")

    @modal.method()
    def transcribe(self, s3_path: str, midi_id: Optional[str] = None) -> str:
        """Takes an S3 path to an audio file and transcribes it to MIDI.
        Returns the S3 path to the generated MIDI file.
        """

        start_time = time.time()
        print(f"Processing audio from: {s3_path}")

        if not s3_path.startswith("s3://"):
            s3_path = f"s3://suno-data-uploads/studio/uploads/{s3_path}.mp3"
            print(f"Using default S3 path: {s3_path}")

        with tempfile.TemporaryDirectory() as td:
            # Download and process audio
            # The transcriber handles stereo to mono conversion.
            audio = Audio.from_s3(s3_path)

            # Generate unique ID for the output file
            output_id = midi_id if midi_id else str(uuid4())
            midi_path = os.path.join(td, f"{output_id}.mid")

            # Transcribe audio to MIDI
            print("Transcribing audio to MIDI...")
            midi = self.transcriber.transcribe(audio, use_tqdm=False)

            # Save MIDI to file
            midi.write(midi_path)
            print(f"MIDI file saved locally to {midi_path}")

            # Upload MIDI to S3
            output_s3_path = f"s3://suno-data-uploads/studio/uploads/{output_id}.mid"
            s3_client.upload_file(
                midi_path,
                "suno-data-uploads",
                f"studio/uploads/{output_id}.mid",
                ExtraArgs={
                    "ContentType": "audio/midi",
                },
            )

            print(
                f"Uploaded MIDI file to {output_s3_path}. Took {round(time.time() - start_time, 2)} seconds"
            )
            return output_s3_path


@app.local_entrypoint()
def main():
    model = MidiTranscriptionStub()
    test_sample_s3_path = (
        "s3://suno-data-uploads/studio/uploads/a5e2198a-f352-4abb-9a24-7f81b143ded3.mp3"  # stone
    )
    print(f"Transcribing {test_sample_s3_path}...")
    output_path = model.transcribe.remote(test_sample_s3_path)
    print(f"MIDI file for test sample saved to: {output_path}")
