#!/usr/bin/env python3
"""
Extract validation metadata from train metas filtered by keepset IDs.

This script:
1. Loads the keepset JSON file containing valid audio IDs
2. Filters the train metadata JSONL to keep only entries with IDs in the keepset
3. Outputs a validation JSONL file with the filtered metadata
"""

import json
import os
from pathlib import Path
from typing import Set

try:
    from tqdm import tqdm
except ImportError:
    # Fallback if tqdm is not available
    def tqdm(iterable, *args, **kwargs):
        return iterable


def load_keepset_ids(keepset_path: str) -> Set[str]:
    """
    Load audio IDs from the keepset JSON file.

    Args:
        keepset_path: Path to the keepset JSON file (e.g., ids_keep_sets_ext_v11_stems.json)

    Returns:
        Set of audio IDs to keep
    """
    print(f"Loading keepset from: {keepset_path}")

    with open(keepset_path, "r") as f:
        keepset_data = json.load(f)

    # Extract audio_ids - the keepset file has structure like {"audio_ids": [...], ...}
    if isinstance(keepset_data, dict):
        if "audio_ids" in keepset_data:
            audio_ids = set(keepset_data["audio_ids"])
        else:
            # Try to find any list in the dict
            for key, val in keepset_data.items():
                if isinstance(val, list):
                    audio_ids = set(val)
                    print(f"Using '{key}' field with {len(audio_ids):,} IDs")
                    break
            else:
                raise ValueError("No list of audio IDs found in keepset file")
    elif isinstance(keepset_data, list):
        audio_ids = set(keepset_data)
    else:
        raise ValueError(f"Unexpected keepset data type: {type(keepset_data)}")

    print(f"Loaded {len(audio_ids):,} audio IDs from keepset")
    return audio_ids


def extract_val_meta(
    train_metas_path: str,
    keepset_path: str,
    output_path: str,
) -> None:
    """
    Extract validation metadata from train metas filtered by keepset IDs.

    Args:
        train_metas_path: Path to the train metadata JSONL file
        keepset_path: Path to the keepset JSON file
        output_path: Path to write the validation metadata JSONL file
    """
    # Load keepset IDs
    valid_ids = load_keepset_ids(keepset_path)

    # Count total lines for progress bar
    print(f"Counting lines in: {train_metas_path}")
    with open(train_metas_path, "r") as f:
        total_lines = sum(1 for _ in f)
    print(f"Total lines: {total_lines:,}")

    # Filter and write validation metadata
    print(f"Filtering metadata and writing to: {output_path}")

    kept_count = 0
    skipped_count = 0

    # Create output directory if needed
    os.makedirs(os.path.dirname(output_path), exist_ok=True)

    with open(train_metas_path, "r") as infile, open(output_path, "w") as outfile:
        for line in tqdm(infile, total=total_lines, desc="Processing metadata"):
            line = line.strip()
            if not line:
                continue

            try:
                meta = json.loads(line)
                meta_id = meta.get("id")

                if meta_id in valid_ids:
                    outfile.write(json.dumps(meta) + "\n")
                    kept_count += 1
                else:
                    skipped_count += 1

            except json.JSONDecodeError as e:
                print(f"Warning: Skipping malformed JSON line: {e}")
                continue

    # Print summary
    print("\n" + "=" * 60)
    print("EXTRACTION SUMMARY")
    print("=" * 60)
    print(f"Total lines processed: {total_lines:,}")
    print(f"Metadata entries kept: {kept_count:,}")
    print(f"Metadata entries skipped: {skipped_count:,}")
    print(f"Keep ratio: {kept_count / total_lines:.2%}")
    print(f"\nOutput file: {output_path}")
    print("=" * 60)


def main():
    """Main function to extract validation metadata."""
    # Configuration
    data_dir = "/app2/suno/data/auk_v0"
    train_metas_filename = "metas_v5_tr.jsonl"
    train_info_filename = "ids_keep_sets_ext_v11_stems.json"

    # Output path - save in RealGen directory since we don't have write access to data_dir
    output_dir = "/home/tony/Work/tony/RealGen"
    output_filename = "metas_v5_val_filtered.jsonl"

    train_metas_path = os.path.join(data_dir, train_metas_filename)
    keepset_path = os.path.join(data_dir, train_info_filename)
    output_path = os.path.join(output_dir, output_filename)

    print("🎵 Validation Metadata Extraction Script")
    print(f"Data directory: {data_dir}")
    print(f"Train metas: {train_metas_filename}")
    print(f"Keepset file: {train_info_filename}")
    print(f"Output file: {output_filename}")
    print()

    # Check if input files exist
    if not os.path.exists(train_metas_path):
        raise FileNotFoundError(f"Train metas file not found: {train_metas_path}")
    if not os.path.exists(keepset_path):
        raise FileNotFoundError(f"Keepset file not found: {keepset_path}")

    # Extract validation metadata
    extract_val_meta(train_metas_path, keepset_path, output_path)

    print("\n✅ Extraction completed successfully!")


if __name__ == "__main__":
    main()
