import os
import gc
import json
import uuid
import tqdm
import time
import torch
import funcy
import random
import tempfile
import collections
import numpy as np
import pandas as pd
from joblib import Parallel, delayed
from matplotlib import pyplot as plt

from scipy.io import wavfile
from suno_utils.audio import Audio
from suno_utils.tasks.data_loader import load_audio_mp
from suno_utils.utils.text import write_jsonl, read_jsonl, write_json, read_json
from suno_utils.utils.s3 import read_from_s3, check_s3_file_exists, open_from_s3
from joblib.externals.loky import get_reusable_executor


def _convert_float_audio(sig):
    dtype = np.int16
    dtype_info = np.iinfo(dtype)
    abs_max = 2 ** (dtype_info.bits - 1)
    offset = dtype_info.min + abs_max
    return (sig * abs_max + offset).clip(dtype_info.min, dtype_info.max).astype(dtype)


def load_metas(filepath, simple=True):
    assert filepath.startswith("s3://")
    data = []
    with open_from_s3(filepath) as f:
        for line in f:
            line = line.strip()
            if len(line) == 0:
                continue
            m = json.loads(line)
            _id = m["id"]
            #             duration_s = m["duration_s"]
            filepath = m.get("s3_filepath", m.get("audio_filepath", m.get("filepath")))
            assert filepath is not None
            m_new = {
                "id": _id,
                "filepath": filepath,
                #                 "duration_s": duration_s,
            }
            if not simple:
                for k in m.keys():
                    #                     if k not in ["id", "duration_s", "s3_filepath", "audio_filepath", "filepath"]:
                    if k not in ["id", "s3_filepath", "audio_filepath", "filepath"]:
                        m_new[k] = m[k]
            data.append(m_new)
    return data


def _write_item(work_item):
    audio_arr, out_filepath = work_item
    try:
        wavfile.write(out_filepath, SAMPLE_RATE, audio_arr.T)
    except:
        return False
    return True


def _mp_write(audio_arr_list, out_filepaths, num_workers=16):
    assert len(audio_arr_list) == len(out_filepaths)
    work_items = list(zip(audio_arr_list, out_filepaths))
    confirmed_list = Parallel(n_jobs=num_workers, prefer="threads", batch_size=1)(
        delayed(_write_item)(work_item) for work_item in work_items
    )
    get_reusable_executor().shutdown(wait=True)
    return confirmed_list


