import unittest
import numpy as np
import random
import sys
import os
from unittest.mock import Mock, patch, MagicMock

# Add parent directory to path for imports
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from data_utils import song_to_samples, get_samples_for_song
from data_types import SamplingParams, SampleData
from block_types import SampleBlockType
from modules.gpt import GPTConfig


class TestToMusic(unittest.TestCase):
    def setUp(self):
        """Set up test fixtures with realistic data structures."""
        random.seed(42)
        np.random.seed(42)

        # Mock configuration
        self.cfg = Mock()
        self.cfg.semantic_rate_hz = 25
        self.cfg.semantic_codebook_size = 4096

        # Mock song data and metadata
        self.song_data = np.random.randint(0, 4096, (1, 1000))  # 40 seconds at 25Hz
        self.song_meta = {
            "duration": 40.0,
            "sample_idx": 123,
            "artist": "Test Artist",
            "tags": "rock, guitar",
        }

        # Mock sampling parameters
        self.sampling_params = SamplingParams(allow_sample=True, prob_sample=0.1)

        # Mock sample data with required parameters
        self.sample_data = SampleData(
            data_row=self.song_data,
            data_meta=self.song_meta,
            sampling_params=self.sampling_params,
            audio_sample_tracks=[np.random.randint(0, 4096, (1, 250))],  # 10 seconds at 25Hz
        )

    def test_song_to_samples_basic(self):
        """Test basic song_to_samples functionality."""
        samples = song_to_samples(
            self.song_data, self.song_meta, self.cfg, n_samples=2, sample_duration_range=(5, 15)
        )

        self.assertEqual(len(samples), 2)
        for sample in samples:
            self.assertIsInstance(sample, dict)
            self.assertIn("data", sample)
            self.assertIn("meta", sample)
            self.assertIsInstance(sample["data"], np.ndarray)
            self.assertEqual(sample["data"].shape[0], 1)  # Single codebook

            # Check sample duration is within range (5-15 seconds at 25Hz)
            duration_tokens = sample["data"].shape[1]
            duration_seconds = duration_tokens / self.cfg.semantic_rate_hz
            self.assertGreaterEqual(duration_seconds, 5)
            self.assertLessEqual(duration_seconds, 15)

    def test_song_to_samples_edge_cases(self):
        """Test song_to_samples with edge cases."""
        # Test with very short song
        short_song = np.random.randint(0, 4096, (1, 100))  # 4 seconds
        short_meta = {"duration": 4.0, "sample_idx": 456}

        samples = song_to_samples(short_song, short_meta, self.cfg, n_samples=1)
        # Short songs might return no samples if they're too short for the minimum duration
        self.assertGreaterEqual(len(samples), 0)

        if len(samples) > 0:
            # Sample should be truncated to song length
            sample_duration = samples[0]["data"].shape[1] / self.cfg.semantic_rate_hz
            self.assertLessEqual(sample_duration, 4.0)

    def test_song_to_samples_no_sample_idx(self):
        """Test song_to_samples when no sample_idx is provided."""
        meta_no_idx = {k: v for k, v in self.song_meta.items() if k != "sample_idx"}

        samples = song_to_samples(self.song_data, meta_no_idx, self.cfg)
        self.assertEqual(len(samples), 1)

        # Should still create a sample, just without sample_idx reference
        self.assertIsInstance(samples[0]["data"], np.ndarray)

    def test_get_samples_for_song(self):
        """Test get_samples_for_song helper function."""
        # Mock data and metas
        data = [self.song_data, np.random.randint(0, 4096, (1, 500))]
        metas = [self.song_meta, {"duration": 20.0, "sample_idx": 789}]

        samples = get_samples_for_song(0, data, metas, self.cfg)

        self.assertIsInstance(samples, list)
        self.assertGreater(len(samples), 0)
        for sample in samples:
            self.assertIn("data", sample)
            self.assertIn("meta", sample)

    def test_sample_block_type(self):
        """Test SampleBlockType properties."""
        self.assertEqual(SampleBlockType.name, "sample")
        self.assertTrue(SampleBlockType.is_causal)

    def test_gpt_config_semantic_sample_token(self):
        """Test GPTConfig semantic_sample_token initialization."""
        # Use realistic values that match the default config
        config = GPTConfig(semantic_codebook_size=4000, semantic_vocab_size=4032)

        # Should auto-initialize to codebook_size + 12
        expected_token = 4000 + 12
        self.assertEqual(config.semantic_sample_token, expected_token)

        # Test explicit setting
        config_explicit = GPTConfig(
            semantic_codebook_size=4000, semantic_vocab_size=4032, semantic_sample_token=4020
        )
        self.assertEqual(config_explicit.semantic_sample_token, 4020)

    def test_sampling_params_defaults(self):
        """Test SamplingParams default values."""
        params = SamplingParams()
        self.assertFalse(params.allow_sample)
        self.assertEqual(params.prob_sample, 0.1)

    def test_sampling_params_custom(self):
        """Test SamplingParams with custom values."""
        params = SamplingParams(allow_sample=True, prob_sample=0.2)
        self.assertTrue(params.allow_sample)
        self.assertEqual(params.prob_sample, 0.2)

    def test_sample_data_structure(self):
        """Test SampleData structure and properties."""
        sample_data = SampleData(
            data_row=self.song_data,
            data_meta=self.song_meta,
            sampling_params=self.sampling_params,
            audio_sample_tracks=self.sample_data.audio_sample_tracks,
        )
        self.assertIsNotNone(sample_data.audio_sample_tracks)
        self.assertEqual(sample_data.audio_sample_tracks[0].shape, (1, 250))

    def test_sample_duration_validation(self):
        """Test sample duration constraints."""
        # Test various sample durations
        durations = [5, 10, 15, 20]  # seconds

        for duration in durations:
            sample_tokens = int(duration * self.cfg.semantic_rate_hz)
            sample_data = np.random.randint(0, 4096, (1, sample_tokens))

            # Verify token count matches expected duration
            calculated_duration = sample_data.shape[1] / self.cfg.semantic_rate_hz
            self.assertAlmostEqual(calculated_duration, duration, places=1)

    def test_semantic_sample_token_validation(self):
        """Test semantic_sample_token is within vocabulary bounds."""
        config = GPTConfig(semantic_codebook_size=4000, semantic_vocab_size=4032)

        # Token should be less than vocab size
        self.assertLess(config.semantic_sample_token, config.semantic_vocab_size)

        # Should be greater than codebook size (reserved range)
        self.assertGreater(config.semantic_sample_token, config.semantic_codebook_size)

    def test_backward_compatibility(self):
        """Test that new features don't break existing functionality."""
        # Test SamplingParams with allow_sample=False (default)
        params = SamplingParams(allow_sample=False)
        self.assertFalse(params.allow_sample)

        # Test SampleData without audio_sample_tracks (but with required params)
        sample_data = SampleData(
            data_row=self.song_data, data_meta=self.song_meta, sampling_params=params
        )
        self.assertIsNone(sample_data.audio_sample_tracks)

        # These should not cause errors in existing code paths

    def test_sample_metadata_preservation(self):
        """Test that sample metadata is correctly preserved."""
        samples = song_to_samples(self.song_data, self.song_meta, self.cfg, n_samples=1)

        sample = samples[0]
        self.assertIn("meta", sample)

        # Original metadata should be preserved in sample
        original_keys = ["duration", "sample_idx", "artist", "tags"]
        for key in original_keys:
            if key in self.song_meta:
                # Sample meta should reference or contain original info
                self.assertIsInstance(sample["meta"], dict)

    def test_integration_sample_to_song_workflow(self):
        """Test the complete sample-to-song workflow integration."""
        # Step 1: Create samples from song
        samples = song_to_samples(self.song_data, self.song_meta, self.cfg, n_samples=2)

        # Step 2: Verify samples can be used for conditioning
        for sample in samples:
            self.assertIsInstance(sample["data"], np.ndarray)
            self.assertEqual(sample["data"].shape[0], 1)  # Single codebook

            # Sample should be suitable for SampleData
            sample_data = SampleData(
                data_row=self.song_data,
                data_meta=self.song_meta,
                sampling_params=self.sampling_params,
                audio_sample_tracks=[sample["data"]],
            )
            self.assertIsNotNone(sample_data.audio_sample_tracks)

            # Verify duration constraints
            duration = sample["data"].shape[1] / self.cfg.semantic_rate_hz
            self.assertGreaterEqual(duration, 5)
            self.assertLessEqual(duration, 15)

    def test_error_handling(self):
        """Test error handling in sample processing."""
        # Test with missing required config attributes
        incomplete_cfg = Mock()
        incomplete_cfg.semantic_rate_hz = "invalid"  # String instead of int
        # This should cause a TypeError when multiplying
        with self.assertRaises((TypeError, AttributeError)):
            song_to_samples(self.song_data, self.song_meta, incomplete_cfg)

        # Test with None data - this should handle gracefully or raise appropriate error
        try:
            result = song_to_samples(None, self.song_meta, self.cfg)
            # If it doesn't raise an error, it should return empty list or similar
            self.assertIsInstance(result, list)
        except (ValueError, TypeError, AttributeError):
            # These are acceptable error types for None input
            pass

    def test_sample_timing_text_functionality(self):
        """Test that sample timing is correctly added to text conditioning."""
        from text_utils import build_text

        # Test case: sample at 1:43.4 (103.4 seconds)
        result = build_text(
            tags=["pop", "energetic"],
            text="Test lyrics here",
            sample_duration_s=180,
            sample_duration_toks=4500,
            inference=True,  # Use inference mode for consistent output
            audio_sample_start_times_s=[103.4],
        )

        # Should contain the timing tag in control tags format (seconds)
        self.assertIn("audio_sample_time_0:103", result)
        # Should also contain token-based timing
        self.assertIn("audio_sample_start_toks_0:2585", result)  # 103.4 * 25 ≈ 2585

        # Test without sample timing
        result_no_timing = build_text(
            tags=["test"],
            text="Test lyrics",
            sample_duration_s=180,
            sample_duration_toks=4500,
            inference=True,
            audio_sample_start_times_s=None,
        )

        # Should not contain audio sample timing tag
        self.assertNotIn("audio_sample_time_", result_no_timing)

        # Test different timing values with zero-based indexing
        test_cases = [
            (12.1, "audio_sample_time_0:12"),
            (176.6, "audio_sample_time_0:177"),
            (31.2, "audio_sample_time_0:31"),
            (3661.5, "audio_sample_time_0:3662"),  # Over 1 hour
        ]

        for start_time, expected_tag in test_cases:
            result = build_text(
                tags=["test"],
                text="lyrics",
                sample_duration_s=30,
                sample_duration_toks=750,
                inference=True,
                audio_sample_start_times_s=[start_time],
            )
            self.assertIn(expected_tag, result, f"Failed for timing {start_time}s")

    def test_multiple_audio_sample_timing(self):
        """Test multiple audio sample timings in control tags."""
        from text_utils import build_text

        # Test case: multiple samples at different times
        result = build_text(
            tags=["electronic", "energetic"],
            text="Multiple samples test",
            sample_duration_s=180,
            sample_duration_toks=4500,
            inference=True,
            audio_sample_start_times_s=[12.3, 67.8, 134.5],
        )

        # Should contain multiple timing tags in seconds format with zero-based indexing
        self.assertIn("audio_sample_time_0:12", result)
        self.assertIn("audio_sample_time_1:68", result)
        self.assertIn("audio_sample_time_2:134", result)  # 134.5 rounds to 134

        # Should also contain token-based timing (25 Hz semantic rate)
        self.assertIn("audio_sample_start_toks_0:308", result)  # 12.3 * 25 ≈ 308
        self.assertIn("audio_sample_start_toks_1:1695", result)  # 67.8 * 25 ≈ 1695
        self.assertIn("audio_sample_start_toks_2:3362", result)  # 134.5 * 25 ≈ 3362

    def test_randomized_sample_count_range(self):
        """Test that max_num_audio_samples creates variable sample counts."""
        import random
        from data_types import SamplingParams

        # Test the randomization logic
        max_samples = 4
        counts = []
        random.seed(42)  # For reproducible test

        for _ in range(20):  # Test multiple iterations
            count = random.randint(1, max_samples)
            counts.append(count)

        # Should have variety in counts
        unique_counts = set(counts)
        self.assertGreater(len(unique_counts), 1, "Should generate different sample counts")
        self.assertGreaterEqual(min(counts), 1, "Minimum should be 1")
        self.assertLessEqual(max(counts), max_samples, f"Maximum should be {max_samples}")

        # Test SamplingParams structure
        params = SamplingParams(max_num_audio_samples=10)
        self.assertEqual(params.max_num_audio_samples, 10)

    def test_audio_sample_source_control_tags(self):
        """Test audio sample source control tags generation."""
        from text_utils import build_text

        # Test single sample with source
        result_single = build_text(
            tags=["electronic"],
            text="Single source test",
            sample_duration_s=30,
            sample_duration_toks=750,
            inference=True,
            audio_sample_start_times_s=[45.2],
            audio_sample_sources=["vocal"],
        )

        # Should contain timing and source tags with zero-based indexing
        self.assertIn("audio_sample_time_0:45", result_single)
        self.assertIn("audio_sample_start_toks_0:1130", result_single)  # 45.2 * 25 ≈ 1130
        self.assertIn("audio_sample_vocal_0", result_single)

        # Test multiple samples with different sources
        result_multiple = build_text(
            tags=["rock"],
            text="Multiple sources test",
            sample_duration_s=120,
            sample_duration_toks=3000,
            inference=True,
            audio_sample_start_times_s=[12.1, 67.8, 134.5],
            audio_sample_sources=["vocal", "drum", "full_mix"],
        )

        # Should contain timing tags
        self.assertIn("audio_sample_time_0:12", result_multiple)
        self.assertIn("audio_sample_time_1:68", result_multiple)
        self.assertIn("audio_sample_time_2:134", result_multiple)  # 134.5 rounds to 134

        # Should contain source tags
        self.assertIn("audio_sample_vocal_0", result_multiple)
        self.assertIn("audio_sample_drum_1", result_multiple)
        self.assertIn("audio_sample_full_mix_2", result_multiple)

        # Test without sources (backward compatibility)
        result_no_sources = build_text(
            tags=["jazz"],
            text="No sources test",
            sample_duration_s=60,
            sample_duration_toks=1500,
            inference=True,
            audio_sample_start_times_s=[23.4],
            audio_sample_sources=None,
        )

        # Should contain timing but no source tags
        self.assertIn("audio_sample_time_0:23", result_no_sources)
        self.assertNotIn("audio_sample_vocal", result_no_sources)
        self.assertNotIn("audio_sample_drum", result_no_sources)
        self.assertNotIn("audio_sample_full_mix", result_no_sources)


if __name__ == "__main__":
    unittest.main()
