import os
import uuid
import torch
import torchaudio
import numpy as np
import matplotlib.pyplot as plt

from suno_utils.audio import Audio
from suno_utils.utils.text import (
    write_jsonl,
    read_jsonl,
    write_json,
    read_json,
    normalize_whitespace,
)

if __name__ == "__main__":

    metas_filepath = "/app/suno/data/chirp_v4_ft_sm/multi/metas_tr_quality.jsonl"
    # metas_filepath = "metadata/deezer_metas_audio_quality.jsonl"
    metas = read_jsonl(metas_filepath)

    subset = os.path.basename(metas_filepath).replace(".jsonl", "")
    os.makedirs("outputs/audio", exist_ok=True)
    os.makedirs(f"outputs/audio/{subset}", exist_ok=True)

    save_audio = False

    results = {}
    scores_list = []
    for meta in metas:
        if "audio_quality" in meta:
            quant = float(meta["audio_quality"]["quantification"])
            pref = float(meta["audio_quality"]["preference"])
            score = float(meta["audio_quality"]["score"])
            # score = ((pref * 2) - 1) * (quant + 1)
            # sign = 1 if pref > 0.5 else -1
            uid = uuid.uuid4()
            results[uid] = {
                "s3_filepath": meta["s3_filepath"],
                "score": score,
                "id": meta["id"],
                # "score": float(meta["audio_quality"]["quantification"]),
                # "start_s": meta["start_s"],
                # "end_s": meta["end_s"],
            }
            scores_list.append(score)

    print(
        f"min: {np.min(scores_list)} max: {np.max(scores_list)} mean: {np.mean(scores_list)}"
    )
    print()
    fig, ax = plt.subplots()
    counts, bins, patches = ax.hist(scores_list, bins=5, edgecolor="black")
    print(bins)

    # Add text annotations
    for count, patch in zip(counts, patches):
        height = patch.get_height()
        ax.text(
            patch.get_x() + patch.get_width() / 2.0,
            height,
            int(height),
            ha="center",
            va="bottom",
        )

    plt.yscale("log")
    plt.savefig(f"outputs/{subset}_scores.png", dpi=300)

    # save the audio from the top N worse and best audio
    for mode in ["worst"]:
        if mode == "worst":
            sorted_results = dict(
                sorted(results.items(), key=lambda item: item[1]["score"])
            )
        else:
            sorted_results = dict(
                sorted(results.items(), key=lambda item: item[1]["score"], reverse=True)
            )

        for idx, (uid, result) in enumerate(sorted_results.items()):
            if idx < 25:
                print(result["score"], result["s3_filepath"])

                # start_frame = int(result["start_s"] * sample_rate)
                # end_frame = int(result["end_s"] * sample_rate)
                # audio = audio[:, start_frame:end_frame]

                if save_audio:
                    s3_id = os.path.basename(result["s3_filepath"]).split(".")[0]
                    audio = Audio.from_s3(result["s3_filepath"], n_channels=2)
                    sample_rate = audio.sample_rate
                    audio = torch.from_numpy(audio.array_float)
                    torchaudio.save(
                        f"""outputs/audio/{subset}/{result["score"]:0.2f}-{s3_id}.mp3""",
                        audio,
                        sample_rate,
                        compression=torchaudio.io.CodecConfig(bit_rate=320_000),
                    )
        print()
