import os

import argparse
import json
from tqdm.contrib.concurrent import thread_map

from suno_utils.audio import Audio
from suno_utils.utils.text import read_jsonl


def parse_args():
    parser = argparse.ArgumentParser()
    parser.add_argument("--dataset", type=str)
    parser.add_argument("--reversed", type=bool, default=False)
    args = parser.parse_args()
    return args


if __name__ == "__main__":
    input_args = parse_args()
    input_args.dataset
    if input_args.dataset == "genius":
        INPUT_PATH = "/home/tony/Data/Hoot/genius_metas.json"
        DOWNLOAD_PATH = "/app/suno/data/hoot/audios/genius"
    elif input_args.dataset == "ytm":
        INPUT_PATH = "/home/tony/Data/Hoot/ytm_metas.json"
        DOWNLOAD_PATH = "/app/suno/data/audios/ytm"
    elif input_args.dataset == "deezer":
        INPUT_PATH = "/app/suno/tmp/clean_deezer_v0_metas.jsonl"
        DOWNLOAD_PATH = "/app/suno/data/audios/deezer"
    elif input_args.dataset == "discogs":
        INPUT_PATH = "/app/suno/tmp/clean_discogs_subset_v0_metas.jsonl"
        DOWNLOAD_PATH = "/app/suno/data/audios/discogs"
    else:
        raise ValueError("Unknown dataset")

    print("processing dataset", input_args.dataset)
    if INPUT_PATH.endswith(".json"):
        with open(INPUT_PATH, "r") as fp:
            metas_map = json.load(fp)
    else:
        metas_map = read_jsonl(INPUT_PATH)
    selected_metas = metas_map[:]
    selected_metas.sort(key=lambda x: x["id"])
    if input_args.reversed:
        selected_metas = selected_metas[::-1]
    os.makedirs(DOWNLOAD_PATH, exist_ok=True)
    print(len(selected_metas), len(set(m["id"] for m in selected_metas)))

    def clean_bad_downloads(index):
        meta = selected_metas[index]
        if "original_id" in meta:
            output_path = os.path.join(DOWNLOAD_PATH, f"{meta['original_id']}.mp3")
        else:
            output_path = os.path.join(DOWNLOAD_PATH, f"{meta['id']}.mp3")
        if os.path.exists(output_path):
            # # check file download status, this is specific for just fixing issues
            file_size = os.stat(output_path).st_size
            if file_size < 1000:
                os.remove(output_path)
                return
            # file_time = datetime.datetime.fromtimestamp(os.path.getmtime(output_path))
            # if file_time > datetime.datetime(2024, 2, 14, 10):
            #     print(output_path, file_time)
            #     os.remove(output_path)
            #     return

    def download_audio_data(index):
        meta = selected_metas[index]
        if "original_id" in meta:
            output_path = os.path.join(DOWNLOAD_PATH, f"{meta['original_id']}.mp3")
        else:
            output_path = os.path.join(DOWNLOAD_PATH, f"{meta['id']}.mp3")
        if os.path.exists(output_path):
            return
        try:
            # create a fake place holder asap
            # with open(output_path, "w") as fp:
            #     fp.write("")
            if "audio_filepath" in meta:
                audio = Audio.from_s3(meta["audio_filepath"])
            else:
                audio = Audio.from_s3(meta["s3_filepath"])
            # for hoot this is fine
            audio = audio.convert(16_000, audio.byte_width, n_channels=1)
            audio.to_mp3(output_path)
            return
        except Exception as e:
            print(e, index)
            # clean the crap up
            if os.path.exists(output_path):
                os.remove(output_path)
            return

    print("START")
    # thread_map(
    #     clean_bad_downloads,
    #     list(range(len(selected_metas))),
    #     max_workers=40,
    #     chunksize=1,
    # )
    # print("Done cleaning")
    thread_map(
        download_audio_data,
        list(range(len(selected_metas))),
        max_workers=100,
        chunksize=1,
    )
    print("DONE")
