import unittest
import random
import sys
import os
import numpy as np

sys.path.insert(0, os.path.join(os.path.dirname(__file__), ".."))

from text_utils import get_control_tags


class TestTextUtils(unittest.TestCase):
    def test_get_control_tags(self):
        random.seed(42)
        assert (
            get_control_tags(121, 121 * 25, sample_vocal_start_s=None, do_augment=False)
            == "{max_duration:410;duration_toks:3025;min_duration:100;duration:121}"
        )
        random.seed(42)
        assert (
            get_control_tags(120.4, int(120.4 * 25), sample_vocal_start_s=None, do_augment=False)
            == "{max_duration:410;duration_toks:3010;min_duration:100;duration:120}"
        )
        random.seed(42)
        assert (
            get_control_tags(120.6, int(120.6 * 25), sample_vocal_start_s=None, do_augment=False)
            == "{max_duration:410;duration_toks:3015;min_duration:100;duration:121}"
        )
        random.seed(42)
        assert (
            get_control_tags(121.1, int(121.1 * 25), sample_vocal_start_s=5, do_augment=False)
            == "{max_duration:410;vocals:early;vocals:normal;duration_toks:3027;min_duration:100;duration:121}"
        )
        random.seed(42)
        assert (
            get_control_tags(121.1, int(121.1 * 25), sample_vocal_start_s=15, do_augment=False)
            == "{max_duration:410;vocals:normal;vocals:intro;duration_toks:3027;min_duration:100;duration:121}"
        )
        random.seed(42)
        assert (
            get_control_tags(121.1, int(121.1 * 25), sample_vocal_start_s=20, do_augment=False)
            == "{duration_toks:3027;max_duration:410;vocals:intro;min_duration:100;duration:121}"
        )
        random.seed(42)
        assert (
            get_control_tags(121.1, int(121.1 * 25), sample_vocal_start_s=50, do_augment=False)
            == "{max_duration:410;duration_toks:3027;min_duration:100;duration:121}"
        )
        random.seed(42)
        assert (
            get_control_tags(121.1, int(121.1 * 25), sample_vocal_start_s=None, do_augment=True)
            == "{duration_toks:3027;max_duration:410;duration:121;min_duration:100}"
        )
        random.seed(42)
        assert (
            get_control_tags(180.1, int(180.1 * 25), sample_vocal_start_s=None, do_augment=True) == None
        )

    def test_get_control_tags_with_spectral_features(self):
        """Test control tags include spectral features in correct format."""
        random.seed(42)
        # Create test spectral features
        centroid_seq = np.array([0.5, 0.6, 0.7])  # Normalized [0,1]
        complexity_seq = np.array([0.3, 0.4, 0.5])  # Normalized [0,1]

        result = get_control_tags(
            sample_duration_s=3.0,
            sample_duration_toks=75,
            do_augment=False,
            spectral_centroid_seq=centroid_seq,
            spectral_complexity_seq=complexity_seq,
        )

        # Should contain spectral tags in 0-100 format with new names
        self.assertIn("spectral_centroid_contour:[50,60,70]", result)
        self.assertIn("spectral_complexity_contour:[30,40,50]", result)
