import time
from beat_processor import BeatProcessor
import subprocess
import tempfile
import boto3
import concurrent.futures
from lib.redis_client import RedisClient, getenv
from audiocraft.models import MusicGen
from audiocraft.data.audio import audio_write
import uuid
import torchaudio
from lib.audio_utils import download_audio, upload_audio, resample_audio, CODEC_SPEEDS
from lib.progress import (
    OpProgressSet,
    OpProgress,
    OpProgressMinSet,
    OpProgressFfmpeg,
)


class RetryableError(Exception):
    pass


class MusicGenRedisClient(RedisClient):
    def __init__(self, model, s3_bucket, *args):
        super().__init__(*args)
        self.model = MusicGen.get_pretrained(model)
        self.s3 = boto3.client("s3")
        self.s3_bucket = s3_bucket
        self.beat_processor = BeatProcessor()

    @classmethod
    def from_env(cls):
        return super().from_env(getenv("MUSICGEN_MODEL"), getenv("MUSICGEN_S3_BUCKET"))

    def warmup(self):
        self.model.set_generation_params(duration=1)
        self.model.generate_unconditional(1)

    def _get_audio_from_s3(self, key):
        with tempfile.TemporaryDirectory() as temp_dir:
            return torchaudio.load(
                download_audio(
                    self.s3,
                    self.s3_bucket,
                    key,
                    temp_dir,
                    # TODO resample?
                )
            )

    def _grid_audio(self, wav, sr, bpm, op_prog):
        with tempfile.TemporaryDirectory() as temp_dir:
            wav = wav.cpu()
            output_path = audio_write(
                f"{temp_dir}/output",
                wav,
                sr,
                strategy="loudness",
                loudness_compressor=True,
                format="wav",
            )
            beats = self.beat_processor.detect_beats(output_path)
            op_prog(1)
            first_bar_time = next((b[0] for b in beats if b[1] == 1), None)
            if first_bar_time is None:
                raise RetryableError("Could not detect first bar start")
            adjusted_beats = [
                b[0] - first_bar_time for b in beats if b[0] > first_bar_time + 1.0e-6
            ]
            if len(adjusted_beats) < 2:
                raise RetryableError("Not enough beats detected")
            cut_wav = wav[:, int(first_bar_time * sr) :]
            average_dt = sum(
                b - a for a, b in zip([0] + adjusted_beats[:-1], adjusted_beats)
            ) / len(adjusted_beats)
            duration_at_avg_bpm = average_dt * len(adjusted_beats)
            duration_at_target_bpm = len(adjusted_beats) * 60 / bpm
            min_duration_change = 1.0e6
            min_duration_factor = None
            for i in [0.25, 0.5, 1.0, 2.0, 4.0]:
                duration_change = abs(
                    1 - duration_at_target_bpm * i / (duration_at_avg_bpm + 1.0e-6)
                )
                if duration_change < min_duration_change:
                    min_duration_change = duration_change
                    min_duration_factor = i
            bpm /= min_duration_factor
            timemap = []
            resample_sr = 48000
            score = 0
            last_orig_time = 0

            for i, orig_time in enumerate(adjusted_beats):
                orig_dt = orig_time - last_orig_time
                last_orig_time = orig_time
                percent_change_this_beat = abs(1 - orig_dt / (60 / bpm))
                score = max(score, percent_change_this_beat)
                stretch_to_time = (i + 1) * 60 / bpm
                timemap.append(
                    (int(orig_time * resample_sr), int(stretch_to_time * resample_sr))
                )

            with open(f"{temp_dir}/timemap.txt", "w") as f:
                for a, b in timemap:
                    f.write(f"{a} {b}\n")
            ts_in_path = audio_write(
                f"{temp_dir}/ts_in",
                cut_wav,
                sr,
                strategy="loudness",
                loudness_compressor=True,
                format="wav",
            )
            if sr != resample_sr:
                output_path = f"{temp_dir}/resampled.wav"
                resample_audio(ts_in_path, output_path, resample_sr)
                ts_in_path = output_path
            op_prog(2)
            subprocess.run(
                [
                    "rubberband",
                    "--fine",
                    "--centre-focus",
                    "--duration",
                    str(float(len(adjusted_beats) * 60 / bpm)),
                    "--timemap",
                    f"{temp_dir}/timemap.txt",
                    ts_in_path,
                    f"{temp_dir}/ts_out.wav",
                ],
                check=True,
                stdout=subprocess.DEVNULL,
                stderr=subprocess.DEVNULL,
            )
            op_prog(3)
            return *torchaudio.load(f"{temp_dir}/ts_out.wav"), score

    def _put_audio_to_s3(self, wav, sr, use_compressor, op_prog=None):
        with tempfile.TemporaryDirectory() as temp_dir:
            output_path = audio_write(
                f"{temp_dir}/output",
                wav.cpu(),
                sr,
                strategy="loudness",
                loudness_compressor=use_compressor,
                format="wav",
            )
            codec = "ogg"
            s3_key = f"{uuid.uuid4()}.{codec}"
            upload_audio(self.s3, self.s3_bucket, s3_key, output_path, codec, op_prog)
            return s3_key

    def _grid_audio_and_upload(self, wav, sr, bpm, op_prog):
        wav, sr, score = self._grid_audio(wav, sr, bpm, op_prog)
        s3_key = self._put_audio_to_s3(wav, sr, use_compressor=False)
        return s3_key, score

    def close(self):
        self.beat_processor.close()

    def handle_request(self, request_data, recorder, notify_progress, times_out_at):
        t0 = time.time()
        print("request: ", request_data)

        want_duration = request_data.get("duration", 10)

        op_set = OpProgressSet(notify_progress, debounce=0.5)
        op_generate = OpProgress(100, 50.0)
        op_set.add(op_generate)
        op_encode = OpProgressFfmpeg(want_duration, CODEC_SPEEDS.get("ogg", 22.0))
        op_set.add(op_encode)

        audio_prompt_key = request_data.get("audioPrompt", None)
        if audio_prompt_key is not None:
            audio_prompt, audio_prompt_sr = self._get_audio_from_s3(audio_prompt_key)
        else:
            audio_prompt, audio_prompt_sr = None, None

        self.model.set_generation_params(duration=want_duration)

        def progress_callback(a, b):
            op_generate.total = b
            op_generate(a)

        self.model.set_custom_progress_callback(progress_callback)
        prompt = request_data.get("textPrompt", None)
        target_bpm = request_data.get("bpm", None)
        batch_size = 1 if target_bpm is None else 4
        if target_bpm is not None:
            op_set_grid = OpProgressMinSet()
            op_set.add(op_set_grid)
            op_grids = []
            for _ in range(batch_size):
                op_grids.append(OpProgress(3, 3.0 / want_duration))
                op_set_grid.add(op_grids[-1])
        if prompt is None and audio_prompt is None:
            wavs = self.model.generate_unconditional(batch_size, progress=True)
        elif audio_prompt is None:
            wavs = self.model.generate([prompt] * batch_size, progress=True)
        else:
            wavs = self.model.generate_with_chroma(
                [prompt if prompt is not None else ""] * batch_size,
                audio_prompt[None].expand(batch_size, -1, -1),
                audio_prompt_sr,
                progress=True,
            )

        if target_bpm is not None:
            results = []
            with concurrent.futures.ThreadPoolExecutor(max_workers=4) as executor:
                for fut in [
                    executor.submit(
                        self._grid_audio_and_upload,
                        wavs[i],
                        self.model.sample_rate,
                        float(target_bpm),
                        op,
                    )
                    for i, op in enumerate(op_grids)
                ]:
                    try:
                        results.append(fut.result())
                    except RetryableError:
                        continue
            if len(results) == 0:
                raise RuntimeError("Sample generation failed")
            results = list(sorted(results, key=lambda x: x[1]))
            print(f"({time.time() - t0:.2f}s) -> ", results)
            return {
                "status": "success",
                "s3Key": results[0][0],
                "results": [{"s3Key": r[0], "score": r[1]} for r in results],
            }
        else:
            s3_key = self._put_audio_to_s3(
                wavs[0], self.model.sample_rate, use_compressor=True, op_prog=op_encode
            )
            print(f"({time.time() - t0:.2f}s) -> ", s3_key)
            return {"status": "success", "s3Key": s3_key}


if __name__ == "__main__":
    client = MusicGenRedisClient.from_env()

    # client.handle_request(
    #     {
    #         "textPrompt": "A happy tune",
    #         "duration": 5,
    #         "bpm": 120,
    #     },
    #     None,
    #     print,
    #     time.time() + 60,
    # )

    client.run()
    client.close()
