"""Audio to MIDI transcription application on modal.

Audio --> BasicPitch --> MIDI
"""

import modal
import tempfile
import time
import os
from uuid import uuid4
import pathlib
from suno_utils.audio import Audio
from suno_utils.worker.settings import s3_client

base_image = (
    modal.Image.from_registry(
        "tensorflow/tensorflow:2.15.0-gpu",
        add_python="3.10",
    )
    .apt_install("curl", "ffmpeg", "sox", "unzip", "libsox-fmt-mp3", "git", "clang")
    .pip_install_from_pyproject(
        str(pathlib.Path(__file__).parent.parent.parent / "pyproject.toml"),
    )
    .pip_install(
        "basic-pitch[tf]",
        "boto3",
        "midiutil",
    )
)

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

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


@app.cls(
    cpu=4,
    gpu="A10G",
    secrets=[aws_secret],
    timeout=240,
    memory=8000,
    allow_concurrent_inputs=4,
    min_containers=1,
)
class MidiTranscriptionStub:
    def __init__(self):
        from basic_pitch import ICASSP_2022_MODEL_PATH
        from basic_pitch.inference import Model

        print("Loading BasicPitch model...")
        self.model = Model(ICASSP_2022_MODEL_PATH)
        print("Model loaded!")

    @modal.method()
    def transcribe(self, s3_path: str, midi_id: 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
            local_audio_path = os.path.join(td, "input.wav")
            audio = Audio.from_s3(s3_path, n_channels=1)
            audio.to_wav(local_audio_path)

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

            # Use basic_pitch's predict_and_save to save MIDI and other outputs
            from basic_pitch.inference import predict_and_save

            # Save MIDI file using the model output
            predict_and_save(
                [local_audio_path],  # List of input audio paths
                td,  # Output directory
                save_midi=True,  # Save MIDI file
                sonify_midi=False,  # Don't save audio rendering of MIDI
                save_model_outputs=False,  # Don't save raw model outputs
                save_notes=False,  # Don't save note events as CSV
                model_or_model_path=self.model,
            )

            # 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 transcription with multiple instrument samples
    test_samples = {
        "stone_vocals": {
            "audio_path": "s3://suno-data-uploads/studio/uploads/45f04b32-8df3-4428-ab7a-47e561ff3847.mp3",
            "instrument": 53,
        },
        "bass": {
            "audio_path": "s3://suno-data-uploads/studio/uploads/ec626427-771d-4b2e-8cd0-7bd2a7f1cc19.mp3",
            "instrument": 32,
        },
        "piano": {
            "audio_path": "s3://suno-data-uploads/studio/uploads/fefebefd-a98b-4167-83ae-0722f8891824.mp3",
            "instrument": 1,
        },
        "electric_guitar": {
            "audio_path": "s3://suno-data-uploads/studio/uploads/bac51035-7732-47b7-b180-7a759edb6bd0.mp3",
            "instrument": 25,
        },
        "flute": {
            "audio_path": "s3://suno-data-uploads/studio/uploads/16dcfe3e-7760-40f5-9d6c-e598010383c9.mp3",
            "instrument": 73,
        },
        "violin": {
            "audio_path": "s3://suno-data-uploads/studio/uploads/6c957b14-66f1-4219-8b83-6001860055d7.mp3",
            "instrument": 40,
        },
    }

    for name, sample in test_samples.items():
        print(f"Transcribing {name}...")
        id = sample["audio_path"].split("/")[-1].split(".")[0]
        output_path = model.transcribe.remote(sample["audio_path"], midi_id=id)
        print(f"MIDI file for {name} saved to: {output_path}")
