#!/usr/bin/env python3
"""
Unified script for downloading NPZ datasets from S3.
Supports different dataset types and configurations.

# For auk_t0 dataset
python download_npz_dataset.py \
    --preset auk_t0 \
    --data_path "/home/tony/Data/Preference/auk_t0/interesting_clips_auk_t0_20250527.pkl"

# For auk_t1 dataset  
python download_npz_dataset.py \
    --preset auk_t1 \
    --data_path "/home/tony/Data/Preference/auk_t1/interesting_clips_auk_t1_20250527.pkl"

# For vae_diffv2_d3 dataset
python download_npz_dataset.py \
    --preset vae_diffv2_d3 \
    --data_path "/home/tony/Data/Preference/up_v2_d3/interesting_clips_ahi_d3_20250527.pkl"
"""

import argparse
import os

import pandas as pd
from suno_utils.utils.s3 import download_s3_files


def get_preset_config(preset_name, data_path):
    """Get predefined configurations for common use cases."""
    presets = {
        "auk_t0": {
            "data_path": data_path,
            "npz_dir": "/app2/suno/data/dpo/auk_t0_npz/",
            "json_dir": "/app2/suno/data/dpo/auk_t0_json/",
            "dataset_type": "standard",
            "model_filter": "auk",
            "download_hoot": True,
            "download_vae": False,
        },
        "auk_t1": {
            "data_path": data_path,
            "npz_dir": "/app2/suno/data/dpo/auk_t1_npz/",
            "json_dir": "/app2/suno/data/dpo/auk_t1_json/",
            "dataset_type": "standard",
            "model_filter": "auk",
            "download_hoot": True,
            "download_vae": False,
        },
        "vae_diffv2_d3": {
            "data_path": data_path,
            "npz_dir": "/app2/suno/data/dpo/diff2_v2_d3",
            "json_dir": None,
            "dataset_type": "upsample",
            "model_filter": "up",
            "download_hoot": False,
            "download_vae": True,
        },
    }
    return presets.get(preset_name)


def parse_args():
    parser = argparse.ArgumentParser(description="Download NPZ datasets from S3")

    # Preset mode - simplified usage
    parser.add_argument(
        "--preset",
        type=str,
        choices=["auk_t0", "auk_t1", "vae_diffv2_d3"],
        help="Use predefined configuration preset",
    )

    # Required arguments
    parser.add_argument(
        "--data_path",
        type=str,
        required=True,
        help="Path to the input data file (.pkl or .csv)",
    )

    # Optional arguments (for manual configuration)
    parser.add_argument(
        "--npz_dir",
        type=str,
        help="Directory to store downloaded NPZ files",
    )
    parser.add_argument(
        "--json_dir",
        type=str,
        default=None,
        help="Directory to store downloaded JSON files (hoot files)",
    )
    parser.add_argument(
        "--dataset_type",
        type=str,
        choices=["standard", "upsample", "vae"],
        default="standard",
        help="Type of dataset processing",
    )
    parser.add_argument(
        "--model_filter",
        type=str,
        default=None,
        help="Filter for model names (e.g., 'auk', 'up')",
    )
    parser.add_argument(
        "--n_cores",
        type=int,
        default=32,
        help="Number of cores for parallel downloading",
    )
    parser.add_argument(
        "--download_vae", action="store_true", help="Download VAE files (_vae.npz)"
    )
    parser.add_argument(
        "--download_hoot", action="store_true", help="Download hoot JSON files"
    )
    parser.add_argument(
        "--skip_deleted",
        action="store_true",
        help="Skip trying to download from deleted folder",
    )

    args = parser.parse_args()

    # Apply preset configuration if specified
    if args.preset:
        preset_config = get_preset_config(args.preset, args.data_path)
        if not preset_config:
            raise ValueError(f"Unknown preset: {args.preset}")

        print(f"Using preset configuration: {args.preset}")

        # Apply preset values, but allow command line overrides
        for key, value in preset_config.items():
            if key == "data_path":
                continue  # Always use the provided data_path

            # Only override if the argument wasn't explicitly set
            current_value = getattr(args, key)
            if key in ["download_vae", "download_hoot"]:
                # For boolean flags, only override if not set
                if not current_value:
                    setattr(args, key, value)
            else:
                # For other arguments, override if None or not set
                if current_value is None:
                    setattr(args, key, value)

    # Validate required arguments
    if not args.npz_dir:
        parser.error("--npz_dir is required (or use --preset)")

    return args


