import os
import json
import random
import time
import pandas as pd
from tqdm import tqdm
from suno_utils.utils.text import write_jsonl
from joblib import Parallel, delayed
from typing import Dict, List, Callable, Any, Optional, Tuple
from dataclasses import dataclass, field


# ============================================================================
# Utility Functions
# ============================================================================


def read_jsonl_lazy(filepath):
    """Lazy generator to read JSONL file line by line without loading everything into memory"""
    with open(filepath, "r", encoding="utf-8") as f:
        for line in f:
            line = line.strip()
            if line:
                try:
                    yield json.loads(line)
                except json.JSONDecodeError:
                    continue


def count_jsonl_lines(filepath):
    """Count lines in a JSONL file efficiently"""
    count = 0
    with open(filepath, "r", encoding="utf-8") as f:
        for _ in f:
            count += 1
    return count


# ============================================================================
# Metadata Loader Functions
# ============================================================================


def load_json_dict(filepath: str, transform: Optional[Callable] = None) -> Dict:
    """Load a JSON file as a dictionary"""
    with open(filepath, "r") as f:
        data = json.load(f)
    if transform:
        data = transform(data)
    return data


def load_jsonl_dict(filepath: str, key_field: str = "id") -> Dict:
    """Load a JSONL file as a dictionary indexed by key_field"""
    return {meta[key_field]: meta for meta in read_jsonl_lazy(filepath)}


def load_csv_dict(filepath: str, key_field: str = "id") -> Dict:
    """Load a CSV file as a dictionary indexed by key_field"""
    df = pd.read_csv(filepath)
    return df.set_index(key_field).to_dict(orient="index")


def load_alignment_map(filepath: str) -> Dict:
    """Load alignment data from JSONL file"""
    result = {}
    for item in tqdm(read_jsonl_lazy(filepath), desc="Building alignment map"):
        k, v, cer = item
        meta = v[0]
        lines, starts, ends = (
            meta["line_text"],
            meta["line_start_s"],
            meta["line_end_s"],
        )

        line_entries = []
        for text, start, end in zip(lines, starts, ends):
            if start is None or end is None:
                continue
            line_entries.append((start, end, text))

        result[k] = {
            "lines": line_entries,
            "cer": cer,
            "text": meta.get("text"),
            "start_s": meta.get("start_s"),
            "end_s": meta.get("end_s"),
            "vocal_start_s": meta.get("vocal_start_s"),
            "vocal_end_s": meta.get("vocal_end_s"),
        }
    return result


# ============================================================================
# Metadata Configuration
# ============================================================================


@dataclass
class MetadataSource:
    """Configuration for a metadata source"""

    name: str  # Name for this metadata source
    filepath: str  # Path to the metadata file
    loader: Callable  # Function to load the metadata
    loader_kwargs: Dict[str, Any] = field(default_factory=dict)  # Arguments for loader
    merge_key: str = "id"  # Key to use for merging (default: "id")
    merge_function: Optional[Callable] = None  # Custom merge function if needed
    output_key: Optional[str] = None  # Key to store in output (default: same as name)
    priority: int = 0  # Priority for merging (lower = higher priority, for conflicts)
    enabled: bool = True  # Whether to enable this metadata source


@dataclass
class MetadataRegistry:
    """Registry of all metadata sources"""

    sources: Dict[str, MetadataSource] = field(default_factory=dict)
    loaded_data: Dict[str, Any] = field(default_factory=dict)

    def register(self, source: MetadataSource):
        """Register a metadata source"""
        self.sources[source.name] = source

    def load_all(self):
        """Load all enabled metadata sources"""
        print("Loading external metadata...")
        for name, source in self.sources.items():
            if not source.enabled:
                continue
            print(f"  Loading {name}...")
            try:
                self.loaded_data[name] = source.loader(
                    source.filepath, **source.loader_kwargs
                )
                size = (
                    len(self.loaded_data[name])
                    if isinstance(self.loaded_data[name], dict)
                    else "N/A"
                )
                print(f"    Loaded {name}: {size:,} entries")
            except Exception as e:
                print(f"    WARNING: Failed to load {name}: {e}")
                self.loaded_data[name] = {}

    def merge_into_meta(self, meta: Dict, meta_id: str) -> Dict:
        """Merge all registered metadata into a meta dictionary"""
        new_meta = meta.copy()

        # Sort sources by priority (lower priority number = higher priority)
        sorted_sources = sorted(
            self.sources.items(), key=lambda x: (x[1].priority, x[0])
        )

        for name, source in sorted_sources:
            if not source.enabled or name not in self.loaded_data:
                continue

            data = self.loaded_data[name]
            if meta_id not in data:
                continue

            value = data[meta_id]

            # Use custom merge function if provided
            if source.merge_function:
                source.merge_function(new_meta, value, name)
            else:
                # Default: merge directly using output_key or name
                output_key = source.output_key or name
                new_meta[output_key] = value

        return new_meta


