import uuid
import demucs.api
import tempfile
import boto3
from concurrent.futures import ThreadPoolExecutor
from lib.redis_client import RedisClient, getenv
import torch
import torchaudio as ta
from lib.audio_utils import download_audio, upload_audio, CODEC_SPEEDS
from lib.progress import (
    OpProgress,
    OpProgressSet,
    OpProgressMinSet,
    OpProgressFfmpeg,
)

SAMPLE_RATE = 44100  # only rate that demucs will ever support


class DemucsRedisClient(RedisClient):
    def __init__(self, model, s3_bucket, *args):
        super().__init__(*args)
        self.separator = demucs.api.Separator(model=model)
        self.s3 = boto3.client("s3")
        self.s3_bucket = s3_bucket

    @classmethod
    def from_env(cls):
        return super().from_env(getenv("DEMUCS_MODEL"), getenv("DEMUCS_S3_BUCKET"))

    def warmup(self):
        self.separator.separate_tensor(torch.randn(2, SAMPLE_RATE, dtype=torch.float32))

    def _separate_s3(self, s3_key, notify_progress, codec, duration):
        duration = max(0.1, float(duration)) if duration is not None else 10
        op_set = OpProgressSet(notify_progress, debounce=0.5)
        decode_op = OpProgressFfmpeg(duration, 300.0)
        op_set.add(decode_op)
        demucs_op = OpProgress(duration, 22.0)
        op_set.add(demucs_op)
        encode_op_set = OpProgressMinSet()
        op_set.add(encode_op_set)
        encode_ops = {}
        for stem in self.separator._model.sources:
            encode_op = OpProgressFfmpeg(duration, CODEC_SPEEDS.get(codec, 22.0))
            encode_op_set.add(encode_op)
            encode_ops[stem] = encode_op

        with tempfile.TemporaryDirectory() as temp_dir, ThreadPoolExecutor(
            max_workers=4
        ) as executor:
            input_path = download_audio(
                self.s3,
                self.s3_bucket,
                s3_key,
                temp_dir,
                progress_handler=decode_op,
                sample_rate=SAMPLE_RATE,
            )
            decode_op.finish()

            def sep_notify(d):
                if d["state"] != "start":
                    return
                true_duration = max(0.1, d["audio_length"] / SAMPLE_RATE)
                demucs_op.total = true_duration
                encode_op_set.total = true_duration
                demucs_op(d["segment_offset"] / SAMPLE_RATE)

            self.separator.update_parameter(callback=sep_notify)

            audio_tensor, loaded_sample_rate = ta.load(input_path)
            assert loaded_sample_rate == SAMPLE_RATE
            was_mono = False
            # demucs crashes on mono audio - convert to stereo
            if audio_tensor.shape[0] == 1:
                was_mono = True
                audio_tensor = audio_tensor.repeat(2, 1).clone()
            _, separated = self.separator.separate_tensor(
                audio_tensor, self.separator.samplerate
            )

            demucs_op.finish()

            futures = []
            res = {}
            for stem, wav in separated.items():
                if was_mono:
                    wav = wav.mean(0, keepdim=True)
                output_path = f"{temp_dir}/{stem}.wav"
                demucs.api.save_audio(
                    wav, output_path, samplerate=self.separator.samplerate
                )
                s3_key = f"{uuid.uuid4()}.{codec}"
                res[stem] = s3_key
                futures.append(
                    executor.submit(
                        upload_audio,
                        self.s3,
                        self.s3_bucket,
                        s3_key,
                        output_path,
                        codec,
                        progress_handler=encode_ops[stem],
                    )
                )
            for future in futures:
                future.result()
            notify_progress(100)
            return res

    def handle_request(self, request_data, recorder, notify_progress, times_out_at):
        print("request: ", request_data)
        stem_dict = self._separate_s3(
            request_data["s3Key"],
            notify_progress,
            request_data.get("codec", "ogg"),
            request_data.get("duration", None),
        )
        return {"status": "success", "stems": stem_dict}


if __name__ == "__main__":
    client = DemucsRedisClient.from_env()
    client.run()
