import random
from torch.utils.data import Dataset, Sampler
from suno_utils.utils.text import read_json, read_jsonl
from ditto_v2.data_loaders.self_sim import SelfSimDataset
from ditto_v2.data_loaders.self_vox_sim import SelfVoxSimDataset
from ditto_v2.data_loaders.self_lyric_sim import SelfLyricSimDataset
from ditto_v2.data_loaders.artist_sim import ArtistSimDataset
from ditto_v2.data_loaders.artist_vox_sim import ArtistVoxSimDataset
from ditto_v2.data_loaders.album_sim import AlbumSimDataset
from ditto_v2.data_loaders.genre_sim import GenreSimDataset
from ditto_v2.data_loaders.lyric_sim import LyricSimDataset


class MultiTaskDataset(Dataset):
    def __init__(
        self,
        split="train",
        self_sim_input_length_s=15.0,
        self_vox_sim_input_length_s=15.0,
        self_lyric_sim_input_length=300,
        artist_sim_input_length_s=15.0,
        artist_vox_sim_input_length_s=15.0,
        album_sim_input_length_s=15.0,
        genre_sim_input_length_s=15.0,
        genre_tag_input_length=300,
        lyric_sim_input_length_s=15.0,
        lyric_sim_input_text_length=300,
        sample_rate=24000,
        num_samples=1000,
        num_self_sim_samples=-1,
        num_vox_sim_samples=-1,
        num_self_lyric_sim_samples=-1,
        num_artist_sim_samples=-1,
        num_artist_vox_sim_samples=-1,
        num_album_sim_samples=-1,
        num_genre_sim_samples=-1,
        num_lyric_sim_samples=-1,
    ):
        assert split in ["train", "valid"]

        # read data
        metadata = read_jsonl("/app/suno/data/v2_audio/metadata/metas_%s.jsonl" % split)
        task_indices = read_json("/app/suno/data/v2_audio/metadata/task_indices.json")

        self.self_sim_dataset = SelfSimDataset(
            metadata,
            task_indices,
            split,
            self_sim_input_length_s,
            sample_rate,
            num_self_sim_samples,
        )
        self.self_vox_sim_dataset = SelfVoxSimDataset(
            metadata,
            task_indices,
            split,
            self_vox_sim_input_length_s,
            sample_rate,
            num_vox_sim_samples,
        )
        self.self_lyric_sim_dataset = SelfLyricSimDataset(
            split,
            input_length=self_lyric_sim_input_length,
            num_samples=num_self_lyric_sim_samples,
        )
        self.artist_sim_dataset = ArtistSimDataset(
            metadata,
            task_indices,
            split,
            artist_sim_input_length_s,
            sample_rate,
            num_artist_sim_samples,
        )
        self.artist_vox_sim_dataset = ArtistVoxSimDataset(
            metadata,
            task_indices,
            split,
            artist_vox_sim_input_length_s,
            sample_rate,
            num_artist_vox_sim_samples,
        )
        self.album_sim_dataset = AlbumSimDataset(
            metadata,
            task_indices,
            split,
            album_sim_input_length_s,
            sample_rate,
            num_album_sim_samples,
        )
        self.genre_sim_dataset = GenreSimDataset(
            metadata,
            task_indices,
            split,
            genre_sim_input_length_s,
            genre_tag_input_length,
            sample_rate,
            num_genre_sim_samples,
        )
        self.lyric_sim_dataset = LyricSimDataset(
            metadata,
            task_indices,
            split,
            lyric_sim_input_length_s,
            lyric_sim_input_text_length,
            sample_rate,
            num_lyric_sim_samples,
        )
        self.split = split
        self.num_samples = num_samples
        self.accumulated_lengths = self.get_accumulated_lengths()

    def get_accumulated_lengths(self):
        accumulated_lengths = []
        current_sum = 0
        for dataset in [
            self.self_sim_dataset,
            self.self_vox_sim_dataset,
            self.self_lyric_sim_dataset,
            self.artist_sim_dataset,
            self.artist_vox_sim_dataset,
            self.album_sim_dataset,
            self.genre_sim_dataset,
            self.lyric_sim_dataset,
        ]:
            current_sum += len(dataset)
            accumulated_lengths.append(current_sum)
        return accumulated_lengths

    def __getitem__(self, index):
        if index < self.accumulated_lengths[0]:
            return self.self_sim_dataset[index], "self_sim"
        elif index < self.accumulated_lengths[1]:
            return self.self_vox_sim_dataset[
                index - self.accumulated_lengths[0]
            ], "self_vox_sim"
        elif index < self.accumulated_lengths[2]:
            return self.self_lyric_sim_dataset[
                index - self.accumulated_lengths[1]
            ], "self_lyric_sim"
        elif index < self.accumulated_lengths[3]:
            return self.artist_sim_dataset[
                index - self.accumulated_lengths[2]
            ], "artist_sim"
        elif index < self.accumulated_lengths[4]:
            return self.artist_vox_sim_dataset[
                index - self.accumulated_lengths[3]
            ], "artist_vox_sim"
        elif index < self.accumulated_lengths[5]:
            return self.album_sim_dataset[
                index - self.accumulated_lengths[4]
            ], "album_sim"
        elif index < self.accumulated_lengths[6]:
            return self.genre_sim_dataset[
                index - self.accumulated_lengths[5]
            ], "genre_sim"
        else:
            return self.lyric_sim_dataset[
                index - self.accumulated_lengths[6]
            ], "lyric_sim"

    def __len__(self):
        return self.accumulated_lengths[-1]


class MultiTaskBatchSampler(Sampler):
    def __init__(self, dataset, batch_size):
        self.dataset = dataset
        self.batch_size = batch_size
        self.task_sizes = [
            len(dataset.self_sim_dataset),
            len(dataset.self_vox_sim_dataset),
            len(dataset.self_lyric_sim_dataset),
            len(dataset.artist_sim_dataset),
            len(dataset.artist_vox_sim_dataset),
            len(dataset.album_sim_dataset),
            len(dataset.genre_sim_dataset),
            len(dataset.lyric_sim_dataset),
        ]
        self.task_indices = list(range(len(self.task_sizes)))

    def __iter__(self):
        if self.dataset.split == "train":
            while True:
                # Randomly choose a task
                task = random.choice(self.task_indices)

                # Calculate the start and end indices for the chosen task
                start_idx = sum(self.task_sizes[:task])
                end_idx = start_idx + self.task_sizes[task]

                # Generate a batch of indices for the chosen task
                batch_indices = random.sample(
                    range(start_idx, end_idx),
                    min(self.batch_size, self.task_sizes[task]),
                )

                yield batch_indices
        else:  # valid
            task_index = 0
            while True:
                # Calculate the start and end indices for the current task
                start_idx = sum(self.task_sizes[:task_index])
                end_idx = start_idx + self.task_sizes[task_index]

                # Generate batches for the current task
                for i in range(start_idx, end_idx, self.batch_size):
                    batch_indices = list(range(i, min(i + self.batch_size, end_idx)))
                    yield batch_indices

                # Move to the next task
                task_index = (task_index + 1) % len(self.task_indices)

    def __len__(self):
        return sum(self.task_sizes) // self.batch_size