def load_data(data_path):
    """Load data from pickle or CSV file."""
    print(f"Loading data from: {data_path}")

    if data_path.endswith(".pkl"):
        df = pd.read_pickle(data_path)
    elif data_path.endswith(".csv"):
        df = pd.read_csv(data_path)
    else:
        raise ValueError("Data file must be .pkl or .csv")

    print(f"Loaded data shape: {df.shape}")
    return df


def filter_dataframe(df, model_filter):
    """Filter dataframe by model name if specified."""
    if model_filter:
        print(f"Filtering by model name containing: {model_filter}")
        df = df[df["model_name"].str.contains(model_filter)]
        print(f"Filtered data shape: {df.shape}")
        print("Model name counts:")
        print(df["model_name"].value_counts())

    return df


def extract_s3_ids(df, dataset_type):
    """Extract S3 IDs based on dataset type."""
    print(f"Extracting S3 IDs for dataset type: {dataset_type}")

    if dataset_type == "upsample":
        # For upsample datasets, extract from metadata
        edit_clip_ids = df["metadata"].apply(lambda x: x.get("upsample_clip_id", ""))
        s3_ids = [s3_id for s3_id in edit_clip_ids if s3_id and len(s3_id) > 0]
        print(f"All upsample clip IDs: {len(s3_ids)}")
        unique_s3_ids = sorted(set(s3_ids))
        print(f"Unique upsample clip IDs: {len(unique_s3_ids)}")
        return unique_s3_ids

    elif dataset_type == "vae":
        # For VAE datasets, use original s3_id
        s3_ids = df["s3_id"].unique()
        print(f"Unique S3 IDs for VAE: {len(s3_ids)}")
        return s3_ids

    else:  # standard
        # For standard datasets, use s3_id values
        s3_ids = df["s3_id"].values
        print(f"Total S3 IDs: {len(s3_ids)}")
        return s3_ids


def create_s3_paths(s3_ids, file_type="npz"):
    """Create S3 and local paths for downloading."""
    if file_type == "npz":
        s3_paths = [
            f"s3://suno-data-uploads/studio/uploads/{s3_id}.npz" for s3_id in s3_ids
        ]
    elif file_type == "vae":
        s3_paths = [
            f"s3://suno-data-uploads/studio/uploads/{s3_id}_vae.npz" for s3_id in s3_ids
        ]
    elif file_type == "hoot":
        s3_paths = [
            f"s3://suno-data-uploads/studio/uploads/{s3_id}_hoot.json"
            for s3_id in s3_ids
        ]
    else:
        raise ValueError(f"Unknown file type: {file_type}")

    return s3_paths


def create_local_paths(s3_ids, local_dir, file_type="npz"):
    """Create local paths for downloaded files."""
    if file_type == "npz":
        local_paths = [f"{local_dir}/{s3_id}.npz" for s3_id in s3_ids]
    elif file_type == "vae":
        local_paths = [f"{local_dir}/{s3_id}_vae.npz" for s3_id in s3_ids]
    elif file_type == "hoot":
        local_paths = [f"{local_dir}/{s3_id}_hoot.json" for s3_id in s3_ids]
    else:
        raise ValueError(f"Unknown file type: {file_type}")

    return local_paths


def get_unfinished_downloads(s3_paths, local_paths, local_dir):
    """Get list of files that haven't been downloaded yet."""
    os.makedirs(local_dir, exist_ok=True)

    finished_paths = os.listdir(local_dir)
    finished_paths_set = set(finished_paths)

    unfinished_s3_paths = [
        path for path in s3_paths if os.path.basename(path) not in finished_paths_set
    ]
    unfinished_local_paths = [
        path for path in local_paths if os.path.basename(path) not in finished_paths_set
    ]

    return unfinished_s3_paths, unfinished_local_paths