# ============================================================================
# Custom Merge Functions
# ============================================================================


def merge_alignment_data(meta: Dict, alignment_data: Dict, source_name: str):
    """Custom merge function for alignment data"""
    meta["text_aligned"] = alignment_data["lines"]
    meta["cer"] = alignment_data["cer"]
    meta["vocal_start_s"] = alignment_data["vocal_start_s"]
    meta["vocal_end_s"] = alignment_data["vocal_end_s"]


def merge_tags(meta: Dict, tags: Any, source_name: str):
    """Custom merge function for tags (combines into a list)
    Accepts either a list of tags or a dict with a 'tags' key
    """
    if "tags" not in meta:
        meta["tags"] = []

    # If tags is a dict with a 'tags' key, extract it
    if isinstance(tags, dict) and "tags" in tags:
        tags = tags["tags"]

    # Now handle tags as a list or single value
    if isinstance(tags, list):
        meta["tags"].extend(tags)
    elif tags:
        meta["tags"].append(tags)


def merge_stem_captions(meta: Dict, stem_captions_dict: Dict, source_name: str):
    """Custom merge function for stem captions"""
    if "stem_captions" in stem_captions_dict:
        meta["stem_captions"] = stem_captions_dict["stem_captions"]
    if "vocal_pitch_range" in stem_captions_dict:
        meta["vocal_pitch_range"] = stem_captions_dict["vocal_pitch_range"]


def merge_all_keys(
    meta: Dict,
    extra_metadata: Dict,
    source_name: str,
    ignore_keys: List[str] = ["artists"],
):
    """Merge all keys from extra_metadata directly into meta, handling tags carefully."""
    if not isinstance(extra_metadata, dict):
        return

    for key, value in extra_metadata.items():
        if key in ignore_keys:
            continue
        if key == "tags":
            if "tags" not in meta:
                # If meta doesn't already have a tags field, create it as a list or single item
                meta["tags"] = (
                    list(value)
                    if isinstance(value, list)
                    else ([value] if value else [])
                )
            else:
                # Merge tags into existing tags list
                if isinstance(value, list):
                    meta["tags"].extend(value)
                elif value:
                    meta["tags"].append(value)
        else:
            meta[key] = value


def merge_priority_alignments(meta: Dict, alignment_data: Dict, source_name: str):
    """Merge alignment data only if not already present (priority-based)"""
    if "text_aligned" not in meta:
        merge_alignment_data(meta, alignment_data, source_name)


# ============================================================================
# Processing Functions
# ============================================================================


def filter_meta(meta: Dict) -> Tuple[Optional[Dict], Optional[str], float]:
    """Filter a single meta entry"""
    # Duration filter
    if meta["duration_s"] < 30 or meta["duration_s"] > 720:
        return None, "duration", meta["duration_s"]

    # check if we have audio_stats, if we dont, skip
    if "audio_stats" not in meta:
        return None, "audio_stats", meta["duration_s"]

    # check if the path exists in the local_filepath_dir
    local_filepath = meta["local_filepath"]
    if not os.path.exists(local_filepath):
        return None, "local_filepath", meta["duration_s"]

    # Loudness filter
    if meta["audio_stats"]["loudness"] < -20 or meta["audio_stats"]["loudness"] > -6:
        return None, "loudness", meta["duration_s"]

    # Silence filter
    if meta["audio_stats"]["silence_percentage"] > 4:
        return None, "silence", meta["duration_s"]

    # CER filter (if alignments exist)
    if meta.get("text_aligned"):
        if meta["cer"] > 0.8:
            return None, "cer", meta["duration_s"]

    # passed all checks, keep meta
    return meta, None, meta["duration_s"]


def normalize_artist_field(meta: Dict) -> Dict:
    """Normalize artist field: rename 'artist' to 'artist_ids' and ensure it's a list of strings"""
    if "artist" in meta:
        artist_value = meta.pop("artist")

        # Convert to list of strings if it's not already
        if isinstance(artist_value, list):
            # Ensure all elements are strings
            artist_ids = [str(item) for item in artist_value if item is not None]
        elif artist_value is not None:
            # Convert single value to list of strings
            artist_ids = [str(artist_value)]
        else:
            # Handle None case
            artist_ids = []

        meta["artist_ids"] = artist_ids

    return meta


