from suno_utils.audio import Audio
import os
from tqdm import tqdm
import random


import os
import pandas as pd
from suno_utils.utils.text import read_jsonl, write_jsonl
from suno_utils.utils.s3 import read_from_s3
from suno_utils.worker.modal_base import get_modal_base_image


import random
import uuid
from tqdm import tqdm
import numpy as np
from collections import defaultdict

import uuid
import os
import tempfile
from suno_utils.utils.s3 import download_s3_files, upload_s3_files
import modal


app = modal.App(image=get_modal_base_image())
aws_secret = modal.Secret.from_name("studio-aws")
SECRETS = [aws_secret]


def mix_stack(bundle_metas, verbose=False):
    stem_metas = bundle_metas["stems"]

    stems_by_category = defaultdict(list)
    for stem in stem_metas:
        stems_by_category[stem["category"]].append(stem)
    for category in stems_by_category:
        random.shuffle(stems_by_category[category])

    category_order = list(stems_by_category.keys())
    random.shuffle(category_order)

    stem_metas = []
    for category in category_order:
        stem_metas.extend(stems_by_category[category])

    stem_audios = [
        Audio.from_s3(mm["s3_filepath"], sample_rate=48_000, n_channels=2)
        for mm in tqdm(stem_metas, desc="Loading stems", disable=not verbose)
    ]

    # make prefix's
    prefix_audios = []
    prefix_metas = []
    for i, stem_audio in tqdm(enumerate(stem_audios), desc="Mixing prefix", disable=not verbose):
        prefix_meta = {
            "id": str(uuid.uuid4()),
            "stems": [m["id"] for m in stem_metas[: i + 1]],
        }
        prefix_metas.append(prefix_meta)
        if len(prefix_audios) == 0:
            prefix_audios.append(stem_audio)
        else:
            prefix_audios.append(Audio.sum((prefix_audios[-1], stem_audio)))
    assert len(prefix_audios) == len(prefix_metas)
    assert len(stem_audios) == len(stem_metas)

    # make suffix's
    suffix_audios = []
    suffix_metas = []
    reversed_stem_audios = list(reversed(stem_audios))
    reversed_stem_metas = list(reversed(stem_metas))
    for i, stem_audio in tqdm(
        enumerate(reversed_stem_audios), desc="Mixing suffix", disable=not verbose
    ):
        suffix_meta = {
            "id": str(uuid.uuid4()),
            "stems": [m["id"] for m in reversed_stem_metas[: i + 1]],
        }
        suffix_metas.append(suffix_meta)
        if len(suffix_audios) == 0:
            suffix_audios.append(stem_audio)
        else:
            suffix_audios.append(Audio.sum((suffix_audios[-1], stem_audio)))
    assert len(suffix_audios) == len(suffix_metas)
    assert len(stem_audios) == len(stem_metas)

    # make groups
    group_audios = []
    group_metas = []
    groups_to_mix = [category for category in category_order if len(stems_by_category[category]) > 1]
    for category in tqdm(groups_to_mix, desc="Mixing groups", disable=not verbose):
        group_meta = {
            "id": str(uuid.uuid4()),
            "stems": [m["id"] for m in stem_metas if m["category"] == category],
            "category": category,
        }
        group_metas.append(group_meta)
        group_audios.append(
            Audio.sum(
                [stem_audios[i] for i in range(len(stem_metas)) if stem_metas[i]["category"] == category]
            )
        )

    # # compute energy for stems
    # for i, stem_audio in enumerate(stem_audios):
    #     energy = stem_audio.get_energy(bin_size_s=1).astype(np.int32).tolist()
    #     stem_metas[i]["energy"] = energy

    return (
        stem_metas,
        stem_audios,
        prefix_metas,
        prefix_audios,
        suffix_metas,
        suffix_audios,
        group_metas,
        group_audios,
    )