def download_files(s3_paths, local_paths, n_cores, skip_deleted=False):
    """Download files from S3."""
    if not s3_paths:
        print("No files to download")
        return []

    print(f"Downloading {len(s3_paths)} files with {n_cores} cores...")

    # Try downloading from uploads folder first
    result = download_s3_files(s3_paths, local_paths, n_cores=n_cores)

    if not skip_deleted:
        # Check for remaining files and try deleted folder
        remaining_s3_paths = []
        remaining_local_paths = []

        for s3_path, local_path in zip(s3_paths, local_paths):
            if not os.path.exists(local_path):
                remaining_s3_paths.append(s3_path)
                remaining_local_paths.append(local_path)

        if remaining_s3_paths:
            print(
                f"Trying to download {len(remaining_s3_paths)} files from deleted folder..."
            )
            deleted_s3_paths = [
                path.replace("/uploads/", "/deleted/") for path in remaining_s3_paths
            ]
            download_s3_files(deleted_s3_paths, remaining_local_paths, n_cores=n_cores)

    return result


def main():
    args = parse_args()

    # Load and filter data
    df = load_data(args.data_path)
    df = filter_dataframe(df, args.model_filter)

    # Extract S3 IDs
    s3_ids = extract_s3_ids(df, args.dataset_type)

    # Download NPZ files
    print("\n=== Downloading NPZ files ===")
    file_type = "vae" if args.download_vae else "npz"
    s3_paths = create_s3_paths(s3_ids, file_type)
    local_paths = create_local_paths(s3_ids, args.npz_dir, file_type)

    unfinished_s3_paths, unfinished_local_paths = get_unfinished_downloads(
        s3_paths, local_paths, args.npz_dir
    )

    print(f"Jobs to be done: {len(unfinished_s3_paths)}")

    if unfinished_s3_paths:
        download_files(
            unfinished_s3_paths, unfinished_local_paths, args.n_cores, args.skip_deleted
        )

    print("Finished downloading NPZ files")

    # Download VAE files if requested and not already done
    if args.download_vae and file_type != "vae":
        print("\n=== Downloading VAE files ===")
        # For VAE files, we need the original s3_id, not upsample_clip_id
        if args.dataset_type == "upsample":
            vae_s3_ids = df["s3_id"].unique()
        else:
            vae_s3_ids = s3_ids

        vae_s3_paths = create_s3_paths(vae_s3_ids, "vae")
        vae_local_paths = create_local_paths(vae_s3_ids, args.npz_dir, "vae")

        unfinished_vae_s3_paths, unfinished_vae_local_paths = get_unfinished_downloads(
            vae_s3_paths, vae_local_paths, args.npz_dir
        )

        print(f"VAE jobs to be done: {len(unfinished_vae_s3_paths)}")

        if unfinished_vae_s3_paths:
            download_files(
                unfinished_vae_s3_paths,
                unfinished_vae_local_paths,
                args.n_cores,
                args.skip_deleted,
            )

        print("Finished downloading VAE files")

    # Download hoot JSON files if requested
    if args.download_hoot and args.json_dir:
        print("\n=== Downloading Hoot JSON files ===")
        hoot_s3_paths = create_s3_paths(s3_ids, "hoot")
        hoot_local_paths = create_local_paths(s3_ids, args.json_dir, "hoot")

        unfinished_hoot_s3_paths, unfinished_hoot_local_paths = (
            get_unfinished_downloads(hoot_s3_paths, hoot_local_paths, args.json_dir)
        )

        print(f"Hoot jobs to be done: {len(unfinished_hoot_s3_paths)}")

        if unfinished_hoot_s3_paths:
            download_files(
                unfinished_hoot_s3_paths,
                unfinished_hoot_local_paths,
                args.n_cores,
                args.skip_deleted,
            )

        print("Finished downloading hoot files")

    print("\n=== All downloads completed ===")


if __name__ == "__main__":
    main()
