import os
import json
import funcy
import boto3
import shutil
import numpy as np

from tqdm import tqdm
from concurrent.futures import ThreadPoolExecutor, as_completed
from suno_utils.utils.text import read_jsonl, write_jsonl
from suno_utils.utils.s3 import read_from_s3, _verify_s3_filepath


def get_s3_files(bucket_name, prefix, max_keys: int = 100000):
    all_files = []
    continuation_token = None

    while True:
        # Prepare the arguments for the request
        list_kwargs = {
            "Bucket": bucket_name,
            "Prefix": prefix,  # List objects under this prefix, or leave blank for all objects
        }

        if continuation_token:
            list_kwargs["ContinuationToken"] = continuation_token

        # Make the request to list objects
        response = s3.list_objects_v2(**list_kwargs)

        # Collect the file keys
        all_files += [obj["Key"] for obj in response.get("Contents", [])]

        # Check if more results are available
        if response.get("IsTruncated"):  # True if there are more results to fetch
            continuation_token = response["NextContinuationToken"]
        else:
            break  # No more results to fetch

    return all_files


def list_s3_directories(bucket_name: str, prefix: str = ""):
    """
    List all directories (prefixes) in an S3 bucket.

    Args:
        bucket_name (str): Name of the S3 bucket
        prefix (str): Optional prefix to filter results (like a directory path)

    Returns:
        List[str]: List of directory paths (prefixes)
    """
    s3_client = boto3.client("s3")
    directories = set()

    # Use paginator to handle buckets with many objects
    paginator = s3_client.get_paginator("list_objects_v2")
    page_iterator = paginator.paginate(Bucket=bucket_name, Prefix=prefix, Delimiter="/")

    # Collect all prefixes (directories)
    for page in page_iterator:
        # Get common prefixes (directories)
        if "CommonPrefixes" in page:
            for prefix_obj in page["CommonPrefixes"]:
                directories.add(prefix_obj["Prefix"])

        # Also check Contents for any directory-like objects
        if "Contents" in page:
            for obj in page["Contents"]:
                key = obj["Key"]
                # If the key contains a slash, add the directory part
                if "/" in key:
                    directory = key.rsplit("/", 1)[0] + "/"
                    directories.add(directory)

    return sorted(list(directories))


if __name__ == "__main__":
    # load the base metas
    VAE_DIM = 128
    VAE_RATE_HZ = 25
    VAE_MEMMAP_SIZE = 750
    SEMANTIC_VOCAB_SIZE = 4000
    SEMANTIC_MEMMAP_SIZE = 750
    VAL_SIZE = 1000
    USE_PAIRS = True  # use pairs instead of single examples
    USE_AUDIO_QUALITY = False
    OUT_DATA_DIR = "/app/suno/data/diffusion_ft/genius_hq_filtered_20k_25hz_20241115_v1"

    if os.path.exists(OUT_DATA_DIR):
        print(f"out data dir {OUT_DATA_DIR} already exists, deleting")
        shutil.rmtree(OUT_DATA_DIR)

    os.makedirs(OUT_DATA_DIR, exist_ok=True)

    # christian/data/upsample_100hz_v4_t_5_20241018

    bucket_name = "suno-data"
    base_dir = "christian/data/genius_hq_filtered_20k"
    output_name = "25hz_20241115_v2/"

    base_metas_path = os.path.join("s3://", bucket_name, base_dir, "metas.jsonl")
    base_metas = read_from_s3(base_metas_path, read_f=read_jsonl)
    print(len(base_metas))
    base_metas_map = {meta["id"]: meta for meta in base_metas}

    # s3 client
    s3 = boto3.client("s3")

    # load quality scores
    # quality_scores_path = "/home/christian/code/christian/notebooks/diff_dpo/upsample_v4_t_5_20241018_25hz_20241031_v1_quality_scores.json"
    # with open(quality_scores_path, "r") as f:
    #    quality_scores = json.load(f)

    # get all dir paths on s3
    dir_paths = list_s3_directories(bucket_name, f"{base_dir}/{output_name}")
    print("total dirs: ", len(dir_paths))

    valid_dir_paths = dir_paths
    # for dir_path in tqdm(dir_paths):
    #    if dir_path in quality_scores:
    #        valid_dir_paths.append(dir_path)

    # print(
    #    f"total dir paths with quality scores: {len(valid_dir_paths)}/{len(dir_paths)}"
    # )

    # split dir paths into train and val
    train_dir_paths = valid_dir_paths[:-VAL_SIZE]
    val_dir_paths = valid_dir_paths[-VAL_SIZE:]

    # now iterate over the val, then train metas
    for dset_type in ["val", "train"]:
        dset_dir_paths = val_dir_paths if dset_type == "val" else train_dir_paths

        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

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

        for chunk_idx, dir_path_chunk in enumerate(tqdm(valid_dir_paths_chunks)):

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

            def process_dir_path(dir_path):

                # get quality scores
                # dir_path_quality_scores = quality_scores[dir_path]
                meta_id = dir_path.strip("/").split("/")[-1]

                # get meta
                meta = base_metas_map[meta_id]

                # select the paths with the highest and lowest quality scores
                # sorted_s3_paths = sorted(
                #    dir_path_quality_scores,
                #    key=lambda x: dir_path_quality_scores[x],
                #    reverse=True,
                # )
                # positive_s3_path = sorted_s3_paths[0]
                # negative_s3_path = sorted_s3_paths[-1]

                # get files in dir
                files = get_s3_files(bucket_name, dir_path)

                npz_filepath = [f for f in files if f.endswith(".npz")][0]
                # load the npz file
                data = read_from_s3(
                    f"s3://{bucket_name}/{npz_filepath}", read_f=np.load
                )

                arr_s = data["semantic_codes"]
                arr_v_neg = data["upsampled_latents"]
                arr_v_pos = data["original_latents"]

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

            with ThreadPoolExecutor(max_workers=16) as executor:
                futures = [
                    executor.submit(process_dir_path, dir_path)
                    for dir_path in dir_path_chunk
                ]
                for future in as_completed(futures):
                    result = future.result()  # list of tuples (arr_s, arr_v, meta)
                    if result:
                        for arr_s, arr_v, meta in 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)