def merge_and_filter_meta(
    meta: Dict,
    registry: MetadataRegistry,
    local_filepath_dir: str,
) -> Tuple[Optional[Dict], Optional[str], float, str]:
    """Merge metadata and filter in a single pass"""
    meta_id = meta["id"]
    new_meta = meta.copy()

    # Merge all registered metadata
    new_meta = registry.merge_into_meta(new_meta, meta_id)

    # Normalize artist field (rename to artist_ids and ensure it's a list of strings)
    new_meta = normalize_artist_field(new_meta)

    # Construct local filepath
    local_filepath = os.path.join(local_filepath_dir, f"{meta_id}.opus")
    new_meta["local_filepath"] = local_filepath

    # Filter immediately after merging
    filtered_meta, reason, dur = filter_meta(new_meta)
    return filtered_meta, reason, dur, meta_id


def process_meta_chunk(
    meta_chunk: List[Dict],
    registry: MetadataRegistry,
    local_filepath_dir: str,
) -> Tuple[List[Tuple], set]:
    """Process a chunk of metas in parallel"""
    results = []
    chunk_ids_processed = set()

    for meta in meta_chunk:
        meta_id = meta["id"]
        # Skip duplicates within this chunk
        if meta_id in chunk_ids_processed:
            continue
        chunk_ids_processed.add(meta_id)

        filtered_meta, reason, dur, meta_id = merge_and_filter_meta(
            meta, registry, local_filepath_dir
        )

        results.append((filtered_meta, reason, dur, meta.get("duration_s", 0.0)))

    return results, chunk_ids_processed


# ============================================================================
# Main Script
# ============================================================================


def setup_metadata_registry(version: str, alignments_version: str) -> MetadataRegistry:
    """Setup and configure all metadata sources"""
    registry = MetadataRegistry()

    # Base directories
    metadata_base = "/app2/suno/data/christian/metadata/"
    codebase_metadata = "/home/christian/code/christian/metadata/"
    alignments_base = "/home/tony/Work/tony/hoot/tmp/"

    # Stem metadata
    registry.register(
        MetadataSource(
            name="stems",
            filepath=f"{metadata_base}stems_metadata_v9.json",
            loader=load_json_dict,
            output_key="stems",
        )
    )

    # Vox stem metadata
    registry.register(
        MetadataSource(
            name="vox_stem",
            filepath=f"{metadata_base}trimmed_vocals_stem_map_v9.json",
            loader=load_json_dict,
            output_key="vox_stem",
        )
    )

    # Stem captions
    registry.register(
        MetadataSource(
            name="stem_captions",
            filepath=f"{metadata_base}stem_captions_v6.json",
            loader=load_json_dict,
            merge_function=merge_stem_captions,
        )
    )

    # Audio production metadata
    registry.register(
        MetadataSource(
            name="audio_production",
            filepath=f"{codebase_metadata}v4/combined_audio_production_features_v2.csv",
            loader=load_csv_dict,
            output_key="audio_stats",
        )
    )

    # Alignments (with priority: genius > discogs > deezer)
    registry.register(
        MetadataSource(
            name="alignments_genius",
            filepath=f"{alignments_base}genius_hq_alignments_h5_t480_v{alignments_version}.jsonl",
            loader=load_alignment_map,
            merge_function=merge_priority_alignments,
            priority=1,  # Highest priority
        )
    )
    registry.register(
        MetadataSource(
            name="alignments_discogs",
            filepath=f"{alignments_base}discogs_hq_alignments_h5_t480_v{alignments_version}.jsonl",
            loader=load_alignment_map,
            merge_function=merge_priority_alignments,
            priority=2,
        )
    )
    registry.register(
        MetadataSource(
            name="alignments_deezer",
            filepath=f"{alignments_base}deezer_hq_alignments_h5_t480_v{alignments_version}.jsonl",
            loader=load_alignment_map,
            merge_function=merge_priority_alignments,
            priority=3,  # Lowest priority
        )
    )

    # Tags - discogs titles
    registry.register(
        MetadataSource(
            name="tags_discogs_titles",
            filepath=f"{codebase_metadata}organized/titles/discogs_subset_title_terms_map.json",
            loader=load_json_dict,
            merge_function=merge_tags,
        )
    )

    # Tags - genius titles
    registry.register(
        MetadataSource(
            name="tags_genius_titles",
            filepath=f"{codebase_metadata}organized/titles/genius_title_terms_map.json",
            loader=load_json_dict,
            merge_function=merge_tags,
        )
    )

    # Tags - GPT tags (with transform to extract gpt_tag)
    def transform_gpt_tags(data):
        return {id: tags["gpt_tag"] for id, tags in data.items()}

    registry.register(
        MetadataSource(
            name="tags_discogs_gpt",
            filepath=f"{codebase_metadata}organized/llm/gpt-4_1-tagging/gpt-4_1-tagging_results_discogs_subset.json",
            loader=load_json_dict,
            loader_kwargs={"transform": transform_gpt_tags},
            merge_function=merge_tags,
        )
    )
    registry.register(
        MetadataSource(
            name="tags_genius_gpt",
            filepath=f"{codebase_metadata}organized/llm/gpt-4_1-tagging/gpt-4_1-tagging_results_genius.json",
            loader=load_json_dict,
            loader_kwargs={"transform": transform_gpt_tags},
            merge_function=merge_tags,
        )
    )
    registry.register(
        MetadataSource(
            name="tags_imslp_gpt",
            filepath=f"{codebase_metadata}organized/llm/gpt-4_1-tagging/gpt-4_1-tagging_results_imslp.json",
            loader=load_json_dict,
            loader_kwargs={"transform": transform_gpt_tags},
            merge_function=merge_tags,
        )
    )

    # Chart metadata - merge into tags
    registry.register(
        MetadataSource(
            name="chart_metas",
            filepath=f"{codebase_metadata}organized/charts/discogs_subset_chart_metas.jsonl",
            loader=load_jsonl_dict,
            merge_function=merge_tags,
        )
    )

    # Grammy metadata - merge into tags
    registry.register(
        MetadataSource(
            name="grammy_metas",
            filepath=f"{codebase_metadata}organized/charts/discogs_subset_grammy_metas.jsonl",
            loader=load_jsonl_dict,
            merge_function=merge_tags,
        )
    )

    registry.register(
        MetadataSource(
            name="extra_metadata",
            filepath=f"{metadata_base}merged_metadata_map.json",
            loader=load_json_dict,
            merge_function=merge_all_keys,
        )
    )

    # add more metadata sources here
    # ...

    return registry