def create_and_upload(bundle_metas, verbose=False):
    (
        stem_metas,
        stem_audios,
        prefix_metas,
        prefix_audios,
        suffix_metas,
        suffix_audios,
        group_metas,
        group_audios,
    ) = mix_stack(bundle_metas, verbose=verbose)
    bundle_id = bundle_metas["id"]

    local_fps = []
    s3_filepaths = []
    with tempfile.TemporaryDirectory() as tempdir:
        for prefix_meta, prefix_audio in zip(prefix_metas, prefix_audios):
            fp = os.path.join(tempdir, f"{prefix_meta['id']}.opus")
            prefix_audio.to_opus(fp)
            local_fps.append(fp)
            s3_path = f"s3://suno-data/datasets/harvest/karaoke_versions/stems_v2/prefix_mix/{bundle_id}/{prefix_meta['id']}.opus"
            s3_filepaths.append(s3_path)
            prefix_meta["s3_filepath"] = s3_path

        for suffix_meta, suffix_audio in zip(suffix_metas, suffix_audios):
            fp = os.path.join(tempdir, f"{suffix_meta['id']}.opus")
            suffix_audio.to_opus(fp)
            local_fps.append(fp)
            s3_path = f"s3://suno-data/datasets/harvest/karaoke_versions/stems_v2/suffix_mix/{bundle_id}/{suffix_meta['id']}.opus"
            s3_filepaths.append(s3_path)
            suffix_meta["s3_filepath"] = s3_path

        for group_meta, group_audio in zip(group_metas, group_audios):
            fp = os.path.join(tempdir, f"{group_meta['id']}.opus")
            group_audio.to_opus(fp)
            local_fps.append(fp)
            s3_path = f"s3://suno-data/datasets/harvest/karaoke_versions/stems_v2/group_mix/{bundle_id}/{group_meta['id']}.opus"
            s3_filepaths.append(s3_path)
            group_meta["s3_filepath"] = s3_path

        upload_s3_files(
            local_fps,
            s3_filepaths,
            chunksize=1000,
            n_cores=20,
            joblib_backend="threads",
            silent=True,
        )

    meta = {
        "id": str(uuid.uuid4()),
        "duration_s": bundle_metas["duration_s"],
        "stems": stem_metas,
        "prefix": prefix_metas,
        "suffix": suffix_metas,
        "group": group_metas,
    }
    return meta


@app.function(secrets=SECRETS, max_containers=1_500)
def safe_create_and_upload(bundle_metas):
    try:
        return create_and_upload(bundle_metas)
    except Exception as e:
        print(e)
        # raise e
        return None


@app.local_entrypoint()
def main():
    # Set up Modal for distributed processing
    stub = modal.Stub("stem-processing")

    S3_AUDIO_DIR = "s3://suno-data/datasets/harvest/karaoke_versions/stems/audio/"

    raw_metas = read_from_s3(
        "s3://suno-data/datasets/harvest/karaoke_versions/stems/karaoke_versions.jsonl",
        read_f=read_jsonl,
    )
    track_descriptions = []
    for m in raw_metas:
        for mm in m["tracks"]:
            track_descriptions.append(mm["description"])
    s = pd.Series(track_descriptions).value_counts()
    print(f"{s.shape[0]} types of stems")
    print(f"{s[s >= 1000].shape[0]} with >=1000")
    print(f"{s[s >= 100].shape[0]} with >=100")

    import json

    # Create a mapping dictionary from category to instrument list
    with open("instrument_categories.json", "r") as f:
        category_to_instruments = json.load(f)

    write_jsonl(category_to_instruments, "instrument_categories.jsonl")
    instrument_to_category = {}
    for category, instruments in category_to_instruments.items():
        for instrument in instruments:
            instrument_to_category[instrument] = category

    def get_category(instrument):
        instrument = instrument.replace(" ", "_")
        return instrument_to_category.get(instrument, "Other")

    metas = []
    for m in raw_metas:
        if any(["file_path" not in mm for mm in m["tracks"]]):
            # only very few are missing this
            continue
        # excluding click (metronome) here
        metas.append(
            {
                "id": m["song_id"],
                "duration_s": round((m["preview_end"] - m["preview_start"]) / 1000, 1),
                "stems": [
                    {
                        "id": f"{m['song_id']}_{i}",
                        "title": mm["description"],
                        "s3_filepath": os.path.join(S3_AUDIO_DIR, mm["file_path"]),
                        "category": get_category(mm["description"]),
                    }
                    for i, mm in enumerate(m["tracks"])
                    if mm["description"] != "Click"
                ],
            }
        )
    assert len([m["id"] for m in metas]) == len(set([m["id"] for m in metas]))
    print(f"{len(metas):,} tracks")
    n_stems = 0
    tot_track_duration_s = 0
    tot_stem_duration_s = 0
    for m in metas:
        n_stems += len(m["stems"])
        tot_track_duration_s += m["duration_s"]
        tot_stem_duration_s += m["duration_s"] * len(m["stems"])
    print(f"{n_stems:,} stems")
    print(f"{round(tot_track_duration_s / 60 / 60):,} hours total tracks")
    print(f"{round(tot_stem_duration_s / 60 / 60):,} hours total stems")

    results = safe_create_and_upload.map(metas[:])
    print(results)
    import json

    # save metas to file
    with open("/tmp/grouped_metas.jsonl", "w") as f:
        for meta in results:
            f.write(json.dumps(meta) + "\n")
