import os
import numpy as np
import shutil
import random
from tqdm import tqdm
from suno_utils.utils.text import read_jsonl, write_jsonl
from concurrent.futures import ThreadPoolExecutor, as_completed

if __name__ == "__main__":
    NPZ_DIR = "/app/suno/christian/data/genius_hq_filtered_raw_10s_1920_npz"
    OUT_DATA_DIR = "/app/suno/christian/data/genius_hq_filtered_raw_10s_1920_memmap"
    METAS_PATH = (
        "/home/christian/code/christian/metadata/genius_hq_metas_filtered.jsonl"
    )

    VAL_SIZE = 10_000
    VAL_ONLY = False
    DURATION_S = 10.0
    VAE_FRAME_RATE = 25
    SEMANTIC_FRAME_RATE = 25
    CHUNK_SIZE = 100

    SEMANTIC_MEMMAP_SIZE = int(DURATION_S * SEMANTIC_FRAME_RATE)
    VAE_MEMMAP_SIZE = int(DURATION_S * VAE_FRAME_RATE * 2)
    VAE_DIM = 1920
    print(f"VAE_MEMMAP_SIZE: {VAE_MEMMAP_SIZE}")
    print(f"SEMANTIC_MEMMAP_SIZE: {SEMANTIC_MEMMAP_SIZE}")

    # load base metas
    base_metas = read_jsonl(METAS_PATH)
    print(f"Found {len(base_metas)} metas")

    # delete the out dir if it exists
    if os.path.exists(OUT_DATA_DIR):
        shutil.rmtree(OUT_DATA_DIR)

    os.makedirs(OUT_DATA_DIR, exist_ok=True)

    # shuffle the metas
    random.seed(42)
    random.shuffle(base_metas)

    # split into train and val
    train_metas = base_metas[:-VAL_SIZE]
    val_metas = base_metas[-VAL_SIZE:]
    # now iterate over the val, then train metas

    if VAL_ONLY:
        dset_types = ["val"]
    else:
        dset_types = ["val", "tr"]

    for dset_type in dset_types:
        dset_metas = val_metas if dset_type == "val" else train_metas

        out_mm_vae_filepath = os.path.join(OUT_DATA_DIR, f"data_vae_{dset_type}.bin")
        out_metas_filepath = os.path.join(OUT_DATA_DIR, f"metas_{dset_type}.jsonl")
        out_mm_semantic_filepath = os.path.join(
            OUT_DATA_DIR, f"data_semantic_{dset_type}.bin"
        )

        # initial write
        out_mm_semantic = np.memmap(
            out_mm_semantic_filepath,
            dtype=np.uint16,
            mode="w+",
            shape=(1),
        )
        out_mm_vae = np.memmap(
            out_mm_vae_filepath,
            dtype=np.float16,
            mode="w+",
            shape=(1),
        )

        n_offs_s = 0
        n_offs_v = 0

        # split valid metas into chunks of CHUNK_SIZE
        dset_metas_chunks = [
            dset_metas[i : i + CHUNK_SIZE]
            for i in range(0, len(dset_metas), CHUNK_SIZE)
        ]
        print("total chunks: ", len(dset_metas_chunks))

        for chunk_idx, dset_meta_chunk in enumerate(tqdm(dset_metas_chunks)):

            arr_s_list = []
            arr_v_list = []
            new_metas = []

            def process_meta(meta):

                # try to load from npz
                npz_path = os.path.join(NPZ_DIR, f"{meta['id']}.npz")
                if os.path.exists(npz_path):
                    data = np.load(npz_path)
                    arr_s = data["semantic_codes"]
                    arr_v = data["frames"]
                else:
                    return None

                # check for nan in arr_v or arr_s
                if not np.all(np.isfinite(arr_v)) or not np.all(np.isfinite(arr_s)):
                    return None

                try:
                    if arr_s.size < SEMANTIC_MEMMAP_SIZE:
                        return None
                    if arr_v.size < VAE_MEMMAP_SIZE * VAE_DIM:
                        return None
                    result = (arr_s, arr_v, meta)
                    return result
                except Exception as e:
                    print(f"error loading {meta['id']}: {e}")
                    return None

            with ThreadPoolExecutor(max_workers=16) as executor:
                futures = [
                    executor.submit(process_meta, meta) for meta in dset_meta_chunk
                ]
                for future in as_completed(futures):
                    result = future.result()
                    if result:
                        arr_s, arr_v, meta = result
                        arr_s_list.append(arr_s)
                        arr_v_list.append(arr_v)
                        new_meta = meta.copy()
                        new_meta["text"] = meta["lyrics"]
                        new_meta["tags"] = meta["tags_text"]
                        new_meta["n_vae_tokens"] = VAE_MEMMAP_SIZE
                        new_metas.append(new_meta)

            print(len(arr_s_list), len(arr_v_list), len(new_metas))
            assert len(arr_s_list) == len(arr_v_list) == len(new_metas)

            # now write to the memmap
            # get a list of all the ids in the id_to_s3_paths
            to_write_len_s = SEMANTIC_MEMMAP_SIZE * len(arr_v_list)
            to_write_len_v = VAE_MEMMAP_SIZE * VAE_DIM * len(arr_v_list)

            out_mm_semantic = np.memmap(
                out_mm_semantic_filepath,
                dtype=np.uint16,
                mode="r+",
                shape=(n_offs_s + to_write_len_s,),
            )

            out_mm_vae = np.memmap(
                out_mm_vae_filepath,
                dtype=np.float16,
                mode="r+",
                shape=(n_offs_v + to_write_len_v,),
            )

            # write to memmap (has to happen sequentially)
            for new_meta, arr_s, arr_v in zip(new_metas, arr_s_list, arr_v_list):
                out_mm_semantic[n_offs_s : n_offs_s + arr_s.size] = arr_s.reshape(
                    -1,
                )
                out_mm_vae[n_offs_v : n_offs_v + arr_v.size] = arr_v.reshape(
                    -1,
                )
                n_offs_s += arr_s.size
                n_offs_v += arr_v.size

            # write it once
            out_mm_semantic.flush()
            out_mm_vae.flush()
            del out_mm_semantic, out_mm_vae

            write_jsonl(new_metas, out_metas_filepath, do_append=True)