if __name__ == "__main__":

    version = "4"
    alignments_version = "5"
    local_filepath_dir = "/app2/suno/data/raw_audio_opus_v0"

    # Setup metadata registry
    registry = setup_metadata_registry(version, alignments_version)

    # Load all metadata
    registry.load_all()

    # ------------------------------------------------------------
    # Source metadata files
    # ------------------------------------------------------------
    print("\nLoading source metadata files...")
    metas_dir = "/app2/suno/data/christian/metadata/clean_metas/"

    source_metas_generators = [
        (
            "discogs_subset",
            read_jsonl_lazy(
                os.path.join(metas_dir, "clean_discogs_subset_v0_metas.jsonl")
            ),
        ),
        (
            "genius",
            read_jsonl_lazy(os.path.join(metas_dir, "clean_genius_v0_metas.jsonl")),
        ),
        (
            "imslp",
            read_jsonl_lazy(os.path.join(metas_dir, "clean_imslp_v0_metas.jsonl")),
        ),
        (
            "deezer",
            read_jsonl_lazy(os.path.join(metas_dir, "clean_deezer_v0_metas.jsonl")),
        ),
        (
            "pond5",
            read_jsonl_lazy(os.path.join(metas_dir, "clean_pond5_v0_metas.jsonl")),
        ),
        (
            "karaoke",
            read_jsonl_lazy(os.path.join(metas_dir, "clean_karaoke_v0_metas.jsonl")),
        ),
    ]

    # ------------------------------------------------------------
    # Merge and filter all metadata in parallel chunks, streaming to file
    # ------------------------------------------------------------
    print("\nMerging and filtering metadata in parallel (streaming)...")

    # Start timing
    start_time = time.time()

    # Track all the ids that we have merged (for deduplication across sources)
    all_ids_merged = set()

    # Statistics tracking
    total_duration_unfiltered = 0.0
    total_duration_filtered = 0.0
    total_unfiltered_count = 0

    # Per-dataset statistics
    dataset_stats = {}

    # Filtering statistics
    filter_counts = {
        "duration": 0,
        "audio_stats": 0,
        "loudness": 0,
        "silence": 0,
        "cer": 0,
        "local_filepath": 0,
    }

    # Stream filtered results directly to file instead of keeping in memory
    filtered_filepath = (
        f"/app2/suno/data/christian/metadata/filtered_metas_v{version}.jsonl"
    )
    filtered_count = 0

    # Configuration for parallel processing
    chunk_size = 2000  # Process metas in chunks of this size
    n_jobs = min(32, os.cpu_count() or 1)
    print(f"Using {n_jobs} parallel workers with chunk size {chunk_size}")

    # Process each source metas generator
    with open(filtered_filepath, "w", encoding="utf-8") as filtered_file:
        for source_name, metas_gen in source_metas_generators:
            print(f"\nProcessing {source_name} in parallel...")

            # Initialize stats for this dataset (must be before processing)
            dataset_stats[source_name] = {
                "unfiltered_count": 0,
                "filtered_count": 0,
                "unfiltered_duration": 0.0,
                "filtered_duration": 0.0,
            }

            # Collect chunks from generator and process in batches
            meta_chunk = []
            chunk_batch = []
            chunk_idx = 0
            batch_size = n_jobs * 2  # Process this many chunks in parallel at once

            def process_batch(batch_chunks):
                """Process a batch of chunks in parallel"""
                return Parallel(n_jobs=n_jobs, backend="threading")(
                    delayed(process_meta_chunk)(
                        chunk,
                        registry,
                        local_filepath_dir,
                    )
                    for chunk in batch_chunks
                )

            for meta in metas_gen:
                meta_chunk.append(meta)

                # When chunk is full, add to batch
                if len(meta_chunk) >= chunk_size:
                    chunk_batch.append(meta_chunk)
                    meta_chunk = []

                    # When batch is full, process it in parallel
                    if len(chunk_batch) >= batch_size:
                        chunk_results_list = process_batch(chunk_batch)

                        # Write results sequentially (thread-safe)
                        for chunk_results, _ in chunk_results_list:
                            for (
                                filtered_meta,
                                reason,
                                dur,
                                unfiltered_dur,
                            ) in chunk_results:
                                total_duration_unfiltered += unfiltered_dur
                                total_unfiltered_count += 1
                                dataset_stats[source_name]["unfiltered_count"] += 1
                                dataset_stats[source_name][
                                    "unfiltered_duration"
                                ] += unfiltered_dur

                                meta_id = filtered_meta["id"] if filtered_meta else None
                                if meta_id and meta_id in all_ids_merged:
                                    continue

                                if filtered_meta is not None:
                                    all_ids_merged.add(meta_id)
                                    filtered_file.write(
                                        json.dumps(filtered_meta) + "\n"
                                    )
                                    filtered_count += 1
                                    total_duration_filtered += dur
                                    dataset_stats[source_name]["filtered_count"] += 1
                                    dataset_stats[source_name][
                                        "filtered_duration"
                                    ] += dur
                                else:
                                    if reason and reason in filter_counts:
                                        filter_counts[reason] += 1

                        chunk_batch = []
                        chunk_idx += len(chunk_results_list)
                        if chunk_idx % (batch_size * 5) == 0:
                            print(
                                f"  Processed {chunk_idx * chunk_size:,} metas, {filtered_count:,} passed filter"
                            )

            # Process remaining chunk
            if meta_chunk:
                chunk_batch.append(meta_chunk)

            # Process remaining batch
            if chunk_batch:
                chunk_results_list = process_batch(chunk_batch)

                for chunk_results, _ in chunk_results_list:
                    for filtered_meta, reason, dur, unfiltered_dur in chunk_results:
                        total_duration_unfiltered += unfiltered_dur
                        total_unfiltered_count += 1
                        dataset_stats[source_name]["unfiltered_count"] += 1
                        dataset_stats[source_name][
                            "unfiltered_duration"
                        ] += unfiltered_dur

                        meta_id = filtered_meta["id"] if filtered_meta else None
                        if meta_id and meta_id in all_ids_merged:
                            continue

                        if filtered_meta is not None:
                            all_ids_merged.add(meta_id)
                            filtered_file.write(json.dumps(filtered_meta) + "\n")
                            filtered_count += 1
                            total_duration_filtered += dur
                            dataset_stats[source_name]["filtered_count"] += 1
                            dataset_stats[source_name]["filtered_duration"] += dur
                        else:
                            if reason and reason in filter_counts:
                                filter_counts[reason] += 1

            print(
                f"  Completed {source_name}: {filtered_count:,} filtered metas so far"
            )

    # End timing
    build_time = time.time() - start_time

    print("\n" + "=" * 80)
    print("BUILD SUMMARY")
    print("=" * 80)
    print(f"\nTotal build time: {build_time:.2f} seconds ({build_time/60:.2f} minutes)")

    print("\nPer-Dataset Statistics:")
    print("-" * 80)
    for dataset_name, stats in sorted(dataset_stats.items()):
        unfiltered_hrs = stats["unfiltered_duration"] / 3600
        filtered_hrs = stats["filtered_duration"] / 3600
        filter_pct = (
            (stats["filtered_duration"] / stats["unfiltered_duration"] * 100)
            if stats["unfiltered_duration"] > 0
            else 0
        )
        print(f"  {dataset_name:20s}:")
        print(
            f"    Unfiltered: {stats['unfiltered_count']:8,} entries ({unfiltered_hrs:8.2f} hours)"
        )
        print(
            f"    Filtered:   {stats['filtered_count']:8,} entries ({filtered_hrs:8.2f} hours, {filter_pct:5.1f}%)"
        )

    print("\nOverall Statistics:")
    print("-" * 80)
    print(f"Total unfiltered entries: {total_unfiltered_count:,}")
    print(f"Total filtered entries: {filtered_count:,}")
    print(f"Total unfiltered duration: {total_duration_unfiltered/3600:.2f} hours")
    print(f"Total filtered duration: {total_duration_filtered/3600:.2f} hours")
    print(
        f"Filtered percentage: {(total_duration_filtered/total_duration_unfiltered*100):.1f}%"
    )

    print("\nFiltering breakdown:")
    print("-" * 80)
    for reason, count in filter_counts.items():
        print(f"  {reason.capitalize():20s} filtered: {count:8,} entries")

    print(f"\nWrote {filtered_count:,} filtered metas to {filtered_filepath}")
    print("=" * 80)

    # ------------------------------------------------------------
    # Shuffle and split (read from file, shuffle, write splits)
    # ------------------------------------------------------------
    print(f"\nShuffling and splitting {filtered_count:,} filtered metas...")

    # Read all filtered metas for shuffling (we need them all for proper shuffle)
    # This is the only time we load all filtered metas into memory
    filtered_metas = []
    for meta in read_jsonl_lazy(filtered_filepath):
        filtered_metas.append(meta)

    # Set random seed for reproducibility
    random.seed(42)
    random.shuffle(filtered_metas)

    # Calculate split sizes (1% for validation)
    val_size = int(len(filtered_metas) * 0.01)
    train_size = len(filtered_metas) - val_size

    print(f"Train size: {train_size:,} ({train_size/len(filtered_metas)*100:.1f}%)")
    print(f"Validation size: {val_size:,} ({val_size/len(filtered_metas)*100:.1f}%)")

    # Split the data
    train_metas = filtered_metas[:train_size]
    val_metas = filtered_metas[train_size:]

    # Free memory
    del filtered_metas

    # Calculate durations for each split
    train_duration = sum(meta["duration_s"] for meta in train_metas)
    val_duration = sum(meta["duration_s"] for meta in val_metas)
    total_duration = train_duration + val_duration

    print("\nDuration Summary:")
    print(
        f"Train duration: {train_duration/3600:.2f} hours ({train_duration/total_duration*100:.1f}%)"
    )
    print(
        f"Validation duration: {val_duration/3600:.2f} hours ({val_duration/total_duration*100:.1f}%)"
    )

    # Define output directory and filenames
    output_dir = "/app2/suno/data/diffusion/v1/"
    train_filename = f"metas_diff_v{version}_tr.jsonl"
    val_filename = f"metas_diff_v{version}_val.jsonl"

    # Ensure output directory exists
    os.makedirs(output_dir, exist_ok=True)

    # Write train and validation files
    train_filepath = os.path.join(output_dir, train_filename)
    val_filepath = os.path.join(output_dir, val_filename)

    print(f"\nWriting train metas to: {train_filepath}")
    write_jsonl(train_metas, train_filepath)

    print(f"Writing validation metas to: {val_filepath}")
    write_jsonl(val_metas, val_filepath)

    print("\n✅ Successfully split and saved:")
    print(f"  Train: {len(train_metas):,} entries ({train_duration/3600:.2f} hours)")
    print(f"  Validation: {len(val_metas):,} entries ({val_duration/3600:.2f} hours)")
    print(f"  Total: {filtered_count:,} entries ({total_duration/3600:.2f} hours)")
