#!/usr/bin/env python3
"""Test script for merge_preference_datasets.

This script creates minimal test datasets and validates the merge operation.
"""
import os
import tempfile
import shutil
import numpy as np
from typing import Dict, Any, List
import json

from merge_preference_datasets import (
    merge_preference_datasets,
    load_sample_from_mmap,
    read_jsonl,
    read_json,
    write_jsonl,
    write_json,
    SEMANTIC_N_CODEBOOKS,
    N_TOKENS_AUDIO,
)


def create_test_dataset(
    output_dir: str,
    n_samples: int,
    is_val: bool = False,
    t_data_memmap: int = 100,  # Smaller for testing
    start_id: int = 0,
) -> None:
    """Create a minimal test dataset.
    
    Args:
        output_dir: Directory to create the test dataset
        n_samples: Number of samples to create
        is_val: Whether this is validation set
        t_data_memmap: Number of tokens per sample (smaller for tests)
        start_id: Starting ID for sample naming
    """
    os.makedirs(output_dir, exist_ok=True)
    
    dset_type = "val" if is_val else "tr"
    mmap_path = os.path.join(output_dir, f"data_{dset_type}.bin")
    meta_path = os.path.join(output_dir, f"meta_{dset_type}.jsonl")
    info_path = os.path.join(output_dir, f"info_{dset_type}.json")
    
    # Create mmap
    sample_size = t_data_memmap * SEMANTIC_N_CODEBOOKS
    total_size = n_samples * sample_size
    mmap = np.memmap(mmap_path, dtype=np.uint16, mode='w+', shape=(total_size,))
    
    # Create metadata and info
    metadata: List[Dict[str, Any]] = []
    info: Dict[str, Dict[str, Any]] = {
        "perference_0": {"idx_list": []},
        "perference_1": {"idx_list": []},
    }
    
    for i in range(n_samples):
        # Create sample data with unique pattern
        sample_data = np.full(
            (t_data_memmap, SEMANTIC_N_CODEBOOKS),
            fill_value=i + start_id,
            dtype=np.uint16
        )
        
        # Write to mmap
        offset = i * sample_size
        mmap[offset:offset + sample_size] = sample_data.flatten()
        
        # Create metadata
        preference = i % 2
        meta = {
            "dataset": f"perference_{preference}",
            "id": f"test_sample_{i + start_id}",
            "start_s": 0.0,
            "vocal_start_s": None,
            "vocal_end_s": None,
            "tags": ["test"],
            "neg_tags": "",
            "control_tags": "",
            "gender": "male" if i % 2 == 0 else "female",
            "control": {},
            "text": f"Test sample {i + start_id}",
            "generated_start_index": 0,
            "user_id": f"user_{i % 3}",
        }
        metadata.append(meta)
        info[f"perference_{preference}"]["idx_list"].append(i)
    
    mmap.flush()
    del mmap
    
    # Write metadata and info
    write_jsonl(metadata, meta_path)
    write_json(info, info_path)
    
    print(f"Created test dataset at {output_dir} with {n_samples} samples")