def build_dataset(
    raw_metas_info,
    sample_rate: int = 48000,
    base_s3_dir: str = "s3://suno-data/datasets",
    is_test: bool = True,
):
    filepath_info = {}
    for dset_key, rel_fp, _, _ in tqdm.tqdm(raw_metas_info):
        metas = load_metas(os.path.join(base_s3_dir, rel_fp))
        random.shuffle(metas)
        filepath_info[dset_key] = metas

    if is_test:
        metas_info = [(a, b, c, 1) for a, b, c, d in raw_metas_info]
    else:
        metas_info = [e for e in raw_metas_info]

    sample_metas_tr = []
    sample_metas_val = []
    is_finished = False
    # for each dset, load in chunks of 1k files, and keep doing until we have enough
    for dset_key, _, (req_min_s, req_max_s), req_duration_h in metas_info:
        dset_duration_h = 0
        for n_iter, filepaths_chunk in enumerate(
            funcy.chunks(5000, filepath_info[dset_key])
        ):
            is_val = n_iter == 0
            if is_val:  # val, so make it smaller
                filepaths_chunk = filepaths_chunk[:500]
            out_dir = os.path.join(out_data_dir, "val" if is_val else "train", dset_key)
            os.makedirs(out_dir, exist_ok=True)
            filepaths_chunk_list = [m["filepath"] for m in filepaths_chunk]
            t0 = time.time()
            audio_arr_list = load_audio_mp(
                filepaths_chunk_list,
                target_sample_rate=sample_rate,
                n_channels=2,
                #             min_duration_s=req_min_s,  # TODO: not here cause we want to skip
                max_duration_s=req_max_s,
                normalize_volume=True,
                num_workers=32,
                force_threads=False,
                #     debug=False,
                silent=True,
            )
            # add random offset just incase
            filtered_audio_arr_list = []
            offset_list = []
            for arr in audio_arr_list:
                if arr is None:
                    filtered_audio_arr_list.append(arr)
                    offset_list.append(0)
                    continue
                offset = int(
                    round(random.uniform(0, arr.shape[-1] // 4 / sample_rate), 1)
                    * sample_rate
                )
                filtered_audio_arr_list.append(arr[:, offset:])
                offset_list.append(offset / sample_rate)
            audio_arr_list = filtered_audio_arr_list
            del filtered_audio_arr_list
            audio_arr_list = [
                _convert_float_audio(arr.numpy()) if arr is not None else None
                for arr in audio_arr_list
            ]
            td_fetch = int(round(time.time() - t0))
            time.sleep(5)  # make sure things close
            # multicore writing
            t0 = time.time()
            n_offset = len(sample_metas_val) if is_val else len(sample_metas_tr)
            new_ids = [str(uuid.uuid4()) for _ in filepaths_chunk]
            out_filepaths = [
                os.path.join(out_dir, f"{new_id}.wav") for new_id in new_ids
            ]
            confirmed_list = _mp_write(audio_arr_list, out_filepaths)
            tot_duration_s = 0
            for is_confirmed, new_id, fp, m, offset_s, arr in zip(
                confirmed_list,
                new_ids,
                out_filepaths,
                filepaths_chunk,
                offset_list,
                audio_arr_list,
            ):
                if not is_confirmed:
                    continue
                duration_s = arr.shape[-1] / sample_rate
                if duration_s < req_min_s:
                    continue
                new_m = {
                    "dataset": dset_key,
                    "id": new_id,
                    "original_id": m["id"],
                    "filepath": fp,
                    "offset_s": offset_s,
                    "duration_s": round(duration_s, 2),
                }
                if is_val:
                    sample_metas_val.append(new_m)
                else:
                    sample_metas_tr.append(new_m)
                tot_duration_s += duration_s
            chunk_duration_h = round(tot_duration_s / 60 / 60, 1)
            del audio_arr_list
            gc.collect()
            td_write = int(round(time.time() - t0))
            time.sleep(5)  # make sure things close
            dset_type = "val" if is_val else "train"
            print(
                f"{dset_key}: {chunk_duration_h:,} hours of data fetched in {td_fetch:,}s"
                f" and written in {td_write:,}s as `{dset_type}`, retained {np.mean(confirmed_list)*100:.1f}%"
            )
            if not is_val:
                dset_duration_h += chunk_duration_h
            if dset_duration_h >= req_duration_h:
                print(
                    f"done with {dset_key}, collected total of {dset_duration_h:,.1f} hours for train"
                )
                break
    is_finished = True

    # TODO: summarize fetched data amounts
    from collections import defaultdict

    val_durations_s = defaultdict(int)
    tr_durations_s = defaultdict(int)
    for m in sample_metas_val:
        val_durations_s[m["dataset"]] += m["duration_s"]
    for m in sample_metas_tr:
        tr_durations_s[m["dataset"]] += m["duration_s"]
    for k, v in tr_durations_s.items():
        print(f"{v/60/60:,.1f} hours of {k} in train")
    print()
    for k, v in val_durations_s.items():
        print(f"{v/60/60:,.1f} hours of {k} in val")

    assert is_finished
    write_jsonl(sample_metas_val, os.path.join(out_data_dir, "metas_val.jsonl"))
    write_jsonl(sample_metas_tr, os.path.join(out_data_dir, "metas_tr.jsonl"))


def make_dac_manifests(out_data_dir: str):
    metas_tr = read_jsonl(os.path.join(out_data_dir, "metas_tr.jsonl"))
    metas_val = read_jsonl(os.path.join(out_data_dir, "metas_val.jsonl"))

    MUSIC_DATASETS = set(
        [
            "podcasts",
            "genius_hq",
            "youtube_music",
            "jamendo",
            "imslp",
            "pond5_music",
            "spot_genres",
            "tency",
            "shutter_music",
            "pond5_sfx",
        ]
    )

    metas_tr = [m for m in metas_tr if m["dataset"] in MUSIC_DATASETS]
    metas_val = [m for m in metas_val if m["dataset"] in MUSIC_DATASETS]
    print(f"{sum(m['duration_s'] for m in metas_tr)/60/60:,.0f} hours in train")
    print(f"{sum(m['duration_s'] for m in metas_val)/60/60:,.0f} hours in val")

    random.seed(7007)
    random.shuffle(metas_tr)
    random.shuffle(metas_val)

    df_tr = pd.DataFrame([{"path": m["filepath"]} for m in metas_tr])
    df_val = pd.DataFrame([{"path": m["filepath"]} for m in metas_val])

    df_tr.to_csv(os.path.join(out_data_dir, "music_tr.csv"), index=False)
    df_val.to_csv(os.path.join(out_data_dir, "music_val.csv"), index=False)


if __name__ == "__main__":

    SAMPLE_RATE = 48_000
    is_test = False

    out_data_dir = "/app/suno/data/audio_2ch_48khz_lg"
    os.makedirs(out_data_dir, exist_ok=True)
    base_s3_dir = "s3://suno-data/datasets"

    # dset_name, metas_filepath, (min_s, max_s), max_hours
    raw_metas_info = [
        # speech
        ("podcasts", "bundles/v0/podcasts/metas.jsonl", (60, 2 * 60), 2_500),
        # music`
        ("genius_hq", "bundles/v1/genius_hq/metas.jsonl", (60, 2 * 60), 5_000),
        ("youtube_music", "bundles/v1/youtube_music/metas.jsonl", (60, 2 * 60), 10_000),
        ("jamendo", "bundles/v1/jamendo/metas.jsonl", (60, 2 * 60), 1_000),
        ("imslp", "bundles/v1/imslp/metas.jsonl", (60, 2 * 60), 1_000),
        ("pond5_music", "bundles/v2/pond5_music/metas.jsonl", (60, 2 * 60), 2_500),
        (
            "spot_genres",
            "harvest/spotify/ytm_spotify_genres_50_simple.jsonl",
            (60, 2 * 60),
            10_000,
        ),
        ("tency", "harvest/tency/tency_plus_fp_flat.jsonl", (20, 2 * 60), 500),
        (
            "shutter_music",
            "bundles/v3/shutter_music/metas_stems_flat.jsonl",
            (20, 2 * 60),
            500,
        ),
        # misc
        ("pond5_sfx", "bundles/v3/pond5_sfx/metas.jsonl", (5, 30), 500),
    ]

    random.seed(6006)

    build_dataset(
        raw_metas_info,
        sample_rate=SAMPLE_RATE,
        base_s3_dir=base_s3_dir,
        is_test=is_test,
    )
    make_dac_manifests(out_data_dir)