def test_merge() -> None:
    """Test the merge_preference_datasets function."""
    print("=" * 70)
    print("TESTING MERGE_PREFERENCE_DATASETS")
    print("=" * 70)
    
    # Create temporary directories
    temp_dir = tempfile.mkdtemp(prefix="test_merge_")
    
    try:
        # Test parameters
        t_data_memmap = 100  # Small for fast testing
        n_samples1 = 20
        n_samples2 = 30
        
        input_dir1 = os.path.join(temp_dir, "dataset1")
        input_dir2 = os.path.join(temp_dir, "dataset2")
        output_dir = os.path.join(temp_dir, "merged")
        
        print(f"\nTest parameters:")
        print(f"  t_data_memmap: {t_data_memmap}")
        print(f"  Dataset 1 samples: {n_samples1}")
        print(f"  Dataset 2 samples: {n_samples2}")
        print(f"  Temp directory: {temp_dir}")
        
        # Create test datasets
        print("\n" + "=" * 70)
        print("Creating test datasets...")
        print("=" * 70)
        create_test_dataset(input_dir1, n_samples1, is_val=True, t_data_memmap=t_data_memmap, start_id=0)
        create_test_dataset(input_dir2, n_samples2, is_val=True, t_data_memmap=t_data_memmap, start_id=n_samples1)
        
        # Merge datasets
        print("\n" + "=" * 70)
        print("Merging datasets...")
        print("=" * 70)
        merge_preference_datasets(
            input_dir1=input_dir1,
            input_dir2=input_dir2,
            output_dir=output_dir,
            is_val=True,
            t_data_memmap=t_data_memmap,
            validate=True,
        )
        
        # Verify merged dataset
        print("\n" + "=" * 70)
        print("Verifying merged dataset...")
        print("=" * 70)
        
        merged_mmap_path = os.path.join(output_dir, "data_val.bin")
        merged_meta_path = os.path.join(output_dir, "meta_val.jsonl")
        merged_info_path = os.path.join(output_dir, "info_val.json")
        
        # Load merged data
        merged_meta = read_jsonl(merged_meta_path)
        merged_info = read_json(merged_info_path)
        
        # Check counts
        assert len(merged_meta) == n_samples1 + n_samples2, \
            f"Metadata count mismatch: {len(merged_meta)} != {n_samples1 + n_samples2}"
        print(f"✓ Metadata count correct: {len(merged_meta)}")
        
        # Check info indices
        total_indices = sum(len(v["idx_list"]) for v in merged_info.values())
        assert total_indices == n_samples1 + n_samples2, \
            f"Info indices count mismatch: {total_indices} != {n_samples1 + n_samples2}"
        print(f"✓ Info indices count correct: {total_indices}")
        
        # Verify samples from dataset 1
        print("\nVerifying samples from dataset 1...")
        for i in range(min(5, n_samples1)):
            sample = load_sample_from_mmap(merged_mmap_path, i, t_data_memmap)
            expected_value = i
            actual_value = sample[0, 0]
            assert actual_value == expected_value, \
                f"Dataset 1 sample {i} mismatch: {actual_value} != {expected_value}"
            
            meta = merged_meta[i]
            assert meta["id"] == f"test_sample_{i}", \
                f"Dataset 1 metadata {i} ID mismatch"
        print(f"✓ Dataset 1 samples verified (checked {min(5, n_samples1)} samples)")
        
        # Verify samples from dataset 2
        print("\nVerifying samples from dataset 2...")
        for i in range(min(5, n_samples2)):
            merged_idx = n_samples1 + i
            sample = load_sample_from_mmap(merged_mmap_path, merged_idx, t_data_memmap)
            expected_value = n_samples1 + i
            actual_value = sample[0, 0]
            assert actual_value == expected_value, \
                f"Dataset 2 sample {i} mismatch: {actual_value} != {expected_value}"
            
            meta = merged_meta[merged_idx]
            assert meta["id"] == f"test_sample_{n_samples1 + i}", \
                f"Dataset 2 metadata {i} ID mismatch"
        print(f"✓ Dataset 2 samples verified (checked {min(5, n_samples2)} samples)")
        
        # Verify info indices are correctly offset
        print("\nVerifying info indices...")
        for dataset_name, dataset_info in merged_info.items():
            idx_list = dataset_info["idx_list"]
            for idx in idx_list:
                meta = merged_meta[idx]
                assert meta["dataset"] == dataset_name, \
                    f"Info index {idx} points to wrong dataset: {meta['dataset']} != {dataset_name}"
        print("✓ All info indices point to correct metadata")
        
        print("\n" + "=" * 70)
        print("✅ ALL TESTS PASSED!")
        print("=" * 70)
        
    finally:
        # Clean up
        if os.path.exists(temp_dir):
            shutil.rmtree(temp_dir)
            print(f"\nCleaned up temp directory: {temp_dir}")


if __name__ == "__main__":
    test_merge()

